diff --git a/.github/workflows/scripts/estimated_times.yaml b/.github/workflows/scripts/estimated_times.yaml index a7ab525c7927..54a2b49aa392 100644 --- a/.github/workflows/scripts/estimated_times.yaml +++ b/.github/workflows/scripts/estimated_times.yaml @@ -49,6 +49,7 @@ estimated_times: tests/e2e/pull_request/one_card/test_completion_with_prompt_embeds.py: 170 tests/e2e/pull_request/one_card/test_cpu_offloading.py: 30 tests/e2e/pull_request/one_card/test_cpu_weight_offload.py: 890 + tests/e2e/pull_request/one_card/test_deepseek_v4_vision_precision.py: 600 tests/e2e/pull_request/one_card/test_guided_decoding.py: 670 tests/e2e/pull_request/one_card/test_minicpm.py: 300 tests/e2e/pull_request/one_card/test_minimax_m3_sparse_attn.py: 440 diff --git a/benchmarks/prepare_indexer_indices.py b/benchmarks/prepare_indexer_indices.py new file mode 100644 index 000000000000..b6adbb29f71a --- /dev/null +++ b/benchmarks/prepare_indexer_indices.py @@ -0,0 +1,112 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Compare fused indexer postprocessing in an NPU graph. + +Run: python benchmarks/prepare_indexer_indices.py +Times exclude compilation, graph capture and host tensor allocation. Repetition +counts adapt to a warmup measurement so slow INT32-sort baselines stay bounded. +""" + +import argparse +import json +import statistics +from functools import partial + +import torch +import torch_npu # noqa: F401 + +from vllm_ascend.ops.triton.prepare_indexer_indices import prepare_indexer_indices +from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton + + +def reference_indices(selected, positions, compress_ratio): + visible = ((positions + 1) // compress_ratio).unsqueeze(-1) + valid = (selected >= 0) & (selected < visible) + sentinel = torch.iinfo(torch.int32).max + selected = torch.where(valid, selected, sentinel).sort(dim=-1).values + return torch.where(selected == sentinel, -1, selected) + + +def graph_latency_us(fn, value): + for _ in range(3): + fn(value) + torch.npu.synchronize() + start = torch.npu.Event(enable_timing=True) + end = torch.npu.Event(enable_timing=True) + start.record() + fn(value) + end.record() + end.synchronize() + estimate_ms = max(start.elapsed_time(end), 0.001) + # Capture at most about 10 ms of work and measure about 50 ms per sample. + batch = min(32, max(1, int(10 / estimate_ms))) + repeats = min(20, max(1, int(50 / (batch * estimate_ms)))) + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True): + for _ in range(batch): + output = fn(value) + for _ in range(3): + graph.replay() + torch.npu.synchronize() + samples = [] + for _ in range(5): + start = torch.npu.Event(enable_timing=True) + end = torch.npu.Event(enable_timing=True) + start.record() + for _ in range(repeats): + graph.replay() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) * 1000 / (batch * repeats)) + # Keep the captured outputs alive until timing completes. + del output + return statistics.median(samples) + + +def benchmark(stage, value, reference, fused, **shape): + expected, actual = reference(value), fused(value) + if isinstance(expected, torch.Tensor): + expected, actual = (expected,), (actual,) + for output, ref in zip(actual, expected): + torch.testing.assert_close(output, ref, rtol=0, atol=0) + original_us = graph_latency_us(reference, value) + fused_us = graph_latency_us(fused, value) + print( + json.dumps( + { + "stage": stage, + **shape, + "reference_us": original_us, + "triton_us": fused_us, + "speedup": original_us / fused_us, + } + ), + flush=True, + ) + + +@torch.inference_mode() +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--tokens", type=int, nargs="+", default=[1, 32, 256, 4096]) + parser.add_argument("--topk", type=int, nargs="+", default=[128, 2048]) + args = parser.parse_args() + torch.npu.set_device(0) + init_device_properties_triton() + torch.manual_seed(41) + for tokens in args.tokens: + for topk in args.topk: + selected = torch.randint(-1, 4096, (tokens, topk), dtype=torch.int32, device="npu") + positions = torch.full((tokens,), 4095, dtype=torch.int64, device="npu") + benchmark( + "indices", + selected, + partial(reference_indices, positions=positions, compress_ratio=2), + partial(prepare_indexer_indices, positions=positions, compress_ratio=2), + tokens=tokens, + topk=topk, + ) + + +if __name__ == "__main__": + main() diff --git a/benchmarks/quantize_indexer_query.py b/benchmarks/quantize_indexer_query.py new file mode 100644 index 000000000000..89bcbfceed6a --- /dev/null +++ b/benchmarks/quantize_indexer_query.py @@ -0,0 +1,83 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Compare indexer query quantization latency inside an NPU graph. + +Run: python benchmarks/quantize_indexer_query.py +Times exclude compilation, graph capture and host tensor allocation. +""" + +import argparse +import json +import statistics + +import torch +import torch_npu # noqa: F401 + +from vllm_ascend.ops.triton.quantize_indexer_query import quantize_indexer_query +from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton + + +def reference(query): + scale = (query.float().abs().amax(-1) / 127.0).half().clamp_min_(2.0**-24) + quantized = (query.float() / scale.float().unsqueeze(-1)).round().clamp(-127, 127).to(torch.int8) + return quantized, scale + + +def graph_latency_us(fn, query, batch=32, repeats=20): + for _ in range(3): + fn(query) + torch.npu.synchronize() + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True): + for _ in range(batch): + output = fn(query) + for _ in range(3): + graph.replay() + torch.npu.synchronize() + samples = [] + for _ in range(5): + start = torch.npu.Event(enable_timing=True) + end = torch.npu.Event(enable_timing=True) + start.record() + for _ in range(repeats): + graph.replay() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) * 1000 / (batch * repeats)) + # Keep the captured outputs alive until timing completes. + assert output[0].shape == query.shape + return statistics.median(samples) + + +@torch.inference_mode() +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--tokens", type=int, nargs="+", default=[1, 32, 256, 4096]) + parser.add_argument("--heads", type=int, nargs="+", default=[32, 64]) + args = parser.parse_args() + torch.npu.set_device(0) + init_device_properties_triton() + torch.manual_seed(41) + for tokens in args.tokens: + for heads in args.heads: + query = torch.randn(tokens, heads, 128, dtype=torch.bfloat16, device="npu") + for actual, expected in zip(quantize_indexer_query(query), reference(query)): + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + original_us = graph_latency_us(reference, query) + fused_us = graph_latency_us(quantize_indexer_query, query) + print( + json.dumps( + { + "tokens": tokens, + "heads": heads, + "reference_us": original_us, + "triton_us": fused_us, + "speedup": original_us / fused_us, + } + ), + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/csrc/attention/common/op_kernel/aicpu_common.h b/csrc/attention/common/op_kernel/aicpu_common.h index 30e06bee6ddc..98104dbbbb32 100644 --- a/csrc/attention/common/op_kernel/aicpu_common.h +++ b/csrc/attention/common/op_kernel/aicpu_common.h @@ -16,6 +16,7 @@ #ifndef AICPU_COMMON_H #define AICPU_COMMON_H +#include #include #include #include "log.h" @@ -103,7 +104,7 @@ inline bool IsTensorExists(const Tensor *tensor) inline std::vector GetTensorDataAsInt64(const Tensor *tensor) { - std::vector result {}; + std::vector result{}; if (!IsTensorExists(tensor)) { return result; diff --git a/csrc/attention/common/op_kernel/arch35/vf/vf_flash_decode_arch35.h b/csrc/attention/common/op_kernel/arch35/vf/vf_flash_decode_arch35.h new file mode 100644 index 000000000000..e0af420ef1d1 --- /dev/null +++ b/csrc/attention/common/op_kernel/arch35/vf/vf_flash_decode_arch35.h @@ -0,0 +1,643 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_flash_decode_arch35.h + * \brief + */ +#ifndef MY_FLASH_DECODE_ARCH35_H +#define MY_FLASH_DECODE_ARCH35_H + +#include "kernel_tensor.h" + +constexpr float FLT_ZERO = 0; +constexpr float FLT_MAX_NEW = 3.402823466e+38F; + +namespace FaVectorApi { +// bf16->fp32 +static constexpr Reg::CastTrait castTraitFp16_32 = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; +// 处理循环splitKVIndex=0的场景,vregDst需要置0 +template +__simd_vf__ void ReduceFinalRes_0_VF(__ubuf__ T *dstUb, __ubuf__ T *lseUb, __ubuf__ T *accumOutUb, uint16_t k, + uint16_t z, uint32_t dealNum1Reg, uint32_t repStride, const uint16_t floatRepSize, + const uint16_t dLoops, uint32_t dealRowCount, uint32_t splitKVIndex) +{ + Reg::RegTensor vregDst; + Reg::RegTensor vregLse; + Reg::RegTensor vregAccumOut; + uint32_t n = dealNum1Reg; + Reg::MaskReg pregTailN = Reg::UpdateMask(n); + + for (k = 0; k < static_cast(dealRowCount); k++) { // repeat g + + Reg::LoadAlign(vregLse, + (__ubuf__ float *&)lseUb + splitKVIndex * dealRowCount * 8 + k * 8); + for (z = 0; z < dLoops; z++) { + // splitKVIndex=0的场景,vregDst不需要load,直接置0 + Reg::Duplicate(vregDst, FLT_ZERO, pregTailN); + Reg::LoadAlign( + vregAccumOut, (__ubuf__ float *&)accumOutUb + k * repStride * 8 + z * floatRepSize); + Reg::Mul(vregAccumOut, vregLse, vregAccumOut, pregTailN); + Reg::Add(vregDst, vregDst, vregAccumOut, pregTailN); + Reg::StoreAlign( + (__ubuf__ float *&)dstUb + k * repStride * 8 + z * floatRepSize, vregDst, pregTailN); + } + } +} + +template +__aicore__ inline void ReduceFinalRes_0(LocalTensor &dstLocal, LocalTensor &lseLocal, + LocalTensor &accumOutLocal, uint32_t dealRowCount, uint64_t headDimAlignFp32, + uint32_t splitKVIndex) +{ + __ubuf__ T *dstUb = (__ubuf__ T *)dstLocal.GetPhyAddr(); + __ubuf__ T *lseUb = (__ubuf__ T *)lseLocal.GetPhyAddr(); + __ubuf__ T *accumOutUb = (__ubuf__ T *)accumOutLocal.GetPhyAddr(); + uint16_t z = 0; + uint16_t k = 0; + const uint16_t floatRepSize = 64; + const uint16_t dLoops = headDimAlignFp32 / floatRepSize; + uint32_t dealNum1Reg = 256 / sizeof(float); + uint32_t repStride = headDimAlignFp32 / 8; + + ReduceFinalRes_0_VF(dstUb, lseUb, accumOutUb, k, z, dealNum1Reg, repStride, floatRepSize, dLoops, dealRowCount, + splitKVIndex); +} + +// 处理循环splitKVIndex>0的场景,reg_dst需要先从dstUb中load之前的结果,再进行add +template +__simd_vf__ void ReduceFinalRes_Rest_VF(__ubuf__ T *dstUb, __ubuf__ T *lseUb, __ubuf__ T *accumOutUb, uint16_t k, + uint16_t z, uint32_t dealNum1Reg, uint32_t repStride, + const uint16_t floatRepSize, const uint16_t dLoops, uint32_t dealRowCount, + uint32_t splitKVIndex) +{ + Reg::RegTensor vregDst; + Reg::RegTensor vregLse; + Reg::RegTensor vregAccumOut; + uint32_t n = dealNum1Reg; + Reg::MaskReg pregTailN = Reg::UpdateMask(n); + uint32_t stride = (0x1 << 16) | 0x8; + + for (k = 0; k < static_cast(dealRowCount); k++) { // repeat g + Reg::LoadAlign(vregLse, + (__ubuf__ float *&)lseUb + splitKVIndex * dealRowCount * 8 + k * 8); + for (z = 0; z < dLoops; z++) { + // splitKVIndex>0的场景,reg_dst需要先从dstUb中load之前的结果,再进行add + Reg::LoadAlign( + vregDst, (__ubuf__ float *&)dstUb + k * repStride * 8 + z * floatRepSize); + Reg::LoadAlign( + vregAccumOut, (__ubuf__ float *&)accumOutUb + k * repStride * 8 + z * floatRepSize); + Reg::Mul(vregAccumOut, vregLse, vregAccumOut, pregTailN); + Reg::Add(vregDst, vregDst, vregAccumOut, pregTailN); + Reg::StoreAlign( + (__ubuf__ float *&)dstUb + k * repStride * 8 + z * floatRepSize, vregDst, pregTailN); + } + } +} + +template +__aicore__ inline void ReduceFinalRes_Rest(LocalTensor &dstLocal, LocalTensor &lseLocal, + LocalTensor &accumOutLocal, uint32_t dealRowCount, + uint64_t headDimAlignFp32, uint32_t splitKVIndex) +{ + __ubuf__ T *dstUb = (__ubuf__ T *)dstLocal.GetPhyAddr(); + __ubuf__ T *lseUb = (__ubuf__ T *)lseLocal.GetPhyAddr(); + __ubuf__ T *accumOutUb = (__ubuf__ T *)accumOutLocal.GetPhyAddr(); + uint16_t k = 0; + uint16_t z = 0; + uint32_t dealNum1Reg = 256 / sizeof(float); + uint32_t repStride = headDimAlignFp32 / 8; + const uint16_t floatRepSize = 64; + const uint16_t dLoops = headDimAlignFp32 / floatRepSize; + + ReduceFinalRes_Rest_VF(dstUb, lseUb, accumOutUb, k, z, dealNum1Reg, repStride, floatRepSize, dLoops, + dealRowCount, splitKVIndex); +} + +template +__aicore__ inline void ReduceFinalRes_VF(LocalTensor &dstLocal, LocalTensor &lseLocal, + LocalTensor &accumOutLocal, uint32_t dealRowCount, + uint64_t headDimAlignFp32, uint32_t splitKVIndex) +{ + if (splitKVIndex == 0) { + ReduceFinalRes_0(dstLocal, lseLocal, accumOutLocal, dealRowCount, headDimAlignFp32, splitKVIndex); + } else { + ReduceFinalRes_Rest(dstLocal, lseLocal, accumOutLocal, dealRowCount, headDimAlignFp32, splitKVIndex); + } +} + +template +__aicore__ inline void ReduceFinalRes_const_VF(LocalTensor &dstLocal, LocalTensor &lseLocal, + LocalTensor &accumOutLocal, uint32_t dealRowCount, + uint32_t splitKVIndex) +{ + if (splitKVIndex == 0) { + ReduceFinalRes_0(dstLocal, lseLocal, accumOutLocal, dealRowCount, headDimAlignFp32, splitKVIndex); + } else { + ReduceFinalRes_Rest(dstLocal, lseLocal, accumOutLocal, dealRowCount, headDimAlignFp32, splitKVIndex); + } +} + +// 处理g<=8的场景 +template +__simd_vf__ void ComputeScaleValue_8_VF(__ubuf__ uint16_t *lseSink, __ubuf__ T *lseMax, __ubuf__ T *lseMaxTmp, + __ubuf__ T *lseSum, __ubuf__ T *lseSumTmp, __ubuf__ T *lseUb, + uint32_t dealCount, uint16_t i, uint32_t dealRowCount, + uint32_t actualCombineLoopSize, bool softmaxLseFlag, bool learnableSinkFlag) +{ + Reg::RegTensor vregLseMax; + Reg::RegTensor vregLseMaxTmp; + Reg::RegTensor vregLseSum; + Reg::RegTensor vregLseSumTmp; + Reg::RegTensor vregLseSink; + Reg::RegTensor vregLseSinkCast; + Reg::RegTensor vregRes; + uint32_t n = dealCount; + Reg::MaskReg pregTailN = Reg::UpdateMask(n); + Reg::MaskReg pregSinkTailN = Reg::UpdateMask(n); + uint16_t blockStride = 0x1; + uint16_t repeatStride = dealRowCount; + + Reg::Duplicate(vregLseMax, -FLT_MAX_NEW, pregTailN); + Reg::Duplicate(vregLseSum, FLT_ZERO, pregTailN); + + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseMaxTmp, (__ubuf__ float *&)lseMaxTmp + i * dealCount); + Reg::Max(vregLseMax, vregLseMax, vregLseMaxTmp, pregTailN); + } + + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseMaxTmp, (__ubuf__ float *&)lseMaxTmp + i * dealCount); + Reg::Sub(vregLseMaxTmp, vregLseMaxTmp, vregLseMax, pregTailN); + Reg::Exp(vregLseMaxTmp, vregLseMaxTmp, pregTailN); + Reg::LoadAlign(vregLseSumTmp, (__ubuf__ float *&)lseSumTmp + i * dealCount); + Reg::Mul(vregLseSumTmp, vregLseSumTmp, vregLseMaxTmp, pregTailN); + Reg::Add(vregLseSum, vregLseSum, vregLseSumTmp, pregTailN); + Reg::StoreAlign((__ubuf__ float *&)lseSumTmp + i * dealCount, vregLseSumTmp, + pregTailN); + } + + if (learnableSinkFlag) { + Reg::LoadAlign((Reg::RegTensor &)vregLseSink, lseSink); + Reg::Cast(vregLseSinkCast, vregLseSink, pregSinkTailN); + + Reg::Sub(vregLseSinkCast, vregLseSinkCast, vregLseMax, pregTailN); + Reg::Exp(vregLseSinkCast, vregLseSinkCast, pregTailN); + Reg::Add(vregLseSum, vregLseSum, vregLseSinkCast, pregTailN); + } + + if (softmaxLseFlag) { + Reg::RegTensor vregMinValue; + Reg::RegTensor vregInfValue; + Reg::MaskReg pregCompare; + constexpr float infValue = 3e+99; // 3e+99 for float inf + constexpr uint32_t tmpMin = 0xFF167699; + float minValue = *((float *)&tmpMin); + Reg::Duplicate(vregMinValue, minValue); + Reg::Duplicate(vregInfValue, infValue); + + Reg::Log(vregRes, vregLseSum, pregTailN); + Reg::Add(vregRes, vregRes, vregLseMax, pregTailN); + // 如果 softmaxMax 等于负无穷,则将 lse 结果置为 inf + Reg::Compare(pregCompare, vregLseMax, vregMinValue, pregTailN); + Reg::Select(vregRes, vregInfValue, vregRes, pregCompare); + Reg::StoreAlign(lseUb, vregRes, pregTailN); + } + + Reg::LocalMemBar(); + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseSumTmp, (__ubuf__ float *&)lseSumTmp + i * dealCount); + Reg::Div(vregLseSumTmp, vregLseSumTmp, vregLseSum, pregTailN); + Reg::StoreAlign( + (__ubuf__ float *&)lseSum, vregLseSumTmp, blockStride, repeatStride, pregTailN); + } +} + +template +__aicore__ inline void ComputeScaleValue_8(const LocalTensor &tmpSinkUb, const LocalTensor &lseMaxUb, + const LocalTensor &lseSumUb, const LocalTensor &lseOutputUb, + uint32_t dealRowCount, uint32_t actualCombineLoopSize, bool softmaxLseFlag, + bool learnableSinkFlag) +{ + uint32_t dealCount = dealRowCount * 8; + uint16_t i = 0; + + __ubuf__ T *lseMax = (__ubuf__ T *)lseMaxUb.GetPhyAddr(); + __ubuf__ T *lseMaxTmp = lseMax; + __ubuf__ T *lseSum = (__ubuf__ T *)lseSumUb.GetPhyAddr(); + __ubuf__ T *lseSumTmp = lseSum; + __ubuf__ T *lseUb = (__ubuf__ T *)lseOutputUb.GetPhyAddr(); + __ubuf__ uint16_t *lseSink = (__ubuf__ uint16_t *)tmpSinkUb.GetPhyAddr(); + + ComputeScaleValue_8_VF(lseSink, lseMax, lseMaxTmp, lseSum, lseSumTmp, lseUb, dealCount, i, dealRowCount, + actualCombineLoopSize, softmaxLseFlag, learnableSinkFlag); +} + +// //lseUb作为scale最终输出 +template +__simd_vf__ void ComputeScaleValue_8_VF_FD(__ubuf__ T *lseSink, __ubuf__ T *lseMax, __ubuf__ T *lseMaxTmp, + __ubuf__ T *lseSum, __ubuf__ T *lseSumTmp, __ubuf__ T *lseOutUb, + __ubuf__ T *lseUb, __ubuf__ T *lseMaxReduce, uint32_t dealCount, uint16_t i, + uint32_t dealRowCount, uint32_t actualCombineLoopSize, + uint16_t softmaxLseFlag, uint16_t learnableSinkFlag) +{ + Reg::RegTensor vregLseMax; + Reg::RegTensor vregLseMaxTmp; + Reg::RegTensor vregRes; + Reg::RegTensor vregLseSum; + Reg::RegTensor vregLseSumTmp; + Reg::RegTensor vregLseSinkCast; + uint16_t blockStride = 0x1; + uint16_t repeatStride = dealRowCount; + uint32_t n = dealCount; + Reg::MaskReg pregTailN = Reg::UpdateMask(n); + Reg::MaskReg pregSinkTailN = Reg::UpdateMask(n); + + Reg::Duplicate(vregLseSum, FLT_ZERO, pregTailN); + Reg::Duplicate(vregLseMax, -FLT_MAX_NEW, pregTailN); + + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseMaxTmp, (__ubuf__ float *&)lseMaxTmp + i * dealCount); + Reg::Max(vregLseMax, vregLseMax, vregLseMaxTmp, pregTailN); + } + Reg::StoreAlign(lseMaxReduce, vregLseMax, pregTailN); + + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseMaxTmp, (__ubuf__ float *&)lseMaxTmp + i * dealCount); + Reg::Sub(vregLseMaxTmp, vregLseMaxTmp, vregLseMax, pregTailN); + Reg::Exp(vregLseMaxTmp, vregLseMaxTmp, pregTailN); + Reg::LoadAlign(vregLseSumTmp, (__ubuf__ float *&)lseSumTmp + i * dealCount); + Reg::Mul(vregLseSumTmp, vregLseSumTmp, vregLseMaxTmp, pregTailN); + Reg::Add(vregLseSum, vregLseSum, vregLseSumTmp, pregTailN); + Reg::StoreAlign((__ubuf__ float *&)lseSumTmp + i * dealCount, vregLseSumTmp, + pregTailN); + } + + for (i = 0; i < static_cast(learnableSinkFlag); ++i) { + Reg::LoadAlign(vregLseSinkCast, (__ubuf__ float *&)lseSink); + Reg::Sub(vregLseSinkCast, vregLseSinkCast, vregLseMax, pregTailN); + Reg::Exp(vregLseSinkCast, vregLseSinkCast, pregTailN); + Reg::Add(vregLseSum, vregLseSum, vregLseSinkCast, pregTailN); + } + + for (i = 0; i < static_cast(softmaxLseFlag); ++i) { + Reg::RegTensor vregMinValue; + Reg::RegTensor vregInfValue; + Reg::MaskReg pregCompare; + constexpr float infValue = 3e+99; // 3e+99 for float inf + constexpr uint32_t tmpMin = 0xFF167699; + float minValue = *((float *)&tmpMin); + Reg::Duplicate(vregInfValue, infValue); + Reg::Duplicate(vregMinValue, minValue); + + Reg::Log(vregRes, vregLseSum, pregTailN); + Reg::Add(vregRes, vregRes, vregLseMax, pregTailN); + // 如果 softmaxMax 等于负无穷,则将 lse 结果置为 inf + Reg::Compare(pregCompare, vregLseMax, vregMinValue, pregTailN); + Reg::Select(vregRes, vregInfValue, vregRes, pregCompare); + Reg::StoreAlign(lseOutUb, vregRes, pregTailN); + } + + Reg::LocalMemBar(); + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseSumTmp, (__ubuf__ float *&)lseSumTmp + i * dealCount); + Reg::Div(vregLseSumTmp, vregLseSumTmp, vregLseSum, pregTailN); + Reg::StoreAlign( + (__ubuf__ float *&)lseUb, vregLseSumTmp, blockStride, repeatStride, pregTailN); + } +} + +template +__aicore__ inline void ComputeScaleValue_8_FD(const LocalTensor &tmpSinkUb, const LocalTensor &lseMaxUb, + const LocalTensor &lseSumUb, const LocalTensor &lseResUb, + const LocalTensor &lseOutputUb, const LocalTensor &lseMaxReduceUb, + uint32_t dealRowCount, uint32_t actualCombineLoopSize, + bool softmaxLseFlag, bool learnableSinkFlag) +{ + uint32_t dealCount = dealRowCount * 8; + uint16_t i = 0; + + __ubuf__ T *lseMax = (__ubuf__ T *)lseMaxUb.GetPhyAddr(); + __ubuf__ T *lseMaxTmp = lseMax; + __ubuf__ T *lseSum = (__ubuf__ T *)lseSumUb.GetPhyAddr(); + __ubuf__ T *lseSumTmp = lseSum; + __ubuf__ T *lseUb = (__ubuf__ T *)lseResUb.GetPhyAddr(); + __ubuf__ T *lseOutUb = (__ubuf__ T *)lseOutputUb.GetPhyAddr(); + __ubuf__ T *lseMaxReduce = (__ubuf__ T *)lseMaxReduceUb.GetPhyAddr(); + // lseSink在前面已经类型转换成T + __ubuf__ T *lseSink = (__ubuf__ T *)tmpSinkUb.GetPhyAddr(); + uint16_t softmaxLseFlagUint = 0; + uint16_t learnableSinkFlagUint = 0; + if (softmaxLseFlag) { + softmaxLseFlagUint = 1; + } + if (learnableSinkFlag) { + learnableSinkFlagUint = 1; + } + ComputeScaleValue_8_VF_FD(lseSink, lseMax, lseMaxTmp, lseSum, lseSumTmp, lseOutUb, lseUb, lseMaxReduce, + dealCount, i, dealRowCount, actualCombineLoopSize, softmaxLseFlagUint, + learnableSinkFlagUint); +} + +// 处理8 +__simd_vf__ void ComputeScaleValue_16_VF(__ubuf__ uint16_t *lseSink, __ubuf__ uint16_t *lseSink2, __ubuf__ T *lseMax, + __ubuf__ T *lseMax2, __ubuf__ T *lseMaxSrc, __ubuf__ T *lseSum, + __ubuf__ T *lseSum2, __ubuf__ T *lseSumSrc, __ubuf__ T *lseUb, + __ubuf__ T *lseUb2, uint32_t dealCountSum, uint32_t dealCount, + uint32_t dealCount2, uint16_t i, uint32_t dealRowCount, + uint32_t actualCombineLoopSize, bool softmaxLseFlag, bool learnableSinkFlag) +{ + Reg::RegTensor vregLseMax; + Reg::RegTensor vregLseMaxTmp; + Reg::RegTensor vregLseMax2; + Reg::RegTensor vregLseMaxTmp2; + Reg::RegTensor vregLseSum; + Reg::RegTensor vregLseSumTmp; + Reg::RegTensor vregLseSum2; + Reg::RegTensor vregLseSumTmp2; + Reg::RegTensor vregLseSink; + Reg::RegTensor vregLseSink2; + Reg::RegTensor vregLseSinkCast; + Reg::RegTensor vregLseSinkCast2; + Reg::RegTensor vregRes; + Reg::RegTensor vregRes2; + uint32_t n = dealCount; + uint32_t n2 = dealCount2; + Reg::MaskReg pregTailN = Reg::UpdateMask(n); + Reg::MaskReg pregTailN2 = Reg::UpdateMask(n2); + Reg::MaskReg pregSinkTailN = Reg::UpdateMask(n); + Reg::MaskReg pregSinkTailN2 = Reg::UpdateMask(n2); + uint16_t blockStride = 0x1; + uint16_t repeatStride = dealRowCount; + + Reg::Duplicate(vregLseMax, -FLT_MAX_NEW, pregTailN); + Reg::Duplicate(vregLseMax2, -FLT_MAX_NEW, pregTailN); + Reg::Duplicate(vregLseSum, FLT_ZERO, pregTailN); + Reg::Duplicate(vregLseSum2, FLT_ZERO, pregTailN); + + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseMaxTmp, lseMaxSrc + i * dealCountSum); + Reg::LoadAlign(vregLseMaxTmp2, lseMaxSrc + i * dealCountSum + dealCount); + Reg::Max(vregLseMax, vregLseMax, vregLseMaxTmp, pregTailN); + Reg::Max(vregLseMax2, vregLseMax2, vregLseMaxTmp2, pregTailN2); + } + + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseMaxTmp, lseMaxSrc + i * dealCountSum); + Reg::LoadAlign(vregLseMaxTmp2, lseMaxSrc + i * dealCountSum + dealCount); + Reg::Sub(vregLseMaxTmp, vregLseMaxTmp, vregLseMax, pregTailN); + Reg::Sub(vregLseMaxTmp2, vregLseMaxTmp2, vregLseMax2, pregTailN2); + Reg::Exp(vregLseMaxTmp, vregLseMaxTmp, pregTailN); + Reg::Exp(vregLseMaxTmp2, vregLseMaxTmp2, pregTailN2); + Reg::LoadAlign(vregLseSumTmp, lseSumSrc + i * dealCountSum); + Reg::LoadAlign(vregLseSumTmp2, lseSumSrc + i * dealCountSum + dealCount); + Reg::Mul(vregLseSumTmp, vregLseSumTmp, vregLseMaxTmp, pregTailN); + Reg::Mul(vregLseSumTmp2, vregLseSumTmp2, vregLseMaxTmp2, pregTailN2); + Reg::Add(vregLseSum, vregLseSum, vregLseSumTmp, pregTailN); + Reg::Add(vregLseSum2, vregLseSum2, vregLseSumTmp2, pregTailN2); + Reg::StoreAlign(lseSumSrc + i * dealCountSum, vregLseSumTmp, pregTailN); + Reg::StoreAlign(lseSumSrc + i * dealCountSum + dealCount, vregLseSumTmp2, + pregTailN2); + } + + if (learnableSinkFlag) { + Reg::LoadAlign((Reg::RegTensor &)vregLseSink, lseSink); + Reg::LoadAlign((Reg::RegTensor &)vregLseSink2, + lseSink + dealCount); + + Reg::Cast(vregLseSinkCast, vregLseSink, pregSinkTailN); + Reg::Cast(vregLseSinkCast2, vregLseSink2, pregSinkTailN2); + + Reg::Sub(vregLseSinkCast, vregLseSinkCast, vregLseMax, pregTailN); + Reg::Sub(vregLseSinkCast2, vregLseSinkCast2, vregLseMax2, pregTailN2); + + Reg::Exp(vregLseSinkCast, vregLseSinkCast, pregTailN); + Reg::Exp(vregLseSinkCast2, vregLseSinkCast2, pregTailN2); + + Reg::Add(vregLseSum, vregLseSum, vregLseSinkCast, pregTailN); + Reg::Add(vregLseSum2, vregLseSum2, vregLseSinkCast2, pregTailN2); + } + + if (softmaxLseFlag) { + Reg::RegTensor vregMinValue; + Reg::RegTensor vregInfValue; + Reg::MaskReg pregCompare; + Reg::MaskReg pregCompare2; + constexpr float infValue = 3e+99; // 3e+99 for float inf + constexpr uint32_t tmpMin = 0xFF167699; + float minValue = *((float *)&tmpMin); + Reg::Duplicate(vregMinValue, minValue); + Reg::Duplicate(vregInfValue, infValue); + + Reg::Log(vregRes, vregLseSum, pregTailN); + Reg::Add(vregRes, vregRes, vregLseMax, pregTailN); + Reg::Log(vregRes2, vregLseSum2, pregTailN2); + Reg::Add(vregRes2, vregRes2, vregLseMax2, pregTailN2); + // 如果 softmaxMax 等于负无穷,则将 lse 结果置为 inf + Reg::Compare(pregCompare, vregLseMax, vregMinValue, pregTailN); + Reg::Compare(pregCompare2, vregLseMax2, vregMinValue, pregTailN2); + Reg::Select(vregRes, vregInfValue, vregRes, pregCompare); + Reg::Select(vregRes2, vregInfValue, vregRes2, pregCompare2); + Reg::StoreAlign(lseUb, vregRes, pregTailN); + Reg::StoreAlign(lseUb2, vregRes2, pregTailN2); + } + + Reg::LocalMemBar(); + for (i = 0; i < static_cast(actualCombineLoopSize); ++i) { + Reg::LoadAlign(vregLseSumTmp, lseSumSrc + i * dealCountSum); + Reg::LoadAlign(vregLseSumTmp2, lseSumSrc + i * dealCountSum + dealCount); + Reg::Div(vregLseSumTmp, vregLseSumTmp, vregLseSum, pregTailN); + Reg::Div(vregLseSumTmp2, vregLseSumTmp2, vregLseSum2, pregTailN2); + Reg::StoreAlign( + lseSum, vregLseSumTmp, blockStride, repeatStride, pregTailN); + Reg::StoreAlign( + lseSum2, vregLseSumTmp2, blockStride, repeatStride, pregTailN2); + } +} + +template +__aicore__ inline void ComputeScaleValue_16(const LocalTensor &tmpSinkUb, const LocalTensor &lseMaxUb, + const LocalTensor &lseSumUb, const LocalTensor &lseOutputUb, + uint32_t dealRowCount, uint32_t actualCombineLoopSize, bool softmaxLseFlag, + bool learnableSinkFlag) +{ + uint32_t dealCountSum = dealRowCount * 8; + uint32_t dealCount = 8 * 8; + uint32_t dealCount2 = dealCountSum - dealCount; + uint16_t i = 0; + + __ubuf__ T *lseMax = (__ubuf__ T *)lseMaxUb.GetPhyAddr(); + __ubuf__ T *lseMax2 = lseMax + 64; + __ubuf__ T *lseMaxSrc = lseMax; + __ubuf__ T *lseSum = (__ubuf__ T *)lseSumUb.GetPhyAddr(); + __ubuf__ T *lseSum2 = lseSum + 64; + __ubuf__ T *lseSumSrc = lseSum; + __ubuf__ T *lseUb = (__ubuf__ T *)lseOutputUb.GetPhyAddr(); + __ubuf__ T *lseUb2 = lseUb + 64; + __ubuf__ uint16_t *lseSink = (__ubuf__ uint16_t *)tmpSinkUb.GetPhyAddr(); + __ubuf__ uint16_t *lseSink2 = lseSink + 64; + + ComputeScaleValue_16_VF(lseSink, lseSink2, lseMax, lseMax2, lseMaxSrc, lseSum, lseSum2, lseSumSrc, lseUb, + lseUb2, dealCountSum, dealCount, dealCount2, i, dealRowCount, + actualCombineLoopSize, softmaxLseFlag, learnableSinkFlag); +} + +template +__aicore__ inline void ComputeScaleValue_VF(const LocalTensor &tmpSinkUb, const LocalTensor &lseMaxUb, + const LocalTensor &lseSumUb, const LocalTensor &lseOutputUb, + uint32_t dealRowCount, uint32_t actualCombineLoopSize, bool softmaxLseFlag, + bool learnableSinkFlag) +{ + if (dealRowCount <= 8) { + ComputeScaleValue_8(tmpSinkUb, lseMaxUb, lseSumUb, lseOutputUb, dealRowCount, actualCombineLoopSize, + softmaxLseFlag, learnableSinkFlag); + } else if (dealRowCount <= 16) { + ComputeScaleValue_16(tmpSinkUb, lseMaxUb, lseSumUb, lseOutputUb, dealRowCount, actualCombineLoopSize, + softmaxLseFlag, learnableSinkFlag); + } +} + +// gqa 非量化走这个模板函数,目前dealRowCount默认为8 +// lseResUb为ScaleValue的计算结果UB +template +__aicore__ inline void ComputeScaleValue_VF_FD(const LocalTensor &tmpSinkUb, const LocalTensor &lseMaxUb, + const LocalTensor &lseSumUb, const LocalTensor &lseResUb, + const LocalTensor &lseOutputUb, const LocalTensor &lseMaxUbTmp, + uint32_t dealRowCount, uint32_t actualCombineLoopSize, + bool softmaxLseFlag, bool learnableSinkFlag) +{ + ComputeScaleValue_8_FD(tmpSinkUb, lseMaxUb, lseSumUb, lseResUb, lseOutputUb, lseMaxUbTmp, dealRowCount, + actualCombineLoopSize, softmaxLseFlag, learnableSinkFlag); +} + +// 处理g<=8的场景 +template +__simd_vf__ void ComputeLogSumExp_8_VF(__ubuf__ T *srcSumLocalInt, __ubuf__ T *srcMaxLocalInt, __ubuf__ T *dstLocalInt, + uint32_t dealCount) +{ + Reg::RegTensor vregSum; + Reg::RegTensor vregMax; + Reg::RegTensor vregRes; + Reg::MaskReg pregTailN = Reg::UpdateMask(dealCount); + Reg::RegTensor vregMinValue; + Reg::RegTensor vregInfValue; + Reg::MaskReg pregCompare; + constexpr float infValue = 3e+99; // 3e+99 for float inf + constexpr uint32_t tmpMin = 0xFF167699; + float minValue = *((float *)&tmpMin); + Reg::Duplicate(vregMinValue, minValue); + Reg::Duplicate(vregInfValue, infValue); + + // 1.load to reg + Reg::LoadAlign(vregSum, (__ubuf__ float *&)srcSumLocalInt); + Reg::LoadAlign(vregMax, (__ubuf__ float *&)srcMaxLocalInt); + + // 2.LogSumExp + Reg::Log(vregRes, vregSum, pregTailN); + Reg::Add(vregRes, vregRes, vregMax, pregTailN); + + // 如果 softmaxMax 等于负无穷,则将 lse 结果置为 inf + Reg::Compare(pregCompare, vregMax, vregMinValue, pregTailN); + Reg::Select(vregRes, vregInfValue, vregRes, pregCompare); + + // 3.copy to ub + Reg::StoreAlign((__ubuf__ float *&)dstLocalInt, vregRes, pregTailN); +} + +template +__aicore__ inline void ComputeLogSumExp_8(const LocalTensor &dstTensor, const LocalTensor &softmaxSumTensor, + const LocalTensor &softmaxMaxTensor, uint32_t dealCount) +{ + __ubuf__ T *srcSumLocalInt = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr(); + __ubuf__ T *srcMaxLocalInt = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr(); + __ubuf__ T *dstLocalInt = (__ubuf__ T *)dstTensor.GetPhyAddr(); + + ComputeLogSumExp_8_VF(srcSumLocalInt, srcMaxLocalInt, dstLocalInt, dealCount); +} + +// 处理8 +__simd_vf__ void ComputeLogSumExp_16_VF(__ubuf__ T *srcSumUb, __ubuf__ T *srcSumUb2, __ubuf__ T *srcMaxUb, + __ubuf__ T *srcMaxUb2, __ubuf__ T *dstUb, __ubuf__ T *dstUb2, + uint32_t dealCount1, uint32_t dealCount2) +{ + Reg::RegTensor vregSum; + Reg::RegTensor vregSum2; + Reg::RegTensor vregMax; + Reg::RegTensor vregMax2; + Reg::RegTensor vregRes; + Reg::RegTensor vregRes2; + Reg::MaskReg pregTailN = Reg::UpdateMask(dealCount1); + Reg::MaskReg pregTailN2 = Reg::UpdateMask(dealCount2); + Reg::RegTensor vregMinValue; + Reg::RegTensor vregInfValue; + Reg::MaskReg pregCompare; + Reg::MaskReg pregCompare2; + constexpr float infValue = 3e+99; // 3e+99 for float inf + constexpr uint32_t tmpMin = 0xFF167699; + float minValue = *((float *)&tmpMin); + Reg::Duplicate(vregMinValue, minValue); + Reg::Duplicate(vregInfValue, infValue); + + // 1.load to reg + Reg::LoadAlign(vregSum, srcSumUb); + Reg::LoadAlign(vregSum2, srcSumUb2); + Reg::LoadAlign(vregMax, srcMaxUb); + Reg::LoadAlign(vregMax2, srcMaxUb2); + + // 2.LogSumExp + Reg::Log(vregRes, vregSum, pregTailN); + Reg::Log(vregRes2, vregSum2, pregTailN2); + Reg::Add(vregRes, vregRes, vregMax, pregTailN); + Reg::Add(vregRes2, vregRes2, vregMax2, pregTailN2); + + // 如果 softmaxMax 等于负无穷,则将 lse 结果置为 inf + Reg::Compare(pregCompare, vregMax, vregMinValue, pregTailN); + Reg::Compare(pregCompare2, vregMax2, vregMinValue, pregTailN2); + Reg::Select(vregRes, vregInfValue, vregRes, pregCompare); + Reg::Select(vregRes2, vregInfValue, vregRes2, pregCompare2); + + // 3.copy to ub + Reg::StoreAlign(dstUb, vregRes, pregTailN); + Reg::StoreAlign(dstUb2, vregRes2, pregTailN2); +} + +template +__aicore__ inline void ComputeLogSumExp_16(const LocalTensor &dstTensor, const LocalTensor &softmaxSumTensor, + const LocalTensor &softmaxMaxTensor, uint32_t dealCount) +{ + __ubuf__ T *srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr(); + __ubuf__ T *srcSumUb2 = srcSumUb + 64; // 一个寄存器最多处理64个数 + __ubuf__ T *srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr(); + __ubuf__ T *srcMaxUb2 = srcMaxUb + 64; + __ubuf__ T *dstUb = (__ubuf__ T *)dstTensor.GetPhyAddr(); + __ubuf__ T *dstUb2 = dstUb + 64; + uint32_t dealCount1 = 8 * 8; + uint32_t dealCount2 = dealCount - dealCount1; + + ComputeLogSumExp_16_VF(srcSumUb, srcSumUb2, srcMaxUb, srcMaxUb2, dstUb, dstUb2, dealCount1, dealCount2); +} + +template +__aicore__ inline void ComputeLogSumExp_VF(const LocalTensor &dstTensor, const LocalTensor &softmaxSumTensor, + const LocalTensor &softmaxMaxTensor, uint32_t dealRowCount) +{ + if (dealRowCount <= 8) { + ComputeLogSumExp_8(dstTensor, softmaxSumTensor, softmaxMaxTensor, dealRowCount * 8); // 8:FP32 in one block + } else if (dealRowCount <= 16) { + ComputeLogSumExp_16(dstTensor, softmaxSumTensor, softmaxMaxTensor, dealRowCount * 8); // 8:FP32 in one block + } +} + +} // namespace FaVectorApi + +#endif // MY_FLASH_DECODE_ARCH35_H diff --git a/csrc/attention/common/op_kernel/attn_buffer.h b/csrc/attention/common/op_kernel/attn_buffer.h new file mode 100644 index 000000000000..783f3ca47424 --- /dev/null +++ b/csrc/attention/common/op_kernel/attn_buffer.h @@ -0,0 +1,454 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file buffer.h + * \brief同步管理 + */ +#ifndef BUFFER_H +#define BUFFER_H +#include +#include "lib/matmul_intf.h" +#if ASC_DEVKIT_MAJOR >= 9 +#include "kernel_basic_intf.h" +#else +#include "kernel_operator.h" +#endif +using namespace AscendC; +namespace fa_base_matmul { +__BLOCK_LOCAL__ __inline__ uint32_t idCounterNum; +#define MAKE_ID ((++idCounterNum) % 11) + +__aicore__ inline void ResetIdCounter() +{ + idCounterNum = 0; +} + +// 核间同步中,AIC(flagId 0-10)对应AIV0(flagId 0-10),对应AIV1(flagId 16-26) +#define AIV0_AIV1_OFFSET 16 + +enum class BufferType { + L1 = 0, + L0A = 1, + L0B = 2, + L0C = 3, + UB = 4, + GM = 5, + C2 = 6, +}; + +enum class SyncType { + NO_SYNC, + INNER_CORE_SYNC, + CROSS_CORE_SYNC_FORWARD, + CROSS_CORE_SYNC_BOTH, + CROSS_CORE_SYNC_BACKWARD, +}; + +enum class SyncMode { + SET_WAIT_FLAG, + LOCK_UNLOCK, +}; + +enum class IdSource { + INTERNAL, // ID由接口内部分配 + EXTERNAL, // ID由接口外部传入(用户指定) +}; + +constexpr uint32_t INVALID_CROSS_CORE_EVENT_ID = 16; +static constexpr uint64_t CROSS_CORE_SYNC_MODE = 4; + +template +struct BufferInfo { + // Cons 消费者,Prod 生产者 + __aicore__ const static constexpr HardEvent ConsWaitProdStatus() + { + if constexpr (Type == BufferType::L1) { + return HardEvent::MTE2_MTE1; + } else if constexpr (Type == BufferType::L0A) { + return HardEvent::MTE1_M; + } else if constexpr (Type == BufferType::L0B) { + return HardEvent::MTE1_M; + } else if constexpr (Type == BufferType::L0C) { + return HardEvent::M_FIX; + } else if constexpr (Type == BufferType::C2) { + return HardEvent::MTE1_M; + } else if constexpr (Type == BufferType::GM) { + return HardEvent::MTE2_S; + } + } + + __aicore__ const static constexpr HardEvent ProdWaitConsStatus() + { + if constexpr (Type == BufferType::L1) { + return HardEvent::MTE1_MTE2; + } else if constexpr (Type == BufferType::L0A) { + return HardEvent::M_MTE1; + } else if constexpr (Type == BufferType::L0B) { + return HardEvent::M_MTE1; + } else if constexpr (Type == BufferType::L0C) { + return HardEvent::FIX_M; + } else if constexpr (Type == BufferType::C2) { + return HardEvent::M_MTE1; + } else if constexpr (Type == BufferType::GM) { + return HardEvent::S_MTE2; + } + } + + __aicore__ const static constexpr pipe_t GetProdPipe() + { + if constexpr (Type == BufferType::L1) { + return PIPE_MTE2; + } else if constexpr (Type == BufferType::L0A) { + return PIPE_MTE1; + } else if constexpr (Type == BufferType::L0B) { + return PIPE_MTE1; + } else if constexpr (Type == BufferType::L0C) { + return PIPE_M; + } + } + + __aicore__ const static constexpr pipe_t GetConsPipe() + { + if constexpr (Type == BufferType::L1) { + return PIPE_MTE1; + } else if constexpr (Type == BufferType::L0A) { + return PIPE_M; + } else if constexpr (Type == BufferType::L0B) { + return PIPE_M; + } else if constexpr (Type == BufferType::L0C) { + return PIPE_FIX; + } + } + + __aicore__ const static constexpr TPosition GetTPosition() + { + if constexpr (Type == BufferType::L1) { + return TPosition::A1; + } else if constexpr (Type == BufferType::L0A) { + return TPosition::A2; + } else if constexpr (Type == BufferType::L0B) { + return TPosition::B2; + } else if constexpr (Type == BufferType::L0C) { + return TPosition::CO1; + } else if constexpr (Type == BufferType::UB) { + return TPosition::VECIN; + } else if constexpr (Type == BufferType::GM) { + return TPosition::GM; + } else if constexpr (Type == BufferType::C2) { + return TPosition::C2; + } + } + + static constexpr HardEvent EventP2C = + ConsWaitProdStatus(); // 生产者到消费者方向的HardEvent:消费者等生产者提供/生产者通知消费者已生成 + static constexpr HardEvent EventC2P = + ProdWaitConsStatus(); // 消费者到生产者方向的HardEvent:生产者等消费者消耗/消费者通知生产者已消耗 + static constexpr pipe_t ProdPipe = GetProdPipe(); // 使用Mutex时的生产者PIPE + static constexpr pipe_t ConsPipe = GetConsPipe(); // 使用Mutex时的消费者PIPE + static constexpr TPosition Position = GetTPosition(); +}; + +// buffer绑定生产者、消费者关系 +// L1 buffer的生产者为MTE2或者MTE3,消费者为MTE1 +// L0A buffer的生产者为MTE1,消费者为M +// L0B buffer的生产者为MTE1,消费者为M +// L0C buffer的生产者为M,消费者为FIX +template +class Buffer { + using TensorType = std::conditional_t, LocalTensor>; + + template + using TargetTensorType = std::conditional_t, LocalTensor>; + +public: + __aicore__ inline Buffer() {} + __aicore__ inline Buffer(TensorType tensor, uint32_t size) + { + tensor_ = tensor; + size_ = size; + if constexpr (syncType == SyncType::CROSS_CORE_SYNC_FORWARD) { + id0_ = MAKE_ID; + id1_ = INVALID_CROSS_CORE_EVENT_ID; + } else if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) { + id0_ = INVALID_CROSS_CORE_EVENT_ID; + id1_ = MAKE_ID; + } else if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) { + id0_ = MAKE_ID; + id1_ = MAKE_ID; + } else { + id0_ = INVALID_CROSS_CORE_EVENT_ID; + id1_ = INVALID_CROSS_CORE_EVENT_ID; + } + } + + template + __aicore__ inline void Init() + { + static_assert(idSource == IdSource::INTERNAL, "idSource should IdSource::INTERNAL."); + if ASCEND_IS_AIC { + if constexpr (syncType == SyncType::INNER_CORE_SYNC && idSource == IdSource::INTERNAL) { + if constexpr (syncMode == SyncMode::SET_WAIT_FLAG) { + p2cEventId_ = GetTPipePtr()->AllocEventID::EventP2C>(); // 确保只能被调用一次 + c2pEventId_ = GetTPipePtr()->AllocEventID::EventC2P>(); + SetFlag::EventC2P>(c2pEventId_); + } else if constexpr (syncMode == SyncMode::LOCK_UNLOCK) { +#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510) + mutexId_ = AllocMutexID(); +#endif + } + } + } + } + + template + __aicore__ inline void Init(uint32_t id) + { + static_assert(idSource == IdSource::EXTERNAL, "idSource should IdSource::EXTERNAL."); + if ASCEND_IS_AIC { + if constexpr (syncType == SyncType::INNER_CORE_SYNC && idSource == IdSource::EXTERNAL) { + if constexpr (syncMode == SyncMode::SET_WAIT_FLAG) { + // 静态Tensor场景, 用户自己指定EventId + p2cEventId_ = id; + c2pEventId_ = id; + SetFlag::EventC2P>(c2pEventId_); + } else if constexpr (syncMode == SyncMode::LOCK_UNLOCK) { +#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510) + mutexId_ = id; +#endif + } + } + } + } + + template + __aicore__ inline void UnInit() + { + if ASCEND_IS_AIC { + if constexpr (syncType == SyncType::INNER_CORE_SYNC) { + if constexpr (syncMode == SyncMode::SET_WAIT_FLAG) { + WaitFlag::EventC2P>(c2pEventId_); + if constexpr (idSource == IdSource::INTERNAL) { + GetTPipePtr()->ReleaseEventID::EventP2C>( + p2cEventId_); // 确保只能被调用一次 + GetTPipePtr()->ReleaseEventID::EventC2P>(c2pEventId_); + } + } else if constexpr (syncMode == SyncMode::LOCK_UNLOCK) { +#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510) + if constexpr (idSource == IdSource::INTERNAL) { + ReleaseMutexID(mutexId_); + } +#endif + } + } + } + } + + template + __aicore__ inline void Lock() + { +#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510) + if ASCEND_IS_AIC { + if constexpr (syncType == SyncType::INNER_CORE_SYNC && syncMode == SyncMode::LOCK_UNLOCK) { + Mutex::Lock(mutexId_); + } + } +#endif + } + + template + __aicore__ inline void Unlock() + { +#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510) + if ASCEND_IS_AIC { + if constexpr (syncType == SyncType::INNER_CORE_SYNC && syncMode == SyncMode::LOCK_UNLOCK) { + Mutex::Unlock(mutexId_); + } + } +#endif + } + + template + __aicore__ inline void Wait() + { + if ASCEND_IS_AIC { + if constexpr (syncType == SyncType::INNER_CORE_SYNC) { + if constexpr (syncMode == SyncMode::SET_WAIT_FLAG) { + if constexpr (EventType == BufferInfo::EventP2C) { + WaitFlag::EventP2C>(p2cEventId_); // 消费者等待生产者完成生产 + } else { + WaitFlag::EventC2P>(c2pEventId_); // 生产者等待消费者完成消费 + } + } else if constexpr (syncMode == SyncMode::LOCK_UNLOCK) { +#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510) + if constexpr (EventType == BufferInfo::EventP2C) { + Mutex::Lock::ConsPipe>(mutexId_); // 消费者加锁 + } else { + Mutex::Lock::ProdPipe>(mutexId_); // 生产者加锁 + } +#endif + } + } + } + } + + template + __aicore__ inline void Set() + { + if ASCEND_IS_AIC { + if constexpr (syncType == SyncType::INNER_CORE_SYNC) { + if constexpr (syncMode == SyncMode::SET_WAIT_FLAG) { + if constexpr (EventType == BufferInfo::EventP2C) { + SetFlag::EventP2C>(p2cEventId_); // 生产者通知消费者已完成生产 + } else { + SetFlag::EventC2P>(c2pEventId_); // 消费者通知生产者已完成消费 + } + } else if constexpr (syncMode == SyncMode::LOCK_UNLOCK) { +#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510) + if constexpr (EventType == BufferInfo::EventP2C) { + Mutex::Unlock::ProdPipe>(mutexId_); // 生产者解锁 + } else { + Mutex::Unlock::ConsPipe>(mutexId_); // 消费者解锁 + } +#endif + } + } + } + } + + template + __aicore__ inline void SetEventID() + { + static_assert(idSource == IdSource::INTERNAL, "idSource should IdSource::INTERNAL."); + if ASCEND_IS_AIC { + if constexpr (idSource == IdSource::INTERNAL && syncMode == SyncMode::SET_WAIT_FLAG) { + p2cEventId_ = GetTPipePtr()->AllocEventID::EventP2C>(); // 确保只能被调用一次 + c2pEventId_ = GetTPipePtr()->AllocEventID::EventC2P>(); + } + } + } + + template + __aicore__ inline TEventID GetEventID() + { + if ASCEND_IS_AIC { + if constexpr (EventType == BufferInfo::EventP2C) { + return p2cEventId_; // 生产者通知消费者已完成生产 + } else { + return c2pEventId_; // 消费者通知生产者已完成消费 + } + } + } + + __aicore__ inline void SetCrossCoreID(uint32_t id0, uint32_t id1) + { + id0_ = id0; + id1_ = id1; + } + + template + __aicore__ inline void WaitCrossCore() + { + if constexpr (bufferType == BufferType::GM && syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) { + // AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV + if ASCEND_IS_AIC { + CrossCoreWaitFlag(id1_); + CrossCoreWaitFlag(id1_ + AIV0_AIV1_OFFSET); + } else { + CrossCoreWaitFlag(id0_); + } + } else if constexpr (bufferType == BufferType::UB || bufferType == BufferType::GM) { + // AIC属于生产者,AIV属于消费者,且一个AIC对应两个AIV + if ASCEND_IS_AIC { + CrossCoreWaitFlag(id1_); + CrossCoreWaitFlag(id1_ + AIV0_AIV1_OFFSET); + } else { + if constexpr (isReuse) { + CrossCoreWaitFlag(id0_); + } else { + CrossCoreWaitFlag(id0_); + } + } + } else if constexpr (bufferType == BufferType::L1) { + // AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV + if ASCEND_IS_AIC { + CrossCoreWaitFlag(id0_); + CrossCoreWaitFlag(id0_ + AIV0_AIV1_OFFSET); + } else { + if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) { + CrossCoreWaitFlag(id1_); + } + } + } + } + + template + __aicore__ inline void SetCrossCore() + { + if constexpr (bufferType == BufferType::GM && syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) { + // AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV + if ASCEND_IS_AIC { + CrossCoreSetFlag(id0_); + CrossCoreSetFlag(id0_ + AIV0_AIV1_OFFSET); + } else { + CrossCoreSetFlag(id1_); + } + } else if constexpr (bufferType == BufferType::UB || bufferType == BufferType::GM) { + // AIC属于生产者,AIV属于消费者,且一个AIC对应两个AIV + if ASCEND_IS_AIC { + CrossCoreSetFlag(id0_); + CrossCoreSetFlag(id0_ + AIV0_AIV1_OFFSET); + } else { + if constexpr (isReuse) { + CrossCoreSetFlag(id1_); + } else { + CrossCoreSetFlag(id1_); + } + } + } else if constexpr (bufferType == BufferType::L1) { + // AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV + if ASCEND_IS_AIC { + if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) { + CrossCoreSetFlag(id1_); + CrossCoreSetFlag(id1_ + AIV0_AIV1_OFFSET); + } + } else { + CrossCoreSetFlag(id0_); + } + } + } + + template + __aicore__ inline TargetTensorType GetTensor() + { + return tensor_.template ReinterpretCast(); + } + + template + __aicore__ inline TargetTensorType GetTensor(uint64_t startindex) + { + TargetTensorType tmpTensor = tensor_.template ReinterpretCast(); + return tmpTensor[startindex]; + } + +private: + TensorType tensor_; + uint32_t size_; + TEventID p2cEventId_; + TEventID c2pEventId_; + uint32_t id0_; // 用作正向同步:生产者通知消费者,或者消费者等待生产者; + uint32_t id1_; // 用作反向同步:消费者通知生产者,或者生产者等待消费者; +#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510) + MutexID mutexId_; +#endif +}; +} // namespace fa_base_matmul +#endif diff --git a/csrc/attention/common/op_kernel/attn_buffer_manager.h b/csrc/attention/common/op_kernel/attn_buffer_manager.h new file mode 100644 index 000000000000..b71200eed08e --- /dev/null +++ b/csrc/attention/common/op_kernel/attn_buffer_manager.h @@ -0,0 +1,73 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file buffer_manager.h + * \brief buffer内存管理 + */ +#ifndef BUFFER_MANAGER_H +#define BUFFER_MANAGER_H + +#if (__NPU_ARCH__ == 5102) +#include "buffer_mix_core.h" +#else +#include "attn_buffer.h" +#endif + +// L1 TPosition::A1 +// L0A TPosition::A2 +// L0B TPosition::B2 +// L0C TPosition::CO1 +// UB TPosition::VECIN +namespace fa_base_matmul { +template +class BufferManager { + using TensorType = std::conditional_t, LocalTensor>; + +public: + __aicore__ inline void Init(TPipe *pipe, uint32_t size) + { + static_assert(bufferType != BufferType::GM, "GM should use workspace."); + TBuf::Position> tbuf; + pipe->InitBuffer(tbuf, size); + mem_ = tbuf.template Get(); + } + + // 静态Tensor使用这个函数 + __aicore__ inline void Init(uint32_t size) + { + static_assert(bufferType != BufferType::GM, "GM should use workspace."); + mem_ = LocalTensor(BufferInfo::Position, 0, size); + } + + __aicore__ inline void Init(__gm__ uint8_t *workspace) + { + static_assert(bufferType == BufferType::GM, "BufferType should be GM."); + mem_.SetGlobalBuffer((__gm__ uint8_t *)workspace); + } + + template + __aicore__ inline Buffer AllocBuffer(uint32_t size) + { + TensorType temp = mem_[offset_]; + offset_ += size; + return Buffer(temp, size); + } + + template + __aicore__ inline void FreeBuffer(Buffer &buffer) + {} + +private: + uint32_t offset_ = 0; + TensorType mem_; +}; +} // namespace fa_base_matmul +#endif diff --git a/csrc/attention/common/op_kernel/init_output.h b/csrc/attention/common/op_kernel/init_output.h new file mode 100644 index 000000000000..6d3d6e7f1061 --- /dev/null +++ b/csrc/attention/common/op_kernel/init_output.h @@ -0,0 +1,90 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file init_output.h + * \brief + */ + +#ifndef INIT_OUTPUT_H +#define INIT_OUTPUT_H + +#if ASC_DEVKIT_MAJOR >= 9 +#include "kernel_vec_intf.h" +#include "kernel_cube_intf.h" +#else +#include "kernel_operator.h" +#endif + +namespace AttentionCommon { + +template +__aicore__ inline void InitOutput(GlobalTensor outGm, uint64_t totalElementNum, uint32_t vecCoreNum, T initValue) +{ + if ASCEND_IS_AIV { + uint64_t singleCoreMaxElementNum = (totalElementNum + vecCoreNum - 1U) / vecCoreNum; + uint32_t tmpBlockIdx = AscendC::GetBlockIdx(); + uint64_t gmOffset = tmpBlockIdx * singleCoreMaxElementNum; + if (gmOffset < totalElementNum) { + uint64_t singleCoreActualElementNum = (gmOffset + singleCoreMaxElementNum > totalElementNum) ? + (totalElementNum - gmOffset) : + singleCoreMaxElementNum; + LocalTensor popBuffer = + AscendC::LocalTensor(TPosition::VECIN, POP_BUF_START_ADDR, POP_BUF_ELE_SIZE * sizeof(T)) + .template ReinterpretCast(); + + if constexpr (ENABLE_LOCK) { + Mutex::Lock(SYNC_ID); + } else { + AscendC::SetFlag(SYNC_ID); + + AscendC::WaitFlag(SYNC_ID); + } + AscendC::Duplicate(popBuffer, initValue, POP_BUF_ELE_SIZE); + if constexpr (ENABLE_LOCK) { + Mutex::Unlock(SYNC_ID); + } else { + AscendC::SetFlag(SYNC_ID); + } + + if constexpr (ENABLE_LOCK) { + Mutex::Lock(SYNC_ID); + } else { + AscendC::WaitFlag(SYNC_ID); + } + uint64_t loopCnt = singleCoreActualElementNum / POP_BUF_ELE_SIZE; + uint64_t tailSize = singleCoreActualElementNum - loopCnt * POP_BUF_ELE_SIZE; + for (uint64_t loop = 0; loop < loopCnt; loop++) { + AscendC::DataCopy(outGm[gmOffset], popBuffer, POP_BUF_ELE_SIZE); + gmOffset += POP_BUF_ELE_SIZE; + } + if (tailSize > 0) { + AscendC::DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1; + dataCopyParams.blockLen = tailSize * sizeof(T); + dataCopyParams.srcStride = 0; + dataCopyParams.dstStride = 0; + AscendC::DataCopyPad(outGm[gmOffset], popBuffer, dataCopyParams); + } + if constexpr (ENABLE_LOCK) { + Mutex::Unlock(SYNC_ID); + } else { + AscendC::SetFlag(SYNC_ID); + + AscendC::WaitFlag(SYNC_ID); + } + } + } +} + +} // namespace AttentionCommon + +#endif // INIT_OUTPUT_H diff --git a/csrc/attention/compressor/op_host/arch35/compressor_tiling.h b/csrc/attention/compressor/op_host/arch35/compressor_tiling.h index fb2be863b9bc..12aa8c3079fc 100644 --- a/csrc/attention/compressor/op_host/arch35/compressor_tiling.h +++ b/csrc/attention/compressor/op_host/arch35/compressor_tiling.h @@ -106,10 +106,11 @@ static const std::string CMP_KV_NAME = "cmp_kv"; static std::string DataTypeToSerialString(ge::DataType type); +// Keep host validation aligned with the selected TH/BF16/FP32-RoPE templates. const std::map> DTYPE_SUPPORT_MAP = { - {X_NAME, {ge::DT_BF16, ge::DT_FLOAT16}}, - {WKV_NAME, {ge::DT_BF16, ge::DT_FLOAT16}}, - {WGATE_NAME, {ge::DT_BF16, ge::DT_FLOAT16}}, + {X_NAME, {ge::DT_BF16}}, + {WKV_NAME, {ge::DT_BF16}}, + {WGATE_NAME, {ge::DT_BF16}}, {STATE_CACHE_NAME, {ge::DT_FLOAT}}, {APE_NAME, {ge::DT_FLOAT}}, {NORM_WEIGHT_NAME, {ge::DT_FLOAT}}, @@ -119,11 +120,11 @@ const std::map> DTYPE_SUPPORT_MAP = { {CU_SEQLENS_NAME, {ge::DT_INT32}}, {SEQUSED_NAME, {ge::DT_INT32}}, {START_POS_NAME, {ge::DT_INT32}}, - {CMP_KV_NAME, {ge::DT_BF16, ge::DT_FLOAT16}} + {CMP_KV_NAME, {ge::DT_BF16}} }; const std::map> DIM_NUM_MAP = { - {X_NAME, {COMPRESSOR_DIM_NUM_2, COMPRESSOR_DIM_NUM_3}}, + {X_NAME, {COMPRESSOR_DIM_NUM_2}}, {WKV_NAME, {COMPRESSOR_DIM_NUM_2}}, {WGATE_NAME, {COMPRESSOR_DIM_NUM_2}}, {STATE_CACHE_NAME, {COMPRESSOR_DIM_NUM_3}}, @@ -223,9 +224,9 @@ struct CompressorBaseShapeInfo { const std::vector ROPE_HEAD_DIM {64}; const std::vector COFF {1, 2}; const std::vector CMP_RATIO {2, 4, 8, 16, 32, 64, 128}; -const std::vector ROTARY_MODE {1, 2}; +const std::vector ROTARY_MODE {2}; const std::vector HEAD_DIM {128, 512}; -const std::vector CACHE_MODE {1, 2}; +const std::vector CACHE_MODE {1}; enum class ROTARY_MODE:uint8_t { HALF = 1, diff --git a/csrc/attention/compressor/op_kernel/arch35/compressor_template_tiling_key.h b/csrc/attention/compressor/op_kernel/arch35/compressor_template_tiling_key.h index b7b423a398f1..56f926635f81 100644 --- a/csrc/attention/compressor/op_kernel/arch35/compressor_template_tiling_key.h +++ b/csrc/attention/compressor/op_kernel/arch35/compressor_template_tiling_key.h @@ -35,7 +35,7 @@ ASCENDC_TPL_ARGS_DECL(compressor, // 算子唯一标识,与opType保持一致 ASCENDC_TPL_UINT_DECL(ROTARY_MODE, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST, 1, 2), // bit:9-10 cache_mode 1:CONTINUOUS 2:cycle ASCENDC_TPL_UINT_DECL(CACHE_MODE, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST, 1, 2), - // bit:11-12 template_id 0:empty_tensor 1:normal 2:full load + // bit:11-12 template_id 0:normal 1:empty_tensor 2:full load ASCENDC_TPL_UINT_DECL(TEMPLATE_ID, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST, 0, 1, 2), ); diff --git a/csrc/attention/compressor/op_kernel/compressor.cpp b/csrc/attention/compressor/op_kernel/compressor.cpp index c6f4a6c4ceed..74615e62cd53 100644 --- a/csrc/attention/compressor/op_kernel/compressor.cpp +++ b/csrc/attention/compressor/op_kernel/compressor.cpp @@ -55,7 +55,7 @@ __global__ __aicore__ void compressor( GET_TILING_DATA_WITH_STRUCT(optiling::CompressorTilingData, tilingDataIn, tiling); if constexpr (static_cast(TemplateId) == TEMPLATE_ID::EMPTY_X) { return; - } + } else { const optiling::CompressorTilingData *__restrict tilingData = &tilingDataIn; TPipe pipe; constexpr auto xLayout = static_cast(XLayout); @@ -76,4 +76,5 @@ __global__ __aicore__ void compressor( INVOKE_COMPRESSOR_GENERAL_OP_IMPL(CompressorKernel, xLayout, xDtype, coff, rotaryMode, cacheMode); } #endif -} \ No newline at end of file + } +} diff --git a/csrc/attention/quant_lightning_indexer_v2/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2/CMakeLists.txt index 93f00822c1c8..e99a153f311b 100644 --- a/csrc/attention/quant_lightning_indexer_v2/CMakeLists.txt +++ b/csrc/attention/quant_lightning_indexer_v2/CMakeLists.txt @@ -1,3 +1,13 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) if(NOT ENABLE_TEST AND NOT BENCHMARK) list(REMOVE_ITEM CURRENT_DIRS tests) diff --git a/csrc/attention/quant_lightning_indexer_v2/docs/aclnnQuantLightningIndexerV2.md b/csrc/attention/quant_lightning_indexer_v2/docs/aclnnQuantLightningIndexerV2.md new file mode 100644 index 000000000000..13d3f67f3464 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/docs/aclnnQuantLightningIndexerV2.md @@ -0,0 +1,898 @@ +# aclnnQuantLightningIndexerV2 + +[📄 查看源码](https://gitcode.com/cann/ops-transformer/tree/master/attention/quant_lightning_indexer_v2) + +## 产品支持情况 + + +- Ascend 950PR/Ascend 950DT:支持 + + +- Atlas A3 训练系列产品/Atlas A3 推理系列产品:支持 + + +- Atlas A2 训练系列产品/Atlas A2 推理系列产品:支持 + + +- Atlas 200I/500 A2 推理产品:不支持 + + +- Atlas 推理系列产品:不支持 + + +- Atlas 训练系列产品:不支持 + + +## 功能说明 + +- 接口功能:`QuantLightningIndexerV2`是推理场景下,稀疏attention前处理的计算,选出关键的稀疏token,并对输入query和key进行量化实现存8算8,获取最大收益。 + +- 版本演进:在QuantLightningIndexer的基础上,新增压缩key场景、分核计算metadata、稀疏value输出等能力。 + +- 计算公式: + +$$ +out = \text{Top-}k\left\{[1]_{1\times g}@\left[(W@[1]_{1\times S_{k}})\odot\text{ReLU}\left(\left(Scale_Q@Scale_K^T\right)\odot\left(Q_{index}^{Quant}@{\left(K_{index}^{Quant}\right)}^T\right)\right)\right]\right\} +$$ + +主要计算过程为: + +1. 将某个token对应的输入参数`query`($Q_{index}^{Quant}\in\R^{g\times d}$)乘以给定上下文`key`($K_{index}^{Quant}\in\R^{S_{k}\times d}$),得到相关性。 +2. 相关性结果与`query`和`key`对应的反量化系数`query_dequant_scale`($Scale_Q$)和`key_dequant_scale`($Scale_K^T$)相乘,通过激活函数$ReLU$过滤无效负相关信号后,得到当前Token与所有前序Token的相关性分数向量。 +3. 将其与权重系数`weights`($W$)相乘后,沿g的方向,选取前$Top-k$个索引值得到输出$out$,作为Attention的输入。 + +## 函数原型 + +每个算子分为[两段式接口](../../../docs/zh/context/two_phase_api.md),必须先调用"aclnnQuantLightningIndexerV2GetWorkspaceSize"接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用"aclnnQuantLightningIndexerV2"接口执行计算。 + +```Cpp +aclnnStatus aclnnQuantLightningIndexerV2GetWorkspaceSize( + const aclTensor *query, + const aclTensor *key, + const aclTensor *weights, + const aclTensor *queryDequantScale, + const aclTensor *keyDequantScale, + const aclTensor *cuSeqLensQOptional, + const aclTensor *cuSeqLensKOptional, + const aclTensor *sequsedQOptional, + const aclTensor *sequsedKOptional, + const aclTensor *cmpResidualKOptional, + const aclTensor *blockTableOptional, + const aclTensor *outputIdxOffsetOptional, + const aclTensor *metadataOptional, + int64_t topk, + int64_t quantMode, + int64_t maxSeqlenQOptional, + char *layoutQOptional, + char *layoutKOptional, + int64_t maskModeOptional, + int64_t cmpRatioOptional, + int64_t returnValueOptional, + const aclTensor *sparseIndicesOut, + const aclTensor *sparseValuesOut, + uint64_t *workspaceSize, + aclOpExecutor **executor) +``` + +```Cpp +aclnnStatus aclnnQuantLightningIndexerV2( + void *workspace, + uint64_t workspaceSize, + aclOpExecutor *executor, + const aclrtStream stream) +``` + +## aclnnQuantLightningIndexerV2GetWorkspaceSize + +- **参数说明:** + +> [!NOTE] +> +> - query、key、weights参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。 +> - S1表示query shape中的S,S2表示key shape中的S,T1表示query shape中的T,N1表示query shape中的N,N2表示key shape中的N。 +> - maxBlockNumPerSeq表示每个Batch中最大sequsedK对应的block数量,S2_MAX表示sequsedK中的最大值 + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensor
query输入公式中量化后的 Query。不支持空tensor。INT8、FLOAT8_e4m3fn、HIFLOAT8、FLOAT4_e2m1ND +
    +
  • layout_query为BSND时,shape为(B,S1,N1,D)。
  • +
  • layout_query为TND时,shape为(T1,N1,D)。
  • +
+
x
key输入公式中量化后的 Key。 +
    +
  • 不支持空tensor。
  • +
  • block_num为PageAttention时block总数,block_size为一个block的token数。
  • +
  • layout_key为PA_BSND时,shape为(block_num, block_size, N2, D)。
  • +
  • layout_key为BSND时,shape为(B, K_S, N2, D),layout_key为TND时,shape为(K_T, N2, D)。
  • +
+
INT8、FLOAT8_e4m3fn、HIFLOAT8、FLOAT4_e2m1ND +
    +
  • layout_key为PA_BSND时,shape为(block_num, block_size, N2, D)。
  • +
+
支持0轴非连续
weights输入公式中的权重系数 W。不支持空tensor。FLOAT16、FLOAT32ND +
    +
  • layout_query为BSND时,shape为(B,S1,N1)。
  • +
  • layout_query为TND时,shape为(T1,N1)。
  • +
+
x
queryDequantScale输入公式中 Query 的反量化系数。不支持空tensor。FLOAT16、FLOAT32、FLOAT8_e8m0ND +
    +
  • quantMode为3/5时,layout_query为BSND时shape为(B,S1,N1,D/64,2),layout_query为TND时shape为(T1,N1,D/64,2)。
  • +
  • quantMode为4时,shape为(1,)。
  • +
  • 其他场景shape与weights保持一致。
  • +
+
x
keyDequantScale输入公式中 Key 的反量化系数。不支持空tensor。FLOAT16、FLOAT32、FLOAT8_e8m0ND +
    +
  • quantMode为3/5时,layout_key为PA_BSND、BSND、TND对应的shape分别为(block_num,block_size,N2,D/64,2)、(B,K_S,N2,D/64,2)、(K_T,N2,D/64,2)。
  • +
  • quantMode为4时,shape为(1,)。
  • +
  • 其他场景下,layout_key为PA_BSND、BSND、TND对应的shape分别为(block_num,block_size,N2)、(B,K_S,N2)、(K_T,N2)。
  • +
+
支持0轴非连续
cuSeqLensQOptional输入每个Batch中,Query的有效token数(TND场景使用cu_seqlens格式)。 +
    +
  • 当layout_query为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和。
  • +
+
INT32ND(B+1,)x
cuSeqLensKOptional输入每个Batch中,Key的有效token数(TND场景使用cu_seqlens格式)。 +
    +
  • 当layout_key为TND时,该入参必须传入。
  • +
+
INT32ND(B+1,)x
sequsedQOptional输入每个Batch中,Query的有效token数(BSND场景使用seqused格式)。该入参中每个Batch的有效token数不超过query中的维度S大小且不小于0。INT32ND(B,)x
sequsedKOptional输入每个Batch中,Key的有效token数(BSND场景使用seqused格式)。 +
    +
  • 该入参中每个Batch的有效token数不超过key中的维度S大小且不小于0。
  • +
  • 当layout_key为PA_BSND时,该入参必须传入。
  • +
+
INT32ND(B,)x
cmpResidualKOptional输入压缩场景下Key的残余长度。需满足0 <= cmpResidualKOptional[i] < cmpRatioOptional。INT32ND(B,)x
blockTableOptional输入表示PageAttention中KV存储使用的block映射表。 +
    +
  • 不支持空tensor。
  • +
  • PageAttention场景下,block_table必须为二维,第一维长度需要等于B,第二维长度不能小于maxBlockNumPerSeq。
  • +
+
INT32ND(B, S2_MAX/block_size)x
outputIdxOffsetOptional输入输出索引的偏移量。-INT32NDlayout_query为BSND时shape为(B,S1,N2),layout_query为TND时shape为(T1,N2)。x
metadataOptional输入QuantLightningIndexerV2Metadata算子传入的分核信息。 +
    +
  • 包含使用核数、分块大小以及每个核处理数据的起始点等内容。
  • +
  • shape大小为[1024],当前不支持传空。
  • +
+
INT32ND(1024,)x
topk输入topK阶段需要保留的Key token索引数量。支持[1, 8192]。INT64---
quantMode输入量化模式。 +
    +
  • 支持传入 1(FLOAT8_e4m3fn量化)、2(Per-Token-Head量化)、3(MXFP8量化)、4(HIFLOAT8量化)、5(MXFP4量化)。
  • +
+
INT64---
maxSeqlenQOptional输入Query的最大序列长度。-INT64---
layoutQOptional输入用于标识输入Query的数据排布格式。 +
    +
  • 支持BSND、TND。
  • +
+
STRING---
layoutKOptional输入用于标识输入Key的数据排布格式。 +
    +
  • 支持 PA_BSND、BSND、TND。
  • +
+
STRING---
maskModeOptional输入表示sparse的模式。 +
    +
  • 0代表defaultMask模式。
  • +
  • 3代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景。
  • +
+
INT64---
cmpRatioOptional输入key的压缩倍数。 +
    +
  • 支持 (0, 128] 内的正整数。
  • +
+
INT64---
returnValueOptional输入表示是否输出sparseValuesOut。 +
    +
  • 1表示输出,0表示不输出。
  • +
+
INT64---
sparseIndicesOut输出公式中的Indices输出。不支持空tensor。INT32ND +
    +
  • layout_query为"BSND"时输出shape为[B, S1, N2, topk]。
  • +
  • layout_query为"TND"时输出shape为[T1, N2, topk]。
  • +
+
x
sparseValuesOut输出公式中的Indices输出对应的value值。 +
    +
  • returnValue为1时输出有效值,无效部分填bf16负无穷;returnValue为0时输出shape为(0,)的空tensor。
  • +
+
BFLOAT16NDreturnValue为1时shape与sparseIndicesOut保持一致;returnValue为0时shape为(0,)。x
workspaceSize输出返回需要在Device侧申请的workspace大小。-----
executor输出返回op执行器,包含了算子计算流程。-----
+ + +- Ascend 950PR/Ascend 950DT: + - `layout_key` 额外支持 BSND 和 TND;支持 PA_BSND、BSND、TND。 + - `quant_mode` 支持 1(FLOAT8_e4m3fn量化)、2(INT8量化)、3(MXFP8量化)、4(HIFLOAT8量化)和 5(MXFP4量化)。 + - `cmp_ratio` 支持 (0, 128] 内任意正整数。 + - 支持 `return_value`。 + - query 和 key:`quant_mode` 为 1/3 时支持 FLOAT8_e4m3fn,`quant_mode` 为 2 时支持 INT8,`quant_mode` 为 4 时支持 HIFLOAT8,`quant_mode` 为 5 时支持 FLOAT4_e2m1。 + - query_dequant_scale 和 key_dequant_scale:`quant_mode` 为 1/4 时支持 FLOAT32,`quant_mode` 为 2 时支持 FLOAT16,`quant_mode` 为 3/5 时支持 FLOAT8_e8m0。 + - weights:`quant_mode` 为 2 时支持 FLOAT16,`quant_mode` 为 1/3/4/5 时支持 FLOAT32。 + - query Q_N 支持 [1, 64]。 + + +- Atlas A3 训练系列产品/Atlas A3 推理系列产品、Atlas A2 训练系列产品/Atlas A2 推理系列产品: + - `layout_key` 仅支持 PA_BSND。 + - `quant_mode` 仅支持 2(Per-Token-Head量化)。 + - `cmp_ratio` 仅支持 2 的幂次方且范围为 [1, 128],即 1/2/4/8/16/32/64/128。 + - 不支持 `outputIdxOffsetOptional`。 + - 不支持 `return_value`。 + - query 和 key:支持 INT8,不支持 FLOAT8_e4m3fn、HIFLOAT8 和 FLOAT4_e2m1。 + - query_dequant_scale 和 key_dequant_scale:支持 FLOAT16,不支持 FLOAT32 和 FLOAT8_e8m0。 + - weights:支持 FLOAT16,不支持 FLOAT32。 + - query Q_N 仅支持 64。 + - topk 仅支持 [1, 2048]。 + + +- **返回值:** + + aclnnStatus:返回状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。 + + 第一段接口会完成入参校验,出现以下场景时报错: + + + + + + + + + + + + + + + + + + + + + + + +
返回值错误码描述
ACLNN_ERR_PARAM_NULLPTR161001如果传入参数是必选输入,输出或者必选属性,且是空指针,则返回161001。
ACLNN_ERR_PARAM_INVALID161002query、key、weights、queryDequantScale、keyDequantScale、cuSeqLensQOptional、cuSeqLensKOptional、sequsedQOptional、sequsedKOptional、cmpResidualKOptional、blockTableOptional、metadataOptional、layoutQOptional、layoutKOptional、topk、quantMode、maskModeOptional、cmpRatioOptional、returnValueOptional、sparseIndicesOut、sparseValuesOut的数据类型和数据格式不在支持的范围内。
+ +## aclnnQuantLightningIndexerV2 + +- **参数说明:** + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出描述
workspace输入在Device侧申请的workspace内存地址。
workspaceSize输入在Device侧申请的workspace大小,由第一段接口aclnnQuantLightningIndexerV2GetWorkspaceSize获取。
executor输入op执行器,包含了算子计算流程。
stream输入指定执行任务的Stream。
+ +- **返回值:** + + aclnnStatus:返回状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。 + +## 约束说明 + +- headdim 支持 128。 +- block_size 取值为 16 的倍数,最大支持 1024。 +- 当 `layout_key` 不为 PA_BSND 时,`layout_query` 和 `layout_key` 必须一致。 +- 当 `quant_mode` 为 3/5 时,`queryDequantScale` 和 `keyDequantScale` 的维数分别比 `query` 和 `key` 多 1,前缀维度保持一致,末两维为(D/64, 2);D必须为64的倍数,每个scale对应D轴上连续32个逻辑元素。 +- 当传入的参数layout_query为TND时,必须传入cuSeqlensQOptional,如果也传入sequsedQOptional,应保证由sequsedQOptional传入的各个batch的query长度不超过根据cuSeqlensQOptional计算出的各个batch的q序列长度。当某个batch由sequsedQOptional传入的q序列长度seqlen1小于由cuSeqlensQOptional计算出的query长度seqlen2时,会启用TND Padding功能,将该batch的seqlen2与seqlen1差值部分的query输出的sparseIndices和sparseValues全部置为无效值。部分长序列场景下,如果需要填充的无效数据过多,由于硬件限制可能会导致aicore执行超时,可以通过(seqlen2 - seqlen1) * topk来计算需要填充的数据量,建议将这个数据量控制在4亿以内。 +- **确定性说明:** aclnnQuantLightningIndexerV2 默认确定性实现。 + +## 调用示例 + +示例代码如下,仅供参考,具体编译和执行过程请参考[编译与运行样例](../../../docs/zh/context/compile_and_run_sample.md)。 + +```Cpp +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file test_aclnn_quant_lightning_indexer_v2.cpp + * \brief + */ +#include +#include +#include +#include +#include "securec.h" +#include "acl/acl.h" +#include "aclnnop/aclnn_quant_lightning_indexer_v2.h" +#include "aclnn/opdev/platform.h" + +using namespace std; + +namespace { + +#define CHECK_RET(cond) ((cond) ? true :(false)) + +#define LOG_PRINT(message, ...) \ + do { \ + (void)printf(message, ##__VA_ARGS__); \ + } while (0) + +int64_t GetShapeSize(const std::vector& shape) { + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} + +int Init(int32_t deviceId, aclrtStream* stream) { + auto ret = aclInit(nullptr); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclInit failed. ERROR: %d\n", ret); + return ret; + } + ret = aclrtSetDevice(deviceId); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); + return ret; + } + ret = aclrtCreateStream(stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); + return ret; + } + return 0; +} + +template +int CreateAclTensor(const std::vector& hostData, const std::vector& shape, void** deviceAddr, + aclDataType dataType, aclTensor** tensor) { + auto size = GetShapeSize(shape) * sizeof(T); + auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); + return ret; + } + + ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); + return ret; + } + + std::vector strides(shape.size(), 1); + for (int64_t i = shape.size() - 2; i >= 0; i--) { + strides[i] = shape[i + 1] * strides[i + 1]; + } + + *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), *deviceAddr); + return 0; +} + +struct TensorResources { + void* queryDeviceAddr = nullptr; + void* keyDeviceAddr = nullptr; + void* weightsDeviceAddr = nullptr; + void* qScaleDeviceAddr = nullptr; + void* kScaleDeviceAddr = nullptr; + void* metadataDeviceAddr = nullptr; + void* sparseIndicesDeviceAddr = nullptr; + void* sparseValuesDeviceAddr = nullptr; + + aclTensor* queryTensor = nullptr; + aclTensor* keyTensor = nullptr; + aclTensor* weightsTensor = nullptr; + aclTensor* qScaleTensor = nullptr; + aclTensor* kScaleTensor = nullptr; + aclTensor* metadataTensor = nullptr; + aclTensor* sparseIndicesTensor = nullptr; + aclTensor* sparseValuesTensor = nullptr; +}; + +int InitializeTensors(TensorResources& resources) { + int64_t B = 2; + int64_t S1 = 4; + int64_t S2 = 8; + int64_t N1 = 64; + int64_t N2 = 1; + int64_t D = 128; + int64_t topk = 512; + + std::vector queryShape = {B, S1, N1, D}; + std::vector keyShape = {B, S2, N2, D}; + std::vector weightsShape = {B, S1, N1}; + std::vector qScaleShape = {B, S1, N1}; + std::vector kScaleShape = {B, S2, N2}; + std::vector metadataShape = {1024}; + std::vector sparseIndicesShape = {B, S1, N2, topk}; + std::vector sparseValuesShape = {B, S1, N2, topk}; + + int64_t queryShapeSize = GetShapeSize(queryShape); + int64_t keyShapeSize = GetShapeSize(keyShape); + int64_t weightsShapeSize = GetShapeSize(weightsShape); + int64_t qScaleShapeSize = GetShapeSize(qScaleShape); + int64_t kScaleShapeSize = GetShapeSize(kScaleShape); + int64_t metadataShapeSize = GetShapeSize(metadataShape); + int64_t sparseIndicesShapeSize = GetShapeSize(sparseIndicesShape); + int64_t sparseValuesShapeSize = GetShapeSize(sparseValuesShape); + + std::vector queryHostData(queryShapeSize, 0x38); + std::vector keyHostData(keyShapeSize, 0x38); + std::vector weightsHostData(weightsShapeSize, 0.01f); + std::vector qScaleHostData(qScaleShapeSize, 1.0f); + std::vector kScaleHostData(kScaleShapeSize, 1.0f); + std::vector metadataHostData(metadataShapeSize, 0); + std::vector sparseIndicesHostData(sparseIndicesShapeSize, 0); + std::vector sparseValuesHostData(sparseValuesShapeSize, 0); + + int ret = CreateAclTensor(queryHostData, queryShape, &resources.queryDeviceAddr, + aclDataType::ACL_FLOAT8_E4M3FN, &resources.queryTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { return ret; } + + ret = CreateAclTensor(keyHostData, keyShape, &resources.keyDeviceAddr, + aclDataType::ACL_FLOAT8_E4M3FN, &resources.keyTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { return ret; } + + ret = CreateAclTensor(weightsHostData, weightsShape, &resources.weightsDeviceAddr, + aclDataType::ACL_FLOAT, &resources.weightsTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { return ret; } + + ret = CreateAclTensor(qScaleHostData, qScaleShape, &resources.qScaleDeviceAddr, + aclDataType::ACL_FLOAT, &resources.qScaleTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { return ret; } + + ret = CreateAclTensor(kScaleHostData, kScaleShape, &resources.kScaleDeviceAddr, + aclDataType::ACL_FLOAT, &resources.kScaleTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { return ret; } + + ret = CreateAclTensor(metadataHostData, metadataShape, &resources.metadataDeviceAddr, + aclDataType::ACL_INT32, &resources.metadataTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { return ret; } + + ret = CreateAclTensor(sparseIndicesHostData, sparseIndicesShape, &resources.sparseIndicesDeviceAddr, + aclDataType::ACL_INT32, &resources.sparseIndicesTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { return ret; } + + ret = CreateAclTensor(sparseValuesHostData, sparseValuesShape, &resources.sparseValuesDeviceAddr, + aclDataType::ACL_BF16, &resources.sparseValuesTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { return ret; } + + return ACL_SUCCESS; +} + +int ExecuteQuantLightningIndexerV2(TensorResources& resources, aclrtStream stream, + void** workspaceAddr, uint64_t* workspaceSize) { + int64_t topk = 512; + int64_t quantMode = 1; + int64_t maskMode = 0; + int64_t cmpRatio = 1; + int64_t returnValue = 1; + constexpr const char layoutQStr[] = "BSND"; + constexpr const char layoutKStr[] = "BSND"; + constexpr size_t layoutQLen = sizeof(layoutQStr); + constexpr size_t layoutKLen = sizeof(layoutKStr); + char layoutQ[layoutQLen]; + char layoutK[layoutKLen]; + errno_t memcpyRet = memcpy_s(layoutQ, sizeof(layoutQ), layoutQStr, layoutQLen); + if (!CHECK_RET(memcpyRet == 0)) { + LOG_PRINT("memcpy_s layoutQ failed. ERROR: %d\n", memcpyRet); + return -1; + } + memcpyRet = memcpy_s(layoutK, sizeof(layoutK), layoutKStr, layoutKLen); + if (!CHECK_RET(memcpyRet == 0)) { + LOG_PRINT("memcpy_s layoutK failed. ERROR: %d\n", memcpyRet); + return -1; + } + aclOpExecutor* executor; + + int ret = aclnnQuantLightningIndexerV2GetWorkspaceSize( + resources.queryTensor, resources.keyTensor, resources.weightsTensor, + resources.qScaleTensor, resources.kScaleTensor, + nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, + resources.metadataTensor, + topk, quantMode, -1, layoutQ, layoutK, maskMode, cmpRatio, returnValue, + resources.sparseIndicesTensor, resources.sparseValuesTensor, + workspaceSize, &executor); + + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclnnQuantLightningIndexerV2GetWorkspaceSize failed. ERROR: %d\n", ret); + return ret; + } + + if (*workspaceSize > 0ULL) { + ret = aclrtMalloc(workspaceAddr, *workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); + return ret; + } + } + + ret = aclnnQuantLightningIndexerV2(*workspaceAddr, *workspaceSize, executor, stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclnnQuantLightningIndexerV2 failed. ERROR: %d\n", ret); + return ret; + } + + return ACL_SUCCESS; +} + +int PrintOutResult(const std::vector& shape, void* deviceAddr) { + auto size = GetShapeSize(shape); + std::vector resultData(size, 0); + auto ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), + deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); + return ret; + } + LOG_PRINT("sparse_indices result (first 10 elements):\n"); + for (int64_t i = 0; i < size && i < 10; i++) { + LOG_PRINT(" [%ld] = %d\n", i, resultData[i]); + } + return ACL_SUCCESS; +} + +void CleanupResources(TensorResources& resources, void* workspaceAddr, + aclrtStream stream, int32_t deviceId) { + if (resources.queryTensor) { aclDestroyTensor(resources.queryTensor); } + if (resources.keyTensor) { aclDestroyTensor(resources.keyTensor); } + if (resources.weightsTensor) { aclDestroyTensor(resources.weightsTensor); } + if (resources.qScaleTensor) { aclDestroyTensor(resources.qScaleTensor); } + if (resources.kScaleTensor) { aclDestroyTensor(resources.kScaleTensor); } + if (resources.metadataTensor) { aclDestroyTensor(resources.metadataTensor); } + if (resources.sparseIndicesTensor) { aclDestroyTensor(resources.sparseIndicesTensor); } + if (resources.sparseValuesTensor) { aclDestroyTensor(resources.sparseValuesTensor); } + + if (resources.queryDeviceAddr) { aclrtFree(resources.queryDeviceAddr); } + if (resources.keyDeviceAddr) { aclrtFree(resources.keyDeviceAddr); } + if (resources.weightsDeviceAddr) { aclrtFree(resources.weightsDeviceAddr); } + if (resources.qScaleDeviceAddr) { aclrtFree(resources.qScaleDeviceAddr); } + if (resources.kScaleDeviceAddr) { aclrtFree(resources.kScaleDeviceAddr); } + if (resources.metadataDeviceAddr) { aclrtFree(resources.metadataDeviceAddr); } + if (resources.sparseIndicesDeviceAddr) { aclrtFree(resources.sparseIndicesDeviceAddr); } + if (resources.sparseValuesDeviceAddr) { aclrtFree(resources.sparseValuesDeviceAddr); } + + if (workspaceAddr) { aclrtFree(workspaceAddr); } + if (stream) { aclrtDestroyStream(stream); } + aclrtResetDevice(deviceId); + aclFinalize(); +} + +} // namespace + +int main() { + if (op::GetCurrentPlatformInfo().GetCurNpuArch() != NpuArch::DAV_3510) { + return 0; + } + int32_t deviceId = 0; + aclrtStream stream = nullptr; + TensorResources resources = {}; + void* workspaceAddr = nullptr; + uint64_t workspaceSize = 0; + int64_t B = 2; + int64_t S1 = 4; + int64_t N2 = 1; + int64_t topk = 512; + std::vector sparseIndicesShape = {B, S1, N2, topk}; + int ret = ACL_SUCCESS; + + ret = Init(deviceId, &stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("Init acl failed. ERROR: %d\n", ret); + return ret; + } + + ret = InitializeTensors(resources); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("InitializeTensors failed. ERROR: %d\n", ret); + CleanupResources(resources, workspaceAddr, stream, deviceId); + return ret; + } + + ret = ExecuteQuantLightningIndexerV2(resources, stream, &workspaceAddr, &workspaceSize); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("ExecuteQuantLightningIndexerV2 failed. ERROR: %d\n", ret); + CleanupResources(resources, workspaceAddr, stream, deviceId); + return ret; + } + + ret = aclrtSynchronizeStream(stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); + CleanupResources(resources, workspaceAddr, stream, deviceId); + return ret; + } + + PrintOutResult(sparseIndicesShape, resources.sparseIndicesDeviceAddr); + + CleanupResources(resources, workspaceAddr, stream, deviceId); + return 0; +} +``` diff --git a/csrc/attention/quant_lightning_indexer_v2/docs/qli_v2_two_level_topk_design.md b/csrc/attention/quant_lightning_indexer_v2/docs/qli_v2_two_level_topk_design.md new file mode 100644 index 000000000000..6f535c63b693 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/docs/qli_v2_two_level_topk_design.md @@ -0,0 +1,667 @@ + + +# QuantLightningIndexerV2 两级 TopK(候选块选择)方案设计 — 仅 arch22 / 910b + +> **Provenance**: 本方案基于 `ops-transformer` 仓库 `ds41` 分支 commit `32d64c27f`(提交MegaMoeWave模板)的代码调研。 +> 参考语义来源: 用户提供的 `select_candidate_blocks` PyTorch 参考实现(口头描述,2026-09-06)。 +> 适用范围: `attention/quant_lightning_indexer_v2`(主算子)+ `attention/quant_lightning_indexer_v2_metadata`(配套算子),**仅 910b (arch22) 路径**,arch35 (950) 不在本次范围。 + +--- + +## 1. 需求分析 + +### 1.1 模型侧语义(来自参考实现) + +```python +# Level One: 每个query选 topk_blocks 个最高分的块(block_size 个位置为一块) +def select_candidate_blocks(logits, compress_lens, topk_blocks=2048, block_size=8): + width = logits.size(-1) # S2(压缩后KV长度) + scores = F.pad(logits, (0, -width % block_size), -inf) # 尾块补 -inf + scores = scores.unflatten(-1, (-1, block_size)).amax(-1) # 块内取max → [.., num_blocks] + num_blocks = scores.size(-1) + last = (compress_lens - 1) // block_size # 最后一个未满块 + scores[last] = +inf # pin:最新token所在块无条件保留 + top = scores.topk(min(topk_blocks, num_blocks)) + keep = scatter(top.indices, top.values > -inf) # 丢弃全 -inf(不可达)块 + return keep.repeat_interleave(block_size, -1)[..., :width] # 位置级 bool mask + + +# 调用侧(跨层共享 shared_attn.candidates): +if self.is_candidate_source: # source 层:产出候选块 + shared_attn.candidates = select_candidate_blocks(index_score, compress_lens, 2048, 8) +elif self.uses_candidates: # consumer 层:候选块外 score 置 -inf 后再走原 topk + index_score = index_score.masked_fill(~shared_attn.candidates, -torch.inf) +``` + +关键点提炼: + +1. **块定义**: 位置 p 属于块 `floor(p / 8)`,块数 `num_blocks = ceil(S2 / 8)`,S2 为**压缩后**的 key 有效长度。 +2. **块得分**: 块内 8 个位置的 score 取 **amax**;不可达位置在 logits 中已是 -inf,不抬升块得分(与 pad 语义一致)。 +3. **pin 规则(compress_lens 形态决定,两种)**: + - **decode / mid-chunk prefill(模型 else 分支,op mask_mode=0)**:`compress_lens = end_pos // ratio` 为标量(= actS2Size),pin 块号 `(actS2Size - 1) // block_size`,batch 级(该 batch 所有行 pin 同一块); + - **prefill 首块(模型 start_pos==0 分支,op mask_mode=3)**:`compress_lens = (arange(1, seqlen+1) // ratio)` 为行级 `[S1, 1]`,行 i 的 pin 块号 `(rowValidLen(i) - 1) // block_size`,其中 `rowValidLen(i) = (actS2SizeOrig - actS1Size + i + 1) // cmpRatio`(与 op mask_mode=3 的行级有效长度公式逐行一致:首块时 actS2SizeOrig=actS1Size=seqlen → `(i+1)//ratio`,与模型代码吻合)→ **每行 pin 各自的最后一块**; + - 保证最新 token 所在块必被选中的语义在两种形态下均成立。 +4. **有效槽**: 可达块数不足 `topk_blocks` 时,未被选中的槽位无效(值 `-inf` 被丢弃)。 +5. **Level-2 动态 topk(模型 gather 侧,无需 op 改动)**: + ```python + topk = min(self.index_topk, end_pos // ratio) # 动态 topk = min(sparse_count, actS2Size) + idxs = index_score.topk(topk, sorted=False).indices.sort(dim=-1) # 升序 + return torch.where(idxs < compress_lens, idxs + offset, -1) + ``` + - `topk` 属性为每次调用传入的动态值([1,2048]),模型传 min 值即可; + - `where(idxs < compress_lens, ·, -1)` 与 kernel 现有语义**等价**:不可达槽位输出 -1(`InitSortOutBuf` 填 (-inf,-1),有效长度外沉底);"topk=sparse_count + 滤 -1" 与 "topk=min 直接输出" 有效索引集合相同; + - `idxs + offset` 由模型侧完成(910b 不支持 output_idx_offset,无需支持); + - **结论:无 kernel/host 改动,仅测试覆盖**(§6.3/§6.5 增加 topk<2048 动态值与 min 等价性用例)。 +6. **source 层自身仍做正常 topk**(if/elif 互斥,source 不做 mask 但照常输出 sparse_indices)。 +7. **候选块跨层共享**,同一 forward 内 KV 状态不变 → 块划分一致;**跨 step 复用无效**(语义约束)。 + +### 1.2 现有算子流水线摘要(arch22, 910b) + +| 阶段 | 执行单元 | 内容 | 数据位置 | +|---|---|---|---| +| S1 | AIC `ComputeMm1` | Q(256=s1×g, 128) @ K^T(128, s2≤2048) → L0C | L1/L0 | +| S2 | AIC `FixpSToL1` | DEQF16(scale 0.001) + **ReLU** → half [s1g, s2] | L1 `sL1_` | +| S3 | AIC `ComputeWs` | w(s1g) @ ReLU(QK)(s1g, s2) → g 维求和 | L0 | +| S4 | AIC `FixpResToGm` | nz2nd → **score[s1, s2]** float | GM `mm1ResGm`(loop%2 双buffer) | +| S5 | AIV `ProcessVec0` | (w × qScale) 广播 → 供 S3 使用 | GM `vec0OutGm` | +| S6 | AIV `ProcessVec1` | score × kScale → SortAll(降序) → MergeSort 进 `globalTopkUb_`(2048 value/idx pair/s1行) → 行末 ExtractIndex 直出 `sparse_indices` | UB/GM | + +> **LD(跨核归约)在 910b 不激活**:`supportFd_` 默认 false 且仅 ASCEND950 分支置 true(metadata aicpu.cpp:329-338, aicpu.h:285)→ `AssignByBlock`(S2 块级跨核切分)被跳过 → 单个 (b, s1) 行的 S2 永不跨核 → kernel `isNeedLD` 恒 false,`ProcessLD` 的 QLD_V2 元数据全 0。**本方案不考虑 LD 路径**(kernel 中的 LD 代码保留不动,仅不激活)。 + +结论:**`index_score`(参考实现中的 logits)= S4 输出 × kScale(S6 第一步之后)**,按 (s1行, s2 tile ≤2048) 流式产生,且每行 topk 结果在本核内独立完成。两级 TopK 的插入点在 S6(ProcessVec1)内部。 + +### 1.3 输入输出定义(新增部分) + +**主算子 QuantLightningIndexerV2 新增(命名与模式编号按用户 2026-09-06 决策)**: + +| 项目 | 名称 | 方向 | 类型/Shape | 约束 | +|---|---|---|---|---| +| 输入 | `candidate_topk_index` | 可选输入 | INT32, `[B, S1, N2, 2048]`(BSND) | `candidate_mode=2` 时必传;值域 `[0, numBlocks)` 或 `-1`(无效槽) | +| 输出 | `candidate_topk_index` | 可选输出 | INT32, 同上 | `candidate_mode=1` 时输出;无序,槽位语义:块号或 -1 | +| 输出 | `sparse_indices`(即 topk_index) | 既有输出 | 不变 | mode=1: 全量 score 的 topk;mode=2: **从 candidate_topk_index 展开的 ≤16384 个候选位置中再挑 topk 个**;mode=3: 现网行为 | +| 属性 | `candidate_mode` | 可选属性 | INT, 默认 **3** | **1**=source(is_candidate_source,输出候选块索引 + 照常输出 topk_index);**2**=consumer(use_candidate,输入候选块索引,在候选内选 topk);**3**=关闭(现网行为,默认值保证既有调用零影响) | +| 属性 | `candidate_topk_blocks` | 可选属性 | INT, 默认 2048 | **(0, 2048] 内 64 的倍数**(2026-09-07 放宽,原为仅 2048;上限受 BASE_TOPK=2048 的 sort/merge 结构约束;`numBlocks > topk_blocks` 时 pin 最新块必须入选,见 §9/P2)| +| 属性 | `candidate_block_size` | 可选属性 | INT, 默认 8 | 范围 [1, 64] 且为 2 的幂(向量 reduce 实现约束);**source 与 consumer 必须一致** | + +**设计决策 — 输出索引而非 bool(用户决策"不再输出 bool,直接输出 topk index")**: + +- `candidate_topk_index` 为**块级索引**(2048 个 int32/行,值=块号),展开后等价于 2048×8=16384 个位置级索引;相比位置级索引(16384 个 int32=64KB/行)**GM 内存压缩 8×**,相比 bool mask `[B,S1,S2]`(prefill 大序列可达 512MB)更是量级差异; +- mode=2 消费时按 `position ∈ [8×blk, 8×blk+8)` 展开,天然还原块内连续 8 位置; +- 既有输出 `sparse_indices` 名称不变(aclnn 兼容),文档语境中的 topk_index 即指它。 + +**int 块索引 vs bool mask 的权衡记录(2026-09-06,结论:本场景无实质不方便)**: + +| 维度 | bool mask | int 块索引(本方案) | +|---|---|---| +| 算子外消费 | 直接可用 | 需先展开(届时在 torch 封装层提供 expand 工具即可,不必改算子接口) | +| 接口约定 | 无 | -1=无效槽、槽位无序(topk 序)、块号×8=起始位置,需文档化(README 约束项) | +| mode=2 kernel 实现 | 读 tile mask 段 + Select,简单 | membership(compare-OR 分批)+ 展开,复杂一档(S4a 已设计) | +| 调试/比对直观性 | 高 | 低(dump 为块号列表;比对已集合化 §6.2) | +| GM 内存/带宽 | S2 字节/行(最坏 512MB) | 8KB/行,**省 8×**(决定性理由) | +语义完备性说明:块内有效性由原 score 兜底(mask_mode=3 下行内不可达位置 score 本为 -inf),mode=2 的 masked score 与 bool 语义严格等价。 + +**metadata 算子:不改接口**(论证见 §3.2)。 + +### 1.4 数学语义(新增部分的形式化) + +对每个 batch b、query 行 s1(n2=1, 910b 上 kHeadNum=1): + +``` +score(b, s1, :) ∈ R^{S2} # 现有 S4×kScale 结果,不可达位置语义为 -inf +numBlocks(b) = ceil(actS2Size(b) / block_size) +blkScore(b, s1, j) = max_{8j ≤ p < min(8j+8, S2)} score(b, s1, p) # 尾块不足8位置按 -inf pad + +# pin 块号(行级,随 mask_mode 取行级/批次级有效长度): +rowValidLen(b, i) = (actS2SizeOrig(b) - actS1Size(b) + i + 1) / cmpRatio # mask_mode=3(行级) +rowValidLen(b, i) = actS2Size(b) # mask_mode=0(batch 级,各行相同) +lastBlk(b, i) = (rowValidLen(b, i) - 1) / block_size # 整除 +blkScore(b, i, lastBlk(b, i)) = +inf # pin:每行 pin 自己的最后一块 + +mode=1 (source): candidateTopkIndex(b, i, :) = topK_blocks 个 blkScore 最大的块号 + (并列值 tie-break 不保证与 PyTorch 一致,见 §6.2;blkScore=-inf 的槽输出 -1) + sparse_indices(b, i, :) = topK(score) # 与现有逻辑完全一致(不 mask) +mode=2 (consumer): score'(b, i, p) = score(b, i, p) 若 floor(p/8) ∈ candidateTopkIndex(b, i, :) + -inf 否则 + # 等价视角:sparse_indices 是"从 candidate_topk_index 展开的 ≤16384 个候选位置中再挑 topk 个" + sparse_indices(b, i, :) = topK(score') # 复用现有 sort/merge 管线 +mode=3 (off): 现网行为,candidate 功能整体关闭(默认值,向后兼容) +``` + +dtype/精度约束:score 全程 fp32(现有路径不变);块化 amax 与 -inf/+inf 处理均在 fp32 域(`NEG_INF=0xFF800000` 已有,新增 `POS_INF=0x7F800000`)。 + +### 1.5 验收标准 + +**功能验收**(mode=3 回归为正确性门禁,非性能): + +- candidate_mode ∈ {1, 2, 3};mode=3 输出与现网 bit 级一致(回归,默认值保证既有调用零影响)。 +- **layout_q:BSND(本轮已实现)+ TND(2026-09-07 需求追加,见 §11)**;layout_k = PA_BBND(BSND 时)或 TND(TND 时,cu_seqlens_k 变长拼接)。**key 0 轴非连续**(PA_BBND 第 0 维 stride > 块逻辑大小)为追加需求,参照 arch35 机制(§11.2)。 +- mask_mode ∈ {0, 3};**cmp_ratio ∈ {1, 2}**(用户指定重点场景,见 §6.6)。 +- **q_head_num 仅测 32**(用户指定:所有 candidate 用例统一 g=32,不再覆盖 64;A1/A2 改造后 g=64 逻辑路径不变,由既有 910b 用例回归兜底)。 +- **topk 动态值**:topk ∈ [1, 2048] 任意值,含 topk = min(index_topk, actS2Size) 语义(等价性验证,无代码改动)。 +- S2 使 numBlocks < / = / > topk_blocks 三类;actS2Size 非 block_size 整除(尾块)。 +- decode(S1=1)与 prefill(S1>1)。 +- 910b 约束继承:quant_mode=2(INT8 + fp16 scale/weight)、无 return_value/output_idx_offset。 + +**精度验收**: + +- `candidate_topk_index`:与参考实现的**选中块集合**一致(无序集合比较 + -1 槽位数一致);允许 tie 场景按 §6.2 策略仲裁。 +- mode=2 `sparse_indices`:与 "参考实现 mask 后走现有 golden topk" 一致(复用现有 compare 框架与 tie 容差策略)。 +- mode=1 `sparse_indices`:与现网 golden 一致(不应受块级计算影响——同一 Vec1 内两套独立 sort buffer)。 + +> 本次交付不含性能验收(用户明确性能暂不关注),但设计保留向量化的块化/mask 实现路径(§3.3/§3.5),避免后续性能优化时返工。 + +**稳定性验收**:多 seed(≥3)、边界 shape(S2%8≠0、actS2Size=1、numBlocks=1、topk_blocks>numBlocks)、B>1 变长 batch。 + +--- + +## 2. 算法拆解(在现有 6 阶段上的增量) + +| # | 阶段 | 现状 | 本次变更 | 执行单元 | +|---|---|---|---|---| +| 0 | 输入解析 | layout/actual seqlen/metadata 分核 | consumer(mode=2): 解析 `candidate_topk_index` GM 指针;新增 TilingData 字段下发 | Scalar | +| 1 | 预处理 | Vec0: w×qScale 广播 | 不变 | Vector | +| 2 | MatMul1 (QK) | AIC | 不变 | Cube | +| 3 | score 后处理 | Fixp DEQF16+ReLU → Ws g-sum → GM | 不变 | Cube | +| 4 | Vec1 score 生成 | mm1Res×kScale | **mode=2: 追加候选块 mask 生成 + Select 置 -inf**(插在 Mul(kScale) 之后、SortAll 之前) | Vector | +| 5 | 归约/块化 | —(无) | **mode=1: 块化 amax(8:1) + pin 尾块置 +inf**(与位置级 SortAll 并行,规模 1/8) | Vector | +| 6 | 选择/TopK + 输出 | SortAll+MergeSort → globalTopkUb_(2048) → 行末 ExtractIndex 直出 | **mode=1: 镜像维护 globalBlockTopkUb_(2048),行末 ExtractIndex 直出 candidate_topk_index GM**;生命周期/重置点与位置级完全同步 | Vector | +| — | 无效清理 | CleanInvalidOutput 填 -1(sparse_indices) | **mode=1: 所有填 -1 的路径同步对 candidate_topk_index 填 -1**(DealActSeqLenIsZero / BSND 无效 S1——BSND 范围内共 2 处,见 §3.3 检查清单;TND padding 路径本轮不适用) | Vector | + +阶段信息表(新增阶段细目): + +**S4a(mode=2, consumer mask)**: + +- 输入:`candidate_topk_index` GM 行 `[topk_blocks]` int32(8KB)、当前 tile score UB `[tileLen]` fp32 +- 输出:masked score UB(原地) +- 计算单元:Vector(compare-OR 分批,见 §3.5) +- 同步:无跨核(每行候选集独立,行内自洽) + +**S5a(mode=1, 块化)**: + +- 输入:score UB `[tileLen]`(kScale 相乘后) +- 输出:`blkScore UB [tileBlkNum=tileLen/8]` fp32 + `blkIdx UB` int32(基址 = tileS2Base/8) +- pin:当前行 pin 块 `lastBlk = (rowValidLen - 1) / block_size`(行级,`rowValidLen` 即 Vec1 已有的 `cuRealAcSeq`,mask_mode=0 时各行相同退化为 batch 级);tile 包含该块时置 +inf;tile 内超出行有效长度(mask_mode=3)的块值为 -inf(pad 语义) +- 计算单元:Vector(8:1 strided ReduceMax + 尾块 -inf 对齐) + +**S6a(mode=1, 块级 sort/merge + 直出)**:镜像 S6 结构,`SortAll(blkNumAligned)` + `MergeSort(globalBlockTopkUb_[row], topk_blocks, ...)` + 行末 `ExtractIndex` → `candidate_topk_index GM`;`AlignS2` 复用(块数对齐到 32/128/512 粒度)。 + +--- + +## 3. Host Tiling 设计(arch22) + +### 3.1 维度建模 + +| 维度 | 含义 | 现值/来源 | +|---|---|---| +| B / N2 / G | batch / kv头(=1) / query头组 | **G = 32(本轮唯一测试规格)** | +| S1 | query 长度 | 不变 | +| S2 | 压缩后 key 长度(actS2Size) | 不变 | +| **numBlocks** | `ceil(actS2Size / block_size)` | **新增,随 batch 变长** | +| **tileBlkNum** | 单 tile 内块数 = `s2BaseSize / block_size = 256` | **新增,常量(block_size=8 时)** | +| topk | 位置级 topk | 动态 [1, 2048],含 min(index_topk, actS2Size) 语义(无代码改动) | +| **topkBlocks** | 块级 topk | **固定 2048**(BASE_TOPK 结构复用;属性保留,host 校验限定,未来扩展只放宽校验) | + +> **q_head_num=32 的支持方式(参照 v1 `quant_lightning_indexer` 实现,全量关键词搜索适配点)**: +> +> - **v1 与 v2 的关键差异**:v1 arch22 kernel 固定 `S1_BASE_SIZE=4`、推导 `mBaseSize = s1BaseSize × gSize`(g=32 → mBase=128);v2 现状反之——固定 `M_BASE_SIZE=256`、推导 `s1BaseSize = 256/gSize`(g=32 → s1BaseSize=8,按行 UB 翻倍)。 +> - **一致性论证**:v2 的 metadata aicpu 本就是 v1 风格(`s1BaseSize_=4` 默认、`mBaseSize_ = s1BaseSize_ × groupSize_`,aicpu.h:307 / aicpu.cpp:247)——g=64 时 4×64=256 与 v2 kernel 固定值恰好重合,**g=32 时若 kernel 不改,分核区间(128 基准)与 kernel 循环(256 基准)失配**。因此本修改同时是 g=32 正确性的必要条件。 +> - **收益**:s1BaseSize 恒为 4 → sortOutBuf_/candidate globalBlockTopkUb_ 等 UB 与 workspace 全部不变,§3.3 的 g=32 UB 紧张场景(原 R3)不存在。 +> - **验证范围(2026-09-07 用户确认)**:candidate 测试统一 **q_head_num=32**(g=64 由既有回归兜底,不在本轮范围);R3 风险随之收窄为"推导变更正确性"而非 UB 容量。 +> +> **适配点全量清单(v1↔v2 归一化 diff + 关键词扫描结论)**: +> +> | # | 位置 | v1 实现 | v2 现状 | 本次动作 | +> |---|---|---|---|---| +> | A1 | kernel `InitTilingData`(v2 kernel_arch22.h:191-193) | `s1BaseSize=4` 固定;`mBaseSize = s1BaseSize × gSize`(v1:171-175) | `mBaseSize=256` 固定;`s1BaseSize = 256/gSize` | **改**:对齐 v1 推导 | +> | A2 | host `GetGSize` 910b 分支(v2 tiling.cpp:857-860) | `gSize > 64` 才拒绝(v1:700-705,即 32 合法) | `gSize != 64` 拒绝 | **改**:放宽允许 32(复用 v2 tiling.h:87 既有常量 `G_SIZE_LIMIT_32_950`) | +> | A3 | cube L0/L1 切分(`ComputeMm1`) | 按 S1 维切:`s1L0LoopCnt = CeilDiv(actM/g, s1Base/2)`,L0 子块 `gSize×s1Base/2`(随 g 缩放,g=32 → 64 行/次) | 固定粒度:`CeilDiv(actM, S1G_BASIC_BLOCK_L0=128)`,L0 子块恒 128 行 | **不改**:g=32 时 actMBaseSize=128 → 单次 L0 循环,mExtension=CeilAlign(128,16)=128 ≤ L1 容量,v2 现有逻辑自动适配(两种切分等价可行,保持 v2 风格改动最小) | +> | A4 | cube `ComputeWs/LoadSToL0b/LoadWeightToL0a/FixpResToGm` | 全部 `gSize` 参数化(`k=gSize`、`s1gOffset += gSize`、`repeatTimes=CeilDiv(gSize,16)`) | 同样已 gSize 参数化 | **不改**(diff 确认一致) | +> | A5 | Vec0 `cuProcEleNum`(v2 service_vector:277) | `CeilAlign(cuS1ProcNum × gSize, 32)`(v1:289,UB 对齐防御) | 无对齐(g=64 时恒为 32 倍数而省略) | **建议同步 v1 的 CeilAlign**:g=32 时 4×32/3×32/2×32/1×32=128/96/64/32 恰好均为 32 倍数(数学上不改动也安全),但对齐写法防御未来 g 非 2 的幂 | +> | A6 | host workspace(v2 DoTiling arch22 分支) | `QliCalcWorkspaceSize` 用保守上界常量(mBase=512×s2Base=512,g 无关) | `M_BASE_SIZE(256) × S2_BASE_SIZE(2048)` 上界 | **不改**:256 是 g∈{32,64} 的 mBaseSize 上界(4×64),g=32 的 128 被覆盖 | +> | A7 | kernel Init workspace 布局(v2 kernel:463-478) | — | `mm1Res: 2×s1BaseSize×s2BaseSize`(s1BaseSize 恒 4 不变);`weightMemSize: 16×mBaseSize×2`(mBaseSize=128 → 减半,变小安全) | **不改**:布局全部由 constInfo 推导,A1 改后自适应 | +> | A8 | metadata aicpu | — | 已是 v1 风格(`mBaseSize_ = 4×groupSize_`) | **不改**(g=32 与改后 kernel 基准一致) | +> | A9 | `GetS2BaseBlockNumOnMask/GetTotalBaseBlockNum/CalcGS1LoopParams` 等 | `s1BaseSize/gSize/mBaseSize` 参数化 | 同 | **不改**(A1 后自动正确) | +> | A10 | UT(tests/ut/op_host/arch22) | v1 有 arch22 tiling 单测 | v2 有同款 | **补**:g=32 的 tiling/infershape 单测 case | + +### 3.2 多核切分 —— 不变性论证 + +- 910b 上 metadata 仅做 batch 级(`AssignByBatch`)与整行级(`AssignByRow`)切分,`AssignByBlock`(S2 块级)因 `supportFd_=false` 跳过 → 每个 (b, s1) 行整体落在一个核上,topk 天然单核完成,无需跨核归约。 +- 块级 topk 与位置级 topk 在**同一 Vec1 调用**内、同一 (bN2, gS1, s2) 任务块上执行,共享现有 metadata 分核(QLI_V2_* AIC 段)。 +- 负载增量:Vec1 每 tile 增加 `O(tileLen/8)` 的块化 + 块级 sort/merge(≈位置级 sort 开销的 1/8)→ 仅放大每块耗时,**不改变任务块数量与划分粒度** → metadata 的分核算法、输出布局、协议(1024 int32, AIC 36×8 + AIV 72×8)均不变,QLD_V2 段在 910b 本就全 0。 + +### 3.3 UB Buffer 规划(AIV,910b UB **实测 192KB**) + +采用 §3.1 的 v1 对齐方案后,s1BaseSize 恒为 4(g∈{32,64} 通用),UB 不随 gSize 变化: + +现有(`InitBuffers`):inQueue 32KB + outQueue 8KB + indexBuf 8KB + tmpBuf 64KB + sortOutBuf 32KB = 144KB。 + +**mode=1 增量**: + +| buffer | 元素数 | 字节 | 生命周期 | 复用 | +|---|---|---|---|---| +| globalBlockTopkUb_ | CeilDiv(s1BaseSize,2) × topkBlocks × 2 (fp32 pair) | 32KB | 整个 gS1 基本块 | 新增 TBuf,重置点与 globalTopkUb_ 同步 | +| 块化临时(实装布局,均在 tmpBuf 64KB 内) | blkScore@tmp[6144] / blkIdx@tmp[6144+blockNumPad] / isPad+pin 链@tmp[14336–15872] / blkSortTmp@tmp[11776] | ≤5KB | 单 tile | 复用 tmpBuf;**blkSortTmp 需容纳 mrgDst+mrgSrc ≤4608 floats,11776+4608=16384 恰为 64KB 末尾(曾置 12288 越界 2KB 触发 aicore)** | + +合计 176KB < 192KB ✓(**实测 UB 为 192KB,非 256KB**;编译期从 PlatformInfo ubSize 校验,**禁止**硬编码常量)。 + +**mode=2 增量**(实装): + +| buffer | 元素数 | 字节 | 说明 | +|---|---|---|---| +| candBuf_(TBuf) | CeilDiv(s1BaseSize,2) 行 × topkBlocks × 2 (fp32 pair) | 32KB(R6 按行分区) | 排序后的候选 [values|idx] 对,每行独立(同核行间 tile0 重排序覆盖是 R6 根因) | +| candConstBuf_(TBuf) | topkBlocks (fp32) | 8KB | -1e30 常量(候选外罚分) | +| 掩码临时(复用 tmpBuf) | candInt@12288 / candSortTmp@4096 / blkIdxF·acc·diff@4352–4864 / posDist@12288 / isOutI32@14336 | ~20KB | **isOutI32 必须避开 [4096,6144)**(该区被 ProcessVec1 的 pen/idxPen 复用,曾重叠致掩码失效) | + +**输出清理检查清单(mode=1 填 -1 的路径,BSND 范围内共 2 条)**: + +1. `DealActSeqLenIsZero`(actS1Size=0 或 actS2Size=0,BSND 分支) +2. BSND 无效 S1 尾部(qSeqSize > actS1Size) + +(causal 下 actS1Size > actS2SizeOrig 的行在 mask_mode=3 时按行有效长度自然处理为块 -inf,非独立清理路径;TND padding 路径本轮不适用,见 §5。) + +### 3.4 Workspace 规划(GM)—— 无增量 + +现有布局(arch22,per AIC 核)不变: + +``` +[0] mm1ResGm : 2 × s1BaseSize × s2BaseSize × 4B +[+off1] vec0OutGm : 16 × mBaseSize × 2 × 2B +[+off2] vec1ResGm(LD) : s1BaseSize × 2 × 2 × BASE_TOPK × 4B ← V1_DECODE 区(910b 不激活,保留) +``` + +**结论:candidate 功能不新增任何 workspace。** 理由:块级 topk 在单核 Vec1 内完成并直出(§3.2),无跨核中间结果落盘;mode=2 的候选集每 tile 从 GM 直接读入 UB。tiling.cpp 的 workspaceSize 计算不变(仍需随本方案回归确认无隐性依赖)。 + +### 3.5 分支策略(host 集中,kernel 单点判断) + +| 分支 | 条件 | 行为 | +|---|---|---| +| mode=3 | candidate_mode=3 | 现有路径原样(默认值,向后兼容;模板内 if 包裹新增段) | +| mode=1 | candidate_mode=1 | S5a/S6a 全开(source) | +| mode=2 | candidate_mode=2 | S4a mask 生效(consumer) | + +mode=2 的 membership 判定(候选列表 → tile 内 256 块 bool,向量化无标量循环): + +- `blkMask[i] = OR_j (candList[j] == tileBlkBase + i)`,compare-OR 分批:每批 32 候选 × 256 块(32KB fp32 view),共 topkBlocks/32 = 64 批/行;随后 `repeat_interleave(8)` 展开 + `Select` 置 -inf。 + +分支条件全部由 host tiling 写入 TilingData 字段(含 `candidateMode`、`candidateTopkBlocks`、`candidateBlockSize`),kernel 内不重复推导。 + +> 注:`numBlocks ≤ topkBlocks` 时 mask 恒为全选(等价于不 mask),实现上**不做专门 fast-path 分支**(性能暂不关注,减少分支与测试组合);该等价性仍作为正确性用例覆盖(§6.3)。 + +--- + +## 4. 契约设计 + +### 4.1 TilingData(QLIV2TilingData 尾部追加,4B 对齐) + +```cpp +TILING_DATA_FIELD_DEF(uint32_t, candidateMode) // 1=source / 2=consumer / 3=off(默认) +TILING_DATA_FIELD_DEF(uint32_t, candidateTopkBlocks) // ≤2048 +TILING_DATA_FIELD_DEF(uint32_t, candidateBlockSize) // 默认8,2的幂 +``` + +host 写入 ↔ kernel 消费对照表随本方案落入 docs(维护接口式追踪)。字段分组注释:`// ---- candidate (two-level topk) ----`。 + +### 4.2 TilingKey —— 不变 + +理由(遵循"只编码影响模板实例的维度"):candidate 路径是纯 Vector 数据通路,不改变 MatMul 形状/dtype/layout/模板实例,仅 kernel 内运行时分支;编码进 key 会造成 ×3 模板实例化与编译时间膨胀,无收益。 + +### 4.3 算子原型 / Infershape / aclnn / torch + +- `quant_lightning_indexer_v2_def.cpp`:新增 optional input `candidate_topk_index`(INT32, ND)、optional output `candidate_topk_index`(INT32, ND)、三个可选属性;仅 `ascend910b` 配置声明(`aicore_config`),950 配置不动。 +- `quant_lightning_indexer_v2_infershape.cpp`:`candidate_mode=1` 时输出 shape = sparse_indices 前缀 + 末维 `candidate_topk_blocks`(**仅 BSND 分支**,q 只考虑 BSND);否则末维 0。InferDataType 补 SetOutputDataType(2, DT_INT32)。 +- aclnn 两段式接口(`aclnnQuantLightningIndexerV2`):GetWorkspaceSize/执行签名追加 optional `candidateBlocks` 输入/输出指针与 3 个属性,向后兼容(指针可空)。 +- torch schema(**已定方案 b——新增 overload,旧接口不动**): + ```python + # 新增变体(source/consumer 统一入口,返回固定三元组) + quant_lightning_indexer.candidate( + query, key, weights, q_descale, k_descale, topk, quant_mode, *, + Tensor? candidate_topk_index=None, # mode=2 时必传 + int candidate_mode=3, # 1=source / 2=consumer / 3=off(默认) + int candidate_topk_blocks=2048, + int candidate_block_size=8, + ...其余参数同旧接口...) -> (Tensor sparse_indices, Tensor sparse_values, Tensor candidate_topk_index) + # 非 source 模式第三元返回 shape (0,) 占位,元信息(block_size 等)随返回对象附带的 sidecar 属性传递 + ``` + - 旧 `quant_lightning_indexer` schema 二元组完全不变,现网调用零影响; + - 模型侧 source/consumer 层统一走 `.candidate` 入口,按 candidate_mode 区分行为; + - block_size 一致性断言(R4)在 python 封装层实现:shared 对象携带 `candidate_block_size` 元信息,consumer 调用时校验。 + +### 4.4 metadata 算子(`quant_lightning_indexer_v2_metadata`) + +**结论:不改**。分核算法/输出协议/属性均不变(论证见 §3.2;910b 无 FD/LD 路径,QLD_V2 段维持全 0)。配套约束: + +- 主算子 mode≠0 时 metadata 照常传入; +- 文档(README)补充:candidate 相关属性不参与分核,因负载模型未变。 + +--- + +## 5. 与既有机制的交互确认 + +### 5.0 mask 计算结论:不修改、不新增 mask mode + +现有 mask_mode 与模型代码分支**逐行数学等价**,无需任何 mask 侧改动。 + +**关键澄清**:模型 if/else 是**推理时间步**的分支,不是一次调用内的分支——每次算子调用仍只有一种 mask_mode: + +``` +prefill 首块 (start_pos==0): 调用 QLI 1 次 (S1=seqlen) → mask_mode=3, 行级 pin +decode step (start_pos>0): 每步调用 QLI 1 次 (S1=1) → mask_mode=0, batch 级 pin +两级 TopK 时间线: + prefill 首块: QLI(mode=1, mask_mode=3) → 产出 candidates 存 shared_attn + decode step: QLI(mode=1/2, mask_mode=0) → source 更新候选 / consumer 消费候选 +``` + +| 模型分支 | 模型行为 | op 对应 | 等价性论证 | +|---|---|---|---| +| `start_pos == 0`(prefill 首块) | `compress_lens[i] = (i+1)//ratio`(`[S1,1]` 行级),`p ≥ compress_lens[i]` 置 -inf(模型侧 masked_fill_;score 下沉算子后由 mask_mode=3 承担) | mask_mode=3 | op 行级有效长度 `(actS2SizeOrig - actS1Size + i + 1)/cmpRatio` 在首块时 `actS2SizeOrig = actS1Size = seqlen` → 退化为 `(i+1)//ratio`,与模型逐行恒等(golden `create_mask`、kernel `cuRealAcSeq`、模型三方公式一致) | +| `else`(decode) | 标量 `compress_lens = end_pos//ratio`,不 mask | mask_mode=0 | op 不 mask,天然匹配;pin 用 batch 级 `(actS2Size-1)/block_size` | + +新增的行级 pin(S5a)**复用同一个 `cuRealAcSeq`**,不构成新 mask——pin 与 mask 同源,不会出现 pin 块落在被 mask 区域的矛盾。 + +**开放问题 O3**:`else` 分支为标量 compress_lens,意味着 mid-chunk prefill(start_pos≠0 且 S1>1)时模型**不做行内因果 mask**。若该调用形态实际不存在,测试矩阵中 "mid-chunk prefill×mode0" 用例删除;若存在,op 照实不 mask(复现模型行为),同样无需改动。待用户确认。 + +> ⚠ **O3 状态(标红保留)**:**此问题未关闭,不允许随本需求悄悄消化。** 当前决策为"暂不测试"——mid-chunk prefill×mode0 用例从本轮测试范围剔除(用例定义保留在矩阵,注释标明暂不执行),但设计上模型语义为"prefill 不做行内 mask",**与常规认知相反**。在模型侧确认该调用形态(chunked prefill / 混布场景是否走 else 分支)之前: +> +> 1. 禁止有人"顺手"在算子里给 mode0+S1>1 补行级 mask(那是语义变更,不是 bug fix); +> 2. 精度测试若碰到 mode0+S1>1 的数据,比对结论一律以"模型不 mask"的 golden 为准; +> 3. 本问题最终去向二选一:确认形态不存在 → 删除用例;确认存在 → 启用用例并补 golden 说明。 + +| 机制 | 交互 | 结论 | +|---|---|---| +| LD / ProcessDecode | 910b 不激活(`supportFd_` 仅 950 置 true,aicpu.cpp:329-338) | **不适用,方案不含 LD 路径**;kernel 现有 LD 代码与 V1_DECODE workspace 保留不动 | +| mask_mode=0(decode / mid-chunk prefill) | 无行 mask,全行有效长度 = actS2Size;pin 为 batch 级 `(actS2Size-1)/block_size`(对应模型 else 分支标量 `compress_lens = end_pos//ratio`) | 兼容 | +| mask_mode=3(prefill 首块) | 行级有效长度 `rowValidLen(i) = (actS2SizeOrig-actS1Size+i+1)/cmpRatio`(块化时行尾块按该行有效长度 pad -inf,等价模型 `masked_fill_`);**pin 为行级** `(rowValidLen(i)-1)/block_size`(对应模型 `[S1,1]` 行级 compress_lens) | 兼容,S5a 内处理;行级 pin 与行级 mask 公式同源(Vec1 `cuRealAcSeq`),天然一致 | +| cmp_ratio 压缩 | 块划分基于压缩后 actS2Size(与 Python 的 S2=logits.size(-1) 一致) | 兼容 | +| Vec1 双缓冲预取(`LI_QUANT_PRELOAD_TASK_CACHE_SIZE=2`) | 块级 sort 结果随 runInfo 双缓冲流转,重置点与位置级 `globalTopkUb_` 完全同步(`info.s2Idx==0` 时重置、行末直出后重置) | 设计强约束:两套 buffer 的重置/输出点成对出现,UT 加不变量断言 | +| TND padding(sequsedQ < cuSeqlensQ) | padding 行的 candidate_topk_index 需填 -1(同 sparse_indices 无效值路径) | **不适用**(仅 TND 布局存在该路径,本轮 q 只考虑 BSND);若未来放开 TND 需同步补上 | +| Sort32/MrgSort 降序假设 | 假设:降序(top-k 语义自洽,-inf 沉底/初始填充行为验证) | 实现期以 API 文档+单测确认(验证点 V1) | + +--- + +## 6. 验证计划 + +### 6.1 Golden 与 provenance(强制) + +- **参考实现独立重建**:golden 不得只搬 Python 片段——用 numpy 按 §1.4 数学定义逐行重建 `select_candidate_blocks`(pad→amax→pin→topk→scatter→expand),与用户 PyTorch 实现交叉验证后再作为 ground truth。 +- **sidecar 强制字段**(`xxx.ref.json`,遵循全局 AGENTS 规则):分支 `ds41` + commit `32d64c27f`、candidate 参数(mode/topk_blocks/block_size/mask_mode/cmp_ratio/layout)、随机种子、生成脚本路径、机器(910b 节点)、时间。 +- **使用前核验**:比对结果与预期矛盾时,第一嫌疑人是 golden 本身(先独立复算,再怀疑 kernel)。 + +### 6.2 精度比对策略(2026-09-07 起:对齐官方 result_compare_method 规则) + +`sparse_indices` 与 `candidate_topk_index` 统一采用与 `tests/pytest/result_compare_method.py::check_result` 相同的两级规则(harness `cmp_indices`,行粒度): + +1. **多重集合门**:整行排序后完全相等(值、`-1`、重复均敏感,**顺序不敏感**——MrgSort 的 tie-break 与 PyTorch 稳定排序不保证一致,实测 r2_m1_prefill_m3 存在精确平分对(score 位型相同)的顺序交换,属合法差异)→ 该行 PASS; +2. **边界容忍回退**(门未过时):直接复用官方 `compare_topk_valid`——gold 有效前缀集合比较,差异元素按**边界值**(gold 前缀最后一个元素的分数)相对误差 ≤ thres(0.001) 容忍; +3. **-1 槽数硬校验**(对官方规则的收紧,仅一处):行级有效计数不一致 → 直接 FAIL。官方仅按 gold 的 valid_len 切片、会放过 npu 多填/少填 -1 槽的情况;candidate 的 -1 槽是硬契约(numBlocks < topk_blocks 时必须填 -1),故收紧。 + +自检(selftest_cmp.py)覆盖五分支:同集不同序 PASS / 边界 0.1% 容忍 PASS / 大差异 FAIL / 重复 FAIL / -1 计数不匹配 FAIL。 +当前状态:15 用例全部行过多重集合门(boundary 0 触发),即输出达到"排序后完全一致"。 + +- 构造 tie 密集用例(常值 score)单独验证 tie 不影响**多重集合**正确性; +- 大 shape(16K/128K/1M)harness 为设备侧独立脚本,CPU 全量 golden 不可行,沿用抽样行集合比较 + 有效性/确定性检查(§10)。 + +### 6.3 测试矩阵(每格至少 1 用例,mode=3 抽样回归) + +| 维度 | 取值 | +|---|---| +| candidate_mode | 0 / 1 / 2 | +| cmp_ratio | **1 / 2**(用户指定两类场景) | +| layout_q × layout_k | BSND×PA(已实现)+ **TND×TND**(§11,cu_seqlens_q/k 变长拼接,q=TND/k=PA 组合不支持——PA 依赖 block_table 与 TND 前缀语义冲突) | +| key 0 轴 stride | **紧凑(stride == 块大小,已实现)+ 非连续(stride > 块大小,§11,key 与 k_scale 均需)** | +| q_head_num | **仅 32**(用户指定;全部 candidate 用例统一 g=32,A1 推导下 mBase=128,重点验分核对齐) | +| topk(动态) | 2048 / **min(index_topk, actS2Size) 场景**(topk<2048,验证 -1 padding 与 min 等价性) | +| 阶段 × mask_mode(按模型分支联动) | decode(S1=1)×mode0;prefill 首块(S1>1)×mode3(行级 pin);~~mid-chunk prefill×mode0~~(**暂不测试,问题保留见 R5/§5.0 O3,问题解决后补测**);**mask0 × 大 shape 全遍历**(2026-09-08 补:{16K,128K,1M}×{BSND,TND}×{m1,m2}×{r1,r2}=24 用例;注意 host 契约 mask_mode=0 时 cmp_residual_k 必须不传,ratio=2 的 act_k=2×s2 无 residual 仍为非整除场景) | +| 大 shape × candidate_mode | m1 全遍历 12(§6.6)+ **m2 补齐**(16K/128K-TND/128K-r2/1M 共 5,含 R6 修复后转正的 big128k_m2) | +| 大 shape × key 非连续 | **{16K,128K,1M}×{m1,m2} 全 6 用例**(stride0=2×紧凑, 生产 pool 翻倍 [15872,128,1,128]) | +| numBlocks vs topkBlocks | <, =, > | +| S2 对齐 | actS2Size%8=0 / ≠0(含 actS2Size=1) | +| B | 1 / 2 / 4(变长 batch;2026-09-08 补 B=2 四件套 + B=4 三件套大 shape:m1/m2/mask0 × BSND/TND,变长 seqused_k,B=4 含非对齐尾 tile 98432) | +| topk_blocks | 2048 全量 + **pin 用例 64**(host 已放宽为 (0,2048] 内 64 的倍数;`numBlocks > topk_blocks` 时 pin 最新块必须入选,单/多 tile 各 1 用例) | +| 大规模全遍历(生产规格) | **q_seq {16K, 128K, 1M} × layout {BSND, TND} × cmp_ratio {1, 2} = 12 用例直接并入 pytest 矩阵**(2026-09-08 用户要求;q_seq=s2 全 prefill,mask_mode=3,mode=1,ratio2 带 cmp_residual_k=1);key pool **[7936,128,1,128]**、block_table **[1,1055]**(置换非恒等映射;1M 场景 pool 默认 8192+4)。**CPU 全量 golden 不可行 → 抽样行官方规则比对 + 全行有效性 + pin(强制保留最新 token 所在块)检查**(§6.6);独立大 shape 脚本降级为深检工具 | + +重点组合:mode1+尾块不对齐+causal、**mode1+prefill 首块(验证行级 pin:不同行 pin 不同块)**、mode2+causal、`numBlocks ≤ topkBlocks` 时 mode=2 与 mode=3 的等价性(mask 全选)、**全部用例统一 g=32(mBase=128 分核对齐,A1 推导的本命规格)**、**topk=min(2048, actS2Size) 与 topk=2048+滤-1 的等价性**、mode1+多核(B>1 触发行间切分,验证各核独立直出正确)、cmp_ratio=2 下块划分与 pin(actS2Size 为压缩后长度,cmp_residual_k 参与原始长度还原)。 + +### 6.4 分阶段验证(中间结果可 dump) + +开发期在 `candidate_topk_index` 输出之外保留 debug 开关:dump 块化得分(S5a 后)与块级 merge 中间结果(S6a 后),用于阶段边界定位(kernel 错 vs golden 错)。合入前移除或进 debug 分支。 + +### 6.5 测试脚本设计(基于 test_quant_lightning_indexer_v2_single.py 修改) + +新增 `tests/pytest/test_quant_lightning_indexer_v2_candidate.py`,结构复制自 `test_quant_lightning_indexer_v2_single.py`(保留 SAVE_PT_DIR/RESULT_PATH 环境变量、run_mode eager/graph 分支、QliV2ResultWriter 落盘机制),做以下修改: + +**(1) 参数扩展**:`param_names` 尾部追加 + +```python +("candidate_mode",) # 1=source / 2=consumer / 3=off(默认) +("candidate_topk_blocks",) # 默认 2048 +("candidate_block_size",) # 默认 8 +``` + +`test_data` 元组同步扩展;`QliV2ResultWriter.case_name/row` 的列随之扩展(sidecar 记录 candidate 参数,满足 §6.1 provenance 要求)。 + +**(2) paramset 用例**(新文件内定义或扩展 `test_quant_lightning_indexer_v2_paramset.py`): + +- 基线模板沿用 910b int8 形态(参照 `quant_li_default_a3`:quant_mode=2、qk_dtype=int8、dequant=float16、**q_head_num=32**、layout_key=PA_BBND),去掉 910b 不支持的 return_value/output_idx_offset。 +- **设备分支新增 `Ascend910B`**(当前服务器为 910B3;现有 paramset 仅有 `Ascend910_93`/`Ascend950` 分支,910B 会 NameError)。 +- 用例集(cmp_ratio × mode 全组合 + 边界): + +| 用例名 | cmp_ratio | mode | 覆盖点 | +|---|---|---|---| +| cand_r1_mode3_decode(回归) | 1 | 3 | decode 基线(mask_mode=0,标量 compress_lens;现网行为 bit 级一致) | +| cand_r1_mode1_decode | 1 | 1 | decode source 基础(batch 级 pin) | +| cand_r1_mode1_prefill_m3 | 1 | 1 | **prefill 首块 + mask_mode=3:行级 pin(不同行 pin 不同块)**+ 行级 mask | +| cand_r1_mode1_prefill_m0 | 1 | 1 | ~~mid-chunk prefill + mask_mode=0~~ **暂不执行**(问题 O3 未决保留,见 §5.0/R5;用例定义保留在矩阵中,待确认后启用) | +| cand_r1_mode1_tail | 1 | 1 | actS2Size%8≠0 + causal + 变长 B | +| cand_r1_mode2_self | 1 | 2 | 自洽候选(见 (3)a) | +| cand_r1_mode2_rand | 1 | 2 | 随机候选子集(见 (3)b)+ numBlockstopkBlocks(S2>16K) | +| cand_r2_mode1_decode | 2 | 1 | cmp_ratio=2 + cmp_residual_k 传入 | +| cand_r2_mode1_prefill_m3 | 2 | 1 | cmp_ratio=2 prefill 首块(行级 pin 含 residual 还原) | +| cand_r2_mode2_self | 2 | 2 | cmp_ratio=2 消费 | +| cand_r1_mode1_g32 | 1 | 1 | **q_head_num=32(本轮统一规格;mBase=128,验证 v1 式推导下分核/循环/输出全链路)**。其余用例同规格执行,不再单列 64 回归 | +| cand_r1_mode1_topk_min | 1 | 1/2/3 | **topk=min(2048, actS2Size)<2048:验证 -1 padding 与 min 等价性(与 topk=2048+滤-1 对比)** | + +**(3) mode=2 的 `candidate_topk_index` 输入生成(golden 侧)**: + +- (a) **自洽候选**:对同一 score 调用参考 `select_candidate_blocks` 的输出作为输入(等价于"source 层与 consumer 层权重相同的退化情形",可校验 mode=2 结果 ⊆ mode=0 结果且含 pin 块); +- (b) **随机子集**:从 `[0, numBlocks)` 随机采样 `min(topk_blocks, numBlocks)` 个块(覆盖任意候选集,含 numBlocks candidate_topk_blocks 的抽样行必须含最新块; +4. **双跑确定性**:大 shape 用例二次运行抽样行逐元素一致; +5. **force_rows 确定性锚点(R10,2026-09-08)**:用例可声明 `force_rows` 行号列表并入抽样集——随机抽样可能漏掉特定 vl(行级有效长度)窗口行(R10 的 stale 窗口:尾 tile cuS2Len∈(64,96],即 mask_mode=3 下行号 i 满足 (vl mod 2048)∈(64,96],每 batch 仅约 32 行);big128k_b2_varlen 锚 [64,73,95]×2 batch,big128k_b4_varlen 锚 12 行覆盖 4 batch。 + +### 7. 风险与开放问题 + +| # | 风险/问题 | 影响 | 缓解 | +|---|---|---|---| +| R1 | tie-break 不一致 | 误报精度问题 | §6.2 集合比较 | +| R2 | Sort32/MrgSort 降序假设不成立 | 整体语义反转 | 验证点 V1:实现前单测确认 API 排序方向 | +| R3 | g=32 改动触碰 `mBaseSize/s1BaseSize` 推导,影响分核与循环边界(kernel 与 metadata 基准必须一致) | g=32 分核错乱/越界 | 对齐 v1 推导(§3.1)+ `cand_r1_mode1_g32` 用例 + g=64 全量回归(推导变更影响既有路径,mode=0 也需回归) | +| R4 | source/consumer 的 block_size 不一致(跨算子约定) | 语义错乱 | 文档强约束 + torch 层断言(`quant_lightning_indexer.candidate` python 封装内校验,§4.3) | +| ~~R6(2026-09-07 实测,**2026-09-08 已解**)~~ | mode=2 S2≥128K prefill candBuf 被中途污染 | **根因:同核多行共享 candBuf**——每 AIV 处理 CeilDiv(s1BaseSize,2)=2 行,s2 内层循环按 gS1 块整体推进,后一行(row+2)的 tile0 重排序覆盖前一行候选,前一行 tile1..63 全部读错(三点快照 SNAP/MID/REF 定位:MID==SNAP 证明主路径无辜,漂移精确在 tile 切换);s1=1(每核单行)不触发,故此前所有小 shape 用例漏检 | **修复:candBuf 按行分区**(innerS1Idx × candBlocks × 2 对索引,2 行 32KB,mode=2 UB 184KB≤192KB);新增 r1_m2_prefill(s1=8 mode=2)作为 R6 小 shape 门禁 + big128k_m2 转正;47 用例全绿 | +| R7(§11) | TND 下 candidateOutOffset 与 cuS1Idx 双重前缀(两者均含 cu_seqlens_q 前缀则行号翻倍) | 偏移结构与主输出 indiceOutOffset 完全同构(前缀在 offset、行号 batch 内),理论无双重;tnd_m1_decode 用例显式验证 GM 行对位 | +| R8(§11) | keyStride0 改造影响现网紧凑场景(PA 现网假设 stride==块大小) | keyStride0==0 或 ==紧凑值时走原公式兜底,现网行为 bit 级不变;pa_gap_m3_regress 回归 | +| R9(§11.6) | **output_idx_offset 在 arch22 为死参数**(入口收指针未绑定未消费,host 校验完整但合法值被静默忽略);调用方传非零偏移时 sparse_indices 不含偏移 → 上层绝对位置还原错误 | A15 对齐 arch35 使能;与 candidate 的契约:offset 仅作用于 sparse_indices,candidate_topk_index 保持相对块号(source 输出取加 offset 前),避免 mode=2 掩蔽整行错位 | +| ~~R10(2026-09-08 实测,**当日已解**)~~ | mode=1 大 shape(总块数>candidate_topk_blocks)vl(行级有效长度)尾 tile 块分数 stale:`brmRepeat = blkLen/64` 整除截断,blkLen=96(AlignS2 在 (64,128] 段唯一非 64 倍数输出)时只归约 [0,64),块 8..11 残留上一 tile/行分数 | 实测 big128k_b2_varlen b1 行 73:块 14344 拿 stale 3.5568(真值 1.1753)虚高挤掉第 2048 名边界块 2760(3.0918);**小 shape 总块数≤2048 时候选集合=全部块,stale 分数不改变集合故漏检**;B=1 大 shape 抽样未踩中窗口行(尾 tile cuS2Len∈(64,96] 即行号 i∈[64,95] 每 batch 仅 32 行) | **修复:brmRepeat 改 `CeilDiv(blkLen, 64)`**(多归约的 [96,128) stale 只落 pad 块槽位,被 -inf 位型链位精确覆盖,无害);harness 新增 force_rows 确定性锚点(§6.6),B=2 锚 6 行、B=4 锚 12 行覆盖各 batch 窗口 | +| **R5(标红,O3 未决)** | **mid-chunk prefill(start_pos≠0 且 S1>1)形态下模型不做行内因果 mask,与常规认知相反**;该调用形态是否存在未确认 | 若误当作 bug "修复"(擅自加行级 mask)将引入语义变更;测试碰 mode0+S1>1 数据可能误报精度问题 | **暂不测试该用例,问题保留**(§5.0 O3 标红段);三不准:不准顺手加 mask / 不准以"常规认知"为 golden / 处理前必须先与模型侧确认调用形态 | + +已决策记录: + +- ~~O1~~ **已定**:torch 接口采用方案 b——新增 overload `quant_lightning_indexer.candidate(...) -> (Tensor, Tensor, Tensor)`,旧 schema 二元组不动(§4.3)。 +- ~~O2~~ **已定**:`candidate_topk_blocks` 当前仅支持 2048;属性保留、host 校验限定 2048,未来扩展只放宽校验不改接口;非默认值测试用例已删除。 + - **2026-09-07 更新**:已放宽为 (0, 2048] 内 64 的倍数(commit a346acdcb,pin 语义验证需要 topk_blocks=64);BASE_TOPK=2048 上限不变。 + +## 8. 实施拆解(文件级) + +| 文件 | 变更 | +|---|---| +| `op_host/quant_lightning_indexer_v2_tiling.h` | TilingData +3 字段;ParaInfo/常量(CANDIDATE_* 索引) | +| `op_host/quant_lightning_indexer_v2_tiling.cpp` | 属性解析与校验(mode 互斥、范围、consumer 必传输入);**`GetGSize` 910b 分支放宽允许 gSize=32(适配点 A2)**;TilingData 写入(workspace 计算不变,A6) | +| `op_kernel/arch22/quant_lightning_indexer_v2_kernel_arch22.h` | **`InitTilingData` 改 v1 式推导(A1):`s1BaseSize=4` 固定、`mBaseSize = s1BaseSize×gSize`**;删除/降级 `M_BASE_SIZE=256` 常量;Init 传参/新 GM 张量(candidateTopkIndexInGm/candidateTopkIndexOutGm)、Vec1 调用点注入 | +| `op_kernel/arch22/quant_lightning_indexer_v2_service_vector_arch22.h` | **Vec0 `cuProcEleNum` 补 `CeilAlign(·, 32)`(A5,同步 v1:289)**;S4a/S5a/S6a 实现、globalBlockTopkUb_、CleanInvalidCandidateOutput(不动 ProcessLD) | +| `tests/ut/op_host/arch22/` | **补 g=32 tiling 单测(A10)** | +| `op_host/quant_lightning_indexer_v2_def.cpp` | 新 input/output/attrs(仅 ascend910b 配置) | +| `op_host/quant_lightning_indexer_v2_infershape.cpp` | candidate_topk_index 输出 shape/dtype | +| `op_kernel/quant_lightning_indexer_v2.cpp` | arch22 入参透传(arch35 分支不动) | +| `examples/` + `docs/aclnnQuantLightningIndexerV2.md` | aclnn 用例与接口文档 | +| `torch_extension/` | schema overload + python 封装 | +| `tests/pytest/test_quant_lightning_indexer_v2_candidate.py` | **新增**,基于 `test_quant_lightning_indexer_v2_single.py` 修改(§6.5) | +| `tests/pytest/quant_lightning_indexer_v2_golden.py` | `select_candidate_blocks_ref`(numpy+torch 双实现)、`GeneralizedQLIV2` 的 mode=1/2 golden 路径 | +| `tests/pytest/result_compare_method.py` | `check_result_candidate`(集合比较 + -1 槽核对) | +| `tests/pytest/qliv2_test_utils.py` | case_name/row 列扩展(candidate 参数入 sidecar) | +| `README.md`(两算子) | 参数表、约束(block_size 一致性、同 forward 有效期、910b 无 LD 说明) | +| **A11** arch22 cube/vector:`KeyNd2NzForPA`/`GetKeyScale` 改用 `keyStride0`/`keyDequantScaleStride0`(对齐 arch35,tiling 字段已存在未消费;0 或紧凑值兜底原公式) | key 0 轴非连续 | +| **A12** op_host tiling.cpp:删除 candidate 的 layout_q=BSND 限制,加 TND 输入校验(cu_seqlens_q 必传) | TND 放行 | +| **A13** arch22 kernel:验证 TND 下 candidateTopkIndexIn/Out GM offset 与 CleanInvalidOutput 的 TND 输出分支(已有 outputLayout==TND 分支) | TND candidate | +| **A14** torch_extension:py/csrc 的 TND 分支透传与输出 shape(ConstructOutputTensor 已有 TND 分支) | TND 封装 | +| **A15** arch22 kernel:使能 output_idx_offset(入口 SetGlobalBuffer + 每行标量读 + 输出前 int32 向量 Adds,参照 arch35 vector:680/856 与 IndicesAddOffset);candidate_topk_index 不加 offset(相对块号契约) | offset 使能 | + +**明确不做**:arch35 (950) 全部路径;LD/ProcessDecode 相关改动;metadata 算子代码;workspace 布局变更;TilingKey 变更。 + +## 9. 实现期实测平台约束(2026-09-07 调试结论,v220 / 910b) + +实现与调试期间实证的平台铁律,均已在代码注释中标注,后续维护必须遵守: + +| # | 约束 | 违反表现 | 正确做法 | +|---|---|---|---| +| P1 | 向量指令 count 必须 64 对齐 | `Duplicate(count=1)`(pin 单元素写)触发 aicore 异常 507015 | 少量元素写入用标量 `SetValue`(须配 P2)或并入 64 对齐向量链 | +| P2 | 标量 UB 写/读与 V 管道互不被 `PipeBarrier` fence | pad 块 -1 填充被 SortAll 抢跑覆盖;排序后 `GetValue` 二分读到 stale 数据(mode=2 IoU=0) | 标量读 V 写结果前 `SetFlag/WaitFlag`(参照 CANN topk_v200 实现);标量写后对称 S_V;或彻底纯向量化 | +| P3 | UB→UB `DataCopy`(走 MTE 管道)不被 `PipeBarrier` 等待 | 跨 tile MergeSort 回拷累加器后,下一 tile 的 MrgSort 读到 stale(tile1 块整段缺失 + -inf 位型+8k 垃圾×4) | 回拷改 int32 视图 `Adds(+0)` 分块(`MergeSortVecCopy`);**float 域 Adds 会把 0xFFFFFFFF(-1 的负符号 NaN 位型)规范化为 0x7FFFFFFF**,必须 int32 | +| P4 | v220 Cast:s32→f32 `CAST_NONE` 合法(vconv_s322f32);**f32→s32 `CAST_NONE` 为 assert no-op**(仅 RINT/FLOOR/CEIL/ROUND/TRUNC 有指令) | fp32 算术链构造的 idx 经 Cast 后保持浮点位型垃圾(score 位型泄漏进索引字段) | 索引构造全程 int32 域(`ArithProgression` + Sub/Maxs/Mins/Mul 链);确需 f32→s32 时用 `CAST_RINT` | +| P5 | kernel `.o` 编译自 `build/binary/ascend910b/src/` 源码副本(cmake configure 时拷贝,make 不刷新),受 `gen/*.done` 门闩控制;且该路径下 `compile_stop.flag` 会阻塞失败后的重试 | 连续三轮"修复无效"实为旧二进制(AGENTS.md 教训复现,部署 .h 源码是新的、编译产物是旧的,标记检查被假象通过) | build_install.sh 流水:rsync 源码副本 → 删 `.done` 门闩 → make → **校验 .o mtime 为本次** → 安装 → `cmp` 部署 .o 与构建 .o 一致 | + +## 10. 当前验证状态(2026-09-07) + +- pytest **87 用例全绿**(2026-09-08 R10 修复 + B=2/B=4 变长补齐后:24 基础 + 12 m1 大遍历 + 5 m2 大补 + 6 pa_gap 大 + 24 mask0 大遍历 + 4 B=2 变长 + 3 B=4 变长 + 其余 TND/offset/门禁;总时长 ~32 分钟): + - 22 基础用例(mode 1/2/3 × cmp_ratio 1/2 × g32/64 × decode/prefill/tail/sparse2048/pin64 单/多 tile/B=4 变长/S1 尾块/tiny/cand64×mode2)全过,零 aicore + - **大 shape 直接入矩阵全过**:big16k_m1 / big128k_m1 / big128k_m1_b2(B=2 变长)/ big1m_m1(131072 块 pin(强制保留最新 token 所在块)抽样验证)/ big1m_m3(回归)——§6.6 抽样行官方两级规则比对 + 全行有效性 + pin 检查,总时长 ~4 分钟 + - 等价回归全过:pa_gap_m3_regress(紧凑场景现网行为不变,R8)、off_zero_regress(零偏移与不传一致) + - **A11/A12/A15 实施后(2026-09-08):tnd_*×7、pa_gap_m1/m2/128k×3、off_m1_decode/m2/off_tnd×3 全部 XPASS 转正;pa_gap_m3_regress(紧凑等价)与 off_zero_regress(零偏移等价)回归保持通过 + - **R6 修复转正(2026-09-08):candBuf 按行分区(同核 2 行共享是根因)**——big128k_m2 通过;新增 r1_m2_prefill(s1=8 mode=2)作为 R6 小 shape 门禁;大 shape 全遍历 12 用例(q_seq×layout×ratio)全过;xfail 清零 + - **R10 修复转正(2026-09-08):blkLen=96 时 brmRepeat 整除截断致尾 tile 块分数 stale(§7 R10)**——B=2 变长四件套(big128k_b2_varlen/tnd_m2/mask0_r2 + big1m_b2_m1)+ B=4 变长三件套([65536, 98432, 131072, 81920],b1=98432 使尾 tile 窗口行移位覆盖非对齐 vl)全过;修复后全量 84 + B=4 新增 3 = 87 用例全绿 +- 大规模生产规格(`/opt/tjj/qli_cand_test/test_big_shape_m1.py`,设备侧 harness):16K(b=1) / 128K(b=2) mode=1、128K(b=2) mode=2 随机半数候选、1M(b=1) mode=1 全部 IoU=1.000000;**1M 场景 pin 块 131071 确认入选**(131072 块中选 top-2048) +- 大 shape 回归:1M mode=3(IoU=1.0 抽样行精确比对 + 双跑确定性)、1M cmp_ratio=2、qseq 256..2048 扫描全过 +- 提交序列:`aede3e62e`(功能实现)→ `a346acdcb`(五处正确性修复:pad 块 int32 向量算术填充 / MergeSortVecCopy 位精确回拷 / pin 纯向量化(全局块号比较)/ mode=2 CountGE 窗口边界 + CAST_RINT + isOutI32 区域重叠 / V_S fence) + +## 11. TND 与 key 0 轴非连续支持(2026-09-07 需求追加,待实施) + +### 11.1 需求(2026-09-08 用户澄清后修正) + +1. **layout_q = TND(仅 Q 侧)**:q 为变长拼接 `[T, G, D]` + cu_seqlens_q;**layout_k 固定 PA_BBND(K 仅支持分页布局,不支持 TND)**——q=TND 时 metadata 强制要求 layout_k=PA_BBND。candidate 三 mode 全支持;现网 kernel 已有 TND 模板分支,本轮打通 host 校验(GetS1Size 的 TND 分支补齐 s1Size=T)、candidate 的 GM offset 与输出布局。 +2. **key 0 轴非连续**(PA_BBND):key 与 k_scale 的第 0 维物理 stride 可大于块逻辑大小(block_table 指向的物理块之间存在间隙,如池化重组后的非紧凑存储)。 + +### 11.2 arch35 参照机制(已实现,直接移植) + +- **key 主体**(cube `KeyNd2NzForPA`):`blkTable.GetValue(bIdx*maxBlockNumPerBatch + s2BlkId) * constInfo_.keyStride0 + s2BlkOffset*headDim`(arch35/quant_lightning_indexer_v2_service_cube_arch35.h:437); +- **k_scale**(vector `GetKeyScale`):`blockId * constInfo_.keyDequantScaleStride0 + startBlockTableOffset`(arch35/..._service_vector_arch35.h:456); +- **host 侧**:tiling.cpp 从 acl tensor `keyStridesVec_[0]` / `keyDequantScaleStridesVec_[0]` 取真实 stride 经 `set_keyStride0`/`set_keyDequantScaleStride0` 下发(tiling.h 字段已存在); +- **arch22 现状差异**:`KeyNd2NzForPA` 硬编码 `blk * kCacheBlockSize * kHeadNum * headDim`(service_cube_arch22.h:317)、`GetKeyScale` 硬编码 `blockId * kCacheBlockSize_`(service_vector_arch22.h:157)——均假设紧凑存储;tiling 字段已下发但 **arch22 kernel 未消费**。 +- **实施补丁(A11)**:kernel 两处寻址改用 `keyStride0`/`keyDequantScaleStride0`(0 时兜底原紧凑公式,现网行为不变);host `CheckKeyContiguous` 的 0 轴放行从仅 arch35 扩展到全部 PA_BBND。**关键发现:aclnn 动态调用下 `GetDynamicInputStride` 恒为空**(仅 TensorV2 图模式携带非连续描述,exe_graph TensorV1 的 GetStride 亦为空)——新增可选属性 `key_stride0`/`key_dequant_scale_stride0`(def.cpp + tiling 下发,csrc 自动取 `key.stride(0)` 传入,调用方无感),GetDynamicInputStride 有值时优先。 + +### 11.3 适配点(A11–A14,见 §8 表) + +**TND 偏移结构(A13 核心)**:`candidateOutOffset = cu_seqlens_q(bIdx) × kHeadNum × candBlocks + n2Idx × candBlocks`(batch 级前缀,CalcRunInfo 已含 TND 分支),行偏移 `cuS1Idx(batch 内行号)× candBlocks`——与主输出 `indiceOutOffset` 完全同构;输出布局 `[T, N2, K]`(csrc ConstructOutputTensor 已有分支),无效行清理走既有 `outputLayout == TND` 分支(kernel_arch22.h:418)。**实施补丁(A12/A13)**:host GetS1Size 补 TND 分支(s1Size = q.shape[0],原仅 BSND 赋值导致 TND 下 s1Size=0、consumer shape 校验 expectSize=0 误报);candidate host 校验放行 layout_q∈{BSND,TND} 且 TND 时强制 layout_k=PA_BBND + cu_seqlens_q 必传。 + +### 11.4 测试用例设计(§6.3 矩阵新增行) + +**TND(q=k=TND,cu_seqlens 变长)**——pytest id:`tnd_m1_decode` / `tnd_m1_prefill` / `tnd_m2_rand` / `tnd_m3_regress` / `tnd_r2_m1` / `tnd_m1_pin64`: + +| 用例 | B | seqs_q | seqs_k | ratio | mask | mode | 要点 | +|---|---|---|---|---|---|---|---| +| tnd_m1_decode | 2 | [1,1] | [1024,2048] | 1 | 3 | 1 | 变长 TND decode;**验证 GM 行对位(R7)** | +| tnd_m1_prefill | 2 | [8,4] | [1024,2048] | 1 | 3 | 1 | TND prefill 行级 pin;seqs_q 非对称 | +| tnd_m2_rand | 1 | [1] | [4096] | 1 | 0 | 2 | TND 消费(候选输入 GM offset 走 TND 前缀) | +| tnd_m3_regress | 2 | [4,4] | [1024,2048] | 1 | 3 | 3 | TND 现网回归 | +| tnd_r2_m1 | 1 | [1] | [1024] | 2+[1] | 3 | 1 | TND + cmp_ratio=2 + residual | +| tnd_m1_pin64 | 1 | [1] | [8192] | 1 | 0 | 1 | TND 多 tile + pin(candBlocks=64) | + +**key 0 轴非连续(PA_BBND,key/k_scale stride(步长)翻倍,块间间隙)**——pytest id:`pa_gap_m1_decode` / `pa_gap_m2` / `pa_gap_m3_regress` / `pa_gap_128k`: + +| 用例 | 要点 | +|---|---| +| pa_gap_m1_decode | `as_strided` 构造 stride=2×紧凑 的 key/k_scale,block_table 只指向偶数物理块;mode=1 | +| pa_gap_m2 | 同上 mode=2(GetKeyScale 的 stride 路径同时覆盖) | +| pa_gap_m3_regress | **紧凑场景回归(R8)**:现网行为 bit 级不变的等价性验证 | +| pa_gap_128k | 大 shape 128K + 非连续 + 官方两级规则抽样比对 | + +**大 shape 交叠**(并入 pytest,§6.6 抽样机制):`tnd_big16k_m1` / `tnd_big128k_m1`(TND 无 block_table,用 cu_seqlens 拼接总池);`pa_gap_128k`(非连续 × 128K);`off_tnd_m1`。 + +**xfail(预期失败)标注**:A11/A12/A15 实施前,tnd_*(host 拒绝 TND)、pa_gap_m1/m2/128k(kernel 紧凑寻址读到 gap 段零值)、off_m1_decode/off_tnd(offset 被静默忽略)以 `pytest.mark.xfail(strict=True)` 入矩阵——实施通过后 XPASS 自动报警提醒摘标记;`pa_gap_m3_regress`(紧凑等价回归)与 `off_zero_regress`(零偏移等价回归)当前实现即应通过。 + +### 11.6 output_idx_offset 使能(A15,2026-09-07 追加调查) + +**现状(A15 已实施,2026-09-08)**:原为死参数(入口收指针未绑定未消费)。现已使能:入口 SetGlobalBuffer + 经 InitVecCandidateTensor 传 tensor 与有效标志(注意 **InitParams 先值拷贝 constInfo**,之后再改 constInfo 标志不生效——标志必须走传参);主输出路径(ProcessVec1 的 needCopyOutGm 分支)与 LD(低延迟归约)路径(ProcessLD 搬出前)各加一处消费:GM 标量读行偏移(读 offset 无 V→S 竞态)+ int32 向量 Adds(64 分块、零偏移零开销、-1 槽位精确不变)。**注**:910b 上 LD 未启用(metadata AICPU 内核 supportFd_ 仅 ASCEND950 置位,fdUsedVecNum=0,isLdCoreEnable 恒 false,ProcessDecode 的 ProcessLD 为死路)——LD 路径消费点为将来 LD 启用的对称实现。 + +**语义(arch35 已实现,直接移植)**:每行一个 int32 偏移,输出拷 GM 前对 sparse_indices 逐元素 `+= offset`(`IndicesAddOffset`:int32 向量 Adds,64 对齐);用于 TND 多请求聚合(各请求 KV cache 起始不同,kernel 内相对 → 输出绝对)。消费点:arch35 vector:681(`outputIdxOffsetGm.GetValue(outputIdxCoreOffset + rowIdx*kHeadNum)`,GM 标量读无 V_S 竞态)与 :856(非零才加,零偏移路径零开销);`outputIdxCoreOffset` 与 indiceOutCoreOffset 同构(TND 前缀)。 + +**与 candidate 的契约(设计决策)**:offset 仅作用于 sparse_indices;`candidate_topk_index`(source 输出/consumer 输入)一律为**加 offset 前的 batch 内相对块号**——否则 mode=2 掩蔽需同步减偏移,且跨层共享(shared_attn.candidates)时双方 offset 可能不同会导致掩蔽错位。跨层传递绝对块号的换算由 torch 封装层提供工具函数(与 §1.3 "expand 工具"同理)。 + +**测试用例**: + +| 用例 | 要点 | +|---|---| +| off_m1_decode | BSND + output_idx_offset=[1000],sparse_indices 每元素 +1000(官方比对:golden 行 + offset),candidate 输出不受影响(仍相对块号) | +| off_tnd_m1 | TND + 每 batch 偏移(cu_seqlens_k 前缀作为 offset),输出绝对 KV 位置 | +| off_zero_regress | offset=0 行为与不传 bit 级一致(零偏移路径零改动) | +| off_m2 | mode=2 + offset:掩蔽用相对候选,输出 sparse_indices 仍加 offset(consumer 场景契约) | + +### 11.5 实施顺序(2026-09-08 全部完成) + +1. ~~A11~~ ✅ kernel 寻址 + host 放行 + stride attr 通路(GetDynamicInputStride aclnn 动态调用恒空); +2. ~~A12/A13/A14~~ ✅ TND(仅 Q 侧):host 放行 + GetS1Size TND 分支 + GM offset 同构验证; +3. ~~A15~~ ✅ output_idx_offset 使能(InitParams 值拷贝陷阱经 InitVecCandidateTensor 传标志); +4. ✅ 大 shape 交叠:pa_gap_128k / tnd_big128k / off_tnd 全过。 +5. ✅ 大 shape 全遍历(2026-09-08):12 用例(q_seq {16K,128K,1M} × layout {BSND,TND} × ratio {1,2})全过(16K/128K 批 65s、1M 批 7m15s)。**修复 harness pin 检查公式**:mask_mode=3 下行级有效长度必须除 cmp_ratio(`(act_k−S1+i+1)//ratio`,与 kernel/golden 一致;原式漏除在小 shape 因 clamp s2 掩盖碰巧通过,大 shape 行 0 暴露——错误的参考检查比没有检查更危险)。 + +**部署要点(踩坑记录)**:attr 变更后编译产物 hash 改变(18e4fb→d566),但 `bin/quant_lightning_indexer_v2.json`(bin 选择配置)残留旧映射——**必须删除该 json 与 binary_info_config.json 强制重生成**,否则运行时按旧映射找不到新 .o(stat file failed),或加载旧产物(新改动静默不生效)。build_install.sh 已加 autogen 头清理,需同步加 bin config 清理。 + +## 12. 新接口规范 py 适配层(2026-09-09) + +面向新芯片接口规范(mxfp4/uint8/e8m0 descale、candidate_block_indices/candidate_block_length 命名)的 **py 分发封装**:按规范签名收参,内部映射到已注册的 `quant_lightning_indexer_candidate`(分别固定 candidate_mode=1/2),**不改 csrc / 算子校验 / 数据类型**;调用侧按芯片选择接口(本后端 910b 走本适配层,新芯片走其自带实现)。 + +**新增入口**(`torch_extension/quant_lightning_indexer.py` 末尾;pip 包 `cann_ops_transformer.ops` 与源仓库 `torch_extension/__init__.py` 同步导出): + +- `quant_lightning_indexer_candidate_source(...)` — 规范接口1(source):返回 4 元组 `(sparse_indices (T1,N2,k), sparse_values, candidate_block_indices (T1,N2,cb), candidate_block_length (T1,N2))`; +- `quant_lightning_indexer_candidate_consumer(...)` — 规范接口2(consumer):输入 `candidate_block_indices (T1,N2,cb)` + `candidate_block_length`,返回 2 元组 `(sparse_indices, sparse_values)`(布局随内部:B=1 BSND 4 维 / TND 3 维)。 + +**映射约定(规范项 → 内部接口)**: + +| 规范项 | 映射 | +|---|---| +| `q_descale` / `k_descale` | `query_dequant_scale` / `key_dequant_scale`(仅改名) | +| `candidate_block_indices` | `candidate_topk_index`(仅改名;块级相对块号;BSND 路径自动补/去 batch 维) | +| `seqused_q`(规范注释为每 batch **key** 截断) | `seqused_k`(歧义点①:若实为截断 query 行数则改一行) | +| `candidate_block_length` | mode=1 输出:py 公式 `vl(i)=clamp((act_k−S1+i+1)//ratio, 0, K)`(act_k=K×ratio+residual,与 golden 同式);mode=2 输入:**忽略**(op 内部已按 mask 规则与 K 取小截断,语义冗余) | +| `candidate_topk_blocks=-1`(无机制默认) | source 场景取 2048;consumer 由 `candidate_block_indices.shape[-1]` 推导 | +| layout 参数消失 | q 恒 3 维 (T1,N1,D):有 `cu_seqlens_q`=TND 直传;无=B=1 BSND(q/w/q_descale 补 batch 维,输出端去掉);**B>1 且无 cu_seqlens_q 显式报错**(batch 维丢失防御) | +| `metadata` 缺省 | py 自动调 metadata 算子生成(shape 推导;含一次 `.item()` 主机同步,调用方可预生成传入绕过) | +| `output_idx_offset`(仅接口2) | 透传(A15 既有支持,仅作用于 sparse_indices) | + +**本后端限制(py 层显式报错,非 op 校验改动)**:`block_table` 必传(key 仅支持 PA_BBND 分页布局);`return_value=True` 不支持(内部入口硬编码 False,不改 csrc)。 + +**验证**:pytest `test_qli_newapi*` 4 用例(BSND B=1 prefill mask3+ratio2+residual / BSND B=1 decode mask0 / TND B=2 变长 mask3 / 错误路径×5)——source/consumer 与旧入口 bit 级一致 + metadata 自动/手传 bit 级一致 + block_length 公式对照 + B>1 无 cu_seqlens_q 防御。既有 87 用例零影响(旧入口抽查回归通过)。 + +**全量矩阵经新接口重跑(2026-09-09,`qli_cand_test/test_qli_newapi_full.py`)**:87 用例中 **76 全过 + 11 规范限制跳过 + 0 失败**(31m54s,机制=monkeypatch `npu_run` → 适配层,主比较逻辑/官方两级规则/pin/全行有效性原样复用)。跳过项两类均为规范表达能力边界而非适配层缺陷:① B>1 且 BSND ×9(新接口 q 恒 3 维,多 batch 必须 cu_seqlens_q/TND——含 big128k_b2_varlen/big1m_b2_m1 等);② mode=1 + 非零 output_idx_offset ×2(off_m1_decode/off_tnd_m1,规范 source 无该参数,仅 consumer 有)。覆盖确认:pa_gap 非连续(大 shape ×6)、pin64、g64、tiny(s2=1)、sparse2048、mask0 大 shape 全遍历、TND B=2/B=4 变长、off_m2(consumer+offset 透传)、off_zero_regress(source 零偏移等价回归)均过。 + +**规范侧待澄清(不阻塞本适配层)**:seqused_q 注释 keys 与命名矛盾;`ori_sparse_indices` 未在输入列表;k_descale 末维 `2` 语义(新芯片侧消费);mode=1 输出 BSND 4 维变体(本适配层统一 3 维 (T1,N2,·),调用侧在新芯片侧由其实现给出)。 diff --git a/csrc/attention/quant_lightning_indexer_v2/examples/test_aclnn_quant_lightning_indexer_v2.cpp b/csrc/attention/quant_lightning_indexer_v2/examples/test_aclnn_quant_lightning_indexer_v2.cpp new file mode 100644 index 000000000000..e5c52f0dce0f --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/examples/test_aclnn_quant_lightning_indexer_v2.cpp @@ -0,0 +1,455 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file test_aclnn_quant_lightning_indexer_v2.cpp + * \brief + */ +#include +#include +#include +#include +#include "securec.h" +#include "acl/acl.h" +#include "aclnnop/aclnn_quant_lightning_indexer_v2.h" +#include "aclnnop/aclnn_quant_lightning_indexer_v2_metadata.h" +#include "aclnn/opdev/platform.h" + +using namespace std; + +namespace { + +#define CHECK_RET(cond) ((cond) ? true : (false)) + +#define LOG_PRINT(message, ...) \ + do { \ + (void)printf(message, ##__VA_ARGS__); \ + } while (0) + +int64_t GetShapeSize(const std::vector &shape) +{ + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} + +int Init(int32_t deviceId, aclrtStream *stream) +{ + auto ret = aclInit(nullptr); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclInit failed. ERROR: %d\n", ret); + return ret; + } + ret = aclrtSetDevice(deviceId); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); + return ret; + } + ret = aclrtCreateStream(stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); + return ret; + } + return 0; +} + +template +int CreateAclTensor(const std::vector &hostData, const std::vector &shape, void **deviceAddr, + aclDataType dataType, aclTensor **tensor) +{ + auto size = GetShapeSize(shape) * sizeof(T); + auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); + return ret; + } + + ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); + return ret; + } + + std::vector strides(shape.size(), 1); + for (int64_t i = shape.size() - 2; i >= 0; i--) { + strides[i] = shape[i + 1] * strides[i + 1]; + } + + *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), *deviceAddr); + return 0; +} + +struct TensorResources { + void *queryDeviceAddr = nullptr; + void *keyDeviceAddr = nullptr; + void *weightsDeviceAddr = nullptr; + void *qScaleDeviceAddr = nullptr; + void *kScaleDeviceAddr = nullptr; + void *metadataDeviceAddr = nullptr; + void *sparseIndicesDeviceAddr = nullptr; + void *sparseValuesDeviceAddr = nullptr; + + aclTensor *queryTensor = nullptr; + aclTensor *keyTensor = nullptr; + aclTensor *weightsTensor = nullptr; + aclTensor *qScaleTensor = nullptr; + aclTensor *kScaleTensor = nullptr; + aclTensor *metadataTensor = nullptr; + aclTensor *sparseIndicesTensor = nullptr; + aclTensor *sparseValuesTensor = nullptr; +}; + +int InitializeTensors(TensorResources &resources) +{ + int64_t B = 2; + int64_t S1 = 4; + int64_t S2 = 8; + int64_t N1 = 64; + int64_t N2 = 1; + int64_t D = 128; + int64_t topk = 512; + + std::vector queryShape = {B, S1, N1, D}; + std::vector keyShape = {B, S2, N2, D}; + std::vector weightsShape = {B, S1, N1}; + std::vector qScaleShape = {B, S1, N1}; + std::vector kScaleShape = {B, S2, N2}; + std::vector metadataShape = {1024}; + std::vector sparseIndicesShape = {B, S1, N2, topk}; + std::vector sparseValuesShape = {B, S1, N2, topk}; + + int64_t queryShapeSize = GetShapeSize(queryShape); + int64_t keyShapeSize = GetShapeSize(keyShape); + int64_t weightsShapeSize = GetShapeSize(weightsShape); + int64_t qScaleShapeSize = GetShapeSize(qScaleShape); + int64_t kScaleShapeSize = GetShapeSize(kScaleShape); + int64_t metadataShapeSize = GetShapeSize(metadataShape); + int64_t sparseIndicesShapeSize = GetShapeSize(sparseIndicesShape); + int64_t sparseValuesShapeSize = GetShapeSize(sparseValuesShape); + + std::vector queryHostData(queryShapeSize, 0x38); + std::vector keyHostData(keyShapeSize, 0x38); + std::vector weightsHostData(weightsShapeSize, 0.01f); + std::vector qScaleHostData(qScaleShapeSize, 1.0f); + std::vector kScaleHostData(kScaleShapeSize, 1.0f); + std::vector metadataHostData(metadataShapeSize, 0); + std::vector sparseIndicesHostData(sparseIndicesShapeSize, 0); + std::vector sparseValuesHostData(sparseValuesShapeSize, 0); + + int ret = CreateAclTensor(queryHostData, queryShape, &resources.queryDeviceAddr, aclDataType::ACL_FLOAT8_E4M3FN, + &resources.queryTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + return ret; + } + + ret = CreateAclTensor(keyHostData, keyShape, &resources.keyDeviceAddr, aclDataType::ACL_FLOAT8_E4M3FN, + &resources.keyTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + return ret; + } + + ret = CreateAclTensor(weightsHostData, weightsShape, &resources.weightsDeviceAddr, aclDataType::ACL_FLOAT, + &resources.weightsTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + return ret; + } + + ret = CreateAclTensor(qScaleHostData, qScaleShape, &resources.qScaleDeviceAddr, aclDataType::ACL_FLOAT, + &resources.qScaleTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + return ret; + } + + ret = CreateAclTensor(kScaleHostData, kScaleShape, &resources.kScaleDeviceAddr, aclDataType::ACL_FLOAT, + &resources.kScaleTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + return ret; + } + + ret = CreateAclTensor(metadataHostData, metadataShape, &resources.metadataDeviceAddr, aclDataType::ACL_INT32, + &resources.metadataTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + return ret; + } + + ret = CreateAclTensor(sparseIndicesHostData, sparseIndicesShape, &resources.sparseIndicesDeviceAddr, + aclDataType::ACL_INT32, &resources.sparseIndicesTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + return ret; + } + + ret = CreateAclTensor(sparseValuesHostData, sparseValuesShape, &resources.sparseValuesDeviceAddr, + aclDataType::ACL_BF16, &resources.sparseValuesTensor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + return ret; + } + + return ACL_SUCCESS; +} + +int GenerateMetadata(TensorResources &resources, aclrtStream stream, int64_t B, int64_t S1, int64_t S2, int64_t N1, + int64_t N2, int64_t D, int64_t topk, int64_t quantMode, int64_t maskMode, int64_t cmpRatio) +{ + constexpr const char layoutQ[] = "BSND"; + constexpr const char layoutK[] = "BSND"; + constexpr size_t layoutLen = sizeof(layoutQ); + char layoutQCopy[layoutLen]; + char layoutKCopy[layoutLen]; + errno_t memcpyRet = memcpy_s(layoutQCopy, sizeof(layoutQCopy), layoutQ, layoutLen); + if (!CHECK_RET(memcpyRet == 0)) { + LOG_PRINT("metadata memcpy_s layoutQ failed. ERROR: %d\n", memcpyRet); + return -1; + } + memcpyRet = memcpy_s(layoutKCopy, sizeof(layoutKCopy), layoutK, layoutLen); + if (!CHECK_RET(memcpyRet == 0)) { + LOG_PRINT("metadata memcpy_s layoutK failed. ERROR: %d\n", memcpyRet); + return -1; + } + + aclOpExecutor *executor; + uint64_t workspaceSize = 0; + int ret = aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize( + nullptr, nullptr, nullptr, nullptr, nullptr, N1, N2, D, topk, quantMode, B, S1, S2, layoutQCopy, layoutKCopy, + maskMode, cmpRatio, resources.metadataTensor, &workspaceSize, &executor); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize failed. ERROR: %d\n", ret); + return ret; + } + + void *metadataWsAddr = nullptr; + if (workspaceSize > 0ULL) { + ret = aclrtMalloc(&metadataWsAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("metadata allocate workspace failed. ERROR: %d\n", ret); + return ret; + } + } + + ret = aclnnQuantLightningIndexerV2Metadata(metadataWsAddr, workspaceSize, executor, stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclnnQuantLightningIndexerV2Metadata failed. ERROR: %d\n", ret); + if (metadataWsAddr) { + (void)aclrtFree(metadataWsAddr); + } + return ret; + } + + ret = aclrtSynchronizeStream(stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("metadata synchronize stream failed. ERROR: %d\n", ret); + if (metadataWsAddr) { + (void)aclrtFree(metadataWsAddr); + } + return ret; + } + + if (metadataWsAddr) { + (void)aclrtFree(metadataWsAddr); + } + return ACL_SUCCESS; +} + +int ExecuteQuantLightningIndexerV2(TensorResources &resources, aclrtStream stream, void **workspaceAddr, + uint64_t *workspaceSize) +{ + int64_t topk = 512; + int64_t quantMode = 1; + int64_t maskMode = 0; + int64_t cmpRatio = 1; + int64_t returnValue = 1; + constexpr const char layoutQStr[] = "BSND"; + constexpr const char layoutKStr[] = "BSND"; + constexpr size_t layoutQLen = sizeof(layoutQStr); + constexpr size_t layoutKLen = sizeof(layoutKStr); + char layoutQ[layoutQLen]; + char layoutK[layoutKLen]; + errno_t memcpyRet = memcpy_s(layoutQ, sizeof(layoutQ), layoutQStr, layoutQLen); + if (!CHECK_RET(memcpyRet == 0)) { + LOG_PRINT("memcpy_s layoutQ failed. ERROR: %d\n", memcpyRet); + return -1; + } + memcpyRet = memcpy_s(layoutK, sizeof(layoutK), layoutKStr, layoutKLen); + if (!CHECK_RET(memcpyRet == 0)) { + LOG_PRINT("memcpy_s layoutK failed. ERROR: %d\n", memcpyRet); + return -1; + } + aclOpExecutor *executor; + + int ret = aclnnQuantLightningIndexerV2GetWorkspaceSize( + resources.queryTensor, resources.keyTensor, resources.weightsTensor, resources.qScaleTensor, + resources.kScaleTensor, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, resources.metadataTensor, + topk, quantMode, -1, layoutQ, layoutK, maskMode, cmpRatio, returnValue, resources.sparseIndicesTensor, + resources.sparseValuesTensor, workspaceSize, &executor); + + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclnnQuantLightningIndexerV2GetWorkspaceSize failed. ERROR: %d\n", ret); + return ret; + } + + if (*workspaceSize > 0ULL) { + ret = aclrtMalloc(workspaceAddr, *workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); + return ret; + } + } + + ret = aclnnQuantLightningIndexerV2(*workspaceAddr, *workspaceSize, executor, stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclnnQuantLightningIndexerV2 failed. ERROR: %d\n", ret); + return ret; + } + + return ACL_SUCCESS; +} + +int PrintOutResult(const std::vector &shape, void *deviceAddr) +{ + auto size = GetShapeSize(shape); + std::vector resultData(size, 0); + auto ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), deviceAddr, + size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); + return ret; + } + LOG_PRINT("sparse_indices result (first 10 elements):\n"); + for (int64_t i = 0; i < size && i < 10; i++) { + LOG_PRINT(" [%ld] = %d\n", i, resultData[i]); + } + return ACL_SUCCESS; +} + +void CleanupResources(TensorResources &resources, void *workspaceAddr, aclrtStream stream, int32_t deviceId) +{ + if (resources.queryTensor) { + aclDestroyTensor(resources.queryTensor); + } + if (resources.keyTensor) { + aclDestroyTensor(resources.keyTensor); + } + if (resources.weightsTensor) { + aclDestroyTensor(resources.weightsTensor); + } + if (resources.qScaleTensor) { + aclDestroyTensor(resources.qScaleTensor); + } + if (resources.kScaleTensor) { + aclDestroyTensor(resources.kScaleTensor); + } + if (resources.metadataTensor) { + aclDestroyTensor(resources.metadataTensor); + } + if (resources.sparseIndicesTensor) { + aclDestroyTensor(resources.sparseIndicesTensor); + } + if (resources.sparseValuesTensor) { + aclDestroyTensor(resources.sparseValuesTensor); + } + + if (resources.queryDeviceAddr) { + aclrtFree(resources.queryDeviceAddr); + } + if (resources.keyDeviceAddr) { + aclrtFree(resources.keyDeviceAddr); + } + if (resources.weightsDeviceAddr) { + aclrtFree(resources.weightsDeviceAddr); + } + if (resources.qScaleDeviceAddr) { + aclrtFree(resources.qScaleDeviceAddr); + } + if (resources.kScaleDeviceAddr) { + aclrtFree(resources.kScaleDeviceAddr); + } + if (resources.metadataDeviceAddr) { + aclrtFree(resources.metadataDeviceAddr); + } + if (resources.sparseIndicesDeviceAddr) { + aclrtFree(resources.sparseIndicesDeviceAddr); + } + if (resources.sparseValuesDeviceAddr) { + aclrtFree(resources.sparseValuesDeviceAddr); + } + + if (workspaceAddr) { + aclrtFree(workspaceAddr); + } + if (stream) { + aclrtDestroyStream(stream); + } + aclrtResetDevice(deviceId); + aclFinalize(); +} + +} // namespace + +int main() +{ + if (op::GetCurrentPlatformInfo().GetCurNpuArch() != NpuArch::DAV_3510) { + return 0; + } + int32_t deviceId = 0; + aclrtStream stream = nullptr; + TensorResources resources = {}; + void *workspaceAddr = nullptr; + uint64_t workspaceSize = 0; + int64_t B = 2; + int64_t S1 = 4; + int64_t S2 = 8; + int64_t N1 = 64; + int64_t N2 = 1; + int64_t D = 128; + int64_t topk = 512; + std::vector sparseIndicesShape = {B, S1, N2, topk}; + int ret = ACL_SUCCESS; + + ret = Init(deviceId, &stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("Init acl failed. ERROR: %d\n", ret); + return ret; + } + + ret = InitializeTensors(resources); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("InitializeTensors failed. ERROR: %d\n", ret); + CleanupResources(resources, workspaceAddr, stream, deviceId); + return ret; + } + + ret = GenerateMetadata(resources, stream, B, S1, S2, N1, N2, D, topk, 1, 0, 1); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("GenerateMetadata failed. ERROR: %d\n", ret); + CleanupResources(resources, workspaceAddr, stream, deviceId); + return ret; + } + + ret = ExecuteQuantLightningIndexerV2(resources, stream, &workspaceAddr, &workspaceSize); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("ExecuteQuantLightningIndexerV2 failed. ERROR: %d\n", ret); + CleanupResources(resources, workspaceAddr, stream, deviceId); + return ret; + } + + ret = aclrtSynchronizeStream(stream); + if (!CHECK_RET(ret == ACL_SUCCESS)) { + LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); + CleanupResources(resources, workspaceAddr, stream, deviceId); + return ret; + } + + PrintOutResult(sparseIndicesShape, resources.sparseIndicesDeviceAddr); + + CleanupResources(resources, workspaceAddr, stream, deviceId); + return 0; +} diff --git a/csrc/attention/quant_lightning_indexer_v2/op_host/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2/op_host/CMakeLists.txt index a80dcbb66a8f..cb73bf477784 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_host/CMakeLists.txt +++ b/csrc/attention/quant_lightning_indexer_v2/op_host/CMakeLists.txt @@ -1,9 +1,19 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- add_op_to_compiled_list() -set(quant_lightning_indexer_v2_depends - "attention/lightning_indexer_v2" - CACHE STRING "Kernel source dependencies for quant_lightning_indexer_v2" FORCE) if (BUILD_OPEN_PROJECT) + # Only the Ascend 950 implementation requires LightningIndexerV2 headers. + if (ASCEND_COMPUTE_UNIT STREQUAL "ascend950") + set(quant_lightning_indexer_v2_depends attention/lightning_indexer_v2 CACHE INTERNAL "Dependencies for quant_lightning_indexer_v2") + endif() target_sources(op_host_aclnn PRIVATE quant_lightning_indexer_v2_def.cpp ) @@ -13,10 +23,11 @@ add_ops_compile_options( OP_NAME QuantLightningIndexerV2 OPTIONS --cce-auto-sync=off -Wno-deprecated-declarations + -Werror -mllvm -cce-vf-remove-membar=false -mllvm -cce-aicore-hoist-movemask=false ) if (NOT BUILD_OPS_RTY_KERNEL) add_modules_sources(OPTYPE quant_lightning_indexer_v2 ACLNNTYPE aclnn) -endif() +endif() \ No newline at end of file diff --git a/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_def.cpp b/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_def.cpp index 93d0fd57e7f4..c14c68734c83 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_def.cpp +++ b/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_def.cpp @@ -73,8 +73,14 @@ class QuantLightningIndexerV2 : public OpDef { .DataTypeList({ge::DT_INT32}) .FormatList({ge::FORMAT_ND}) .AutoContiguous(); + this->Input("candidate_topk_index") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); this->Output("sparse_indices").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND}); this->Output("sparse_values").ParamType(REQUIRED).DataTypeList({ge::DT_BF16}).FormatList({ge::FORMAT_ND}); + this->Output("candidate_topk_index_out").ParamType(OPTIONAL).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND}); this->Attr("topk").AttrType(REQUIRED).Int(2048); // 2048: 筛选前2048个作为输出index this->Attr("quant_mode").AttrType(REQUIRED).Int(1); // 1: per-token-head this->Attr("max_seqlen_q").AttrType(OPTIONAL).Int(-1); // -1: 默认值,表示任意可能长度 @@ -83,6 +89,11 @@ class QuantLightningIndexerV2 : public OpDef { this->Attr("mask_mode").AttrType(OPTIONAL).Int(0); // 0: 默认值,无mask this->Attr("cmp_ratio").AttrType(OPTIONAL).Int(1); this->Attr("return_value").AttrType(OPTIONAL).Int(0); // 0: 默认值 + this->Attr("candidate_mode").AttrType(OPTIONAL).Int(3); // 3: 默认关闭candidate + this->Attr("candidate_topk_blocks").AttrType(OPTIONAL).Int(2048); // 块级topk个数, 当前仅支持2048 + this->Attr("candidate_block_size").AttrType(OPTIONAL).Int(8); // 候选块大小(位置数) + this->Attr("key_stride0").AttrType(OPTIONAL).Int(0); // A11: key 第0维 stride (0=紧凑; aclnn 动态调用下 tiling 拿不到 tensor stride, 由 csrc 自动传入) + this->Attr("key_dequant_scale_stride0").AttrType(OPTIONAL).Int(0); // A11: k_scale 第0维 stride OpAICoreConfig aicore_config; aicore_config.DynamicCompileStaticFlag(true) .DynamicFormatFlag(true) diff --git a/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_infershape.cpp b/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_infershape.cpp index fac1a3f761da..60d46a5be420 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_infershape.cpp +++ b/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_infershape.cpp @@ -27,6 +27,11 @@ constexpr uint32_t ATTR_SPARSE_COUNT_INDEX = 0; constexpr uint32_t ATTR_QUERY_LAYOUT_INDEX = 3; constexpr uint32_t ATTR_KV_LAYOUT_INDEX = 4; constexpr uint32_t ATTR_RETURN_VALUE_INDEX = 7; +constexpr uint32_t ATTR_CANDIDATE_MODE_INDEX = 8; +constexpr uint32_t ATTR_CANDIDATE_TOPK_BLOCKS_INDEX = 9; +constexpr uint32_t CANDIDATE_MODE_SOURCE = 1; +constexpr uint32_t CANDIDATE_TOPK_INDEX_OUTPUT_INDEX = 2; +constexpr uint32_t CANDIDATE_TOPK_BLOCKS_FIX = 2048; constexpr uint32_t DIM_NUM_3 = 3; constexpr uint32_t DIM_NUM_4 = 4; @@ -83,6 +88,26 @@ static ge::graphStatus InferShapeQuantLightningIndexerV2(gert::InferShapeContext sparseValuesShape->SetDim(0, 0); } + // candidate_topk_index: 仅 candidate_mode=1(source) 时输出, BSND 布局 [B, S1, N2, candidate_topk_blocks] + gert::Shape *candidateTopkIndexShape = context->GetOutputShape(CANDIDATE_TOPK_INDEX_OUTPUT_INDEX); + OP_CHECK_NULL_WITH_CONTEXT(context, candidateTopkIndexShape); + const int32_t *candidate_mode = attrs->GetAttrPointer(ATTR_CANDIDATE_MODE_INDEX); + uint32_t candidateMode = (candidate_mode != nullptr) ? static_cast(*candidate_mode) : 3U; + if (candidateMode == CANDIDATE_MODE_SOURCE) { + OP_CHECK_IF(inputLayoutQueryPtrStr != "BSND", + OP_LOGE("QuantLightningIndexerV2", + "candidate_mode=1 only supports layout_q=BSND, but got %s.", + inputLayoutQueryPtrStr.c_str()), + return GRAPH_FAILED); + const int64_t *candidate_topk_blocks = attrs->GetAttrPointer(ATTR_CANDIDATE_TOPK_BLOCKS_INDEX); + int64_t candBlocks = (candidate_topk_blocks != nullptr) ? *candidate_topk_blocks : CANDIDATE_TOPK_BLOCKS_FIX; + *candidateTopkIndexShape = *sparseIndicesShape; + candidateTopkIndexShape->SetDim(sparseIndicesShape->GetDimNum() - 1, candBlocks); + } else { + candidateTopkIndexShape->SetDimNum(1); + candidateTopkIndexShape->SetDim(0, 0); + } + OP_LOGD(context->GetNodeName(), "QuantLightningIndexerV2 InferShape end."); return ge::GRAPH_SUCCESS; } @@ -97,6 +122,7 @@ static ge::graphStatus InferDataTypeQuantLightningIndexerV2(gert::InferDataTypeC // default index data type is int32 ge::DataType outputType = ge::DT_INT32; context->SetOutputDataType(0, outputType); + context->SetOutputDataType(CANDIDATE_TOPK_INDEX_OUTPUT_INDEX, outputType); OP_LOGD(context->GetNodeName(), "QuantLightningIndexerV2 InferDataType end."); return GRAPH_SUCCESS; } diff --git a/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_tiling.cpp b/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_tiling.cpp index 00d1c45ebd47..63cb943ffd52 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_tiling.cpp +++ b/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_tiling.cpp @@ -97,12 +97,6 @@ static std::string ToStringRaw(const gert::Shape &shape) return oss.str(); } -// Keep shape diagnostics consistent without an external opsbase formatting call. -static std::string ShapeToStringForLog(const gert::Shape &shape) -{ - return "[" + ToStringRaw(shape) + "]"; -} - // --------------------------QLIV2InfoParser类成员函数定义------------------------------------- ge::graphStatus QLIV2InfoParser::CheckRequiredInOutExistence() const { @@ -238,6 +232,8 @@ void QLIV2InfoParser::GetOptionalInputParaInfo() opParamInfo_.outputIdxOffset.desc = context_->GetOptionalInputDesc(OUTPUT_IDX_OFFSET_INDEX); opParamInfo_.metadata.tensor = context_->GetOptionalInputTensor(METADATA_INDEX); opParamInfo_.metadata.desc = context_->GetOptionalInputDesc(METADATA_INDEX); + opParamInfo_.candidateTopkIndex.tensor = context_->GetOptionalInputTensor(CANDIDATE_TOPK_INDEX_INPUT_INDEX); + opParamInfo_.candidateTopkIndex.desc = context_->GetOptionalInputDesc(CANDIDATE_TOPK_INDEX_INPUT_INDEX); } void QLIV2InfoParser::GetInputParaInfo() @@ -279,6 +275,15 @@ ge::graphStatus QLIV2InfoParser::GetAttrParaInfo() opParamInfo_.sparseMode = attrs->GetAttrPointer(ATTR_MASK_MODE_INDEX); opParamInfo_.cmpRatio = attrs->GetAttrPointer(ATTR_CMP_RATIO_INDEX); opParamInfo_.returnValue = attrs->GetAttrPointer(ATTR_RETURN_VALUE_INDEX); + opParamInfo_.candidateMode = attrs->GetAttrPointer(ATTR_CANDIDATE_MODE_INDEX); + opParamInfo_.candidateTopkBlocks = attrs->GetAttrPointer(ATTR_CANDIDATE_TOPK_BLOCKS_INDEX); + opParamInfo_.candidateBlockSize = attrs->GetAttrPointer(ATTR_CANDIDATE_BLOCK_SIZE_INDEX); + // A11: key 0 轴非连续 — tiling 侧 stride 来源优先级: + // 1) GetDynamicInputStride (仅 TensorV2/图模式携带非连续描述时有值) + // 2) 显式属性 key_stride0/key_dequant_scale_stride0 (aclnn 动态调用下 1) 恒为空, csrc 从 tensor.stride() 自动传入) + // 两路都空 = 紧凑存储 + opParamInfo_.keyStride0Attr = attrs->GetAttrPointer(ATTR_KEY_STRIDE0_INDEX); + opParamInfo_.keyDequantScaleStride0Attr = attrs->GetAttrPointer(ATTR_KEY_DEQUANT_SCALE_STRIDE0_INDEX); auto keyStrides = context_->GetDynamicInputStride(KEY_INDEX, 0); auto keyDequantScaleStrides = context_->GetDynamicInputStride(KEY_DEQUANT_SCALE_INDEX, 0); if (keyStrides != nullptr && keyStrides->GetDimNum() > 0) { @@ -412,6 +417,61 @@ ge::graphStatus QLIV2InfoParser::CheckAttrParaInfo() "Max_seqlen_q must >= -1"), return ge::GRAPH_FAILED); + // -------------------candidate (two-level topk) 校验------------------- + uint32_t candidateMode = (opParamInfo_.candidateMode != nullptr) ? + static_cast(*opParamInfo_.candidateMode) : CANDIDATE_MODE_OFF; + OP_CHECK_IF((candidateMode != CANDIDATE_MODE_SOURCE) && (candidateMode != CANDIDATE_MODE_CONSUMER) && + (candidateMode != CANDIDATE_MODE_OFF), + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "candidate_mode", + std::to_string(candidateMode), + "Candidate_mode only supports 1(source), 2(consumer) or 3(off)"), + return ge::GRAPH_FAILED); + if (candidateMode != CANDIDATE_MODE_OFF) { + // candidate 功能当前仅在 arch22 (910b/910_93) 实现 + OP_CHECK_IF(npuArch_ == NpuArch::DAV_3510, + OP_LOGE(opName_, "candidate_mode only supported on ascend910b/ascend910_93."), + return ge::GRAPH_FAILED); + // A12: layout_q 支持 BSND 与 TND (TND 需 cu_seqlens_q; layout_k 固定 PA_BBND — K 仅支持分页布局) + OP_CHECK_IF(std::string(opParamInfo_.layOutQuery) != "BSND" && std::string(opParamInfo_.layOutQuery) != "TND", + OP_LOGE(opName_, "candidate_mode only supports layout_q=BSND/TND, but got %s.", + layout_query.c_str()), + return ge::GRAPH_FAILED); + if (std::string(opParamInfo_.layOutQuery) == "TND") { + OP_CHECK_IF(std::string(opParamInfo_.layOutKey) != "PA_BBND", + OP_LOGE(opName_, "candidate_mode with layout_q=TND requires layout_k=PA_BBND, but got %s.", + layout_key.c_str()), + return ge::GRAPH_FAILED); + OP_CHECK_IF(opParamInfo_.cuSeqLensQ.tensor == nullptr, + OP_LOGE(opName_, "candidate_mode with layout_q=TND requires cu_seqlens_q."), + return ge::GRAPH_FAILED); + } + uint32_t candBlocks = (opParamInfo_.candidateTopkBlocks != nullptr) ? + static_cast(*opParamInfo_.candidateTopkBlocks) : + CANDIDATE_TOPK_BLOCKS_FIX; + // 放宽为 (0, 2048] 内 64 的倍数: 累加器/抽取/拷出逻辑均按 64 对齐设计 + OP_CHECK_IF(candBlocks == 0 || candBlocks > CANDIDATE_TOPK_BLOCKS_FIX || (candBlocks % 64) != 0, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "candidate_topk_blocks", + std::to_string(candBlocks), + "Candidate_topk_blocks must be a multiple of 64 in (0, 2048]"), + return ge::GRAPH_FAILED); + uint32_t candBlkSize = (opParamInfo_.candidateBlockSize != nullptr) ? + static_cast(*opParamInfo_.candidateBlockSize) : + CANDIDATE_BLOCK_SIZE_DEFAULT; + // 当前仅支持 8: BlockReduceMax 以 32B 块 (8 fp32) 为归约粒度 + OP_CHECK_IF(candBlkSize != CANDIDATE_BLOCK_SIZE_DEFAULT, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "candidate_block_size", + std::to_string(candBlkSize), + "Candidate_block_size only supports 8 currently"), + return ge::GRAPH_FAILED); + if (candidateMode == CANDIDATE_MODE_CONSUMER) { + OP_CHECK_IF(opParamInfo_.candidateTopkIndex.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "candidate_topk_index", + "candidate_topk_index input is required when " + "candidate_mode=2(consumer)"), + return ge::GRAPH_FAILED); + } + } + return ge::GRAPH_SUCCESS; } @@ -686,8 +746,8 @@ ge::graphStatus QLIV2InfoParser::GetAndCheckOptionalInput() opParamInfo_.cuSeqLensK.tensor->GetStorageShape().GetShapeSize(), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "cu_seqlens_q and cu_seqlens_k", - ShapeToStringForLog(opParamInfo_.cuSeqLensQ.tensor->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.cuSeqLensK.tensor->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.cuSeqLensQ.tensor->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.cuSeqLensK.tensor->GetStorageShape()), "When layout_q is TND and layout_k is TND, " "the shape of cu_seqlens_q must equal the shape of cu_seqlens_k"), return ge::GRAPH_FAILED); @@ -844,9 +904,9 @@ ge::graphStatus QLIV2InfoParser::GetGSize() { if (n1Size_ % n2Size_ != 0) { OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(opName_, "q and k", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.key.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.key.shape->GetStorageShape()), "The head num of q can not be a multiple of the head num of k"); return ge::GRAPH_FAILED; } @@ -856,13 +916,14 @@ ge::graphStatus QLIV2InfoParser::GetGSize() OP_CHECK_IF(gSize_ > G_SIZE_LIMIT, OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q and k", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.key.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.key.shape->GetStorageShape()), "The value of (the head num of q divided by the head num of k) must <= 64"), return ge::GRAPH_FAILED); } else { - OP_CHECK_IF(gSize_ != G_SIZE_LIMIT, - OP_LOGE(opName_, "N1 is %u, N2 is %u, N1 divided by N2 must equal 64.", n1Size_, n2Size_), + // 910b/910_93: 支持 gSize=64/32 (32 参照 v1 quant_lightning_indexer, mBaseSize=4*gSize 推导) + OP_CHECK_IF((gSize_ != G_SIZE_LIMIT) && (gSize_ != G_SIZE_LIMIT_32_950), + OP_LOGE(opName_, "N1 is %u, N2 is %u, N1 divided by N2 must equal 64 or 32.", n1Size_, n2Size_), return ge::GRAPH_FAILED); } @@ -906,8 +967,8 @@ ge::graphStatus QLIV2InfoParser::GetBatchSize() OP_CHECK_IF((cuSeqLensKSize - 1) != bSize_, OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "cu_seqlens_q and cu_seqlens_k", - ShapeToStringForLog(opParamInfo_.cuSeqLensK.tensor->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.cuSeqLensK.tensor->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.cuSeqLensK.tensor->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.cuSeqLensK.tensor->GetStorageShape()), "The batch sizes derived from cu_seqlens_q and cu_seqlens_k must be same"), return ge::GRAPH_FAILED); } @@ -947,6 +1008,12 @@ ge::graphStatus QLIV2InfoParser::GetS1Size() { if (qLayout_ == DataLayout::BSND) { s1Size_ = opParamInfo_.query.shape->GetStorageShape().GetDim(1); + } else if (qLayout_ == DataLayout::TND) { + // A12: TND 的 q 为 [T, G, D] 拼接, s1Size = 总行数 T。 + // 注意: TND 主路径的批前缀/行数均由 kernel 从 cu_seqlens_q(GM) 逐批读取, 不消费 s1Size; + // tiling 阶段禁止读 tensor 数据 (gert::Tensor data 未就绪, 读取直接段错误); + // consumer shape 校验式已用 TND 专分支 (query.shape[0] x N2 x candBlocks), 与 s1Size 解耦 + s1Size_ = opParamInfo_.query.shape->GetStorageShape().GetDim(0); } return ge::GRAPH_SUCCESS; } @@ -1043,9 +1110,9 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() opParamInfo_.blockTable.tensor->GetStorageShape().GetDim(0) != bSize_)), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "cu_seqlens_q, seqused_k and block_table", - ShapeToStringForLog(opParamInfo_.cuSeqLensQ.tensor->GetStorageShape()) + ", " + - ShapeToStringForLog(opParamInfo_.sequsedK.tensor->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.blockTable.tensor->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.cuSeqLensQ.tensor->GetStorageShape()) + ", " + + Ops::Base::ToString(opParamInfo_.sequsedK.tensor->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.blockTable.tensor->GetStorageShape()), "TND case, the dim 0 of cu_seqlens_q, seqused_k and block_table must be same"), return ge::GRAPH_FAILED); OP_CHECK_IF((kLayout_ == DataLayout::TND) && (opParamInfo_.cuSeqLensK.tensor->GetShapeSize() != bSize_ + 1), @@ -1061,9 +1128,9 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() (opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0) != qTsize), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q, w and sparse_indices", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + ", " + - ShapeToStringForLog(opParamInfo_.weights.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.attenOut.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + ", " + + Ops::Base::ToString(opParamInfo_.weights.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.attenOut.shape->GetStorageShape()), "TND case q, w, sparse_values dim 0 are " + std::to_string(qTsize) + ", " + std::to_string(opParamInfo_.weights.shape->GetStorageShape().GetDim(0)) + ", " + std::to_string(opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0)) + @@ -1074,8 +1141,8 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() OP_CHECK_IF((opParamInfo_.sparseValues.shape->GetStorageShape().GetDim(0) != qTsize), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q and sparse_values", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.sparseValues.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.sparseValues.shape->GetStorageShape()), "TND case q and sparse_values dim 0 are " + std::to_string(qTsize) + ", " + std::to_string(opParamInfo_.sparseValues.shape->GetStorageShape().GetDim(0)) + " respectively, they must be same"), @@ -1085,8 +1152,8 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() OP_CHECK_IF((opParamInfo_.outputIdxOffset.tensor->GetStorageShape().GetDim(0) != qTsize), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q and output_idx_offset", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.outputIdxOffset.tensor->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.outputIdxOffset.tensor->GetStorageShape()), "TND case q and output_idx_offset dim 0 are " + std::to_string(qTsize) + " and " + std::to_string(opParamInfo_.outputIdxOffset.tensor->GetStorageShape().GetDim(0)) + " respectively, they must be same"), @@ -1104,11 +1171,11 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() (opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0) != bSize_)), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q, w, seqused_k, block_table and sparse_indices", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + ", " + - ShapeToStringForLog(opParamInfo_.weights.shape->GetStorageShape()) + ", " + - ShapeToStringForLog(opParamInfo_.sequsedK.tensor->GetStorageShape()) + ", " + - ShapeToStringForLog(opParamInfo_.blockTable.tensor->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.attenOut.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + ", " + + Ops::Base::ToString(opParamInfo_.weights.shape->GetStorageShape()) + ", " + + Ops::Base::ToString(opParamInfo_.sequsedK.tensor->GetStorageShape()) + ", " + + Ops::Base::ToString(opParamInfo_.blockTable.tensor->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.attenOut.shape->GetStorageShape()), "BSND case q, w, seqused_k, block_table, sparse_indices dim 0 are " + std::to_string(bSize_) + ", " + std::to_string(opParamInfo_.weights.shape->GetStorageShape().GetDim(0)) + ", " + std::to_string(opParamInfo_.sequsedK.tensor->GetStorageShape().GetDim(0)) + ", " + @@ -1123,10 +1190,10 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() (opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0) != bSize_)), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q, w, seqused_k and sparse_indices", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + ", " + - ShapeToStringForLog(opParamInfo_.weights.shape->GetStorageShape()) + ", " + - ShapeToStringForLog(opParamInfo_.sequsedK.tensor->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.attenOut.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + ", " + + Ops::Base::ToString(opParamInfo_.weights.shape->GetStorageShape()) + ", " + + Ops::Base::ToString(opParamInfo_.sequsedK.tensor->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.attenOut.shape->GetStorageShape()), "BSND case q, w, seqused_k, sparse_indices dim 0 are " + std::to_string(bSize_) + ", " + std::to_string(opParamInfo_.weights.shape->GetStorageShape().GetDim(0)) + ", " + std::to_string(opParamInfo_.sequsedK.tensor->GetStorageShape().GetDim(0)) + ", " + @@ -1137,8 +1204,8 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() (opParamInfo_.sequsedQ.tensor != nullptr) && (opParamInfo_.sequsedQ.tensor->GetShapeSize() != bSize_), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q and seqused_q", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.sequsedQ.tensor->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.sequsedQ.tensor->GetStorageShape()), "BSND case q, seqused_q dim 0 are " + std::to_string(bSize_) + ", " + std::to_string(opParamInfo_.sequsedQ.tensor->GetStorageShape().GetDim(0)) + " respectively, they must be same"), @@ -1148,9 +1215,9 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() (opParamInfo_.attenOut.shape->GetStorageShape().GetDim(1) != s1Size_), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q, w and sparse_indices", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + ", " + - ShapeToStringForLog(opParamInfo_.weights.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.attenOut.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + ", " + + Ops::Base::ToString(opParamInfo_.weights.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.attenOut.shape->GetStorageShape()), "BSND case q, w and sparse_indices dim 1 are " + std::to_string(s1Size_) + ", " + std::to_string(opParamInfo_.weights.shape->GetStorageShape().GetDim(1)) + ", " + std::to_string(opParamInfo_.attenOut.shape->GetStorageShape().GetDim(1)) + @@ -1163,8 +1230,8 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() OP_CHECK_IF((opParamInfo_.weights.shape->GetStorageShape().GetDim(queryWeightsN1Dim) != n1Size_), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q and w", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.weights.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.weights.shape->GetStorageShape()), "BSND case the head num of q, w are " + std::to_string(n1Size_) + ", " + std::to_string(opParamInfo_.weights.shape->GetStorageShape().GetDim(queryWeightsN1Dim)) + " respectively, they must be same"), @@ -1174,9 +1241,9 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() ((kLayout_ != DataLayout::TND && opParamInfo_.key.shape->GetStorageShape().GetDim(DIM_IDX_THREE) != headDim_) || (kLayout_ == DataLayout::TND && opParamInfo_.key.shape->GetStorageShape().GetDim(DIM_IDX_TWO) != headDim_)), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(opName_, "q and k", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.key.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.key.shape->GetStorageShape()), "BSND case q, k last dim are " + std::to_string(headDim_) + ", " + std::to_string(opParamInfo_.key.shape->GetStorageShape().GetDim( (kLayout_ == DataLayout::TND) ? DIM_IDX_TWO : DIM_IDX_THREE)) + @@ -1186,8 +1253,8 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() OP_CHECK_IF((opParamInfo_.attenOut.shape->GetStorageShape().GetDim(outN2Dim) != n2Size_), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "k and sparse_indices", - ShapeToStringForLog(opParamInfo_.key.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.attenOut.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.key.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.attenOut.shape->GetStorageShape()), "BSND case the head num of k, sparse_indices are " + std::to_string(n2Size_) + ", " + std::to_string(opParamInfo_.attenOut.shape->GetStorageShape().GetDim(outN2Dim)) + " respectively, they must be same"), @@ -1196,8 +1263,8 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() OP_CHECK_IF((opParamInfo_.attenOut.shape->GetStorageShape().GetDim(outN2Dim + 1) != *opParamInfo_.sparseCount), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "sparse_count and sparse_indices", - ShapeToStringForLog(opParamInfo_.key.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.attenOut.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.key.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.attenOut.shape->GetStorageShape()), "BSND case sparse_count, sparse_indices last dim are " + std::to_string(*opParamInfo_.sparseCount) + ", " + std::to_string(opParamInfo_.attenOut.shape->GetStorageShape().GetDim(outN2Dim + 1)) + " respectively, they must be same"), @@ -1217,8 +1284,8 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() OP_CHECK_IF((opParamInfo_.sparseValues.shape->GetStorageShape().GetDim(outN2Dim) != n2Size_), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "k and sparse_values", - ShapeToStringForLog(opParamInfo_.key.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.sparseValues.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.key.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.sparseValues.shape->GetStorageShape()), "The head num of k and sparse_values must be same"), return ge::GRAPH_FAILED); OP_CHECK_IF( @@ -1226,7 +1293,7 @@ ge::graphStatus QLIV2InfoParser::ValidateInputShapesMatch() OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "topk and sparse_values", std::to_string(*opParamInfo_.sparseCount) + " and " + - ShapeToStringForLog(opParamInfo_.sparseValues.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.sparseValues.shape->GetStorageShape()), "The last dim of sparse_values must be same as topk"), return ge::GRAPH_FAILED); } @@ -1300,8 +1367,8 @@ ge::graphStatus QLIV2InfoParser::CheckScaleShape() dimValueQueryScale != dimValueQuery, OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q and q_descale", - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.query_dequant_scale.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.query_dequant_scale.shape->GetStorageShape()), "Q_descale's shape[" + std::to_string(i) + "] " + std::to_string(dimValueQueryScale) + " and q's shape[" + std::to_string(i) + "] " + std::to_string(dimValueQuery) + " are not same"), return ge::GRAPH_FAILED); @@ -1312,8 +1379,8 @@ ge::graphStatus QLIV2InfoParser::CheckScaleShape() OP_CHECK_IF(dimValueKeyScale != dimValueKey, OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "k and k_descale", - ShapeToStringForLog(opParamInfo_.key.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.key_dequant_scale.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.key.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.key_dequant_scale.shape->GetStorageShape()), "K_descale's shape[" + std::to_string(i) + "] " + std::to_string(dimValueKeyScale) + " and k's shape[" + std::to_string(i) + "] " + std::to_string(dimValueKey) + " are not the same"), @@ -1324,7 +1391,7 @@ ge::graphStatus QLIV2InfoParser::CheckScaleShape() (opParamInfo_.query_dequant_scale.shape->GetStorageShape().GetDim(qShapeDim - 1) != expectScaleD) || (opParamInfo_.query_dequant_scale.shape->GetStorageShape().GetDim(qShapeDim) != MX_E8M0_SCALE_PACK_NUM), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( - opName_, "q_descale", ShapeToStringForLog(opParamInfo_.query_dequant_scale.shape->GetStorageShape()), + opName_, "q_descale", Ops::Base::ToString(opParamInfo_.query_dequant_scale.shape->GetStorageShape()), "When quant_mode is " + std::to_string(*opParamInfo_.quantMode) + ", q_descale's last dims should be [" + std::to_string(expectScaleD) + ", " + std::to_string(MX_E8M0_SCALE_PACK_NUM) + "]"), @@ -1333,7 +1400,7 @@ ge::graphStatus QLIV2InfoParser::CheckScaleShape() (opParamInfo_.key_dequant_scale.shape->GetStorageShape().GetDim(kShapeDim - 1) != expectScaleD) || (opParamInfo_.key_dequant_scale.shape->GetStorageShape().GetDim(kShapeDim) != MX_E8M0_SCALE_PACK_NUM), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( - opName_, "k_descale", ShapeToStringForLog(opParamInfo_.key_dequant_scale.shape->GetStorageShape()), + opName_, "k_descale", Ops::Base::ToString(opParamInfo_.key_dequant_scale.shape->GetStorageShape()), "When quant_mode is " + std::to_string(*opParamInfo_.quantMode) + ", k_descale's last dims should be [" + std::to_string(expectScaleD) + ", " + std::to_string(MX_E8M0_SCALE_PACK_NUM) + "]"), @@ -1354,8 +1421,8 @@ ge::graphStatus QLIV2InfoParser::CheckScaleShape() OP_CHECK_IF(dimValueQueryScale != dimValueQuery, OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "q_descale and q", - ShapeToStringForLog(opParamInfo_.query_dequant_scale.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.query.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.query_dequant_scale.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.query.shape->GetStorageShape()), "q_descale's shape[" + std::to_string(i) + "] and q's shape[" + std::to_string(i) + "] are not the same"), return ge::GRAPH_FAILED); @@ -1367,8 +1434,8 @@ ge::graphStatus QLIV2InfoParser::CheckScaleShape() OP_CHECK_IF(dimValueKeyScale != dimValueKey, OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( opName_, "k_descale and k", - ShapeToStringForLog(opParamInfo_.key_dequant_scale.shape->GetStorageShape()) + " and " + - ShapeToStringForLog(opParamInfo_.key.shape->GetStorageShape()), + Ops::Base::ToString(opParamInfo_.key_dequant_scale.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.key.shape->GetStorageShape()), "k_descale's shape[" + std::to_string(i) + "] and k's shape[" + std::to_string(i) + "] are not the same"), return ge::GRAPH_FAILED); @@ -1385,8 +1452,10 @@ ge::graphStatus QLIV2InfoParser::CheckKeyContiguous() const { bool keyNonContiguous = false; bool scaleNonContiguous = false; - // PA_BBND: axis 0 may be non-contiguous (paged cache stride0); remaining axes must be contiguous. - // Non-PA_BBND: every axis must be contiguous. + // A5/A11 PA_BBND: 0轴允许非连续,从1轴开始检查;非PA_BBND: 从0轴开始检查 + // (A11 起不限 arch35 — arch22 kernel 已按 keyStride0/keyDequantScaleStride0 寻址, 紧凑值兜底) + // PA_BBND: axis 0 allows non-contiguous, check starts from axis 1 + // Non-PA_BBND: check starts from axis 0 size_t checkStartIdx = (kLayout_ == DataLayout::PA_BBND) ? 1 : 0; if (!keyStridesVec_.empty() && opParamInfo_.key.shape != nullptr) { auto &shape = opParamInfo_.key.shape->GetStorageShape(); @@ -1449,9 +1518,29 @@ ge::graphStatus QLIV2InfoParser::CheckKeyContiguous() const opName_, "k", "When layout_k is PA_BBND, key stride0 must be positive, but got " + std::to_string(keyStridesVec_[0])), - return ge::GRAPH_FAILED); + return ge::GRAPH_FAILED); } } + // A11 校验: 显式属性 stride0 不得小于紧凑值 (0 轴只允许 padding 型非连续) + if (opParamInfo_.keyStride0Attr != nullptr && *opParamInfo_.keyStride0Attr > 0) { + uint64_t compactKey0 = static_cast(blockSize_) * n2Size_ * headDim_; + OP_CHECK_IF(static_cast(*opParamInfo_.keyStride0Attr) < compactKey0, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "key_stride0", + "key_stride0 (" + std::to_string(*opParamInfo_.keyStride0Attr) + + ") is smaller than the compact value (" + std::to_string(compactKey0) + + "); only 0-axis padding (stride >= compact) is supported"), + return ge::GRAPH_FAILED); + } + if (opParamInfo_.keyDequantScaleStride0Attr != nullptr && *opParamInfo_.keyDequantScaleStride0Attr > 0) { + uint64_t compactScale0 = static_cast(blockSize_) * n2Size_; + OP_CHECK_IF(static_cast(*opParamInfo_.keyDequantScaleStride0Attr) < compactScale0, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "key_dequant_scale_stride0", + "key_dequant_scale_stride0 (" + std::to_string(*opParamInfo_.keyDequantScaleStride0Attr) + + ") is smaller than the compact value (" + std::to_string(compactScale0) + ")"), + return ge::GRAPH_FAILED); + } if (isMxQuantMode && !keyDequantScaleStridesVec_.empty()) { OP_CHECK_IF( keyDequantScaleStridesVec_[0] <= 0 || keyDequantScaleStridesVec_[0] % MX_E8M0_SCALE_PACK_NUM != 0, @@ -1499,6 +1588,14 @@ void QLIV2InfoParser::GenerateInfo(QLIV2TilingInfo &QLIV2Info) QLIV2Info.cmpRatio = *opParamInfo_.cmpRatio; QLIV2Info.returnValue = *opParamInfo_.returnValue; QLIV2Info.maxSeqlenQ = (opParamInfo_.maxSeqlenQ != nullptr) ? *opParamInfo_.maxSeqlenQ : -1; + QLIV2Info.candidateMode = (opParamInfo_.candidateMode != nullptr) ? + static_cast(*opParamInfo_.candidateMode) : CANDIDATE_MODE_OFF; + QLIV2Info.candidateTopkBlocks = (opParamInfo_.candidateTopkBlocks != nullptr) ? + static_cast(*opParamInfo_.candidateTopkBlocks) : + CANDIDATE_TOPK_BLOCKS_FIX; + QLIV2Info.candidateBlockSize = (opParamInfo_.candidateBlockSize != nullptr) ? + static_cast(*opParamInfo_.candidateBlockSize) : + CANDIDATE_BLOCK_SIZE_DEFAULT; QLIV2Info.keyStridesVec = keyStridesVec_; QLIV2Info.keyDequantScaleStridesVec = keyDequantScaleStridesVec_; @@ -1509,16 +1606,24 @@ void QLIV2InfoParser::GenerateInfo(QLIV2TilingInfo &QLIV2Info) keyStride0 /= MXFP4_PACK_NUM; } QLIV2Info.keyStride0 = keyStride0; + } else if (opParamInfo_.keyStride0Attr != nullptr && *opParamInfo_.keyStride0Attr > 0) { + // A11: aclnn 动态调用下 GetDynamicInputStride 恒空, 由显式属性兜底 (csrc 自动取 tensor.stride(0) 传入) + QLIV2Info.keyStride0 = static_cast(*opParamInfo_.keyStride0Attr); } else { - QLIV2Info.keyStride0 = 0; // 非PA无需使用stride + QLIV2Info.keyStride0 = 0; // 紧凑存储 } if (!keyDequantScaleStridesVec_.empty()) { QLIV2Info.keyDequantScaleStride0 = static_cast(keyDequantScaleStridesVec_[0]); + } else if (opParamInfo_.keyDequantScaleStride0Attr != nullptr && + *opParamInfo_.keyDequantScaleStride0Attr > 0) { + QLIV2Info.keyDequantScaleStride0 = static_cast(*opParamInfo_.keyDequantScaleStride0Attr); } else if ((*opParamInfo_.quantMode == QUANT_MODE_MXFP8) || (*opParamInfo_.quantMode == QUANT_MODE_MXFP4)) { QLIV2Info.keyDequantScaleStride0 = static_cast(blockSize_) * (headDim_ / MX_SCALE_GROUP_SIZE); } else { QLIV2Info.keyDequantScaleStride0 = 0; } + // A11 校验: 显式/描述的 stride0 不得小于紧凑值 (1 轴起必须连续, 由 CheckKeyContiguous 保证; + // 0 轴 stride < 块紧凑值意味着块内跨块, 不支持) — 见 CheckKeyContiguous 内 PA_BBND 段 QLIV2Info.inputQLayout = qLayout_; QLIV2Info.inputKLayout = kLayout_; @@ -1610,6 +1715,27 @@ ge::graphStatus QuantLightningIndexerV2Tiling::DoTiling(QLIV2TilingInfo *tilingI workSpaces[0] = workspaceSize; // -------------set tilingdata----------------- + // candidate (two-level topk) 输入校验: mode=2 时 candidate_topk_index 必须为 [B, S1, N2, candBlocks] int32 + if (tilingInfo->candidateMode == CANDIDATE_MODE_CONSUMER) { + OP_CHECK_IF(tilingInfo->opParamInfo.candidateTopkIndex.desc == nullptr || + tilingInfo->opParamInfo.candidateTopkIndex.desc->GetDataType() != ge::DT_INT32, + OP_LOGE("QuantLightningIndexerV2", "candidate_topk_index dtype only supports int32."), + return ge::GRAPH_FAILED); + int64_t expectSize = 0; + if (tilingInfo->inputQLayout == DataLayout::TND) { + // A12 修正: TND 输出布局为 [T, N2, K], T = query.shape[0], 非 B x s1Size (会双重计数) + expectSize = tilingInfo->opParamInfo.query.shape->GetStorageShape().GetDim(0) * + tilingInfo->n2Size * tilingInfo->candidateTopkBlocks; + } else { + expectSize = static_cast(tilingInfo->bSize) * tilingInfo->s1Size * tilingInfo->n2Size * + tilingInfo->candidateTopkBlocks; + } + int64_t actualSize = tilingInfo->opParamInfo.candidateTopkIndex.tensor->GetShapeSize(); + OP_CHECK_IF(actualSize != expectSize, + OP_LOGE("QuantLightningIndexerV2", + "candidate_topk_index shape size must be %ld, but got %ld.", expectSize, actualSize), + return ge::GRAPH_FAILED); + } tilingData_.set_bSize(tilingInfo->bSize); tilingData_.set_s2Size(tilingInfo->s2Size); tilingData_.set_s1Size(tilingInfo->s1Size); @@ -1625,6 +1751,10 @@ ge::graphStatus QuantLightningIndexerV2Tiling::DoTiling(QLIV2TilingInfo *tilingI tilingData_.set_keyDequantScaleStride0(tilingInfo->keyDequantScaleStride0); tilingData_.set_quantMode(*tilingInfo->opParamInfo.quantMode); tilingData_.set_usedCoreNum(blockDim); + // ---- candidate (two-level topk) ---- + tilingData_.set_candidateMode(tilingInfo->candidateMode); + tilingData_.set_candidateTopkBlocks(tilingInfo->candidateTopkBlocks); + tilingData_.set_candidateBlockSize(tilingInfo->candidateBlockSize); tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity()); context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize()); diff --git a/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_tiling.h b/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_tiling.h index c61002d3f780..8837cceb507a 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_tiling.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_host/quant_lightning_indexer_v2_tiling.h @@ -36,7 +36,11 @@ struct TilingOptionalParaInfo { const gert::Tensor *tensor; }; -enum class DataLayout : uint32_t { BSND = 0, TND = 1, PA_BBND = 2 }; +enum class DataLayout : uint32_t { + BSND = 0, + TND = 1, + PA_BBND = 2 +}; // ------------------算子原型索引常量定义---------------- // Inputs Index @@ -53,8 +57,10 @@ constexpr uint32_t CMP_RESIDUAL_K_INDEX = 9; constexpr uint32_t BLOCK_TABLE_INDEX = 10; constexpr uint32_t OUTPUT_IDX_OFFSET_INDEX = 11; constexpr uint32_t METADATA_INDEX = 12; +constexpr uint32_t CANDIDATE_TOPK_INDEX_INPUT_INDEX = 13; constexpr uint32_t SPARSE_INDICES_INDEX = 0; constexpr uint32_t SPARSE_VALUES_INDEX = 1; +constexpr uint32_t CANDIDATE_TOPK_INDEX_OUTPUT_INDEX = 2; // Attributes Index constexpr uint32_t ATTR_TOPK_INDEX = 0; constexpr uint32_t ATTR_QUANT_MODE_INDEX = 1; @@ -64,6 +70,11 @@ constexpr uint32_t ATTR_KEY_LAYOUT_INDEX = 4; constexpr uint32_t ATTR_MASK_MODE_INDEX = 5; constexpr uint32_t ATTR_CMP_RATIO_INDEX = 6; constexpr uint32_t ATTR_RETURN_VALUE_INDEX = 7; +constexpr uint32_t ATTR_CANDIDATE_MODE_INDEX = 8; +constexpr uint32_t ATTR_CANDIDATE_TOPK_BLOCKS_INDEX = 9; +constexpr uint32_t ATTR_CANDIDATE_BLOCK_SIZE_INDEX = 10; +constexpr uint32_t ATTR_KEY_STRIDE0_INDEX = 11; // A11: key 第0维 stride +constexpr uint32_t ATTR_KEY_DEQUANT_SCALE_STRIDE0_INDEX = 12; // A11: k_scale 第0维 stride // Dim Index constexpr uint32_t DIM_IDX_ZERO = 0; @@ -94,6 +105,14 @@ constexpr uint32_t MX_E8M0_SCALE_PACK_NUM = 2; // MX的E8M0 scale形状最后一 constexpr uint32_t MXFP4_PACK_NUM = 2; // 每个uint8承载2个FP4 E2M1逻辑元素 constexpr uint32_t MX_SCALE_GROUP_SIZE = 32; // MX量化每32个D维元素对应1个E8M0 scale +// ------------------candidate 两级TopK 常量------------------ +constexpr uint32_t CANDIDATE_MODE_SOURCE = 1; // is_candidate_source: 输出候选块索引 +constexpr uint32_t CANDIDATE_MODE_CONSUMER = 2; // use_candidate: 输入候选块索引, 候选内选topk +constexpr uint32_t CANDIDATE_MODE_OFF = 3; // 关闭candidate功能(默认) +constexpr uint32_t CANDIDATE_TOPK_BLOCKS_FIX = 2048; // O2决策: 当前仅支持2048 +constexpr uint32_t CANDIDATE_BLOCK_SIZE_DEFAULT = 8; +constexpr uint32_t CANDIDATE_BLOCK_SIZE_MAX = 64; + // -----------算子TilingData定义--------------- BEGIN_TILING_DATA_DEF(QLIV2TilingData) TILING_DATA_FIELD_DEF(uint32_t, bSize) @@ -112,6 +131,10 @@ TILING_DATA_FIELD_DEF(int32_t, maxSeqlenQ) TILING_DATA_FIELD_DEF(uint32_t, keyStride0) TILING_DATA_FIELD_DEF(uint32_t, keyDequantScaleStride0) TILING_DATA_FIELD_DEF(uint32_t, quantMode) +// ---- candidate (two-level topk) ---- +TILING_DATA_FIELD_DEF(uint32_t, candidateMode) +TILING_DATA_FIELD_DEF(uint32_t, candidateTopkBlocks) +TILING_DATA_FIELD_DEF(uint32_t, candidateBlockSize) END_TILING_DATA_DEF REGISTER_TILING_DATA_CLASS(QuantLightningIndexerV2, QLIV2TilingData) @@ -133,6 +156,7 @@ struct QLIV2ParaInfo { TilingOptionalParaInfo blockTable = {nullptr, nullptr}; TilingOptionalParaInfo outputIdxOffset = {nullptr, nullptr}; TilingOptionalParaInfo metadata = {nullptr, nullptr}; + TilingOptionalParaInfo candidateTopkIndex = {nullptr, nullptr}; TilingRequiredParaInfo attenOut = {nullptr, nullptr}; TilingRequiredParaInfo sparseValues = {nullptr, nullptr}; @@ -145,6 +169,11 @@ struct QLIV2ParaInfo { const int32_t *sparseCount = nullptr; const int32_t *cmpRatio = nullptr; const int32_t *returnValue = nullptr; + const int32_t *candidateMode = nullptr; + const int32_t *candidateTopkBlocks = nullptr; + const int32_t *candidateBlockSize = nullptr; + const int32_t *keyStride0Attr = nullptr; // A11: key 第0维 stride 显式属性 + const int32_t *keyDequantScaleStride0Attr = nullptr; // A11: k_scale 第0维 stride 显式属性 }; // -----------算子Tiling入参信息类--------------- @@ -178,6 +207,10 @@ class QLIV2TilingInfo { std::vector keyStridesVec; std::vector keyDequantScaleStridesVec; int32_t maxSeqlenQ = -1; + // candidate (two-level topk) + uint32_t candidateMode = CANDIDATE_MODE_OFF; + uint32_t candidateTopkBlocks = CANDIDATE_TOPK_BLOCKS_FIX; + uint32_t candidateBlockSize = CANDIDATE_BLOCK_SIZE_DEFAULT; // DType ge::DataType inputQType = ge::DT_FLOAT16; ge::DataType inputKType = ge::DT_FLOAT16; @@ -190,7 +223,9 @@ class QLIV2TilingInfo { // -----------算子Tiling入参信息解析及Check类--------------- class QLIV2InfoParser { public: - explicit QLIV2InfoParser(gert::TilingContext *context) : context_(context) {} + explicit QLIV2InfoParser(gert::TilingContext *context) + : context_(context) + {} ~QLIV2InfoParser() = default; ge::graphStatus CheckRequiredInOutExistence() const; @@ -266,7 +301,8 @@ class QLIV2InfoParser { // ---------------算子Tiling类--------------- class QuantLightningIndexerV2Tiling { public: - explicit QuantLightningIndexerV2Tiling(gert::TilingContext *context) : context_(context) {}; + explicit QuantLightningIndexerV2Tiling(gert::TilingContext *context) + : context_(context) {}; ge::graphStatus DoTiling(QLIV2TilingInfo *tilingInfo); private: diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_common_arch22.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_common_arch22.h index cf5ec8ef4509..e231adf937a4 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_common_arch22.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_common_arch22.h @@ -17,6 +17,11 @@ namespace QLIV2Common { +// candidate (two-level topk) 模式, 与 host 侧 tiling.h 定义保持一致 +inline constexpr uint32_t CANDIDATE_MODE_SOURCE = 1; // is_candidate_source: 输出候选块索引 +inline constexpr uint32_t CANDIDATE_MODE_CONSUMER = 2; // use_candidate: 输入候选块索引, 候选内选topk +inline constexpr uint32_t CANDIDATE_MODE_OFF = 3; // 关闭candidate功能(默认) + // 与tiling的layout保持一致 enum class LI_LAYOUT : uint32_t { BSND = 0, @@ -25,8 +30,8 @@ enum class LI_LAYOUT : uint32_t { }; template + const bool PAGE_ATTENTION = false, LI_LAYOUT Q_LAYOUT_T = LI_LAYOUT::BSND, + LI_LAYOUT K_LAYOUT_T = LI_LAYOUT::PA_BBND, typename... Args> struct QLIV2Type { using queryType = Q_T; using keyType = K_T; @@ -56,6 +61,8 @@ struct RunInfo { uint64_t tensorKeyScaleOffset; uint64_t tensorWeightsOffset; uint64_t indiceOutOffset; + uint64_t candidateOutOffset; + uint64_t outputIdxOffsetCoreOffset; // A15: output_idx_offset 的 batch 级前缀 (行 x kHeadNum 布局) bool isFirstS2InnerLoop; bool isLastS2InnerLoop; @@ -97,22 +104,27 @@ struct ConstInfo { uint64_t qHeadNum = 0ULL; uint64_t kHeadNum; uint64_t headDim; - uint64_t sparseCount; // topK选取大小 - uint64_t kSeqSize = 0ULL; // kv最大S长度 - uint64_t qSeqSize = 1ULL; // q最大S长度 - uint32_t kCacheBlockSize = 0; // PA场景的block size - uint32_t maxBlockNumPerBatch = 0; // PA场景的最大单batch block number - uint32_t keyStride0 = 0; // PA 0-axis stride of key (elements) - uint32_t keyDequantScaleStride0 = 0; // PA 0-axis stride of key scale (elements) - LI_LAYOUT outputLayout; // 输出的格式 + uint64_t sparseCount; // topK选取大小 + uint64_t kSeqSize = 0ULL; // kv最大S长度 + uint64_t qSeqSize = 1ULL; // q最大S长度 + uint32_t kCacheBlockSize = 0; // PA场景的block size + uint32_t maxBlockNumPerBatch = 0; // PA场景的最大单batch block number + uint32_t keyStride0 = 0; // key 第0维真实 stride (0=非PA/旧调用, 紧凑兜底; 对齐 arch35 A11) + uint32_t keyDequantScaleStride0 = 0; // k_scale 第0维真实 stride (同上) + LI_LAYOUT outputLayout; // 输出的格式 bool attenMaskFlag = false; - uint32_t cmpRatio = 1; // 压缩率 + uint32_t cmpRatio = 1; // 压缩率 + + uint32_t actualLenQDims = 0U; // query的actualSeqLength 的维度 + uint32_t actualLenDims = 0U; // KV 的actualSeqLength 的维度 + uint32_t cmpResiduaKLenDims = 0U; // cmpResidualK的维度 + bool isAccumSeqS1 = false; // 是否累加模式 + bool isAccumSeqS2 = false; // 是否累加模式 - uint32_t actualLenQDims = 0U; // query的actualSeqLength 的维度 - uint32_t actualLenDims = 0U; // KV 的actualSeqLength 的维度 - uint32_t cmpResiduaKLenDims = 0U; // cmpResidualK的维度 - bool isAccumSeqS1 = false; // 是否累加模式 - bool isAccumSeqS2 = false; // 是否累加模式 + // candidate (two-level topk) + uint32_t candidateMode = 3U; // 1=source 2=consumer 3=off + uint32_t candidateTopkBlocks = 2048U; + uint32_t candidateBlockSize = 8U; uint32_t s2Start = 0U; uint32_t s2End = 0U; @@ -124,23 +136,23 @@ struct ConstInfo { }; struct LdSplitCoreInfo { - bool isLdCoreEnable = false; // 当前核是否参与规约任务 - uint32_t saveWorkSpaceIdx = 0U; // 存放LD参数的地址 - uint32_t bn2Idx = 0U; // 归约任务 - uint32_t bIdx = 0U; - uint32_t n2Idx = 0U; - uint32_t mIdx = 0U; - uint32_t workspaceIdx = 0U; // 当前AIV核上规约任务的索引 - uint32_t workspaceNum = 0U; // 当前AIV核上规约任务的S2切分数量 - uint32_t mStart = 0U; - uint32_t mNum = 0U; - uint64_t indiceOutCoreOffset = 0U; // 最终输出索引搬出Topk的初始偏移地址 - }; + bool isLdCoreEnable = false; // 当前核是否参与规约任务 + uint32_t saveWorkSpaceIdx = 0U; // 存放LD参数的地址 + uint32_t bn2Idx = 0U; // 归约任务 + uint32_t bIdx = 0U; + uint32_t n2Idx = 0U; + uint32_t mIdx = 0U; + uint32_t workspaceIdx = 0U; // 当前AIV核上规约任务的索引 + uint32_t workspaceNum = 0U; // 当前AIV核上规约任务的S2切分数量 + uint32_t mStart = 0U; + uint32_t mNum = 0U; + uint64_t indiceOutCoreOffset = 0U; // 最终输出索引搬出Topk的初始偏移地址 +}; template __aicore__ inline T1 Align(T1 num, T2 rnd) { - return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd) * (rnd))); + return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd) * (rnd))); } template @@ -160,6 +172,6 @@ __aicore__ inline T CeilDiv(T num, T rnd) { return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd))); } -} // namespace QLIV2Common +} // namespace QLIV2Common -#endif // QUANT_LIGHTNING_INDEXER_V2_COMMON_H \ No newline at end of file +#endif // QUANT_LIGHTNING_INDEXER_V2_COMMON_H diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_kernel_arch22.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_kernel_arch22.h index 3d99a0432187..5a4a8b22c6ff 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_kernel_arch22.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_kernel_arch22.h @@ -43,31 +43,31 @@ struct TempLoopInfo { uint32_t bIdx = 0U; uint32_t n2Idx = 0U; uint32_t gS1Idx = 0U; - uint32_t gS1LoopEnd = 0U; // gS1方向循环的结束Idx - uint32_t s2LoopEnd = 0U; // S2方向循环的结束Idx - uint32_t actS1Size = 1ULL; // 当前Batch循环处理的S1轴的实际大小 + uint32_t gS1LoopEnd = 0U; // gS1方向循环的结束Idx + uint32_t s2LoopEnd = 0U; // S2方向循环的结束Idx + uint32_t actS1Size = 1ULL; // 当前Batch循环处理的S1轴的实际大小 uint32_t actS2Size = 0ULL; uint32_t actS2SizeOrig = 0ULL; bool curActSeqLenIsZero = false; - bool needDealActS1LessThanS1 = false; // S1的实际长度小于shape的S1长度时,是否需要清理输出 - bool isNeedLD = false; // 该基本块是否需要LD - uint32_t actMBaseSize = 0U; // m轴(gS1)方向实际大小 - uint32_t mBasicSizeTail = 0U; // gS1方向循环的尾基本块大小 - uint32_t s2BasicSizeTail = 0U; // S2方向循环的尾基本块大小 + bool needDealActS1LessThanS1 = false; // S1的实际长度小于shape的S1长度时,是否需要清理输出 + bool isNeedLD = false; // 该基本块是否需要LD + uint32_t actMBaseSize = 0U; // m轴(gS1)方向实际大小 + uint32_t mBasicSizeTail = 0U; // gS1方向循环的尾基本块大小 + uint32_t s2BasicSizeTail = 0U; // S2方向循环的尾基本块大小 uint32_t validS2Len = 0U; }; template class QLIV2Preload { public: - __aicore__ inline QLIV2Preload() {}; + __aicore__ inline QLIV2Preload(){}; __aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights, - __gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale, - __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensK, - __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sequsedK, + __gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale, __gm__ uint8_t *cuSeqlensQ, + __gm__ uint8_t *cuSeqlensK, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sequsedK, __gm__ uint8_t *cmpResidualK, __gm__ uint8_t *blockTable, __gm__ uint8_t *outputIdxOffset, __gm__ uint8_t *metadata, - __gm__ uint8_t *sparseIndices, __gm__ uint8_t *sparseValues, + __gm__ uint8_t *candidateTopkIndex, __gm__ uint8_t *sparseIndices, + __gm__ uint8_t *sparseValues, __gm__ uint8_t *candidateTopkIndexOut, __gm__ uint8_t *workspace, const QLIV2TilingData *__restrict tiling, TPipe *tPipe); __aicore__ inline void Process(); @@ -88,7 +88,8 @@ class QLIV2Preload { static constexpr uint32_t SYNC_C1_V1_FLAG = 4; static constexpr uint32_t SYNC_V1_C1_FLAG = 5; - static constexpr uint32_t M_BASE_SIZE = 256; + // 参照 v1: s1BaseSize 固定, mBaseSize = s1BaseSize * gSize (支持 gSize=32/64) + static constexpr uint32_t S1_BASE_SIZE = 4; static constexpr uint32_t S2_BASE_SIZE = 2048; static constexpr uint32_t HEAD_DIM = 128; static constexpr uint32_t K_HEAD_NUM = 1; @@ -110,6 +111,8 @@ class QLIV2Preload { uint64_t keyScaleCoreOffset = 0ULL; uint64_t weightsCoreOffset = 0ULL; uint64_t indiceOutCoreOffset = 0ULL; + uint64_t candidateOutCoreOffset = 0ULL; + uint64_t outputIdxOffsetCoreOffset = 0ULL; // A15 uint32_t coreZeroEnable = 1U; // ================================Global Buffer区================================= @@ -118,7 +121,10 @@ class QLIV2Preload { GlobalTensor weightsGm; GlobalTensor metadataGm; GlobalTensor indiceOutGm; + GlobalTensor candidateTopkIndexInGm; + GlobalTensor candidateTopkIndexOutGm; GlobalTensor blockTableGm; + GlobalTensor outputIdxOffsetGm; // A15: 每行输出索引偏移 GlobalTensor actualSeqLengthsGmQ; GlobalTensor actualSeqLengthsGm; @@ -147,8 +153,7 @@ class QLIV2Preload { // ================================Process functions================================ __aicore__ inline void ProcessMain(); __aicore__ inline void ProcessBaseBlock(uint32_t loop, uint64_t s2LoopIdx, - QLIV2Common::RunInfo - runInfo[LI_QUANT_PRELOAD_TASK_CACHE_SIZE]); + QLIV2Common::RunInfo runInfo[LI_QUANT_PRELOAD_TASK_CACHE_SIZE]); __aicore__ inline void ProcessInvalid(); __aicore__ inline void ProcessDecode(); // ================================Params Calc===================================== @@ -156,13 +161,12 @@ class QLIV2Preload { __aicore__ inline void GetBN2Idx(uint32_t bN2Idx); __aicore__ inline uint32_t GetActualSeqLen(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq, GlobalTensor &actualSeqLengthsGm, uint32_t defaultSeqLen); - __aicore__ inline uint32_t GetActualSeqLenKey(uint32_t bIdx, uint32_t actualLenDims, - uint32_t cmpResiduaKLenDims, bool isAccumSeq, - GlobalTensor &actualSeqLengthsGm, - GlobalTensor &cmpResidualKGm, - uint32_t defaultSeqLen, uint32_t cmpRatio); + __aicore__ inline uint32_t GetActualSeqLenKey(uint32_t bIdx, uint32_t actualLenDims, uint32_t cmpResiduaKLenDims, + bool isAccumSeq, GlobalTensor &actualSeqLengthsGm, + GlobalTensor &cmpResidualKGm, uint32_t defaultSeqLen, + uint32_t cmpRatio); __aicore__ inline void GetS1S2ActualSeqLen(uint32_t bIdx, uint32_t &actS1Size, uint32_t &actS2Size, - uint32_t &actS2SizeOrig); + uint32_t &actS2SizeOrig); __aicore__ inline void CalcS2LoopParams(uint32_t bN2LoopIdx, uint32_t gS1LoopIdx); __aicore__ inline void CalcRunInfo(uint32_t loop, uint32_t s2LoopIdx, QLIV2Common::RunInfo &runInfo); __aicore__ inline void DealActSeqLenIsZero(uint32_t bIdx, uint32_t n2Idx, uint32_t s1Start); @@ -178,9 +182,11 @@ __aicore__ inline void QLIV2Preload::InitTilingData(const QLIV2TilingDat constInfo.attenMaskFlag = (tilingData->sparseMode == 3); constInfo.kCacheBlockSize = tilingData->blockSize; constInfo.maxBlockNumPerBatch = tilingData->maxBlockNumPerBatch; + constInfo.keyStride0 = tilingData->keyStride0; // A11: key 0轴非连续 + constInfo.keyDequantScaleStride0 = tilingData->keyDequantScaleStride0; constInfo.sparseCount = tilingData->sparseCount; constInfo.cmpRatio = tilingData->cmpRatio; - constInfo.outputLayout = Q_LAYOUT_T; // 输出和输入形状一致 + constInfo.outputLayout = Q_LAYOUT_T; // 输出和输入形状一致 if constexpr (Q_LAYOUT_T == LI_LAYOUT::TND) { constInfo.isAccumSeqS1 = true; } @@ -190,18 +196,15 @@ __aicore__ inline void QLIV2Preload::InitTilingData(const QLIV2TilingDat constInfo.kHeadNum = K_HEAD_NUM; constInfo.headDim = HEAD_DIM; - constInfo.keyStride0 = tilingData->keyStride0; - if (constInfo.keyStride0 == 0) { - constInfo.keyStride0 = constInfo.kCacheBlockSize * constInfo.kHeadNum * constInfo.headDim; - } - constInfo.keyDequantScaleStride0 = tilingData->keyDequantScaleStride0; - if (constInfo.keyDequantScaleStride0 == 0) { - constInfo.keyDequantScaleStride0 = constInfo.kCacheBlockSize; - } - constInfo.mBaseSize = M_BASE_SIZE; + constInfo.s1BaseSize = S1_BASE_SIZE; constInfo.s2BaseSize = S2_BASE_SIZE; - constInfo.s1BaseSize = (constInfo.mBaseSize + constInfo.gSize - 1) / constInfo.gSize; + constInfo.mBaseSize = constInfo.s1BaseSize * constInfo.gSize; + + // candidate (two-level topk) + constInfo.candidateMode = tilingData->candidateMode; + constInfo.candidateTopkBlocks = tilingData->candidateTopkBlocks; + constInfo.candidateBlockSize = tilingData->candidateBlockSize; } template @@ -262,8 +265,8 @@ __aicore__ inline void QLIV2Preload::InitActualSeqLen(__gm__ uint8_t *cu template __aicore__ inline uint32_t QLIV2Preload::GetActualSeqLen(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq, - GlobalTensor &actualSeqLengthsGm, - uint32_t defaultSeqLen) + GlobalTensor &actualSeqLengthsGm, + uint32_t defaultSeqLen) { if (actualLenDims == 0) { return defaultSeqLen; @@ -283,7 +286,7 @@ __aicore__ inline uint32_t QLIV2Preload::GetActualSeqLenKey(uint32_t bId GlobalTensor &cmpResidualKGm, uint32_t defaultSeqLen, uint32_t cmpRatio) { - uint32_t cmpResidualK; // 当前bidx对应的cmpResidualK + uint32_t cmpResidualK; // 当前bidx对应的cmpResidualK if (cmpResiduaKLenDims == 0) { cmpResidualK = 0; } else { @@ -301,31 +304,29 @@ __aicore__ inline uint32_t QLIV2Preload::GetActualSeqLenKey(uint32_t bId template __aicore__ inline void QLIV2Preload::GetS1S2ActualSeqLen(uint32_t bIdx, uint32_t &actS1Size, - uint32_t &actS2Size, uint32_t &actS2SizeOrig) + uint32_t &actS2Size, uint32_t &actS2SizeOrig) { actS1Size = GetActualSeqLen(bIdx, constInfo.actualLenQDims, constInfo.isAccumSeqS1, actualSeqLengthsGmQ, constInfo.qSeqSize); - actS2SizeOrig = - GetActualSeqLenKey(bIdx, constInfo.actualLenDims, constInfo.cmpResiduaKLenDims, constInfo.isAccumSeqS2, - actualSeqLengthsGm, cmpResidualKGm, constInfo.kSeqSize, constInfo.cmpRatio); // 压缩前的actS2Size - actS2Size = actS2SizeOrig / constInfo.cmpRatio; // 真实使用的压缩后S2长度 + actS2SizeOrig = GetActualSeqLenKey(bIdx, constInfo.actualLenDims, constInfo.cmpResiduaKLenDims, + constInfo.isAccumSeqS2, actualSeqLengthsGm, cmpResidualKGm, constInfo.kSeqSize, + constInfo.cmpRatio); // 压缩前的actS2Size + actS2Size = actS2SizeOrig / constInfo.cmpRatio; // 真实使用的压缩后S2长度 } template __aicore__ inline uint32_t QLIV2Preload::GetS2BaseBlockNumOnMask(uint32_t s1gIdx, uint32_t actS1Size, - uint32_t actS2SizeOrig, uint32_t &validS2Len) + uint32_t actS2SizeOrig, uint32_t &validS2Len) { if (actS2SizeOrig / constInfo.cmpRatio == 0) { validS2Len = 0; return 0; } uint32_t s1Offset = constInfo.s1BaseSize * s1gIdx; - int32_t validS2LenBase = static_cast(actS2SizeOrig) - - static_cast(actS1Size); // 压缩前的validS2LenBase - validS2Len = - (static_cast(s1Offset) + validS2LenBase + - static_cast(constInfo.s1BaseSize)) / - static_cast(constInfo.cmpRatio); + int32_t validS2LenBase = + static_cast(actS2SizeOrig) - static_cast(actS1Size); // 压缩前的validS2LenBase + validS2Len = (static_cast(s1Offset) + validS2LenBase + static_cast(constInfo.s1BaseSize)) / + static_cast(constInfo.cmpRatio); validS2Len = Min(validS2Len, static_cast(actS2SizeOrig) / constInfo.cmpRatio); validS2Len = Max(validS2Len, 1); return (validS2Len + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize; @@ -367,7 +368,7 @@ __aicore__ void inline QLIV2Preload::SplitCore() } constInfo.bN2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, QLI_V2_BN2_END_INDEX, false)); constInfo.gS1End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, QLI_V2_M_END_INDEX, false)); - constInfo.s2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, QLI_V2_S2_END_INDEX, false)); + constInfo.s2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, QLI_V2_S2_END_INDEX, false)); // 如果0核都没有启动,说明所有核都没启动 coreZeroEnable = metadataGm.GetValue(GetAttrAbsIndex(0, QLI_V2_CORE_ENABLE_INDEX, false)); @@ -407,11 +408,11 @@ __aicore__ void inline QLIV2Preload::SplitCore() actualSeqQPrefixSum = (ldInfo.bIdx <= 0) ? 0 : ldInfo.bIdx * constInfo.qSeqSize; } // 搬出Topk的初始偏移地址 - ldInfo.indiceOutCoreOffset = actualSeqQPrefixSum * constInfo.kHeadNum * constInfo.sparseCount + - static_cast(ldInfo.n2Idx) * constInfo.sparseCount + - static_cast(ldInfo.mIdx) * constInfo.s1BaseSize * - constInfo.kHeadNum * constInfo.sparseCount; - } + ldInfo.indiceOutCoreOffset = + actualSeqQPrefixSum * constInfo.kHeadNum * constInfo.sparseCount + + static_cast(ldInfo.n2Idx) * constInfo.sparseCount + + static_cast(ldInfo.mIdx) * constInfo.s1BaseSize * constInfo.kHeadNum * constInfo.sparseCount; + } } template @@ -425,18 +426,17 @@ __aicore__ inline void QLIV2Preload::DealActSeqLenIsZero(uint32_t bIdx, for (uint32_t s1Idx = s1Start; s1Idx < s1Count; s1Idx++) { uint64_t indiceOutOffset = - (tBase + s1Idx) * constInfo.kHeadNum * constInfo.sparseCount + // T轴、s1轴偏移 - n2Idx * constInfo.sparseCount; // N2轴偏移 + (tBase + s1Idx) * constInfo.kHeadNum * constInfo.sparseCount + // T轴、s1轴偏移 + n2Idx * constInfo.sparseCount; // N2轴偏移 vectorService.CleanInvalidOutput(indiceOutOffset); } } else if (constInfo.outputLayout == LI_LAYOUT::BSND) { for (uint32_t s1Idx = s1Start; s1Idx < constInfo.qSeqSize; s1Idx++) { // B,S1,N2,K - uint64_t indiceOutOffset = static_cast(bIdx) * constInfo.qSeqSize * - constInfo.kHeadNum * constInfo.sparseCount + - static_cast(s1Idx) * constInfo.kHeadNum * - constInfo.sparseCount + // B轴、S1轴偏移 - static_cast(n2Idx) * constInfo.sparseCount; // N2轴偏移 + uint64_t indiceOutOffset = + static_cast(bIdx) * constInfo.qSeqSize * constInfo.kHeadNum * constInfo.sparseCount + + static_cast(s1Idx) * constInfo.kHeadNum * constInfo.sparseCount + // B轴、S1轴偏移 + static_cast(n2Idx) * constInfo.sparseCount; // N2轴偏移 vectorService.CleanInvalidOutput(indiceOutOffset); } } @@ -444,21 +444,19 @@ __aicore__ inline void QLIV2Preload::DealActSeqLenIsZero(uint32_t bIdx, } template -__aicore__ inline void QLIV2Preload::Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights, - __gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale, - __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensK, - __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sequsedK, - __gm__ uint8_t *cmpResidualK, __gm__ uint8_t *blockTable, - __gm__ uint8_t *outputIdxOffset, __gm__ uint8_t *metadata, - __gm__ uint8_t *sparseIndices, __gm__ uint8_t *sparseValues, - __gm__ uint8_t *workspace, const QLIV2TilingData *__restrict tiling, - TPipe *tPipe) +__aicore__ inline void QLIV2Preload::Init( + __gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights, __gm__ uint8_t *queryScale, + __gm__ uint8_t *keyScale, __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensK, __gm__ uint8_t *sequsedQ, + __gm__ uint8_t *sequsedK, __gm__ uint8_t *cmpResidualK, __gm__ uint8_t *blockTable, __gm__ uint8_t *outputIdxOffset, + __gm__ uint8_t *metadata, __gm__ uint8_t *candidateTopkIndex, __gm__ uint8_t *sparseIndices, + __gm__ uint8_t *sparseValues, __gm__ uint8_t *candidateTopkIndexOut, __gm__ uint8_t *workspace, + const QLIV2TilingData *__restrict tiling, TPipe *tPipe) { if ASCEND_IS_AIV { - tmpBlockIdx = GetBlockIdx(); // vec:0-47 + tmpBlockIdx = GetBlockIdx(); // vec:0-47 aiCoreIdx = tmpBlockIdx / 2; } else { - tmpBlockIdx = GetBlockIdx(); // cube:0-23 + tmpBlockIdx = GetBlockIdx(); // cube:0-23 aiCoreIdx = tmpBlockIdx; } @@ -477,12 +475,12 @@ __aicore__ inline void QLIV2Preload::Init(__gm__ uint8_t *query, __gm__ uint64_t offset = 0; // mm1开DoubleBuffer - GlobalTensor mm1ResGm; // 存放S + GlobalTensor mm1ResGm; // 存放S uint64_t singleCoreMm1ResSize = WS_DOUBLE * constInfo.s1BaseSize * constInfo.s2BaseSize * sizeof(MM1_OUT_T); mm1ResGm.SetGlobalBuffer((__gm__ MM1_OUT_T *)(workspace + aiCoreIdx * singleCoreMm1ResSize)); offset += GetBlockNum() * singleCoreMm1ResSize; - GlobalTensor weightWorkspaceGm; // v1阶段处理w*scale后的结果 + GlobalTensor weightWorkspaceGm; // v1阶段处理w*scale后的结果 uint64_t weightMemSize = BLOCK_CUBE * constInfo.mBaseSize * WS_DOUBLE * sizeof(half); weightWorkspaceGm.SetGlobalBuffer((__gm__ half *)(workspace + offset + aiCoreIdx * weightMemSize)); offset += GetBlockNum() * weightMemSize; @@ -490,7 +488,7 @@ __aicore__ inline void QLIV2Preload::Init(__gm__ uint8_t *query, __gm__ // ld流程需要ws大小: [aicnum, 2, s1BaseSize, topkOut_*2] // (aic, 8, 2, 2, 2048) // (aic, s1_cube, 头尾, idx/value, K) - GlobalTensor vec1ResGm; // 存放TopK计算中间结果 + GlobalTensor vec1ResGm; // 存放TopK计算中间结果 vec1ResGm.SetGlobalBuffer((__gm__ float *)(workspace + offset)); offset += GetBlockNum() * constInfo.s1BaseSize * WS_DOUBLE * WS_DOUBLE * BASE_TOPK * sizeof(float); @@ -503,7 +501,20 @@ __aicore__ inline void QLIV2Preload::Init(__gm__ uint8_t *query, __gm__ qScaleGm.SetGlobalBuffer((__gm__ half *)queryScale); kScaleGm.SetGlobalBuffer((__gm__ half *)keyScale); blockTableGm.SetGlobalBuffer((__gm__ int32_t *)blockTable); + if (constInfo.candidateMode == CANDIDATE_MODE_CONSUMER && candidateTopkIndex != nullptr) { + candidateTopkIndexInGm.SetGlobalBuffer((__gm__ int32_t *)candidateTopkIndex); + } + if (constInfo.candidateMode == CANDIDATE_MODE_SOURCE && candidateTopkIndexOut != nullptr) { + candidateTopkIndexOutGm.SetGlobalBuffer((__gm__ int32_t *)candidateTopkIndexOut); + } + if (outputIdxOffset != nullptr) { // A15: 使能 output_idx_offset (参照 arch35) + outputIdxOffsetGm.SetGlobalBuffer((__gm__ int32_t *)outputIdxOffset); + } + // A15: 有效标志须与 tensor 一并传给 vector (InitParams 已先值拷贝 constInfo, 此处再改标志不生效) + bool outputIdxOffsetValid = (outputIdxOffset != nullptr); vectorService.InitVecInputTensor(weightsGm, qScaleGm, kScaleGm, indiceOutGm, blockTableGm); + vectorService.InitVecCandidateTensor(candidateTopkIndexInGm, candidateTopkIndexOutGm, outputIdxOffsetGm, + outputIdxOffsetValid); vectorService.InitVecWorkspaceTensor(weightWorkspaceGm, mm1ResGm, vec1ResGm); } else { matmulService.InitParams(constInfo); @@ -552,8 +563,8 @@ __aicore__ inline void QLIV2Preload::CalcS2LoopParams(uint32_t bN2LoopId tempLoopInfo.isNeedLD = false; } tempLoopInfo.s2BasicSizeTail = tempLoopInfo.validS2Len % constInfo.s2BaseSize; - tempLoopInfo.s2BasicSizeTail = (tempLoopInfo.s2BasicSizeTail == 0) ? - constInfo.s2BaseSize : tempLoopInfo.s2BasicSizeTail; + tempLoopInfo.s2BasicSizeTail = + (tempLoopInfo.s2BasicSizeTail == 0) ? constInfo.s2BaseSize : tempLoopInfo.s2BasicSizeTail; } template @@ -571,8 +582,8 @@ __aicore__ inline void QLIV2Preload::CalcGS1LoopParams(uint32_t bN2LoopI (tempLoopInfo.mBasicSizeTail == 0) ? constInfo.mBaseSize : tempLoopInfo.mBasicSizeTail; uint32_t gS1SplitNum = (tempLoopInfo.actS1Size * constInfo.gSize + constInfo.mBaseSize - 1) / constInfo.mBaseSize; - tempLoopInfo.gS1LoopEnd = (bN2LoopIdx + 1 == constInfo.bN2End && constInfo.gS1End != 0) - ? constInfo.gS1End : gS1SplitNum; + tempLoopInfo.gS1LoopEnd = + (bN2LoopIdx + 1 == constInfo.bN2End && constInfo.gS1End != 0) ? constInfo.gS1End : gS1SplitNum; if constexpr (Q_LAYOUT_T == LI_LAYOUT::BSND) { if (tempLoopInfo.gS1LoopEnd == gS1SplitNum && constInfo.qSeqSize > tempLoopInfo.actS1Size) { tempLoopInfo.needDealActS1LessThanS1 = true; @@ -582,7 +593,7 @@ __aicore__ inline void QLIV2Preload::CalcGS1LoopParams(uint32_t bN2LoopI template __aicore__ inline void QLIV2Preload::CalcRunInfo(uint32_t loop, uint32_t s2LoopIdx, - QLIV2Common::RunInfo &runInfo) + QLIV2Common::RunInfo &runInfo) { runInfo.loop = loop; runInfo.bIdx = tempLoopInfo.bIdx; @@ -610,9 +621,8 @@ __aicore__ inline void QLIV2Preload::CalcRunInfo(uint32_t loop, uint32_t if (runInfo.s2Idx == s2SplitNum - 1) { runInfo.actualSingleProcessSInnerSize = tempLoopInfo.s2BasicSizeTail; } - runInfo.actualSingleProcessSInnerSizeAlign = - QLIV2Common::Align((uint32_t)runInfo.actualSingleProcessSInnerSize, - QLIV2Common::ConstInfo::BUFFER_SIZE_BYTE_32B); + runInfo.actualSingleProcessSInnerSizeAlign = QLIV2Common::Align((uint32_t)runInfo.actualSingleProcessSInnerSize, + QLIV2Common::ConstInfo::BUFFER_SIZE_BYTE_32B); runInfo.isFirstS2InnerLoop = s2LoopIdx == constInfo.s2Start; runInfo.isLastS2InnerLoop = (s2LoopIdx + 1 == tempLoopInfo.s2LoopEnd); @@ -621,7 +631,7 @@ __aicore__ inline void QLIV2Preload::CalcRunInfo(uint32_t loop, uint32_t uint64_t actualSeqQPrefixSum; if constexpr (Q_LAYOUT_T == LI_LAYOUT::TND) { actualSeqQPrefixSum = (runInfo.bIdx <= 0) ? 0 : actualSeqLengthsGmQ.GetValue(runInfo.bIdx); - } else { // BSND + } else { // BSND actualSeqQPrefixSum = (runInfo.bIdx <= 0) ? 0 : runInfo.bIdx * constInfo.qSeqSize; } uint64_t tndBIdxOffset = actualSeqQPrefixSum * constInfo.qHeadNum * constInfo.headDim; @@ -632,6 +642,10 @@ __aicore__ inline void QLIV2Preload::CalcRunInfo(uint32_t loop, uint32_t // B,S1,N2,k/T,N2,k indiceOutCoreOffset = actualSeqQPrefixSum * constInfo.kHeadNum * constInfo.sparseCount + runInfo.n2Idx * constInfo.sparseCount; + candidateOutCoreOffset = actualSeqQPrefixSum * constInfo.kHeadNum * constInfo.candidateTopkBlocks + + runInfo.n2Idx * constInfo.candidateTopkBlocks; + // A15: output_idx_offset 布局 [行, N2] (每行每 key 头一个 int32), batch 级前缀 + outputIdxOffsetCoreOffset = actualSeqQPrefixSum * constInfo.kHeadNum + runInfo.n2Idx; } uint64_t actualSeqKPrefixSum; if constexpr (K_LAYOUT_T == LI_LAYOUT::TND) { // T N2 D, cu_seqlens_k @@ -647,6 +661,8 @@ __aicore__ inline void QLIV2Preload::CalcRunInfo(uint32_t loop, uint32_t runInfo.tensorKeyScaleOffset = keyScaleCoreOffset; runInfo.tensorWeightsOffset = weightsCoreOffset; runInfo.indiceOutOffset = indiceOutCoreOffset; + runInfo.candidateOutOffset = candidateOutCoreOffset; + runInfo.outputIdxOffsetCoreOffset = outputIdxOffsetCoreOffset; } template @@ -666,7 +682,7 @@ template __aicore__ inline void QLIV2Preload::ProcessInvalid() { if ASCEND_IS_AIV { - uint32_t aivCoreNum = GetBlockNum() * 2; // 2 means c:v = 1:2 + uint32_t aivCoreNum = GetBlockNum() * 2; // 2 means c:v = 1:2 uint64_t totalOutputSize = constInfo.batchSize * constInfo.qSeqSize * constInfo.kHeadNum * constInfo.sparseCount; uint64_t singleCoreSize = @@ -727,7 +743,7 @@ __aicore__ inline void QLIV2Preload::ProcessMain() CrossCoreWaitFlag(constInfo.syncC1V1); vectorService.ProcessVec1(runInfo[1 - gloop % LI_QUANT_PRELOAD_TASK_CACHE_SIZE]); CrossCoreSetFlag( - constInfo.syncV1C1); // 反向同步 1 + constInfo.syncV1C1); // 反向同步 1 } } continue; @@ -735,10 +751,9 @@ __aicore__ inline void QLIV2Preload::ProcessMain() for (uint32_t gS1LoopIdx = constInfo.gS1Start; gS1LoopIdx < tempLoopInfo.gS1LoopEnd; gS1LoopIdx++) { CalcS2LoopParams(bN2LoopIdx, gS1LoopIdx); bool isEnd = (bN2LoopIdx + 1 == constInfo.bN2End) && (gS1LoopIdx + 1 == tempLoopInfo.gS1LoopEnd); - uint32_t extraLoop = isEnd ? LI_QUANT_PRELOAD_TASK_CACHE_SIZE - 1 : 0; // 只preload一轮 + uint32_t extraLoop = isEnd ? LI_QUANT_PRELOAD_TASK_CACHE_SIZE - 1 : 0; // 只preload一轮 - for (uint32_t s2LoopIdx = constInfo.s2Start; - s2LoopIdx < (tempLoopInfo.s2LoopEnd + extraLoop); + for (uint32_t s2LoopIdx = constInfo.s2Start; s2LoopIdx < (tempLoopInfo.s2LoopEnd + extraLoop); s2LoopIdx++) { ProcessBaseBlock(gloop, s2LoopIdx, runInfo); ++gloop; @@ -764,8 +779,7 @@ __aicore__ inline void QLIV2Preload::ProcessMain() template __aicore__ inline void QLIV2Preload::ProcessBaseBlock( - uint32_t loop, uint64_t s2LoopIdx, - QLIV2Common::RunInfo runInfo[LI_QUANT_PRELOAD_TASK_CACHE_SIZE]) + uint32_t loop, uint64_t s2LoopIdx, QLIV2Common::RunInfo runInfo[LI_QUANT_PRELOAD_TASK_CACHE_SIZE]) { int32_t curTaskId = loop % LI_QUANT_PRELOAD_TASK_CACHE_SIZE; QLIV2Common::RunInfo &curRunInfo = runInfo[curTaskId]; @@ -778,17 +792,16 @@ __aicore__ inline void QLIV2Preload::ProcessBaseBlock( if (curRunInfo.isFirstS2InnerLoop) { CrossCoreWaitFlag(constInfo.syncV0C1); } - CrossCoreWaitFlag(constInfo.syncV1C1); // 反向同步 1 + CrossCoreWaitFlag(constInfo.syncV1C1); // 反向同步 1 matmulService.ComputeMm1(curRunInfo); CrossCoreSetFlag(constInfo.syncC1V1); if (curRunInfo.isLastS2InnerLoop) { // 反向同步 0 - CrossCoreSetFlag(constInfo.syncC1V0); + CrossCoreSetFlag(constInfo.syncC1V0); } } else { if (curRunInfo.isFirstS2InnerLoop) { - CrossCoreWaitFlag(constInfo.syncC1V0); // 反向同步 0 + CrossCoreWaitFlag(constInfo.syncC1V0); // 反向同步 0 vectorService.ProcessVec0(curRunInfo); CrossCoreSetFlag(constInfo.syncV0C1); } @@ -799,7 +812,7 @@ __aicore__ inline void QLIV2Preload::ProcessBaseBlock( if ASCEND_IS_AIV { CrossCoreWaitFlag(constInfo.syncC1V1); vectorService.ProcessVec1(lastRunInfo); - CrossCoreSetFlag(constInfo.syncV1C1); // 反向同步 1 + CrossCoreSetFlag(constInfo.syncV1C1); // 反向同步 1 } lastRunInfo.isValid = false; } @@ -818,5 +831,5 @@ __aicore__ inline void QLIV2Preload::ProcessDecode() } } -} // namespace QLIV2Kernel -#endif // QUANT_LIGHTNING_INDEXER_V2_KERNEL_H +} // namespace QLIV2Kernel +#endif // QUANT_LIGHTNING_INDEXER_V2_KERNEL_H diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_service_cube_arch22.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_service_cube_arch22.h index 968e6fe84b54..a810d738190d 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_service_cube_arch22.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_service_cube_arch22.h @@ -37,7 +37,7 @@ class QLIV2Matmul { using Q_T = typename QLIV2T::queryType; using K_T = typename QLIV2T::keyType; - __aicore__ inline QLIV2Matmul() {}; + __aicore__ inline QLIV2Matmul(){}; __aicore__ inline void InitBuffers(TPipe *pipe); __aicore__ inline void InitMm1GlobalTensor(const GlobalTensor &blkTableGm, const GlobalTensor &keyGm, const GlobalTensor &queryGm, const GlobalTensor &mm1ResGm, @@ -47,12 +47,12 @@ class QLIV2Matmul { __aicore__ inline void FreeEventID(); __aicore__ inline void ComputeMm1(const QLIV2Common::RunInfo &runInfo); - static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding; + static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding; static constexpr uint64_t DOUBLE_BUF_NUM = 2; static constexpr uint64_t L0AB_BUF_NUM = 4; static constexpr uint32_t KEY_MTE1_MTE2_EVENT = EVENT_ID2; - static constexpr uint32_t QW_MTE1_MTE2_EVENT = EVENT_ID5; // KEY_MTE1_MTE2_EVENT + DOUBLE_BUF_NUM; + static constexpr uint32_t QW_MTE1_MTE2_EVENT = EVENT_ID5; // KEY_MTE1_MTE2_EVENT + DOUBLE_BUF_NUM; static constexpr uint32_t M_MTE1_EVENT = EVENT_ID3; static constexpr uint32_t M_FIX_EVENT = EVENT_ID0; static constexpr uint32_t FIX_M_EVENT = EVENT_ID2; @@ -83,7 +83,7 @@ class QLIV2Matmul { __aicore__ inline void LoadQueryToL0a(uint64_t s1gL1Offset, uint64_t s1gL1RealSize, uint64_t s1gL0RealSize); __aicore__ inline void QueryNd2Nz(uint64_t s1gL1RealSize, const QLIV2Common::RunInfo &runInfo); __aicore__ inline void KeyNd2NzForPA(uint64_t s2L1RealSize, uint64_t s2GmOffset, - const QLIV2Common::RunInfo &runInfo); + const QLIV2Common::RunInfo &runInfo); __aicore__ inline void KeyNd2Nz(uint64_t s2L1RealSize, const MmInfo &mmInfo, const QLIV2Common::RunInfo &runInfo); __aicore__ inline void FixpSToL1(uint64_t s1gL0RealSize, uint64_t s2L0RealSize); __aicore__ inline void LoadSToL0b(uint64_t s1gL1RealSize, uint64_t s2L0RealSize, uint64_t sL1BufIdx, @@ -163,10 +163,10 @@ __aicore__ inline void QLIV2Matmul::InitBuffers(TPipe *pipe) template __aicore__ inline void QLIV2Matmul::InitMm1GlobalTensor(const GlobalTensor &blkTableGm, - const GlobalTensor &keyGm, - const GlobalTensor &queryGm, - const GlobalTensor &mm1ResGm, - const GlobalTensor &weightWorkspaceGm) + const GlobalTensor &keyGm, + const GlobalTensor &queryGm, + const GlobalTensor &mm1ResGm, + const GlobalTensor &weightWorkspaceGm) { blkTableGm_ = blkTableGm; keyGm_ = keyGm; @@ -177,7 +177,7 @@ __aicore__ inline void QLIV2Matmul::InitMm1GlobalTensor(const GlobalTens template __aicore__ inline void QLIV2Matmul::ProcessWs(uint64_t s1gL0RealSize, uint64_t s1gL1Offset, uint64_t sL1BufIdx, - const MmInfo &mmInfo, const QLIV2Common::RunInfo &runInfo) + const MmInfo &mmInfo, const QLIV2Common::RunInfo &runInfo) { WaitFlag(FIX_M_EVENT + l0cBufIdx_ % DOUBLE_BUF_NUM); for (int64_t s1gOffset = 0; s1gOffset < s1gL0RealSize; s1gOffset += constInfo_.gSize) { @@ -199,8 +199,8 @@ __aicore__ inline void QLIV2Matmul::ProcessWs(uint64_t s1gL0RealSize, ui template __aicore__ inline void QLIV2Matmul::ProcessQk(uint64_t s1gL0RealSize, uint64_t s1gL1Offset, - uint64_t s1L0LoopCnt, - const MmInfo &mmInfo, const QLIV2Common::RunInfo &runInfo) + uint64_t s1L0LoopCnt, const MmInfo &mmInfo, + const QLIV2Common::RunInfo &runInfo) { if (mmInfo.s1gL0LoopId == 0) { WaitFlag(KEY_MTE1_MTE2_EVENT + keyL1BufIdx_ % DOUBLE_BUF_NUM); @@ -235,16 +235,16 @@ __aicore__ inline void QLIV2Matmul::ProcessQk(uint64_t s1gL0RealSize, ui template __aicore__ inline void QLIV2Matmul::CalcMmInfo(MmInfo &mmInfo, uint64_t loopIdx, uint64_t s1L0LoopCnt, - const MmInfo &lastMmInfo, const QLIV2Common::RunInfo &runInfo) + const MmInfo &lastMmInfo, const QLIV2Common::RunInfo &runInfo) { mmInfo.s2L0LoopId = loopIdx / s1L0LoopCnt; mmInfo.s1gL0LoopId = loopIdx % s1L0LoopCnt; if (mmInfo.s1gL0LoopId == 0) { mmInfo.s2GmOffset = mmInfo.s2L0LoopId * S2_BASIC_BLOCK_L0; - mmInfo.s2L0RealSize = mmInfo.s2GmOffset + S2_BASIC_BLOCK_L0 > runInfo.actualSingleProcessSInnerSize - ? runInfo.actualSingleProcessSInnerSize - mmInfo.s2GmOffset - : S2_BASIC_BLOCK_L0; + mmInfo.s2L0RealSize = mmInfo.s2GmOffset + S2_BASIC_BLOCK_L0 > runInfo.actualSingleProcessSInnerSize ? + runInfo.actualSingleProcessSInnerSize - mmInfo.s2GmOffset : + S2_BASIC_BLOCK_L0; } else { mmInfo.s2L0RealSize = lastMmInfo.s2L0RealSize; } @@ -255,12 +255,12 @@ __aicore__ inline void QLIV2Matmul::ComputeMm1(const QLIV2Common::RunInf { if (runInfo.isFirstS2InnerLoop) { WaitFlag(QW_MTE1_MTE2_EVENT + qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM); - QueryNd2Nz(runInfo.actMBaseSize, runInfo); // 256 * 128 // L1BasicBlock + QueryNd2Nz(runInfo.actMBaseSize, runInfo); // 256 * 128 // L1BasicBlock WeightDmaCopy(runInfo.actMBaseSize, runInfo); } int64_t loopIdx = 0; - int64_t s2L0LoopCnt = CeilDiv(runInfo.actualSingleProcessSInnerSize, S2_BASIC_BLOCK_L0); // 2048取128 - int64_t s1L0LoopCnt = CeilDiv(runInfo.actMBaseSize, S1G_BASIC_BLOCK_L0); // 256取128 + int64_t s2L0LoopCnt = CeilDiv(runInfo.actualSingleProcessSInnerSize, S2_BASIC_BLOCK_L0); // 2048取128 + int64_t s1L0LoopCnt = CeilDiv(runInfo.actMBaseSize, S1G_BASIC_BLOCK_L0); // 256取128 int64_t s1gL1Offset[2] = {0, static_cast(S1G_BASIC_BLOCK_L0)}; int64_t s1gL0RealSize[2] = {s1L0LoopCnt > 1 ? static_cast(S1G_BASIC_BLOCK_L0) : runInfo.actMBaseSize, runInfo.actMBaseSize - s1gL1Offset[1]}; @@ -268,8 +268,7 @@ __aicore__ inline void QLIV2Matmul::ComputeMm1(const QLIV2Common::RunInf CalcMmInfo(mmInfo[loopIdx & 1], loopIdx, s1L0LoopCnt, mmInfo[(loopIdx + 1) & 1], runInfo); ProcessQk(s1gL0RealSize[mmInfo[loopIdx & 1].s1gL0LoopId % s1L0LoopCnt], - s1gL1Offset[mmInfo[loopIdx & 1].s1gL0LoopId % s1L0LoopCnt], s1L0LoopCnt, mmInfo[loopIdx & 1], - runInfo); + s1gL1Offset[mmInfo[loopIdx & 1].s1gL0LoopId % s1L0LoopCnt], s1L0LoopCnt, mmInfo[loopIdx & 1], runInfo); SetFlag(FIX_MTE1_EVENT + sL1BufIdx_ % DOUBLE_BUF_NUM); sL1BufIdx_++; @@ -288,8 +287,8 @@ __aicore__ inline void QLIV2Matmul::ComputeMm1(const QLIV2Common::RunInf WaitFlag(FIX_MTE1_EVENT + sL1BufIdx_ % DOUBLE_BUF_NUM); ProcessWs(s1gL0RealSize[mmInfo[(loopIdx + 1) & 1].s1gL0LoopId % s1L0LoopCnt], - s1gL1Offset[mmInfo[(loopIdx + 1) & 1].s1gL0LoopId % s1L0LoopCnt], sL1BufIdx_, - mmInfo[(loopIdx + 1) & 1], runInfo); + s1gL1Offset[mmInfo[(loopIdx + 1) & 1].s1gL0LoopId % s1L0LoopCnt], sL1BufIdx_, + mmInfo[(loopIdx + 1) & 1], runInfo); loopIdx++; } @@ -308,24 +307,30 @@ __aicore__ inline void QLIV2Matmul::ComputeMm1(const QLIV2Common::RunInf // blkNum, blkSize, N2, D template __aicore__ inline void QLIV2Matmul::KeyNd2NzForPA(uint64_t s2L1RealSize, uint64_t s2GmOffset, - const QLIV2Common::RunInfo &runInfo) + const QLIV2Common::RunInfo &runInfo) { uint64_t s2L1Offset = 0; + // A11: key 0 轴非连续寻址 — 块基址 = 表值 x keyStride0 (host 下发的真实 stride)。 + // keyStride0 == 0 (旧调用/非 PA 紧凑) 时兜底原紧凑公式, 现网行为 bit 级不变 (R8) + uint64_t blkStride = constInfo_.keyStride0 != 0 ? + static_cast(constInfo_.keyStride0) : + static_cast(constInfo_.kCacheBlockSize) * constInfo_.kHeadNum * + constInfo_.headDim; while (s2L1Offset < s2L1RealSize) { uint64_t s2BlkId = (s2L1Offset + s2GmOffset) / constInfo_.kCacheBlockSize; uint64_t s2BlkOffset = (s2L1Offset + s2GmOffset) % constInfo_.kCacheBlockSize; uint64_t keyGmOffset = blkTableGm_.GetValue(runInfo.bIdx * constInfo_.maxBlockNumPerBatch + s2BlkId) * - constInfo_.keyStride0 + + blkStride + s2BlkOffset * constInfo_.headDim; uint64_t s2Mte2Size = s2L1RealSize - s2L1Offset; - s2Mte2Size = s2BlkOffset + s2Mte2Size >= constInfo_.kCacheBlockSize ? constInfo_.kCacheBlockSize - s2BlkOffset - : s2Mte2Size; + s2Mte2Size = s2BlkOffset + s2Mte2Size >= constInfo_.kCacheBlockSize ? constInfo_.kCacheBlockSize - s2BlkOffset : + s2Mte2Size; Nd2NzParams nd2nzPara; nd2nzPara.ndNum = 1; - nd2nzPara.nValue = s2Mte2Size; // 行数 + nd2nzPara.nValue = s2Mte2Size; // 行数 nd2nzPara.dValue = constInfo_.headDim; nd2nzPara.srcDValue = constInfo_.headDim; - nd2nzPara.dstNzC0Stride = CeilAlign(s2L1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block + nd2nzPara.dstNzC0Stride = CeilAlign(s2L1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block nd2nzPara.dstNzNStride = 1; nd2nzPara.srcNdMatrixStride = 0; nd2nzPara.dstNzMatrixStride = 0; @@ -338,7 +343,7 @@ __aicore__ inline void QLIV2Matmul::KeyNd2NzForPA(uint64_t s2L1RealSize, template __aicore__ inline void QLIV2Matmul::KeyNd2Nz(uint64_t s2L1RealSize, const MmInfo &mmInfo, - const QLIV2Common::RunInfo &runInfo) + const QLIV2Common::RunInfo &runInfo) { uint64_t dStride = constInfo_.headDim; if constexpr (K_LAYOUT_T == LI_LAYOUT::BSND || K_LAYOUT_T == LI_LAYOUT::TND) { @@ -346,10 +351,10 @@ __aicore__ inline void QLIV2Matmul::KeyNd2Nz(uint64_t s2L1RealSize, cons } Nd2NzParams nd2nzPara; nd2nzPara.ndNum = 1; - nd2nzPara.nValue = s2L1RealSize; // 行数 + nd2nzPara.nValue = s2L1RealSize; // 行数 nd2nzPara.dValue = constInfo_.headDim; nd2nzPara.srcDValue = dStride; - nd2nzPara.dstNzC0Stride = CeilAlign(s2L1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block + nd2nzPara.dstNzC0Stride = CeilAlign(s2L1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block nd2nzPara.dstNzNStride = 1; nd2nzPara.srcNdMatrixStride = 0; nd2nzPara.dstNzMatrixStride = 0; @@ -377,10 +382,10 @@ __aicore__ inline void QLIV2Matmul::QueryNd2Nz(uint64_t s1gL1RealSize, c { Nd2NzParams nd2nzPara; nd2nzPara.ndNum = 1; - nd2nzPara.nValue = s1gL1RealSize; // 行数 + nd2nzPara.nValue = s1gL1RealSize; // 行数 nd2nzPara.dValue = constInfo_.headDim; nd2nzPara.srcDValue = constInfo_.headDim; - nd2nzPara.dstNzC0Stride = CeilAlign(s1gL1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block + nd2nzPara.dstNzC0Stride = CeilAlign(s1gL1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block nd2nzPara.dstNzNStride = 1; nd2nzPara.srcNdMatrixStride = 0; nd2nzPara.dstNzMatrixStride = 0; @@ -392,22 +397,22 @@ __aicore__ inline void QLIV2Matmul::QueryNd2Nz(uint64_t s1gL1RealSize, c // s1g, d template __aicore__ inline void QLIV2Matmul::LoadQueryToL0a(uint64_t s1gL1Offset, uint64_t s1gL1RealSize, - uint64_t s1gL0RealSize) + uint64_t s1gL0RealSize) { LoadData3DParamsV2 loadData3DParams; // SetFmatrixParams - loadData3DParams.l1H = CeilDiv(s1gL1RealSize, BLOCK_CUBE); // Hin=M1=8 - loadData3DParams.l1W = BLOCK_CUBE; // Win=M0 - loadData3DParams.channelSize = constInfo_.headDim; // Cin=K + loadData3DParams.l1H = CeilDiv(s1gL1RealSize, BLOCK_CUBE); // Hin=M1=8 + loadData3DParams.l1W = BLOCK_CUBE; // Win=M0 + loadData3DParams.channelSize = constInfo_.headDim; // Cin=K loadData3DParams.padList[0] = 0; loadData3DParams.padList[1] = 0; loadData3DParams.padList[2] = 0; - loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果 + loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果 // SetLoadToA0Params - loadData3DParams.mExtension = CeilAlign(s1gL0RealSize, BLOCK_CUBE); // M height维度目的 - loadData3DParams.kExtension = constInfo_.headDim; // K width维度目的 + loadData3DParams.mExtension = CeilAlign(s1gL0RealSize, BLOCK_CUBE); // M height维度目的 + loadData3DParams.kExtension = constInfo_.headDim; // K width维度目的 loadData3DParams.mStartPt = s1gL1Offset; loadData3DParams.kStartPt = 0; loadData3DParams.strideW = 1; @@ -429,23 +434,22 @@ __aicore__ inline void QLIV2Matmul::LoadQueryToL0a(uint64_t s1gL1Offset, // s1, g, s2 --> 2 * 64* 128 template __aicore__ inline void QLIV2Matmul::LoadSToL0b(uint64_t s1gL1RealSize, uint64_t s2L0RealSize, - uint64_t sL1BufIdx, - int64_t mStartPt) + uint64_t sL1BufIdx, int64_t mStartPt) { LoadData3DParamsV2 loadData3DParams; // SetFmatrixParams - loadData3DParams.l1H = S1G_BASIC_BLOCK_L0 / BLOCK_CUBE; // Hin=M1=8 - loadData3DParams.l1W = BLOCK_CUBE; // Win=M0 - loadData3DParams.channelSize = CeilAlign(s2L0RealSize, BLOCK_CUBE); // Cin=K + loadData3DParams.l1H = S1G_BASIC_BLOCK_L0 / BLOCK_CUBE; // Hin=M1=8 + loadData3DParams.l1W = BLOCK_CUBE; // Win=M0 + loadData3DParams.channelSize = CeilAlign(s2L0RealSize, BLOCK_CUBE); // Cin=K loadData3DParams.padList[0] = 0; loadData3DParams.padList[1] = 0; loadData3DParams.padList[2] = 0; - loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果 + loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果 // SetLoadToA0Params - loadData3DParams.mExtension = constInfo_.gSize; // M height维度目的 - loadData3DParams.kExtension = CeilAlign(s2L0RealSize, BLOCK_CUBE); // K width维度目的 + loadData3DParams.mExtension = constInfo_.gSize; // M height维度目的 + loadData3DParams.kExtension = CeilAlign(s2L0RealSize, BLOCK_CUBE); // K width维度目的 loadData3DParams.kStartPt = 0; loadData3DParams.strideW = 1; loadData3DParams.strideH = 1; @@ -475,7 +479,7 @@ __aicore__ inline void QLIV2Matmul::LoadWeightToL0a(uint64_t s1gL1Offset loadData2DParams.dstGap = 0; loadData2DParams.ifTranspose = true; LoadData(l0a_.template ReinterpretCast()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K], - weightL1_[(qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM) * WEIGHT_BUFFER_OFFSET + s1gL1Offset* BLOCK_CUBE], + weightL1_[(qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM) * WEIGHT_BUFFER_OFFSET + s1gL1Offset * BLOCK_CUBE], loadData2DParams); } @@ -507,9 +511,8 @@ __aicore__ inline void QLIV2Matmul::ComputeWs(uint64_t s1gL0RealSize, ui mmadParams.cmatrixSource = false; Mmad(cL0_.template ReinterpretCast()[(l0cBufIdx_ % DOUBLE_BUF_NUM) * L0C_BUFFER_OFFSET + s1gOffset * S2_BASIC_BLOCK_L0], - l0a_.template ReinterpretCast()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K], - l0b_.template ReinterpretCast()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K], - mmadParams); + l0a_.template ReinterpretCast()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K], + l0b_.template ReinterpretCast()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K], mmadParams); } template @@ -553,8 +556,8 @@ __aicore__ inline void QLIV2Matmul::FixpSToL1(uint64_t s1gL0RealSize, ui template __aicore__ inline void QLIV2Matmul::FixpResToGm(uint64_t s1L0RealCount, uint64_t s2L0RealSize, - uint64_t s1GmOffset, - uint64_t s2GmOffset, const QLIV2Common::RunInfo &runInfo) + uint64_t s1GmOffset, uint64_t s2GmOffset, + const QLIV2Common::RunInfo &runInfo) { SetFlag(M_FIX_EVENT); WaitFlag(M_FIX_EVENT); @@ -613,5 +616,5 @@ __aicore__ inline void QLIV2Matmul::FreeEventID() WaitFlag(FIX_M_EVENT + 0); WaitFlag(FIX_M_EVENT + 1); } -} // namespace QLIV2Kernel -#endif // QUANT_LIGHTNING_INDEXER_V2_SERVICE_CUBE_H \ No newline at end of file +} // namespace QLIV2Kernel +#endif // QUANT_LIGHTNING_INDEXER_V2_SERVICE_CUBE_H diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_service_vector_arch22.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_service_vector_arch22.h index d95faf15d31a..42ab0aa74ac4 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_service_vector_arch22.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_service_vector_arch22.h @@ -42,7 +42,7 @@ class QLIV2Vector { // MM输出数据类型, 当前只支持float using MM1_OUT_T = float; - __aicore__ inline QLIV2Vector() {}; + __aicore__ inline QLIV2Vector(){}; __aicore__ inline void ProcessVec0(const QLIV2Common::RunInfo &info); __aicore__ inline void ProcessVec1(const QLIV2Common::RunInfo &info); __aicore__ inline void InitBuffers(TPipe *pipe); @@ -52,6 +52,10 @@ class QLIV2Vector { __aicore__ inline void ProcessLD(); __aicore__ inline void InitVecWorkspaceTensor(GlobalTensor vec0OutGm, GlobalTensor mm1ResGm, GlobalTensor vec1ResGm); + __aicore__ inline void InitVecCandidateTensor(GlobalTensor candidateTopkIndexInGm, + GlobalTensor candidateTopkIndexOutGm, + GlobalTensor outputIdxOffsetGm, + bool outputIdxOffsetValid); __aicore__ inline void InitVecInputTensor(GlobalTensor weightsGm, GlobalTensor qScaleGm, GlobalTensor kScaleGm, GlobalTensor indiceOutGm, GlobalTensor blockTableGm); @@ -69,12 +73,37 @@ class QLIV2Vector { GlobalTensor kScaleGm; GlobalTensor vec0OutGm; GlobalTensor indiceOutGm; + GlobalTensor candidateTopkIndexInGm; + GlobalTensor candidateTopkIndexOutGm; GlobalTensor blockTableGm; // =================================常量区================================= private: __aicore__ inline void GetKeyScale(const QLIV2Common::RunInfo &runInfo, const LocalTensor &resUb, int64_t batchId, int64_t startS2, int64_t getLen); + // candidate (two-level topk) + __aicore__ inline int32_t CountGE(const LocalTensor &sortedDesc, int32_t n, float x); + __aicore__ inline void BuildCandidateMask(const QLIV2Common::RunInfo &info, int32_t cuS1Idx, + int32_t cuBaseS2Idx, int32_t innerS1Idx); + __aicore__ inline void ProcessCandBlockTopk(const QLIV2Common::RunInfo &info, int32_t cuS1Idx, int32_t cuS2Len, + int32_t cuS2LenVecAlign, int32_t cuRealAcSeq, int32_t innerS1Idx); + __aicore__ inline void CopyOutCandTopkIndex(const QLIV2Common::RunInfo &info, int32_t cuS1Idx, int32_t innerS1Idx); + // candidate 新增向量 op 的 64 元素分块封装 (本版本 Level-2 count > 64 的 vmax 等指令会触发 aicore 异常) + __aicore__ inline void VecMaxPair(const LocalTensor &dst, const LocalTensor &src0, + const LocalTensor &src1, int32_t n); + __aicore__ inline void VecMinPair(const LocalTensor &dst, const LocalTensor &src0, + const LocalTensor &src1, int32_t n); + __aicore__ inline void VecAbs(const LocalTensor &dst, const LocalTensor &src, int32_t n); + __aicore__ inline void VecMinsScalar(const LocalTensor &dst, const LocalTensor &src, float v, + int32_t n); + __aicore__ inline void VecMaxsScalar(const LocalTensor &dst, const LocalTensor &src, float v, + int32_t n); + __aicore__ inline void VecSubPair(const LocalTensor &dst, const LocalTensor &src0, + const LocalTensor &src1, int32_t n); + __aicore__ inline void VecAddsScalar(const LocalTensor &dst, const LocalTensor &src, float v, + int32_t n); + __aicore__ inline void VecMulPair(const LocalTensor &dst, const LocalTensor &src0, + const LocalTensor &src1, int32_t n); // ================================Local Buffer区==================================== // queue TQue inQueue_; @@ -91,8 +120,18 @@ class QLIV2Vector { TBuf<> ldOutValueBuf_; TBuf<> ldOutIdxBuf_; + // tmp buff for candidate (two-level topk) + TBuf blockSortOutBuf_; // mode=1: 块级 topk 累加器 [CeilDiv(s1BaseSize,2), 2048, 2] + TBuf candBuf_; // mode=2: 排序后的候选块索引 [2048, 2] (value+idx) + TBuf candConstBuf_; // mode=1/2: negHuge 常量 [2048] + LocalTensor globalTopkIndice_; LocalTensor globalTopkUb_; + LocalTensor globalBlockTopkUb_; + LocalTensor candIsOut_; // mode=2: position 级 0/1 (1=候选外), 仅用于分数降级 (R11 leak) + LocalTensor candNegHuge_; // -1e30 常量 + GlobalTensor outputIdxOffsetGm_; // A15: 每行输出索引偏移 (仅 sparse_indices) + bool isOutputIdxOffsetValid_ = false; // A15: offset 是否传入 (经 InitVecCandidateTensor 传递) int32_t blockId_ = -1; // para for vector @@ -117,13 +156,17 @@ class QLIV2Vector { template __aicore__ inline void QLIV2Vector::GetKeyScale(const QLIV2Common::RunInfo &runInfo, - const LocalTensor &resUb, - int64_t batchId, int64_t startS2, int64_t getLen) + const LocalTensor &resUb, int64_t batchId, + int64_t startS2, int64_t getLen) { // startS2一定能整除kCacheBlockSize_ AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; AscendC::DataCopyExtParams copyInParams; if constexpr (PAGE_ATTENTION) { + // A11: k_scale 0 轴非连续 — 块基址 = 表值 x keyDequantScaleStride0 (真实 stride); + // 0 时兜底原紧凑公式 (kCacheBlockSize_), 现网行为不变 (R8) + int32_t kScaleBlkStride = constInfo_.keyDequantScaleStride0 != 0 ? + static_cast(constInfo_.keyDequantScaleStride0) : kCacheBlockSize_; int32_t startBlockTableIdx = startS2 / kCacheBlockSize_; int32_t startBlockTableOffset = startS2 % kCacheBlockSize_; int32_t blockTableBatchOffset = batchId * maxBlockNumPerBatch_; @@ -139,9 +182,7 @@ __aicore__ inline void QLIV2Vector::GetKeyScale(const QLIV2Common::RunIn int32_t blockId = blockTableGm.GetValue(blockTableBatchOffset + startBlockTableIdx); SetWaitFlag(HardEvent::S_MTE2); AscendC::DataCopyPad(resUb, - kScaleGm[static_cast(blockId) * constInfo_.keyDequantScaleStride0 + - startBlockTableOffset], - copyInParams, padParams); + kScaleGm[blockId * kScaleBlkStride + startBlockTableOffset], copyInParams, padParams); startBlockTableIdx++; getLen = getLen - firstPartLen; resUbBaseOffset = firstPartLen; @@ -155,8 +196,7 @@ __aicore__ inline void QLIV2Vector::GetKeyScale(const QLIV2Common::RunIn int32_t blockId = blockTableGm.GetValue(blockTableBatchOffset + startBlockTableIdx + i); SetWaitFlag(HardEvent::S_MTE2); AscendC::DataCopyPad(resUb[resUbBaseOffset + i * kCacheBlockSize_], - kScaleGm[static_cast(blockId) * constInfo_.keyDequantScaleStride0], - copyInParams, padParams); + kScaleGm[blockId * kScaleBlkStride], copyInParams, padParams); } } else { copyInParams.blockCount = 1; @@ -171,16 +211,32 @@ __aicore__ inline void QLIV2Vector::GetKeyScale(const QLIV2Common::RunIn template __aicore__ inline void QLIV2Vector::InitBuffers(TPipe *pipe) { - pipe->InitBuffer(inQueue_, 2, s2BaseSize_ * sizeof(float) * 2); // 32KB - pipe->InitBuffer(outQueue_, 1, BASE_TOPK * sizeof(float)); // 8 KB - pipe->InitBuffer(indexBuf_, s2BaseSize_ * sizeof(int32_t)); // 8 KB - pipe->InitBuffer(tmpBuf_, 64 * 1024); // 64KB - pipe->InitBuffer(sortOutBuf_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE * sizeof(float)); // 32KB + pipe->InitBuffer(inQueue_, 2, s2BaseSize_ * sizeof(float) * 2); // 32KB + pipe->InitBuffer(outQueue_, 1, BASE_TOPK * sizeof(float)); // 8 KB + pipe->InitBuffer(indexBuf_, s2BaseSize_ * sizeof(int32_t)); // 8 KB + pipe->InitBuffer(tmpBuf_, 64 * 1024); // 64KB + pipe->InitBuffer(sortOutBuf_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE * sizeof(float)); // 32KB globalTopkIndice_ = indexBuf_.Get(); globalTopkUb_ = sortOutBuf_.Get(); globalTopkNum_ = 0; + // candidate (two-level topk) 按模式分配, mode=3 不占用额外 UB + // UB 预算 (192KB): 基础 144KB + mode=1 blockSortOutBuf 32KB = 176KB / mode=2 candBuf 16KB + candConstBuf 8KB = 168KB + if (constInfo_.candidateMode == CANDIDATE_MODE_SOURCE) { + pipe->InitBuffer(blockSortOutBuf_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE * sizeof(float)); // 32KB + globalBlockTopkUb_ = blockSortOutBuf_.Get(); + InitSortOutBuf(globalBlockTopkUb_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE); + } else if (constInfo_.candidateMode == CANDIDATE_MODE_CONSUMER) { + // R6 修复: candBuf 按行分区 (每 AIV 处理 CeilDiv(s1BaseSize,2)=2 行, 行间 tile0 重排序 + // 会互相覆盖) — 2 行 x 2048 对 x 8B = 32KB; mode=2 UB 预算 144+32+8=184KB <= 192KB + pipe->InitBuffer(candBuf_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK * 2 * sizeof(float)); + pipe->InitBuffer(candConstBuf_, BASE_TOPK * sizeof(float)); // 8KB: negHuge + candNegHuge_ = candConstBuf_.Get(); + Duplicate(candNegHuge_.template ReinterpretCast(), QLIV2ServiceVec::NEG_HUGE_F32, BASE_TOPK); + PipeBarrier(); + } + // 基本块执行前初始化UB和GM // step1. 初始化一个有序索引 0 - s2BaseSize_ ArithProgression(globalTopkIndice_, 0, 1, s2BaseSize_); @@ -200,8 +256,8 @@ __aicore__ inline void QLIV2Vector::InitLDBuffers(TPipe *pipe) template __aicore__ inline void QLIV2Vector::InitParams(const struct QLIV2Common::ConstInfo &constInfo, - const struct QLIV2Common::LdSplitCoreInfo &ldInfo, - const QLIV2TilingData *__restrict tilingData) + const struct QLIV2Common::LdSplitCoreInfo &ldInfo, + const QLIV2TilingData *__restrict tilingData) { this->constInfo_ = constInfo; this->ldInfo_ = ldInfo; @@ -212,8 +268,8 @@ __aicore__ inline void QLIV2Vector::InitParams(const struct QLIV2Common: kHeadNum_ = constInfo.kHeadNum; qHeadNum_ = constInfo.qHeadNum; // define MMBase para - s1BaseSize_ = constInfo.s1BaseSize; // 4 - s2BaseSize_ = constInfo.s2BaseSize; // 2048 + s1BaseSize_ = constInfo.s1BaseSize; // 4 + s2BaseSize_ = constInfo.s2BaseSize; // 2048 kCacheBlockSize_ = constInfo.kCacheBlockSize; maxBlockNumPerBatch_ = constInfo.maxBlockNumPerBatch; blockId_ = GetBlockIdx(); @@ -221,10 +277,9 @@ __aicore__ inline void QLIV2Vector::InitParams(const struct QLIV2Common: template __aicore__ inline void QLIV2Vector::InitVecInputTensor(GlobalTensor weightsGm, - GlobalTensor qScaleGm, - GlobalTensor kScaleGm, - GlobalTensor indiceOutGm, - GlobalTensor blockTableGm) + GlobalTensor qScaleGm, GlobalTensor kScaleGm, + GlobalTensor indiceOutGm, + GlobalTensor blockTableGm) { this->weightsGm = weightsGm; this->qScaleGm = qScaleGm; @@ -235,8 +290,8 @@ __aicore__ inline void QLIV2Vector::InitVecInputTensor(GlobalTensor __aicore__ inline void QLIV2Vector::InitVecWorkspaceTensor(GlobalTensor vec0OutGm, - GlobalTensor mm1ResGm, - GlobalTensor vec1ResGm) + GlobalTensor mm1ResGm, + GlobalTensor vec1ResGm) { this->mm1ResGm = mm1ResGm; this->vec1ResGm = vec1ResGm; @@ -244,14 +299,24 @@ __aicore__ inline void QLIV2Vector::InitVecWorkspaceTensor(GlobalTensor< } template -__aicore__ inline void QLIV2Vector::AllocEventID() +__aicore__ inline void QLIV2Vector::InitVecCandidateTensor(GlobalTensor candidateTopkIndexInGm, + GlobalTensor candidateTopkIndexOutGm, + GlobalTensor outputIdxOffsetGm, + bool outputIdxOffsetValid) { + this->candidateTopkIndexInGm = candidateTopkIndexInGm; + this->candidateTopkIndexOutGm = candidateTopkIndexOutGm; + this->outputIdxOffsetGm_ = outputIdxOffsetGm; // A15 + this->isOutputIdxOffsetValid_ = outputIdxOffsetValid; } +template +__aicore__ inline void QLIV2Vector::AllocEventID() +{} + template __aicore__ inline void QLIV2Vector::FreeEventID() -{ -} +{} template __aicore__ inline void QLIV2Vector::CleanInvalidOutput(int64_t invalidS1offset) @@ -264,6 +329,383 @@ __aicore__ inline void QLIV2Vector::CleanInvalidOutput(int64_t invalidS1 valueULocal = outQueue_.DeQue(); QLIV2ServiceVec::CopyOut(indiceOutGm[invalidS1offset], idxULocal1, constInfo_.sparseCount); outQueue_.FreeTensor(valueULocal); + // mode=1: candidate_topk_index 同步填 -1 (行偏移 = sparse 偏移 / sparseCount * candidateTopkBlocks) + if (constInfo_.candidateMode == CANDIDATE_MODE_SOURCE) { + uint64_t candOffset = static_cast(invalidS1offset) / constInfo_.sparseCount * + constInfo_.candidateTopkBlocks; + LocalTensor candULocal = outQueue_.AllocTensor(); + LocalTensor candIdxLocal = candULocal.template ReinterpretCast(); + Duplicate(candIdxLocal, constInfo_.INVALID_IDX, constInfo_.candidateTopkBlocks); + outQueue_.EnQue(candULocal); + candULocal = outQueue_.DeQue(); + QLIV2ServiceVec::CopyOut(candidateTopkIndexOutGm[candOffset], candIdxLocal, constInfo_.candidateTopkBlocks); + outQueue_.FreeTensor(candULocal); + } +} + + +// 64 元素分块封装: 规避 Level-2 大 count 的 vmax/vmin/vabs 等指令异常 +template +__aicore__ inline void QLIV2Vector::VecMaxPair(const LocalTensor &dst, const LocalTensor &src0, + const LocalTensor &src1, int32_t n) +{ + for (int32_t off = 0; off < n; off += 64) { + Max(dst[off], src0[off], src1[off], 64); + } + PipeBarrier(); +} + +template +__aicore__ inline void QLIV2Vector::VecMinPair(const LocalTensor &dst, const LocalTensor &src0, + const LocalTensor &src1, int32_t n) +{ + for (int32_t off = 0; off < n; off += 64) { + Min(dst[off], src0[off], src1[off], 64); + } + PipeBarrier(); +} + +template +__aicore__ inline void QLIV2Vector::VecAbs(const LocalTensor &dst, const LocalTensor &src, + int32_t n) +{ + for (int32_t off = 0; off < n; off += 64) { + Abs(dst[off], src[off], 64); + } + PipeBarrier(); +} + +template +__aicore__ inline void QLIV2Vector::VecMinsScalar(const LocalTensor &dst, + const LocalTensor &src, float v, int32_t n) +{ + for (int32_t off = 0; off < n; off += 64) { + Mins(dst[off], src[off], v, 64); + } + PipeBarrier(); +} + +template +__aicore__ inline void QLIV2Vector::VecMaxsScalar(const LocalTensor &dst, + const LocalTensor &src, float v, int32_t n) +{ + for (int32_t off = 0; off < n; off += 64) { + Maxs(dst[off], src[off], v, 64); + } + PipeBarrier(); +} + + +template +__aicore__ inline void QLIV2Vector::VecSubPair(const LocalTensor &dst, const LocalTensor &src0, + const LocalTensor &src1, int32_t n) +{ + for (int32_t off = 0; off < n; off += 64) { + Sub(dst[off], src0[off], src1[off], 64); + } + PipeBarrier(); +} + +template +__aicore__ inline void QLIV2Vector::VecAddsScalar(const LocalTensor &dst, + const LocalTensor &src, float v, int32_t n) +{ + for (int32_t off = 0; off < n; off += 64) { + Adds(dst[off], src[off], v, 64); + } + PipeBarrier(); +} + +template +__aicore__ inline void QLIV2Vector::VecMulPair(const LocalTensor &dst, const LocalTensor &src0, + const LocalTensor &src1, int32_t n) +{ + for (int32_t off = 0; off < n; off += 64) { + Mul(dst[off], src0[off], src1[off], 64); + } + PipeBarrier(); +} + +// 统计降序序列 (偶数 lane 为 value) 中 >= x 的元素个数, 二分实现 +template +__aicore__ inline int32_t QLIV2Vector::CountGE(const LocalTensor &sortedDesc, int32_t n, float x) +{ + int32_t lo = 0; + int32_t hi = n; + while (lo < hi) { + int32_t mid = (lo + hi) / 2; + if (sortedDesc.GetValue(2 * mid) >= x) { + lo = mid + 1; + } else { + hi = mid; + } + } + return lo; +} + +// mode=2 (use_candidate): 加载/排序候选行 (每行首个S2分片), 判定 tile 内块归属, +// 产出 position 级 isOut (fp32 0/1, 1=候选外) 与其 int32 形式 +template +__aicore__ inline void QLIV2Vector::BuildCandidateMask(const QLIV2Common::RunInfo &info, int32_t cuS1Idx, + int32_t cuBaseS2Idx, int32_t innerS1Idx) +{ + int32_t candBlocks = static_cast(constInfo_.candidateTopkBlocks); + int32_t blockSize = static_cast(constInfo_.candidateBlockSize); + int32_t tileBlkNum = s2BaseSize_ / blockSize; + int32_t tileBlockBase = cuBaseS2Idx / blockSize; + // R6 修复: 同核多行 (每 AIV 处理 CeilDiv(s1BaseSize,2)=2 行) 共享 candBuf 时, + // 后一行的 tile0 重排序会覆盖前行候选 (s2 内层循环按 gS1 块整体推进) — candBuf 按行分区 + LocalTensor tmp = tmpBuf_.Get(); + LocalTensor candPairs = candBuf_.Get()[innerS1Idx * candBlocks * 2]; + if (info.isFirstS2InnerLoop) { + // 加载候选行到 tmp 尾部, 转 fp32 后与有序索引组成 [values | idx] 对并降序排序 + LocalTensor candInt = tmp[12288].template ReinterpretCast(); + AscendC::DataCopyExtParams copyInParams; + copyInParams.blockCount = 1; + copyInParams.blockLen = candBlocks * sizeof(int32_t); + copyInParams.srcStride = 0; + copyInParams.dstStride = 0; + copyInParams.rsv = 0; + AscendC::DataCopyPadExtParams padParams{true, 0, 0, 0}; + SetWaitFlag(HardEvent::V_MTE2); + AscendC::DataCopyPad(candInt, + candidateTopkIndexInGm[info.candidateOutOffset + + cuS1Idx * constInfo_.candidateTopkBlocks], + copyInParams, padParams); + SetWaitFlag(HardEvent::MTE2_V); + PipeBarrier(); + Cast(candPairs, candInt, RoundMode::CAST_NONE, candBlocks); + PipeBarrier(); + DataCopy(candPairs[candBlocks].template ReinterpretCast(), globalTopkIndice_, candBlocks); + PipeBarrier(); + LocalTensor candSortTmp = tmp[4096]; + QLIV2ServiceVec::SortAll(candPairs, candSortTmp, candBlocks); + PipeBarrier(); + // V->S 围栏: 排序结果对后续标量 GetValue (CountGE 二分) 可见 + // (PipeBarrier 不足以保证标量读到 V 写数据; 参照 CANN topk_v200 的 SetFlag/WaitFlag 模式) + SetWaitFlag(HardEvent::V_S); + } + // 候选与 tile 块号 [tileBlockBase, tileBlockBase+tileBlkNum) 的最小距离, 0 即命中 + // (candBuf 按行分区, R6: 同核行间覆盖已修复) + LocalTensor candPairs2 = candBuf_.Get()[innerS1Idx * candBlocks * 2]; + int32_t lo = CountGE(candPairs2, candBlocks, static_cast(tileBlockBase)); + int32_t hi = CountGE(candPairs2, candBlocks, static_cast(tileBlockBase + tileBlkNum)); + LocalTensor blkIdxF = tmp[4096]; // [tileBlkNum] + LocalTensor acc = tmp[4352]; // [tileBlkNum] + LocalTensor diff = tmp[4608]; // [tileBlkNum] + Cast(blkIdxF, globalTopkIndice_, RoundMode::CAST_NONE, tileBlkNum); + PipeBarrier(); + Adds(blkIdxF, blkIdxF, static_cast(tileBlockBase), tileBlkNum); + PipeBarrier(); + Duplicate(acc.ReinterpretCast(), QLIV2ServiceVec::POS_INF_F32, tileBlkNum); + PipeBarrier(); + // 降序排序下: 值 >= base+tileBlkNum 占 [0, hi), 落在 tile 内的候选占 [hi, lo), < base (含 -1 pad) 占 [lo, ...) + for (int32_t j = hi; j < lo; j++) { + float v = candPairs2.GetValue(2 * j); + Adds(diff, blkIdxF, -v, tileBlkNum); + PipeBarrier(); + VecAbs(diff, diff, tileBlkNum); + VecMinPair(acc, acc, diff, tileBlkNum); + } + // 距离 0 → 候选块; Brcb 展开到位置级; 距离 clamp 到 1 得 isOut (0/1) + LocalTensor posDist = tmp[12288]; // candInt 已释放, 可复用 + Brcb(posDist, acc, tileBlkNum / 8, {1, 8}); + PipeBarrier(); + VecMinsScalar(posDist, posDist, 1.0f, s2BaseSize_); + PipeBarrier(); + candIsOut_ = posDist; + // R11 (leak 语义, §1.6): isOut 仅用于分数降级 (pen 链), 索引保留真实位置号 — + // 候选外可达位置作为 topk 填充泄漏成有效索引 (对齐模型 where(idxs < compress_lens) 语义), + // 不再把候选外索引改 -1 (原 isOutI32/CAST_RINT 链已删)。 +} + +// mode=1 (is_candidate_source): 块内 amax (log2(blockSize) 轮 Max 树) + pin 尾块 + 块级排序/归并 +template +__aicore__ inline void QLIV2Vector::ProcessCandBlockTopk(const QLIV2Common::RunInfo &info, int32_t cuS1Idx, + int32_t cuS2Len, int32_t cuS2LenVecAlign, + int32_t cuRealAcSeq, int32_t innerS1Idx) +{ + int32_t blkLen = cuS2LenVecAlign; + int32_t blockSize = static_cast(constInfo_.candidateBlockSize); + int32_t blockNum = blkLen / blockSize; // 含尾块 (pad -inf 后按对齐长度计) + int32_t realBlockNum = (cuS2Len + blockSize - 1) / blockSize; // 含有效位置的块数 + // 所有向量 count 按 64 对齐 (非对齐 count 的向量指令在本架构触发 aicore 异常); + // 且块数必须过 AlignS2 (32*(4^n)*m, m<=3): SortAll 的 MrgSort 循环按 mrgGroups/4 缩减, + // 非 4 幂组数 (如 192 块=6 组) 会整组丢失 (实测 s2=1536 丢块 128..191 段) + int32_t blockNumPad = (blockNum < 64) ? 64 : AlignS2(blockNum); + int32_t tileBlockBase = info.s2Idx * s2BaseSize_ / blockSize; + LocalTensor tmp = tmpBuf_.Get(); + // 块内归约: vcgmax 每 32B 块 (8 fp32) 出 1 个 max, 紧凑输出 [blockNum] + LocalTensor blkScore = tmp[6144]; + // R10: blkLen=96 是 AlignS2 输出中唯一非 64 倍数的值 (≤128 段对齐到 32 的倍数), + // 96/64 整除截断为 1 只归约 [0,64) — 块 8..11 残留上一 tile/行的 stale 分数 + // (实测 big128k_b2_varlen b1 行 73 tile 56: 块 14344 拿到 tile 55 块 14088 的 + // 3.5568, 虚高挤掉 2048 名边界块 2760; 小 shape 总块数≤2048 集合不变故未暴露)。 + // 改 CeilDiv 覆盖全部块; 多归约的 [96,128) stale 只落在 pad 块槽位, 被 -inf + // 位型链位精确覆盖, 无害。 + int32_t brmRepeat = CeilDiv(blkLen, 64); + BlockReduceMax(blkScore, tmp[0], brmRepeat, 64, 1, 1, 8); + PipeBarrier(); + int32_t lastBlk = (cuRealAcSeq - 1) / blockSize; + int32_t pinLocal = lastBlk - tileBlockBase; + // 块级 [scores | idx] 对: idx = tileBlockBase + j; j >= realBlockNum 的 pad 块置 -1 + // (禁止标量 SetValue: 标量写不受 PipeBarrier 围栏, 与后续 V 管道读存在确定性竞态 + // (实测 innerS1Idx=1 行的填充被 SortAll 抢跑覆盖); 禁止 s32->f32 Cast (v220 无此组合)。 + // 改为纯 int32 向量算术: idx' = idx - (idx+1)*isPad, isPad = clamp(idx - thr, 0, 1), + // thr = base + realBlockNum - 1; 所有指令 count=blockNumPad (64 对齐), 按 64 分块避免 mask 寄存器限制。 + // scratch 用 mode=1 独占区 [14336,15360) (mode=2 的 isOutI32 已随 R11 leak 改造删除), + // 避开主路径 tmpSortBuf [4096,12288), 消除跨迭代残留读的隐患) + LocalTensor idxScr = tmp[14336].template ReinterpretCast(); // 等差源 + LocalTensor thrI = tmp[14592].template ReinterpretCast(); // 阈值 + LocalTensor offI = tmp[14848].template ReinterpretCast(); // idx-thr -> isPad + LocalTensor tI = tmp[15104].template ReinterpretCast(); // (idx+1)*isPad + LocalTensor blkIdx = tmp[6144 + blockNumPad].template ReinterpretCast(); // 最终 idx + ArithProgression(idxScr, static_cast(tileBlockBase), 1, blockNumPad); + Duplicate(thrI, tileBlockBase + realBlockNum - 1, blockNumPad); + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Sub(offI[off], idxScr[off], thrI[off], 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Maxs(offI[off], offI[off], 0, 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Mins(offI[off], offI[off], 1, 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Adds(tI[off], idxScr[off], 1, 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Mul(tI[off], tI[off], offI[off], 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Sub(blkIdx[off], idxScr[off], tI[off], 64); + } + PipeBarrier(); + // pad 块 score 归一为 -inf 位型: 覆盖 [realBlockNum, blockNumPad), 含 stale 区 [blockNum, blockNumPad) + // (尾 tile cuS2Len 非整块时 BlockReduceMax 只写 [0, blockNum), 其后是脏数据 — + // 脏 score 的 (-inf,-1) 不同构对会挤进 topk 挤掉真实块, 实测 128K 场景 -1 槽超标) + // score' = score + (score - NINF_bits)*(-isPad), isPad 于 pad 位为 1; 全 int32 向量 + if (blockNumPad > realBlockNum) { + LocalTensor pRaw = tmp[15360].template ReinterpretCast(); // -isPad + LocalTensor pNeg = tmp[15616].template ReinterpretCast(); // (score-NINF)*(-isPad) + LocalTensor scoreI = blkScore.template ReinterpretCast(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Duplicate(pRaw[off], -1, 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Mul(pRaw[off], pRaw[off], offI[off], 64); // -isPad + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Adds(pNeg[off], scoreI[off], 8388608, 64); // score - NINF_bits (NINF=0xFF800000=-8388608) + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Mul(pNeg[off], pNeg[off], pRaw[off], 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Add(scoreI[off], scoreI[off], pNeg[off], 64); // pad 块 -> -inf 位型 + } + PipeBarrier(); + } + // pin 最新 token 所在块 (+inf 位型注入该块 score): 纯 int32 向量实现, 在 isPad 链之后 + // (复用其 scratch; 标量 SetValue 写不受 PipeBarrier fence, 与 SortAll 多轮读存在竞态, 已弃用)。 + // flag = -1 于 pin 位, 0 于其余: score' = score + (score - PINF_bits) * flag, + // pin 位得 PINF_bits (0x7F800000 = +inf), 其余不变 (flag=0 使 pad 位 -inf 的中间溢出无害)。 + if (pinLocal >= 0 && pinLocal < realBlockNum) { + LocalTensor pRaw = tmp[15360].template ReinterpretCast(); // [15360,15616) + LocalTensor pNeg = tmp[15616].template ReinterpretCast(); // [15616,15872) + LocalTensor pA = thrI; // 复用 isPad 链已释放 scratch + LocalTensor pB = offI; + LocalTensor pFlag = tI; + LocalTensor scoreI = blkScore.template ReinterpretCast(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + // idxScr 为全局块号, pin 也必须用全局 id (tileBlockBase + pinLocal) 比较 + Adds(pRaw[off], idxScr[off], -(tileBlockBase + pinLocal), 64); // 0 于 pin 块 + } + PipeBarrier(); + Duplicate(pNeg, 0, blockNumPad); + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Sub(pNeg[off], pNeg[off], pRaw[off], 64); // -(idx-pinLocal) + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Maxs(pA[off], pRaw[off], 0, 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Mins(pA[off], pA[off], 1, 64); // 1 若 idx > pin + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Maxs(pB[off], pNeg[off], 0, 64); + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Mins(pB[off], pB[off], 1, 64); // 1 若 idx < pin + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Add(pFlag[off], pA[off], pB[off], 64); // isNotPin (0/1, 互斥) + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Adds(pFlag[off], pFlag[off], -1, 64); // -1 于 pin, 0 其余 + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Adds(pRaw[off], scoreI[off], -2139095040, 64); // score - 0x7F800000 (复用 pRaw) + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Mul(pRaw[off], pRaw[off], pFlag[off], 64); // (score-PINF)*flag + } + PipeBarrier(); + for (int32_t off = 0; off < blockNumPad; off += 64) { + Add(scoreI[off], scoreI[off], pRaw[off], 64); // pin -> +inf 位型 + } + PipeBarrier(); + } + // 块级排序 + 归并到块级累加器 + LocalTensor blkPairs = tmp[6144]; // [scores blockNumPad | idx blockNumPad] + // MrgSort tmp 需容纳 mrgDst+mrgSrc = 2048*2 + blockNumPad*2 <= 4608 floats; + // 11776 + 4608 = 16384 恰为 tmpBuf_(64KB) 末尾, 12288 起会越界 2KB (aicore 异常) + LocalTensor blkSortTmp = tmp[11776]; // 与 blkPairs 不重叠 (MrgSort src/tmp 禁止重叠) + QLIV2ServiceVec::SortAll(blkPairs, blkSortTmp, blockNumPad); + PipeBarrier(); + // 候选专用合并: 纯 V 回拷 (MergeSort 尾部 UB->UB DataCopy 不受 PipeBarrier fence, + // 多 tile 下一次 MrgSort 读 acc 存在调度敏感竞态, 实测 tile1 读到 stale 数据) + QLIV2ServiceVec::MergeSortVecCopy(globalBlockTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], + constInfo_.candidateTopkBlocks, blkPairs, blockNumPad, blkSortTmp); + PipeBarrier(); +} + +// mode=1: 行末直出 candidate_topk_index 并复位块级累加器 +template +__aicore__ inline void QLIV2Vector::CopyOutCandTopkIndex(const QLIV2Common::RunInfo &info, int32_t cuS1Idx, + int32_t innerS1Idx) +{ + LocalTensor candIdxULocal = outQueue_.AllocTensor(); + ExtractIndex(candIdxULocal, + globalBlockTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE].template ReinterpretCast(), + constInfo_.candidateTopkBlocks); + PipeBarrier(); + InitSortOutBuf(globalBlockTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], BASE_TOPK_VALUE_IDX_SIZE); + outQueue_.EnQue(candIdxULocal); + candIdxULocal = outQueue_.DeQue(); + QLIV2ServiceVec::CopyOut(candidateTopkIndexOutGm[info.candidateOutOffset + + cuS1Idx * constInfo_.candidateTopkBlocks], + candIdxULocal.template ReinterpretCast(), constInfo_.candidateTopkBlocks); + outQueue_.FreeTensor(candIdxULocal); } template @@ -280,14 +722,15 @@ __aicore__ inline void QLIV2Vector::ProcessVec0(const QLIV2Common::RunIn int64_t weightGmOffset = info.tensorWeightsOffset + cuBaseS1Idx * qHeadNum_; // 当前需要计算的S1行数,处理尾块场景 int32_t cuS1ProcNum = cuBaseS1Idx + s1BaseSize_ > info.actS1Size ? info.actS1Size % s1BaseSize_ : s1BaseSize_; - int32_t cuProcEleNum = cuS1ProcNum * gSize_; + int32_t cuProcRealNum = cuS1ProcNum * gSize_; + int32_t cuProcEleNum = QLIV2Common::Align(cuProcRealNum, 32); // 32: UB对齐, 参照v1 LocalTensor inWeightsUb = inQueue_.AllocTensor(); LocalTensor inQScaleUb = inWeightsUb[cuProcEleNum]; AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; AscendC::DataCopyExtParams copyInParams; copyInParams.blockCount = 1; - copyInParams.blockLen = cuProcEleNum * sizeof(half); + copyInParams.blockLen = cuProcRealNum * sizeof(half); copyInParams.srcStride = 0; copyInParams.dstStride = 0; copyInParams.rsv = 0; @@ -306,7 +749,7 @@ __aicore__ inline void QLIV2Vector::ProcessVec0(const QLIV2Common::RunIn resUb = outQueue_.DeQue(); AscendC::DataCopyParams copyOutParams; copyOutParams.blockCount = 1; - copyOutParams.blockLen = cuProcEleNum * BLOCK_CUBE * sizeof(half); + copyOutParams.blockLen = cuProcRealNum * BLOCK_CUBE * sizeof(half); copyOutParams.srcStride = 0; copyOutParams.dstStride = 0; AscendC::DataCopyPad(vec0OutGm[vec0OutGmOffset], resUb, copyOutParams); @@ -350,6 +793,9 @@ __aicore__ inline void QLIV2Vector::ProcessVec1(const QLIV2Common::RunIn // globalTopkUb_ value,index=-inf,-1 InitSortOutBuf(globalTopkUb_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE); blockS2StartIdx_ = 0; + if (constInfo_.candidateMode == CANDIDATE_MODE_SOURCE) { + InitSortOutBuf(globalBlockTopkUb_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE); + } } else if (info.loop == 0) { blockS2StartIdx_ = info.s2Idx; } @@ -395,6 +841,19 @@ __aicore__ inline void QLIV2Vector::ProcessVec1(const QLIV2Common::RunIn PipeBarrier(); AscendC::Mul(mmInUb, mmInUb, kScaleUb, cuS2Len); PipeBarrier(); + // mode=2 (use_candidate, R11 leak): 候选块外 score 降级 NEG_HUGE (S4a, 纯算术无 NaN); + // 仅降级排序资格、不取消入选资格 — 不足 topk 时候选外可达位置作为填充泄漏 (§1.6) + if (constInfo_.candidateMode == CANDIDATE_MODE_CONSUMER) { + BuildCandidateMask(info, cuS1Idx, cuBaseS2Idx, innerS1Idx); + PipeBarrier(); + LocalTensor pen = tmpBuf_.Get()[4096]; // blkIdxF/acc/d 已释放, 复用 + Sub(pen, candNegHuge_, mmInUb, cuS2Len); + PipeBarrier(); + Mul(pen, pen, candIsOut_, cuS2Len); + PipeBarrier(); + Add(mmInUb, mmInUb, pen, cuS2Len); + PipeBarrier(); + } LocalTensor sortBuff = tmpBuf_.Get(); LocalTensor sortScoreUb = sortBuff; LocalTensor sortIndiceUb = sortBuff[cuS2LenVecAlign]; @@ -412,11 +871,19 @@ __aicore__ inline void QLIV2Vector::ProcessVec1(const QLIV2Common::RunIn } Adds(sortIndiceUbInt, globalTopkIndice_, static_cast(cuBaseS2Idx), cuS2Len); PipeBarrier(); + // R11 (leak 语义, §1.6): 候选块外索引不再置 -1 — + // 分数已在 S4a 降级 NEG_HUGE (仍高于 -inf: 不足 topk 时作为填充入选, 排候选内之后), + // 索引保留真实位置号, 与模型 where(idxs < compress_lens, idxs + offset, -1) 等价: + // 仅不可达位置 (score -inf 沉底 + 尾部 -1 填充) 输出 -1。 + // mode=1 (is_candidate_source): 块化 amax + pin + 块级排序归并 (S5a/S6a) + if (constInfo_.candidateMode == CANDIDATE_MODE_SOURCE) { + ProcessCandBlockTopk(info, cuS1Idx, cuS2Len, cuS2LenVecAlign, cuRealAcSeq, innerS1Idx); + } LocalTensor tmpSortBuf = sortBuff[2 * cuS2LenVecAlign]; QLIV2ServiceVec::SortAll(sortBuff, tmpSortBuf, cuS2LenVecAlign); PipeBarrier(); QLIV2ServiceVec::MergeSort(globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], BASE_TOPK, sortBuff, - cuS2LenVecAlign, tmpSortBuf); + cuS2LenVecAlign, tmpSortBuf); PipeBarrier(); bool isS2End = cuBaseS2Idx + s2BaseSize_ >= cuRealAcSeq; bool needCopyOutGm = blockS2StartIdx_ == 0 && isS2End; @@ -427,23 +894,41 @@ __aicore__ inline void QLIV2Vector::ProcessVec1(const QLIV2Common::RunIn globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE].template ReinterpretCast(), BASE_TOPK); PipeBarrier(); + // A15: output_idx_offset — 输出拷 GM 前逐元素加行偏移 (对齐 arch35 IndicesAddOffset; + // 零偏移零开销; GM 标量读无 V_S 竞态; int32 整数域 Adds 位精确, -1 槽 +0 不变。 + // candidate_topk_index 不加 (相对块号契约, §11.6)) + if (isOutputIdxOffsetValid_) { + int32_t rowOff = outputIdxOffsetGm_.GetValue(info.outputIdxOffsetCoreOffset + + cuS1Idx * kHeadNum_); + if (rowOff != 0) { + LocalTensor idxI32 = idxULocal.template ReinterpretCast(); + for (int32_t off = 0; off < constInfo_.sparseCount; off += 64) { + Adds(idxI32[off], idxI32[off], rowOff, 64); + } + PipeBarrier(); + } + } InitSortOutBuf(globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], BASE_TOPK_VALUE_IDX_SIZE); outQueue_.EnQue(idxULocal); idxULocal = outQueue_.DeQue(); QLIV2ServiceVec::CopyOut(indiceOutGm[info.indiceOutOffset + cuS1Idx * constInfo_.sparseCount], - idxULocal.template ReinterpretCast(), constInfo_.sparseCount); + idxULocal.template ReinterpretCast(), constInfo_.sparseCount); outQueue_.FreeTensor(idxULocal); + // mode=1: 块级 topk 直出 candidate_topk_index + if (constInfo_.candidateMode == CANDIDATE_MODE_SOURCE) { + CopyOutCandTopkIndex(info, cuS1Idx, innerS1Idx); + } } // LD拷贝到当前vector对应S1的位置 - if (info.isNeedLD && info.isLastS2InnerLoop) { // 当前核存在归约任务 且是最后处理的一段 + if (info.isNeedLD && info.isLastS2InnerLoop) { // 当前核存在归约任务 且是最后处理的一段 AscendC::DataCopyExtParams copyWsParams; copyWsParams.blockLen = BASE_TOPK_VALUE_IDX_SIZE * sizeof(float); copyWsParams.srcStride = 0; copyWsParams.dstStride = 0; copyWsParams.blockCount = 1; SetWaitFlag(HardEvent::V_MTE3); - AscendC::DataCopyPad(vec1ResGm[wsOffset], - globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], copyWsParams); + AscendC::DataCopyPad(vec1ResGm[wsOffset], globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], + copyWsParams); SetWaitFlag(HardEvent::MTE3_V); InitSortOutBuf(globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], BASE_TOPK_VALUE_IDX_SIZE); PipeBarrier(); @@ -455,17 +940,17 @@ __aicore__ inline void QLIV2Vector::ProcessVec1(const QLIV2Common::RunIn InitSortOutBuf(globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], BASE_TOPK_VALUE_IDX_SIZE); SetWaitFlag(HardEvent::V_MTE3); QLIV2ServiceVec::CopyOut(vec1ResGm[wsOffset], globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], - BASE_TOPK_VALUE_IDX_SIZE); + BASE_TOPK_VALUE_IDX_SIZE); SetWaitFlag(HardEvent::MTE3_V); } else { CleanInvalidOutput(info.indiceOutOffset + cuS1Idx * constInfo_.sparseCount); } } else if (cuS2Len <= 0) { // LD拷贝到当前vector对应S1的位置 - if (info.isNeedLD && info.isLastS2InnerLoop) { // 当前核存在归约任务 且是最后处理的一段 + if (info.isNeedLD && info.isLastS2InnerLoop) { // 当前核存在归约任务 且是最后处理的一段 SetWaitFlag(HardEvent::V_MTE3); QLIV2ServiceVec::CopyOut(vec1ResGm[wsOffset], globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], - BASE_TOPK_VALUE_IDX_SIZE); + BASE_TOPK_VALUE_IDX_SIZE); SetWaitFlag(HardEvent::MTE3_V); InitSortOutBuf(globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], BASE_TOPK_VALUE_IDX_SIZE); PipeBarrier(); @@ -526,9 +1011,9 @@ __aicore__ inline void QLIV2Vector::ProcessLD() indexValueParams.dstStride = 0; uint32_t ldProWorkspaceNum = ldInfo_.workspaceNum; - uint32_t ldProcessLen = 4; // 4: 4块归约任务做一次Merge + uint32_t ldProcessLen = 4; // 4: 4块归约任务做一次Merge uint32_t ldProcessNum = (ldProWorkspaceNum - 1) / (ldProcessLen - 1); - uint32_t ldTailLen = ldProWorkspaceNum - (ldProcessNum * (ldProcessLen - 1) + 1); // 尾块长度 + uint32_t ldTailLen = ldProWorkspaceNum - (ldProcessNum * (ldProcessLen - 1) + 1); // 尾块长度 for (uint32_t j = 0; j < ldInfo_.mNum; j++) { SetWaitFlag(HardEvent::V_MTE2); @@ -543,15 +1028,15 @@ __aicore__ inline void QLIV2Vector::ProcessLD() // 处理等于4的部分 for (uint32_t i = 0; i < ldProcessNum; i++) { // LD处理偏移 - uint64_t wsOffset = static_cast(ldInfo_.workspaceIdx) * s1BaseSize_ * BASE_TOPK_VALUE_IDX_SIZE + - static_cast(ldInfo_.mStart + j) * BASE_TOPK_VALUE_IDX_SIZE+ - static_cast(i * (ldProcessLen - 1) + 1) * - s1BaseSize_ * BASE_TOPK_VALUE_IDX_SIZE; - indexValueParams.blockCount = ldProcessLen - 1; // 拷贝4块进行merge + uint64_t wsOffset = + static_cast(ldInfo_.workspaceIdx) * s1BaseSize_ * BASE_TOPK_VALUE_IDX_SIZE + + static_cast(ldInfo_.mStart + j) * BASE_TOPK_VALUE_IDX_SIZE + + static_cast(i * (ldProcessLen - 1) + 1) * s1BaseSize_ * BASE_TOPK_VALUE_IDX_SIZE; + indexValueParams.blockCount = ldProcessLen - 1; // 拷贝4块进行merge SetWaitFlag(HardEvent::V_MTE2); - AscendC::DataCopyPad(curValueIdxUb[valueOffset], vec1ResGm[wsOffset], - indexValueParams, indexValuePadParams); + AscendC::DataCopyPad(curValueIdxUb[valueOffset], vec1ResGm[wsOffset], indexValueParams, + indexValuePadParams); // merge参数 AscendC::MrgSort4Info params; params.elementLengths[0] = BASE_TOPK; @@ -577,15 +1062,14 @@ __aicore__ inline void QLIV2Vector::ProcessLD() // 处理不等于4的部分 if (ldTailLen != 0) { // 搬运尾块 - uint64_t wsOffsetTail = static_cast(ldInfo_.workspaceIdx) * s1BaseSize_ - * BASE_TOPK_VALUE_IDX_SIZE + - static_cast(ldInfo_.mStart + j) * BASE_TOPK_VALUE_IDX_SIZE + - static_cast(ldProcessNum * (ldProcessLen - 1) + 1) * s1BaseSize_ - * BASE_TOPK_VALUE_IDX_SIZE; + uint64_t wsOffsetTail = + static_cast(ldInfo_.workspaceIdx) * s1BaseSize_ * BASE_TOPK_VALUE_IDX_SIZE + + static_cast(ldInfo_.mStart + j) * BASE_TOPK_VALUE_IDX_SIZE + + static_cast(ldProcessNum * (ldProcessLen - 1) + 1) * s1BaseSize_ * BASE_TOPK_VALUE_IDX_SIZE; indexValueParams.blockCount = ldTailLen; SetWaitFlag(HardEvent::V_MTE2); - AscendC::DataCopyPad(curValueIdxUb[valueOffset], vec1ResGm[wsOffsetTail], - indexValueParams, indexValuePadParams); + AscendC::DataCopyPad(curValueIdxUb[valueOffset], vec1ResGm[wsOffsetTail], indexValueParams, + indexValuePadParams); SetWaitFlag(HardEvent::MTE2_V); AscendC::MrgSort4Info params; params.elementLengths[0] = BASE_TOPK; @@ -617,14 +1101,27 @@ __aicore__ inline void QLIV2Vector::ProcessLD() PipeBarrier(); InitSortOutBuf(curValueIdxUb, BASE_TOPK_VALUE_IDX_SIZE); LocalTensor idxULocal1 = outIdxUb.template ReinterpretCast(); + // A15: LD(decode) 路径同样加 output_idx_offset (行前缀 = indiceOutCoreOffset/sparseCount, 即 + // batch 前缀 x kHeadNum + n2; 与 ProcessVec1 消费点同语义) + if (isOutputIdxOffsetValid_) { + uint32_t rowGlobal = ldInfo_.mStart + j; + int32_t rowOff = outputIdxOffsetGm_.GetValue( + ldInfo_.indiceOutCoreOffset / constInfo_.sparseCount + rowGlobal * constInfo_.kHeadNum); + if (rowOff != 0) { + for (int32_t off = 0; off < constInfo_.sparseCount; off += 64) { + Adds(idxULocal1[off], idxULocal1[off], rowOff, 64); + } + PipeBarrier(); + } + } SetWaitFlag(HardEvent::V_MTE3); - uint64_t outOffset = ldInfo_.indiceOutCoreOffset + - (ldInfo_.mStart + j) * constInfo_.kHeadNum * constInfo_.sparseCount; + uint64_t outOffset = + ldInfo_.indiceOutCoreOffset + (ldInfo_.mStart + j) * constInfo_.kHeadNum * constInfo_.sparseCount; AscendC::DataCopyPad(indiceOutGm[outOffset], idxULocal1, copyOutParams); SetWaitFlag(HardEvent::MTE3_V); } SetWaitFlag(HardEvent::MTE3_V); } -} // namespace QLIV2Kernel -#endif // QUANT_LIGHTNING_INDEXER_V2_SERVICE_VECTOR_H \ No newline at end of file +} // namespace QLIV2Kernel +#endif // QUANT_LIGHTNING_INDEXER_V2_SERVICE_VECTOR_H diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_vector.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_vector.h index 39a36af3a118..3c5fa3737d36 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_vector.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch22/quant_lightning_indexer_v2_vector.h @@ -22,6 +22,8 @@ namespace QLIV2ServiceVec { using namespace AscendC; constexpr int32_t NEG_INF = 0xFF800000; +constexpr int32_t POS_INF_F32 = 0x7F800000; // +inf (fp32 bits) +constexpr int32_t NEG_HUGE_F32 = 0xF14A3E31; // -1e30 (fp32 bits) constexpr int32_t INVALID_INDEX = -1; constexpr uint8_t VEC_REPEAT_MAX = 255; constexpr uint8_t B32_VEC_ELM_NUM = 64; @@ -175,12 +177,45 @@ __aicore__ inline void ExtractIndex(const LocalTensor &idxULocal, cons gatherMaskParams.src0BlockStride = 1; gatherMaskParams.src0RepeatStride = B32_VEC_REPEAT_STRIDE; gatherMaskParams.src1RepeatStride = 0; - uint64_t rsvdCnt = 0; // 用于保存筛选后保留下来的元素个数 - uint8_t src1Pattern = 2; // 固定模式2,表示筛选出奇数索引的数 + uint64_t rsvdCnt = 0; // 用于保存筛选后保留下来的元素个数 + uint8_t src1Pattern = 2; // 固定模式2,表示筛选出奇数索引的数 AscendC::GatherMask(idxULocal, sortLocal, src1Pattern, false, static_cast(0), gatherMaskParams, rsvdCnt); AscendC::PipeBarrier(); } +/** + 候选路径专用合并: 与 MergeSort 语义相同, 但回拷 mrgDst 用纯向量 Adds (64 分块)。 + MergeSort 尾部的 UB->UB DataCopy 若走 MTE 管道, 其后的 PipeBarrier 不构成 fence, + 跨 tile/跨行对累加器的下一次 MrgSort 读存在调度敏感竞态 (实测多 tile 场景 tile1 读到 tile0 stale + 数据, 累加器出现 -inf 位型+8k 垃圾且 tile1 块整段缺失)。纯 V 回拷全程受 PipeBarrier 约束。 + */ +__aicore__ inline void MergeSortVecCopy(const LocalTensor &mrgDst, int32_t mrgDstNum, + LocalTensor &mrgSrc, int32_t mrgSrcNum, LocalTensor &tmpTensor) +{ + AscendC::MrgSort4Info params; + params.elementLengths[0] = mrgSrcNum; + params.elementLengths[1] = mrgDstNum; + params.ifExhaustedSuspension = false; + params.validBit = 0b0011; + params.repeatTimes = 1; + + AscendC::MrgSortSrcList srcList; + srcList.src1 = mrgSrc; + srcList.src2 = mrgDst; + + AscendC::MrgSort(tmpTensor, srcList, params); + AscendC::PipeBarrier(); + // 回拷必须位精确: -1 (0xFFFFFFFF) 作为 float 是负符号 NaN, vadd 会规范化为 0x7FFFFFFF; + // 用 int32 视图 Adds(+0) 保证 bit-exact 且全程 V 管道 (受 PipeBarrier fence) + int32_t copyNum = mrgDstNum * VALUE_AND_INDEX_NUM; + LocalTensor dstI = mrgDst.template ReinterpretCast(); + LocalTensor srcI = tmpTensor.template ReinterpretCast(); + for (int32_t off = 0; off < copyNum; off += 64) { + AscendC::Adds(dstI[off], srcI[off], 0, 64); + } + AscendC::PipeBarrier(); +} + template __aicore__ inline void SetWaitFlag(HardEvent evt) { @@ -189,5 +224,5 @@ __aicore__ inline void SetWaitFlag(HardEvent evt) AscendC::WaitFlag(eventId); } -} // namespace QLIV2ServiceVec -#endif // QUANT_LIGHTNING_INDEXER_V2_VECTOR_H +} // namespace QLIV2ServiceVec +#endif // QUANT_LIGHTNING_INDEXER_V2_VECTOR_H diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/quant_lightning_indexer_v2_service_cube_arch35.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/quant_lightning_indexer_v2_service_cube_arch35.h index d63962ecf70d..558f1fbc9a24 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/quant_lightning_indexer_v2_service_cube_arch35.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/quant_lightning_indexer_v2_service_cube_arch35.h @@ -142,8 +142,8 @@ class QLIV2Matmul { uint64_t keyGmStart_ = 0; // L1数据对应的GM S2偏移 uint64_t keyLoadedSize_ = 0; // L1中实际加载的S2元素数量 uint64_t s2BasicBlock_ = 128; - uint64_t qkHeadDim_ = 128; // Q/K单行GM搬入宽度;MXFP4为打包后的headDim/2,其他场景为headDim - uint64_t scaleHeadDim_ = 4; // MX scale单行元素数,即headDim/32,MXFP8/MXFP4共用 + uint64_t qkHeadDim_ = 128; // Q/K单行GM搬入宽度;MXFP4为打包后的headDim/2,其他场景为headDim + uint64_t scaleHeadDim_ = 4; // MX scale单行元素数,即headDim/32,MXFP8/MXFP4共用 uint64_t keyBufferOffset_ = 16384; // Key L1乒乓缓冲区步长,s2BasicBlock_ * D_BASIC_BLOCK uint64_t keyScaleBufferOffset_ = 512; // Key scale L1乒乓缓冲区步长,s2BasicBlock_ * D_BASIC_BLOCK / 32 @@ -214,7 +214,7 @@ __aicore__ inline void QLIV2Matmul::InitMm1GlobalTensor(const GlobalTens keyGm_ = keyGm; queryGm_ = queryGm; if constexpr (IS_MX) { - mxKeyScaleGmBf16_ = keyScaleGmBf16; // gitleaks:allow + mxKeyScaleGmBf16_ = keyScaleGmBf16; mxQueryScaleGmBf16_ = queryScaleGmBf16; } } @@ -712,9 +712,9 @@ __aicore__ inline void QLIV2Matmul::Fixp(uint64_t s1gGmOffset, uint64_t cL0_[(l0BufIdx_ % L0_BUF_NUM) * l0cBufferOffset_], fixpipeParams); fixpipeParams.subBlockId = 1; - Fixpipe( - mm1ResUB_[(runInfo.loop % 2) * (UB_BANK_STRIDE / sizeof(QK_T))], - cL0_[(l0BufIdx_ % L0_BUF_NUM) * l0cBufferOffset_ + mSize / 2 * 16], fixpipeParams); + Fixpipe(mm1ResUB_[(runInfo.loop % 2) * (UB_BANK_STRIDE / sizeof(QK_T))], + cL0_[(l0BufIdx_ % L0_BUF_NUM) * l0cBufferOffset_ + mSize / 2 * 16], + fixpipeParams); } } diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/quant_lightning_indexer_v2_service_vector_arch35.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/quant_lightning_indexer_v2_service_vector_arch35.h index 33700cf83796..6859318e18c4 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/quant_lightning_indexer_v2_service_vector_arch35.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/quant_lightning_indexer_v2_service_vector_arch35.h @@ -176,7 +176,7 @@ class QLIV2Vector { float globalQScale_ = 1.0f; // quantMode=4时全局query scale float globalKScale_ = 1.0f; // quantMode=4时全局key scale uint32_t trunkLen_ = 0; // ProcessTopK(非LD路径)每次处理的s2长度 - uint32_t trunkLenLd_ = 0; // ProcessLD路径每次处理的s2长度,根据sparseCount动态计算以填满UB + uint32_t trunkLenLd_ = 0; // ProcessLD路径每次处理的s2长度,根据sparseCount动态计算以填满UB bool returnValueFlag = false; struct QLIV2Common::ConstInfo constInfo_; @@ -199,8 +199,7 @@ __aicore__ inline void QLIV2Vector::InitBuffers(TPipe *pipe) pipe->InitBuffer(kScaleBuf_, 2 * s2BaseSize_ * 16 * sizeof(SCALE_T)); kScaleUB_ = kScaleBuf_.Get(); // kScale // 大小:2(开dB) * 2 * 64 * 4 = 1KB - pipe->InitBuffer(qScaleBuf_, - 2 * CeilDiv(s1BaseSize_, 2) * UB_BANK_DEPTH_STRIDE); + pipe->InitBuffer(qScaleBuf_, 2 * CeilDiv(s1BaseSize_, 2) * UB_BANK_DEPTH_STRIDE); qScaleUB_ = qScaleBuf_.Get(); // qScale // 大小:2(开dB) * 2 * 128 * 4 = 2KB pipe->InitBuffer(outBuf_, 2 * CeilDiv(s1BaseSize_, 2) * s2BaseSize_ * sizeof(SCORE_T)); @@ -431,9 +430,8 @@ __aicore__ inline void QLIV2Vector::DoTndPadding(const QLIV2Common::RunI } template __aicore__ inline void QLIV2Vector::GetKeyScale(const QLIV2Common::RunInfo &runInfo, - LocalTensor &kScaleUB, - int64_t batchId, int64_t startS2, - int64_t getLen) + LocalTensor &kScaleUB, int64_t batchId, + int64_t startS2, int64_t getLen) { // startS2一定能整除kCacheBlockSize_ AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; @@ -574,15 +572,13 @@ __aicore__ inline void QLIV2Vector::ProcessVec1(const QLIV2Common::RunIn curAivS1ProcNum); } else if constexpr (IS_WEIGHT_FP16) { auto qScaleBase = qScaleUB_[qScalepingpong * (UB_BANK_STRIDE / sizeof(WEIGHT_T))]; - auto kScaleBase = kScaleUB_[kScalepingpong * 16 * s2BaseSize_ + - ((info.s2Idx - info.s2Start) % 16) * s2BaseSize_]; - vector1::BatchMulWeightAndReduceSum(outBase, UB_BANK_DEPTH_STRIDE / sizeof(SCORE_T), - qkBase, qkVLstride, - (uint32_t)(gSize_ * UB_BANK_DEPTH_STRIDE / sizeof(QK_T)), - weightBase, UB_BANK_DEPTH_STRIDE / sizeof(WEIGHT_T), weightTempBase, - kScaleBase, (uint32_t)0, - qScaleBase, UB_BANK_DEPTH_STRIDE / sizeof(SCALE_T), - gSize_, curAivS1ProcNum); + auto kScaleBase = + kScaleUB_[kScalepingpong * 16 * s2BaseSize_ + ((info.s2Idx - info.s2Start) % 16) * s2BaseSize_]; + vector1::BatchMulWeightAndReduceSum(outBase, UB_BANK_DEPTH_STRIDE / sizeof(SCORE_T), qkBase, qkVLstride, + (uint32_t)(gSize_ * UB_BANK_DEPTH_STRIDE / sizeof(QK_T)), weightBase, + UB_BANK_DEPTH_STRIDE / sizeof(WEIGHT_T), weightTempBase, kScaleBase, + (uint32_t)0, qScaleBase, UB_BANK_DEPTH_STRIDE / sizeof(SCALE_T), gSize_, + curAivS1ProcNum); } else if (constInfo_.quantMode == 4) { // 4: per_tensor量化 // quantMode为4时不适用sacle的UB float kScaleValue = globalKScale_; diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/quant_lightning_indexer_v2_topk.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/quant_lightning_indexer_v2_topk.h index c715bd318611..99417592c02a 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/quant_lightning_indexer_v2_topk.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/quant_lightning_indexer_v2_topk.h @@ -20,17 +20,15 @@ #include "vf_topk_16_gather_quant_v2.h" namespace topk { -template +template class LITopk { public: - __aicore__ inline void operator()(LocalTensor& outputIdxLocal, - LocalTensor& inputLocal, + __aicore__ inline void operator()(LocalTensor &outputIdxLocal, LocalTensor &inputLocal, uint32_t s2SeqLen) - { - } + {} }; -template<> +template <> class LITopk { public: static __aicore__ inline uint32_t GetSharedTmpBufferSize(uint32_t topK) @@ -49,7 +47,7 @@ class LITopk { this->topK = topK; } - __aicore__ inline void InitBuffers(LocalTensor& sharedTmpBuffer) + __aicore__ inline void InitBuffers(LocalTensor &sharedTmpBuffer) { tmpIdxLocal = sharedTmpBuffer[0]; tmpValueLocal = tmpIdxLocal[topK]; @@ -62,47 +60,47 @@ class LITopk { outputValueLocal = nkValueLocal[64]; } - __aicore__ inline void operator()(LocalTensor& outputIdxLocal, - LocalTensor& inputLocal, + __aicore__ inline void operator()(LocalTensor &outputIdxLocal, LocalTensor &inputLocal, uint32_t s2SeqLen) { - topkb32::LiTopKVF(outputIdxLocal, // filter阶段使用输出value Buf topK * 4B + topkb32::LiTopKVF(outputIdxLocal, // filter阶段使用输出value Buf topK * 4B outputValueLocal, // filter阶段使用输出 Idx Buf topK * 4B - inputLocal, // 输入 s2SeqLen * 4B - tmpIdxLocal, // filter阶段使用暂存index Buf topK * 4B - tmpValueLocal, // filter阶段使用暂存value Buf topK * 4B - histogramsLocal, // 直方图的临时Buf 256 * 4B - idx0Local, // 输入数据第1个8位Buf 256 * 4B - idx1Local, // 输入数据第2个8位Buf 256 * 4B - idx2Local, // 输入数据第3个8位Buf 256 * 4B - idx3Local, // 输入数据第4个8位Buf 256 * 4B - nkValueLocal, // next_k 暂存Buf 64 * 4B - topK, // topk数量 - s2SeqLen); // 输入元素总数 + inputLocal, // 输入 s2SeqLen * 4B + tmpIdxLocal, // filter阶段使用暂存index Buf topK * 4B + tmpValueLocal, // filter阶段使用暂存value Buf topK * 4B + histogramsLocal, // 直方图的临时Buf 256 * 4B + idx0Local, // 输入数据第1个8位Buf 256 * 4B + idx1Local, // 输入数据第2个8位Buf 256 * 4B + idx2Local, // 输入数据第3个8位Buf 256 * 4B + idx3Local, // 输入数据第4个8位Buf 256 * 4B + nkValueLocal, // next_k 暂存Buf 64 * 4B + topK, // topk数量 + s2SeqLen); // 输入元素总数 } + private: - LocalTensor tmpIdxLocal; // filter阶段使用暂存index Buf topK * 4B - LocalTensor tmpValueLocal; // filter阶段使用暂存value Buf topK * 4B - LocalTensor histogramsLocal; // 直方图的临时Buf 256 * 4B - LocalTensor idx0Local; // 输入数据第1个8位Buf 256 * 4B - LocalTensor idx1Local; // 输入数据第2个8位Buf 256 * 4B - LocalTensor idx2Local; // 输入数据第3个8位Buf 256 * 4B - LocalTensor idx3Local; // 输入数据第4个8位Buf 256 * 4B - LocalTensor nkValueLocal; // next_k 暂存Buf 64 * 4B + LocalTensor tmpIdxLocal; // filter阶段使用暂存index Buf topK * 4B + LocalTensor tmpValueLocal; // filter阶段使用暂存value Buf topK * 4B + LocalTensor histogramsLocal; // 直方图的临时Buf 256 * 4B + LocalTensor idx0Local; // 输入数据第1个8位Buf 256 * 4B + LocalTensor idx1Local; // 输入数据第2个8位Buf 256 * 4B + LocalTensor idx2Local; // 输入数据第3个8位Buf 256 * 4B + LocalTensor idx3Local; // 输入数据第4个8位Buf 256 * 4B + LocalTensor nkValueLocal; // next_k 暂存Buf 64 * 4B LocalTensor outputValueLocal; // 输出value tensor uint32_t topK; }; -template<> +template <> class LITopk { public: __aicore__ inline uint32_t GetSharedTmpBufferSize() { // 2 * QLIV2Common::Align(topK, (uint32_t)256): 两块hisIndexLocal; // 3 * 256: histogramsLocal idxHighLocal idxLowLocal; 64: nkValueLocal - uint64_t bufferSize1 = (2 * QLIV2Common::Align(topK, (uint32_t)256) + 3 * 256 + 64) * sizeof(uint32_t); + uint64_t bufferSize1 = (2 * QLIV2Common::Align(topK, (uint32_t)256) + 3 * 256 + 64) * sizeof(uint32_t); // QLIV2Common::Align(topK, (uint32_t)256) + trunkLen:tmpIndexLocal - uint64_t bufferSize2 = (QLIV2Common::Align(topK, (uint32_t)256) + trunkLen) * sizeof(uint16_t); + uint64_t bufferSize2 = (QLIV2Common::Align(topK, (uint32_t)256) + trunkLen) * sizeof(uint16_t); uint64_t reuseBufferSize = QLIV2Common::Align(topK, (uint32_t)256) * sizeof(uint32_t); return bufferSize1 + bufferSize2 - reuseBufferSize; } @@ -110,10 +108,10 @@ class LITopk { __aicore__ inline void Init(uint32_t topK, uint32_t trunkLen) { this->topK = topK; - this->trunkLen = trunkLen; + this->trunkLen = trunkLen; } - __aicore__ inline void InitBuffers(LocalTensor& sharedTmpBuffer, LocalTensor& indicesOutLocal) + __aicore__ inline void InitBuffers(LocalTensor &sharedTmpBuffer, LocalTensor &indicesOutLocal) { LocalTensor hisIndexLocal1 = indicesOutLocal; LocalTensor hisIndexLocal2 = sharedTmpBuffer[0]; @@ -127,10 +125,9 @@ class LITopk { tmpIndexLocal = tmpIndexLocalTmp.template ReinterpretCast(); } - __aicore__ inline void TopK(LocalTensor& mrgValueLocal, LocalTensor& indicesOutLocal, - LocalTensor& hisValueLocal, uint32_t s2SeqLen, uint32_t loopIdx, - uint32_t s2LoopNum, bool isNeedLD, bool returnValueFlag, - uint32_t outputIdxOffset) + __aicore__ inline void TopK(LocalTensor &mrgValueLocal, LocalTensor &indicesOutLocal, + LocalTensor &hisValueLocal, uint32_t s2SeqLen, uint32_t loopIdx, + uint32_t s2LoopNum, bool isNeedLD, bool returnValueFlag, uint32_t outputIdxOffset) { // true: 开启返回hisValueLocal if (s2LoopNum == 1) { @@ -158,10 +155,9 @@ class LITopk { idxLowLocal, nkValueLocal, topK, s2SeqLen); PipeBarrier(); uint32_t curProcess = topK < trunkLen ? loopIdx * trunkLen - QLIV2Common::Align(topK, (uint32_t)256) : - (loopIdx - 1) * trunkLen; - topkb16gather::LiTopKGatherVF(hisIndexLocal[(loopIdx + 1) % 2], hisValueLocal, mrgValueLocal, - tmpIndexLocal, hisIndexLocal[loopIdx % 2], topK, - curProcess, s2SeqLen); + (loopIdx - 1) * trunkLen; + topkb16gather::LiTopKGatherVF(hisIndexLocal[(loopIdx + 1) % 2], hisValueLocal, mrgValueLocal, tmpIndexLocal, + hisIndexLocal[loopIdx % 2], topK, curProcess, s2SeqLen); if (loopIdx == s2LoopNum - 1) { PipeBarrier(); if ((loopIdx + 1) % 2 == 1) { // 2:pingpong @@ -173,7 +169,7 @@ class LITopk { if (loopIdx == 0 && isNeedLD) { topkb16gather::LiTopKVF(tmpIndexLocal, hisValueLocal, mrgValueLocal, histogramsLocal, idxHighLocal, - idxLowLocal, nkValueLocal, topK, s2SeqLen); + idxLowLocal, nkValueLocal, topK, s2SeqLen); PipeBarrier(); Cast(hisIndexLocal[(loopIdx + 1) % 2], tmpIndexLocal, RoundMode::CAST_NONE, topK); PipeBarrier(); @@ -184,10 +180,9 @@ class LITopk { idxLowLocal, nkValueLocal, topK, s2SeqLen); PipeBarrier(); uint32_t curProcess = topK < trunkLen ? loopIdx * trunkLen - QLIV2Common::Align(topK, (uint32_t)256) : - (loopIdx - 1) * trunkLen; - topkb16gather::LiTopKGatherVF(hisIndexLocal[(loopIdx + 1) % 2], hisValueLocal, mrgValueLocal, - tmpIndexLocal, hisIndexLocal[loopIdx % 2], - topK, curProcess, s2SeqLen); + (loopIdx - 1) * trunkLen; + topkb16gather::LiTopKGatherVF(hisIndexLocal[(loopIdx + 1) % 2], hisValueLocal, mrgValueLocal, tmpIndexLocal, + hisIndexLocal[loopIdx % 2], topK, curProcess, s2SeqLen); PipeBarrier(); AscendC::DataCopy(indicesOutLocal, hisIndexLocal[(loopIdx + 1) % 2], QLIV2Common::Align(topK, (uint32_t)256)); @@ -196,10 +191,9 @@ class LITopk { topkb16gather::IndicesAddOffset(indicesOutLocal, outputIdxOffset, topK); } } - __aicore__ inline void LdTopK(LocalTensor& mrgValueLocal, LocalTensor indexLocal, - LocalTensor& indicesOutLocal, - LocalTensor& hisValueLocal, uint32_t s2SeqLen, uint32_t loopIdx, - uint32_t s2LoopNum) + __aicore__ inline void LdTopK(LocalTensor &mrgValueLocal, LocalTensor indexLocal, + LocalTensor &indicesOutLocal, LocalTensor &hisValueLocal, + uint32_t s2SeqLen, uint32_t loopIdx, uint32_t s2LoopNum) { topkb16gather::LiTopKVF(tmpIndexLocal, hisValueLocal, mrgValueLocal, histogramsLocal, idxHighLocal, idxLowLocal, nkValueLocal, topK, s2SeqLen); @@ -208,14 +202,14 @@ class LITopk { } private: - LocalTensor hisIndexLocal[2]; // 每trunkLen长度的s2选出的topK个索引 - LocalTensor histogramsLocal; // 直方图的临时Buf 256 * 4B - LocalTensor idxHighLocal; // 输入数据高8位Buf 256 * 4B - LocalTensor idxLowLocal; // 输入数据低8位Buf 256 * 4B - LocalTensor nkValueLocal; // next_k 暂存Buf 64 * 4B - LocalTensor tmpIndexLocal; // 每trunkLen + topK的临时index + LocalTensor hisIndexLocal[2]; // 每trunkLen长度的s2选出的topK个索引 + LocalTensor histogramsLocal; // 直方图的临时Buf 256 * 4B + LocalTensor idxHighLocal; // 输入数据高8位Buf 256 * 4B + LocalTensor idxLowLocal; // 输入数据低8位Buf 256 * 4B + LocalTensor nkValueLocal; // next_k 暂存Buf 64 * 4B + LocalTensor tmpIndexLocal; // 每trunkLen + topK的临时index uint32_t topK = 512; uint32_t trunkLen = 16384; }; -} -#endif // QUANT_LIGHTNING_INDEXER_V2_TOPK_H +} // namespace topk +#endif // QUANT_LIGHTNING_INDEXER_V2_TOPK_H diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/quant_lightning_indexer_v2_vector1.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/quant_lightning_indexer_v2_vector1.h index 3d06a214dfca..08b7cb9e4b13 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/quant_lightning_indexer_v2_vector1.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/quant_lightning_indexer_v2_vector1.h @@ -25,19 +25,19 @@ namespace vector1 { __simd_vf__ void UIntToFloatReturnValueVF(__ubuf__ bfloat16_t *outBuf, __ubuf__ uint16_t *inBuf, uint16_t vfLoop) { - MicroAPI::RegTensor regIn; - MicroAPI::RegTensor regOut; - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regIn; + Reg::RegTensor regOut; + Reg::MaskReg maskAllB16 = Reg::CreateMask(); for (uint16_t i = 0; i < vfLoop; ++i) { - MicroAPI::LoadAlign(regIn, inBuf + i * 128); + Reg::LoadAlign(regIn, inBuf + i * 128); liV2Vector1::UIntSortConstCtx uint16Ctx; liV2Vector1::InitUIntSortConstCtx(uint16Ctx, maskAllB16); liV2Vector1::UIntToSortableKey(regOut, regIn, uint16Ctx, maskAllB16); - MicroAPI::StoreAlign(outBuf + i * 128, regOut, maskAllB16); + Reg::StoreAlign(outBuf + i * 128, regOut, maskAllB16); } } @@ -55,30 +55,29 @@ __aicore__ inline void UIntToFloatReturnValue(const LocalTensor &out __simd_vf__ void UIntToFloatReturnValueWithInfMaskVF(__ubuf__ bfloat16_t *valueOutBuf, __ubuf__ uint16_t *scoreOutBuf, uint16_t vfLoop, uint16_t negInfBits) { - MicroAPI::RegTensor regIn; - MicroAPI::RegTensor regOut; - MicroAPI::RegTensor regNegInf; - MicroAPI::RegTensor regZero; - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskInvalid; + Reg::RegTensor regIn; + Reg::RegTensor regOut; + Reg::RegTensor regNegInf; + Reg::RegTensor regZero; + Reg::MaskReg maskAllB16 = Reg::CreateMask(); + Reg::MaskReg maskInvalid; // 常量寄存器初始化:-inf(0xFF80) 与 0(用于识别无效位) - MicroAPI::Duplicate(regNegInf, negInfBits, maskAllB16); - MicroAPI::Duplicate(regZero, (uint16_t)0, maskAllB16); + Reg::Duplicate(regNegInf, negInfBits, maskAllB16); + Reg::Duplicate(regZero, (uint16_t)0, maskAllB16); liV2Vector1::UIntSortConstCtx uint16Ctx; liV2Vector1::InitUIntSortConstCtx(uint16Ctx, maskAllB16); for (uint16_t i = 0; i < vfLoop; ++i) { - MicroAPI::LoadAlign(regIn, scoreOutBuf + i * 128); + Reg::LoadAlign(regIn, scoreOutBuf + i * 128); // 比较得无效位掩码:score==0 即无效位(可排序键性质,真实值永不为0) - MicroAPI::Compare(maskInvalid, regIn, regZero, maskAllB16); + Reg::Compare(maskInvalid, regIn, regZero, maskAllB16); // 逆变换:可排序键 → bf16 值 liV2Vector1::UIntToSortableKey(regOut, regIn, uint16Ctx, maskAllB16); // 无效位覆盖为 -inf,有效位保留还原值 - MicroAPI::Select((MicroAPI::RegTensor &)regOut, regNegInf, (MicroAPI::RegTensor &)regOut, - maskInvalid); - MicroAPI::StoreAlign(valueOutBuf + i * 128, regOut, maskAllB16); + Reg::Select((Reg::RegTensor &)regOut, regNegInf, (Reg::RegTensor &)regOut, maskInvalid); + Reg::StoreAlign(valueOutBuf + i * 128, regOut, maskAllB16); } } @@ -93,93 +92,80 @@ __aicore__ inline void UIntToFloatReturnValueWithInfMask(const LocalTensor (®KScaleFP16)[2], - AscendC::MicroAPI::RegTensor (®KScale)[2], - AscendC::MicroAPI::MaskReg &maskAllB16, - __ubuf__ half *kScale_) +__simd_callee__ inline void LoadKScaleFP16(AscendC::Reg::RegTensor (®KScaleFP16)[2], + AscendC::Reg::RegTensor (®KScale)[2], + AscendC::Reg::MaskReg &maskAllB16, __ubuf__ half *kScale_) { - constexpr static MicroAPI::CastTrait castTraitFP16ToFP32 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - AscendC::MicroAPI::LoadAlign(regKScaleFP16[0], kScale_); - AscendC::MicroAPI::LoadAlign(regKScaleFP16[1], kScale_ + 64); - AscendC::MicroAPI::Cast(regKScale[0], regKScaleFP16[0], maskAllB16); - AscendC::MicroAPI::Cast(regKScale[1], regKScaleFP16[1], maskAllB16); + constexpr static Reg::CastTrait castTraitFP16ToFP32 = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + AscendC::Reg::LoadAlign(regKScaleFP16[0], kScale_); + AscendC::Reg::LoadAlign(regKScaleFP16[1], kScale_ + 64); + AscendC::Reg::Cast(regKScale[0], regKScaleFP16[0], maskAllB16); + AscendC::Reg::Cast(regKScale[1], regKScaleFP16[1], maskAllB16); } -__simd_callee__ inline void CastFP32ToFP16ToFP32(AscendC::MicroAPI::RegTensor (®QK0)[2], - AscendC::MicroAPI::RegTensor (®QK0Half)[2], - AscendC::MicroAPI::MaskReg &maskAllB32) +__simd_callee__ inline void CastFP32ToFP16ToFP32(AscendC::Reg::RegTensor (®QK0)[2], + AscendC::Reg::RegTensor (®QK0Half)[2], + AscendC::Reg::MaskReg &maskAllB32) { - AscendC::MicroAPI::MaskReg maskAllB16 = AscendC::MicroAPI::CreateMask(); - constexpr static MicroAPI::CastTrait castTraitFP32ToFP16 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; - constexpr static MicroAPI::CastTrait castTraitFP16ToFP32 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + AscendC::Reg::MaskReg maskAllB16 = AscendC::Reg::CreateMask(); + constexpr static Reg::CastTrait castTraitFP32ToFP16 = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; + constexpr static Reg::CastTrait castTraitFP16ToFP32 = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; float mulsScalar = 1.0 / 1024; - MicroAPI::Muls(regQK0[0], regQK0[0], mulsScalar, maskAllB32); - MicroAPI::Muls(regQK0[1], regQK0[1], mulsScalar, maskAllB32); + Reg::Muls(regQK0[0], regQK0[0], mulsScalar, maskAllB32); + Reg::Muls(regQK0[1], regQK0[1], mulsScalar, maskAllB32); - MicroAPI::Cast(regQK0Half[0], regQK0[0], maskAllB32); - MicroAPI::Cast(regQK0Half[1], regQK0[1], maskAllB32); + Reg::Cast(regQK0Half[0], regQK0[0], maskAllB32); + Reg::Cast(regQK0Half[1], regQK0[1], maskAllB32); - MicroAPI::Cast(regQK0[0], regQK0Half[0], maskAllB16); - MicroAPI::Cast(regQK0[1], regQK0Half[1], maskAllB16); + Reg::Cast(regQK0[0], regQK0Half[0], maskAllB16); + Reg::Cast(regQK0[1], regQK0Half[1], maskAllB16); } // int32 in uint16 out -__simd_vf__ void MulWeightAndReduceSumInt32GSizeOddVF(__ubuf__ uint16_t *out, __ubuf__ int32_t *qk, - uint32_t qkVLStride, __ubuf__ half *weight, - __ubuf__ half *kScale, __ubuf__ half *qScale, - uint16_t gSize) +__simd_vf__ void MulWeightAndReduceSumInt32GSizeOddVF(__ubuf__ uint16_t *out, __ubuf__ int32_t *qk, uint32_t qkVLStride, + __ubuf__ half *weight, __ubuf__ half *kScale, + __ubuf__ half *qScale, uint16_t gSize) { - MicroAPI::RegTensor regwBrc; - MicroAPI::RegTensor regQK[2]; - MicroAPI::RegTensor regQKHalf[2]; - MicroAPI::RegTensor regQKInt32[2]; - MicroAPI::RegTensor regW; - MicroAPI::RegTensor regWFP16; - MicroAPI::RegTensor regWFP16Temp; - MicroAPI::RegTensor regQScale; - MicroAPI::RegTensor regQScaleFP16; - MicroAPI::RegTensor regKScale[2]; - MicroAPI::RegTensor regKScaleFP16[2]; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regwBrc; + Reg::RegTensor regQK[2]; + Reg::RegTensor regQKHalf[2]; + Reg::RegTensor regQKInt32[2]; + Reg::RegTensor regW; + Reg::RegTensor regWFP16; + Reg::RegTensor regWFP16Temp; + Reg::RegTensor regQScale; + Reg::RegTensor regQScaleFP16; + Reg::RegTensor regKScale[2]; + Reg::RegTensor regKScaleFP16[2]; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::MERGING, - RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = {MicroAPI::RegLayout::ONE, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF16ToF32 = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::UNKNOWN}; - constexpr static MicroAPI::CastTrait castTraitInt32ToFP32 = {MicroAPI::RegLayout::UNKNOWN, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16 = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_RINT}; - MicroAPI::LoadAlign(regWFP16, weight); - MicroAPI::LoadAlign(regQScaleFP16, qScale); - MicroAPI::Cast(regW, regWFP16, maskAllB16); - MicroAPI::Cast(regQScale, regQScaleFP16, maskAllB16); - MicroAPI::Mul(regW, regW, regQScale, maskAllB32); - MicroAPI::Cast(regWFP16Temp, regW, maskAllB32); - MicroAPI::Cast(regW, regWFP16Temp, maskAllB16); + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF16ToF32 = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + constexpr static Reg::CastTrait castTraitInt32ToFP32 = {Reg::RegLayout::UNKNOWN, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16 = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; + Reg::LoadAlign(regWFP16, weight); + Reg::LoadAlign(regQScaleFP16, qScale); + Reg::Cast(regW, regWFP16, maskAllB16); + Reg::Cast(regQScale, regQScaleFP16, maskAllB16); + Reg::Mul(regW, regW, regQScale, maskAllB32); + Reg::Cast(regWFP16Temp, regW, maskAllB32); + Reg::Cast(regW, regWFP16Temp, maskAllB16); liV2Vector1::DuplicateZero(regSum0, maskAllB32); liV2Vector1::DuplicateZero(regSum1, maskAllB32); @@ -187,20 +173,20 @@ __simd_vf__ void MulWeightAndReduceSumInt32GSizeOddVF(__ubuf__ uint16_t *out, __ // float mulsScalar = 1.0f / 1024; // unroll2 for (uint16_t i = (uint16_t)(0); (uint16_t)(i + 1) < gSize; i += 2) { - MicroAPI::LoadAlign(regQKInt32[0], qk + 128 * i); - MicroAPI::LoadAlign(regQKInt32[1], qk + 128 * i + qkVLStride); - MicroAPI::Cast(regQK[0], regQKInt32[0], maskAllB32); - MicroAPI::Cast(regQK[1], regQKInt32[1], maskAllB32); + Reg::LoadAlign(regQKInt32[0], qk + 128 * i); + Reg::LoadAlign(regQKInt32[1], qk + 128 * i + qkVLStride); + Reg::Cast(regQK[0], regQKInt32[0], maskAllB32); + Reg::Cast(regQK[1], regQKInt32[1], maskAllB32); CastFP32ToFP16ToFP32(regQK, regQKHalf, maskAllB32); liV2Vector1::BroadcastLane(regwBrc, regW, i); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::LoadAlign(regQKInt32[0], qk + 128 * i + 128); - MicroAPI::LoadAlign(regQKInt32[1], qk + 128 * i + 128 + qkVLStride); - MicroAPI::Cast(regQK[0], regQKInt32[0], maskAllB32); - MicroAPI::Cast(regQK[1], regQKInt32[1], maskAllB32); + Reg::LoadAlign(regQKInt32[0], qk + 128 * i + 128); + Reg::LoadAlign(regQKInt32[1], qk + 128 * i + 128 + qkVLStride); + Reg::Cast(regQK[0], regQKInt32[0], maskAllB32); + Reg::Cast(regQK[1], regQKInt32[1], maskAllB32); CastFP32ToFP16ToFP32(regQK, regQKHalf, maskAllB32); @@ -208,156 +194,145 @@ __simd_vf__ void MulWeightAndReduceSumInt32GSizeOddVF(__ubuf__ uint16_t *out, __ liV2Vector1::WeightedAccum(regSum1, regQK, regwBrc, maskAllB32); } - MicroAPI::LoadAlign(regQKInt32[0], qk + 128 * (gSize - 1)); - MicroAPI::LoadAlign(regQKInt32[1], qk + 128 * (gSize - 1) + qkVLStride); - MicroAPI::Cast(regQK[0], regQKInt32[0], maskAllB32); - MicroAPI::Cast(regQK[1], regQKInt32[1], maskAllB32); + Reg::LoadAlign(regQKInt32[0], qk + 128 * (gSize - 1)); + Reg::LoadAlign(regQKInt32[1], qk + 128 * (gSize - 1) + qkVLStride); + Reg::Cast(regQK[0], regQKInt32[0], maskAllB32); + Reg::Cast(regQK[1], regQKInt32[1], maskAllB32); CastFP32ToFP16ToFP32(regQK, regQKHalf, maskAllB32); liV2Vector1::BroadcastLane(regwBrc, regW, gSize - 1); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); - MicroAPI::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); + Reg::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); + Reg::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); - MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); + Reg::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); + Reg::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); - MicroAPI::RegTensor regSumBF16; + Reg::RegTensor regSumBF16; // interleave cast ==> regSum[1] high regSum[0] low - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::Cast(regSumBF16, regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16, regSum0[0], maskAllB32); + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::Cast(regSumBF16, regSum0[1], maskAllB32); + Reg::Cast(regSumBF16, regSum0[0], maskAllB32); - MicroAPI::RegTensor regOut; + Reg::RegTensor regOut; liV2Vector1::FloatToSortableKey(regOut, regSumBF16, bf16Ctx, maskAllB16); // normal store - MicroAPI::StoreAlign(out, regOut, maskAllB16); + Reg::StoreAlign(out, regOut, maskAllB16); } // float in uint16 out __simd_vf__ void MulWeightAndReduceSumF32GSizeOddVF(__ubuf__ uint16_t *out, __ubuf__ float *qk, uint32_t qkVLStride, - __ubuf__ float *weight, __ubuf__ float *kScale, - __ubuf__ float *qScale, uint16_t gSize) + __ubuf__ float *weight, __ubuf__ float *kScale, + __ubuf__ float *qScale, uint16_t gSize) { - MicroAPI::RegTensor regwBrc; - MicroAPI::RegTensor regQK[2]; - MicroAPI::RegTensor regW; + Reg::RegTensor regwBrc; + Reg::RegTensor regQK[2]; + Reg::RegTensor regW; - MicroAPI::RegTensor regQScale; - MicroAPI::RegTensor regKScale[2]; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regQScale; + Reg::RegTensor regKScale[2]; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = { - MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = { - MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; - MicroAPI::LoadAlign(regW, weight); - MicroAPI::LoadAlign(regQScale, qScale); - MicroAPI::Mul(regW, regW, regQScale, maskAllB32); + Reg::LoadAlign(regW, weight); + Reg::LoadAlign(regQScale, qScale); + Reg::Mul(regW, regW, regQScale, maskAllB32); liV2Vector1::DuplicateZero(regSum0, maskAllB32); liV2Vector1::DuplicateZero(regSum1, maskAllB32); - MicroAPI::LoadAlign(regKScale[0], kScale); - MicroAPI::LoadAlign(regKScale[1], kScale + 64); + Reg::LoadAlign(regKScale[0], kScale); + Reg::LoadAlign(regKScale[1], kScale + 64); // unroll2 for (uint16_t i = (uint16_t)(0); (uint16_t)(i + 1) < gSize; i += 2) { - MicroAPI::LoadAlign(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上 - MicroAPI::LoadAlign(regQK[1], qk + 128 * i + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上 + Reg::LoadAlign(regQK[1], qk + 128 * i + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, i); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::LoadAlign(regQK[0], qk + 128 * i + 128); - MicroAPI::LoadAlign(regQK[1], qk + 128 * i + 128 + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * i + 128); + Reg::LoadAlign(regQK[1], qk + 128 * i + 128 + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, i + 1); liV2Vector1::WeightedAccum(regSum1, regQK, regwBrc, maskAllB32); } - MicroAPI::LoadAlign(regQK[0], qk + 128 * (gSize - 1)); // RowStride是128, 行都落在一个bank上 - MicroAPI::LoadAlign(regQK[1], qk + 128 * (gSize - 1) + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * (gSize - 1)); // RowStride是128, 行都落在一个bank上 + Reg::LoadAlign(regQK[1], qk + 128 * (gSize - 1) + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, (gSize - 1)); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); - MicroAPI::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); + Reg::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); + Reg::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); - MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); + Reg::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); + Reg::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); - MicroAPI::RegTensor regSumBF16; + Reg::RegTensor regSumBF16; // interleave cast ==> regSum[1] high regSum[0] low - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::Cast(regSumBF16, regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16, regSum0[0], maskAllB32); + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::Cast(regSumBF16, regSum0[1], maskAllB32); + Reg::Cast(regSumBF16, regSum0[0], maskAllB32); - MicroAPI::RegTensor regOut; + Reg::RegTensor regOut; liV2Vector1::FloatToSortableKey(regOut, regSumBF16, bf16Ctx, maskAllB16); // normal store - MicroAPI::StoreAlign(out, regOut, maskAllB16); + Reg::StoreAlign(out, regOut, maskAllB16); } // int32 in uint16 out __simd_vf__ void MulWeightAndReduceSumInt32GSizeEvenVF(__ubuf__ uint16_t *out, __ubuf__ int32_t *qk, uint32_t qkVLStride, __ubuf__ half *weight, - __ubuf__ half *kScale, __ubuf__ half *qScale, - uint16_t gSize) + __ubuf__ half *kScale, __ubuf__ half *qScale, uint16_t gSize) { - MicroAPI::RegTensor regwBrc; - MicroAPI::RegTensor regQK[2]; - MicroAPI::RegTensor regQKHalf[2]; - MicroAPI::RegTensor regQKInt32[2]; - MicroAPI::RegTensor regW; - MicroAPI::RegTensor regWFP16; - MicroAPI::RegTensor regWFP16Temp; - MicroAPI::RegTensor regQScale; - MicroAPI::RegTensor regQScaleFP16; - MicroAPI::RegTensor regKScale[2]; - MicroAPI::RegTensor regKScaleFP16[2]; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regwBrc; + Reg::RegTensor regQK[2]; + Reg::RegTensor regQKHalf[2]; + Reg::RegTensor regQKInt32[2]; + Reg::RegTensor regW; + Reg::RegTensor regWFP16; + Reg::RegTensor regWFP16Temp; + Reg::RegTensor regQScale; + Reg::RegTensor regQScaleFP16; + Reg::RegTensor regKScale[2]; + Reg::RegTensor regKScaleFP16[2]; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::MERGING, - RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = {MicroAPI::RegLayout::ONE, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF16ToF32 = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::UNKNOWN}; - constexpr static MicroAPI::CastTrait castTraitInt32ToFP32 = {MicroAPI::RegLayout::UNKNOWN, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16 = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_RINT}; - MicroAPI::LoadAlign(regWFP16, weight); - MicroAPI::LoadAlign(regQScaleFP16, qScale); - MicroAPI::Cast(regW, regWFP16, maskAllB16); - MicroAPI::Cast(regQScale, regQScaleFP16, maskAllB16); - MicroAPI::Mul(regW, regW, regQScale, maskAllB32); - MicroAPI::Cast(regWFP16Temp, regW, maskAllB32); - MicroAPI::Cast(regW, regWFP16Temp, maskAllB16); + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF16ToF32 = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + constexpr static Reg::CastTrait castTraitInt32ToFP32 = {Reg::RegLayout::UNKNOWN, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16 = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; + Reg::LoadAlign(regWFP16, weight); + Reg::LoadAlign(regQScaleFP16, qScale); + Reg::Cast(regW, regWFP16, maskAllB16); + Reg::Cast(regQScale, regQScaleFP16, maskAllB16); + Reg::Mul(regW, regW, regQScale, maskAllB32); + Reg::Cast(regWFP16Temp, regW, maskAllB32); + Reg::Cast(regW, regWFP16Temp, maskAllB16); liV2Vector1::DuplicateZero(regSum0, maskAllB32); liV2Vector1::DuplicateZero(regSum1, maskAllB32); @@ -365,20 +340,20 @@ __simd_vf__ void MulWeightAndReduceSumInt32GSizeEvenVF(__ubuf__ uint16_t *out, _ // float mulsScalar = 1.0f / 1024; // unroll2 for (uint16_t i = (uint16_t)(0); i < gSize; i += 2) { - MicroAPI::LoadAlign(regQKInt32[0], qk + 128 * i); - MicroAPI::LoadAlign(regQKInt32[1], qk + 128 * i + qkVLStride); - MicroAPI::Cast(regQK[0], regQKInt32[0], maskAllB32); - MicroAPI::Cast(regQK[1], regQKInt32[1], maskAllB32); + Reg::LoadAlign(regQKInt32[0], qk + 128 * i); + Reg::LoadAlign(regQKInt32[1], qk + 128 * i + qkVLStride); + Reg::Cast(regQK[0], regQKInt32[0], maskAllB32); + Reg::Cast(regQK[1], regQKInt32[1], maskAllB32); CastFP32ToFP16ToFP32(regQK, regQKHalf, maskAllB32); liV2Vector1::BroadcastLane(regwBrc, regW, i); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::LoadAlign(regQKInt32[0], qk + 128 * i + 128); - MicroAPI::LoadAlign(regQKInt32[1], qk + 128 * i + 128 + qkVLStride); - MicroAPI::Cast(regQK[0], regQKInt32[0], maskAllB32); - MicroAPI::Cast(regQK[1], regQKInt32[1], maskAllB32); + Reg::LoadAlign(regQKInt32[0], qk + 128 * i + 128); + Reg::LoadAlign(regQKInt32[1], qk + 128 * i + 128 + qkVLStride); + Reg::Cast(regQK[0], regQKInt32[0], maskAllB32); + Reg::Cast(regQK[1], regQKInt32[1], maskAllB32); CastFP32ToFP16ToFP32(regQK, regQKHalf, maskAllB32); @@ -386,86 +361,86 @@ __simd_vf__ void MulWeightAndReduceSumInt32GSizeEvenVF(__ubuf__ uint16_t *out, _ liV2Vector1::WeightedAccum(regSum1, regQK, regwBrc, maskAllB32); } - MicroAPI::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); - MicroAPI::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); + Reg::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); + Reg::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); - MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); + Reg::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); + Reg::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); - MicroAPI::RegTensor regSumBF16; + Reg::RegTensor regSumBF16; // interleave cast ==> regSum[1] high regSum[0] low - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::Cast(regSumBF16, regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16, regSum0[0], maskAllB32); + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::Cast(regSumBF16, regSum0[1], maskAllB32); + Reg::Cast(regSumBF16, regSum0[0], maskAllB32); - MicroAPI::RegTensor regOut; + Reg::RegTensor regOut; liV2Vector1::FloatToSortableKey(regOut, regSumBF16, bf16Ctx, maskAllB16); // normal store - MicroAPI::StoreAlign(out, regOut, maskAllB16); + Reg::StoreAlign(out, regOut, maskAllB16); } __simd_vf__ void MulWeightAndReduceSumF32GSizeEvenVF(__ubuf__ uint16_t *out, __ubuf__ float *qk, uint32_t qkVLStride, __ubuf__ float *weight, __ubuf__ float *kScale, __ubuf__ float *qScale, uint16_t gSize) { - MicroAPI::RegTensor regwBrc; - MicroAPI::RegTensor regQK[2]; - MicroAPI::RegTensor regW; + Reg::RegTensor regwBrc; + Reg::RegTensor regQK[2]; + Reg::RegTensor regW; - MicroAPI::RegTensor regQScale; - MicroAPI::RegTensor regKScale[2]; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regQScale; + Reg::RegTensor regKScale[2]; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = { - MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = { - MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; - MicroAPI::LoadAlign(regW, weight); - MicroAPI::LoadAlign(regQScale, qScale); - MicroAPI::Mul(regW, regW, regQScale, maskAllB32); + Reg::LoadAlign(regW, weight); + Reg::LoadAlign(regQScale, qScale); + Reg::Mul(regW, regW, regQScale, maskAllB32); liV2Vector1::DuplicateZero(regSum0, maskAllB32); liV2Vector1::DuplicateZero(regSum1, maskAllB32); - MicroAPI::LoadAlign(regKScale[0], kScale); - MicroAPI::LoadAlign(regKScale[1], kScale + 64); + Reg::LoadAlign(regKScale[0], kScale); + Reg::LoadAlign(regKScale[1], kScale + 64); // unroll2 for (uint16_t i = (uint16_t)(0); i < gSize; i += 2) { - MicroAPI::LoadAlign(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上 - MicroAPI::LoadAlign(regQK[1], qk + 128 * i + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上 + Reg::LoadAlign(regQK[1], qk + 128 * i + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, i); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::LoadAlign(regQK[0], qk + 128 * i + 128); - MicroAPI::LoadAlign(regQK[1], qk + 128 * i + 128 + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * i + 128); + Reg::LoadAlign(regQK[1], qk + 128 * i + 128 + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, i + 1); liV2Vector1::WeightedAccum(regSum1, regQK, regwBrc, maskAllB32); } - MicroAPI::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); - MicroAPI::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); + Reg::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); + Reg::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); - MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); + Reg::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); + Reg::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); - MicroAPI::RegTensor regSumBF16; + Reg::RegTensor regSumBF16; // interleave cast ==> regSum[1] high regSum[0] low - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::Cast(regSumBF16, regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16, regSum0[0], maskAllB32); + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::Cast(regSumBF16, regSum0[1], maskAllB32); + Reg::Cast(regSumBF16, regSum0[0], maskAllB32); - MicroAPI::RegTensor regOut; + Reg::RegTensor regOut; liV2Vector1::FloatToSortableKey(regOut, regSumBF16, bf16Ctx, maskAllB16); // normal store - MicroAPI::StoreAlign(out, regOut, maskAllB16); + Reg::StoreAlign(out, regOut, maskAllB16); } __aicore__ inline void MulWeightAndReduceSum(const LocalTensor &out_, // out [S2Base] [128 ] @@ -512,67 +487,67 @@ __aicore__ inline void MulWeightAndReduceSum(const LocalTensor &out_, __simd_vf__ void MulWeightAndReduceSumB16VF(__ubuf__ uint16_t *out, __ubuf__ bfloat16_t *qk, __ubuf__ float *weight, __ubuf__ float *kScale, __ubuf__ float *qScale, uint16_t gSize) { - MicroAPI::RegTensor regQK[4]; - MicroAPI::RegTensor regQKB16[2]; - MicroAPI::RegTensor regW; - MicroAPI::RegTensor regwBrc[2]; - MicroAPI::RegTensor regQScale; - MicroAPI::RegTensor regKScale[2]; - MicroAPI::RegTensor regSum[2]; + Reg::RegTensor regQK[4]; + Reg::RegTensor regQKB16[2]; + Reg::RegTensor regW; + Reg::RegTensor regwBrc[2]; + Reg::RegTensor regQScale; + Reg::RegTensor regKScale[2]; + Reg::RegTensor regSum[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); - MicroAPI::RegTensor regSumBF16; + Reg::RegTensor regSumBF16; liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - using CastTrait = MicroAPI::CastTrait; - static constexpr CastTrait castTraitB162B32_EVEN = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - static constexpr CastTrait castTraitB162B32_ODD = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + using CastTrait = Reg::CastTrait; + static constexpr CastTrait castTraitB162B32_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr CastTrait castTraitB162B32_ODD = {Reg::RegLayout::ONE, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - constexpr static CastTrait castTraitF32ToF16_EVEN = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; - constexpr static CastTrait castTraitF32ToF16_ODD = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; - MicroAPI::LoadAlign(regW, weight); - MicroAPI::LoadAlign(regQScale, qScale); - MicroAPI::Mul(regW, regW, regQScale, maskAllB32); - MicroAPI::StoreAlign(weight, regW, maskAllB32); - MicroAPI::LocalMemBar(); + Reg::LoadAlign(regW, weight); + Reg::LoadAlign(regQScale, qScale); + Reg::Mul(regW, regW, regQScale, maskAllB32); + Reg::StoreAlign(weight, regW, maskAllB32); + Reg::LocalMemBar(); liV2Vector1::DuplicateZero(regSum, maskAllB32); // interleave load - MicroAPI::LoadAlign(regKScale[0], regKScale[1], kScale); + Reg::LoadAlign(regKScale[0], regKScale[1], kScale); // Duplicate + Gather方法劣化 // Relu在cube随路做 for (uint16_t i = (uint16_t)(0); i < gSize; i++) { // RowStride是256, 行都落在一个bank上 - MicroAPI::LoadAlign(regQKB16[0], qk + 256 * i); - MicroAPI::LoadAlign(regwBrc[0], weight + i); + Reg::LoadAlign(regQKB16[0], qk + 256 * i); + Reg::LoadAlign(regwBrc[0], weight + i); // interleave cast - MicroAPI::Cast(regQK[0], regQKB16[0], maskAllB16); - MicroAPI::Cast(regQK[1], regQKB16[0], maskAllB16); - MicroAPI::MulAddDst(regSum[0], regQK[0], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum[1], regQK[1], regwBrc[0], maskAllB32); + Reg::Cast(regQK[0], regQKB16[0], maskAllB16); + Reg::Cast(regQK[1], regQKB16[0], maskAllB16); + Reg::MulAddDst(regSum[0], regQK[0], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum[1], regQK[1], regwBrc[0], maskAllB32); } - MicroAPI::Mul(regSum[0], regSum[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum[1], regSum[1], regKScale[1], maskAllB32); + Reg::Mul(regSum[0], regSum[0], regKScale[0], maskAllB32); + Reg::Mul(regSum[1], regSum[1], regKScale[1], maskAllB32); // interleave cast back - MicroAPI::Cast(regSumBF16, regSum[1], maskAllB32); - MicroAPI::Cast(regSumBF16, regSum[0], maskAllB32); + Reg::Cast(regSumBF16, regSum[1], maskAllB32); + Reg::Cast(regSumBF16, regSum[0], maskAllB32); - MicroAPI::RegTensor regOut; + Reg::RegTensor regOut; liV2Vector1::FloatToSortableKey(regOut, regSumBF16, bf16Ctx, maskAllB16); // norm load - MicroAPI::StoreAlign(out, regOut, maskAllB16); + Reg::StoreAlign(out, regOut, maskAllB16); } __aicore__ inline void MulWeightAndReduceSum(const LocalTensor &out_, // out [S2Base] [128 ] @@ -599,153 +574,142 @@ __simd_vf__ void MulWeightAndReduceSum2F32VF(__ubuf__ uint16_t *out0, __ubuf__ u __ubuf__ float *qScale0, __ubuf__ float *qScale1, __ubuf__ float *kScale0, uint16_t gSize) { - MicroAPI::RegTensor regwBrc[2]; - MicroAPI::RegTensor regQK0[2]; - MicroAPI::RegTensor regQK1[2]; - MicroAPI::RegTensor regW[2]; - - MicroAPI::RegTensor regQScale[2]; - MicroAPI::RegTensor regKScale[2]; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regwBrc[2]; + Reg::RegTensor regQK0[2]; + Reg::RegTensor regQK1[2]; + Reg::RegTensor regW[2]; + + Reg::RegTensor regQScale[2]; + Reg::RegTensor regKScale[2]; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = { - MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = { - MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; - - MicroAPI::LoadAlign(regW[0], weight0); - MicroAPI::LoadAlign(regW[1], weight1); - MicroAPI::LoadAlign(regQScale[0], qScale0); - MicroAPI::LoadAlign(regQScale[1], qScale1); - MicroAPI::Mul(regW[0], regW[0], regQScale[0], maskAllB32); - MicroAPI::Mul(regW[1], regW[1], regQScale[1], maskAllB32); + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + + Reg::LoadAlign(regW[0], weight0); + Reg::LoadAlign(regW[1], weight1); + Reg::LoadAlign(regQScale[0], qScale0); + Reg::LoadAlign(regQScale[1], qScale1); + Reg::Mul(regW[0], regW[0], regQScale[0], maskAllB32); + Reg::Mul(regW[1], regW[1], regQScale[1], maskAllB32); // regW[0]与weight1混合使用 - MicroAPI::StoreAlign(weightTemp, regW[1], maskAllB32); - MicroAPI::LocalMemBar(); + Reg::StoreAlign(weightTemp, regW[1], maskAllB32); + Reg::LocalMemBar(); liV2Vector1::DuplicateZero(regSum0, maskAllB32); liV2Vector1::DuplicateZero(regSum1, maskAllB32); - MicroAPI::LoadAlign(regKScale[0], kScale0); - MicroAPI::LoadAlign(regKScale[1], kScale0 + 64); + Reg::LoadAlign(regKScale[0], kScale0); + Reg::LoadAlign(regKScale[1], kScale0 + 64); for (uint16_t i = (uint16_t)(0); i < gSize; i++) { - MicroAPI::LoadAlign(regQK0[0], qk0 + 128 * i); - MicroAPI::LoadAlign(regQK0[1], qk0 + 128 * i + qkVLStride); - MicroAPI::LoadAlign(regQK1[0], qk1 + 128 * i); - MicroAPI::LoadAlign(regQK1[1], qk1 + 128 * i + qkVLStride); + Reg::LoadAlign(regQK0[0], qk0 + 128 * i); + Reg::LoadAlign(regQK0[1], qk0 + 128 * i + qkVLStride); + Reg::LoadAlign(regQK1[0], qk1 + 128 * i); + Reg::LoadAlign(regQK1[1], qk1 + 128 * i + qkVLStride); // 混合使用对整体性能更好 liV2Vector1::BroadcastLane(regwBrc[0], regW[0], i); // Weight无bank冲突,用LoadAlign来提取weight标量 // 地址空间处理:原 BroadcastLane(ptr) 内联为 LoadAlign BRC,避免 __ubuf__ 传给 __local_mem__ 参数 - MicroAPI::LoadAlign(regwBrc[1], weightTemp + i); - MicroAPI::Relu(regQK0[0], regQK0[0], maskAllB32); - MicroAPI::Relu(regQK0[1], regQK0[1], maskAllB32); - MicroAPI::Relu(regQK1[0], regQK1[0], maskAllB32); - MicroAPI::Relu(regQK1[1], regQK1[1], maskAllB32); - MicroAPI::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32); - MicroAPI::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32); + Reg::LoadAlign(regwBrc[1], weightTemp + i); + Reg::Relu(regQK0[0], regQK0[0], maskAllB32); + Reg::Relu(regQK0[1], regQK0[1], maskAllB32); + Reg::Relu(regQK1[0], regQK1[0], maskAllB32); + Reg::Relu(regQK1[1], regQK1[1], maskAllB32); + Reg::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32); + Reg::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32); } // Apply kScale scaling - MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); - MicroAPI::Mul(regSum1[0], regSum1[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum1[1], regSum1[1], regKScale[1], maskAllB32); + Reg::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); + Reg::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); + Reg::Mul(regSum1[0], regSum1[0], regKScale[0], maskAllB32); + Reg::Mul(regSum1[1], regSum1[1], regKScale[1], maskAllB32); // Convert to bfloat16 and store output channel - MicroAPI::RegTensor regSumBF16[2]; - MicroAPI::RegTensor regOut[2]; - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::DeInterleave(regSum1[0], regSum1[1], regSum1[0], regSum1[1]); - MicroAPI::Cast(regSumBF16[0], regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16[1], regSum1[1], maskAllB32); - MicroAPI::Cast(regSumBF16[0], regSum0[0], maskAllB32); - MicroAPI::Cast(regSumBF16[1], regSum1[0], maskAllB32); + Reg::RegTensor regSumBF16[2]; + Reg::RegTensor regOut[2]; + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::DeInterleave(regSum1[0], regSum1[1], regSum1[0], regSum1[1]); + Reg::Cast(regSumBF16[0], regSum0[1], maskAllB32); + Reg::Cast(regSumBF16[1], regSum1[1], maskAllB32); + Reg::Cast(regSumBF16[0], regSum0[0], maskAllB32); + Reg::Cast(regSumBF16[1], regSum1[0], maskAllB32); liV2Vector1::FloatX2ToSortableKey(regOut[0], regOut[1], regSumBF16[0], regSumBF16[1], bf16Ctx, maskAllB16); - MicroAPI::StoreAlign(out0, regOut[0], maskAllB16); - MicroAPI::StoreAlign(out1, regOut[1], maskAllB16); + Reg::StoreAlign(out0, regOut[0], maskAllB16); + Reg::StoreAlign(out1, regOut[1], maskAllB16); } // 计算S1=2 // int32 in uint16 out -__simd_vf__ void MulWeightAndReduceSum2Int32VF(__ubuf__ uint16_t *out0, __ubuf__ uint16_t *out1, - __ubuf__ int32_t *qk0, __ubuf__ int32_t *qk1, - uint32_t qkVLStride, __ubuf__ half *weight0, +__simd_vf__ void MulWeightAndReduceSum2Int32VF(__ubuf__ uint16_t *out0, __ubuf__ uint16_t *out1, __ubuf__ int32_t *qk0, + __ubuf__ int32_t *qk1, uint32_t qkVLStride, __ubuf__ half *weight0, __ubuf__ half *weight1, __ubuf__ float *weightTemp, - __ubuf__ half *qScale0, __ubuf__ half *qScale1, - __ubuf__ half *kScale0, uint16_t gSize) + __ubuf__ half *qScale0, __ubuf__ half *qScale1, __ubuf__ half *kScale0, + uint16_t gSize) { - MicroAPI::RegTensor regwBrc[2]; - MicroAPI::RegTensor regQK0[2]; - MicroAPI::RegTensor regQK1[2]; - MicroAPI::RegTensor regQK0Half[2]; - MicroAPI::RegTensor regQK1Half[2]; - MicroAPI::RegTensor regQK0Int32[2]; - MicroAPI::RegTensor regQK1Int32[2]; - MicroAPI::RegTensor regW[2]; - MicroAPI::RegTensor regWFP16[2]; - MicroAPI::RegTensor regWFP16Temp[2]; - MicroAPI::RegTensor regQScale[2]; - MicroAPI::RegTensor regQScaleFP16[2]; - MicroAPI::RegTensor regKScale[2]; - MicroAPI::RegTensor regKScaleFP16[2]; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regwBrc[2]; + Reg::RegTensor regQK0[2]; + Reg::RegTensor regQK1[2]; + Reg::RegTensor regQK0Half[2]; + Reg::RegTensor regQK1Half[2]; + Reg::RegTensor regQK0Int32[2]; + Reg::RegTensor regQK1Int32[2]; + Reg::RegTensor regW[2]; + Reg::RegTensor regWFP16[2]; + Reg::RegTensor regWFP16Temp[2]; + Reg::RegTensor regQScale[2]; + Reg::RegTensor regQScaleFP16[2]; + Reg::RegTensor regKScale[2]; + Reg::RegTensor regKScaleFP16[2]; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::MERGING, - RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = {MicroAPI::RegLayout::ONE, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF16ToF32 = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::UNKNOWN}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16 = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_RINT}; - constexpr static MicroAPI::CastTrait castTraitInt32ToFP32 = {MicroAPI::RegLayout::UNKNOWN, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_ROUND}; - MicroAPI::LoadAlign(regWFP16[0], weight0); - MicroAPI::LoadAlign(regWFP16[1], weight1); - MicroAPI::LoadAlign(regQScaleFP16[0], qScale0); - MicroAPI::LoadAlign(regQScaleFP16[1], qScale1); - MicroAPI::Cast(regW[0], regWFP16[0], maskAllB16); - MicroAPI::Cast(regW[1], regWFP16[1], maskAllB16); - MicroAPI::Cast(regQScale[0], regQScaleFP16[0], maskAllB16); - MicroAPI::Cast(regQScale[1], regQScaleFP16[1], maskAllB16); - - MicroAPI::Mul(regW[0], regW[0], regQScale[0], maskAllB32); - MicroAPI::Mul(regW[1], regW[1], regQScale[1], maskAllB32); - - MicroAPI::Cast(regWFP16Temp[0], regW[0], maskAllB32); - MicroAPI::Cast(regW[0], regWFP16Temp[0], maskAllB16); - MicroAPI::Cast(regWFP16Temp[1], regW[1], maskAllB32); - MicroAPI::Cast(regW[1], regWFP16Temp[1], maskAllB16); + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF16ToF32 = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + constexpr static Reg::CastTrait castTraitF32ToF16 = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; + constexpr static Reg::CastTrait castTraitInt32ToFP32 = {Reg::RegLayout::UNKNOWN, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + Reg::LoadAlign(regWFP16[0], weight0); + Reg::LoadAlign(regWFP16[1], weight1); + Reg::LoadAlign(regQScaleFP16[0], qScale0); + Reg::LoadAlign(regQScaleFP16[1], qScale1); + Reg::Cast(regW[0], regWFP16[0], maskAllB16); + Reg::Cast(regW[1], regWFP16[1], maskAllB16); + Reg::Cast(regQScale[0], regQScaleFP16[0], maskAllB16); + Reg::Cast(regQScale[1], regQScaleFP16[1], maskAllB16); + + Reg::Mul(regW[0], regW[0], regQScale[0], maskAllB32); + Reg::Mul(regW[1], regW[1], regQScale[1], maskAllB32); + + Reg::Cast(regWFP16Temp[0], regW[0], maskAllB32); + Reg::Cast(regW[0], regWFP16Temp[0], maskAllB16); + Reg::Cast(regWFP16Temp[1], regW[1], maskAllB32); + Reg::Cast(regW[1], regWFP16Temp[1], maskAllB16); // regW[0]与weight1混合使用 - MicroAPI::StoreAlign(weightTemp, regW[1], maskAllB32); - MicroAPI::LocalMemBar(); + Reg::StoreAlign(weightTemp, regW[1], maskAllB32); + Reg::LocalMemBar(); liV2Vector1::DuplicateZero(regSum0, maskAllB32); liV2Vector1::DuplicateZero(regSum1, maskAllB32); @@ -753,53 +717,53 @@ __simd_vf__ void MulWeightAndReduceSum2Int32VF(__ubuf__ uint16_t *out0, __ubuf__ // float mulsScalar = 1.0 / 1024; for (uint16_t i = (uint16_t)(0); i < gSize; i++) { - MicroAPI::LoadAlign(regQK0Int32[0], qk0 + 128 * i); - MicroAPI::Cast(regQK0[0], regQK0Int32[0], maskAllB32); - MicroAPI::LoadAlign(regQK0Int32[1], qk0 + 128 * i + qkVLStride); - MicroAPI::Cast(regQK0[1], regQK0Int32[1], maskAllB32); - MicroAPI::LoadAlign(regQK1Int32[0], qk1 + 128 * i); - MicroAPI::Cast(regQK1[0], regQK1Int32[0], maskAllB32); - MicroAPI::LoadAlign(regQK1Int32[1], qk1 + 128 * i + qkVLStride); - MicroAPI::Cast(regQK1[1], regQK1Int32[1], maskAllB32); + Reg::LoadAlign(regQK0Int32[0], qk0 + 128 * i); + Reg::Cast(regQK0[0], regQK0Int32[0], maskAllB32); + Reg::LoadAlign(regQK0Int32[1], qk0 + 128 * i + qkVLStride); + Reg::Cast(regQK0[1], regQK0Int32[1], maskAllB32); + Reg::LoadAlign(regQK1Int32[0], qk1 + 128 * i); + Reg::Cast(regQK1[0], regQK1Int32[0], maskAllB32); + Reg::LoadAlign(regQK1Int32[1], qk1 + 128 * i + qkVLStride); + Reg::Cast(regQK1[1], regQK1Int32[1], maskAllB32); // 混合使用对整体性能更好 liV2Vector1::BroadcastLane(regwBrc[0], regW[0], i); // Weight无bank冲突,用LoadAlign来提取weight标量 // 地址空间处理:原 BroadcastLane(ptr) 内联为 LoadAlign BRC,避免 __ubuf__ 传给 __local_mem__ 参数 - MicroAPI::LoadAlign(regwBrc[1], weightTemp + i); - MicroAPI::Relu(regQK0[0], regQK0[0], maskAllB32); - MicroAPI::Relu(regQK0[1], regQK0[1], maskAllB32); - MicroAPI::Relu(regQK1[0], regQK1[0], maskAllB32); - MicroAPI::Relu(regQK1[1], regQK1[1], maskAllB32); + Reg::LoadAlign(regwBrc[1], weightTemp + i); + Reg::Relu(regQK0[0], regQK0[0], maskAllB32); + Reg::Relu(regQK0[1], regQK0[1], maskAllB32); + Reg::Relu(regQK1[0], regQK1[0], maskAllB32); + Reg::Relu(regQK1[1], regQK1[1], maskAllB32); CastFP32ToFP16ToFP32(regQK0, regQK0Half, maskAllB32); CastFP32ToFP16ToFP32(regQK1, regQK1Half, maskAllB32); - MicroAPI::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32); - MicroAPI::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32); + Reg::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32); + Reg::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32); } // Apply kScale scaling - MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); - MicroAPI::Mul(regSum1[0], regSum1[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum1[1], regSum1[1], regKScale[1], maskAllB32); + Reg::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); + Reg::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); + Reg::Mul(regSum1[0], regSum1[0], regKScale[0], maskAllB32); + Reg::Mul(regSum1[1], regSum1[1], regKScale[1], maskAllB32); // Convert to bfloat16 and store output channel - MicroAPI::RegTensor regSumBF16[2]; - MicroAPI::RegTensor regOut[2]; - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::DeInterleave(regSum1[0], regSum1[1], regSum1[0], regSum1[1]); - MicroAPI::Cast(regSumBF16[0], regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16[1], regSum1[1], maskAllB32); - MicroAPI::Cast(regSumBF16[0], regSum0[0], maskAllB32); - MicroAPI::Cast(regSumBF16[1], regSum1[0], maskAllB32); - - liV2Vector1::FloatX2ToSortableKey(regOut[0], regOut[1], - regSumBF16[0], regSumBF16[1], bf16Ctx, maskAllB16); - MicroAPI::StoreAlign(out0, regOut[0], maskAllB16); - MicroAPI::StoreAlign(out1, regOut[1], maskAllB16); + Reg::RegTensor regSumBF16[2]; + Reg::RegTensor regOut[2]; + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::DeInterleave(regSum1[0], regSum1[1], regSum1[0], regSum1[1]); + Reg::Cast(regSumBF16[0], regSum0[1], maskAllB32); + Reg::Cast(regSumBF16[1], regSum1[1], maskAllB32); + Reg::Cast(regSumBF16[0], regSum0[0], maskAllB32); + Reg::Cast(regSumBF16[1], regSum1[0], maskAllB32); + + liV2Vector1::FloatX2ToSortableKey(regOut[0], regOut[1], regSumBF16[0], regSumBF16[1], bf16Ctx, + maskAllB16); + Reg::StoreAlign(out0, regOut[0], maskAllB16); + Reg::StoreAlign(out1, regOut[1], maskAllB16); } __aicore__ inline void MulWeightAndReduceSum2(const LocalTensor &out_, // out [2, S2Base] [128 ] @@ -834,11 +798,9 @@ __aicore__ inline void MulWeightAndReduceSum2(const LocalTensor &out_, __aicore__ inline void MulWeightAndReduceSum2(const LocalTensor &out_, // out [2, S2Base] [128 ] uint32_t outStride, const LocalTensor &qk_, // q*k^t [2, G, S2Base] [64 128] - uint32_t qkVLStride, - uint32_t qkStride, + uint32_t qkVLStride, uint32_t qkStride, const LocalTensor &weight_, // w [2, G] [64 ] - uint32_t weightStride, - const LocalTensor &weightTemp_, + uint32_t weightStride, const LocalTensor &weightTemp_, const LocalTensor &kScale_, // kScale [S2Base] [128 ] uint32_t kScaleStride, const LocalTensor &qScale_, // qScale [2, G] [64 ] @@ -858,98 +820,96 @@ __aicore__ inline void MulWeightAndReduceSum2(const LocalTensor &out_, // kScaleStride is zero __ubuf__ uint16_t *out1 = out0 + outStride; - MulWeightAndReduceSum2Int32VF(out0, out1, qk0, qk1, qkVLStride, weight0, weight1, weightTemp, - qScale0, qScale1, kScale0, (uint16_t)gSize); + MulWeightAndReduceSum2Int32VF(out0, out1, qk0, qk1, qkVLStride, weight0, weight1, weightTemp, qScale0, qScale1, + kScale0, (uint16_t)gSize); } // 计算S1=2 // bfloat16 in uint16 out -__simd_vf__ void MulWeightAndReduceSum2B16VF(__ubuf__ uint16_t *out0, __ubuf__ uint16_t *out1, - __ubuf__ bfloat16_t *qk0, - __ubuf__ bfloat16_t *qk1, __ubuf__ float *weight0, - __ubuf__ float *weight1, +__simd_vf__ void MulWeightAndReduceSum2B16VF(__ubuf__ uint16_t *out0, __ubuf__ uint16_t *out1, __ubuf__ bfloat16_t *qk0, + __ubuf__ bfloat16_t *qk1, __ubuf__ float *weight0, __ubuf__ float *weight1, __ubuf__ float *weightTemp0, __ubuf__ float *weightTemp1, - __ubuf__ float *qScale0, __ubuf__ float *qScale1, - __ubuf__ float *kScale0, uint16_t gSize) + __ubuf__ float *qScale0, __ubuf__ float *qScale1, __ubuf__ float *kScale0, + uint16_t gSize) { - MicroAPI::RegTensor regwBrc[2]; - MicroAPI::RegTensor regQK0[2]; - MicroAPI::RegTensor regQK1[2]; - MicroAPI::RegTensor regW[2]; - MicroAPI::RegTensor regQKB16[2]; - - MicroAPI::RegTensor regQScale[2]; - MicroAPI::RegTensor regKScale[2]; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regwBrc[2]; + Reg::RegTensor regQK0[2]; + Reg::RegTensor regQK1[2]; + Reg::RegTensor regW[2]; + Reg::RegTensor regQKB16[2]; + + Reg::RegTensor regQScale[2]; + Reg::RegTensor regKScale[2]; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - using CastTrait = MicroAPI::CastTrait; - static constexpr CastTrait castTraitB162B32_EVEN = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - static constexpr CastTrait castTraitB162B32_ODD = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = { - MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = { - MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; - - MicroAPI::LoadAlign(regW[0], weight0); - MicroAPI::LoadAlign(regW[1], weight1); - MicroAPI::LoadAlign(regQScale[0], qScale0); - MicroAPI::LoadAlign(regQScale[1], qScale1); - MicroAPI::Mul(regW[0], regW[0], regQScale[0], maskAllB32); - MicroAPI::Mul(regW[1], regW[1], regQScale[1], maskAllB32); + using CastTrait = Reg::CastTrait; + static constexpr CastTrait castTraitB162B32_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr CastTrait castTraitB162B32_ODD = {Reg::RegLayout::ONE, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + + Reg::LoadAlign(regW[0], weight0); + Reg::LoadAlign(regW[1], weight1); + Reg::LoadAlign(regQScale[0], qScale0); + Reg::LoadAlign(regQScale[1], qScale1); + Reg::Mul(regW[0], regW[0], regQScale[0], maskAllB32); + Reg::Mul(regW[1], regW[1], regQScale[1], maskAllB32); // 读写依赖,寄存器可以保序 - MicroAPI::StoreAlign(weightTemp0, regW[0], maskAllB32); - MicroAPI::StoreAlign(weightTemp1, regW[1], maskAllB32); + Reg::StoreAlign(weightTemp0, regW[0], maskAllB32); + Reg::StoreAlign(weightTemp1, regW[1], maskAllB32); liV2Vector1::DuplicateZero(regSum0, maskAllB32); liV2Vector1::DuplicateZero(regSum1, maskAllB32); // interleave load - MicroAPI::LoadAlign(regKScale[0], regKScale[1], kScale0); + Reg::LoadAlign(regKScale[0], regKScale[1], kScale0); for (uint16_t i = (uint16_t)(0); i < gSize; i++) { // RowStride是256, 行都落在一个bank上 - MicroAPI::LoadAlign(regQKB16[0], qk0 + 256 * i); + Reg::LoadAlign(regQKB16[0], qk0 + 256 * i); // RowStride是256, 行都落在一个bank上 - MicroAPI::LoadAlign(regQKB16[1], qk1 + 256 * i); - MicroAPI::LoadAlign(regwBrc[0], weightTemp0 + i); - MicroAPI::LoadAlign(regwBrc[1], weightTemp1 + i); + Reg::LoadAlign(regQKB16[1], qk1 + 256 * i); + Reg::LoadAlign(regwBrc[0], weightTemp0 + i); + Reg::LoadAlign(regwBrc[1], weightTemp1 + i); // interleave cast - MicroAPI::Cast(regQK0[0], regQKB16[0], maskAllB32); - MicroAPI::Cast(regQK0[1], regQKB16[0], maskAllB32); - MicroAPI::Cast(regQK1[0], regQKB16[1], maskAllB32); - MicroAPI::Cast(regQK1[1], regQKB16[1], maskAllB32); - MicroAPI::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32); - MicroAPI::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32); + Reg::Cast(regQK0[0], regQKB16[0], maskAllB32); + Reg::Cast(regQK0[1], regQKB16[0], maskAllB32); + Reg::Cast(regQK1[0], regQKB16[1], maskAllB32); + Reg::Cast(regQK1[1], regQKB16[1], maskAllB32); + Reg::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32); + Reg::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32); } // Apply kScale scaling - MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); - MicroAPI::Mul(regSum1[0], regSum1[0], regKScale[0], maskAllB32); - MicroAPI::Mul(regSum1[1], regSum1[1], regKScale[1], maskAllB32); + Reg::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32); + Reg::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32); + Reg::Mul(regSum1[0], regSum1[0], regKScale[0], maskAllB32); + Reg::Mul(regSum1[1], regSum1[1], regKScale[1], maskAllB32); // Convert to bfloat16 and store output channel - MicroAPI::RegTensor regSumBF16[2]; - MicroAPI::RegTensor regOut[2]; - MicroAPI::Cast(regSumBF16[0], regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16[1], regSum1[1], maskAllB32); - MicroAPI::Cast(regSumBF16[0], regSum0[0], maskAllB32); - MicroAPI::Cast(regSumBF16[1], regSum1[0], maskAllB32); + Reg::RegTensor regSumBF16[2]; + Reg::RegTensor regOut[2]; + Reg::Cast(regSumBF16[0], regSum0[1], maskAllB32); + Reg::Cast(regSumBF16[1], regSum1[1], maskAllB32); + Reg::Cast(regSumBF16[0], regSum0[0], maskAllB32); + Reg::Cast(regSumBF16[1], regSum1[0], maskAllB32); liV2Vector1::FloatX2ToSortableKey(regOut[0], regOut[1], regSumBF16[0], regSumBF16[1], bf16Ctx, maskAllB16); - MicroAPI::StoreAlign(out0, regOut[0], maskAllB16); - MicroAPI::StoreAlign(out1, regOut[1], maskAllB16); + Reg::StoreAlign(out0, regOut[0], maskAllB16); + Reg::StoreAlign(out1, regOut[1], maskAllB16); } __aicore__ inline void MulWeightAndReduceSum2(const LocalTensor &out_, // out [2, S2Base] [128 ] @@ -1021,25 +981,25 @@ __simd_vf__ void MulWeightAndReduceSumOptionalScaleGSizeEvenVF(__ubuf__ uint16_t (void)qScaleValue; } - MicroAPI::RegTensor regwBrc; - MicroAPI::RegTensor regQK[2]; - MicroAPI::RegTensor regW; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regwBrc; + Reg::RegTensor regQK[2]; + Reg::RegTensor regW; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = { - MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = { - MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; - MicroAPI::LoadAlign(regW, weight); + Reg::LoadAlign(regW, weight); if constexpr (WITH_SCALE) { - MicroAPI::Muls(regW, regW, qScaleValue, maskAllB32); + Reg::Muls(regW, regW, qScaleValue, maskAllB32); } liV2Vector1::DuplicateZero(regSum0, maskAllB32); @@ -1047,35 +1007,35 @@ __simd_vf__ void MulWeightAndReduceSumOptionalScaleGSizeEvenVF(__ubuf__ uint16_t // unroll2 for (uint16_t i = (uint16_t)(0); i < gSize; i += 2) { - MicroAPI::LoadAlign(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上 - MicroAPI::LoadAlign(regQK[1], qk + 128 * i + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上 + Reg::LoadAlign(regQK[1], qk + 128 * i + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, i); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::LoadAlign(regQK[0], qk + 128 * i + 128); - MicroAPI::LoadAlign(regQK[1], qk + 128 * i + 128 + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * i + 128); + Reg::LoadAlign(regQK[1], qk + 128 * i + 128 + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, i + 1); liV2Vector1::WeightedAccum(regSum1, regQK, regwBrc, maskAllB32); } - MicroAPI::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); - MicroAPI::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); + Reg::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); + Reg::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); if constexpr (WITH_SCALE) { - MicroAPI::Muls(regSum0[0], regSum0[0], kScaleValue, maskAllB32); - MicroAPI::Muls(regSum0[1], regSum0[1], kScaleValue, maskAllB32); + Reg::Muls(regSum0[0], regSum0[0], kScaleValue, maskAllB32); + Reg::Muls(regSum0[1], regSum0[1], kScaleValue, maskAllB32); } - MicroAPI::RegTensor regSumBF16; + Reg::RegTensor regSumBF16; // interleave cast ==> regSum[1] high regSum[0] low - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::Cast(regSumBF16, regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16, regSum0[0], maskAllB32); + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::Cast(regSumBF16, regSum0[1], maskAllB32); + Reg::Cast(regSumBF16, regSum0[0], maskAllB32); - MicroAPI::RegTensor regOut; + Reg::RegTensor regOut; liV2Vector1::FloatToSortableKey(regOut, regSumBF16, bf16Ctx, maskAllB16); // normal store - MicroAPI::StoreAlign(out, regOut, maskAllB16); + Reg::StoreAlign(out, regOut, maskAllB16); } template @@ -1088,25 +1048,25 @@ __simd_vf__ void MulWeightAndReduceSumOptionalScaleGSizeOddVF(__ubuf__ uint16_t (void)qScaleValue; } - MicroAPI::RegTensor regwBrc; - MicroAPI::RegTensor regQK[2]; - MicroAPI::RegTensor regW; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regwBrc; + Reg::RegTensor regQK[2]; + Reg::RegTensor regW; + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = { - MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = { - MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; - MicroAPI::LoadAlign(regW, weight); + Reg::LoadAlign(regW, weight); if constexpr (WITH_SCALE) { - MicroAPI::Muls(regW, regW, qScaleValue, maskAllB32); + Reg::Muls(regW, regW, qScaleValue, maskAllB32); } liV2Vector1::DuplicateZero(regSum0, maskAllB32); @@ -1114,40 +1074,40 @@ __simd_vf__ void MulWeightAndReduceSumOptionalScaleGSizeOddVF(__ubuf__ uint16_t // unroll2 for (uint16_t i = (uint16_t)(0); (uint16_t)(i + 1) < gSize; i += 2) { - MicroAPI::LoadAlign(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上 - MicroAPI::LoadAlign(regQK[1], qk + 128 * i + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上 + Reg::LoadAlign(regQK[1], qk + 128 * i + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, i); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::LoadAlign(regQK[0], qk + 128 * i + 128); - MicroAPI::LoadAlign(regQK[1], qk + 128 * i + 128 + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * i + 128); + Reg::LoadAlign(regQK[1], qk + 128 * i + 128 + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, i + 1); liV2Vector1::WeightedAccum(regSum1, regQK, regwBrc, maskAllB32); } - MicroAPI::LoadAlign(regQK[0], qk + 128 * (gSize - 1)); // RowStride是128, 行都落在一个bank上 - MicroAPI::LoadAlign(regQK[1], qk + 128 * (gSize - 1) + qkVLStride); + Reg::LoadAlign(regQK[0], qk + 128 * (gSize - 1)); // RowStride是128, 行都落在一个bank上 + Reg::LoadAlign(regQK[1], qk + 128 * (gSize - 1) + qkVLStride); liV2Vector1::BroadcastLane(regwBrc, regW, gSize - 1); liV2Vector1::WeightedAccum(regSum0, regQK, regwBrc, maskAllB32); - MicroAPI::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); - MicroAPI::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); + Reg::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32); + Reg::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32); if constexpr (WITH_SCALE) { - MicroAPI::Muls(regSum0[0], regSum0[0], kScaleValue, maskAllB32); - MicroAPI::Muls(regSum0[1], regSum0[1], kScaleValue, maskAllB32); + Reg::Muls(regSum0[0], regSum0[0], kScaleValue, maskAllB32); + Reg::Muls(regSum0[1], regSum0[1], kScaleValue, maskAllB32); } - MicroAPI::RegTensor regSumBF16; + Reg::RegTensor regSumBF16; // interleave cast ==> regSum[1] high regSum[0] low - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::Cast(regSumBF16, regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16, regSum0[0], maskAllB32); + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::Cast(regSumBF16, regSum0[1], maskAllB32); + Reg::Cast(regSumBF16, regSum0[0], maskAllB32); - MicroAPI::RegTensor regOut; + Reg::RegTensor regOut; liV2Vector1::FloatToSortableKey(regOut, regSumBF16, bf16Ctx, maskAllB16); // normal store - MicroAPI::StoreAlign(out, regOut, maskAllB16); + Reg::StoreAlign(out, regOut, maskAllB16); } template @@ -1206,77 +1166,77 @@ __simd_vf__ void MulWeightAndReduceSumOptionalScale2VF(__ubuf__ uint16_t *out0, (void)qScaleValue; } - MicroAPI::RegTensor regwBrc[2]; - MicroAPI::RegTensor regQK0[2]; - MicroAPI::RegTensor regQK1[2]; - MicroAPI::RegTensor regW[2]; + Reg::RegTensor regwBrc[2]; + Reg::RegTensor regQK0[2]; + Reg::RegTensor regQK1[2]; + Reg::RegTensor regW[2]; - MicroAPI::RegTensor regSum0[2]; - MicroAPI::RegTensor regSum1[2]; - MicroAPI::MaskReg maskAllB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg maskAllB16 = MicroAPI::CreateMask(); + Reg::RegTensor regSum0[2]; + Reg::RegTensor regSum1[2]; + Reg::MaskReg maskAllB32 = Reg::CreateMask(); + Reg::MaskReg maskAllB16 = Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); - constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = { - MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; - constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = { - MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT, MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::MERGING, RoundMode::CAST_ROUND}; + constexpr static Reg::CastTrait castTraitF32ToF16_ODD = {Reg::RegLayout::ONE, Reg::SatMode::NO_SAT, + Reg::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; - MicroAPI::LoadAlign(regW[0], weight0); - MicroAPI::LoadAlign(regW[1], weight1); + Reg::LoadAlign(regW[0], weight0); + Reg::LoadAlign(regW[1], weight1); if constexpr (WITH_SCALE) { - MicroAPI::Muls(regW[0], regW[0], qScaleValue, maskAllB32); - MicroAPI::Muls(regW[1], regW[1], qScaleValue, maskAllB32); + Reg::Muls(regW[0], regW[0], qScaleValue, maskAllB32); + Reg::Muls(regW[1], regW[1], qScaleValue, maskAllB32); } // regW[0]与weight1混合使用 - MicroAPI::StoreAlign(weightTemp, regW[1], maskAllB32); - MicroAPI::LocalMemBar(); + Reg::StoreAlign(weightTemp, regW[1], maskAllB32); + Reg::LocalMemBar(); liV2Vector1::DuplicateZero(regSum0, maskAllB32); liV2Vector1::DuplicateZero(regSum1, maskAllB32); for (uint16_t i = (uint16_t)(0); i < gSize; i++) { - MicroAPI::LoadAlign(regQK0[0], qk0 + 128 * i); - MicroAPI::LoadAlign(regQK0[1], qk0 + 128 * i + qkVLStride); - MicroAPI::LoadAlign(regQK1[0], qk1 + 128 * i); - MicroAPI::LoadAlign(regQK1[1], qk1 + 128 * i + qkVLStride); + Reg::LoadAlign(regQK0[0], qk0 + 128 * i); + Reg::LoadAlign(regQK0[1], qk0 + 128 * i + qkVLStride); + Reg::LoadAlign(regQK1[0], qk1 + 128 * i); + Reg::LoadAlign(regQK1[1], qk1 + 128 * i + qkVLStride); // 混合使用对整体性能更好 liV2Vector1::BroadcastLane(regwBrc[0], regW[0], i); // Weight无bank冲突,用LoadAlign来提取weight标量 // 地址空间处理:原 BroadcastLane(ptr) 内联为 LoadAlign BRC,避免 __ubuf__ 传给 __local_mem__ 参数 - MicroAPI::LoadAlign(regwBrc[1], weightTemp + i); - MicroAPI::Relu(regQK0[0], regQK0[0], maskAllB32); - MicroAPI::Relu(regQK0[1], regQK0[1], maskAllB32); - MicroAPI::Relu(regQK1[0], regQK1[0], maskAllB32); - MicroAPI::Relu(regQK1[1], regQK1[1], maskAllB32); - MicroAPI::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32); - MicroAPI::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32); - MicroAPI::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32); + Reg::LoadAlign(regwBrc[1], weightTemp + i); + Reg::Relu(regQK0[0], regQK0[0], maskAllB32); + Reg::Relu(regQK0[1], regQK0[1], maskAllB32); + Reg::Relu(regQK1[0], regQK1[0], maskAllB32); + Reg::Relu(regQK1[1], regQK1[1], maskAllB32); + Reg::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32); + Reg::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32); + Reg::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32); } if constexpr (WITH_SCALE) { - MicroAPI::Muls(regSum0[0], regSum0[0], kScaleValue, maskAllB32); - MicroAPI::Muls(regSum0[1], regSum0[1], kScaleValue, maskAllB32); - MicroAPI::Muls(regSum1[0], regSum1[0], kScaleValue, maskAllB32); - MicroAPI::Muls(regSum1[1], regSum1[1], kScaleValue, maskAllB32); + Reg::Muls(regSum0[0], regSum0[0], kScaleValue, maskAllB32); + Reg::Muls(regSum0[1], regSum0[1], kScaleValue, maskAllB32); + Reg::Muls(regSum1[0], regSum1[0], kScaleValue, maskAllB32); + Reg::Muls(regSum1[1], regSum1[1], kScaleValue, maskAllB32); } // Convert to bfloat16 and store output channel - MicroAPI::RegTensor regSumBF16[2]; - MicroAPI::RegTensor regOut[2]; - MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); - MicroAPI::DeInterleave(regSum1[0], regSum1[1], regSum1[0], regSum1[1]); - MicroAPI::Cast(regSumBF16[0], regSum0[1], maskAllB32); - MicroAPI::Cast(regSumBF16[1], regSum1[1], maskAllB32); - MicroAPI::Cast(regSumBF16[0], regSum0[0], maskAllB32); - MicroAPI::Cast(regSumBF16[1], regSum1[0], maskAllB32); + Reg::RegTensor regSumBF16[2]; + Reg::RegTensor regOut[2]; + Reg::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]); + Reg::DeInterleave(regSum1[0], regSum1[1], regSum1[0], regSum1[1]); + Reg::Cast(regSumBF16[0], regSum0[1], maskAllB32); + Reg::Cast(regSumBF16[1], regSum1[1], maskAllB32); + Reg::Cast(regSumBF16[0], regSum0[0], maskAllB32); + Reg::Cast(regSumBF16[1], regSum1[0], maskAllB32); liV2Vector1::FloatX2ToSortableKey(regOut[0], regOut[1], regSumBF16[0], regSumBF16[1], bf16Ctx, maskAllB16); - MicroAPI::StoreAlign(out0, regOut[0], maskAllB16); - MicroAPI::StoreAlign(out1, regOut[1], maskAllB16); + Reg::StoreAlign(out0, regOut[0], maskAllB16); + Reg::StoreAlign(out1, regOut[1], maskAllB16); } template @@ -1331,48 +1291,46 @@ __aicore__ inline void MulWeightAndReduceSumMX2(const LocalTensor &out weightTemp_, 1.0f, 1.0f, gSize); } -__simd_callee__ inline void CastWeightToBf16(AscendC::MicroAPI::RegTensor &dst, __ubuf__ float *src, - AscendC::MicroAPI::MaskReg &maskAllB32) +__simd_callee__ inline void CastWeightToBf16(AscendC::Reg::RegTensor &dst, __ubuf__ float *src, + AscendC::Reg::MaskReg &maskAllB32) { - using CastTrait = AscendC::MicroAPI::CastTrait; - static constexpr CastTrait castTraitF32ToBf16 = {AscendC::MicroAPI::RegLayout::ZERO, - AscendC::MicroAPI::SatMode::NO_SAT, - AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; - AscendC::MicroAPI::RegTensor regWeightF32; - AscendC::MicroAPI::LoadAlign(regWeightF32, src); - AscendC::MicroAPI::Cast(dst, regWeightF32, maskAllB32); + using CastTrait = AscendC::Reg::CastTrait; + static constexpr CastTrait castTraitF32ToBf16 = {AscendC::Reg::RegLayout::ZERO, AscendC::Reg::SatMode::NO_SAT, + AscendC::Reg::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; + AscendC::Reg::RegTensor regWeightF32; + AscendC::Reg::LoadAlign(regWeightF32, src); + AscendC::Reg::Cast(dst, regWeightF32, maskAllB32); // Cast结果按B32 lane落位,在目标寄存器内压紧为连续BF16,供BroadcastLane按元素索引。 - AscendC::MicroAPI::Pack( - (AscendC::MicroAPI::RegTensor &)dst, (AscendC::MicroAPI::RegTensor &)dst); + AscendC::Reg::Pack((AscendC::Reg::RegTensor &)dst, + (AscendC::Reg::RegTensor &)dst); } __simd_vf__ void MulWeightAndReduceSumMXFP4VF(__ubuf__ uint16_t *out, __ubuf__ bfloat16_t *qk, __ubuf__ float *weight, uint16_t gSize) { constexpr uint32_t BF16_QK_ROW_STRIDE = UB_BANK_DEPTH_STRIDE / sizeof(bfloat16_t); - AscendC::MicroAPI::RegTensor regQK; - AscendC::MicroAPI::RegTensor regWeight; - AscendC::MicroAPI::RegTensor regWeightBrc; - AscendC::MicroAPI::RegTensor regSum; - AscendC::MicroAPI::MaskReg maskAllB16 = - AscendC::MicroAPI::CreateMask(); - AscendC::MicroAPI::MaskReg maskAllB32 = AscendC::MicroAPI::CreateMask(); + AscendC::Reg::RegTensor regQK; + AscendC::Reg::RegTensor regWeight; + AscendC::Reg::RegTensor regWeightBrc; + AscendC::Reg::RegTensor regSum; + AscendC::Reg::MaskReg maskAllB16 = AscendC::Reg::CreateMask(); + AscendC::Reg::MaskReg maskAllB32 = AscendC::Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); CastWeightToBf16(regWeight, weight, maskAllB32); - AscendC::MicroAPI::Duplicate(regSum, bfloat16_t(0.0f), maskAllB16); + AscendC::Reg::Duplicate(regSum, bfloat16_t(0.0f), maskAllB16); for (uint16_t i = 0; i < gSize; i++) { - AscendC::MicroAPI::LoadAlign(regQK, qk + BF16_QK_ROW_STRIDE * i); + AscendC::Reg::LoadAlign(regQK, qk + BF16_QK_ROW_STRIDE * i); liV2Vector1::BroadcastLane(regWeightBrc, regWeight, i); - AscendC::MicroAPI::MulAddDst(regSum, regQK, regWeightBrc, maskAllB16); + AscendC::Reg::MulAddDst(regSum, regQK, regWeightBrc, maskAllB16); } - AscendC::MicroAPI::RegTensor regOut; + AscendC::Reg::RegTensor regOut; liV2Vector1::FloatToSortableKey(regOut, regSum, bf16Ctx, maskAllB16); - AscendC::MicroAPI::StoreAlign(out, regOut, maskAllB16); + AscendC::Reg::StoreAlign(out, regOut, maskAllB16); } __aicore__ inline void MulWeightAndReduceSumMXFP4(const LocalTensor &out_, const LocalTensor &qk_, @@ -1391,37 +1349,35 @@ __simd_vf__ void MulWeightAndReduceSumMXFP4TwoRowsVF(__ubuf__ uint16_t *out0, __ __ubuf__ float *weight0, __ubuf__ float *weight1, uint16_t gSize) { constexpr uint32_t BF16_QK_ROW_STRIDE = UB_BANK_DEPTH_STRIDE / sizeof(bfloat16_t); - AscendC::MicroAPI::RegTensor regQK0; - AscendC::MicroAPI::RegTensor regQK1; - AscendC::MicroAPI::RegTensor regWeight[2]; - AscendC::MicroAPI::RegTensor regWeightBrc[2]; - AscendC::MicroAPI::RegTensor regSum[2]; - AscendC::MicroAPI::MaskReg maskAllB16 = - AscendC::MicroAPI::CreateMask(); - AscendC::MicroAPI::MaskReg maskAllB32 = - AscendC::MicroAPI::CreateMask(); + AscendC::Reg::RegTensor regQK0; + AscendC::Reg::RegTensor regQK1; + AscendC::Reg::RegTensor regWeight[2]; + AscendC::Reg::RegTensor regWeightBrc[2]; + AscendC::Reg::RegTensor regSum[2]; + AscendC::Reg::MaskReg maskAllB16 = AscendC::Reg::CreateMask(); + AscendC::Reg::MaskReg maskAllB32 = AscendC::Reg::CreateMask(); liV2Vector1::FloatSortConstCtx bf16Ctx; liV2Vector1::InitFloatSortConstCtx(bf16Ctx, maskAllB16); CastWeightToBf16(regWeight[0], weight0, maskAllB32); CastWeightToBf16(regWeight[1], weight1, maskAllB32); - AscendC::MicroAPI::Duplicate(regSum[0], bfloat16_t(0.0f), maskAllB16); - AscendC::MicroAPI::Duplicate(regSum[1], bfloat16_t(0.0f), maskAllB16); + AscendC::Reg::Duplicate(regSum[0], bfloat16_t(0.0f), maskAllB16); + AscendC::Reg::Duplicate(regSum[1], bfloat16_t(0.0f), maskAllB16); for (uint16_t i = 0; i < gSize; i++) { - AscendC::MicroAPI::LoadAlign(regQK0, qk0 + BF16_QK_ROW_STRIDE * i); - AscendC::MicroAPI::LoadAlign(regQK1, qk1 + BF16_QK_ROW_STRIDE * i); + AscendC::Reg::LoadAlign(regQK0, qk0 + BF16_QK_ROW_STRIDE * i); + AscendC::Reg::LoadAlign(regQK1, qk1 + BF16_QK_ROW_STRIDE * i); liV2Vector1::BroadcastLane(regWeightBrc[0], regWeight[0], i); liV2Vector1::BroadcastLane(regWeightBrc[1], regWeight[1], i); - AscendC::MicroAPI::MulAddDst(regSum[0], regQK0, regWeightBrc[0], maskAllB16); - AscendC::MicroAPI::MulAddDst(regSum[1], regQK1, regWeightBrc[1], maskAllB16); + AscendC::Reg::MulAddDst(regSum[0], regQK0, regWeightBrc[0], maskAllB16); + AscendC::Reg::MulAddDst(regSum[1], regQK1, regWeightBrc[1], maskAllB16); } - AscendC::MicroAPI::RegTensor regOut[2]; + AscendC::Reg::RegTensor regOut[2]; liV2Vector1::FloatX2ToSortableKey(regOut[0], regOut[1], regSum[0], regSum[1], bf16Ctx, maskAllB16); - AscendC::MicroAPI::StoreAlign(out0, regOut[0], maskAllB16); - AscendC::MicroAPI::StoreAlign(out1, regOut[1], maskAllB16); + AscendC::Reg::StoreAlign(out0, regOut[0], maskAllB16); + AscendC::Reg::StoreAlign(out1, regOut[1], maskAllB16); } __aicore__ inline void MulWeightAndReduceSumMXFP4TwoRows(const LocalTensor &out_, uint32_t outStride, @@ -1468,8 +1424,7 @@ __aicore__ inline void BatchMulWeightAndReduceSumMX(const LocalTensor & return; } if (batch == 2) { - MulWeightAndReduceSumMX2(out_, outStride, qk_, qkVLStride, qkStride, - weight_, weightStride, weightTemp_, gSize); + MulWeightAndReduceSumMX2(out_, outStride, qk_, qkVLStride, qkStride, weight_, weightStride, weightTemp_, gSize); } else { MulWeightAndReduceSumMX(out_, qk_, qkVLStride, weight_, gSize); } diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/vf_topk.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/vf_topk.h index aab1e2b926b9..dc1af9872834 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/vf_topk.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/vf_topk.h @@ -9,547 +9,503 @@   */ /*! -* \file vf_topk.h -* \brief -*/ + * \file vf_topk.h + * \brief + */ #ifndef VF_TOP_K_H #define VF_TOP_K_H namespace topkb32 { -__simd_callee__ inline void StoreHistogramResult(__ubuf__ uint32_t* histogramsBuf, - MicroAPI::RegTensor& cout0, - MicroAPI::RegTensor& cout1, - MicroAPI::MaskReg& pregB16, - MicroAPI::MaskReg& pregB32) +__simd_callee__ inline void StoreHistogramResult(__ubuf__ uint32_t *histogramsBuf, Reg::RegTensor &cout0, + Reg::RegTensor &cout1, Reg::MaskReg &pregB16, + Reg::MaskReg &pregB32) { - MicroAPI::RegTensor cout0U32Even; - MicroAPI::RegTensor cout0U32Odd; - MicroAPI::RegTensor cout1U32Even; - MicroAPI::RegTensor cout1U32Odd; - - static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - - static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE, - MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - - MicroAPI::Cast(cout0U32Even, cout0, pregB16); - MicroAPI::Cast(cout0U32Odd, cout0, pregB16); - MicroAPI::Cast(cout1U32Even, cout1, pregB16); - MicroAPI::Cast(cout1U32Odd, cout1, pregB16); - - MicroAPI::StoreAlign(histogramsBuf, - cout0U32Even, cout0U32Odd, pregB32); - MicroAPI::StoreAlign(histogramsBuf + 128, - cout1U32Even, cout1U32Odd, pregB32); + Reg::RegTensor cout0U32Even; + Reg::RegTensor cout0U32Odd; + Reg::RegTensor cout1U32Even; + Reg::RegTensor cout1U32Odd; + + static constexpr Reg::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + + static constexpr Reg::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {Reg::RegLayout::ONE, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + + Reg::Cast(cout0U32Even, cout0, pregB16); + Reg::Cast(cout0U32Odd, cout0, pregB16); + Reg::Cast(cout1U32Even, cout1, pregB16); + Reg::Cast(cout1U32Odd, cout1, pregB16); + + Reg::StoreAlign(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32); + Reg::StoreAlign(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32); } -__simd_callee__ inline void FindTargetBinAndUpdateNextK(__ubuf__ uint32_t* idxBuf, - __ubuf__ uint32_t* nkValueBuf, - __ubuf__ uint32_t* histogramsBuf, - MicroAPI::RegTensor& btmK, - MicroAPI::MaskReg& pregB32) +__simd_callee__ inline void FindTargetBinAndUpdateNextK(__ubuf__ uint32_t *idxBuf, __ubuf__ uint32_t *nkValueBuf, + __ubuf__ uint32_t *histogramsBuf, + Reg::RegTensor &btmK, Reg::MaskReg &pregB32) { - MicroAPI::ClearSpr(); + Reg::ClearSpr(); - MicroAPI::UnalignRegForStore alignIdx; + Reg::UnalignRegForStore alignIdx; for (uint16_t i = 0; i < (uint16_t)(4); ++i) { - MicroAPI::RegTensor idxC; - MicroAPI::RegTensor cout; - MicroAPI::RegTensor sqzIdx; - - MicroAPI::MaskReg pregGE = MicroAPI::CreateMask(); - - MicroAPI::Arange(idxC, i * 64); - MicroAPI::LoadAlign(cout, histogramsBuf + i * 64); - MicroAPI::Compare(pregGE, cout, btmK, pregB32); - MicroAPI::Squeeze( - sqzIdx, (MicroAPI::RegTensor&)idxC, pregGE); - MicroAPI::StoreUnAlign(idxBuf, sqzIdx, alignIdx); + Reg::RegTensor idxC; + Reg::RegTensor cout; + Reg::RegTensor sqzIdx; + + Reg::MaskReg pregGE = Reg::CreateMask(); + + Reg::Arange(idxC, i * 64); + Reg::LoadAlign(cout, histogramsBuf + i * 64); + Reg::Compare(pregGE, cout, btmK, pregB32); + Reg::Squeeze(sqzIdx, (Reg::RegTensor &)idxC, pregGE); + Reg::StoreUnAlign(idxBuf, sqzIdx, alignIdx); } - MicroAPI::StoreUnAlignPost(idxBuf, alignIdx); + Reg::StoreUnAlignPost(idxBuf, alignIdx); - MicroAPI::LocalMemBar(); + Reg::LocalMemBar(); - MicroAPI::RegTensor idx; - MicroAPI::LoadAlign(idx, idxBuf); + Reg::RegTensor idx; + Reg::LoadAlign(idx, idxBuf); - MicroAPI::RegTensor idxAll1; - MicroAPI::RegTensor idxPrev; - MicroAPI::RegTensor prevBinValue; - MicroAPI::Duplicate(idxAll1, 1); + Reg::RegTensor idxAll1; + Reg::RegTensor idxPrev; + Reg::RegTensor prevBinValue; + Reg::Duplicate(idxAll1, 1); - MicroAPI::RegTensor zeroAll; - MicroAPI::Duplicate(zeroAll, 0); + Reg::RegTensor zeroAll; + Reg::Duplicate(zeroAll, 0); - MicroAPI::MaskReg pregZero = MicroAPI::CreateMask(); - MicroAPI::Compare(pregZero, idx, zeroAll, pregB32); - MicroAPI::Sub(idxPrev, idx, (MicroAPI::RegTensor&)idxAll1, pregB32); - MicroAPI::ShiftRights(idxPrev, idxPrev, (int16_t)24, pregB32); + Reg::MaskReg pregZero = Reg::CreateMask(); + Reg::Compare(pregZero, idx, zeroAll, pregB32); + Reg::Sub(idxPrev, idx, (Reg::RegTensor &)idxAll1, pregB32); + Reg::ShiftRights(idxPrev, idxPrev, (int16_t)24, pregB32); - MicroAPI::Gather(prevBinValue, histogramsBuf, idxPrev, pregB32); - MicroAPI::Select(prevBinValue, zeroAll, prevBinValue, pregZero); + Reg::Gather(prevBinValue, histogramsBuf, idxPrev, pregB32); + Reg::Select(prevBinValue, zeroAll, prevBinValue, pregZero); - MicroAPI::RegTensor nextK; - MicroAPI::Sub(nextK, btmK, prevBinValue, pregB32); - MicroAPI::StoreAlign(nkValueBuf, nextK, pregB32); + Reg::RegTensor nextK; + Reg::Sub(nextK, btmK, prevBinValue, pregB32); + Reg::StoreAlign(nkValueBuf, nextK, pregB32); } -template -__simd_vf__ void HistogramsFirstVFImpl(__ubuf__ uint32_t* histogramsBuf, - __ubuf__ uint32_t* inputBuf, - uint16_t vfLoop, bool init) +template +__simd_vf__ void HistogramsFirstVFImpl(__ubuf__ uint32_t *histogramsBuf, __ubuf__ uint32_t *inputBuf, uint16_t vfLoop, + bool init) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); + Reg::MaskReg pregB8 = Reg::CreateMask(); // 计算直方图cout0 0-127 cout1 128-255 - MicroAPI::RegTensor cout0; - MicroAPI::RegTensor cout1; - MicroAPI::Duplicate(cout0, 0); - MicroAPI::Duplicate(cout1, 0); + Reg::RegTensor cout0; + Reg::RegTensor cout1; + Reg::Duplicate(cout0, 0); + Reg::Duplicate(cout1, 0); - MicroAPI::RegTensor vreg0; - MicroAPI::RegTensor vreg1; - MicroAPI::RegTensor vreg2; - MicroAPI::RegTensor vreg3; + Reg::RegTensor vreg0; + Reg::RegTensor vreg1; + Reg::RegTensor vreg2; + Reg::RegTensor vreg3; // 32bit 高16bit - MicroAPI::RegTensor vreg0U16; + Reg::RegTensor vreg0U16; // 32bit 低16bit - MicroAPI::RegTensor vreg1U16; - MicroAPI::RegTensor vreg2U16; - MicroAPI::RegTensor vreg3U16; + Reg::RegTensor vreg1U16; + Reg::RegTensor vreg2U16; + Reg::RegTensor vreg3U16; for (uint16_t i = 0; i < vfLoop; ++i) { - MicroAPI::LoadAlign(vreg1U16, vreg0U16, inputBuf + i * 256); - MicroAPI::LoadAlign( - vreg3U16, vreg2U16, inputBuf + (i * 256) + 128); - - MicroAPI::DeInterleave(vreg1, vreg0, - (MicroAPI::RegTensor&)vreg0U16, - (MicroAPI::RegTensor&)vreg2U16); - - MicroAPI::Histograms(cout0, vreg0, pregB8); - MicroAPI::Histograms(cout1, vreg0, pregB8); + Reg::LoadAlign(vreg1U16, vreg0U16, inputBuf + i * 256); + Reg::LoadAlign(vreg3U16, vreg2U16, inputBuf + (i * 256) + 128); + + Reg::DeInterleave(vreg1, vreg0, (Reg::RegTensor &)vreg0U16, (Reg::RegTensor &)vreg2U16); + + Reg::Histograms(cout0, vreg0, + pregB8); + Reg::Histograms(cout1, vreg0, + pregB8); } StoreHistogramResult(histogramsBuf, cout0, cout1, pregB16, pregB32); } -__simd_vf__ void FindFirstTargetBinVFImpl(__ubuf__ uint32_t* idx0Buf, - __ubuf__ uint32_t* nkValueBuf, __ubuf__ uint32_t* - histogramsBuf, uint32_t bottomK) +__simd_vf__ void FindFirstTargetBinVFImpl(__ubuf__ uint32_t *idx0Buf, __ubuf__ uint32_t *nkValueBuf, + __ubuf__ uint32_t *histogramsBuf, uint32_t bottomK) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::RegTensor btmK; - MicroAPI::Duplicate(btmK, bottomK); + Reg::RegTensor btmK; + Reg::Duplicate(btmK, bottomK); FindTargetBinAndUpdateNextK(idx0Buf, nkValueBuf, histogramsBuf, btmK, pregB32); } -template -__simd_vf__ void HistogramsSecondVFImpl(__ubuf__ uint32_t* histogramsBuf, - __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* idx0Buf, - uint16_t vfLoop, bool init) +template +__simd_vf__ void HistogramsSecondVFImpl(__ubuf__ uint32_t *histogramsBuf, __ubuf__ uint32_t *inputBuf, + __ubuf__ uint32_t *idx0Buf, uint16_t vfLoop, bool init) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); + Reg::MaskReg pregB8 = Reg::CreateMask(); // 计算直方图0-127 128-255 - MicroAPI::RegTensor cout0; - MicroAPI::RegTensor cout1; - MicroAPI::Duplicate(cout0, 0); - MicroAPI::Duplicate(cout1, 0); + Reg::RegTensor cout0; + Reg::RegTensor cout1; + Reg::Duplicate(cout0, 0); + Reg::Duplicate(cout1, 0); - MicroAPI::RegTensor idx0; + Reg::RegTensor idx0; // 0x000000fc -> 0xfcfcfcfc - MicroAPI::LoadAlign(idx0, idx0Buf); + Reg::LoadAlign(idx0, idx0Buf); - MicroAPI::RegTensor vreg0U16; - MicroAPI::RegTensor vreg1U16; - MicroAPI::RegTensor vreg2U16; - MicroAPI::RegTensor vreg3U16; + Reg::RegTensor vreg0U16; + Reg::RegTensor vreg1U16; + Reg::RegTensor vreg2U16; + Reg::RegTensor vreg3U16; - MicroAPI::RegTensor vreg0; - MicroAPI::RegTensor vreg1; - MicroAPI::RegTensor vreg2; - MicroAPI::RegTensor vreg3; + Reg::RegTensor vreg0; + Reg::RegTensor vreg1; + Reg::RegTensor vreg2; + Reg::RegTensor vreg3; for (uint16_t i = 0; i < vfLoop; ++i) { - MicroAPI::LoadAlign(vreg1U16, - vreg0U16, inputBuf + i * 256); - MicroAPI::LoadAlign(vreg3U16, - vreg2U16, inputBuf + (i * 256) + 128); - - MicroAPI::DeInterleave(vreg1, vreg0, - (MicroAPI::RegTensor&)vreg0U16, - (MicroAPI::RegTensor&)vreg2U16); - - MicroAPI::MaskReg pregEQ = MicroAPI::CreateMask(); - MicroAPI::Compare(pregEQ, vreg0, (MicroAPI::RegTensor&)idx0, pregB8); - - MicroAPI::Histograms(cout0, vreg1, pregEQ); - MicroAPI::Histograms(cout1, vreg1, pregEQ); + Reg::LoadAlign(vreg1U16, vreg0U16, inputBuf + i * 256); + Reg::LoadAlign(vreg3U16, vreg2U16, inputBuf + (i * 256) + 128); + + Reg::DeInterleave(vreg1, vreg0, (Reg::RegTensor &)vreg0U16, (Reg::RegTensor &)vreg2U16); + + Reg::MaskReg pregEQ = Reg::CreateMask(); + Reg::Compare(pregEQ, vreg0, (Reg::RegTensor &)idx0, pregB8); + + Reg::Histograms(cout0, vreg1, + pregEQ); + Reg::Histograms(cout1, vreg1, + pregEQ); } StoreHistogramResult(histogramsBuf, cout0, cout1, pregB16, pregB32); } // kValue新的bottomK -__simd_vf__ void FindSecondTargetBinVFImpl(__ubuf__ uint32_t* idx1Buf, - __ubuf__ uint32_t* nkValueBuf, __ubuf__ uint32_t* kValue, - __ubuf__ uint32_t* histogramsBuf) +__simd_vf__ void FindSecondTargetBinVFImpl(__ubuf__ uint32_t *idx1Buf, __ubuf__ uint32_t *nkValueBuf, + __ubuf__ uint32_t *kValue, __ubuf__ uint32_t *histogramsBuf) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::RegTensor btmK1; - MicroAPI::LoadAlign(btmK1, kValue); + Reg::RegTensor btmK1; + Reg::LoadAlign(btmK1, kValue); FindTargetBinAndUpdateNextK(idx1Buf, nkValueBuf, histogramsBuf, btmK1, pregB32); } -template -__simd_vf__ void HistogramsThirdVFImpl(__ubuf__ uint32_t* histogramsBuf, - __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* idx0Buf, - __ubuf__ uint32_t* idx1Buf, uint16_t vfLoop, bool init) +template +__simd_vf__ void HistogramsThirdVFImpl(__ubuf__ uint32_t *histogramsBuf, __ubuf__ uint32_t *inputBuf, + __ubuf__ uint32_t *idx0Buf, __ubuf__ uint32_t *idx1Buf, uint16_t vfLoop, + bool init) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); + Reg::MaskReg pregB8 = Reg::CreateMask(); // 计算直方图0-127 128-255 - MicroAPI::RegTensor cout0; - MicroAPI::RegTensor cout1; - MicroAPI::Duplicate(cout0, 0); - MicroAPI::Duplicate(cout1, 0); + Reg::RegTensor cout0; + Reg::RegTensor cout1; + Reg::Duplicate(cout0, 0); + Reg::Duplicate(cout1, 0); - MicroAPI::RegTensor idx0; - MicroAPI::RegTensor idx1; + Reg::RegTensor idx0; + Reg::RegTensor idx1; // 0x000000fc -> 0xfcfcfcfc - MicroAPI::LoadAlign(idx0, idx0Buf); - MicroAPI::LoadAlign(idx1, idx1Buf); + Reg::LoadAlign(idx0, idx0Buf); + Reg::LoadAlign(idx1, idx1Buf); - MicroAPI::RegTensor vreg0; - MicroAPI::RegTensor vreg1; - MicroAPI::RegTensor vreg2; - MicroAPI::RegTensor vreg3; + Reg::RegTensor vreg0; + Reg::RegTensor vreg1; + Reg::RegTensor vreg2; + Reg::RegTensor vreg3; - MicroAPI::RegTensor vreg0U16; - MicroAPI::RegTensor vreg1U16; - MicroAPI::RegTensor vreg2U16; - MicroAPI::RegTensor vreg3U16; + Reg::RegTensor vreg0U16; + Reg::RegTensor vreg1U16; + Reg::RegTensor vreg2U16; + Reg::RegTensor vreg3U16; for (uint16_t i = 0; i < vfLoop; ++i) { - MicroAPI::LoadAlign(vreg1U16, - vreg0U16, inputBuf + i * 256); - MicroAPI::LoadAlign(vreg3U16, - vreg2U16, inputBuf + (i * 256) + 128); - - MicroAPI::DeInterleave(vreg1, vreg0, (MicroAPI::RegTensor&)vreg0U16, - (MicroAPI::RegTensor&)vreg2U16); - MicroAPI::DeInterleave(vreg3, vreg2, (MicroAPI::RegTensor&)vreg1U16, - (MicroAPI::RegTensor&)vreg3U16); - - MicroAPI::MaskReg pregEQ0 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregEQ1 = MicroAPI::CreateMask(); - MicroAPI::Compare(pregEQ0, vreg0, (MicroAPI::RegTensor&)idx0, pregB8); - MicroAPI::Compare(pregEQ1, vreg1, (MicroAPI::RegTensor&)idx1, pregB8); - - MicroAPI::MaskReg pregEQ = MicroAPI::CreateMask(); - MicroAPI::And(pregEQ, pregEQ0, pregEQ1, pregB8); - - MicroAPI::Histograms(cout0, vreg2, pregEQ); - MicroAPI::Histograms(cout1, vreg2, pregEQ); + Reg::LoadAlign(vreg1U16, vreg0U16, inputBuf + i * 256); + Reg::LoadAlign(vreg3U16, vreg2U16, inputBuf + (i * 256) + 128); + + Reg::DeInterleave(vreg1, vreg0, (Reg::RegTensor &)vreg0U16, (Reg::RegTensor &)vreg2U16); + Reg::DeInterleave(vreg3, vreg2, (Reg::RegTensor &)vreg1U16, (Reg::RegTensor &)vreg3U16); + + Reg::MaskReg pregEQ0 = Reg::CreateMask(); + Reg::MaskReg pregEQ1 = Reg::CreateMask(); + Reg::Compare(pregEQ0, vreg0, (Reg::RegTensor &)idx0, pregB8); + Reg::Compare(pregEQ1, vreg1, (Reg::RegTensor &)idx1, pregB8); + + Reg::MaskReg pregEQ = Reg::CreateMask(); + Reg::And(pregEQ, pregEQ0, pregEQ1, pregB8); + + Reg::Histograms(cout0, vreg2, + pregEQ); + Reg::Histograms(cout1, vreg2, + pregEQ); } StoreHistogramResult(histogramsBuf, cout0, cout1, pregB16, pregB32); } -__simd_vf__ void FindThirdTargetBinVFImpl(__ubuf__ uint32_t* idx2Buf, - __ubuf__ uint32_t* nkValueBuf, __ubuf__ uint32_t* kValue, - __ubuf__ uint32_t* histogramsBuf) +__simd_vf__ void FindThirdTargetBinVFImpl(__ubuf__ uint32_t *idx2Buf, __ubuf__ uint32_t *nkValueBuf, + __ubuf__ uint32_t *kValue, __ubuf__ uint32_t *histogramsBuf) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::RegTensor btmK2; - MicroAPI::LoadAlign(btmK2, kValue); + Reg::RegTensor btmK2; + Reg::LoadAlign(btmK2, kValue); FindTargetBinAndUpdateNextK(idx2Buf, nkValueBuf, histogramsBuf, btmK2, pregB32); } -template -__simd_vf__ void HistogramsLastVFImpl(__ubuf__ uint32_t* histogramsBuf, - __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* idx0Buf, - __ubuf__ uint32_t* idx1Buf, __ubuf__ uint32_t* idx2Buf, - uint16_t vfLoop, bool init) +template +__simd_vf__ void HistogramsLastVFImpl(__ubuf__ uint32_t *histogramsBuf, __ubuf__ uint32_t *inputBuf, + __ubuf__ uint32_t *idx0Buf, __ubuf__ uint32_t *idx1Buf, + __ubuf__ uint32_t *idx2Buf, uint16_t vfLoop, bool init) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); + Reg::MaskReg pregB8 = Reg::CreateMask(); - MicroAPI::RegTensor idx0; - MicroAPI::RegTensor idx1; - MicroAPI::RegTensor idx2; + Reg::RegTensor idx0; + Reg::RegTensor idx1; + Reg::RegTensor idx2; // 0x000000fc -> 0xfcfcfcfc - MicroAPI::LoadAlign(idx0, idx0Buf); - MicroAPI::LoadAlign(idx1, idx1Buf); - MicroAPI::LoadAlign(idx2, idx2Buf); + Reg::LoadAlign(idx0, idx0Buf); + Reg::LoadAlign(idx1, idx1Buf); + Reg::LoadAlign(idx2, idx2Buf); // 计算直方图0-127 128-255 - MicroAPI::RegTensor cout0; - MicroAPI::RegTensor cout1; - MicroAPI::Duplicate(cout0, 0); - MicroAPI::Duplicate(cout1, 0); + Reg::RegTensor cout0; + Reg::RegTensor cout1; + Reg::Duplicate(cout0, 0); + Reg::Duplicate(cout1, 0); - MicroAPI::RegTensor vreg0U16; - MicroAPI::RegTensor vreg1U16; - MicroAPI::RegTensor vreg2U16; - MicroAPI::RegTensor vreg3U16; + Reg::RegTensor vreg0U16; + Reg::RegTensor vreg1U16; + Reg::RegTensor vreg2U16; + Reg::RegTensor vreg3U16; - MicroAPI::RegTensor vreg0; - MicroAPI::RegTensor vreg1; - MicroAPI::RegTensor vreg2; - MicroAPI::RegTensor vreg3; + Reg::RegTensor vreg0; + Reg::RegTensor vreg1; + Reg::RegTensor vreg2; + Reg::RegTensor vreg3; for (uint16_t i = 0; i < vfLoop; ++i) { - MicroAPI::LoadAlign(vreg1U16, vreg0U16, inputBuf + i * 256); - MicroAPI::LoadAlign(vreg3U16, - vreg2U16, inputBuf + (i * 256) + 128); - - MicroAPI::DeInterleave(vreg1, vreg0, - (MicroAPI::RegTensor&)vreg0U16, - (MicroAPI::RegTensor&)vreg2U16); - MicroAPI::DeInterleave(vreg3, vreg2, - (MicroAPI::RegTensor&)vreg1U16, - (MicroAPI::RegTensor&)vreg3U16); - - MicroAPI::MaskReg pregEQ0 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregEQ1 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregEQ2 = MicroAPI::CreateMask(); - MicroAPI::Compare(pregEQ0, vreg0, (MicroAPI::RegTensor&)idx0, pregB8); - MicroAPI::Compare(pregEQ1, vreg1, (MicroAPI::RegTensor&)idx1, pregB8); - MicroAPI::Compare(pregEQ2, vreg2, (MicroAPI::RegTensor&)idx2, pregB8); - - MicroAPI::MaskReg pregEQ0And1 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregEQAll = MicroAPI::CreateMask(); - MicroAPI::And(pregEQ0And1, pregEQ0, pregEQ1, pregB8); - MicroAPI::And(pregEQAll, pregEQ0And1, pregEQ2, pregB8); - - MicroAPI::Histograms(cout0, vreg3, pregEQAll); - MicroAPI::Histograms(cout1, vreg3, pregEQAll); + Reg::LoadAlign(vreg1U16, vreg0U16, inputBuf + i * 256); + Reg::LoadAlign(vreg3U16, vreg2U16, inputBuf + (i * 256) + 128); + + Reg::DeInterleave(vreg1, vreg0, (Reg::RegTensor &)vreg0U16, (Reg::RegTensor &)vreg2U16); + Reg::DeInterleave(vreg3, vreg2, (Reg::RegTensor &)vreg1U16, (Reg::RegTensor &)vreg3U16); + + Reg::MaskReg pregEQ0 = Reg::CreateMask(); + Reg::MaskReg pregEQ1 = Reg::CreateMask(); + Reg::MaskReg pregEQ2 = Reg::CreateMask(); + Reg::Compare(pregEQ0, vreg0, (Reg::RegTensor &)idx0, pregB8); + Reg::Compare(pregEQ1, vreg1, (Reg::RegTensor &)idx1, pregB8); + Reg::Compare(pregEQ2, vreg2, (Reg::RegTensor &)idx2, pregB8); + + Reg::MaskReg pregEQ0And1 = Reg::CreateMask(); + Reg::MaskReg pregEQAll = Reg::CreateMask(); + Reg::And(pregEQ0And1, pregEQ0, pregEQ1, pregB8); + Reg::And(pregEQAll, pregEQ0And1, pregEQ2, pregB8); + + Reg::Histograms(cout0, vreg3, + pregEQAll); + Reg::Histograms(cout1, vreg3, + pregEQAll); } StoreHistogramResult(histogramsBuf, cout0, cout1, pregB16, pregB32); } -__simd_vf__ void FindKthVFImpl(__ubuf__ uint32_t* kValue, - __ubuf__ uint32_t* histogramsBuf, __ubuf__ uint32_t* idx0Buf, - __ubuf__ uint32_t* idx1Buf, __ubuf__ uint32_t* idx2Buf, - __ubuf__ uint32_t* idx3Buf) +__simd_vf__ void FindKthVFImpl(__ubuf__ uint32_t *kValue, __ubuf__ uint32_t *histogramsBuf, __ubuf__ uint32_t *idx0Buf, + __ubuf__ uint32_t *idx1Buf, __ubuf__ uint32_t *idx2Buf, __ubuf__ uint32_t *idx3Buf) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::ClearSpr(); + Reg::ClearSpr(); - MicroAPI::UnalignRegForStore alignIdx3; + Reg::UnalignRegForStore alignIdx3; - MicroAPI::RegTensor btmK3; - MicroAPI::LoadAlign(btmK3, kValue); + Reg::RegTensor btmK3; + Reg::LoadAlign(btmK3, kValue); for (uint16_t i = 0; i < (uint16_t)(4); ++i) { - MicroAPI::RegTensor idxC; - MicroAPI::RegTensor cout; - MicroAPI::RegTensor sqzIdx3; - - MicroAPI::MaskReg pregGE = MicroAPI::CreateMask(); - - MicroAPI::Arange(idxC, i * 64); - MicroAPI::LoadAlign(cout, histogramsBuf + i * 64); - MicroAPI::Compare(pregGE, cout, btmK3, pregB32); - MicroAPI::Squeeze(sqzIdx3, - (MicroAPI::RegTensor&)idxC, pregGE); - MicroAPI::StoreUnAlign(idx3Buf, sqzIdx3, alignIdx3); + Reg::RegTensor idxC; + Reg::RegTensor cout; + Reg::RegTensor sqzIdx3; + + Reg::MaskReg pregGE = Reg::CreateMask(); + + Reg::Arange(idxC, i * 64); + Reg::LoadAlign(cout, histogramsBuf + i * 64); + Reg::Compare(pregGE, cout, btmK3, pregB32); + Reg::Squeeze(sqzIdx3, (Reg::RegTensor &)idxC, pregGE); + Reg::StoreUnAlign(idx3Buf, sqzIdx3, alignIdx3); } - MicroAPI::StoreUnAlignPost(idx3Buf, alignIdx3); + Reg::StoreUnAlignPost(idx3Buf, alignIdx3); - MicroAPI::LocalMemBar(); + Reg::LocalMemBar(); - MicroAPI::RegTensor idx0; - MicroAPI::RegTensor idx1; - MicroAPI::RegTensor idx2; - MicroAPI::RegTensor idx3; - MicroAPI::LoadAlign(idx0, idx0Buf); - MicroAPI::LoadAlign(idx1, idx1Buf); - MicroAPI::LoadAlign(idx2, idx2Buf); - MicroAPI::LoadAlign(idx3, idx3Buf); + Reg::RegTensor idx0; + Reg::RegTensor idx1; + Reg::RegTensor idx2; + Reg::RegTensor idx3; + Reg::LoadAlign(idx0, idx0Buf); + Reg::LoadAlign(idx1, idx1Buf); + Reg::LoadAlign(idx2, idx2Buf); + Reg::LoadAlign(idx3, idx3Buf); - MicroAPI::ShiftLefts(idx0, idx0, (int16_t)24, pregB32); - MicroAPI::ShiftLefts(idx1, idx1, (int16_t)16, pregB32); - MicroAPI::ShiftLefts(idx2, idx2, (int16_t)8, pregB32); + Reg::ShiftLefts(idx0, idx0, (int16_t)24, pregB32); + Reg::ShiftLefts(idx1, idx1, (int16_t)16, pregB32); + Reg::ShiftLefts(idx2, idx2, (int16_t)8, pregB32); // ADD - MicroAPI::Add(idx0, idx0, idx1, pregB32); - MicroAPI::Add(idx0, idx0, idx2, pregB32); - MicroAPI::Add(idx0, idx0, idx3, pregB32); + Reg::Add(idx0, idx0, idx1, pregB32); + Reg::Add(idx0, idx0, idx2, pregB32); + Reg::Add(idx0, idx0, idx3, pregB32); - MicroAPI::StoreAlign(kValue, idx0, pregB32); + Reg::StoreAlign(kValue, idx0, pregB32); } -__simd_vf__ void FindIdxGTOutputVFImpl(__ubuf__ uint32_t* outputIdxBuf, - __ubuf__ uint32_t* inputBuf, uint32_t beginIdx, - __ubuf__ uint32_t* kValue, uint16_t vfLoop) +__simd_vf__ void FindIdxGTOutputVFImpl(__ubuf__ uint32_t *outputIdxBuf, __ubuf__ uint32_t *inputBuf, uint32_t beginIdx, + __ubuf__ uint32_t *kValue, uint16_t vfLoop) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::ClearSpr(); + Reg::ClearSpr(); - MicroAPI::UnalignRegForStore alignIdx; + Reg::UnalignRegForStore alignIdx; - MicroAPI::RegTensor kthValue; - MicroAPI::LoadAlign(kthValue, kValue); + Reg::RegTensor kthValue; + Reg::LoadAlign(kthValue, kValue); - MicroAPI::RegTensor vregInput; + Reg::RegTensor vregInput; for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) { - MicroAPI::RegTensor idxC; - MicroAPI::Arange(idxC, beginIdx + i * 64); + Reg::RegTensor idxC; + Reg::Arange(idxC, beginIdx + i * 64); - MicroAPI::LoadAlign(vregInput, inputBuf + i * 64); + Reg::LoadAlign(vregInput, inputBuf + i * 64); - MicroAPI::MaskReg poutGT = MicroAPI::CreateMask(); + Reg::MaskReg poutGT = Reg::CreateMask(); - MicroAPI::RegTensor sqzIdxOut; - MicroAPI::Compare(poutGT, vregInput, kthValue, pregB32); + Reg::RegTensor sqzIdxOut; + Reg::Compare(poutGT, vregInput, kthValue, pregB32); - MicroAPI::Squeeze(sqzIdxOut, - (MicroAPI::RegTensor&)idxC, poutGT); - MicroAPI::StoreUnAlign(outputIdxBuf, sqzIdxOut, alignIdx); + Reg::Squeeze(sqzIdxOut, (Reg::RegTensor &)idxC, poutGT); + Reg::StoreUnAlign(outputIdxBuf, sqzIdxOut, alignIdx); } - MicroAPI::StoreUnAlignPost(outputIdxBuf, alignIdx); + Reg::StoreUnAlignPost(outputIdxBuf, alignIdx); } -__simd_vf__ void FindIdxEQOutputVFImpl(__ubuf__ uint32_t* outputIdxBuf, - __ubuf__ uint32_t* inputBuf, uint32_t beginIdx, - __ubuf__ uint32_t* kValue) +__simd_vf__ void FindIdxEQOutputVFImpl(__ubuf__ uint32_t *outputIdxBuf, __ubuf__ uint32_t *inputBuf, uint32_t beginIdx, + __ubuf__ uint32_t *kValue) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::UnalignRegForStore alignIdx; + Reg::UnalignRegForStore alignIdx; - MicroAPI::RegTensor kthValue; - MicroAPI::LoadAlign(kthValue, kValue); + Reg::RegTensor kthValue; + Reg::LoadAlign(kthValue, kValue); - MicroAPI::RegTensor vregInput; + Reg::RegTensor vregInput; - MicroAPI::RegTensor idxC; - MicroAPI::Arange(idxC, beginIdx); + Reg::RegTensor idxC; + Reg::Arange(idxC, beginIdx); - MicroAPI::LoadAlign(vregInput, inputBuf); + Reg::LoadAlign(vregInput, inputBuf); - MicroAPI::MaskReg poutEQ = MicroAPI::CreateMask(); + Reg::MaskReg poutEQ = Reg::CreateMask(); - MicroAPI::RegTensor sqzIdxOut; - MicroAPI::Compare(poutEQ, vregInput, kthValue, pregB32); + Reg::RegTensor sqzIdxOut; + Reg::Compare(poutEQ, vregInput, kthValue, pregB32); - MicroAPI::Squeeze(sqzIdxOut, - (MicroAPI::RegTensor&)idxC, poutEQ); - MicroAPI::StoreUnAlign(outputIdxBuf, sqzIdxOut, alignIdx); - MicroAPI::StoreUnAlignPost(outputIdxBuf, alignIdx); + Reg::Squeeze(sqzIdxOut, (Reg::RegTensor &)idxC, poutEQ); + Reg::StoreUnAlign(outputIdxBuf, sqzIdxOut, alignIdx); + Reg::StoreUnAlignPost(outputIdxBuf, alignIdx); } -__simd_vf__ void FindValueGTOutputVFImpl(__ubuf__ uint32_t* outputValueBuf, - __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* kValue, - uint16_t vfLoop) +__simd_vf__ void FindValueGTOutputVFImpl(__ubuf__ uint32_t *outputValueBuf, __ubuf__ uint32_t *inputBuf, + __ubuf__ uint32_t *kValue, uint16_t vfLoop) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::ClearSpr(); + Reg::ClearSpr(); - MicroAPI::UnalignRegForStore alignValue; + Reg::UnalignRegForStore alignValue; - MicroAPI::RegTensor kthValue; - MicroAPI::LoadAlign(kthValue, kValue); + Reg::RegTensor kthValue; + Reg::LoadAlign(kthValue, kValue); - MicroAPI::RegTensor vregInput; + Reg::RegTensor vregInput; for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) { - MicroAPI::LoadAlign(vregInput, inputBuf + i * 64); + Reg::LoadAlign(vregInput, inputBuf + i * 64); - MicroAPI::MaskReg poutGT = MicroAPI::CreateMask(); + Reg::MaskReg poutGT = Reg::CreateMask(); - MicroAPI::RegTensor sqzValueOut; - MicroAPI::Compare(poutGT, vregInput, kthValue, pregB32); + Reg::RegTensor sqzValueOut; + Reg::Compare(poutGT, vregInput, kthValue, pregB32); - MicroAPI::Squeeze(sqzValueOut, vregInput, poutGT); - MicroAPI::StoreUnAlign(outputValueBuf, - sqzValueOut, alignValue); + Reg::Squeeze(sqzValueOut, vregInput, poutGT); + Reg::StoreUnAlign(outputValueBuf, sqzValueOut, alignValue); } - MicroAPI::StoreUnAlignPost(outputValueBuf, alignValue); + Reg::StoreUnAlignPost(outputValueBuf, alignValue); } -__simd_vf__ void FindValueEQOutputVFImpl(__ubuf__ uint32_t* outputValueBuf, - __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* kValue) +__simd_vf__ void FindValueEQOutputVFImpl(__ubuf__ uint32_t *outputValueBuf, __ubuf__ uint32_t *inputBuf, + __ubuf__ uint32_t *kValue) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::UnalignRegForStore alignValue; + Reg::UnalignRegForStore alignValue; - MicroAPI::RegTensor kthValue; - MicroAPI::LoadAlign(kthValue, kValue); + Reg::RegTensor kthValue; + Reg::LoadAlign(kthValue, kValue); - MicroAPI::RegTensor vregInput; + Reg::RegTensor vregInput; - MicroAPI::LoadAlign(vregInput, inputBuf); + Reg::LoadAlign(vregInput, inputBuf); - MicroAPI::MaskReg poutEQ = MicroAPI::CreateMask(); + Reg::MaskReg poutEQ = Reg::CreateMask(); - MicroAPI::RegTensor sqzValueOut; - MicroAPI::Compare(poutEQ, vregInput, kthValue, pregB32); + Reg::RegTensor sqzValueOut; + Reg::Compare(poutEQ, vregInput, kthValue, pregB32); - MicroAPI::Squeeze(sqzValueOut, vregInput, poutEQ); - MicroAPI::StoreUnAlign(outputValueBuf, sqzValueOut, alignValue); - MicroAPI::StoreUnAlignPost(outputValueBuf, alignValue); + Reg::Squeeze(sqzValueOut, vregInput, poutEQ); + Reg::StoreUnAlign(outputValueBuf, sqzValueOut, alignValue); + Reg::StoreUnAlignPost(outputValueBuf, alignValue); } -__aicore__ inline void LiTopKVF(const LocalTensor& outputIdxLocal, - const LocalTensor& outputValueLocal, - const LocalTensor& inputLocal, - const LocalTensor& tmpIdxLocal, - const LocalTensor& tmpValueLocal, - const LocalTensor& histogramsLocal, - const LocalTensor& idx0Local, - const LocalTensor& idx1Local, - const LocalTensor& idx2Local, - const LocalTensor& idx3Local, - const LocalTensor& nkValueLocal, - uint32_t topK, - uint32_t s2SeqLen) +__aicore__ inline void LiTopKVF(const LocalTensor &outputIdxLocal, + const LocalTensor &outputValueLocal, const LocalTensor &inputLocal, + const LocalTensor &tmpIdxLocal, const LocalTensor &tmpValueLocal, + const LocalTensor &histogramsLocal, const LocalTensor &idx0Local, + const LocalTensor &idx1Local, const LocalTensor &idx2Local, + const LocalTensor &idx3Local, const LocalTensor &nkValueLocal, + uint32_t topK, uint32_t s2SeqLen) { - __ubuf__ uint32_t* outputIdxBuf = (__ubuf__ uint32_t*)outputIdxLocal.GetPhyAddr(); - __ubuf__ uint32_t* outputValueBuf = (__ubuf__ uint32_t*)outputValueLocal.GetPhyAddr(); - __ubuf__ uint32_t* inputBuf = (__ubuf__ uint32_t*)inputLocal.GetPhyAddr(); - __ubuf__ uint32_t* tmpIdxBuf = (__ubuf__ uint32_t*)tmpIdxLocal.GetPhyAddr(); - __ubuf__ uint32_t* tmpValueBuf = (__ubuf__ uint32_t*)tmpValueLocal.GetPhyAddr(); - __ubuf__ uint32_t* histogramsBuf = (__ubuf__ uint32_t*)histogramsLocal.GetPhyAddr(); - __ubuf__ uint32_t* idx0Buf = (__ubuf__ uint32_t*)idx0Local.GetPhyAddr(); - __ubuf__ uint32_t* idx1Buf = (__ubuf__ uint32_t*)idx1Local.GetPhyAddr(); - __ubuf__ uint32_t* idx2Buf = (__ubuf__ uint32_t*)idx2Local.GetPhyAddr(); - __ubuf__ uint32_t* idx3Buf = (__ubuf__ uint32_t*)idx3Local.GetPhyAddr(); - __ubuf__ uint32_t* nkValueBuf = (__ubuf__ uint32_t*)nkValueLocal.GetPhyAddr(); + __ubuf__ uint32_t *outputIdxBuf = (__ubuf__ uint32_t *)outputIdxLocal.GetPhyAddr(); + __ubuf__ uint32_t *outputValueBuf = (__ubuf__ uint32_t *)outputValueLocal.GetPhyAddr(); + __ubuf__ uint32_t *inputBuf = (__ubuf__ uint32_t *)inputLocal.GetPhyAddr(); + __ubuf__ uint32_t *tmpIdxBuf = (__ubuf__ uint32_t *)tmpIdxLocal.GetPhyAddr(); + __ubuf__ uint32_t *tmpValueBuf = (__ubuf__ uint32_t *)tmpValueLocal.GetPhyAddr(); + __ubuf__ uint32_t *histogramsBuf = (__ubuf__ uint32_t *)histogramsLocal.GetPhyAddr(); + __ubuf__ uint32_t *idx0Buf = (__ubuf__ uint32_t *)idx0Local.GetPhyAddr(); + __ubuf__ uint32_t *idx1Buf = (__ubuf__ uint32_t *)idx1Local.GetPhyAddr(); + __ubuf__ uint32_t *idx2Buf = (__ubuf__ uint32_t *)idx2Local.GetPhyAddr(); + __ubuf__ uint32_t *idx3Buf = (__ubuf__ uint32_t *)idx3Local.GetPhyAddr(); + __ubuf__ uint32_t *nkValueBuf = (__ubuf__ uint32_t *)nkValueLocal.GetPhyAddr(); uint32_t bottomK = s2SeqLen - topK + 1; uint32_t beginIdx = 0; @@ -598,12 +554,12 @@ __aicore__ inline void LiTopKVF(const LocalTensor& outputIdxLocal, int64_t arIdxNumPerLoop = AscendC::GetSpr(); if (((arIdxNumPerLoop - arIdxNum) / sizeof(uint32_t)) < remainIdxNum) { // 调用一次查找等于k-value情况的过程 - beginIdx = i * 64; // 64: 块起始偏移量 + beginIdx = i * 64; // 64: 块起始偏移量 FindIdxEQOutputVFImpl(outputIdxBuf, inputBuf + i * 64, beginIdx, nkValueBuf); // 64: 块起始偏移量 } else { break; } } } -} +} // namespace topkb32 #endif diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/vf_topk_16_gather_quant_v2.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/vf_topk_16_gather_quant_v2.h index 34b7753ed636..8f51c73efbf3 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/vf_topk_16_gather_quant_v2.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/arch35/vf/vf_topk_16_gather_quant_v2.h @@ -9,365 +9,353 @@  */ /*! -* \file vf_topk_16_gather_quant_v2.h -* \brief -*/ + * \file vf_topk_16_gather_quant_v2.h + * \brief + */ #ifndef VF_TOPK_16_GATHER_QUANT_V2_H #define VF_TOPK_16_GATHER_QUANT_V2_H namespace topkb16gather { -template -__simd_vf__ void HistogramsHighVFImpl(__ubuf__ uint32_t* histogramsBuf, __ubuf__ uint16_t* inputBuf, uint16_t vfLoop, - bool init) +template +__simd_vf__ void HistogramsHighVFImpl(__ubuf__ uint32_t *histogramsBuf, __ubuf__ uint16_t *inputBuf, uint16_t vfLoop, + bool init) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); + Reg::MaskReg pregB8 = Reg::CreateMask(); // 计算直方图cout0 0-127 cout1 128-255 - MicroAPI::RegTensor cout0; - MicroAPI::RegTensor cout1; - MicroAPI::Duplicate(cout0, 0); - MicroAPI::Duplicate(cout1, 0); + Reg::RegTensor cout0; + Reg::RegTensor cout1; + Reg::Duplicate(cout0, 0); + Reg::Duplicate(cout1, 0); - MicroAPI::RegTensor cout0U32Even; - MicroAPI::RegTensor cout0U32Odd; - MicroAPI::RegTensor cout1U32Even; - MicroAPI::RegTensor cout1U32Odd; + Reg::RegTensor cout0U32Even; + Reg::RegTensor cout0U32Odd; + Reg::RegTensor cout1U32Even; + Reg::RegTensor cout1U32Odd; - MicroAPI::RegTensor vregHigh; - MicroAPI::RegTensor vregLow; + Reg::RegTensor vregHigh; + Reg::RegTensor vregLow; - static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr Reg::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE, - MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr Reg::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {Reg::RegLayout::ONE, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; for (uint16_t i = 0; i < vfLoop; ++i) { - MicroAPI::LoadAlign(vregLow, vregHigh, inputBuf + i * 256); - - MicroAPI::Histograms(cout0, (MicroAPI::RegTensor&)vregHigh, - pregB8); - MicroAPI::Histograms(cout1, (MicroAPI::RegTensor&)vregHigh, - pregB8); + Reg::LoadAlign(vregLow, vregHigh, inputBuf + i * 256); + + Reg::Histograms( + cout0, (Reg::RegTensor &)vregHigh, pregB8); + Reg::Histograms( + cout1, (Reg::RegTensor &)vregHigh, pregB8); } - MicroAPI::Cast(cout0U32Even, cout0, pregB16); - MicroAPI::Cast(cout0U32Odd, cout0, pregB16); - MicroAPI::Cast(cout1U32Even, cout1, pregB16); - MicroAPI::Cast(cout1U32Odd, cout1, pregB16); + Reg::Cast(cout0U32Even, cout0, pregB16); + Reg::Cast(cout0U32Odd, cout0, pregB16); + Reg::Cast(cout1U32Even, cout1, pregB16); + Reg::Cast(cout1U32Odd, cout1, pregB16); - MicroAPI::StoreAlign(histogramsBuf, cout0U32Even, cout0U32Odd, - pregB32); - MicroAPI::StoreAlign(histogramsBuf + 128, cout1U32Even, cout1U32Odd, - pregB32); + Reg::StoreAlign(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32); + Reg::StoreAlign(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32); } -__simd_vf__ void FindHighTargetBinVFImpl(__ubuf__ uint32_t* idxHighBuf, __ubuf__ uint32_t* nkValueBuf, - __ubuf__ uint32_t* histogramsBuf, uint32_t bottomK) +__simd_vf__ void FindHighTargetBinVFImpl(__ubuf__ uint32_t *idxHighBuf, __ubuf__ uint32_t *nkValueBuf, + __ubuf__ uint32_t *histogramsBuf, uint32_t bottomK) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::MaskReg pregGE; + Reg::MaskReg pregGE; - MicroAPI::ClearSpr(); + Reg::ClearSpr(); - MicroAPI::UnalignRegForStore alignIdxHigh; + Reg::UnalignRegForStore alignIdxHigh; - MicroAPI::RegTensor btmK; - MicroAPI::Duplicate(btmK, bottomK); + Reg::RegTensor btmK; + Reg::Duplicate(btmK, bottomK); - MicroAPI::RegTensor idxC; - MicroAPI::RegTensor cout; - MicroAPI::RegTensor sqzIdxHigh; + Reg::RegTensor idxC; + Reg::RegTensor cout; + Reg::RegTensor sqzIdxHigh; for (uint16_t i = 0; i < (uint16_t)(4); ++i) { - MicroAPI::Arange(idxC, i * 64); + Reg::Arange(idxC, i * 64); - MicroAPI::LoadAlign(cout, histogramsBuf + i * 64); + Reg::LoadAlign(cout, histogramsBuf + i * 64); - MicroAPI::Compare(pregGE, cout, btmK, pregB32); + Reg::Compare(pregGE, cout, btmK, pregB32); - MicroAPI::Squeeze(sqzIdxHigh, - (MicroAPI::RegTensor&)idxC, pregGE); - MicroAPI::StoreUnAlign(idxHighBuf, sqzIdxHigh, alignIdxHigh); + Reg::Squeeze(sqzIdxHigh, (Reg::RegTensor &)idxC, pregGE); + Reg::StoreUnAlign(idxHighBuf, sqzIdxHigh, alignIdxHigh); } - MicroAPI::StoreUnAlignPost(idxHighBuf, alignIdxHigh); + Reg::StoreUnAlignPost(idxHighBuf, alignIdxHigh); - MicroAPI::LocalMemBar(); + Reg::LocalMemBar(); - MicroAPI::RegTensor idxHigh; - MicroAPI::LoadAlign(idxHigh, idxHighBuf); + Reg::RegTensor idxHigh; + Reg::LoadAlign(idxHigh, idxHighBuf); - MicroAPI::RegTensor idxAll1; - MicroAPI::RegTensor idxPrev0; - MicroAPI::RegTensor prevBinValue; - MicroAPI::Duplicate(idxAll1, 1); + Reg::RegTensor idxAll1; + Reg::RegTensor idxPrev0; + Reg::RegTensor prevBinValue; + Reg::Duplicate(idxAll1, 1); - MicroAPI::RegTensor zeroAll; - MicroAPI::Duplicate(zeroAll, 0); + Reg::RegTensor zeroAll; + Reg::Duplicate(zeroAll, 0); - MicroAPI::MaskReg preg0 = MicroAPI::CreateMask(); - MicroAPI::Compare(preg0, idxHigh, zeroAll, pregB32); - MicroAPI::Sub(idxPrev0, idxHigh, (MicroAPI::RegTensor&)idxAll1, pregB32); - MicroAPI::ShiftRights(idxPrev0, idxPrev0, (int16_t)24, pregB32); + Reg::MaskReg preg0 = Reg::CreateMask(); + Reg::Compare(preg0, idxHigh, zeroAll, pregB32); + Reg::Sub(idxPrev0, idxHigh, (Reg::RegTensor &)idxAll1, pregB32); + Reg::ShiftRights(idxPrev0, idxPrev0, (int16_t)24, pregB32); - MicroAPI::Gather(prevBinValue, histogramsBuf, idxPrev0, pregB32); - MicroAPI::Select(prevBinValue, zeroAll, prevBinValue, preg0); + Reg::Gather(prevBinValue, histogramsBuf, idxPrev0, pregB32); + Reg::Select(prevBinValue, zeroAll, prevBinValue, preg0); - MicroAPI::RegTensor nextK; - MicroAPI::Sub(nextK, btmK, prevBinValue, pregB32); - MicroAPI::StoreAlign(nkValueBuf, nextK, pregB32); + Reg::RegTensor nextK; + Reg::Sub(nextK, btmK, prevBinValue, pregB32); + Reg::StoreAlign(nkValueBuf, nextK, pregB32); } -template -__simd_vf__ void HistogramsLowVFImpl(__ubuf__ uint32_t* histogramsBuf, __ubuf__ uint16_t* inputBuf, - __ubuf__ uint32_t* idxHighBuf, uint16_t vfLoop, bool init) +template +__simd_vf__ void HistogramsLowVFImpl(__ubuf__ uint32_t *histogramsBuf, __ubuf__ uint16_t *inputBuf, + __ubuf__ uint32_t *idxHighBuf, uint16_t vfLoop, bool init) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); + Reg::MaskReg pregB8 = Reg::CreateMask(); - MicroAPI::MaskReg pregEQ; + Reg::MaskReg pregEQ; // 计算直方图0-127 128-255 - MicroAPI::RegTensor cout0; - MicroAPI::RegTensor cout1; - MicroAPI::Duplicate(cout0, 0); - MicroAPI::Duplicate(cout1, 0); + Reg::RegTensor cout0; + Reg::RegTensor cout1; + Reg::Duplicate(cout0, 0); + Reg::Duplicate(cout1, 0); - MicroAPI::RegTensor cout0U32Even; - MicroAPI::RegTensor cout0U32Odd; - MicroAPI::RegTensor cout1U32Even; - MicroAPI::RegTensor cout1U32Odd; + Reg::RegTensor cout0U32Even; + Reg::RegTensor cout0U32Odd; + Reg::RegTensor cout1U32Even; + Reg::RegTensor cout1U32Odd; - MicroAPI::RegTensor idxHigh; - MicroAPI::LoadAlign(idxHigh, idxHighBuf); + Reg::RegTensor idxHigh; + Reg::LoadAlign(idxHigh, idxHighBuf); - MicroAPI::RegTensor vregHigh; - MicroAPI::RegTensor vregLow; + Reg::RegTensor vregHigh; + Reg::RegTensor vregLow; - static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr Reg::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {Reg::RegLayout::ZERO, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; - static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE, - MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr Reg::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {Reg::RegLayout::ONE, Reg::SatMode::UNKNOWN, + Reg::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; for (uint16_t i = 0; i < vfLoop; ++i) { - MicroAPI::LoadAlign(vregLow, vregHigh, inputBuf + i * 256); + Reg::LoadAlign(vregLow, vregHigh, inputBuf + i * 256); - MicroAPI::Compare(pregEQ, (MicroAPI::RegTensor&)vregHigh, - (MicroAPI::RegTensor&)idxHigh, pregB8); + Reg::Compare(pregEQ, (Reg::RegTensor &)vregHigh, + (Reg::RegTensor &)idxHigh, pregB8); - MicroAPI::Histograms(cout0, (MicroAPI::RegTensor&)vregLow, - pregEQ); - MicroAPI::Histograms(cout1, (MicroAPI::RegTensor&)vregLow, - pregEQ); + Reg::Histograms( + cout0, (Reg::RegTensor &)vregLow, pregEQ); + Reg::Histograms( + cout1, (Reg::RegTensor &)vregLow, pregEQ); } - MicroAPI::Cast(cout0U32Even, cout0, pregB16); - MicroAPI::Cast(cout0U32Odd, cout0, pregB16); - MicroAPI::Cast(cout1U32Even, cout1, pregB16); - MicroAPI::Cast(cout1U32Odd, cout1, pregB16); + Reg::Cast(cout0U32Even, cout0, pregB16); + Reg::Cast(cout0U32Odd, cout0, pregB16); + Reg::Cast(cout1U32Even, cout1, pregB16); + Reg::Cast(cout1U32Odd, cout1, pregB16); - MicroAPI::StoreAlign(histogramsBuf, cout0U32Even, cout0U32Odd, - pregB32); - MicroAPI::StoreAlign(histogramsBuf + 128, cout1U32Even, cout1U32Odd, - pregB32); + Reg::StoreAlign(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32); + Reg::StoreAlign(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32); } -__simd_vf__ void FindKthVFImpl(__ubuf__ uint32_t* kValue, __ubuf__ uint32_t* histogramsBuf, - __ubuf__ uint32_t* idxHighBuf, __ubuf__ uint32_t* idxLowBuf) +__simd_vf__ void FindKthVFImpl(__ubuf__ uint32_t *kValue, __ubuf__ uint32_t *histogramsBuf, + __ubuf__ uint32_t *idxHighBuf, __ubuf__ uint32_t *idxLowBuf) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); - MicroAPI::MaskReg pregGE; + Reg::MaskReg pregGE; - MicroAPI::ClearSpr(); + Reg::ClearSpr(); - MicroAPI::UnalignRegForStore alignIdxLow; + Reg::UnalignRegForStore alignIdxLow; - MicroAPI::RegTensor btmK; - MicroAPI::LoadAlign(btmK, kValue); + Reg::RegTensor btmK; + Reg::LoadAlign(btmK, kValue); - MicroAPI::RegTensor idxC; - MicroAPI::RegTensor cout; - MicroAPI::RegTensor sqzIdxLow; + Reg::RegTensor idxC; + Reg::RegTensor cout; + Reg::RegTensor sqzIdxLow; for (uint16_t i = 0; i < (uint16_t)(4); ++i) { - MicroAPI::Arange(idxC, i * 64); + Reg::Arange(idxC, i * 64); - MicroAPI::LoadAlign(cout, histogramsBuf + i * 64); + Reg::LoadAlign(cout, histogramsBuf + i * 64); - MicroAPI::Compare(pregGE, cout, btmK, pregB32); + Reg::Compare(pregGE, cout, btmK, pregB32); - MicroAPI::Squeeze(sqzIdxLow, - (MicroAPI::RegTensor&)idxC, pregGE); - MicroAPI::StoreUnAlign(idxLowBuf, sqzIdxLow, alignIdxLow); + Reg::Squeeze(sqzIdxLow, (Reg::RegTensor &)idxC, pregGE); + Reg::StoreUnAlign(idxLowBuf, sqzIdxLow, alignIdxLow); } - MicroAPI::StoreUnAlignPost(idxLowBuf, alignIdxLow); + Reg::StoreUnAlignPost(idxLowBuf, alignIdxLow); - MicroAPI::LocalMemBar(); + Reg::LocalMemBar(); - MicroAPI::RegTensor idxHigh; - MicroAPI::RegTensor idxLow; - MicroAPI::LoadAlign(idxHigh, idxHighBuf); - MicroAPI::LoadAlign(idxLow, idxLowBuf); + Reg::RegTensor idxHigh; + Reg::RegTensor idxLow; + Reg::LoadAlign(idxHigh, idxHighBuf); + Reg::LoadAlign(idxLow, idxLowBuf); - MicroAPI::RegTensor idxTmp; - MicroAPI::Duplicate(idxTmp, 0xff00); + Reg::RegTensor idxTmp; + Reg::Duplicate(idxTmp, 0xff00); - MicroAPI::And(idxHigh, idxHigh, (MicroAPI::RegTensor&)idxTmp, pregB32); + Reg::And(idxHigh, idxHigh, (Reg::RegTensor &)idxTmp, pregB32); - MicroAPI::RegTensor idxK; - MicroAPI::Add(idxK, idxHigh, idxLow, pregB16); + Reg::RegTensor idxK; + Reg::Add(idxK, idxHigh, idxLow, pregB16); - MicroAPI::StoreAlign(kValue, idxK, pregB32); + Reg::StoreAlign(kValue, idxK, pregB32); } /** 输出所有大于的kth-value的Index */ -__simd_vf__ void FindIdxGTOutputVFImpl(__ubuf__ uint16_t* outputIdxBuf, __ubuf__ uint16_t* inputValueBuf, - uint16_t beginIdx, __ubuf__ uint32_t* kValue, uint16_t vfLoop) +__simd_vf__ void FindIdxGTOutputVFImpl(__ubuf__ uint16_t *outputIdxBuf, __ubuf__ uint16_t *inputValueBuf, + uint16_t beginIdx, __ubuf__ uint32_t *kValue, uint16_t vfLoop) { - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); - MicroAPI::MaskReg poutGT; + Reg::MaskReg poutGT; - MicroAPI::ClearSpr(); + Reg::ClearSpr(); - MicroAPI::UnalignRegForStore alignIdx; + Reg::UnalignRegForStore alignIdx; - MicroAPI::RegTensor kthValue; - MicroAPI::LoadAlign(kthValue, kValue); + Reg::RegTensor kthValue; + Reg::LoadAlign(kthValue, kValue); - MicroAPI::RegTensor vregInput; - MicroAPI::RegTensor idxC; - MicroAPI::RegTensor sqzIdxOut; + Reg::RegTensor vregInput; + Reg::RegTensor idxC; + Reg::RegTensor sqzIdxOut; for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) { - MicroAPI::Arange(idxC, beginIdx + i * 128); + Reg::Arange(idxC, beginIdx + i * 128); - MicroAPI::LoadAlign(vregInput, inputValueBuf + i * 128); + Reg::LoadAlign(vregInput, inputValueBuf + i * 128); - MicroAPI::Compare(poutGT, vregInput, (MicroAPI::RegTensor&)kthValue, pregB16); + Reg::Compare(poutGT, vregInput, (Reg::RegTensor &)kthValue, pregB16); - MicroAPI::Squeeze(sqzIdxOut, - (MicroAPI::RegTensor&)idxC, poutGT); - MicroAPI::StoreUnAlign(outputIdxBuf, sqzIdxOut, alignIdx); + Reg::Squeeze(sqzIdxOut, (Reg::RegTensor &)idxC, poutGT); + Reg::StoreUnAlign(outputIdxBuf, sqzIdxOut, alignIdx); } - MicroAPI::StoreUnAlignPost(outputIdxBuf, alignIdx); + Reg::StoreUnAlignPost(outputIdxBuf, alignIdx); } /** 输出所有等于的kth-value的Index */ -__simd_vf__ void FindIdxEQOutputVFImpl(__ubuf__ uint16_t* outputIdxBuf, __ubuf__ uint16_t* inputValueBuf, - uint16_t beginIdx, __ubuf__ uint32_t* kValue, uint16_t vfLoop) +__simd_vf__ void FindIdxEQOutputVFImpl(__ubuf__ uint16_t *outputIdxBuf, __ubuf__ uint16_t *inputValueBuf, + uint16_t beginIdx, __ubuf__ uint32_t *kValue, uint16_t vfLoop) { - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); - MicroAPI::MaskReg poutEQ; + Reg::MaskReg poutEQ; - MicroAPI::UnalignRegForStore alignIdx; + Reg::UnalignRegForStore alignIdx; - MicroAPI::RegTensor kthValue; - MicroAPI::LoadAlign(kthValue, kValue); + Reg::RegTensor kthValue; + Reg::LoadAlign(kthValue, kValue); - MicroAPI::RegTensor vregInput; - MicroAPI::RegTensor idxC; - MicroAPI::RegTensor sqzIdxOut; + Reg::RegTensor vregInput; + Reg::RegTensor idxC; + Reg::RegTensor sqzIdxOut; for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) { - MicroAPI::Arange(idxC, beginIdx + i * 128); + Reg::Arange(idxC, beginIdx + i * 128); - MicroAPI::LoadAlign(vregInput, inputValueBuf + i * 128); + Reg::LoadAlign(vregInput, inputValueBuf + i * 128); - MicroAPI::Compare(poutEQ, vregInput, (MicroAPI::RegTensor&)kthValue, pregB16); + Reg::Compare(poutEQ, vregInput, (Reg::RegTensor &)kthValue, pregB16); - MicroAPI::Squeeze(sqzIdxOut, - (MicroAPI::RegTensor&)idxC, poutEQ); - MicroAPI::StoreUnAlign(outputIdxBuf, sqzIdxOut, alignIdx); + Reg::Squeeze(sqzIdxOut, (Reg::RegTensor &)idxC, poutEQ); + Reg::StoreUnAlign(outputIdxBuf, sqzIdxOut, alignIdx); } - MicroAPI::StoreUnAlignPost(outputIdxBuf, alignIdx); + Reg::StoreUnAlignPost(outputIdxBuf, alignIdx); } /** 输出最终的Value */ -__simd_vf__ void FindValueOutputVFImpl(__ubuf__ uint16_t* outputValueBuf, __ubuf__ uint16_t* inputValueBuf, - __ubuf__ uint16_t* tmpIdxBuf, uint16_t vfLoop) +__simd_vf__ void FindValueOutputVFImpl(__ubuf__ uint16_t *outputValueBuf, __ubuf__ uint16_t *inputValueBuf, + __ubuf__ uint16_t *tmpIdxBuf, uint16_t vfLoop) { - MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask(); + Reg::MaskReg pregB16 = Reg::CreateMask(); - MicroAPI::RegTensor tmpIdx; - MicroAPI::RegTensor outputValue; + Reg::RegTensor tmpIdx; + Reg::RegTensor outputValue; for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) { - MicroAPI::LoadAlign(tmpIdx, tmpIdxBuf + i * 128); + Reg::LoadAlign(tmpIdx, tmpIdxBuf + i * 128); - MicroAPI::Gather(outputValue, inputValueBuf, tmpIdx, pregB16); + Reg::Gather(outputValue, inputValueBuf, tmpIdx, pregB16); - MicroAPI::StoreAlign(outputValueBuf + i * 128, outputValue, pregB16); + Reg::StoreAlign(outputValueBuf + i * 128, outputValue, pregB16); } } /** 输出最终的Idx */ -__simd_vf__ void FindRealIndexVFImpl(__ubuf__ uint32_t* outputIdxBuf, __ubuf__ uint16_t* tmpIdxBuf, - __ubuf__ uint32_t* hisIdxBuf, uint32_t topK, uint32_t loopIndex, uint16_t vfLoop) +__simd_vf__ void FindRealIndexVFImpl(__ubuf__ uint32_t *outputIdxBuf, __ubuf__ uint16_t *tmpIdxBuf, + __ubuf__ uint32_t *hisIdxBuf, uint32_t topK, uint32_t loopIndex, uint16_t vfLoop) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::MaskReg pregNow; - MicroAPI::MaskReg pregHis; + Reg::MaskReg pregNow; + Reg::MaskReg pregHis; - MicroAPI::RegTensor tmpIdx; - MicroAPI::RegTensor outputGatherIdx; - MicroAPI::RegTensor outputAddsIdx; + Reg::RegTensor tmpIdx; + Reg::RegTensor outputGatherIdx; + Reg::RegTensor outputAddsIdx; for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) { - MicroAPI::LoadAlign(tmpIdx, tmpIdxBuf + i * 64); + Reg::LoadAlign(tmpIdx, tmpIdxBuf + i * 64); - MicroAPI::Compares(pregNow, (MicroAPI::RegTensor&)tmpIdx, topK - 1, pregB32); - MicroAPI::Xor(pregHis, pregNow, pregB32, pregB32); + Reg::Compares(pregNow, (Reg::RegTensor &)tmpIdx, topK - 1, pregB32); + Reg::Xor(pregHis, pregNow, pregB32, pregB32); - MicroAPI::Gather(outputGatherIdx, hisIdxBuf, (MicroAPI::RegTensor&)tmpIdx, pregHis); - MicroAPI::Adds(outputAddsIdx, (MicroAPI::RegTensor&)tmpIdx, loopIndex, pregNow); + Reg::Gather(outputGatherIdx, hisIdxBuf, (Reg::RegTensor &)tmpIdx, pregHis); + Reg::Adds(outputAddsIdx, (Reg::RegTensor &)tmpIdx, loopIndex, pregNow); - MicroAPI::Add(outputGatherIdx, outputGatherIdx, outputAddsIdx, pregB32); + Reg::Add(outputGatherIdx, outputGatherIdx, outputAddsIdx, pregB32); - MicroAPI::StoreAlign(outputIdxBuf + i * 64, outputGatherIdx, pregB32); + Reg::StoreAlign(outputIdxBuf + i * 64, outputGatherIdx, pregB32); } } -__simd_vf__ void IndicesAddOffsetVF(__ubuf__ uint32_t* indicesOutBuf, uint32_t outputIdxOffset, uint32_t vfLoop) +__simd_vf__ void IndicesAddOffsetVF(__ubuf__ uint32_t *indicesOutBuf, uint32_t outputIdxOffset, uint32_t vfLoop) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::RegTensor outIndices; + Reg::RegTensor outIndices; for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) { - MicroAPI::LoadAlign(outIndices, indicesOutBuf + i * 64); - MicroAPI::Adds(outIndices, outIndices, outputIdxOffset, pregB32); - MicroAPI::StoreAlign(indicesOutBuf + i * 64, outIndices, pregB32); + Reg::LoadAlign(outIndices, indicesOutBuf + i * 64); + Reg::Adds(outIndices, outIndices, outputIdxOffset, pregB32); + Reg::StoreAlign(indicesOutBuf + i * 64, outIndices, pregB32); } } -__aicore__ inline void IndicesAddOffset(const LocalTensor& indicesOutLocal, - uint32_t outputIdxOffset, uint32_t topK) +__aicore__ inline void IndicesAddOffset(const LocalTensor &indicesOutLocal, uint32_t outputIdxOffset, + uint32_t topK) { - __ubuf__ uint32_t* indicesOutBuf = (__ubuf__ uint32_t*)indicesOutLocal.GetPhyAddr(); + __ubuf__ uint32_t *indicesOutBuf = (__ubuf__ uint32_t *)indicesOutLocal.GetPhyAddr(); const uint16_t repeatSize32 = 64; uint16_t topkLoopNum32 = (topK + repeatSize32 - 1) / repeatSize32; IndicesAddOffsetVF(indicesOutBuf, outputIdxOffset, topkLoopNum32); @@ -385,24 +373,20 @@ __aicore__ inline void IndicesAddOffset(const LocalTensor& indicesOutL * @param topK topK元素 * @param validLen 有效元素个数:QLIV2Common::Align(topkCountAlign256_ + validTrunkLen, (uint32_t)256) */ -template // 是否输出VALUE -__aicore__ inline void LiTopKVF(const LocalTensor& tmpIdxLocal, - const LocalTensor& outputValueLocal, - const LocalTensor& inputValueLocal, - const LocalTensor& histogramsLocal, - const LocalTensor& idxHighLocal, - const LocalTensor& idxLowLocal, - const LocalTensor& nkValueLocal, - uint32_t topK, - uint32_t validLen) +template // 是否输出VALUE +__aicore__ inline void LiTopKVF(const LocalTensor &tmpIdxLocal, const LocalTensor &outputValueLocal, + const LocalTensor &inputValueLocal, + const LocalTensor &histogramsLocal, const LocalTensor &idxHighLocal, + const LocalTensor &idxLowLocal, const LocalTensor &nkValueLocal, + uint32_t topK, uint32_t validLen) { - __ubuf__ uint16_t* tmpIdxBuf = (__ubuf__ uint16_t*)tmpIdxLocal.GetPhyAddr(); - __ubuf__ uint16_t* outputValueBuf = (__ubuf__ uint16_t*)outputValueLocal.GetPhyAddr(); - __ubuf__ uint16_t* inputValueBuf = (__ubuf__ uint16_t*)inputValueLocal.GetPhyAddr(); - __ubuf__ uint32_t* histogramsBuf = (__ubuf__ uint32_t*)histogramsLocal.GetPhyAddr(); - __ubuf__ uint32_t* idxHighBuf = (__ubuf__ uint32_t*)idxHighLocal.GetPhyAddr(); - __ubuf__ uint32_t* idxLowBuf = (__ubuf__ uint32_t*)idxLowLocal.GetPhyAddr(); - __ubuf__ uint32_t* nkValueBuf = (__ubuf__ uint32_t*)nkValueLocal.GetPhyAddr(); + __ubuf__ uint16_t *tmpIdxBuf = (__ubuf__ uint16_t *)tmpIdxLocal.GetPhyAddr(); + __ubuf__ uint16_t *outputValueBuf = (__ubuf__ uint16_t *)outputValueLocal.GetPhyAddr(); + __ubuf__ uint16_t *inputValueBuf = (__ubuf__ uint16_t *)inputValueLocal.GetPhyAddr(); + __ubuf__ uint32_t *histogramsBuf = (__ubuf__ uint32_t *)histogramsLocal.GetPhyAddr(); + __ubuf__ uint32_t *idxHighBuf = (__ubuf__ uint32_t *)idxHighLocal.GetPhyAddr(); + __ubuf__ uint32_t *idxLowBuf = (__ubuf__ uint32_t *)idxLowLocal.GetPhyAddr(); + __ubuf__ uint32_t *nkValueBuf = (__ubuf__ uint32_t *)nkValueLocal.GetPhyAddr(); uint32_t bottomK = validLen - topK + 1; uint32_t beginIdx = 0; @@ -440,20 +424,20 @@ __aicore__ inline void LiTopKVF(const LocalTensor& tmpIdxLocal, /** LD:输出最终的Idx */ -__simd_vf__ void FindLDRealIndexVFImpl(__ubuf__ uint32_t* outputIdxBuf, __ubuf__ uint16_t* tmpIdxBuf, - __ubuf__ uint32_t* hisIdxBuf, uint16_t vfLoop) +__simd_vf__ void FindLDRealIndexVFImpl(__ubuf__ uint32_t *outputIdxBuf, __ubuf__ uint16_t *tmpIdxBuf, + __ubuf__ uint32_t *hisIdxBuf, uint16_t vfLoop) { - MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask(); + Reg::MaskReg pregB32 = Reg::CreateMask(); - MicroAPI::RegTensor tmpIdx; - MicroAPI::RegTensor outputIdx; + Reg::RegTensor tmpIdx; + Reg::RegTensor outputIdx; for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) { - MicroAPI::LoadAlign(tmpIdx, tmpIdxBuf + i * 64); + Reg::LoadAlign(tmpIdx, tmpIdxBuf + i * 64); - MicroAPI::Gather(outputIdx, hisIdxBuf, (MicroAPI::RegTensor&)tmpIdx, pregB32); + Reg::Gather(outputIdx, hisIdxBuf, (Reg::RegTensor &)tmpIdx, pregB32); - MicroAPI::StoreAlign(outputIdxBuf + i * 64, outputIdx, pregB32); + Reg::StoreAlign(outputIdxBuf + i * 64, outputIdx, pregB32); } } @@ -468,20 +452,18 @@ __simd_vf__ void FindLDRealIndexVFImpl(__ubuf__ uint32_t* outputIdxBuf, __ubuf__ * @param loopBasicIdx 当前循环需要加上得基准Index * @param validLen 有效元素个数 */ -__aicore__ inline void LiTopKGatherVF(const LocalTensor& outputIdxLocal, - const LocalTensor& outputValueLocal, - const LocalTensor& inputValueLocal, - const LocalTensor& tmpIdxLocal, - const LocalTensor& hisIdxLocal, - uint32_t topK, - uint32_t loopBasicIdx, +__aicore__ inline void LiTopKGatherVF(const LocalTensor &outputIdxLocal, + const LocalTensor &outputValueLocal, + const LocalTensor &inputValueLocal, + const LocalTensor &tmpIdxLocal, + const LocalTensor &hisIdxLocal, uint32_t topK, uint32_t loopBasicIdx, uint32_t validLen) { - __ubuf__ uint32_t* outputIdxBuf = (__ubuf__ uint32_t*)outputIdxLocal.GetPhyAddr(); - __ubuf__ uint16_t* outputValueBuf = (__ubuf__ uint16_t*)outputValueLocal.GetPhyAddr(); - __ubuf__ uint16_t* inputValueBuf = (__ubuf__ uint16_t*)inputValueLocal.GetPhyAddr(); - __ubuf__ uint16_t* tmpIdxBuf = (__ubuf__ uint16_t*)tmpIdxLocal.GetPhyAddr(); - __ubuf__ uint32_t* hisIdxBuf = (__ubuf__ uint32_t*)hisIdxLocal.GetPhyAddr(); + __ubuf__ uint32_t *outputIdxBuf = (__ubuf__ uint32_t *)outputIdxLocal.GetPhyAddr(); + __ubuf__ uint16_t *outputValueBuf = (__ubuf__ uint16_t *)outputValueLocal.GetPhyAddr(); + __ubuf__ uint16_t *inputValueBuf = (__ubuf__ uint16_t *)inputValueLocal.GetPhyAddr(); + __ubuf__ uint16_t *tmpIdxBuf = (__ubuf__ uint16_t *)tmpIdxLocal.GetPhyAddr(); + __ubuf__ uint32_t *hisIdxBuf = (__ubuf__ uint32_t *)hisIdxLocal.GetPhyAddr(); const uint16_t repeatSize32 = 64; const uint16_t repeatSize16 = 128; @@ -494,14 +476,14 @@ __aicore__ inline void LiTopKGatherVF(const LocalTensor& outputIdxLoca /** LD:gather最终的Idx */ -__aicore__ inline void LiTopKLDGatherVF(const LocalTensor& outputIdxLocal, // 输出Idx topK * 2B - const LocalTensor& tmpIdxLocal, // 本轮tmpIdx输入 validLen * 2B - const LocalTensor& hisIdxLocal, // 上一轮Idx输入 topK * 4B - uint32_t topK) // topK元素个数 +__aicore__ inline void LiTopKLDGatherVF(const LocalTensor &outputIdxLocal, // 输出Idx topK * 2B + const LocalTensor &tmpIdxLocal, // 本轮tmpIdx输入 validLen * 2B + const LocalTensor &hisIdxLocal, // 上一轮Idx输入 topK * 4B + uint32_t topK) // topK元素个数 { - __ubuf__ uint32_t* outputIdxBuf = (__ubuf__ uint32_t*)outputIdxLocal.GetPhyAddr(); - __ubuf__ uint16_t* tmpIdxBuf = (__ubuf__ uint16_t*)tmpIdxLocal.GetPhyAddr(); - __ubuf__ uint32_t* hisIdxBuf = (__ubuf__ uint32_t*)hisIdxLocal.GetPhyAddr(); + __ubuf__ uint32_t *outputIdxBuf = (__ubuf__ uint32_t *)outputIdxLocal.GetPhyAddr(); + __ubuf__ uint16_t *tmpIdxBuf = (__ubuf__ uint16_t *)tmpIdxLocal.GetPhyAddr(); + __ubuf__ uint32_t *hisIdxBuf = (__ubuf__ uint32_t *)hisIdxLocal.GetPhyAddr(); const uint16_t repeatSize32 = 64; const uint16_t repeatSize16 = 128; @@ -510,5 +492,5 @@ __aicore__ inline void LiTopKLDGatherVF(const LocalTensor& outputIdxLo FindLDRealIndexVFImpl(outputIdxBuf, tmpIdxBuf, hisIdxBuf, topkLoopNum32); } -} -#endif // VF_TOPK_16_GATHER_QUANT_V2_H \ No newline at end of file +} // namespace topkb16gather +#endif // VF_TOPK_16_GATHER_QUANT_V2_H diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2.cpp b/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2.cpp index 8c9e7fcb7db9..211e8ab3a888 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2.cpp +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2.cpp @@ -32,12 +32,23 @@ using namespace optiling::detail; op.Process(); \ } while (0) +// arch22: 含 candidate_topk_index 输入/输出 (两级TopK) +#define INVOKE_LI_CANDIDATE_OP_IMPL(templateClass, ...) \ + do { \ + templateClass> op; \ + op.Init(query, key, weights, queryScale, keyScale, cuSeqlensQ, cuSeqlensK, sequsedQ, sequsedK, cmpResidualK, \ + blockTable, outputIdxOffset, metadata, candidateTopkIndex, sparseIndices, sparseValues, \ + candidateTopkIndexOut, user, tiling_data, &tPipe); \ + op.Process(); \ + } while (0) + template __global__ __aicore__ void quant_lightning_indexer_v2( __gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights, __gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale, __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensK, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sequsedK, __gm__ uint8_t *cmpResidualK, __gm__ uint8_t *blockTable, __gm__ uint8_t *outputIdxOffset, - __gm__ uint8_t *metadata, __gm__ uint8_t *sparseIndices, __gm__ uint8_t *sparseValues, __gm__ uint8_t *workspace, + __gm__ uint8_t *metadata, __gm__ uint8_t *candidateTopkIndex, __gm__ uint8_t *sparseIndices, + __gm__ uint8_t *sparseValues, __gm__ uint8_t *candidateTopkIndexOut, __gm__ uint8_t *workspace, __gm__ uint8_t *tiling) { TPipe tPipe; @@ -72,7 +83,7 @@ __global__ __aicore__ void quant_lightning_indexer_v2( } #else - INVOKE_LI_NO_KFC_OP_IMPL(QLIV2Preload, int8_t, int8_t, float, uint16_t, int32_t, PAGE_ATTENTION, - LI_LAYOUT(Q_LAYOUT_T), LI_LAYOUT(K_LAYOUT_T)); + INVOKE_LI_CANDIDATE_OP_IMPL(QLIV2Preload, int8_t, int8_t, float, uint16_t, int32_t, PAGE_ATTENTION, + LI_LAYOUT(Q_LAYOUT_T), LI_LAYOUT(K_LAYOUT_T)); #endif } diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_metadata.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_metadata.h index 17a97a122639..85b4f31e5467 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_metadata.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_metadata.h @@ -9,9 +9,9 @@ */ /*! -* \file quant_lightning_indexer_v2_metadata.h -* \brief -*/ + * \file quant_lightning_indexer_v2_metadata.h + * \brief + */ #ifndef QUANT_LIGHTNING_INDEXER_V2_METADATA_H #define QUANT_LIGHTNING_INDEXER_V2_METADATA_H @@ -48,7 +48,7 @@ inline constexpr uint32_t QLD_V2_WORKSPACE_NUM_INDEX = 4; inline constexpr uint32_t QLD_V2_M_START_INDEX = 5; inline constexpr uint32_t QLD_V2_M_NUM_INDEX = 6; - /** +/** * @brief 获取属性的绝对索引 * @param coreIdx 核索引 * @param metaIdx 元数据索引 @@ -67,13 +67,13 @@ __aicore__ inline uint32_t GetAttrAbsIndex(uint32_t coreIdx, uint32_t metaIdx, b #endif namespace detail { - struct QliV2Metadata { - uint32_t qliV2Metadata[AIC_CORE_MAX_NUM][QLI_V2_METADATA_SIZE]; - uint32_t qldV2Metadata[AIV_CORE_MAX_NUM][QLD_V2_METADATA_SIZE]; - }; +struct QliV2Metadata { + uint32_t qliV2Metadata[AIC_CORE_MAX_NUM][QLI_V2_METADATA_SIZE]; + uint32_t qldV2Metadata[AIV_CORE_MAX_NUM][QLD_V2_METADATA_SIZE]; }; +}; // namespace detail static_assert(QLI_V2_METADATA_TOTAL_SIZE * sizeof(QLI_V2_METADATA_T) >= sizeof(detail::QliV2Metadata)); -}; +}; // namespace optiling -#endif // QUANT_LIGHTNING_INDEXER_V2_METADATA_H +#endif // QUANT_LIGHTNING_INDEXER_V2_METADATA_H diff --git a/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_template_tiling_key.h b/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_template_tiling_key.h index 03c8dd244620..00c692fcb341 100644 --- a/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_template_tiling_key.h +++ b/csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_template_tiling_key.h @@ -31,16 +31,14 @@ // 模板参数支持的范围定义 #if (__CCE_AICORE__ == 310) -ASCENDC_TPL_ARGS_DECL(QuantLightningIndexerV2, // 算子OpType - ASCENDC_TPL_DTYPE_DECL(DT_Q, QLIV2_TPL_FLOAT8_E4M3FN, QLIV2_TPL_HIFLOAT8, - QLIV2_TPL_FLOAT4_E2M1, QLIV2_TPL_INT8), - ASCENDC_TPL_DTYPE_DECL(DT_K, QLIV2_TPL_FLOAT8_E4M3FN, QLIV2_TPL_HIFLOAT8, - QLIV2_TPL_FLOAT4_E2M1, QLIV2_TPL_INT8), - ASCENDC_TPL_DTYPE_DECL(DT_OUT, QLIV2_TPL_INT32), ASCENDC_TPL_BOOL_DECL(PAGE_ATTENTION, 1, 0), - ASCENDC_TPL_UINT_DECL(Q_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_BSND, - QLIV2_LAYOUT_TND), - ASCENDC_TPL_UINT_DECL(K_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_BSND, - QLIV2_LAYOUT_TND, QLIV2_LAYOUT_PA_BBND), ); +ASCENDC_TPL_ARGS_DECL( + QuantLightningIndexerV2, // 算子OpType + ASCENDC_TPL_DTYPE_DECL(DT_Q, QLIV2_TPL_FLOAT8_E4M3FN, QLIV2_TPL_HIFLOAT8, QLIV2_TPL_FLOAT4_E2M1, QLIV2_TPL_INT8), + ASCENDC_TPL_DTYPE_DECL(DT_K, QLIV2_TPL_FLOAT8_E4M3FN, QLIV2_TPL_HIFLOAT8, QLIV2_TPL_FLOAT4_E2M1, QLIV2_TPL_INT8), + ASCENDC_TPL_DTYPE_DECL(DT_OUT, QLIV2_TPL_INT32), ASCENDC_TPL_BOOL_DECL(PAGE_ATTENTION, 1, 0), + ASCENDC_TPL_UINT_DECL(Q_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_BSND, QLIV2_LAYOUT_TND), + ASCENDC_TPL_UINT_DECL(K_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_BSND, QLIV2_LAYOUT_TND, + QLIV2_LAYOUT_PA_BBND), ); // 支持的模板参数组合 // 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法 ASCENDC_TPL_SEL( @@ -91,21 +89,18 @@ ASCENDC_TPL_SEL( ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLIV2_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0), ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_TND), ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_TND), ), - ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLIV2_TPL_INT8), - ASCENDC_TPL_DTYPE_SEL(DT_K, QLIV2_TPL_INT8), + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLIV2_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_K, QLIV2_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLIV2_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 1), ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_BSND, QLIV2_LAYOUT_TND), ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_PA_BBND), ), - ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLIV2_TPL_INT8), - ASCENDC_TPL_DTYPE_SEL(DT_K, QLIV2_TPL_INT8), + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLIV2_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_K, QLIV2_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLIV2_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0), ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_BSND), ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_BSND), ), - ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLIV2_TPL_INT8), - ASCENDC_TPL_DTYPE_SEL(DT_K, QLIV2_TPL_INT8), + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLIV2_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_K, QLIV2_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLIV2_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0), ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_TND), - ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_TND),)); + ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLIV2_LAYOUT_TND), )); #else ASCENDC_TPL_ARGS_DECL(QuantLightningIndexerV2, // 算子OpType ASCENDC_TPL_DTYPE_DECL(DT_Q, QLIV2_TPL_INT8), ASCENDC_TPL_DTYPE_DECL(DT_K, QLIV2_TPL_INT8), diff --git a/csrc/attention/quant_lightning_indexer_v2/quant_lightning_indexer_v2_torch_adpt.h b/csrc/attention/quant_lightning_indexer_v2/quant_lightning_indexer_v2_torch_adpt.h new file mode 100644 index 000000000000..f529923c7cc2 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/quant_lightning_indexer_v2_torch_adpt.h @@ -0,0 +1,164 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#pragma once + +namespace vllm_ascend::qli_v2 { +constexpr int SIZE = 8; +constexpr int DIM_0 = 0; +constexpr int DIM_1 = 1; +constexpr int DIM_2 = 2; +inline at::Tensor valid_tensor(const c10::optional& value, const at::Device& device) { + return value.has_value() ? *value : at::empty({0}, at::TensorOptions().dtype(at::kInt).device(device)); +} +constexpr int64_t QLI_V2_METADATA_SIZE = 1024; + +at::Tensor QuantLightningIndexerMetadata(int64_t numHeadsQ, int64_t numHeadsK, int64_t headDim, int64_t topk, + int64_t quantMode, const c10::optional &cuSeqlensQ, + const c10::optional &cuSeqlensK, + const c10::optional &sequsedQ, + const c10::optional &sequsedK, + const c10::optional &cmpResidualK, int64_t batchSize, + int64_t maxSeqlenQ, int64_t maxSeqlenK, c10::string_view layoutQ, + c10::string_view layoutK, int64_t maskMode, int64_t cmpRatio) +{ + at::Device outputDevice = at::Device(std::string("npu")); + if (cuSeqlensQ.has_value()) { + outputDevice = cuSeqlensQ.value().device(); + } else if (cuSeqlensK.has_value()) { + outputDevice = cuSeqlensK.value().device(); + } else if (sequsedQ.has_value()) { + outputDevice = sequsedQ.value().device(); + } else if (sequsedK.has_value()) { + outputDevice = sequsedK.value().device(); + } else if (cmpResidualK.has_value()) { + outputDevice = cmpResidualK.value().device(); + } + + at::Tensor output = torch::empty({QLI_V2_METADATA_SIZE}, torch::dtype(torch::kInt32).device(outputDevice)); + auto cuSeqlensQVal = valid_tensor(cuSeqlensQ, outputDevice); + auto cuSeqlensKVal = valid_tensor(cuSeqlensK, outputDevice); + auto sequsedQVal = valid_tensor(sequsedQ, outputDevice); + auto sequsedKVal = valid_tensor(sequsedK, outputDevice); + auto cmpResidualKVal = valid_tensor(cmpResidualK, outputDevice); + + std::string layoutQStr = std::string(layoutQ); + std::string layoutKStr = std::string(layoutK); + char *layoutQPtr = const_cast(layoutQStr.c_str()); + char *layoutKPtr = const_cast(layoutKStr.c_str()); + + if (output.device().is_meta()) return output; + EXEC_NPU_CMD(aclnnQuantLightningIndexerV2Metadata, cuSeqlensQVal, cuSeqlensKVal, sequsedQVal, sequsedKVal, + cmpResidualKVal, numHeadsQ, numHeadsK, headDim, topk, quantMode, batchSize, maxSeqlenQ, maxSeqlenK, + layoutQPtr, layoutKPtr, maskMode, cmpRatio, output); + return output; +} + +std::tuple ConstructQuantLightningIndexerOutputTensor( + const at::Tensor &query, const at::Tensor &key, int64_t sparseCount, std::string queryLayoutStr, + std::string keyLayoutStr, int64_t returnValue) +{ + at::SmallVector outputSize; + for (size_t i = 0; i < query.sizes().size(); i++) { + TORCH_CHECK(query.size(i) > 0, + "All values within query's shape should be greater " + "than 0, but shape[", + i, "] is ", query.size(i)); + } + for (size_t i = 0; i < key.sizes().size(); i++) { + TORCH_CHECK(key.size(i) > 0, + "All values within key's shape should be greater " + "than 0, but shape[", + i, "] is ", key.size(i)); + } + TORCH_CHECK(sparseCount > 0, "sparse count should be greater than 0, but now is ", sparseCount); + int64_t keyHeadNum = (keyLayoutStr == "TND") ? key.size(DIM_1) : key.size(DIM_2); + if (queryLayoutStr == "BSND") { + outputSize = {query.size(DIM_0), query.size(DIM_1), keyHeadNum, sparseCount}; + } else { + int nDimIndex = 0; + nDimIndex = (keyLayoutStr == "TND") ? DIM_1 : DIM_2; + outputSize = {query.size(DIM_0), key.size(nDimIndex), sparseCount}; + } + at::Tensor sparseIndicesOut = at::empty(outputSize, query.options().dtype(at::kInt)); + at::Tensor sparseValuesOut; + if (returnValue) { + sparseValuesOut = at::empty(outputSize, query.options().dtype(at::kBFloat16)); + } else { + sparseValuesOut = at::empty({0}, query.options().dtype(at::kBFloat16)); + } + + return std::tuple(sparseIndicesOut, sparseValuesOut); +} + +std::tuple QuantLightningIndexerCandidate( + const at::Tensor &query, const at::Tensor &key, const at::Tensor &weights, const at::Tensor &queryDequantScale, + const at::Tensor &keyDequantScale, int64_t topk, int64_t quantMode, + const c10::optional &candidateTopkIndexIn, const c10::optional &cuSeqlensQ, + const c10::optional &cuSeqlensK, const c10::optional &sequsedQ, + const c10::optional &sequsedK, const c10::optional &cmpResidualK, + const c10::optional &blockTable, const c10::optional &outputIdxOffset, + const c10::optional &metadata, int64_t maxSeqlenQ, c10::string_view layoutQ, + c10::string_view layoutK, int64_t maskMode, int64_t cmpRatio, int64_t candidateMode, + int64_t candidateTopkBlocks, int64_t candidateBlockSize) +{ + TORCH_CHECK(query.numel() > 0, "Tensor query is empty.") + TORCH_CHECK(key.numel() > 0, "Tensor key is empty.") + + std::string queryLayoutStr = std::string(layoutQ); + std::string keyLayoutStr = std::string(layoutK); + + std::tuple quantLightningIndexerOutput = + ConstructQuantLightningIndexerOutputTensor(query, key, topk, queryLayoutStr, keyLayoutStr, 0); + at::Tensor sparseIndicesOut = std::get<0>(quantLightningIndexerOutput); + at::Tensor sparseValuesOut = std::get<1>(quantLightningIndexerOutput); + + int64_t keyHeadNum = (keyLayoutStr == "TND") ? key.size(DIM_1) : key.size(DIM_2); + at::Tensor candidateTopkIndexOut; + if (candidateMode == 1) { + at::SmallVector candSize; + if (queryLayoutStr == "BSND") { + candSize = {query.size(DIM_0), query.size(DIM_1), keyHeadNum, candidateTopkBlocks}; + } else { + candSize = {query.size(DIM_0), keyHeadNum, candidateTopkBlocks}; + } + candidateTopkIndexOut = at::empty(candSize, query.options().dtype(at::kInt)); + } else { + candidateTopkIndexOut = at::empty({0}, query.options().dtype(at::kInt)); + } + + char *queryLayoutPtr = const_cast(queryLayoutStr.c_str()); + char *keyLayoutPtr = const_cast(keyLayoutStr.c_str()); + int64_t returnValue = 0; + + TORCH_CHECK(quantMode == 2, "Aurora QLI V2 currently supports INT8 quant_mode=2"); + TORCH_CHECK(query.scalar_type() == at::kChar && key.scalar_type() == at::kChar, + "QLI V2 query/key must be INT8"); + TORCH_CHECK(weights.scalar_type() == at::kHalf && queryDequantScale.scalar_type() == at::kHalf && + keyDequantScale.scalar_type() == at::kHalf, "QLI V2 weights/scales must be FP16"); + TORCH_CHECK(candidateMode >= 1 && candidateMode <= 3, "Invalid candidate_mode"); + TORCH_CHECK(candidateMode != 2 || candidateTopkIndexIn.has_value(), "Consumer requires candidate blocks"); + if (query.device().is_meta()) return {sparseIndicesOut, sparseValuesOut, candidateTopkIndexOut}; + + // A11: key 0 轴非连续 — aclnn 动态调用下 tiling 拿不到 tensor stride (仅 TensorV2/图模式可见), + // 从 key/k_scale 的 torch stride(0) 自动显式传入 (紧凑存储时等于紧凑值, 走 kernel 兜底语义) + int64_t keyStride0 = key.stride(0); + int64_t keyScaleStride0 = keyDequantScale.stride(0); + + EXEC_NPU_CMD(aclnnQuantLightningIndexerV2, query, key, weights, queryDequantScale, keyDequantScale, + cuSeqlensQ, cuSeqlensK, sequsedQ, sequsedK, cmpResidualK, blockTable, outputIdxOffset, + metadata, candidateTopkIndexIn, topk, quantMode, maxSeqlenQ, queryLayoutPtr, keyLayoutPtr, maskMode, + cmpRatio, returnValue, candidateMode, candidateTopkBlocks, candidateBlockSize, keyStride0, + keyScaleStride0, sparseIndicesOut, sparseValuesOut, candidateTopkIndexOut); + + return std::tuple(sparseIndicesOut, sparseValuesOut, candidateTopkIndexOut); +} + +} // namespace vllm_ascend::qli_v2 diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2/tests/CMakeLists.txt new file mode 100644 index 000000000000..d5a84231c22b --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/CMakeLists.txt @@ -0,0 +1,16 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/compare_batch_outputs.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/compare_batch_outputs.py new file mode 100644 index 000000000000..5e670016e668 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/compare_batch_outputs.py @@ -0,0 +1,444 @@ +#!/usr/bin/env python3 +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +"""Compare LI_V2/QLI_V2 output-0 dumps across TTK batch cases.""" + +import argparse +import ast +import csv +import hashlib +import importlib.util +import json +import sys +from collections import Counter, defaultdict +from pathlib import Path + +import numpy as np + + +def load_batch_protocol(): + name = "qli_v2_ttk_batch_consistency" + path = Path(__file__).with_name("impl") / "batch_consistency.py" + if name in sys.modules: + return sys.modules[name] + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +class IndexerBatchDumpComparator: + """Validate exact output indices after all enabled cases finish.""" + + def __init__( + self, + case_csv, + result_csv, + dump_dir, + report_path, + excel_path, + require_intra_case, + require_cross_case, + ): + self.case_csv = Path(case_csv) + self.result_csv = Path(result_csv) + self.dump_dir = Path(dump_dir) + self.report_path = Path(report_path) + self.excel_path = Path(excel_path) if excel_path else None + self.require_intra_case = require_intra_case + self.require_cross_case = require_cross_case + + @staticmethod + def read_rows(path): + with path.open("r", encoding="utf-8-sig", newline="") as stream: + return list(csv.DictReader(stream)) + + @staticmethod + def parse_cell(row, name, default=None): + value = row.get(name) + if value is None or value == "": + return default + try: + return ast.literal_eval(value) + except (SyntaxError, ValueError) as error: + raise ValueError(f"{row.get('testcase_name', '')}: invalid {name}: {error}") from error + + @classmethod + def case_enabled(cls, row): + value = row.get("is_enabled") + if value is None or value.strip() == "": + return True + normalized = value.strip().title() + try: + return bool(ast.literal_eval(normalized)) + except (SyntaxError, ValueError) as error: + raise ValueError(f"{row.get('testcase_name', '')}: invalid is_enabled") from error + + @staticmethod + def vector(attributes, name, batch_size, default): + value = attributes.get(f"{name}_values") + if value is None: + return [default] * batch_size + value = [int(item) for item in value] + if len(value) != batch_size: + raise ValueError(f"{name}_values length does not equal B={batch_size}") + return value + + @staticmethod + def prefix_lengths(attributes, name, batch_size): + value = attributes.get(f"{name}_values") + if value is None: + return None + value = [int(item) for item in value] + if len(value) != batch_size + 1 or value[0] != 0: + raise ValueError(f"{name}_values must contain B + 1 prefix values") + if any(right <= left for left, right in zip(value, value[1:])): + raise ValueError(f"{name}_values must be strictly increasing") + return [right - left for left, right in zip(value, value[1:])] + + def geometry(self, row): + shapes = self.parse_cell(row, "tensor_view_shapes") + dtypes = self.parse_cell(row, "tensor_dtypes") + attributes = self.parse_cell(row, "attributes", {}) + if len(shapes) not in (11, 13): + raise ValueError(f"{row['testcase_name']}: expected LI/QLI direct input slots") + q_shape = tuple(shapes[0]) + k_shape = tuple(shapes[1]) + layout_q = attributes.get("layout_q", attributes.get("layout_query", "BSND")) + layout_k = attributes.get("layout_k", attributes.get("layout_key", "BSND")) + if layout_q == "BSND": + batch_size, q_extent, q_heads, head_dim = q_shape + q_prefix = None + q_lengths = self.vector(attributes, "seqused_q", batch_size, q_extent) + output_prefix = None + elif layout_q == "TND": + q_extent, q_heads, head_dim = q_shape + q_values = attributes.get("cu_seqlens_q_values") + if q_values is None or q_values[0] != 0 or q_values[-1] != q_extent: + raise ValueError(f"{row['testcase_name']}: TND q prefix must span the q tensor") + batch_size = len(q_values) - 1 + q_prefix = [int(item) for item in q_values] + q_lengths = [right - left for left, right in zip(q_prefix, q_prefix[1:])] + output_prefix = q_prefix + else: + raise ValueError(f"{row['testcase_name']}: unsupported layout_q={layout_q}") + + if layout_k == "BSND": + if int(k_shape[0]) != batch_size: + raise ValueError(f"{row['testcase_name']}: key B does not match q B") + k_extent, key_heads = int(k_shape[1]), int(k_shape[2]) + k_lengths = self.vector(attributes, "seqused_k", batch_size, k_extent) + block_size = None + elif layout_k == "TND": + key_heads = int(k_shape[1]) + k_lengths = self.prefix_lengths(attributes, "cu_seqlens_k", batch_size) + if k_lengths is None or sum(k_lengths) != int(k_shape[0]): + raise ValueError(f"{row['testcase_name']}: invalid TND k prefix") + block_size = None + elif layout_k == "PA_BBND": + block_size, key_heads = int(k_shape[1]), int(k_shape[2]) + k_lengths = self.vector(attributes, "seqused_k", batch_size, 0) + else: + raise ValueError(f"{row['testcase_name']}: unsupported layout_k={layout_k}") + + topk = int(attributes.get("topk", attributes.get("sparse_count"))) + output_shape = (batch_size, q_extent, key_heads, topk) if layout_q == "BSND" else (q_extent, key_heads, topk) + return { + "attributes": attributes, + "input_dtypes": tuple(dtypes), + "batch_size": batch_size, + "q_heads": int(q_heads), + "key_heads": key_heads, + "head_dim": int(head_dim), + "q_lengths": q_lengths, + "k_lengths": k_lengths, + "residual": self.vector(attributes, "cmp_residual_k", batch_size, 0), + "layout_q": layout_q, + "layout_k": layout_k, + "block_size": block_size, + "output_prefix": output_prefix, + "output_shape": output_shape, + "topk": topk, + } + + @staticmethod + def output_selector(relation, geometry): + axes, slices, _seed = relation + batch_slice = slices[0] + sequence_slice = slices[1] if axes == (0, 1) else None + batch_start, batch_stop, _ = batch_slice + if batch_stop > geometry["batch_size"]: + raise ValueError("logical B slice exceeds output batch") + if geometry["layout_q"] == "BSND": + selector = [slice(*batch_slice)] + if sequence_slice is not None: + if sequence_slice[1] > geometry["q_lengths"][batch_start]: + raise ValueError("logical S slice exceeds BSND output") + selector.append(slice(*sequence_slice)) + else: + active_lengths = geometry["q_lengths"][batch_start:batch_stop] + if ( + not active_lengths + or len(set(active_lengths)) != 1 + or active_lengths[0] <= 0 + or active_lengths[0] > geometry["output_shape"][1] + ): + raise ValueError("invalid effective q lengths for BSND output") + selector.append(slice(0, active_lengths[0], 1)) + else: + prefix = geometry["output_prefix"] + if sequence_slice is None: + start, stop = prefix[batch_start], prefix[batch_stop] + else: + start = prefix[batch_start] + sequence_slice[0] + stop = prefix[batch_start] + sequence_slice[1] + if stop > prefix[batch_start + 1]: + raise ValueError("logical S slice exceeds TND output") + selector = [slice(start, stop, 1)] + selector.extend([slice(None)] * (len(geometry["output_shape"]) - len(selector))) + return tuple(selector) + + @staticmethod + def context(row, relation, geometry): + axes, slices, _seed = relation + batch_start, batch_stop, _ = slices[0] + sequence_slice = slices[1] if axes == (0, 1) else None + attributes = geometry["attributes"] + q_lengths = geometry["q_lengths"][batch_start:batch_stop] + if sequence_slice is not None: + q_lengths = [sequence_slice[1] - sequence_slice[0]] + ignored = { + "seqused_q_values", + "seqused_k_values", + "cu_seqlens_q_values", + "cu_seqlens_k_values", + "cmp_residual_k_values", + "batch_deterministic_level", + } + scalar_attributes = { + key: value + for key, value in attributes.items() + if key not in ignored and not isinstance(value, (list, tuple, dict)) + } + return { + "api_name": row.get("api_name"), + "input_dtypes": geometry["input_dtypes"], + "layout_q": geometry["layout_q"], + "layout_k": geometry["layout_k"], + "q_heads": geometry["q_heads"], + "key_heads": geometry["key_heads"], + "head_dim": geometry["head_dim"], + "topk": geometry["topk"], + "block_size": geometry["block_size"], + "q_lengths": tuple(q_lengths), + "k_lengths": tuple(geometry["k_lengths"][batch_start:batch_stop]), + "residual": tuple(geometry["residual"][batch_start:batch_stop]), + "attributes": scalar_attributes, + } + + def build_samples(self, case_rows, result_rows): + results = {row["testcase_name"]: row for row in result_rows if row.get("testcase_name")} + protocol_class = load_batch_protocol().BatchRelationProtocol + samples = [] + for row in case_rows: + testcase_name = row.get("testcase_name") + if not testcase_name: + raise ValueError("case CSV contains an empty testcase_name") + if row.get("batch_axis") in (None, ""): + continue + result = results.get(testcase_name) + if result is None: + raise ValueError(f"{testcase_name}: result CSV has no matching row") + if result.get("precision_status") != "PASS": + raise ValueError(f"{testcase_name}: precision_status is not PASS") + eager_precision = result.get("eager_precision") or "" + if "NO_OUTPU" in eager_precision: + raise ValueError( + f"{testcase_name}: eager_precision reports no output; raw-byte batch comparison is invalid" + ) + if self.require_intra_case and "batch_intra=PASS" not in eager_precision: + raise ValueError(f"{testcase_name}: same-case batch check did not PASS") + + geometry = self.geometry(row) + output_path = self.dump_dir / f"{testcase_name}_output_0.bin" + if not output_path.is_file(): + raise ValueError(f"{testcase_name}: missing output dump {output_path}") + output_bytes = output_path.read_bytes() + expected_bytes = int(np.prod(geometry["output_shape"])) * 4 + if len(output_bytes) != expected_bytes: + raise ValueError(f"{testcase_name}: output bytes={len(output_bytes)}, expected={expected_bytes}") + output = np.frombuffer(output_bytes, dtype=np.int32).reshape(geometry["output_shape"]) + protocol = protocol_class("QLI_V2" if len(geometry["input_dtypes"]) == 13 else "LI_V2") + relations = protocol.parse( + self.parse_cell(row, "batch_axis"), + self.parse_cell(row, "batch_slice_info"), + self.parse_cell(row, "batch_seed"), + ) + for relation in relations: + selected = np.ascontiguousarray(output[self.output_selector(relation, geometry)]) + axes, slices, seed = relation + relation_size = tuple(stop - start for start, stop, _step in slices) + value = selected.view(np.uint8).tobytes() + samples.append( + { + "testcase_name": testcase_name, + "relation": (axes, seed, relation_size), + "slice": { + "B": slices[0], + "S": slices[1] if len(slices) > 1 else None, + }, + "shape": tuple(selected.shape), + "context": self.context(row, relation, geometry), + "sha256": hashlib.sha256(value).hexdigest(), + "value": value, + } + ) + if not samples: + raise ValueError("case CSV contains no enabled batch relation") + return samples + + def compare_group(self, relation, samples): + reference = samples[0] + errors = [] + for sample in samples[1:]: + if sample["shape"] != reference["shape"]: + errors.append(f"{sample['testcase_name']}: output slice shape differs") + if sample["context"] != reference["context"]: + errors.append(f"{sample['testcase_name']}: relation context differs") + if sample["value"] != reference["value"]: + errors.append(f"{sample['testcase_name']}: raw output bytes differ") + case_counts = Counter(sample["testcase_name"] for sample in samples) + if self.require_intra_case: + missing = sorted(name for name, count in case_counts.items() if count < 2) + if missing: + errors.append("fewer than two same-case samples: " + ", ".join(missing)) + if self.require_cross_case and len(case_counts) < 2: + errors.append("relation occurs in fewer than two testcases") + return { + "relation": relation, + "status": "PASS" if not errors else "FAIL", + "case_count": len(case_counts), + "sample_count": len(samples), + "errors": errors, + "samples": [ + {key: sample[key] for key in ("testcase_name", "slice", "shape", "sha256")} for sample in samples + ], + } + + def write_excel(self, report): + from openpyxl import Workbook + + workbook = Workbook() + sheet = workbook.active + sheet.title = "Relations" + sheet.append( + ( + "Relation", + "Status", + "Cases", + "Samples", + "Testcase", + "Slice", + "SHA-256", + "Errors", + ) + ) + for group in report["groups"]: + errors = "\n".join(group["errors"]) + for sample in group["samples"]: + sheet.append( + ( + repr(group["relation"]), + group["status"], + group["case_count"], + group["sample_count"], + sample["testcase_name"], + repr(sample["slice"]), + sample["sha256"], + errors, + ) + ) + self.excel_path.parent.mkdir(parents=True, exist_ok=True) + workbook.save(self.excel_path) + + def run(self): + case_rows = self.read_rows(self.case_csv) + result_rows = self.read_rows(self.result_csv) + disabled = [row.get("testcase_name", "") for row in case_rows if not self.case_enabled(row)] + enabled = [row for row in case_rows if self.case_enabled(row)] + samples = self.build_samples(enabled, result_rows) + grouped = defaultdict(list) + for sample in samples: + grouped[sample["relation"]].append(sample) + groups = [self.compare_group(relation, values) for relation, values in sorted(grouped.items())] + passed = all(group["status"] == "PASS" for group in groups) + report = { + "status": "PASS" if passed else "FAIL", + "case_csv": str(self.case_csv.resolve()), + "result_csv": str(self.result_csv.resolve()), + "dump_dir": str(self.dump_dir.resolve()), + "group_count": len(groups), + "sample_count": len(samples), + "disabled_case_count": len(disabled), + "disabled_cases": disabled, + "groups": groups, + } + self.report_path.parent.mkdir(parents=True, exist_ok=True) + self.report_path.write_text(json.dumps(report, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") + if self.excel_path is not None: + self.write_excel(report) + return passed, report + + +def build_parser(): + parser = argparse.ArgumentParser(description="Compare LI_V2/QLI_V2 raw output-0 batch relations.") + parser.add_argument("--csv", required=True) + parser.add_argument("--result", required=True) + parser.add_argument("--dump-dir", required=True) + parser.add_argument("--report", required=True) + parser.add_argument("--excel") + parser.add_argument("--output-index", type=int, default=0) + parser.add_argument("--require-intra-case", action="store_true") + parser.add_argument("--require-cross-case", action="store_true") + return parser + + +def main(): + args = build_parser().parse_args() + if args.output_index != 0: + print("LI_V2/QLI_V2 batch comparison supports output index 0 only") + return 1 + comparator = IndexerBatchDumpComparator( + args.csv, + args.result, + args.dump_dir, + args.report, + args.excel, + args.require_intra_case, + args.require_cross_case, + ) + try: + passed, report = comparator.run() + except (OSError, ValueError, csv.Error) as error: + print(f"batch consistency comparison failed: {error}") + return 1 + print( + f"batch consistency {report['status']}: groups={report['group_count']}, " + f"samples={report['sample_count']}, report={args.report}" + ) + return 0 if passed else 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/batch_consistency.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/batch_consistency.py new file mode 100644 index 000000000000..15c0d572771d --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/batch_consistency.py @@ -0,0 +1,646 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +"""Small batch-consistency protocol shared by LI_V2 and QLI_V2 assets.""" + +import hashlib +import random +from numbers import Integral + +import numpy as np +import torch + +HIFLOAT8_QUANT_MODE = 4 +SUPPORTED_QUANT_MODES = (1, 2, 3, 4, 5) +MXFP4_DECODE_VALUES = torch.tensor( + ( + 0.0, + 0.5, + 1.0, + 1.5, + 2.0, + 3.0, + 4.0, + 6.0, + -0.0, + -0.5, + -1.0, + -1.5, + -2.0, + -3.0, + -4.0, + -6.0, + ), + dtype=torch.float32, +) + + +class CaseRandomContext: + """Give batch cases distinct backgrounds without changing normal cases.""" + + def __init__(self, attributes): + fields = tuple(attributes.get(name) for name in ("batch_axis", "batch_slice_info", "batch_seed")) + self.enabled = any(value is not None for value in fields) + self.testcase_name = attributes.get("testcase_name", "") + self.python_state = None + self.numpy_state = None + self.torch_state = None + + def __enter__(self): + if not self.enabled: + return self + digest = hashlib.sha256(str(self.testcase_name).encode("utf-8")).digest() + seed = int.from_bytes(digest[:8], "big") % ((1 << 32) - 1) + self.python_state = random.getstate() + self.numpy_state = np.random.get_state() + self.torch_state = torch.random.get_rng_state() + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + return self + + def __exit__(self, exc_type, exc_value, traceback): + if self.enabled: + random.setstate(self.python_state) + np.random.set_state(self.numpy_state) + torch.random.set_rng_state(self.torch_state) + return False + + +class BatchRelationProtocol: + """Parse the q-only logical B/S relation contract used by both indexers.""" + + def __init__(self, operator_name): + self.operator_name = operator_name + + @staticmethod + def relation_slices_overlap(first, second): + """Return whether two relation samples select the same q output region.""" + first_axes, first_slices = first[0], first[1] + second_axes, second_slices = second[0], second[1] + if first_axes != second_axes: + return False + first_batch = first_slices[0] + second_batch = second_slices[0] + if first_batch[1] <= second_batch[0] or second_batch[1] <= first_batch[0]: + return False + if first_axes == (0,): + return True + first_sequence = first_slices[1] + second_sequence = second_slices[1] + return not (first_sequence[1] <= second_sequence[0] or second_sequence[1] <= first_sequence[0]) + + def validate_disjoint_relations(self, relations): + """Reject duplicate or overlapping samples that would self-compare.""" + for index, relation in enumerate(relations): + for candidate in relations[index + 1 :]: + if self.relation_slices_overlap(relation, candidate): + raise ValueError(f"{self.operator_name} relation samples must not overlap") + + def parse(self, batch_axis, batch_slice_info, batch_seed): + fields = (batch_axis, batch_slice_info, batch_seed) + if all(value is None for value in fields): + return None + if any(value is None for value in fields): + raise ValueError(f"{self.operator_name} batch_axis, batch_slice_info and batch_seed must be set together") + if not (len(batch_axis) == len(batch_slice_info) == len(batch_seed)): + raise ValueError(f"{self.operator_name} batch metadata counts differ") + if not batch_axis or tuple(batch_axis[0]) not in ((0,), (0, 1)): + raise ValueError(f"{self.operator_name} supports q logical axes (0,) or (0, 1)") + if any(value is not None for value in batch_axis[1:]): + raise ValueError(f"{self.operator_name} supports q relations only") + if batch_slice_info[0] is None or batch_seed[0] is None: + raise ValueError(f"{self.operator_name} requires q slices and q seeds") + if any(value is not None for value in batch_slice_info[1:]): + raise ValueError(f"{self.operator_name} supports q relations only") + if any(value is not None for value in batch_seed[1:]): + raise ValueError(f"{self.operator_name} supports q relation seeds only") + + axes = tuple(batch_axis[0]) + axis_slices = batch_slice_info[0] + axis_seeds = batch_seed[0] + if len(axis_slices) != len(axes) or len(axis_seeds) != len(axes): + raise ValueError(f"{self.operator_name} q groups do not match q axes") + sample_count = len(axis_slices[0]) + if not sample_count or any(len(values) != sample_count for values in (*axis_slices, *axis_seeds)): + raise ValueError(f"{self.operator_name} q sample counts differ or are empty") + + relations = [] + for sample_index in range(sample_count): + slices = [] + relation_seed = None + for axis_group, axis in enumerate(axes): + value = axis_slices[axis_group][sample_index] + if not isinstance(value, (tuple, list)) or len(value) != 3: + raise ValueError(f"{self.operator_name} invalid q axis {axis} slice: {value!r}") + if not all(isinstance(item, Integral) for item in value): + raise ValueError(f"{self.operator_name} slices must contain integers") + start, stop, step = (int(item) for item in value) + if step != 1 or start < 0 or start >= stop: + raise ValueError(f"{self.operator_name} slices must be non-empty and contiguous") + seed = axis_seeds[axis_group][sample_index] + if not isinstance(seed, Integral): + raise ValueError(f"{self.operator_name} batch seed must be an integer") + seed = int(seed) + if relation_seed is not None and seed != relation_seed: + raise ValueError(f"{self.operator_name} logical B and S must use the same seed") + relation_seed = seed + slices.append((start, stop, step)) + if axes == (0, 1) and slices[0][1] - slices[0][0] != 1: + raise ValueError(f"{self.operator_name} logical (B,S) requires one B per sample") + relations.append((axes, tuple(slices), relation_seed)) + self.validate_disjoint_relations(relations) + return relations + + def validate_id(self, batch_consistency_id, relations): + """Check the framework ID fields that identify a logical relation. + + Framework versions may encode slice bounds or lengths differently. The + seed and axis are the stable identity; slice ranges remain validated by + ``parse`` and the relation/output checks below. + """ + if not isinstance(batch_consistency_id, (tuple, list)) or len(batch_consistency_id) != 1: + raise ValueError("batch_consistency_id must contain one q relation group") + axes = relations[0][0] + id_groups = batch_consistency_id[0] + if not isinstance(id_groups, (tuple, list)) or len(id_groups) != len(axes): + raise ValueError("batch_consistency_id axis groups do not match q axes") + for group_index, axis in enumerate(axes): + ids = id_groups[group_index] + if not isinstance(ids, (tuple, list)): + raise ValueError("batch_consistency_id samples must be sequences") + if len(ids) != len(relations): + raise ValueError("batch_consistency_id sample count does not match q relations") + for relation_id, relation in zip(ids, relations): + parts = str(relation_id).split("_", 2) + if len(parts) < 2: + raise ValueError(f"invalid batch_consistency_id relation: {relation_id!r}") + try: + id_seed, id_axis = int(parts[0]), int(parts[1]) + except ValueError as error: + raise ValueError(f"invalid batch_consistency_id relation: {relation_id!r}") from error + if id_seed != int(relation[2]) or id_axis != int(axis): + raise ValueError("batch_consistency_id seed/axis does not match q relation") + + +class IndexerBatchInputNormalizer: + """Materialize equal logical inputs for declared LI/QLI relations.""" + + def __init__(self, data, attributes, operator_name, quantized): + self.data = data + self.attributes = attributes + self.operator_name = operator_name + self.quantized = quantized + self.quant_mode = int(attributes.get("quant_mode", 1)) if quantized else None + self.layout_q = attributes.get("layout_q", attributes.get("layout_query", "BSND")) + self.layout_k = attributes.get("layout_k", attributes.get("layout_key", "BSND")) + self.query = data["query"] + self.key = data["key"] + self.weights = data["weights"] + self.query_scale = data.get("query_dequant_scale") + self.key_scale = data.get("key_dequant_scale") + self.offset = data.get("output_idx_offset") + self.block_table = data.get("block_table") + self.q_prefix = self.tensor_values(data.get("cu_seqlens_query", data.get("cu_seqlens_q"))) + self.k_prefix = self.tensor_values(data.get("cu_seqlens_key", data.get("cu_seqlens_k"))) + self.batch_size = self.resolve_batch_size() + self.q_lengths = self.resolve_lengths("q") + self.k_lengths = self.resolve_lengths("k") + self.residual = self.resolve_vector("cmp_residual_k", 0) + self.assigned_blocks = {} + + @staticmethod + def tensor_values(value): + if value is None: + return None + if torch.is_tensor(value): + value = value.detach().cpu().reshape(-1).tolist() + return [int(item) for item in value] + + def resolve_batch_size(self): + if self.layout_q == "BSND": + return int(self.query.shape[0]) + if self.layout_q == "TND" and self.q_prefix is not None: + return len(self.q_prefix) - 1 + raise ValueError(f"{self.operator_name} batch consistency requires BSND or explicit TND prefix") + + def resolve_vector(self, name, default): + value = self.attributes.get(f"{name}_values") + if value is None: + value = self.data.get(name) + value = self.tensor_values(value) + if value is None: + return [default] * self.batch_size + if len(value) != self.batch_size: + raise ValueError(f"{self.operator_name} {name} length must equal B={self.batch_size}") + return value + + def resolve_lengths(self, target): + prefix = self.q_prefix if target == "q" else self.k_prefix + tensor = self.query if target == "q" else self.key + layout = self.layout_q if target == "q" else self.layout_k + if prefix is not None: + if ( + len(prefix) != self.batch_size + 1 + or prefix[0] != 0 + or prefix[-1] != int(tensor.shape[0]) + or any(right <= left for left, right in zip(prefix, prefix[1:])) + ): + raise ValueError(f"{self.operator_name} {target} prefix must strictly span its tensor") + lengths = [right - left for left, right in zip(prefix, prefix[1:])] + else: + if layout == "BSND": + lengths = [int(tensor.shape[1])] * self.batch_size + else: + lengths = self.resolve_vector(f"seqused_{target}", 0) + actual = self.attributes.get(f"seqused_{target}_values") + if actual is not None: + actual = [int(item) for item in actual] + if len(actual) != self.batch_size: + raise ValueError(f"{self.operator_name} seqused_{target} length must equal B") + if any(length <= 0 for length in actual): + raise ValueError(f"{self.operator_name} seqused_{target} must be positive") + if layout in ("BSND", "TND") and any( + actual_length > physical_length for actual_length, physical_length in zip(actual, lengths) + ): + raise ValueError(f"{self.operator_name} seqused_{target} exceeds its tensor extent") + return actual + return lengths + + @staticmethod + def derived_seed(seed, relative_batch, slot): + value = f"{int(seed)}:{int(relative_batch)}:{int(slot)}".encode("ascii") + return int.from_bytes(hashlib.sha256(value).digest()[:8], "big") % ((1 << 63) - 1) + + @classmethod + def random_tensor(cls, shape, template, seed, relative_batch, slot, positive=False): + generator = torch.Generator(device="cpu") + generator.manual_seed(cls.derived_seed(seed, relative_batch, slot)) + dtype = template.dtype + if dtype == torch.bool: + value = torch.randint(0, 2, shape, generator=generator, dtype=torch.int64) + elif "float4" in str(dtype): + # CPU cannot cast into Float4, but its packed byte view is writable. + packed = torch.randint(0, 16, shape, generator=generator, dtype=torch.uint8) + return packed.view(dtype) + elif dtype.is_floating_point: + value = torch.rand(shape, generator=generator, dtype=torch.float32) + value = value * 0.75 + 0.25 if positive else value - 0.5 + elif dtype == torch.uint8: + value = torch.randint(0, 16, shape, generator=generator, dtype=torch.int64) + else: + value = torch.randint(-8, 9, shape, generator=generator, dtype=torch.int64) + return value.to(dtype=dtype) + + @staticmethod + def copy_selection(tensor, selector, value): + if tensor is None: + return + source = value + if torch.is_tensor(source) and "float4" in str(source.dtype) and "float4" not in str(tensor.dtype): + # PyTorch has no Float4 CPU cast kernel. The pytest CPU golden + # stores unpacked values, while the device input uses packed bytes. + packed = source.view(torch.uint8) + decode_values = MXFP4_DECODE_VALUES.to(device=source.device) + low = decode_values[(packed & 0x0F).to(torch.long)] + high = decode_values[(packed >> 4).to(torch.long)] + source = torch.stack((low, high), dim=-1).reshape(*source.shape[:-1], source.shape[-1] * 2) + tensor[selector].copy_(source.to(dtype=tensor.dtype, device=tensor.device)) + + def query_selector(self, batch_index, sequence_slice): + if self.layout_q == "BSND": + start, stop = (0, self.q_lengths[batch_index]) + if sequence_slice is not None: + start, stop = sequence_slice[:2] + return (batch_index, slice(start, stop, 1)), stop - start + token_start = self.q_prefix[batch_index] + token_stop = self.q_prefix[batch_index + 1] + if sequence_slice is not None: + token_start += sequence_slice[0] + token_stop = self.q_prefix[batch_index] + sequence_slice[1] + return (slice(token_start, token_stop, 1),), token_stop - token_start + + def query_capacity(self, batch_index): + """Return the physical q span used by the raw-byte output comparator.""" + if self.layout_q == "BSND": + return int(self.query.shape[1]) + return self.q_prefix[batch_index + 1] - self.q_prefix[batch_index] + + def query_comparison_length(self, batch_index): + """Match the q span that the output comparator will actually select.""" + if self.layout_q == "TND": + return self.query_capacity(batch_index) + return self.q_lengths[batch_index] + + def validate_relations(self, relations): + mask_mode = int(self.attributes.get("mask_mode", self.attributes.get("sparse_mode", 0))) + grouped_signatures = {} + occupied = [] + for axes, slices, seed in relations: + batch_start, batch_stop, _ = slices[0] + if batch_stop > self.batch_size: + raise ValueError(f"{self.operator_name} logical B slice exceeds B={self.batch_size}") + sequence_slice = slices[1] if axes == (0, 1) else None + if sequence_slice is not None and mask_mode != 0: + raise ValueError(f"{self.operator_name} shifted logical S relations require mask_mode=0") + if ( + sequence_slice is None + and len({self.query_capacity(batch_index) for batch_index in range(batch_start, batch_stop)}) != 1 + ): + raise ValueError(f"{self.operator_name} one B-only relation requires equal q output spans") + signature = [] + for batch_index in range(batch_start, batch_stop): + selector, q_count = self.query_selector(batch_index, sequence_slice) + if sequence_slice is not None and sequence_slice[1] > self.q_lengths[batch_index]: + raise ValueError(f"{self.operator_name} logical S slice exceeds effective q length") + occupied.append((selector, seed)) + signature.append( + ( + q_count if sequence_slice is not None else self.query_comparison_length(batch_index), + self.k_lengths[batch_index], + self.residual[batch_index], + ) + ) + relation_size = tuple(stop - start for start, stop, _step in slices) + key = (axes, seed, relation_size) + value = tuple(signature) + previous = grouped_signatures.setdefault(key, value) + if previous != value: + raise ValueError( + f"{self.operator_name} relation requires equal q output spans, K lengths and residuals" + ) + + for index, (left, left_seed) in enumerate(occupied): + for right, right_seed in occupied[index + 1 :]: + if left_seed == right_seed or len(left) != len(right): + continue + if self.selectors_overlap(left, right): + raise ValueError(f"{self.operator_name} relations with different seeds overlap") + + @staticmethod + def selectors_overlap(left, right): + for left_item, right_item in zip(left, right): + if isinstance(left_item, int) or isinstance(right_item, int): + if left_item != right_item: + return False + continue + if left_item.stop <= right_item.start or right_item.stop <= left_item.start: + return False + return True + + def query_references(self, name): + references = [self.data.get(name)] + state = self.data.get("golden_state", {}).get("forward_inputs", {}) + references.append(state.get(name)) + return [value for value in references if value is not None] + + def fill_query_inputs(self, batch_index, sequence_slice, seed, relative_batch): + selector, _count = self.query_selector(batch_index, sequence_slice) + targets = ( + (self.query_references("query"), 0, False), + (self.query_references("weights"), 1, False), + (self.query_references("output_idx_offset"), 3, True), + ) + if self.quant_mode != HIFLOAT8_QUANT_MODE: + targets += ((self.query_references("query_dequant_scale"), 2, True),) + for references, slot, positive in targets: + if not references: + continue + value = self.random_tensor( + tuple(references[0][selector].shape), + references[0], + seed, + relative_batch, + slot, + positive, + ) + for tensor in references: + self.copy_selection(tensor, selector, value) + + def key_references(self, name): + references = [] + if name == "key": + references.append(self.data.get("cpu_key")) + state = self.data.get("golden_state", {}).get("forward_inputs", {}) + references.append(state.get(name)) + return [value for value in references if value is not None] + + def input_references(self, name): + state = self.data.get("golden_state", {}).get("forward_inputs", {}) + return [value for value in (self.data.get(name), state.get(name)) if value is not None] + + def fill_hifloat8_scales(self): + """Use one stable global scale because mode 4 scales have shape ``(1,)``.""" + for name in ("query_dequant_scale", "key_dequant_scale"): + for tensor in self.input_references(name): + if torch.is_tensor(tensor): + tensor.fill_(1) + else: + np.asarray(tensor).fill(1) + + def scatter_paged(self, tensor, batch_index, value, seed, relative_batch): + table = self.tensor_values(self.block_table[batch_index]) + block_size = int(tensor.shape[1]) + copied = 0 + owner = (seed, relative_batch) + for block_id in table: + if block_id < 0 or copied >= value.shape[0]: + continue + count = min(block_size, int(value.shape[0]) - copied) + assignment = self.assigned_blocks.setdefault(block_id, owner) + if assignment != owner: + raise ValueError( + f"{self.operator_name} paged relations share block {block_id} between different logical batches" + ) + tensor[block_id, :count].copy_(value[copied : copied + count].to(tensor.device, tensor.dtype)) + copied += count + if copied != value.shape[0]: + raise ValueError(f"{self.operator_name} block table has insufficient capacity") + + def fill_key_tensor(self, name, batch_index, seed, relative_batch, slot, positive): + tensor = self.data.get(name) + if tensor is None: + return + key_length = self.k_lengths[batch_index] + if self.layout_k == "BSND": + selector = (batch_index, slice(0, key_length, 1)) + shape = tuple(tensor[selector].shape) + elif self.layout_k == "TND": + start, stop = self.k_prefix[batch_index : batch_index + 2] + selector = (slice(start, stop, 1),) + shape = tuple(tensor[selector].shape) + elif self.layout_k == "PA_BBND": + if self.block_table is None: + raise ValueError(f"{self.operator_name} PA_BBND requires block_table") + selector = None + shape = (key_length, *tuple(tensor.shape[2:])) + else: + raise ValueError(f"{self.operator_name} unsupported key layout {self.layout_k!r}") + value = self.random_tensor(shape, tensor, seed, relative_batch, slot, positive) + if selector is None: + self.scatter_paged(tensor, batch_index, value, seed, relative_batch) + else: + self.copy_selection(tensor, selector, value) + + for reference in self.key_references(name): + if reference is tensor: + continue + if self.layout_k == "PA_BBND": + permutation = (1, 0, *range(2, value.ndim)) + reference[batch_index, :, :key_length].copy_( + value.permute(permutation).to(reference.device, reference.dtype) + ) + else: + self.copy_selection(reference, selector, value) + + def apply(self, relations): + self.validate_relations(relations) + if self.quantized and self.quant_mode not in SUPPORTED_QUANT_MODES: + raise ValueError(f"{self.operator_name} batch consistency supports quant_mode 1 through 5") + for axes, slices, seed in relations: + batch_start, batch_stop, _ = slices[0] + sequence_slice = slices[1] if axes == (0, 1) else None + for relative_batch, batch_index in enumerate(range(batch_start, batch_stop)): + self.fill_query_inputs(batch_index, sequence_slice, seed, relative_batch) + self.fill_key_tensor("key", batch_index, seed, relative_batch, 10, False) + if self.quant_mode != HIFLOAT8_QUANT_MODE: + self.fill_key_tensor( + "key_dequant_scale", + batch_index, + seed, + relative_batch, + 11, + True, + ) + if self.quant_mode == HIFLOAT8_QUANT_MODE: + self.fill_hifloat8_scales() + + +class IndexerBatchOutputComparator: + """Compare exact output-0 slices for relations inside one testcase.""" + + def __init__(self, operator_name): + self.operator_name = operator_name + self.protocol = BatchRelationProtocol(operator_name) + + @staticmethod + def storage_bytes(value): + if torch.is_tensor(value): + tensor = value.detach().cpu().contiguous() + return ( + tuple(tensor.shape), + str(tensor.dtype), + tensor.view(torch.uint8).numpy().tobytes(), + ) + array = np.ascontiguousarray(np.asarray(value)) + return tuple(array.shape), array.dtype.str, array.view(np.uint8).tobytes() + + def output_selector(self, output, relation, attributes): + axes, slices, _seed = relation + batch_slice = slices[0] + sequence_slice = slices[1] if axes == (0, 1) else None + layout_q = attributes.get("layout_q", attributes.get("layout_query", "BSND")) + batch_start, batch_stop, _ = batch_slice + if layout_q == "BSND": + if batch_stop > output.shape[0]: + raise ValueError(f"{self.operator_name} logical B slice exceeds output") + selector = [slice(*batch_slice)] + if sequence_slice is not None: + if sequence_slice[1] > output.shape[1]: + raise ValueError(f"{self.operator_name} logical S slice exceeds output") + selector.append(slice(*sequence_slice)) + else: + q_lengths = attributes.get("seqused_q_values") + if q_lengths is not None: + selected_lengths = [int(value) for value in q_lengths[batch_start:batch_stop]] + if ( + len(selected_lengths) != batch_stop - batch_start + or len(set(selected_lengths)) != 1 + or selected_lengths[0] <= 0 + or selected_lengths[0] > output.shape[1] + ): + raise ValueError(f"{self.operator_name} invalid effective q lengths for output") + selector.append(slice(0, selected_lengths[0], 1)) + elif layout_q == "TND": + prefix = attributes.get("cu_seqlens_q_values") + if prefix is None or prefix[0] != 0 or prefix[-1] != output.shape[0]: + raise ValueError(f"{self.operator_name} TND output requires q prefix values") + if batch_stop >= len(prefix): + raise ValueError(f"{self.operator_name} logical B slice exceeds q prefix") + if sequence_slice is None: + token_start, token_stop = prefix[batch_start], prefix[batch_stop] + else: + token_start = prefix[batch_start] + sequence_slice[0] + token_stop = prefix[batch_start] + sequence_slice[1] + if token_stop > prefix[batch_start + 1]: + raise ValueError(f"{self.operator_name} logical S slice exceeds TND interval") + selector = [slice(token_start, token_stop, 1)] + else: + raise ValueError(f"{self.operator_name} unsupported query layout {layout_q!r}") + selector.extend([slice(None)] * (output.ndim - len(selector))) + return tuple(selector) + + def compare( + self, + output, + batch_consistency_id, + batch_axis, + batch_slice_info, + batch_seed, + compare_context, + ): + try: + relations = self.protocol.parse(batch_axis, batch_slice_info, batch_seed) + if relations is None: + return None + self.protocol.validate_id(batch_consistency_id, relations) + if output is None: + raise ValueError(f"{self.operator_name} batch output is None") + value = output.detach().cpu() if torch.is_tensor(output) else np.asarray(output) + attributes = dict(compare_context.attributes) if compare_context else {} + groups = {} + for relation in relations: + selected = value[self.output_selector(value, relation, attributes)] + axes, slices, seed = relation + stored = self.storage_bytes(selected) + relation_size = tuple(stop - start for start, stop, _step in slices) + groups.setdefault((axes, seed, relation_size), []).append(stored) + compared = 0 + for key, values in groups.items(): + if len(values) < 2: + continue + compared += 1 + if any(values[0] != item for item in values[1:]): + return { + "pass": False, + "precision": "batch_intra=FAIL", + "error_info": f"{self.operator_name} relation {key} differs", + } + if compared == 0: + return {"pass": True, "precision": "batch_intra=NOT_APPLICABLE"} + return {"pass": True, "precision": "batch_intra=PASS"} + except (IndexError, TypeError, ValueError) as error: + return { + "pass": False, + "precision": "batch_config=FAIL", + "error_info": str(error), + } + + +def normalize_indexer_inputs(data, attributes, operator_name, quantized=False): + protocol = BatchRelationProtocol(operator_name) + relations = protocol.parse( + attributes.get("batch_axis"), + attributes.get("batch_slice_info"), + attributes.get("batch_seed"), + ) + if relations is not None: + IndexerBatchInputNormalizer(data, attributes, operator_name, quantized).apply(relations) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/compare.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/compare.py new file mode 100644 index 000000000000..693505ba7022 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/compare.py @@ -0,0 +1,222 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +"""TTK result adapter for the QuantLightningIndexer V2 pytest TopK comparison.""" + +import importlib.util +import logging +import sys +import threading +from pathlib import Path + +import numpy as np +import torch + + +class PytestV2TopKComparator: + """Run the pytest V2 TopK compare with replay-safe data from the TestSpec.""" + + def __init__(self): + self.module = None + self.lock = threading.Lock() + + def load_module(self): + if self.module is not None: + return self.module + with self.lock: + if self.module is not None: + return self.module + pytest_dir = Path(__file__).resolve().parents[2] / "pytest" + module_path = pytest_dir / "result_compare_method.py" + module_name = "qli_v2_ttk_pytest_compare" + inserted = str(pytest_dir) not in sys.path + original_basic_config = logging.basicConfig + if inserted: + sys.path.insert(0, str(pytest_dir)) + try: + logging.basicConfig = lambda *args, **kwargs: None + spec = importlib.util.spec_from_file_location(module_name, module_path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {module_path}") + module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = module + spec.loader.exec_module(module) + self.module = module + except Exception as exc: + sys.modules.pop(module_name, None) + raise RuntimeError( + "Failed to load QuantLightningIndexerV2 pytest compare; " + f"module={module_path.resolve()}; " + f"original error: {type(exc).__name__}: {exc}" + ) from exc + finally: + logging.basicConfig = original_basic_config + if inserted and str(pytest_dir) in sys.path: + sys.path.remove(str(pytest_dir)) + return self.module + + @staticmethod + def to_torch(value): + if value is None: + return None + if torch.is_tensor(value): + return value.detach().cpu().clone() + array = np.array(value, copy=True, order="C") + dtype_name = str(array.dtype) + custom_dtypes = { + "bfloat16": (np.uint16, torch.bfloat16), + "float8_e4m3fn": (np.uint8, torch.float8_e4m3fn), + "float8_e5m2": (np.uint8, torch.float8_e5m2), + } + if dtype_name in custom_dtypes: + storage_dtype, torch_dtype = custom_dtypes[dtype_name] + storage = np.ascontiguousarray(array).view(storage_dtype) + return torch.from_numpy(storage).view(torch_dtype).reshape(array.shape) + return torch.from_numpy(array) + + @staticmethod + def result_dict(result, stage): + if not isinstance(result, (list, tuple)) or len(result) < 2: + raise ValueError(f"pytest {stage} returned invalid result: {result!r}") + status, precision = result[:2] + passed = str(status).strip().lower() == "pass" + return { + "pass": passed, + "precision": float(precision), + "error_info": None if passed else (f"pytest QuantLightningIndexerV2 {stage} returned {status!r}"), + } + + def compare(self, *outputs, compare_data=None): + if compare_data is None: + raise ValueError("QuantLightningIndexerV2 pytest compare data is unavailable") + if len(outputs) < 2 or len(outputs) % 2 != 0: + return { + "pass": False, + "precision": "invalid", + "error_info": "compare expects NPU outputs followed by golden outputs", + } + params = compare_data.get("params") + topk_value = compare_data.get("topk_value") + if params is None or topk_value is None: + raise ValueError("QuantLightningIndexerV2 pytest compare data lacks params or topk_value") + half = len(outputs) // 2 + npu_outputs = outputs[:half] + golden_outputs = outputs[half:] + if tuple(getattr(npu_outputs[0], "shape", ())) != tuple(getattr(golden_outputs[0], "shape", ())): + return { + "pass": False, + "precision": "shape_mismatch", + "error_info": ( + "index output shape mismatch: " + f"npu={getattr(npu_outputs[0], 'shape', None)}, " + f"golden={getattr(golden_outputs[0], 'shape', None)}" + ), + } + return_value = bool(params[-2]) + if return_value and half < 2: + return { + "pass": False, + "precision": "missing_output", + "error_info": "return_value is enabled but the NPU sparse-value output is missing", + } + npu_values = npu_outputs[1] if half > 1 else torch.empty(0) + golden_values = golden_outputs[1] if half > 1 else torch.empty(0) + if return_value and tuple(getattr(npu_values, "shape", ())) != tuple(getattr(golden_values, "shape", ())): + return { + "pass": False, + "precision": "shape_mismatch", + "error_info": ( + "sparse-value output shape mismatch: " + f"npu={getattr(npu_values, 'shape', None)}, " + f"golden={getattr(golden_values, 'shape', None)}" + ), + } + npu_indices = self.to_torch(npu_outputs[0]) + npu_values = self.to_torch(npu_values) + golden_values = self.to_torch(golden_values) + if return_value: + npu_values, sort_order = npu_values.sort(dim=-1, descending=True) + npu_indices = torch.gather(npu_indices, dim=-1, index=sort_order) + golden_indices = self.to_torch(golden_outputs[0]) + topk_value = self.to_torch(topk_value) + output_idx_offset = self.to_torch(compare_data.get("output_idx_offset")) + golden_values_for_index = golden_values.detach().cpu().float().numpy() + npu_values_for_index = npu_values.detach().cpu().float().numpy() + module = self.load_module() + index_result = module.check_result( + golden_indices, + npu_indices, + topk_value, + output_idx_offset, + params, + golden_values_for_index, + npu_values_for_index, + ) + results = [self.result_dict(index_result, "index compare")] + if return_value: + value_result = module.check_result_return_value( + golden_values, + npu_values, + params, + golden_indices, + npu_indices, + topk_value, + output_idx_offset, + ) + results.append(self.result_dict(value_result, "sparse-value compare")) + return results + + +def load_batch_comparator(): + name = "qli_v2_ttk_batch_consistency" + path = Path(__file__).with_name("batch_consistency.py") + if name in sys.modules: + return sys.modules[name] + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +COMPARATOR = PytestV2TopKComparator() + + +def compare( + *outputs, + compare_data=None, + batch_consistency_id=None, + batch_axis=None, + batch_slice_info=None, + batch_seed=None, + compare_context=None, +): + """Compare V2 TopK outputs with the canonical pytest policy.""" + results = COMPARATOR.compare(*outputs, compare_data=compare_data) + if not isinstance(results, list) or not all(result["pass"] for result in results): + return results + batch_result = ( + load_batch_comparator() + .IndexerBatchOutputComparator("QLI_V2") + .compare( + outputs[0], + batch_consistency_id, + batch_axis, + batch_slice_info, + batch_seed, + compare_context, + ) + ) + if batch_result is not None: + results.append(batch_result) + return results diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/golden.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/golden.py new file mode 100644 index 000000000000..2ec356958933 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/golden.py @@ -0,0 +1,207 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +"""CPU Golden adapter for QuantLightningIndexer V2 TTK cases.""" + +import importlib.util +import sys +from pathlib import Path + +import numpy as np +import torch + +PYTEST_MODULE_NAME = "qli_v2_pytest_golden" +PYTEST_MODULE_FILE = "quant_lightning_indexer_v2_golden.py" + + +class CaseDataStore: + """Share pytest data in-process and return compact metadata API inputs.""" + + def __init__(self): + self.case_data = {} + self.active_testcase_name = None + + def clear(self): + self.case_data.clear() + self.active_testcase_name = None + + def put(self, testcase_name, data): + if testcase_name is not None: + self.case_data[str(testcase_name)] = data + + def get(self, testcase_name): + if testcase_name is None: + return None + return self.case_data.get(str(testcase_name)) + + def discard(self, data): + for testcase_name, stored in tuple(self.case_data.items()): + if stored is data: + self.case_data.pop(testcase_name, None) + + +CASE_DATA = CaseDataStore() + + +def load_pytest_golden(): + """Load the pytest CPU reference only when the Golden stage needs it.""" + if PYTEST_MODULE_NAME in sys.modules: + return sys.modules[PYTEST_MODULE_NAME] + pytest_dir = Path(__file__).resolve().parents[2] / "pytest" + path = pytest_dir / PYTEST_MODULE_FILE + inserted = str(pytest_dir) not in sys.path + if inserted: + sys.path.insert(0, str(pytest_dir)) + try: + spec = importlib.util.spec_from_file_location(PYTEST_MODULE_NAME, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[PYTEST_MODULE_NAME] = module + spec.loader.exec_module(module) + return module + except Exception as exc: + sys.modules.pop(PYTEST_MODULE_NAME, None) + raise RuntimeError( + "Failed to load QuantLightningIndexer V2 pytest Golden module; " + f"module={path.resolve()}; original error: {type(exc).__name__}: {exc}" + ) from exc + finally: + if inserted: + sys.path.remove(str(pytest_dir)) + + +def get_case_data(testcase_name): + return CASE_DATA.get(testcase_name) + + +def materialize_golden(data): + if data.get("cpu_result") is None: + load_pytest_golden().generate_cpu_golden(data) + return data + + +def activate_case_data(testcase_name): + data = CASE_DATA.get(testcase_name) + if data is None: + raise RuntimeError("QuantLightningIndexer V2 Golden requires pytest data from the input stage") + CASE_DATA.active_testcase_name = str(testcase_name) + return materialize_golden(data) + + +def get_compare_data(testcase_name): + if testcase_name is None: + testcase_name = CASE_DATA.active_testcase_name + if testcase_name is None: + return None + data = CASE_DATA.get(testcase_name) + return None if data is None else materialize_golden(data) + + +def set_compare_data(testcase_name, data): + name = str(testcase_name) + CASE_DATA.active_testcase_name = name + CASE_DATA.case_data[name] = data + + +def discard_compare_data(data): + CASE_DATA.discard(data) + CASE_DATA.active_testcase_name = None + + +def cpu_quant_lightning_indexer_v2( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + *, + return_value=0, + testcase_name=None, + **kwargs, +): + """Materialize Golden from the exact pytest data produced by input.""" + del query, key, weights, query_dequant_scale, key_dequant_scale + del topk, quant_mode, kwargs + data = activate_case_data(testcase_name) + if int(return_value): + sparse_value = data["cpu_topk_value"] + else: + sparse_value = torch.zeros(0, dtype=data["topk_value"].dtype) + return data["cpu_result"], sparse_value + + +def cpu_aclnn_qli_v2( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + block_table, + output_idx_offset, + metadata, + topk, + quant_mode, + max_seqlen_q, + layout_q, + layout_k, + mask_mode, + cmp_ratio, + return_value, + sparse_indices_out, + sparse_values_out, + testcase_name=None, + **kwargs, +): + """Return the pytest Golden for the ACLNN C API parameter order.""" + del ( + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + block_table, + output_idx_offset, + metadata, + max_seqlen_q, + layout_q, + layout_k, + mask_mode, + cmp_ratio, + sparse_indices_out, + ) + sparse_indices, sparse_values = cpu_quant_lightning_indexer_v2( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + return_value=return_value, + testcase_name=testcase_name, + **kwargs, + ) + if not int(return_value): + if sparse_values_out is None: + raise ValueError("ACLNN QLI_V2 requires the sparseValuesOut tensor slot") + if torch.is_tensor(sparse_values_out): + sparse_values = torch.zeros(tuple(sparse_values_out.shape), dtype=sparse_values_out.dtype) + else: + sparse_values = np.zeros_like(np.asarray(sparse_values_out)) + return sparse_indices, sparse_values diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/inputs.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/inputs.py new file mode 100644 index 000000000000..14a0937e0cb1 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/inputs.py @@ -0,0 +1,909 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +"""Input customization for QuantLightningIndexer V2 TTK cases.""" + +import importlib.util +import inspect +import sys +from pathlib import Path + +import numpy as np +import torch + +QUANT_MODE_MXFP4 = 5 +QUANT_MODE_MXFP8 = 3 +QUANT_MODE_HIF4 = 6 + + +def restore_mx_input_dtypes(query, key, query_scale, key_scale, quant_mode): + """Restore packed MX dtypes needed by the existing replay compare path.""" + quant_mode = int(quant_mode) + if quant_mode == QUANT_MODE_MXFP8: + qk_dtype = getattr(torch, "float8_e4m3fn", None) + elif quant_mode == QUANT_MODE_MXFP4: + qk_dtype = getattr(torch, "float4_e2m1fn_x2", None) + else: + return query, key, query_scale, key_scale + scale_dtype = getattr(torch, "float8_e8m0fnu", None) + if qk_dtype is None or scale_dtype is None: + raise RuntimeError("current PyTorch does not provide the requested MX dtype") + if query.dtype == torch.uint8: + query = query.view(qk_dtype) + if key.dtype == torch.uint8: + key = key.view(qk_dtype) + if query_scale.dtype == torch.uint8: + query_scale = query_scale.view(scale_dtype) + if key_scale.dtype == torch.uint8: + key_scale = key_scale.view(scale_dtype) + return query, key, query_scale, key_scale + + +class QuantLightningIndexerV2InputAdapter: + """Translate a TTK case and reuse the pytest input/golden generator.""" + + @staticmethod + def module_load_error(stage, path, exc): + return RuntimeError( + "Failed to load QuantLightningIndexerV2 module; " + f"stage={stage}; module={path.resolve()}; " + f"original error: {type(exc).__name__}: {exc}" + ) + + def __init__(self): + self.pytest_golden = None + self.pytest_normalizer = None + self.batch_consistency = None + + def load_batch_consistency(self): + if self.batch_consistency is not None: + return self.batch_consistency + name = "qli_v2_ttk_batch_consistency" + path = Path(__file__).with_name("batch_consistency.py") + try: + if name in sys.modules: + module = sys.modules[name] + else: + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + self.batch_consistency = module + return module + except Exception as exc: + sys.modules.pop(name, None) + raise self.module_load_error("assets batch consistency", path, exc) from exc + + @staticmethod + def load_golden_store(): + name = "qli_v2_ttk_golden" + if name in sys.modules: + return sys.modules[name] + path = Path(__file__).with_name("golden.py") + try: + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + except Exception as exc: + sys.modules.pop(name, None) + raise QuantLightningIndexerV2InputAdapter.module_load_error("assets Golden store", path, exc) from exc + return module + + @staticmethod + def load_metadata_protocol(): + name = "qli_v2_ttk_metadata_protocol" + if name in sys.modules: + return sys.modules[name] + path = Path(__file__).with_name("metadata_protocol.py") + try: + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + except Exception as exc: + sys.modules.pop(name, None) + raise QuantLightningIndexerV2InputAdapter.module_load_error("assets metadata protocol", path, exc) from exc + return module + + def load_pytest_golden(self): + if self.pytest_golden is not None: + return self.pytest_golden + pytest_dir = Path(__file__).resolve().parents[2] / "pytest" + path = pytest_dir / "quant_lightning_indexer_v2_golden.py" + name = "qli_v2_pytest_golden" + inserted = str(pytest_dir) not in sys.path + if inserted: + sys.path.insert(0, str(pytest_dir)) + try: + if name in sys.modules: + module = sys.modules[name] + else: + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + self.pytest_golden = module + return module + except Exception as exc: + sys.modules.pop(name, None) + raise self.module_load_error("pytest Golden", path, exc) from exc + finally: + if inserted: + sys.path.remove(str(pytest_dir)) + + def load_pytest_normalizer(self): + if self.pytest_normalizer is not None: + return self.pytest_normalizer + pytest_dir = Path(__file__).resolve().parents[2] / "pytest" + path = pytest_dir / "qliv2_parameter_normalization.py" + name = "qli_v2_pytest_normalizer" + try: + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + self.pytest_normalizer = module + return module + except Exception as exc: + sys.modules.pop(name, None) + raise self.module_load_error("pytest parameter normalizer", path, exc) from exc + + @staticmethod + def list_value(kwargs, name): + value = kwargs.get(f"{name}_values") + if value is None: + return None + if torch.is_tensor(value): + value = value.detach().cpu().reshape(-1).tolist() + elif isinstance(value, np.ndarray): + value = value.reshape(-1).tolist() + return [int(item) for item in value] + + @staticmethod + def tensor_dtype(tensor): + if torch.is_tensor(tensor): + return tensor.dtype + if "hifloat8" in str(tensor.dtype): + return torch.uint8 + return torch.from_numpy(np.asarray(tensor)).dtype + + @staticmethod + def to_cpu_tensor(tensor): + if tensor is None: + return None + if torch.is_tensor(tensor): + return tensor.detach().cpu() + array = np.asarray(tensor) + if "hifloat8" in str(array.dtype): + array = array.view(np.uint8) + return torch.from_numpy(np.array(array, copy=True)) + + @staticmethod + def prefix_lengths(value): + if not value: + return [] + return [int(value[index + 1]) - int(value[index]) for index in range(len(value) - 1)] + + @staticmethod + def data_range(input_ranges, index): + if input_ranges and index < len(input_ranges) and input_ranges[index] is not None: + return repr(list(input_ranges[index])) + return None + + @staticmethod + def qk_dtype_name(tensor, quant_mode): + if quant_mode == QUANT_MODE_MXFP8: + return "FLOAT8_E4M3FN" + if quant_mode == QUANT_MODE_MXFP4: + return "FLOAT4_E2M1FN_X2" + if quant_mode == QUANT_MODE_HIF4: + return "HIF4" + dtype = QuantLightningIndexerV2InputAdapter.tensor_dtype(tensor) + dtype_name = str(tensor.dtype) + if dtype == torch.int8: + return "INT8" + if "float8_e4m3fn" in dtype_name: + return "FLOAT8_E4M3FN" + if dtype == torch.uint8: + return "HIFLOAT8" + raise ValueError(f"unsupported QLI_V2 q/k dtype: {tensor.dtype}") + + @staticmethod + def dequant_dtype_name(tensor, quant_mode): + if quant_mode in (QUANT_MODE_MXFP8, QUANT_MODE_MXFP4): + return "FLOAT8_E8M0FNU" + dtype = QuantLightningIndexerV2InputAdapter.tensor_dtype(tensor) + if dtype == torch.float16: + return "FP16" + if dtype == torch.float32: + return "FP32" + raise ValueError(f"unsupported QLI_V2 dequant dtype: {tensor.dtype}") + + @staticmethod + def weight_dtype_name(tensor): + mapping = { + torch.int8: "INT8", + torch.uint8: "UINT8", + torch.float16: "FP16", + torch.float32: "FP32", + torch.bfloat16: "BF16", + } + dtype = QuantLightningIndexerV2InputAdapter.tensor_dtype(tensor) + if dtype in mapping: + return mapping[dtype] + if "float8_e4m3fn" in str(tensor.dtype): + return "FLOAT8_E4M3FN" + raise ValueError(f"unsupported QLI_V2 weight dtype: {tensor.dtype}") + + @staticmethod + def pytest_uses_weight_dtype(pytest_golden): + parameters = inspect.signature(pytest_golden.GeneralizedQLIV2.__init__).parameters + return "weight_dtype" in parameters + + def geometry(self, query, key, layout_query, layout_key, cu_q, cu_k, seq_q, seq_k, quant_mode): + q_lengths = self.prefix_lengths(cu_q) or (seq_q or []) + k_lengths = self.prefix_lengths(cu_k) or (seq_k or []) + if layout_query == "BSND": + batch_size, q_seq, q_head_num, head_dim = [int(item) for item in query.shape] + q_t_size = 0 + elif layout_query == "TND": + q_t_size, q_head_num, head_dim = [int(item) for item in query.shape] + batch_size = len(q_lengths) + q_seq = max(q_lengths, default=q_t_size) + else: + raise ValueError(f"unsupported QLI_V2 query layout: {layout_query}") + + is_aclnn_float4 = isinstance(query, np.ndarray) and "float4" in str(query.dtype) + if quant_mode in (QUANT_MODE_MXFP4, QUANT_MODE_HIF4) and not is_aclnn_float4: + head_dim *= 2 + + if layout_key == "BSND": + _, k_seq, k_head_num, _ = [int(item) for item in key.shape] + k_t_size = 0 + block_size = 0 + block_num = 0 + elif layout_key == "TND": + k_t_size, k_head_num, _ = [int(item) for item in key.shape] + k_seq = max(k_lengths, default=k_t_size) + block_size = 0 + block_num = 0 + elif layout_key == "PA_BBND": + block_num, block_size, k_head_num, _ = [int(item) for item in key.shape] + capacity = block_num * block_size + per_batch_capacity = capacity // batch_size if batch_size > 0 else block_size + k_seq = max(k_lengths, default=per_batch_capacity) + k_t_size = 0 + else: + raise ValueError(f"unsupported QLI_V2 key layout: {layout_key}") + return ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + ) + + def build_case_params( + self, + query, + key, + weights, + query_dequant_scale, + layout_query, + layout_key, + kwargs, + include_weight_dtype=False, + ): + cu_q = self.list_value(kwargs, "cu_seqlens_q") + cu_k = self.list_value(kwargs, "cu_seqlens_k") + seq_q = self.list_value(kwargs, "seqused_q") + seq_k = self.list_value(kwargs, "seqused_k") + residual = self.list_value(kwargs, "cmp_residual_k") + quant_mode = kwargs.get("quant_mode") + quant_mode = None if quant_mode is None else int(quant_mode) + geometry = self.geometry( + query, + key, + layout_query, + layout_key, + cu_q, + cu_k, + seq_q, + seq_k, + quant_mode, + ) + input_ranges = kwargs.get("qli_input_ranges") or kwargs.get("input_ranges") or () + output_range = kwargs.get("output_idx_offset_range") + output_range = None if output_range is None else repr(list(output_range)) + max_seqlen_q = kwargs.get("max_seqlen_q") + max_seqlen_q = None if max_seqlen_q is None else int(max_seqlen_q) + dtype_params = (self.qk_dtype_name(query, quant_mode),) + if include_weight_dtype: + dtype_params += (self.weight_dtype_name(weights),) + dtype_params += ( + self.dequant_dtype_name(query_dequant_scale, quant_mode), + "INT32", + ) + params = ( + geometry + + dtype_params + + ( + cu_q, + cu_k, + seq_q, + seq_k, + residual, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + kwargs.get("sparse_count"), + kwargs.get("sparse_mode"), + self.data_range(input_ranges, 0), + self.data_range(input_ranges, 1), + self.data_range(input_ranges, 2), + self.data_range(input_ranges, 3), + self.data_range(input_ranges, 4), + kwargs.get("cmp_ratio"), + kwargs.get("return_value"), + output_range, + ) + ) + return self.load_pytest_normalizer().normalize_qliv2_params(params) + + @staticmethod + def copy_tensor(dst, src, name, packed_dtype=None): + if dst is None: + if src is not None: + import logging + + logging.warning("%s is absent from CSV but pytest generator produced a tensor. Skipping copy.", name) + return + if src is None: + raise ValueError(f"{name} is present in CSV but pytest generator returned None") + src_cpu = QuantLightningIndexerV2InputAdapter.to_cpu_tensor(src) + src_cpu = src_cpu.contiguous() + if packed_dtype is not None: + src_cpu = src_cpu.to(packed_dtype) + if torch.is_tensor(dst) and dst.dtype == torch.uint8 and src_cpu.element_size() == 1: + src_cpu = src_cpu.view(torch.uint8) + if tuple(dst.shape) != tuple(src_cpu.shape): + raise ValueError(f"{name} shape mismatch: TTK={tuple(dst.shape)} pytest={tuple(src_cpu.shape)}") + if torch.is_tensor(dst): + src_tensor = torch.as_tensor(src_cpu) + dst.copy_(src_tensor.to(dtype=dst.dtype, device=dst.device)) + return + + dst_array = np.asarray(dst) + if "float8" in str(src_cpu.dtype): + # ACLNN keeps this storage as a NumPy custom dtype. Copy its + # encoded bytes directly so assets do not depend on TTK's dtype helpers. + src_array = src_cpu.view(torch.uint8).numpy() + np.copyto(dst_array.view(np.uint8), src_array) + return + else: + src_array = np.asarray(src_cpu) + if "hifloat8" in str(dst_array.dtype): + np.copyto(dst_array.view(np.uint8), src_array.view(np.uint8)) + else: + np.copyto(dst_array, src_array.astype(dst_array.dtype, copy=False)) + + @staticmethod + def tensor_values(tensor): + if tensor is None: + return None + tensor = QuantLightningIndexerV2InputAdapter.to_cpu_tensor(tensor) + return [int(value) for value in tensor.reshape(-1).tolist()] + + @staticmethod + def unpack_mxfp4(tensor, fp4_values): + """Unpack two FP4 E2M1 values stored in each uint8 byte.""" + packed = tensor.view(torch.uint8).contiguous() + codes = torch.stack( + (packed & 0x0F, packed >> 4), + dim=-1, + ).flatten(-2) + return fp4_values[codes.to(torch.long)] + + @staticmethod + def restore_paged_tensor(tensor, block_table, batch_size, sequence_length): + """Restore a paged key or scale tensor for the pytest compare model.""" + physical = QuantLightningIndexerV2InputAdapter.to_cpu_tensor(tensor) + table = QuantLightningIndexerV2InputAdapter.to_cpu_tensor(block_table).to(torch.int64) + if physical.ndim < 3: + raise ValueError(f"paged tensor must have at least 3 dimensions, got shape {tuple(physical.shape)}") + block_size, head_num = int(physical.shape[1]), int(physical.shape[2]) + trailing = tuple(int(dim) for dim in physical.shape[3:]) + logical = torch.zeros( + (batch_size, head_num, sequence_length, *trailing), + dtype=physical.dtype, + ) + for batch_idx in range(batch_size): + for logical_block, block_id_value in enumerate(table[batch_idx].tolist()): + if block_id_value < 0: + continue + if block_id_value >= physical.shape[0]: + raise ValueError(f"block id {block_id_value} exceeds paged block count") + start = logical_block * block_size + if start >= sequence_length: + break + count = min(block_size, sequence_length - start) + block = physical[block_id_value, :count] + permutation = (1, 0, *range(2, block.ndim)) + logical[batch_idx, :, start : start + count] = block.permute(*permutation) + return logical + + @staticmethod + def normalize_compare_attributes(compare_context): + attributes = dict(compare_context.attributes) + aliases = { + "topk": "sparse_count", + "mask_mode": "sparse_mode", + "layout_q": "layout_query", + "layout_k": "layout_key", + "quantMode": "quant_mode", + "maxSeqlenQ": "max_seqlen_q", + "layoutQOptional": "layout_query", + "layoutKOptional": "layout_key", + "maskMode": "sparse_mode", + "cmpRatio": "cmp_ratio", + "returnValue": "return_value", + } + for source, target in aliases.items(): + if target not in attributes and source in attributes: + attributes[target] = attributes[source] + return attributes + + def rebuild_compare_data(self, compare_context): + """Rebuild only the pytest TopK compare context from replayed inputs.""" + tensors = tuple(compare_context.input_tensors or ()) + if len(tensors) < 12: + raise ValueError("QuantLightningIndexerV2 compare context requires twelve input slots") + ( + query, + key, + weights, + query_scale, + key_scale, + cu_q, + cu_k, + seq_q, + seq_k, + residual, + block_table, + offset, + ) = tensors[:12] + attributes = self.normalize_compare_attributes(compare_context) + for name, tensor in ( + ("cu_seqlens_q", cu_q), + ("cu_seqlens_k", cu_k), + ("seqused_q", seq_q), + ("seqused_k", seq_k), + ("cmp_residual_k", residual), + ): + values = self.tensor_values(tensor) + if values is not None: + attributes[f"{name}_values"] = values + + layout_q = attributes.get("layout_query") + layout_k = attributes.get("layout_key") + if layout_q is None or layout_k is None: + raise ValueError("QLI_V2 replay compare requires layout_query and layout_key from attributes") + pytest_golden = self.load_pytest_golden() + uses_weight_dtype = self.pytest_uses_weight_dtype(pytest_golden) + params = self.build_case_params( + query, + key, + weights, + query_scale, + layout_q, + layout_k, + attributes, + include_weight_dtype=uses_weight_dtype, + ) + if uses_weight_dtype: + ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + _, + _, + _, + _, + cu_q_values, + cu_k_values, + seq_q_values, + seq_k_values, + residual_values, + max_seqlen_q, + quant_mode, + _, + _, + sparse_count, + sparse_mode, + _, + _, + _, + _, + _, + cmp_ratio, + return_value, + _, + ) = params + else: + ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + _, + _, + _, + cu_q_values, + cu_k_values, + seq_q_values, + seq_k_values, + residual_values, + max_seqlen_q, + quant_mode, + _, + _, + sparse_count, + sparse_mode, + _, + _, + _, + _, + _, + cmp_ratio, + return_value, + _, + ) = params + + q_lengths = self.prefix_lengths(cu_q_values) if layout_q == "TND" else (seq_q_values or [q_seq] * batch_size) + k_lengths = self.prefix_lengths(cu_k_values) if layout_k == "TND" else (seq_k_values or [k_seq] * batch_size) + residual_for_cpu = [0] * batch_size if cmp_ratio == 1 or sparse_mode == 0 else list(residual_values) + + def as_int_tensor(value): + return None if value is None else torch.tensor(value, dtype=torch.int32) + + cu_q_cpu = as_int_tensor(cu_q_values) + cu_k_cpu = as_int_tensor(cu_k_values) + seq_q_cpu = as_int_tensor(seq_q_values) + seq_k_cpu = as_int_tensor(seq_k_values) + query, key, query_scale, key_scale = restore_mx_input_dtypes( + self.to_cpu_tensor(query), + self.to_cpu_tensor(key), + self.to_cpu_tensor(query_scale), + self.to_cpu_tensor(key_scale), + quant_mode, + ) + weights = self.to_cpu_tensor(weights) + block_table = self.to_cpu_tensor(block_table) + offset = self.to_cpu_tensor(offset) + qk_dtype = query.dtype + if quant_mode == QUANT_MODE_MXFP4: + query = self.unpack_mxfp4(query, pytest_golden.FP4_E2M1_VALUES) + key = self.unpack_mxfp4(key, pytest_golden.FP4_E2M1_VALUES) + model_args = [ + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + ] + if uses_weight_dtype: + model_args.append(weights.dtype) + model_args.extend( + [ + query_scale.dtype, + torch.int32, + cu_q_cpu, + cu_k_cpu, + q_lengths, + k_lengths, + residual_for_cpu, + max_seqlen_q, + quant_mode, + layout_q, + layout_k, + sparse_count, + sparse_mode, + None, + None, + None, + None, + None, + cmp_ratio, + return_value, + ] + ) + model = pytest_golden.GeneralizedQLIV2(*model_args) + + key_for_cpu = key + key_scale_for_cpu = key_scale + if layout_k == "PA_BBND": + if block_table is None or not k_lengths: + raise ValueError("PA_BBND compare context requires block_table and seqused_k") + sequence_length = max(k_lengths) + key_for_cpu = self.restore_paged_tensor(key, block_table, batch_size, sequence_length) + if quant_mode != 4: + key_scale_for_cpu = self.restore_paged_tensor(key_scale, block_table, batch_size, sequence_length) + + query_scale_for_cpu = query_scale + if quant_mode == 4: + query_scale_for_cpu = torch.full( + tuple(query.shape[:-1]), + query_scale.reshape(-1)[0].item(), + dtype=query_scale.dtype, + ) + if layout_k == "PA_BBND": + key_scale_shape = (batch_size, k_head_num, max(k_lengths)) + else: + key_scale_shape = tuple(key.shape[:-1]) + key_scale_for_cpu = torch.full( + key_scale_shape, + key_scale.reshape(-1)[0].item(), + dtype=key_scale.dtype, + ) + + _, scores, _ = model.forward( + query, + key_for_cpu, + weights, + query_scale_for_cpu, + key_scale_for_cpu, + cu_q_cpu, + cu_k_cpu, + seq_q_cpu, + seq_k_cpu, + block_table, + offset, + ) + return { + "params": params, + "scores": scores, + "topk_value": scores, + "output_idx_offset": None if offset is None else offset.detach().cpu(), + "score_layout": layout_q, + "cu_seqlens_q": cu_q_cpu, + } + + def customize( + self, + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + block_table, + output_idx_offset, + layout_query, + layout_key, + kwargs, + ): + pytest_golden = self.load_pytest_golden() + quant_mode = int(kwargs.get("quant_mode") or 1) + params = self.build_case_params( + query, + key, + weights, + query_dequant_scale, + layout_query, + layout_key, + kwargs, + include_weight_dtype=self.pytest_uses_weight_dtype(pytest_golden), + ) + batch_module = self.load_batch_consistency() + with batch_module.CaseRandomContext(kwargs): + data = pytest_golden.generate_qliv2_test_data(params, generate_golden=False) + batch_module.normalize_indexer_inputs(data, kwargs, "QLI_V2", quantized=True) + for name, dst, src_name in ( + ("query", query, "query"), + ("key", key, "key"), + ("weights", weights, "weights"), + ("query_dequant_scale", query_dequant_scale, "query_dequant_scale"), + ("key_dequant_scale", key_dequant_scale, "key_dequant_scale"), + ("cu_seqlens_q", cu_seqlens_q, "cu_seqlens_query"), + ("cu_seqlens_k", cu_seqlens_k, "cu_seqlens_key"), + ("seqused_q", seqused_q, "seqused_q"), + ("seqused_k", seqused_k, "seqused_k"), + ("cmp_residual_k", cmp_residual_k, "cmp_residual_k_for_npu"), + ("block_table", block_table, "block_table"), + ("output_idx_offset", output_idx_offset, "output_idx_offset"), + ): + src = data.get(src_name) + if ( + quant_mode == QUANT_MODE_MXFP4 + and name in ("query", "key") + and isinstance(dst, np.ndarray) + and "float4" in str(dst.dtype) + ): + src = self.unpack_mxfp4(src, pytest_golden.FP4_E2M1_VALUES) + packed_dtype = torch.float8_e4m3fn if quant_mode == QUANT_MODE_MXFP8 and name in ("query", "key") else None + self.copy_tensor(dst, src, name, packed_dtype) + return data + + +INPUT_ADAPTER = QuantLightningIndexerV2InputAdapter() + + +def rebuild_qli_v2_compare_data(compare_context): + return INPUT_ADAPTER.rebuild_compare_data(compare_context) + + +def zero_metadata(metadata): + if torch.is_tensor(metadata): + metadata.zero_() + else: + metadata[...] = 0 + + +def generate_qli_v2_inputs( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + *, + cu_seqlens_q=None, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=None, + cmp_residual_k=None, + block_table=None, + output_idx_offset=None, + metadata=None, + max_seqlen_q=-1, + layout_q="BSND", + layout_k="BSND", + mask_mode=0, + cmp_ratio=1, + return_value=0, + **kwargs, +): + """Populate pytest-derived inputs; metadata is filled by npu_preprocess.""" + if metadata is None: + raise ValueError("QLI_V2 direct API CSV must reserve the metadata tensor slot") + params = dict(kwargs) + params.update( + { + "sparse_count": topk, + "quant_mode": quant_mode, + "max_seqlen_q": max_seqlen_q, + "layout_query": layout_q, + "layout_key": layout_k, + "sparse_mode": mask_mode, + "cmp_ratio": cmp_ratio, + "return_value": return_value, + } + ) + data = INPUT_ADAPTER.customize( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + block_table, + output_idx_offset, + layout_q, + layout_k, + params, + ) + zero_metadata(metadata) + case_data = INPUT_ADAPTER.load_golden_store().CASE_DATA + testcase_name = params.get("testcase_name") + case_data.put(testcase_name, data) + INPUT_ADAPTER.load_metadata_protocol().save_metadata_inputs( + "quant_lightning_indexer_v2", testcase_name, data.get("metadata_input") + ) + return data + + +def generate_aclnn_qli_v2_inputs( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + block_table, + output_idx_offset, + metadata, + topk, + quant_mode, + max_seqlen_q, + layout_q, + layout_k, + mask_mode, + cmp_ratio, + return_value, + sparse_indices_out, + sparse_values_out, + **kwargs, +): + """Map the ACLNN C signature to the canonical pytest input adapter.""" + del sparse_indices_out, sparse_values_out + return generate_qli_v2_inputs( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + seqused_q=seqused_q, + seqused_k=seqused_k, + cmp_residual_k=cmp_residual_k, + block_table=block_table, + output_idx_offset=output_idx_offset, + metadata=metadata, + max_seqlen_q=max_seqlen_q, + layout_q=layout_q, + layout_k=layout_k, + mask_mode=mask_mode, + cmp_ratio=cmp_ratio, + return_value=return_value, + **kwargs, + ) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/metadata_inputs.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/metadata_inputs.py new file mode 100644 index 000000000000..da4f287d6695 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/metadata_inputs.py @@ -0,0 +1,98 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""Input customization for standalone QuantLightningIndexerMetadata cases.""" + +import importlib.util +import sys +from pathlib import Path + +import numpy as np + +_VECTOR_NAMES = ( + "cu_seqlens_q", + "cu_seqlens_k", + "seqused_q", + "seqused_k", + "cmp_residual_k", +) + + +def load_metadata_protocol(): + name = "qli_v2_ttk_metadata_protocol" + if name in sys.modules: + return sys.modules[name] + path = Path(__file__).with_name("metadata_protocol.py") + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + try: + spec.loader.exec_module(module) + except Exception: + sys.modules.pop(name, None) + raise + return module + + +def load_sidecar_values(kwargs): + metadata_input = load_metadata_protocol().load_metadata_inputs( + "quant_lightning_indexer_v2", kwargs.get("testcase_name") + ) + return {} if metadata_input is None else metadata_input + + +def copy_values(target, values, name): + if target is None: + return + if values is None: + raise ValueError(f"QuantLightningIndexerMetadata requires {name}_values") + source = np.asarray(values, dtype=np.int32) + if tuple(source.shape) != tuple(target.shape): + raise ValueError( + f"QuantLightningIndexerMetadata {name} shape mismatch: " + f"CSV={tuple(target.shape)}, values={tuple(source.shape)}" + ) + if hasattr(target, "copy_"): + import torch + + target.copy_(torch.as_tensor(source, dtype=target.dtype, device=target.device)) + else: + np.copyto(target, source.astype(target.dtype, copy=False)) + + +def generate_quant_lightning_indexer_metadata_inputs( + num_heads_q, + num_heads_k, + head_dim, + topk, + quant_mode, + *, + cu_seqlens_q=None, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=None, + cmp_residual_k=None, + **kwargs, +): + """Copy explicit descriptor vectors into metadata API input tensors.""" + del num_heads_q, num_heads_k, head_dim, topk, quant_mode + sidecar = load_sidecar_values(kwargs) + tensors = ( + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + ) + for name, tensor in zip(_VECTOR_NAMES, tensors): + values = sidecar[name] if name in sidecar else kwargs.get(f"{name}_values") + copy_values(tensor, values, name) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/metadata_protocol.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/metadata_protocol.py new file mode 100644 index 000000000000..5769d6f864ff --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/metadata_protocol.py @@ -0,0 +1,210 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""Private five-operator bridge for TTK manual-data prepare and replay. + +The protocol is intentionally environment- and file-based. It does not import +TTK or depend on a runner/context object. TTK continues to own its case files; +the assets only keep a compact sidecar outside each case directory and replace +an existing metadata input atomically after a zero-placeholder replay. +""" + +import hashlib +import os +import tempfile +from pathlib import Path + +import numpy as np +import regex as re +import torch + +ENV_NAME = "MANUAL_DATA_DIRS" +PROTOCOL_VERSION = 1 +SIDECAR_DIRECTORY = ".ttk_asset_metadata" +_SAFE_CASE_NAME = re.compile(r"[^A-Za-z0-9_.-]+") +_SAFE_OPERATOR = re.compile(r"[A-Za-z0-9_.-]+") +_FORMATS = ("bin", "npy", "pt") + + +def manual_data_roots(): + """Resolve the optional private roots from one shell environment value.""" + value = os.getenv(ENV_NAME) + if not value: + return () + roots = [] + seen = set() + for item in value.split(os.pathsep): + if not item.strip(): + continue + root = Path(item).expanduser().resolve() + if root not in seen: + roots.append(root) + seen.add(root) + return tuple(roots) + + +def case_directory_name(testcase_name): + """Mirror TTK's stable case-directory rule without importing TTK.""" + name = str(testcase_name) + safe = _SAFE_CASE_NAME.sub("_", name).strip("._") or "case" + if safe == name and len(safe) <= 120: + return safe + digest = hashlib.sha256(name.encode("utf-8")).hexdigest()[:12] + return f"{safe[:96]}-{digest}" + + +def validate_operator(operator): + if _SAFE_OPERATOR.fullmatch(str(operator)) is None: + raise ValueError(f"invalid metadata protocol operator: {operator!r}") + + +def clone_to_cpu(value): + if torch.is_tensor(value): + return value.detach().cpu().clone() + if isinstance(value, np.ndarray): + return torch.from_numpy(np.array(value, copy=True)) + if isinstance(value, dict): + return {name: clone_to_cpu(item) for name, item in value.items()} + if isinstance(value, list): + return [clone_to_cpu(item) for item in value] + if isinstance(value, tuple): + return tuple(clone_to_cpu(item) for item in value) + if value is None or isinstance(value, (str, bool, int, float)): + return value + if hasattr(value, "item"): + return value.item() + raise TypeError(f"unsupported metadata sidecar value: {type(value).__name__}") + + +def build_sidecar_path(root, operator, testcase_name): + validate_operator(operator) + return root / SIDECAR_DIRECTORY / str(operator) / f"{case_directory_name(testcase_name)}.pt" + + +def save_metadata_inputs(operator, testcase_name, metadata_input): + """Atomically save exact CPU metadata arguments during input preparation.""" + roots = manual_data_roots() + if not roots or testcase_name is None: + return None + if not isinstance(metadata_input, dict): + raise ValueError(f"{operator} pytest data lacks metadata_input") + + path = build_sidecar_path(roots[0], operator, testcase_name) + if path.parent.is_symlink(): + raise ValueError(f"metadata sidecar directory must not be a symlink: {path.parent}") + path.parent.mkdir(parents=True, exist_ok=True) + payload = { + "version": PROTOCOL_VERSION, + "operator": str(operator), + "testcase_name": str(testcase_name), + "metadata_input": clone_to_cpu(metadata_input), + } + temporary = path.with_name(f".{path.name}.{os.getpid()}.tmp") + try: + torch.save(payload, temporary) + os.replace(temporary, path) + finally: + temporary.unlink(missing_ok=True) + return path + + +def load_pt(path): + try: + return torch.load(path, map_location="cpu", weights_only=True) + except TypeError: + return torch.load(path, map_location="cpu") + + +def load_metadata_inputs(operator, testcase_name): + """Load the first matching sidecar, or return None for fallback derivation.""" + if testcase_name is None: + return None + for root in manual_data_roots(): + path = build_sidecar_path(root, operator, testcase_name) + if not path.exists(): + continue + if path.is_symlink() or not path.is_file(): + raise ValueError(f"metadata sidecar must be a regular file: {path}") + payload = load_pt(path) + if ( + not isinstance(payload, dict) + or payload.get("version") != PROTOCOL_VERSION + or payload.get("operator") != str(operator) + or payload.get("testcase_name") != str(testcase_name) + or not isinstance(payload.get("metadata_input"), dict) + ): + raise ValueError(f"incompatible metadata sidecar: {path}") + return payload["metadata_input"] + return None + + +def metadata_is_materialized(metadata): + """Treat a nonzero metadata slot as authoritative in every execution mode.""" + if metadata is None: + return False + if torch.is_tensor(metadata): + return bool(torch.count_nonzero(metadata).item()) + return bool(np.count_nonzero(np.asarray(metadata))) + + +def metadata_to_array(metadata): + if torch.is_tensor(metadata): + return metadata.detach().cpu().contiguous().numpy() + return np.ascontiguousarray(np.asarray(metadata)) + + +def find_metadata_file(root, testcase_name, metadata_index): + case_dir = root / case_directory_name(testcase_name) + if not case_dir.exists(): + return None + if case_dir.is_symlink() or not case_dir.is_dir(): + raise ValueError(f"manual-data testcase path must be a regular directory: {case_dir}") + for file_format in _FORMATS: + matches = tuple(case_dir.glob(f"input_{int(metadata_index)}_*.{file_format}")) + if len(matches) > 1: + raise RuntimeError(f"expected one metadata input[{metadata_index}], found {len(matches)} in {case_dir}") + if matches: + path = matches[0] + if path.is_symlink() or not path.is_file(): + raise ValueError(f"metadata input must be a regular file: {path}") + return path + return None + + +def write_array(path, array): + with tempfile.NamedTemporaryFile(prefix=f".{path.name}.", suffix=".tmp", dir=path.parent, delete=False) as stream: + temporary = Path(stream.name) + try: + if path.suffix == ".bin": + array.tofile(temporary) + elif path.suffix == ".npy": + with temporary.open("wb") as stream: + np.save(stream, array, allow_pickle=False) + elif path.suffix == ".pt": + torch.save(torch.from_numpy(np.array(array, copy=True)), temporary) + else: + raise ValueError(f"unsupported metadata input format: {path.suffix}") + os.replace(temporary, path) + finally: + temporary.unlink(missing_ok=True) + + +def rewrite_metadata_input(operator, testcase_name, metadata_index, metadata): + """Replace the existing zero placeholder after successful metadata execution.""" + validate_operator(operator) + if testcase_name is None: + return None + for root in manual_data_roots(): + path = find_metadata_file(root, testcase_name, metadata_index) + if path is None: + continue + write_array(path, metadata_to_array(metadata)) + return path + return None diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/npu_preprocess.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/npu_preprocess.py new file mode 100644 index 000000000000..8aa9ff40ecce --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/impl/npu_preprocess.py @@ -0,0 +1,323 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""Populate QuantLightningIndexer V2 metadata without TTK state coupling.""" + +import importlib.util +import logging +import sys +from pathlib import Path + +import torch + +OPERATOR = "quant_lightning_indexer_v2" +METADATA_INDEX = 12 +QUANT_MODE_MXFP8 = 3 +QUANT_MODE_MXFP4 = 5 +QUANT_MODE_HIF4 = 6 +ACLNN_PARAMETER_NAMES = ( + "query", + "key", + "weights", + "query_dequant_scale", + "key_dequant_scale", + "cu_seqlens_q", + "cu_seqlens_k", + "seqused_q", + "seqused_k", + "cmp_residual_k", + "block_table", + "output_idx_offset", + "metadata", + "topk", + "quant_mode", + "max_seqlen_q", + "layout_q", + "layout_k", + "mask_mode", + "cmp_ratio", + "return_value", + "sparse_indices_out", + "sparse_values_out", +) + + +def load_metadata_protocol(): + name = "qli_v2_ttk_metadata_protocol" + if name in sys.modules: + return sys.modules[name] + path = Path(__file__).with_name("metadata_protocol.py") + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + try: + spec.loader.exec_module(module) + except Exception: + sys.modules.pop(name, None) + raise + return module + + +def get_attribute(kwargs, name, default=None, aliases=()): + for key in (name, f"pytest_{name}", *aliases): + value = kwargs.get(key) + if value is not None: + return value + return default + + +def get_values(kwargs, name, tensor): + value = kwargs.get(f"{name}_values") + if value is None: + value = tensor + if value is None: + return None + if torch.is_tensor(value): + return value.detach().cpu().reshape(-1).tolist() + return [int(item) for item in value] + + +def max_sequence(prefix, used, fallback): + if used: + return max(int(value) for value in used) + if prefix and len(prefix) > 1: + return max(int(prefix[index + 1]) - int(prefix[index]) for index in range(len(prefix) - 1)) + return int(fallback) + + +def restore_mx_dtypes(query, key, query_scale, key_scale, quant_mode): + """Support older uint8-storage CSVs; native TTK dtypes pass through.""" + quant_mode = int(quant_mode) + if quant_mode == QUANT_MODE_MXFP8: + qk_dtype = getattr(torch, "float8_e4m3fn", None) + elif quant_mode == QUANT_MODE_MXFP4: + qk_dtype = getattr(torch, "float4_e2m1fn_x2", None) + else: + return + scale_dtype = getattr(torch, "float8_e8m0fnu", None) + if qk_dtype is None or scale_dtype is None: + raise RuntimeError("current PyTorch does not provide the requested MX dtype") + for name, tensor, dtype in ( + ("query", query, qk_dtype), + ("key", key, qk_dtype), + ("query_dequant_scale", query_scale, scale_dtype), + ("key_dequant_scale", key_scale, scale_dtype), + ): + if not torch.is_tensor(tensor): + continue + if tensor.dtype == dtype: + continue + if tensor.dtype != torch.uint8: + raise TypeError(f"QLI_V2 {name} must use {dtype} or uint8 storage, got {tensor.dtype}") + tensor.data = tensor.data.view(dtype) + + +def build_metadata_arguments(query, key, topk, quant_mode, layout_q, layout_k, mask_mode, cmp_ratio, kwargs): + q_shape = tuple(int(value) for value in query.shape) + k_shape = tuple(int(value) for value in key.shape) + num_heads_q = q_shape[2] if layout_q == "BSND" else q_shape[1] + num_heads_k = k_shape[1] if layout_k == "TND" else k_shape[2] + head_dim = q_shape[-1] * (2 if int(quant_mode) in (QUANT_MODE_MXFP4, QUANT_MODE_HIF4) else 1) + cu_q = kwargs.get("cu_seqlens_q") + cu_k = kwargs.get("cu_seqlens_k") + seq_q = kwargs.get("seqused_q") + seq_k = kwargs.get("seqused_k") + cu_q_values = get_values(kwargs, "cu_seqlens_q", cu_q) + cu_k_values = get_values(kwargs, "cu_seqlens_k", cu_k) + seq_q_values = get_values(kwargs, "seqused_q", seq_q) + seq_k_values = get_values(kwargs, "seqused_k", seq_k) + + batch_size = get_attribute(kwargs, "batch_size") + if batch_size is None: + if seq_q_values is not None: + batch_size = len(seq_q_values) + elif cu_q_values is not None: + batch_size = len(cu_q_values) - 1 + elif layout_q == "BSND": + batch_size = q_shape[0] + else: + batch_size = 0 + + q_fallback = q_shape[1] if layout_q == "BSND" else q_shape[0] + if layout_k == "BSND": + k_fallback = k_shape[1] + elif layout_k == "TND": + k_fallback = k_shape[0] + else: + k_fallback = int(get_attribute(kwargs, "max_seqlen_k", k_shape[1])) + + return { + "num_heads_q": int(get_attribute(kwargs, "num_heads_q", num_heads_q, ("pytest_q_head_num",))), + "num_heads_k": int(get_attribute(kwargs, "num_heads_k", num_heads_k, ("pytest_k_head_num",))), + "head_dim": int(get_attribute(kwargs, "head_dim", head_dim)), + "topk": int(topk), + "quant_mode": int(quant_mode), + "cu_seqlens_q": cu_q, + "cu_seqlens_k": cu_k, + "seqused_q": seq_q, + "seqused_k": seq_k, + "cmp_residual_k": kwargs.get("cmp_residual_k"), + "batch_size": int(batch_size), + "max_seqlen_q": int( + get_attribute( + kwargs, + "metadata_max_seqlen_q", + max_sequence(cu_q_values, seq_q_values, q_fallback), + ) + ), + "max_seqlen_k": int( + get_attribute( + kwargs, + "metadata_max_seqlen_k", + max_sequence(cu_k_values, seq_k_values, k_fallback), + ) + ), + "layout_q": str(layout_q), + "layout_k": str(layout_k), + "mask_mode": int(mask_mode), + "cmp_ratio": int(cmp_ratio), + } + + +def move_to_device(value, target): + if value is None: + return None + if torch.is_tensor(value): + return value.to(device=target.device) + return torch.as_tensor(value, device=target.device) + + +def run_metadata(arguments, metadata): + return torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata( + int(arguments["num_heads_q"]), + int(arguments["num_heads_k"]), + int(arguments["head_dim"]), + int(arguments["topk"]), + int(arguments["quant_mode"]), + cu_seqlens_q=move_to_device(arguments.get("cu_seqlens_q"), metadata), + cu_seqlens_k=move_to_device(arguments.get("cu_seqlens_k"), metadata), + seqused_q=move_to_device(arguments.get("seqused_q"), metadata), + seqused_k=move_to_device(arguments.get("seqused_k"), metadata), + cmp_residual_k=move_to_device(arguments.get("cmp_residual_k"), metadata), + batch_size=int(arguments["batch_size"]), + max_seqlen_q=int(arguments["max_seqlen_q"]), + max_seqlen_k=int(arguments["max_seqlen_k"]), + layout_q=str(arguments["layout_q"]), + layout_k=str(arguments["layout_k"]), + mask_mode=int(arguments["mask_mode"]), + cmp_ratio=int(arguments["cmp_ratio"]), + ) + + +def run( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + *, + cu_seqlens_q=None, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=None, + cmp_residual_k=None, + metadata=None, + layout_q="BSND", + layout_k="BSND", + mask_mode=0, + cmp_ratio=1, + **kwargs, +): + """Generate metadata once, or reuse a nonzero manual-data input.""" + del weights + if metadata is None: + raise ValueError("QuantLightningIndexer V2 npu_preprocess requires metadata") + restore_mx_dtypes(query, key, query_dequant_scale, key_dequant_scale, quant_mode) + arguments_kwargs = dict(kwargs) + arguments_kwargs.update( + { + "cu_seqlens_q": cu_seqlens_q, + "cu_seqlens_k": cu_seqlens_k, + "seqused_q": seqused_q, + "seqused_k": seqused_k, + "cmp_residual_k": cmp_residual_k, + } + ) + protocol = load_metadata_protocol() + testcase_name = kwargs.get("testcase_name") + force_metadata_refresh = bool(get_attribute(kwargs, "metadata_refresh", False)) + if protocol.metadata_is_materialized(metadata) and not force_metadata_refresh: + logging.info("[%s] reuse nonzero QLI_V2 metadata input", testcase_name) + return None + arguments = protocol.load_metadata_inputs(OPERATOR, testcase_name) + if arguments is not None: + source = "manual-data sidecar" + else: + arguments = build_metadata_arguments( + query, + key, + topk, + quant_mode, + layout_q, + layout_k, + mask_mode, + cmp_ratio, + arguments_kwargs, + ) + source = "main API fallback (sidecar unavailable)" + logging.info( + "[%s] build QLI_V2 metadata from %s; forced=%s", + testcase_name, + source, + force_metadata_refresh, + ) + generated = run_metadata(arguments, metadata) + if tuple(metadata.shape) != tuple(generated.shape): + raise ValueError( + f"QLI_V2 metadata shape mismatch: placeholder={tuple(metadata.shape)}, generated={tuple(generated.shape)}" + ) + metadata.copy_(generated.to(dtype=metadata.dtype, device=metadata.device)) + rewritten = protocol.rewrite_metadata_input(OPERATOR, testcase_name, METADATA_INDEX, metadata) + if rewritten is not None: + logging.info("[%s] rewrote QLI_V2 metadata input: %s", testcase_name, rewritten) + return None + + +def run_aclnn(*args, **kwargs): + """Adapt the ACLNN main API order to the shared Torch metadata hook.""" + if len(args) != len(ACLNN_PARAMETER_NAMES): + raise ValueError( + f"QuantLightningIndexerV2 ACLNN hook expects {len(ACLNN_PARAMETER_NAMES)} arguments, got {len(args)}" + ) + values = dict(zip(ACLNN_PARAMETER_NAMES, args)) + host_metadata = values["metadata"] + metadata = move_to_device(host_metadata, torch.empty(0, device="npu")) + values["metadata"] = metadata + result = run( + values.pop("query"), + values.pop("key"), + values.pop("weights"), + values.pop("query_dequant_scale"), + values.pop("key_dequant_scale"), + values.pop("topk"), + values.pop("quant_mode"), + **values, + **kwargs, + ) + if torch.is_tensor(host_metadata): + if host_metadata.device != metadata.device: + host_metadata.copy_(metadata.to(dtype=host_metadata.dtype, device=host_metadata.device)) + else: + host_metadata[...] = metadata.detach().cpu().numpy() + return result diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/assets/spec.py b/csrc/attention/quant_lightning_indexer_v2/tests/assets/spec.py new file mode 100644 index 000000000000..c05988088d0c --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/assets/spec.py @@ -0,0 +1,93 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +"""TestSpec adapter for QuantLightningIndexer V2 TTK assets.""" + +import importlib.util +import sys +from pathlib import Path + +ASSET_IMPL_DIR = Path(__file__).with_name("impl") + + +def load_impl_module(stem): + name = f"qli_v2_ttk_{stem}" + if name in sys.modules: + return sys.modules[name] + path = ASSET_IMPL_DIR / f"{stem}.py" + try: + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"cannot create import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + except Exception as exc: + sys.modules.pop(name, None) + raise RuntimeError( + "Failed to load QuantLightningIndexerV2 assets module; " + f"stage=impl/{stem}; module={path.resolve()}; " + f"original error: {type(exc).__name__}: {exc}" + ) from exc + return module + + +npu_preprocess_module = load_impl_module("npu_preprocess") +golden_module = load_impl_module("golden") +inputs_module = load_impl_module("inputs") +metadata_inputs_module = load_impl_module("metadata_inputs") +compare_module = load_impl_module("compare") + + +class QuantLightningIndexerV2Spec: + golden = golden_module.cpu_quant_lightning_indexer_v2 + customize_inputs = inputs_module.generate_qli_v2_inputs + npu_preprocess = npu_preprocess_module.run + tolerance = { + "float16": {"standard": "stat_rel_err"}, + "bfloat16": {"standard": "stat_rel_err"}, + "float8_e4m3fn": {"standard": "stat_rel_err"}, + } + + def compare(*outputs, compare_context=None, **kwargs): + testcase_name = None if compare_context is None else compare_context.testcase_name + data = golden_module.get_compare_data(testcase_name) + if data is None: + if compare_context is None: + raise RuntimeError("QuantLightningIndexerV2 pytest compare requires compare_context") + data = inputs_module.rebuild_qli_v2_compare_data(compare_context) + golden_module.set_compare_data(compare_context.testcase_name, data) + try: + return compare_module.compare( + *outputs, + compare_data=data, + compare_context=compare_context, + **kwargs, + ) + finally: + golden_module.discard_compare_data(data) + + +class AclnnQuantLightningIndexerV2Spec(QuantLightningIndexerV2Spec): + golden = golden_module.cpu_aclnn_qli_v2 + customize_inputs = inputs_module.generate_aclnn_qli_v2_inputs + npu_preprocess = npu_preprocess_module.run_aclnn + + +class QuantLightningIndexerMetadataSpec: + customize_inputs = metadata_inputs_module.generate_quant_lightning_indexer_metadata_inputs + + +__spec__ = { + "torch.ops.cann_ops_transformer.quant_lightning_indexer": "QuantLightningIndexerV2Spec", + "torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata": ("QuantLightningIndexerMetadataSpec"), + "aclnnQuantLightningIndexerV2": "AclnnQuantLightningIndexerV2Spec", +} diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/README.md b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/README.md new file mode 100644 index 000000000000..784b5f82933b --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/README.md @@ -0,0 +1,259 @@ +# quant_lightning_indexer_v2算子测试框架 + +## 功能说明 + +基于pytest测试框架,实现quant_lightning_indexer_v2算子的功能验证: + +- **CPU侧**:复现算子功能用以生成golden数据 +- **NPU侧**:通过TorchNPU进行算子直调获取实际数据 +- **精度对比**:进行CPU与NPU结果的精度对比验证算子功能 +- **双模式执行隔离**:支持直接pytest多进程执行和shell层进程隔离两种批量模式 +- **PT 复跑模式**:`batch -P <目录>` 直接执行已有 PT;增加 `-E` 才会先生成 PT +- **可控批跑**:`-P` 统一指定 PT 生成和读取目录,支持结果路径、case 名和 1-based 序号选择 +- **性能采集**:支持挂载msprof采集算子性能数据并汇总输出 +- **运行模式切换**:支持eager直接调用和graph(torch.compile + torchair)两种算子调用模式 + +## 当前实现范围 + +### 参数限制 + +- **数据格式**: + - **query_layout**:BSND、TND + - **key_layout**: PA_BBND、BSND、TND + +- **数据类型**: + - **qk_dtype**: FLOAT8_E4M3FN、INT8、HIFLOAT8、FLOAT4_E2M1 + - **dequant_dtype**: FP32(Ascend950)、FP16(Ascend910_93)、FLOAT8_E8M0(MXFP8/MXFP4) + - **actual_seq_dtype**: INT32 + +### PyTorch MX类型约定 + +- `quant_mode`为3/5时,`query_dequant_scale`和`key_dequant_scale`必须为`torch.float8_e8m0fnu`。 +- `quant_mode`为5时,`query`和`key`的算子逻辑数据类型为`FLOAT4_E2M1`(文档简写为`float4_e2m1`): + - PyTorch提供`torch.float4_e2m1fn_x2`时,优先使用该原生打包类型;名称中的`x2`表示每个物理字节打包两个E2M1逻辑元素。 + - PyTorch未提供该类型时,使用`torch.uint8`承载已打包数据,封装层会将其按`ACL_FLOAT4_E2M1`传入算子。 + +- **运行模式**: + - **eager**:直接调用 `torch.ops.cann_ops_transformer.quant_lightning_indexer` + - **graph**:通过 `torch.compile` + `torchair` 后端编译执行(需torchair支持) + +### 环境配置 + +#### 前置要求 + +1、 确认TorchNPU为最新版本 +2、 激活CANN包和自定义算子包 +3、 graph模式需要安装torchair编译器后端 + +#### custom包调用 + +支持custom包调用 + +## 文件结构 + +### pytest文件结构说明 + +- test_run.sh # 执行脚本,支持single/batch/batch_exec三种命令 +- batch_isolated_run.sh # 批量隔离执行脚本(shell层进程隔离+msprof性能采集) +- quant_lightning_indexer_v2_golden.py # cpu侧算子golden实现 +- quant_lightning_indexer_v2_acl_graph.py # graph模式torchair后端实现 +- result_compare_method.py # cpu golden与npu输出精度对比 +- qliv2_test_utils.py # case选择、稳定命名和结果表公共逻辑 +- collect_perf_data.py # msprof性能数据收集与汇总 +- pytest.ini # 创建测试标记 + +单用例测试: + +- test_quant_lightning_indexer_v2_single.py # pytest测试单用例运行主程序 +- test_quant_lightning_indexer_v2_paramset.py # 单用例入参配置,按芯片型号自动选择用例 + +批量测试: + +- test_quant_lightning_indexer_v2_batch.py # 用例批量测试主程序并生成excel文件保存结果 +- ./batch/quant_lightning_indexer_v2_pt_loadprocess.py # 读取pt文件并调用算子获取npu输出 +- ./batch/quant_lightning_indexer_v2_pt_save.py # 读取excel表格批量生成用例pt文件 +- ./batch/list_pt_from_excel.py # 从Excel提取Testcase_Name并按名匹配pt文件(batch_exec模式用) + +## 架构说明 + +- **数据生成入口**:`generate_qliv2_test_data` 复用原有 batch 数据链生成输入和 CPU golden,不调用 metadata 或主算子;参数准备仍可查询设备信息 +- **single 模式**:直接执行配置用例;指定 `--save-pt` 时保存并执行同一份实际输入,避免二次随机生成 +- **batch 模式**:`-P` 是唯一 PT 目录;有 `-E` 时先生成再执行,无 `-E` 时直接执行已有 PT +- **batch_exec 模式**:按 Excel 的 `Testcase_Name` 筛选已有 PT,仅执行 NPU 和精度对比 +- 两路共用 `_qliv2_prepare_tensors_and_metadata` 和 `_qliv2_run_compiled_graph`,统一使用 `fullgraph=False` +- 结果表会先落盘;index 或 return value 精度结果为 `Failed` 时,pytest 随后以非零状态退出 + +## 使用方法 + +在pytest文件夹路径下执行: + +### 运行测试用例 + +#### 单用例调测 + +1、手动配置test_quant_lightning_indexer_v2_paramset.py的ENABLED_PARAMS参数 + +2、执行指令: + +``` bash +bash test_run.sh single +bash test_run.sh single --save-pt ./single_pt -O ./result/single.xlsx +bash test_run.sh single -M graph --save-pt ./single_pt +``` + +#### 用例的批量生成与测试 + +##### 方式A:test_run.sh 批量执行 + +1、excel路径下存放用例excel表格 + +`-P` 同时指定 PT 的保存目录和读取目录,默认是当前 pytest 目录下的 `pt_path`。 + +##### 直接执行已有 PT + +不传 `-E` 时不读取 Excel、不重新生成 PT,只执行 NPU 和 compare: + +3、执行指令: + +``` bash +bash test_run.sh batch -P ./pt_path +bash test_run.sh batch -P ./pt_path -O ./result/rerun.xlsx +bash test_run.sh batch -P ./pt_path -C case_b,case_a # 按名称和给定顺序 +bash test_run.sh batch -P ./pt_path -I 3,1,5-7 # 按自然排序后的序号 +bash test_run.sh batch -P ./pt_path -M graph +``` + +4、配置区默认值: + +| 变量 | 默认值 | 命令行参数 | 说明 | +|---|---|---|---| +| DEFAULT_EXCEL | `./excel/test_cases.xlsx` | `-E` | Excel 用例表格路径(**必须指定具体文件名**,不支持通配符如 `./excel/*`) | +| DEFAULT_PT_PATH | `./pt_path` | `-P` | pt 文件存放目录 | +| (无) | `Sheet1` | `-S` | Excel Sheet 页名 | +| (无) | `eager` | `-M` | 运行模式(eager/graph) | + +#### 根据 Excel 表格筛选已有 pt 批量执行(batch_exec 模式) +> +> 仅重新执行 NPU 测试和精度对比,不重新生成 pt 文件。适用于已有 pt 文件、只需更新精度结果的场景。 + +增加 `-E` 后,脚本先把 Excel 用例生成到 `-P`,再从同一个目录执行: + +2、执行指令: + +``` bash +bash test_run.sh batch -E ./excel/test_cases.xlsx -P ./pt_path +bash test_run.sh batch -E ./excel/test_cases.xlsx -S Sheet1 -P ./pt_path +bash test_run.sh batch -E ./excel/test_cases.xlsx -P ./pt_path -O ./result/batch.xlsx +``` + +3、执行流程: + +- 从 Excel 表格读取 `Testcase_Name` 列 +- 按 `.pt` 在 pt_path 下匹配对应的 .pt 文件 +- 仅对匹配到的 .pt 文件执行 NPU 测试和精度对比 +- 生成 `result.xlsx` 测试结果表格 +- 如果 Excel 中某条用例无对应的 .pt 文件,会输出警告并跳过该用例 + +4、与 `batch` 模式的区别: + +| | `batch` | `batch_exec` | +|---|---|---| +| pt 生成 | 每次重新生成 | 跳过 | +| 执行速度 | 较慢(含 pt 生成) | 较快 | +| 适用场景 | 首次运行 / 参数变更 | 精度复测 / 仅 NPU 结果更新 | + +##### 方式B:手工分步执行 + +1、生成pt文件: + +``` bash +python3 batch/quant_lightning_indexer_v2_pt_save.py excel/test_cases.xlsx pt_path +python3 batch/quant_lightning_indexer_v2_pt_save.py excel/test_cases.xlsx pt_path --sheet Sheet1 # 指定 Sheet 页 +``` + +2、替换测试脚本路径: + +``` bash +QLIV2_TESTCASE_DIR=pt_path QLIV2_RESULT_PATH=result.xlsx \ +python3 -m pytest -rA -s test_quant_lightning_indexer_v2_batch.py -v -m ci +``` + +3、执行测试: + +``` bash +python3 -m pytest -rA -s test_quant_lightning_indexer_v2_batch.py -v -m ci -W ignore::UserWarning -W ignore::DeprecationWarning +``` + +4、恢复测试脚本: + +``` bash +cp test_quant_lightning_indexer_v2_batch.py.bak test_quant_lightning_indexer_v2_batch.py +``` + +##### 方式C:批量隔离执行(推荐用于性能采集) + +对每条用例单独拉起一个pytest进程,实现进程间完全隔离,避免单条用例崩溃影响其他用例。 + +``` bash +bash batch_isolated_run.sh ./pt_path 0 # 不采集性能 +bash batch_isolated_run.sh ./pt_path 1 # 采集性能(挂载msprof) +bash batch_isolated_run.sh ./pt_path 0 graph # graph模式 + 不采集性能 +bash batch_isolated_run.sh ./pt_path 1 graph # graph模式 + 性能采集 +``` + +## Excel 用例表格式 + +`excel/test_cases.xlsx` 需包含以下列(Sheet1): + +| 列名 | 类型 | 示例 | +|---|---|---| +| Testcase_Name | str | `test_case_01` | +| batch_size | int | `8` | +| q_seq | int | `15` | +| k_seq | int | `111` | +| q_t_size | int | `8` | +| k_t_size | int | `15` | +| q_head_num | int | `64` | +| k_head_num | int | `1` | +| head_dim | int | `128` | +| block_size | int | `512` | +| block_num | int | `8` | +| qk_dtype | str | `FLOAT8_E4M3FN` / `INT8` / `HIFLOAT8` / `FLOAT4_E2M1` | +| dequant_dtype | str | `FP32` / `FP16` / `FLOAT8_E8M0` | +| actual_seq_dtype | str | `INT32` | +| cu_seqlens_q | None/str | `None` 或 `"[0, 1]"` | +| cu_seqlens_k | None/str | `None` 或 `"[0, 1]"` | +| seqused_q | None/str | `None` 或 `"[3,3,3,3,3,3,3,3]"` | +| seqused_k | str | `"[28,24,80,96,47,76,0,111]"` | +| cmp_residual_k | None/str | `None` 或 `"[0,0,0,0,0,0,0,0]"`(cmp_ratio>1时必填)| +| max_seqlen_q | int | `-1` | +| quant_mode | int | `1` / `2` / `4` | +| layout_query | str | `BSND` / `TND` | +| layout_key | str | `PA_BBND` | +| sparse_count | int | `512` | +| sparse_mode | int | `0` / `3` | +| query_datarange | str | `"[-448,448]"` | +| key_datarange | str | `"[-20,20]"` | +| weights_datarange | str | `"[-123,123]"` | +| q_scale_datarange | str | `"[0,255]"` | +| k_scale_datarange | str | `"[0,65504]"` | +| cmp_ratio | int | `1` / `4` | +| return_value | int | `0` / `1` | +| output_idx_offset | None/str | `None` 或列表字符串 | + +**注意事项**: + +- `dequant_dtype`:Ascend950的`quant_mode=3/5`仅支持`FLOAT8_E8M0`,其他量化模式支持`FP32`;Ascend910_93支持`FP16` +- `cmp_ratio > 1`且`sparse_mode != 0`时,`cmp_residual_k`必填(长度=batch_size的列表) +- `return_value=1`时,`output_idx_offset`需提供有效值 +- Ascend910_93要求`quant_mode=2` + +## 输出文件 + +| 文件 | 说明 | +|---|---| +| `result.xlsx` | 测试结果(精度、参数等) | +| `result_perf.xlsx` | 测试结果 + 性能数据(仅msprof模式) | +| `batch_summary.log` | 批量执行详细日志 | +| `batch_fail_list.log` | 失败用例清单 | +| `PROF_*/` | msprof性能原始数据目录 | diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/list_pt_from_excel.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/list_pt_from_excel.py new file mode 100644 index 000000000000..f5cec0030c04 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/list_pt_from_excel.py @@ -0,0 +1,63 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import argparse +import os +import sys + +import pandas as pd + + +def list_pt_from_excel(excel_path, sheet_name, pt_dir): + if not os.path.exists(excel_path): + print(f"ERROR: Excel file not found: {excel_path}", file=sys.stderr) + sys.exit(1) + + if not os.path.isdir(pt_dir): + print(f"ERROR: pt directory not found: {pt_dir}", file=sys.stderr) + sys.exit(1) + + df = pd.read_excel(excel_path, sheet_name=sheet_name) + if "Testcase_Name" not in df.columns: + print(f"ERROR: Column 'Testcase_Name' not found in sheet '{sheet_name}'", file=sys.stderr) + sys.exit(1) + + pt_files = [] + missing = [] + for case_name in df["Testcase_Name"]: + pt_path = os.path.join(pt_dir, f"{case_name}.pt") + if os.path.isfile(pt_path): + pt_files.append(pt_path) + else: + missing.append(case_name) + + if missing: + print(f"WARNING: {len(missing)} cases have no matching .pt file: {missing}", file=sys.stderr) + + if not pt_files: + print("ERROR: No matching .pt files found for any case in Excel", file=sys.stderr) + sys.exit(1) + + print(",".join(pt_files)) + + +def main(): + parser = argparse.ArgumentParser(description="Extract Testcase_Name from Excel and map to .pt files in pt_dir") + parser.add_argument("excel_path", type=str, help="Path to Excel file") + parser.add_argument("pt_dir", type=str, help="Directory containing .pt files") + parser.add_argument("--sheet", "-S", type=str, default="Sheet1", help="Sheet name (default: Sheet1)") + args = parser.parse_args() + + list_pt_from_excel(args.excel_path, args.sheet, args.pt_dir) + + +if __name__ == "__main__": + main() diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/quant_lightning_indexer_v2_pt_loadprocess.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/quant_lightning_indexer_v2_pt_loadprocess.py new file mode 100644 index 000000000000..839e13a6a9bc --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/quant_lightning_indexer_v2_pt_loadprocess.py @@ -0,0 +1,208 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + + +import torch +import torch_npu +from qliv2_parameter_normalization import normalize_qliv2_params + +QUANT_MODE_MXFP4 = 5 + + +def test_qliv2_process(filepath, device_id=0): + # 加载测试数据 + test_data = torch.load(filepath, map_location="cpu") + + params = normalize_qliv2_params(test_data["params"]) + cpu_result = test_data["cpu_result"] + topk_value = test_data["topk_value"] + cpu_topk_value = test_data["cpu_topk_value"] + print("执行用例:", filepath) + torch_npu.npu.set_device(device_id) + + quant_mode = test_data["quant_mode"] + if quant_mode == QUANT_MODE_MXFP4: + query = test_data["query"].view(torch.uint8).npu() + key = test_data["key"].view(torch.uint8).npu() + if "blockFusion" in test_data and test_data["blockFusion"] is not None: + blockFusion = test_data["blockFusion"].view(torch.uint8).npu() + elif params[10] == "FLOAT8_E4M3FN" or params[10] == torch.float8_e4m3fn: + query = test_data["query"].to(dtype=torch.float8_e4m3fn).npu() + key = test_data["key"].to(dtype=torch.float8_e4m3fn).npu() + if "blockFusion" in test_data and test_data["blockFusion"] is not None: + blockFusion = test_data["blockFusion"] + if blockFusion.dtype == torch.uint8: + blockFusion = blockFusion.view(torch.float8_e4m3fn) + else: + blockFusion = blockFusion.to(dtype=torch.float8_e4m3fn) + blockFusion = blockFusion.npu() + else: + query = test_data["query"].npu() + key = test_data["key"].npu() + if "blockFusion" in test_data and test_data["blockFusion"] is not None: + blockFusion = test_data["blockFusion"].npu() + + max_seqlen_q = params[19] + return_value = params[31] + weights = test_data["weights"].npu() + query_dequant_scale = test_data["query_dequant_scale"].npu() + key_dequant_scale = test_data["key_dequant_scale"].npu() + if "blockFusion" in test_data and test_data["blockFusion"] is not None: + block_num = params[9] + block_size = params[8] + head_dim = params[7] + k_head_num = params[6] + k_head_num = int(k_head_num) + head_dim = int(head_dim) + block_size = int(block_size) + block_num = int(block_num) + dequant_dtype_str = params[12] + if dequant_dtype_str == "FP16" or dequant_dtype_str == torch.float16: + dequant_dtype = torch.float16 + elif dequant_dtype_str == "FP32" or dequant_dtype_str == torch.float32: + dequant_dtype = torch.float32 + else: + dequant_dtype = torch.float16 + key = blockFusion[:, : block_size * k_head_num * head_dim].view(block_num, block_size, k_head_num, head_dim) + key_dequant_scale = ( + blockFusion[:, block_size * k_head_num * head_dim :] + .view(dequant_dtype) + .view(block_num, block_size, k_head_num) + ) + if test_data["seqused_q"] is not None: + seqused_q = test_data["seqused_q"].npu() + else: + seqused_q = None + if test_data["seqused_k"] is not None: + seqused_k = test_data["seqused_k"].npu() + else: + seqused_k = None + if test_data["output_idx_offset"] is not None: + output_idx_offset = test_data["output_idx_offset"].npu() + else: + output_idx_offset = None + if test_data["cu_seqlens_query"] is not None: + cu_seqlens_query = test_data["cu_seqlens_query"].npu() + else: + cu_seqlens_query = None + if test_data["cu_seqlens_key"] is not None: + cu_seqlens_key = test_data["cu_seqlens_key"].npu() + else: + cu_seqlens_key = None + if test_data["block_table"] is not None: + block_table = test_data["block_table"].npu() + else: + block_table = None + layout_query = test_data["layout_query"] + layout_key = test_data["layout_key"] + sparse_count = test_data["sparse_count"] + sparse_mode = test_data["sparse_mode"] + cmp_ratio = test_data["cmp_ratio"] + if test_data["cmp_residual_k_for_npu"] is not None: + cmp_residual_k_for_npu = test_data["cmp_residual_k_for_npu"].npu() + else: + cmp_residual_k_for_npu = None + + max_seqlen_q_meta = test_data["max_seqlen_q_meta"] + max_seqlen_k_meta = test_data["max_seqlen_k_meta"] + metadata = torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata( + cu_seqlens_q=cu_seqlens_query, + cu_seqlens_k=cu_seqlens_key, + seqused_q=seqused_q, + seqused_k=seqused_k, + cmp_residual_k=cmp_residual_k_for_npu, + batch_size=params[0], + max_seqlen_q=max_seqlen_q_meta, + max_seqlen_k=max_seqlen_k_meta, + num_heads_q=params[5], + num_heads_k=params[6], + head_dim=params[7], + topk=sparse_count, + quant_mode=quant_mode, + mask_mode=sparse_mode, + layout_q=layout_query, + layout_k=layout_key, + cmp_ratio=cmp_ratio, + ) + metadata = metadata.npu() + + # 调用qli算子 + npu_result, npu_value = torch.ops.cann_ops_transformer.quant_lightning_indexer( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + cu_seqlens_q=cu_seqlens_query, + cu_seqlens_k=cu_seqlens_key, + seqused_q=seqused_q, + seqused_k=seqused_k, + cmp_residual_k=cmp_residual_k_for_npu, + output_idx_offset=output_idx_offset, + max_seqlen_q=max_seqlen_q, + block_table=block_table, + metadata=metadata, + quant_mode=quant_mode, + layout_q=layout_query, + layout_k=layout_key, + topk=sparse_count, + mask_mode=sparse_mode, + cmp_ratio=cmp_ratio, + return_value=return_value, + ) + + torch.npu.synchronize() + npu_topk_value = npu_value + if return_value: + if npu_topk_value.shape != npu_result.shape: + raise RuntimeError( + "sparse_values and sparse_indices must have the same shape when return_value=1, " + f"but got {tuple(npu_topk_value.shape)} and {tuple(npu_result.shape)}" + ) + npu_topk_value, npu_sort_order = npu_topk_value.sort(dim=-1, descending=True) + npu_result = torch.gather(npu_result, dim=-1, index=npu_sort_order) + return ( + cpu_result, + npu_result, + topk_value, + cpu_topk_value, + npu_topk_value, + output_idx_offset, + params, + ) + + +def test_qliv2_process_graph(filepath, device_id=0): + """ + graph 模式:从 .pt 文件加载 pre-computed tensor,走 torch.compile + torchair 后端执行算子, + 跳过 generate_qliv2_test_data 的随机数据重新生成和 CPU golden 重算。 + 与 eager 模式共用相同的 .pt 数据,仅算子调用路径不同(compile vs eager)。 + """ + import quant_lightning_indexer_v2_acl_graph + + test_data = torch.load(filepath, map_location="cpu") + params = normalize_qliv2_params(test_data["params"]) + output_idx_offset = test_data.get("output_idx_offset", None) + + torch_npu.npu.set_device(device_id) + cpu_result, npu_result, topk_value, cpu_topk_value, npu_topk_value = ( + quant_lightning_indexer_v2_acl_graph.qliv2_output_acl_graph_from_pt(params, test_data) + ) + + return ( + cpu_result, + npu_result, + topk_value, + cpu_topk_value, + npu_topk_value, + output_idx_offset, + params, + ) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/quant_lightning_indexer_v2_pt_save.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/quant_lightning_indexer_v2_pt_save.py new file mode 100644 index 000000000000..8a4ae75e5655 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/quant_lightning_indexer_v2_pt_save.py @@ -0,0 +1,184 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import os + +import numpy as np +import pandas as pd +import torch +from quant_lightning_indexer_v2_golden import generate_qliv2_test_data + +try: + import torch_npu +except ImportError: + torch_npu = None +import argparse + +import pytest + + +def load_excel_test_cases(excel_file_path: str, sheetname: str): + """ + 从 Excel 文件加载测试用例。 + + 参数: + excel_file_path (str): Excel 文件的路径。 + sheetname (str, optional): 工作表名称。若未提供,则默认 'Sheet1'。 + + 返回: + list[tuple]: 测试用例元组列表,每个元组包含 20+ 个字段。 + 若失败或跳过,则返回空列表。 + """ + # 优先使用传入的 sheetname,否则尝试从环境变量获取 + if sheetname is None: + sheetname = "Sheet1" + + # 检查文件是否存在 + if not os.path.exists(excel_file_path): + pytest.skip(f"Excel file not found: {excel_file_path}", allow_module_level=True) + + try: + # 读取 Excel 文件的指定 sheet + df = pd.read_excel(excel_file_path, sheet_name=sheetname) + df = df.replace({np.nan: None, pd.NA: None}) + + # 定义必需的列名 + required_columns = [ + "Testcase_Name", + "batch_size", + "q_seq", + "k_seq", + "q_t_size", + "k_t_size", + "q_head_num", + "k_head_num", + "head_dim", + "block_size", + "block_num", + "qk_dtype", + "dequant_dtype", + "actual_seq_dtype", + "cu_seqlens_q", + "cu_seqlens_k", + "seqused_q", + "seqused_k", + "cmp_residual_k", + "max_seqlen_q", + "quant_mode", + "layout_query", + "layout_key", + "sparse_count", + "sparse_mode", + "query_datarange", + "key_datarange", + "weights_datarange", + "q_scale_datarange", + "k_scale_datarange", + "cmp_ratio", + "return_value", + "output_idx_offset", + ] + + # 检查是否缺少必要列 + missing_cols = [col for col in required_columns if col not in df.columns] + if missing_cols: + pytest.skip( + f"Missing required columns in Excel: {missing_cols}", + allow_module_level=True, + ) + + # 构建测试用例列表 + test_cases = [] + for _, row in df.iterrows(): + test_cases.append( + ( + row["Testcase_Name"], + row["batch_size"], + row["q_seq"], + row["k_seq"], + row["q_t_size"], + row["k_t_size"], + row["q_head_num"], + row["k_head_num"], + row["head_dim"], + row["block_size"], + row["block_num"], + row["qk_dtype"], + row["weight_dtype"] + if "weight_dtype" in row and row["weight_dtype"] is not None + else row["dequant_dtype"], + row["dequant_dtype"], + row["actual_seq_dtype"], + row["cu_seqlens_q"], + row["cu_seqlens_k"], + row["seqused_q"], + row["seqused_k"], + row["cmp_residual_k"], + row["max_seqlen_q"], + row["quant_mode"], + row["layout_query"], + row["layout_key"], + row["sparse_count"], + row["sparse_mode"], + row["query_datarange"], + row["key_datarange"], + row["weights_datarange"], + row["q_scale_datarange"], + row["k_scale_datarange"], + row["cmp_ratio"], + row["return_value"], + row["output_idx_offset"], + ) + ) + + return test_cases + + except Exception as e: + pytest.skip(f"Failed to read Excel file: {e}", allow_module_level=True) + return None + + +def save_test_case(test_cases, file_path): + print("正在保存pt文件...") + # 创建输出目录 + os.makedirs(file_path, exist_ok=True) + + for idx, case in enumerate(test_cases): + try: + case_name = case[0] + output_tensors = generate_qliv2_test_data(case[1:]) + # 生成文件名 + input_filename = f"{case_name}.pt" + input_filepath = os.path.join(file_path, input_filename) + + # 保存数据 + torch.save(output_tensors, input_filepath) + print(f"测试用例已保存到: {input_filepath}") + + except Exception as e: + print(f"[失败] 生成 pt 文件失败: {case[0]} (索引: {idx})") + print(f"错误详情: {e}") + + +def main(): + parser = argparse.ArgumentParser(description="qliv2_pt_save.py 接收路径参数") + parser.add_argument("path1", type=str, help="第一个路径") + parser.add_argument("path2", type=str, help="第二个路径") + parser.add_argument("--sheet", "-S", type=str, default="Sheet1", help="Sheet 页名(默认: Sheet1)") + args = parser.parse_args() + path1 = args.path1 + path2 = args.path2 + testcase = load_excel_test_cases(path1, args.sheet) + save_test_case(testcase, path2) + + +if __name__ == "__main__": + main() diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/replace_path.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/replace_path.py new file mode 100644 index 000000000000..bf5ddf640fe2 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch/replace_path.py @@ -0,0 +1,39 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import fileinput +import sys + + +def replace_paths_in_test_file(test_file_path, path): + """ + 替换test_quant_lightning_indexer_v2_batch.py中的占位符为实际路径 + :param test_file_path: test_quant_lightning_indexer_v2_batch.py的路径 + :param path1: 实际路径 + """ + try: + # 逐行替换占位符 + with fileinput.FileInput(test_file_path, inplace=True, backup=".bak") as f: + for line in f: + # 替换__PATH__为实际路径 + line = line.replace("__PATH__", path) + # 输出替换后的行(inplace=True会自动写回文件) + print(line, end="") + print(f" 已成功替换 {test_file_path} 中的路径") + except Exception as e: + print(f" 替换路径失败:{e}") + sys.exit(1) + + +if __name__ == "__main__": + test_file = sys.argv[1] + path = sys.argv[2] + replace_paths_in_test_file(test_file, path) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch_isolated_run.sh b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch_isolated_run.sh new file mode 100644 index 000000000000..689d5edd7488 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/batch_isolated_run.sh @@ -0,0 +1,142 @@ +#!/bin/bash +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2025 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +# ----------------------------------------------------------------------------------------------------------- + +# 批量隔离执行脚本: +# 1. 获取指定路径下的所有 .pt 用例 +# 2. 对每条用例单独拉起一个 pytest 进程执行, 进程间完全隔离 +# - 某条用例 device 越界/崩溃不会影响后续用例 +# - msprof 挂载在单用例进程上, 性能数据采集互不干扰 +# 用法: +# bash batch_isolated_run.sh [用例目录] [是否msprof采集: 0|1] [运行模式: eager|graph] +# 示例: +# bash batch_isolated_run.sh # 默认 pt_path 目录, 不采集性能, eager模式 +# bash batch_isolated_run.sh pt_path 1 # 启用 msprof 采集性能, eager模式 +# bash batch_isolated_run.sh pt_path 0 graph # graph模式, 不采集性能 +# ----------------------------------------------------------------------------------------------------------- + +set -o pipefail + +# Ctrl+C / SIGTERM 中断处理: 递归杀所有子进程后退出 +_cleanup_on_interrupt() { + echo -e "\n\n[中断] 收到终止信号,正在清理子进程..." | tee -a "$SUMMARY_LOG" + pkill -TERM -P $$ 2>/dev/null + sleep 2 + pkill -KILL -P $$ 2>/dev/null + exit 130 +} +trap _cleanup_on_interrupt SIGINT SIGTERM + +TEST_SCRIPT="test_quant_lightning_indexer_v2_batch.py" +TESTCASE_DIR="${1:-./pt_path}" +USE_MSPROF="${2:-0}" +RUN_MODE="${3:-eager}" +RESULT_XLSX="result.xlsx" +SUMMARY_LOG="batch_summary.log" +FAIL_LOG="batch_fail_list.log" + +# 清理旧文件 +[ -f "$RESULT_XLSX" ] && rm -f "$RESULT_XLSX" +[ -f "${RESULT_XLSX%.xlsx}_perf.xlsx" ] && rm -f "${RESULT_XLSX%.xlsx}_perf.xlsx" +rm -f "${RESULT_XLSX%.xlsx}_perf.xlsx.tmp.xlsx" +: > "$SUMMARY_LOG" +: > "$FAIL_LOG" + +# 清理旧的 PROF 文件夹, 避免与本次运行的数据混淆 +# 匹配 PROF_* 和 PROF_*_test_case_* 两种命名格式 +for _prof_dir in PROF_*/; do + [ -d "$_prof_dir" ] && rm -rf "$_prof_dir" +done + +# 1. 获取指定路径下的所有用例路径 +if [ ! -d "$TESTCASE_DIR" ]; then + echo "错误: 用例目录不存在: $TESTCASE_DIR" + exit 1 +fi + +mapfile -t CASE_FILES < <(find "$TESTCASE_DIR" -maxdepth 1 -name "*.pt" | sort) +TOTAL=${#CASE_FILES[@]} +if [ "$TOTAL" -eq 0 ]; then + echo "错误: 目录 $TESTCASE_DIR 下未找到 .pt 用例" + exit 1 +fi + +echo "共发现 $TOTAL 条用例, 目录: $TESTCASE_DIR , msprof采集: $USE_MSPROF , 运行模式: $RUN_MODE" +echo "开始隔离批量执行..." | tee -a "$SUMMARY_LOG" + +PASS=0 +FAIL=0 +FAIL_LIST=() + +# 2. 对每条用例单独调用一次测试脚本, 独立进程 +i=0 +for case_file in "${CASE_FILES[@]}"; do + i=$((i+1)) + case_name=$(basename "$case_file") + echo -e "\n===== [$i/$TOTAL] 执行用例: $case_name =====" | tee -a "$SUMMARY_LOG" + + if [ "$USE_MSPROF" = "1" ]; then + RUN_CMD="QLIV2_TESTCASE_PATH=\"${case_file}\" QLIV2_RUN_MODE=\"${RUN_MODE}\" msprof python3 -m pytest -rA -s ${TEST_SCRIPT} -v -m ci -W ignore::UserWarning -W ignore::DeprecationWarning" + else + RUN_CMD="QLIV2_TESTCASE_PATH=\"${case_file}\" QLIV2_RUN_MODE=\"${RUN_MODE}\" python3 -m pytest -rA -s ${TEST_SCRIPT} -v -m ci -W ignore::UserWarning -W ignore::DeprecationWarning" + fi + + eval "$RUN_CMD" 2>&1 | grep -v "^ninja: no work to do\.$" | tee -a "$SUMMARY_LOG" + status=${PIPESTATUS[0]} + + if [ "$status" -eq 0 ]; then + PASS=$((PASS+1)) + echo "[PASS] $case_name" | tee -a "$SUMMARY_LOG" + # 增量收集性能数据(每条用例跑完立即写入 result_perf.xlsx) + if [ "$USE_MSPROF" = "1" ]; then + sync + python3 collect_perf_data.py --incremental --test_result_path "$RESULT_XLSX" 2>&1 | tee -a "$SUMMARY_LOG" + # 重命名 PROF 文件夹,防止下一条用例的 msprof 覆盖 + _latest_prof=$(ls -dt PROF_*/ 2>/dev/null | head -1) + if [ -n "$_latest_prof" ]; then + _new_name="${_latest_prof%/}_${case_name%.pt}" + mv "$_latest_prof" "$_new_name" 2>/dev/null + fi + fi + else + FAIL=$((FAIL+1)) + FAIL_LIST+=("$case_name") + echo "[FAIL] $case_name" | tee -a "$SUMMARY_LOG" + echo "$case_name" >> "$FAIL_LOG" + fi +done + +# 3. 最终汇总(批量模式兜底,确保所有用例的性能数据都已收集) +if [ "$USE_MSPROF" = "1" ]; then + echo -e "\n========== 性能数据汇总校验 ==========" | tee -a "$SUMMARY_LOG" + python3 collect_perf_data.py --test_result_path "$RESULT_XLSX" 2>&1 | tee -a "$SUMMARY_LOG" +fi + +# 汇总 +echo -e "\n========== 批量执行汇总 ==========" | tee -a "$SUMMARY_LOG" +echo "总计: $TOTAL 通过: $PASS 失败: $FAIL" | tee -a "$SUMMARY_LOG" +if [ "$FAIL" -gt 0 ]; then + echo "失败用例:" | tee -a "$SUMMARY_LOG" + for f in "${FAIL_LIST[@]}"; do + echo " - $f" | tee -a "$SUMMARY_LOG" + done +fi +echo "详细日志: $SUMMARY_LOG" +echo "失败清单: $FAIL_LOG" +echo "结果表格: $RESULT_XLSX" +if [ "$USE_MSPROF" = "1" ]; then + echo "性能表格: ${RESULT_XLSX%.xlsx}_perf.xlsx" +fi + +if [ "$FAIL" -gt 0 ]; then + exit 1 +fi +exit 0 diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/collect_perf_data.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/collect_perf_data.py new file mode 100644 index 000000000000..27d79adc9753 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/collect_perf_data.py @@ -0,0 +1,253 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import argparse +import glob +import os + +import pandas as pd + + +def extract_qliv2_row(prof_folder): + prof_output = os.path.join(prof_folder, "mindstudio_profiler_output") + if not os.path.isdir(prof_output): + return None + + csv_files = glob.glob(os.path.join(prof_output, "op_summary*.csv")) + if not csv_files: + return None + + df = pd.read_csv(csv_files[0]) + target = df[df["Op Name"] == "QuantLightningIndexerV2"] + if target.empty: + return None + + row = target.iloc[0].to_dict() + return row + + +def collect_and_save_incremental(prof_folder, test_result_path): + """增量模式:读取单条用例的性能数据,追加到 result_perf.xlsx""" + if not os.path.exists(test_result_path): + print(f"结果文件不存在: {test_result_path}") + return None + + df_result = pd.read_excel(test_result_path) + if df_result.empty: + return None + + row_data = extract_qliv2_row(prof_folder) + if row_data is None: + return None + + last_idx = df_result.shape[0] - 1 + case_name = df_result.iloc[last_idx]["case_name"] + + perf_dict = {last_idx: row_data} + perf_df = pd.DataFrame.from_dict(perf_dict, orient="index") + + overlap = set(df_result.columns) & set(perf_df.columns) + if overlap: + rename_map = {col: f"op_{col}" for col in overlap} + perf_df = perf_df.rename(columns=rename_map) + + perf_path = test_result_path.replace(".xlsx", "_perf.xlsx") + tmp_path = perf_path + ".tmp.xlsx" + + if os.path.exists(perf_path): + df_perf = pd.read_excel(perf_path) + # 以 df_result 为基准重建,保留已有 perf 列,追加新行 + perf_cols = [c for c in df_perf.columns if c not in df_result.columns] + df_merged = df_result.copy() + for col in perf_cols: + df_merged[col] = None + for idx in range(min(len(df_result), len(df_perf))): + if col in df_perf.columns and idx < len(df_perf): + val = df_perf.at[idx, col] + if not (isinstance(val, float) and pd.isna(val)): + df_merged.at[idx, col] = val + for col in perf_df.columns: + if col not in df_merged.columns: + df_merged[col] = None + df_merged.loc[last_idx, col] = perf_df.loc[last_idx, col] + df_merged.to_excel(tmp_path, index=False) + os.replace(tmp_path, perf_path) + else: + df_result_with_perf = pd.concat([df_result, perf_df], axis=1) + df_result_with_perf.to_excel(tmp_path, index=False) + os.replace(tmp_path, perf_path) + + print(f" [perf] {case_name} Task Duration: {row_data['Task Duration(us)']}us -> {perf_path}") + return row_data + + +def collect_all(test_result_path, is_compare=False, perf_golden_path="perf_golden.xlsx"): + """批量模式:收集所有 PROF 文件夹的性能数据""" + if not os.path.exists(test_result_path): + print(f"结果文件不存在: {test_result_path}") + return + + df_b = pd.read_excel(test_result_path) + valid_mask = df_b["result"] != "NPU ERROR" + valid_count = valid_mask.sum() + + if valid_count == 0: + print("没有有效用例,跳过性能数据收集") + return + + prof_folders = sorted( + [d for d in os.listdir(".") if os.path.isdir(d) and d.startswith("PROF")], key=lambda x: os.path.getmtime(x) + ) + + print("============= 开始收集性能数据 =============") + print(f"有效用例数: {valid_count}, PROF文件夹数: {len(prof_folders)}") + + if len(prof_folders) == 0: + print("未找到PROF文件夹, 跳过性能数据收集") + return + + if len(prof_folders) != valid_count: + print(f"警告: PROF文件夹数量({len(prof_folders)})与有效用例数({valid_count})不一致") + + perf_rows = {} + prof_idx = 0 + + for i in range(df_b.shape[0]): + if not valid_mask.iloc[i]: + continue + + if prof_idx >= len(prof_folders): + print(f" [{i}] PROF文件夹不足, 跳过剩余用例") + break + + prof = prof_folders[prof_idx] + case_name = df_b.iloc[i]["case_name"] + row_data = extract_qliv2_row(prof) + + if row_data is not None: + perf_rows[i] = row_data + print(f" [{prof_idx + 1}] {case_name} -> {prof} (Task Duration: {row_data['Task Duration(us)']}us)") + else: + print(f" [{prof_idx + 1}] {case_name}: 未找到QuantLightningIndexer数据") + + prof_idx += 1 + + if not perf_rows: + print("未收集到任何性能数据") + return + + perf_df = pd.DataFrame.from_dict(perf_rows, orient="index") + overlap = set(df_b.columns) & set(perf_df.columns) + if overlap: + rename_map = {col: f"op_{col}" for col in overlap} + perf_df = perf_df.rename(columns=rename_map) + print(f"op_summary列名冲突已重命名: {list(rename_map.values())}") + + df_b = pd.concat([df_b, perf_df], axis=1) + + if is_compare: + try: + df_c = pd.read_excel(perf_golden_path) + except Exception as e: + print(f"读取基线数据失败: {e}") + is_compare = False + + if is_compare: + perf_threshold = 10 + perf_fail_list = [] + df_b["qliv2_perf_diff"] = "" + df_b["perf_result"] = "" + + for i in perf_rows: + try: + cur_sas = df_b.at[i, "Task Duration(us)"] + golden_sas = df_c.iloc[i]["Task Duration(us)"] + sas_diff = float(cur_sas) - float(golden_sas) + df_b.at[i, "qliv2_perf_diff"] = sas_diff + + if abs(sas_diff) > perf_threshold: + df_b.at[i, "perf_result"] = "Failed" + perf_fail_list.append(df_b.iloc[i]["case_name"]) + else: + df_b.at[i, "perf_result"] = "Pass" + except Exception as e: + print(f" 基线对比出错 (行{i}): {e}") + + if perf_fail_list: + print(f"性能不达标用例: {perf_fail_list}") + + new_path = test_result_path.replace(".xlsx", "_perf.xlsx") + tmp_path = new_path + ".tmp.xlsx" + if os.path.exists(new_path): + existing = pd.read_excel(new_path) + existing = existing.reindex(range(df_b.shape[0])) + for col in df_b.columns: + for idx in range(df_b.shape[0]): + if pd.isna(existing.at[idx, col]): + existing.at[idx, col] = df_b.at[idx, col] + op_cols = [c for c in existing.columns if c.startswith("op_") or c == "Task Duration(us)"] + filled_count = existing[op_cols].dropna(how="all").shape[0] if op_cols else 0 + if filled_count >= len(perf_rows): + print(f"增量模式已收集 {filled_count}/{len(perf_rows)} 条性能数据,跳过批量补采") + return + print(f"增量模式仅 {filled_count}/{len(perf_rows)} 条,批量补采覆盖") + for col in perf_df.columns: + if col in existing.columns: + for idx in perf_rows: + if pd.isna(existing.at[idx, col]): + existing.at[idx, col] = perf_df.at[idx, col] + else: + existing[col] = None + for idx in perf_rows: + existing.at[idx, col] = perf_df.at[idx, col] + existing.to_excel(tmp_path, index=False) + os.replace(tmp_path, new_path) + return + + df_b.to_excel(tmp_path, index=False) + os.replace(tmp_path, new_path) + + print(f"\n性能数据已保存: {new_path}") + print(f"共拼接 {len(perf_rows)} 条用例的 op_summary 全字段数据") + print("============= 性能数据收集完成 =============") + + +def main(): + parser = argparse.ArgumentParser(description="收集性能数据") + parser.add_argument("--test_result_path", type=str, default="result.xlsx") + parser.add_argument( + "--incremental", action="store_true", default=False, help="增量模式:只收集最新一条用例的性能数据" + ) + parser.add_argument("--prof_folder", type=str, default=None, help="增量模式下指定 PROF 文件夹路径") + parser.add_argument("--is_compare", action="store_true", default=False) + parser.add_argument("--perf_golden_path", type=str, default="perf_golden.xlsx") + args = parser.parse_args() + + if args.incremental: + if args.prof_folder: + prof = args.prof_folder + else: + prof_folders = sorted( + [d for d in os.listdir(".") if os.path.isdir(d) and d.startswith("PROF")], + key=lambda x: os.path.getmtime(x), + ) + if not prof_folders: + print("未找到PROF文件夹") + return + prof = prof_folders[-1] + + collect_and_save_incremental(prof, args.test_result_path) + else: + collect_all(args.test_result_path, args.is_compare, args.perf_golden_path) + + +if __name__ == "__main__": + main() diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/pytest.ini b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/pytest.ini new file mode 100644 index 000000000000..c880a3e3eb75 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/pytest.ini @@ -0,0 +1,4 @@ +[pytest] +markers = + ci: mark a test as a CI test + graph: marks tests as graph mode compilation tests diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/qliv2_parameter_normalization.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/qliv2_parameter_normalization.py new file mode 100644 index 000000000000..13205de33b3e --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/qliv2_parameter_normalization.py @@ -0,0 +1,138 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +"""Pure QLI_V2 parameter normalization shared by pytest and TTK adapters.""" + + +def normalize_qliv2_params(params): + """Apply the batch pytest scalar conversions without generating test data.""" + values = tuple(params) + if len(values) not in (32, 33): + raise ValueError(f"QLI_V2 parameter count mismatch: got {len(values)}, expected 32 or 33") + + has_weight_dtype = len(values) == 33 + if has_weight_dtype: + ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + weight_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + output_idx_offset, + ) = values + else: + ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + output_idx_offset, + ) = values + weight_dtype = dequant_dtype + + q_t_size = 0 if q_t_size is None else int(q_t_size) + k_t_size = 0 if k_t_size is None else int(k_t_size) + block_size = 0 if block_size is None else int(block_size) + block_num = 0 if block_num is None else int(block_num) + max_seqlen_q = -1 if max_seqlen_q is None else int(max_seqlen_q) + quant_mode = 1 if quant_mode is None else int(quant_mode) + if has_weight_dtype and weight_dtype is None: + weight_dtype = dequant_dtype + + normalized = ( + int(batch_size), + int(q_seq), + int(k_seq), + q_t_size, + k_t_size, + int(q_head_num), + int(k_head_num), + int(head_dim), + block_size, + block_num, + qk_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + int(sparse_count), + int(sparse_mode), + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + int(cmp_ratio), + int(return_value), + output_idx_offset, + ) + return normalized[:11] + (weight_dtype,) + normalized[11:] diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/qliv2_test_utils.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/qliv2_test_utils.py new file mode 100644 index 000000000000..18456cf4c9b1 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/qliv2_test_utils.py @@ -0,0 +1,207 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +"""NPU-independent case selection, naming, and result helpers for QLI_V2 tests.""" + +from pathlib import Path + +import pandas as pd +import regex as re + +PARAM_NAMES = ( + "batch_size", + "q_seq", + "k_seq", + "q_t_size", + "k_t_size", + "q_head_num", + "k_head_num", + "head_dim", + "block_size", + "block_num", + "qk_dtype", + "weight_dtype", + "dequant_dtype", + "actual_seq_dtype", + "cu_seq_q", + "cu_seq_k", + "act_seq_q", + "act_seq_k", + "cmp_residual_k", + "max_seqlen_q", + "quant_mode", + "layout_query", + "layout_key", + "sparse_count", + "sparse_mode", + "query_datarange", + "key_datarange", + "weights_datarange", + "q_scale_datarange", + "k_scale_datarange", + "cmp_ratio", + "return_value", + "output_idx_offset", +) + + +def ensure_comparison_passed( + case_name, + result, + fulfill_percent, + result_return_value="N/A", + fulfill_percent_return_value=0, +): + """Raise a serializable error when an accuracy comparison fails.""" + failures = [] + if result != "Pass": + failures.append(f"index result={result}, fulfill_percent={fulfill_percent}") + if result_return_value not in ("N/A", "Pass"): + failures.append(f"value result={result_return_value}, fulfill_percent={fulfill_percent_return_value}") + if failures: + raise AssertionError(f"accuracy comparison failed for {case_name}: " + "; ".join(failures)) + + +class QliV2CaseSelector: + """Resolve an ordered subset of PT cases by explicit name or one-based index.""" + + @staticmethod + def natural_key(path): + return [int(part) if part.isdigit() else part.lower() for part in re.split(r"(\d+)", Path(path).name)] + + @staticmethod + def parse_indexes(expression, total): + if not expression: + return [] + indexes = [] + for token in str(expression).split(","): + token = token.strip() + if not token: + continue + if "-" in token: + start_text, end_text = token.split("-", 1) + start, end = int(start_text), int(end_text) + if end < start: + raise ValueError(f"invalid descending case index range: {token}") + indexes.extend(range(start, end + 1)) + else: + indexes.append(int(token)) + invalid = [index for index in indexes if index < 1 or index > total] + if invalid: + raise ValueError(f"case indexes out of range 1..{total}: {invalid}") + return indexes + + @classmethod + def resolve(cls, pt_dir, explicit_files="", case_names="", case_indexes=""): + if explicit_files: + candidates = [Path(item.strip()) for item in explicit_files.split(",") if item.strip()] + else: + directory = Path(pt_dir) + if not directory.is_dir(): + raise ValueError(f"PT directory does not exist: {directory}") + candidates = sorted(directory.glob("*.pt"), key=cls.natural_key) + + missing = [str(path) for path in candidates if not path.is_file()] + if missing: + raise ValueError(f"PT files do not exist: {missing}") + if not candidates: + raise ValueError(f"no PT cases found in: {pt_dir}") + if case_names and case_indexes: + raise ValueError("case names and case indexes cannot be specified together") + + if case_names: + by_name = {path.stem: path for path in candidates} + selected = [] + unknown = [] + for item in case_names.split(","): + name = Path(item.strip()).stem + if not name: + continue + if name not in by_name: + unknown.append(name) + else: + selected.append(by_name[name]) + if unknown: + raise ValueError(f"unknown case names: {unknown}") + candidates = selected + elif case_indexes: + indexes = cls.parse_indexes(case_indexes, len(candidates)) + candidates = [candidates[index - 1] for index in indexes] + + return [str(path) for path in candidates] + + +class QliV2ResultWriter: + """Build stable case names and append rows using the batch result schema.""" + + @staticmethod + def case_name(params, explicit_name=None): + if explicit_name: + normalized = re.sub(r"[^A-Za-z0-9_.-]+", "_", str(explicit_name)) + normalized = normalized.strip("._-") + if not normalized: + raise ValueError("case name has no usable filename characters") + return normalized + + values = list(params) + readable = ( + f"QLI_B{values[0]}_S1{values[1]}_S2{values[2]}_" + f"N1{values[5]}_N2{values[6]}_D{values[7]}_" + f"{values[21]}_{values[22]}_{values[10]}_" + f"QM{values[20]}_SM{values[24]}_CR{values[30]}_" + f"K{values[23]}_RV{values[31]}" + ) + return re.sub(r"[^A-Za-z0-9_.-]+", "_", readable) + + @staticmethod + def row( + case_name, + params, + result, + fulfill_percent, + result_return_value="N/A", + fulfill_percent_return_value=0, + ): + values = list(params) + if len(values) != len(PARAM_NAMES): + raise ValueError(f"QLI_V2 parameter count mismatch: got {len(values)}, expected {len(PARAM_NAMES)}") + row = {"case_name": case_name} + row.update(dict(zip(PARAM_NAMES, values))) + row.update( + { + "result": result, + "fulfill_percent": fulfill_percent, + "result_return_value": result_return_value, + "fulfill_percent_return_value": fulfill_percent_return_value, + } + ) + return row + + @staticmethod + def append(path, row): + output = Path(path) + output.parent.mkdir(parents=True, exist_ok=True) + if output.exists(): + frame = pd.read_excel(output) + expected_columns = list(row.keys()) + legacy_columns = [name for name in expected_columns if name != "return_value"] + if list(frame.columns) == legacy_columns: + frame["return_value"] = None + frame = frame[expected_columns] + elif list(frame.columns) != expected_columns: + raise ValueError( + "result columns do not match existing Excel: " + f"existing={list(frame.columns)}, current={list(row.keys())}" + ) + frame = pd.concat([frame, pd.DataFrame([row])], ignore_index=True) + else: + frame = pd.DataFrame([row]) + frame.to_excel(output, index=False) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/quant_lightning_indexer_v2_acl_graph.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/quant_lightning_indexer_v2_acl_graph.py new file mode 100644 index 000000000000..5f8184edcd23 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/quant_lightning_indexer_v2_acl_graph.py @@ -0,0 +1,248 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import quant_lightning_indexer_v2_golden +import torch +import torch.nn as nn +import torchair +from torchair.configs.compiler_config import CompilerConfig + +QUANT_MODE_MXFP4 = 5 + + +class QLIV2Network(nn.Module): + def __init__(self): + super().__init__() + + def forward( + self, + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + output_idx_offset, + max_seqlen_q, + block_table, + metadata, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + cmp_ratio, + return_value, + ): + return torch.ops.cann_ops_transformer.quant_lightning_indexer( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + seqused_q=seqused_q, + seqused_k=seqused_k, + cmp_residual_k=cmp_residual_k, + output_idx_offset=output_idx_offset, + max_seqlen_q=max_seqlen_q, + block_table=block_table, + metadata=metadata, + quant_mode=quant_mode, + layout_q=layout_query, + layout_k=layout_key, + topk=sparse_count, + mask_mode=sparse_mode, + cmp_ratio=cmp_ratio, + return_value=return_value, + ) + + +def _qliv2_prepare_tensors_and_metadata(params, tensor_dict): + """ + 统一处理 tensor 准备和 metadata 构造(共用逻辑)。 + 兼容两个来源:generate_qliv2_test_data 返回值和 .pt 文件加载,二者都在 CPU。 + """ + qk_dtype = params[10] + quant_mode = tensor_dict["quant_mode"] + + if quant_mode == QUANT_MODE_MXFP4: + # TorchAir通过foreach_copy搬运图输入,当前不支持FP4 shell dtype;使用相同存储的 + # packed uint8视图,C++入口会根据quant_mode恢复ACL_FLOAT4_E2M1语义。 + query = tensor_dict["query"].view(torch.uint8).npu() + key = tensor_dict["key"].view(torch.uint8).npu() + if "blockFusion" in tensor_dict and tensor_dict["blockFusion"] is not None: + blockFusion = tensor_dict["blockFusion"].view(torch.uint8).npu() + elif qk_dtype == "FLOAT8_E4M3FN" or qk_dtype == torch.float8_e4m3fn: + query = tensor_dict["query"].to(dtype=torch.float8_e4m3fn).npu() + key = tensor_dict["key"].to(dtype=torch.float8_e4m3fn).npu() + if "blockFusion" in tensor_dict and tensor_dict["blockFusion"] is not None: + blockFusion = tensor_dict["blockFusion"] + if blockFusion.dtype == torch.uint8: + blockFusion = blockFusion.view(torch.float8_e4m3fn) + else: + blockFusion = blockFusion.to(dtype=torch.float8_e4m3fn) + blockFusion = blockFusion.npu() + else: + query = tensor_dict["query"].npu() + key = tensor_dict["key"].npu() + if "blockFusion" in tensor_dict and tensor_dict["blockFusion"] is not None: + blockFusion = tensor_dict["blockFusion"].npu() + + weights = tensor_dict["weights"].npu() + query_dequant_scale = tensor_dict["query_dequant_scale"].npu() + key_dequant_scale = tensor_dict["key_dequant_scale"].npu() + + if "blockFusion" in tensor_dict and tensor_dict["blockFusion"] is not None: + block_num = int(params[9]) + block_size = int(params[8]) + head_dim = int(params[7]) + k_head_num = int(params[6]) + dequant_dtype_str = params[12] + if dequant_dtype_str == "FP16" or dequant_dtype_str == torch.float16: + dequant_dtype = torch.float16 + elif dequant_dtype_str == "FP32" or dequant_dtype_str == torch.float32: + dequant_dtype = torch.float32 + else: + dequant_dtype = torch.float16 + key = blockFusion[:, : block_size * k_head_num * head_dim].view(block_num, block_size, k_head_num, head_dim) + key_dequant_scale = ( + blockFusion[:, block_size * k_head_num * head_dim :] + .view(dequant_dtype) + .view(block_num, block_size, k_head_num) + ) + + cu_seqlens_query = tensor_dict["cu_seqlens_query"].npu() if tensor_dict["cu_seqlens_query"] is not None else None + cu_seqlens_key = tensor_dict["cu_seqlens_key"].npu() if tensor_dict["cu_seqlens_key"] is not None else None + seqused_q = tensor_dict["seqused_q"].npu() if tensor_dict["seqused_q"] is not None else None + seqused_k = tensor_dict["seqused_k"].npu() if tensor_dict["seqused_k"] is not None else None + output_idx_offset = tensor_dict["output_idx_offset"].npu() if tensor_dict["output_idx_offset"] is not None else None + block_table = tensor_dict["block_table"].npu() if tensor_dict["block_table"] is not None else None + cmp_residual_k_for_npu = ( + tensor_dict["cmp_residual_k_for_npu"].npu() if tensor_dict.get("cmp_residual_k_for_npu") is not None else None + ) + + layout_query = tensor_dict["layout_query"] + layout_key = tensor_dict["layout_key"] + sparse_count = tensor_dict["sparse_count"] + sparse_mode = tensor_dict["sparse_mode"] + cmp_ratio = tensor_dict["cmp_ratio"] + max_seqlen_q_meta = tensor_dict["max_seqlen_q_meta"] + max_seqlen_k_meta = tensor_dict["max_seqlen_k_meta"] + + q_head_num = int(params[5]) + k_head_num = int(params[6]) + head_dim = int(params[7]) + batch_size = int(params[0]) + + metadata = torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata( + cu_seqlens_q=cu_seqlens_query, + cu_seqlens_k=cu_seqlens_key, + seqused_q=seqused_q, + seqused_k=seqused_k, + cmp_residual_k=cmp_residual_k_for_npu, + batch_size=batch_size, + max_seqlen_q=max_seqlen_q_meta, + max_seqlen_k=max_seqlen_k_meta, + num_heads_q=q_head_num, + num_heads_k=k_head_num, + head_dim=head_dim, + topk=sparse_count, + quant_mode=quant_mode, + mask_mode=sparse_mode, + layout_q=layout_query, + layout_k=layout_key, + cmp_ratio=cmp_ratio, + ) + metadata = metadata.npu() + + run_args = { + "query": query, + "key": key, + "weights": weights, + "query_dequant_scale": query_dequant_scale, + "key_dequant_scale": key_dequant_scale, + "cu_seqlens_q": cu_seqlens_query, + "cu_seqlens_k": cu_seqlens_key, + "seqused_q": seqused_q, + "seqused_k": seqused_k, + "cmp_residual_k": cmp_residual_k_for_npu, + "output_idx_offset": output_idx_offset, + "max_seqlen_q": params[19], + "block_table": block_table, + "metadata": metadata, + "quant_mode": quant_mode, + "layout_query": layout_query, + "layout_key": layout_key, + "sparse_count": sparse_count, + "sparse_mode": sparse_mode, + "cmp_ratio": cmp_ratio, + "return_value": params[31], + } + return run_args + + +def _qliv2_run_compiled_graph(run_args): + """ + 通过 torch.compile + torchair 后端执行算子(共用逻辑)。 + """ + config = CompilerConfig() + config.mode = "reduce-overhead" + npu_backend = torchair.get_npu_backend(compiler_config=config) + torch._dynamo.reset() + npu_mode = torch.compile(QLIV2Network().npu(), fullgraph=False, backend=npu_backend, dynamic=False) + npu_result, npu_topk_value = npu_mode(**run_args) + torch.npu.synchronize() + if run_args["return_value"]: + if npu_topk_value.shape != npu_result.shape: + raise RuntimeError( + "sparse_values and sparse_indices must have the same shape when return_value=1, " + f"but got {tuple(npu_topk_value.shape)} and {tuple(npu_result.shape)}" + ) + npu_topk_value, npu_sort_order = npu_topk_value.sort(dim=-1, descending=True) + npu_result = torch.gather(npu_result, dim=-1, index=npu_sort_order) + return npu_result, npu_topk_value + + +def qliv2_output_acl_graph(params): + """ + graph 模式入口(single 用例使用):即时生成随机 tensor + CPU golden,再走 torch.compile 执行。 + """ + print("acl_graph") + tensor_dict = quant_lightning_indexer_v2_golden.generate_qliv2_test_data(params) + cpu_result = tensor_dict["cpu_result"] + topk_value = tensor_dict["topk_value"] + cpu_topk_value = tensor_dict["cpu_topk_value"] + + run_args = _qliv2_prepare_tensors_and_metadata(params, tensor_dict) + npu_result, npu_topk_value = _qliv2_run_compiled_graph(run_args) + + return cpu_result, npu_result, topk_value, cpu_topk_value, npu_topk_value + + +def qliv2_output_acl_graph_from_pt(params, tensor_dict): + """ + graph 模式入口(batch 用例使用):从 .pt 文件加载 pre-computed tensor,走 torch.compile 执行。 + 跳过 generate_qliv2_test_data 的随机数据重新生成和 CPU golden 重算。 + """ + cpu_result = tensor_dict["cpu_result"] + topk_value = tensor_dict["topk_value"] + cpu_topk_value = tensor_dict["cpu_topk_value"] + + run_args = _qliv2_prepare_tensors_and_metadata(params, tensor_dict) + npu_result, npu_topk_value = _qliv2_run_compiled_graph(run_args) + + return cpu_result, npu_result, topk_value, cpu_topk_value, npu_topk_value diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/quant_lightning_indexer_v2_golden.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/quant_lightning_indexer_v2_golden.py new file mode 100644 index 000000000000..388606a75da2 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/quant_lightning_indexer_v2_golden.py @@ -0,0 +1,2861 @@ +#!/usr/bin/python +# ruff: noqa +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import torch + +try: + import torch_npu +except ImportError: + torch_npu = None +import ast +import copy +import math +import random + +import numpy as np +from qliv2_parameter_normalization import normalize_qliv2_params + +try: + import cann_ops_transformer +except ImportError: + cann_ops_transformer = None + +DISCONTINUOUS_KEYS = True # key非连续 +DEFAULT_SPLIT_S1 = False # golden切分S1Flag +DEFAULT_S1SIZE = 2048 # s1切分基本块大小 + +FP32_FRACTION_BITS = 23 # fp32尾数位数 + +HIF8_EXP_ZERO_THRESHOLD = -23 # 边界值 +HIF8_EXP_DML_MIN = -22 # DML最小指数 +HIF8_EXP_DML_MAX = -15 # DML最大指数 +HIF8_EXP_D0 = 0 # D0指数值 +HIF8_EXP_D1_BOUNDARY = 1 # D1指数值 +HIF8_EXP_D2_MIN, HIF8_EXP_D2_MAX = 2, 3 # D2指数范围 +HIF8_EXP_D3_MIN, HIF8_EXP_D3_MAX = 4, 7 # D3指数范围 +HIF8_EXP_D4_MIN, HIF8_EXP_D4_MAX = 8, 15 # D4指数范围 + +HIF8_DOT_DML = 0 # DML: Denormal Low, 指数范围 -22 ~ -16, 0位尾数 +HIF8_DOT_D0 = 1 # D0: 指数为0,3位尾数(最高精度) +HIF8_DOT_D1 = 2 # D1: 指数为±1,3位尾数 +HIF8_DOT_D2 = 4 # D2: 指数为±2 ~ ±3,3位尾数 +HIF8_DOT_D3 = 8 # D3: 指数为±4 ~ ±7,2位尾数 +HIF8_DOT_D4 = 12 # D4: 指数为±8 ~ ±15,1位尾数(最低精度) +HIF8_DOT_INVALID = -1 # 无效状态 + +HIF8_FRAC_BITS_DML = 0 # DML档位尾数位数 +HIF8_FRAC_BITS_D0 = 3 # D0档位尾数位数 +HIF8_FRAC_BITS_D1 = 3 # D1档位尾数位数 +HIF8_FRAC_BITS_D2 = 3 # D2档位尾数位数 +HIF8_FRAC_BITS_D3 = 2 # D3档位尾数位数 +HIF8_FRAC_BITS_D4 = 1 # D4档位尾数位数 + +HIF8_EXP_BITS_DML = 3 # DML档位指数位数 +HIF8_EXP_BITS_D0 = 0 # D0档位指数位数 +HIF8_EXP_BITS_D1 = 1 # D1档位指数位数 +HIF8_EXP_BITS_D2 = 2 # D2档位指数位数 +HIF8_EXP_BITS_D3 = 3 # D3档位指数位数 +HIF8_EXP_BITS_D4 = 4 # D4档位指数位数 + +HIF8_ZERO = 0 +HIF8_NAN = 128 # 0b10000000, NaN +HIF8_NEG_INF = 239 # 0b11101111, -inf +HIF8_NEG_MAX = 238 # 0b11101110, 负极大值 +HIF8_POS_INF = 111 # 0b01101111, +inf +HIF8_POS_MAX = 110 # 0b01101110, 正极大值 + +HIF8_SIGN_MASK = 128 # 0b10000000, 符号位掩码 +HIF8_DOT_MASK = 120 # 0b01110000, dot值掩码 +HIF8_FRAC_MASK_3BIT = 7 # 0b00000111, 3位尾数掩码(D0/D1/D2) +HIF8_FRAC_MASK_2BIT = 3 # 0b00000011, 2位尾数掩码(D3) +HIF8_FRAC_MASK_1BIT = 1 # 0b00000001, 1位尾数掩码(D4) +HIF8_EXP_MASK_DML = 7 # 0b00000111, DML指数掩码(bit0-2) +HIF8_EXP_MASK_D4 = 30 # 0b00011110, D4指数掩码(bit1-4) +HIF8_EXP_MASK_D3 = 28 # 0b00011100, D3指数掩码(bit2-4) +HIF8_EXP_MASK_D2 = 24 # 0b00011000, D2指数掩码(bit3-4) +HIF8_EXP_SIGN_MASK_D1 = 8 # 0b00001000, D1指数掩码(bit3) + +HIF8_DOT_BIT_SHIFT = 3 # Dot值在HiF8中的起始位置(bit3) +HIF8_DML_EXP_OFFSET = 23 # DML指数偏移值 +HIF8_OVERFLOW_SCALE = 1.25 # 溢出阈值缩放因子 +HIF8_MAX_FINITE_VALUE = 32768 # 最大有限值(非饱和模式下的边界值, 2^15 + +SSR_T14_MASK = 16383 # 0b0011 1111 1111 1111, 14位低位掩码 +SSR_F14_OFFSET = 8192 # 0b0010 0000 0000 0000, F14偏移值 +SSR_DML_SHIFT = 10 # SSR舍入移位值 +SSR_RESERVED_BITS = 14 # SSR舍入保留位数 +HYBRID_ROUND_EXP_THRESHOLD = 4 # 混合舍入的指数分界点 + +QUANT_MODE_MXFP8 = 3 +QUANT_MODE_HIFLOAT8 = 4 +QUANT_MODE_MXFP4 = 5 +MX_SCALE_SHAPE_ALIGN = 64 +MX_SCALE_PACK_NUM = 2 +MX_SCALE_GROUP_SIZE = MX_SCALE_SHAPE_ALIGN // MX_SCALE_PACK_NUM +FP4_PACK_NUM = 2 +E8M0_ONE_VALUE = 127 +BF16_SIGNIFICAND_BITS = 8 +BF16_MIN_NORMAL = 2.0**-126 +BF16_MIN_SUBNORMAL = 2.0**-133 +FP4_E2M1_VALUES = torch.tensor( + [ + 0.0, + 0.5, + 1.0, + 1.5, + 2.0, + 3.0, + 4.0, + 6.0, + -0.0, + -0.5, + -1.0, + -1.5, + -2.0, + -3.0, + -4.0, + -6.0, + ], + dtype=torch.float32, +) +MXFP4_TORCH_DTYPE = torch.float4_e2m1fn_x2 + + +def is_mx_quant_mode(quant_mode): + return quant_mode in (QUANT_MODE_MXFP8, QUANT_MODE_MXFP4) + + +def is_mxfp4_quant_mode(quant_mode): + return quant_mode == QUANT_MODE_MXFP4 + + +def round_fp64_to_bf16_rne(value): + if value.dtype != torch.float64: + raise TypeError(f"round_fp64_to_bf16_rne expects float64 input, got {value.dtype}") + + mantissa, exponent = torch.frexp(value) + rounded_normal = torch.ldexp( + torch.round(mantissa * (1 << BF16_SIGNIFICAND_BITS)), + exponent - BF16_SIGNIFICAND_BITS, + ) + rounded_subnormal = torch.round(value / BF16_MIN_SUBNORMAL) * BF16_MIN_SUBNORMAL + rounded = torch.where(value.abs() < BF16_MIN_NORMAL, rounded_subnormal, rounded_normal) + rounded = torch.where(torch.isfinite(value), rounded, value) + + # The value is already on the BF16 grid, so both casts below are exact. + return rounded.to(torch.float32).to(torch.bfloat16) + + +def reduce_mxfp4_weighted_qk(weight_matrix, qk_matrix): + output_shape = (weight_matrix.shape[0], weight_matrix.shape[1], qk_matrix.shape[2]) + acc_bf16 = torch.zeros(output_shape, dtype=torch.bfloat16, device=weight_matrix.device) + + for g_idx in range(weight_matrix.shape[2]): + weight_fp64 = weight_matrix[:, :, g_idx : g_idx + 1].to(torch.float64) + qk_fp64 = qk_matrix[:, g_idx : g_idx + 1, :].to(torch.float64) + # Model a BF16-destination FMA with one BF16 rounding per G. + fma_fp64 = acc_bf16.to(torch.float64) + weight_fp64 * qk_fp64 + acc_bf16 = round_fp64_to_bf16_rne(fma_fp64) + return acc_bf16 + + +def reduce_mxfp8_weighted_qk(weight_matrix, qk_matrix): + output_shape = (weight_matrix.shape[0], weight_matrix.shape[1], qk_matrix.shape[2]) + acc_fp32 = torch.zeros(output_shape, dtype=torch.float32, device=weight_matrix.device) + + for g_idx in range(weight_matrix.shape[2]): + weight_fp64 = weight_matrix[:, :, g_idx : g_idx + 1].to(torch.float64) + qk_fp64 = qk_matrix[:, g_idx : g_idx + 1, :].to(torch.float64) + # Model an FP32-destination MulAddDst: compute the fused multiply-add in + # FP64, then round once to the FP32 destination after every G. + fma_fp64 = acc_fp32.to(torch.float64) + weight_fp64 * qk_fp64 + acc_fp32 = fma_fp64.to(torch.float32) + return acc_fp32 + + +def get_qk_physical_head_dim(head_dim, quant_mode): + return head_dim // FP4_PACK_NUM if is_mxfp4_quant_mode(quant_mode) else head_dim + + +def e8m0_raw_to_float(raw): + raw_float = raw.to(torch.float32) + scale = torch.pow(2.0, raw_float - E8M0_ONE_VALUE) + scale = torch.where(raw == 255, torch.full_like(scale, float("nan")), scale) + return scale + + +def validate_mx_scale_dtype(scale_dtype): + if scale_dtype != torch.float8_e8m0fnu: + raise TypeError(f"MXFP8/MXFP4 require dequant_dtype=torch.float8_e8m0fnu, but got {scale_dtype}") + + +def make_e8m0_tensor_from_raw(raw, scale_dtype, to_npu=True): + validate_mx_scale_dtype(scale_dtype) + scale = raw.contiguous().view(scale_dtype) + return scale.npu() if to_npu else scale + + +def make_mx_e8m0_scale_pair(base_shape, tail_shape, scale_range, scale_dtype, to_npu=True): + validate_mx_scale_dtype(scale_dtype) + range_min = float(scale_range[0]) + range_max = float(scale_range[1]) + if not math.isfinite(range_min) or not math.isfinite(range_max): + raise ValueError(f"E8M0 scale range must be finite, got {scale_range}") + if range_min <= 0 or range_min > range_max: + raise ValueError(f"E8M0 scale range must satisfy 0 < min <= max, got {scale_range}") + + # datarange表示真实scale值;在生成器内部转换为E8M0编码(编码e对应2^(e-127))。 + range_min_code = max(0, math.ceil(math.log2(range_min)) + E8M0_ONE_VALUE) + range_max_code = min(254, math.floor(math.log2(range_max)) + E8M0_ONE_VALUE) + while range_min_code <= 254 and math.ldexp(1.0, range_min_code - E8M0_ONE_VALUE) < range_min: + range_min_code += 1 + while range_max_code >= 0 and math.ldexp(1.0, range_max_code - E8M0_ONE_VALUE) > range_max: + range_max_code -= 1 + if range_min_code > range_max_code: + raise ValueError(f"E8M0 scale range contains no representable value: {scale_range}") + + # 沿连续存储顺序循环取值,保证不同token和不同D group使用可区分的scale。 + scale_shape = tuple(base_shape) + tuple(tail_shape) + raw_count = range_max_code - range_min_code + 1 + scale_raw = torch.arange(math.prod(scale_shape), dtype=torch.int64) + scale_raw = (scale_raw % raw_count + range_min_code).reshape(scale_shape).to(torch.uint8) + + cpu_scale = e8m0_raw_to_float(scale_raw) + return make_e8m0_tensor_from_raw(scale_raw, scale_dtype, to_npu), cpu_scale + + +def validate_mxfp4_dtype(data_dtype): + if data_dtype != MXFP4_TORCH_DTYPE: + raise TypeError(f"MXFP4 requires qk_dtype=torch.float4_e2m1fn_x2, but got {data_dtype}") + + +def make_mxfp4_tensor_pair(logical_shape, data_range, data_dtype, to_npu=True): + validate_mxfp4_dtype(data_dtype) + range_min = float(data_range[0]) + range_max = float(data_range[1]) + if not math.isfinite(range_min) or not math.isfinite(range_max) or range_min > range_max: + raise ValueError(f"invalid FP4 E2M1 numeric range: {data_range}") + if logical_shape[-1] % FP4_PACK_NUM != 0: + raise ValueError(f"MXFP4 head_dim must be even, got {logical_shape[-1]}") + + valid_codes = ( + torch.nonzero( + (range_min <= FP4_E2M1_VALUES) & (range_max >= FP4_E2M1_VALUES), + as_tuple=False, + ) + .flatten() + .to(torch.uint8) + ) + if valid_codes.numel() == 0: + raise ValueError(f"FP4 E2M1 range contains no representable value: {data_range}") + + raw_indices = torch.randint(valid_codes.numel(), tuple(logical_shape), dtype=torch.long) + raw = valid_codes[raw_indices] + raw_rows = raw.reshape(-1, logical_shape[-1]) + coverage_count = min(valid_codes.numel(), logical_shape[-1]) + coverage = torch.arange(coverage_count).unsqueeze(0) + torch.arange(raw_rows.shape[0]).unsqueeze(1) + raw_rows[:, :coverage_count] = valid_codes[coverage % valid_codes.numel()] + packed = raw[..., 0::2] | (raw[..., 1::2] << 4) + packed_mxfp4 = packed.contiguous().view(data_dtype) + cpu_ref = FP4_E2M1_VALUES[raw.to(torch.long)] + return (packed_mxfp4.npu() if to_npu else packed_mxfp4), cpu_ref + + +def make_e8m0_zero(shape, scale_dtype): + validate_mx_scale_dtype(scale_dtype) + return torch.zeros(shape, dtype=torch.uint8).view(scale_dtype) + + +class GeneralizedQLIV2: + def __init__( + self, + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + weight_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + act_seq_q, + act_seq_k, + cmp_residual_k, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + split_s1=DEFAULT_SPLIT_S1, + s1size=DEFAULT_S1SIZE, + ): + self.batch_size = batch_size + self.q_seq = q_seq + self.k_seq = k_seq + self.q_t_size = q_t_size + self.k_t_size = k_t_size + self.q_head_num = q_head_num + self.k_head_num = k_head_num + self.group_size = q_head_num // k_head_num + self.head_dim = head_dim + self.block_size = block_size + self.block_num = block_num + self.qk_dtype = qk_dtype + self.weight_dtype = weight_dtype + self.dequant_dtype = dequant_dtype + self.actual_seq_dtype = actual_seq_dtype + self.cu_seqlens_q = cu_seqlens_q + self.cu_seqlens_k = cu_seqlens_k + self.seqused_q = act_seq_q + self.seqused_k = act_seq_k + self.act_seq_q = act_seq_q + self.act_seq_k = act_seq_k + self.cmp_residual_k = cmp_residual_k + self.max_seqlen_q = max_seqlen_q + self.quant_mode = quant_mode + self.layout_query = layout_query + self.layout_key = layout_key + self.sparse_count = sparse_count + self.sparse_mode = sparse_mode + self.cmp_ratio = cmp_ratio + self.w_dtype = weight_dtype + self.return_value = return_value + self.split_s1 = split_s1 # 是否切分S1轴 / Whether to split the S1 axis + self.s1size = s1size # S1轴切分块大小 / S1 axis chunk size + + if layout_query == "BSND": + self.q_shape = [batch_size, q_seq, q_head_num, head_dim] + self.w_shape = [batch_size, q_seq, q_head_num] + self.q_tnd_flag = 0 + elif layout_query == "TND": + self.q_shape = [q_t_size, q_head_num, head_dim] + self.w_shape = [q_t_size, q_head_num] + self.q_tnd_flag = 1 + + if layout_key == "BSND": + self.k_shape = [batch_size, k_seq, k_head_num, head_dim] + elif layout_key == "TND": + self.k_shape = [k_t_size, k_head_num, head_dim] + + if layout_query == "BSND": + self.out_shape = [batch_size, q_seq, k_head_num, sparse_count] + self.output_idx_offset_shape = [batch_size, q_seq, k_head_num] + elif layout_query == "TND": + self.out_shape = [q_t_size, k_head_num, sparse_count] + self.output_idx_offset_shape = [q_t_size, k_head_num] + + def cal_atten_bnsd(self, output_idx_offset): + batch_size = self.batch_size + qs = self.q_seq + ks = self.k_seq + n1 = self.q_head_num + n2 = self.k_head_num + cu_seqlens_q = self.cu_seqlens_q + cu_seqlens_k = self.cu_seqlens_k + seqused_q = self.seqused_q + seqused_k = self.seqused_k + cmp_residual_k = self.cmp_residual_k + q_bnsd_tensor = self.q_bnsd_tensor + k_bnsd_tensor = self.k_bnsd_tensor + wt_bnsd_tensor = self.wt_bnsd_tensor + mask_tensor = self.m_tensor + q_scale_bnsd_tensor = self.q_scale_bnsd_tensor + k_scale_bnsd_tensor = self.k_scale_bnsd_tensor + cmp_ratio = self.cmp_ratio + + out_shape_bnsd = copy.deepcopy(self.q_bnsd_shape) + out_shape_bnsd[1] = n2 + out_shape_bnsd[-1] = self.sparse_count + + out_shape_bnss = copy.deepcopy(self.q_bnsd_shape) + out_shape_bnss[1] = n2 + # out_shape_bnss[-1] = math.floor(max(actualSeqLengths_k)) + out_shape_bnss[-1] = math.floor(max(seqused_k)) if seqused_k is not None else ks + + y = torch.full(out_shape_bnsd, -1, dtype=torch.int32) + y_value = torch.full(out_shape_bnss, -float("inf"), dtype=torch.float32) + y_value_np = np.full(out_shape_bnsd, -np.inf, dtype=np.float32) + + prefix = 0 + for b_idx in range(batch_size): + if self.layout_query == "TND": + if seqused_q is not None: + curr_actualSeq_q = seqused_q[b_idx] + else: + # 已被处理为shape为(B,)的tensor + curr_actualSeq_q = cu_seqlens_q[b_idx] + elif self.layout_query == "BSND": + if seqused_q is not None: + curr_actualSeq_q = seqused_q[b_idx] + else: + curr_actualSeq_q = qs + + if self.layout_key == "TND": + if seqused_k is not None: + curr_actualSeq_k = seqused_k[b_idx] + else: + curr_actualSeq_k = cu_seqlens_k[b_idx] + elif self.layout_key == "PA_BBND": + curr_actualSeq_k = seqused_k[b_idx] + elif self.layout_key == "BSND": + if seqused_k is not None: + curr_actualSeq_k = seqused_k[b_idx] + else: + curr_actualSeq_k = ks + self.cur_actseq_q = curr_actualSeq_q + self.cur_actseq_k = curr_actualSeq_k + + self.cur_b_idx = b_idx + + if self.split_s1: + # 切分S1轴以减小中间结果内存占用 + # Split S1 axis to reduce intermediate result memory usage + num_s1_chunks = math.ceil(curr_actualSeq_q / self.s1size) if curr_actualSeq_q > 0 else 1 + for s1_chunk_idx in range(num_s1_chunks): + s1_start = s1_chunk_idx * self.s1size + s1_end = min(s1_start + self.s1size, curr_actualSeq_q) + cur_chunk_s1 = s1_end - s1_start + + self.cur_q = q_bnsd_tensor[b_idx : (b_idx + 1), :, s1_start:s1_end, :] + self.cur_k = k_bnsd_tensor[b_idx : (b_idx + 1), :, :curr_actualSeq_k, :] + self.cur_wt = wt_bnsd_tensor[b_idx : (b_idx + 1), :, s1_start:s1_end, :] + self.cur_q_scale = q_scale_bnsd_tensor[b_idx : (b_idx + 1), :, s1_start:s1_end, :] + self.cur_k_scale = k_scale_bnsd_tensor[b_idx : (b_idx + 1), :, :curr_actualSeq_k] + if self.sparse_mode != 0: + self.cur_m = mask_tensor[b_idx : (b_idx + 1), s1_start:s1_end, :curr_actualSeq_k] + + if cur_chunk_s1 != 0: + actual_selected_count = min(curr_actualSeq_k, self.sparse_count) + if is_mx_quant_mode(self.quant_mode): + ( + y[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :actual_selected_count, + ], + y_value[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :curr_actualSeq_k, + ], + ) = self.cal_atten_per_batch_mx(b_idx) + elif self.qk_dtype == torch.int8: + ( + y[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :actual_selected_count, + ], + y_value[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :curr_actualSeq_k, + ], + ) = self.cal_atten_per_batch_int8(b_idx) + elif self.qk_dtype == torch.float8_e4m3fn: + ( + y[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :actual_selected_count, + ], + y_value[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :curr_actualSeq_k, + ], + ) = self.cal_atten_per_batch_fp8(b_idx) + elif self.qk_dtype == torch.uint8: + ( + y[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :actual_selected_count, + ], + y_value[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :curr_actualSeq_k, + ], + ) = self.cal_atten_per_batch_hifp8(b_idx) + if output_idx_offset is not None: + if self.layout_query == "TND": + offset = output_idx_offset.flatten()[prefix + s1_start : prefix + s1_end].reshape( + 1, -1, 1 + ) + else: + offset = output_idx_offset.flatten()[ + b_idx * qs + s1_start : b_idx * qs + s1_end + ].reshape(1, -1, 1) + offset_mask = ( + y[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :actual_selected_count, + ] + != -1 + ) + y[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :actual_selected_count, + ] += offset * offset_mask + y_value_np[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :actual_selected_count, + ] = -np.sort(-y_value.numpy())[ + b_idx : (b_idx + 1), + :, + s1_start:s1_end, + :actual_selected_count, + ] + y[ + b_idx : (b_idx + 1), + :, + curr_actualSeq_q:, + : min(curr_actualSeq_k, self.sparse_count), + ] = -1 + else: + self.cur_q = q_bnsd_tensor[b_idx : (b_idx + 1), :, :curr_actualSeq_q, :] + self.cur_k = k_bnsd_tensor[b_idx : (b_idx + 1), :, :curr_actualSeq_k, :] + self.cur_wt = wt_bnsd_tensor[b_idx : (b_idx + 1), :, :curr_actualSeq_q, :] + self.cur_q_scale = q_scale_bnsd_tensor[b_idx : (b_idx + 1), :, :curr_actualSeq_q, :] + self.cur_k_scale = k_scale_bnsd_tensor[b_idx : (b_idx + 1), :, :curr_actualSeq_k] + if self.sparse_mode != 0: + self.cur_m = mask_tensor[b_idx : (b_idx + 1), :curr_actualSeq_q, :curr_actualSeq_k] + + if curr_actualSeq_q != 0: + actual_selected_count = min(curr_actualSeq_k, self.sparse_count) + if is_mx_quant_mode(self.quant_mode): + ( + y[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :actual_selected_count, + ], + y_value[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :curr_actualSeq_k, + ], + ) = self.cal_atten_per_batch_mx(b_idx) + elif self.qk_dtype == torch.int8: + ( + y[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :actual_selected_count, + ], + y_value[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :curr_actualSeq_k, + ], + ) = self.cal_atten_per_batch_int8(b_idx) + elif self.qk_dtype == torch.float8_e4m3fn: + ( + y[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :actual_selected_count, + ], + y_value[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :curr_actualSeq_k, + ], + ) = self.cal_atten_per_batch_fp8(b_idx) + elif self.qk_dtype == torch.uint8: + ( + y[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :actual_selected_count, + ], + y_value[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :curr_actualSeq_k, + ], + ) = self.cal_atten_per_batch_hifp8(b_idx) + y[ + b_idx : (b_idx + 1), + :, + curr_actualSeq_q:, + :actual_selected_count, + ] = -1 + if output_idx_offset is not None: + if self.layout_query == "TND": + offset = output_idx_offset.flatten()[prefix : prefix + curr_actualSeq_q].reshape(1, -1, 1) + else: + offset = output_idx_offset.flatten()[b_idx * qs : b_idx * qs + curr_actualSeq_q].reshape( + 1, -1, 1 + ) + offset_mask = ( + y[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :actual_selected_count, + ] + != -1 + ) + y[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :actual_selected_count, + ] += offset * offset_mask + y_value_np[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :actual_selected_count, + ] = -np.sort(-y_value.numpy())[ + b_idx : (b_idx + 1), + :, + :curr_actualSeq_q, + :actual_selected_count, + ] + else: + pass + if self.layout_query == "TND": + prefix += cu_seqlens_q[b_idx] + return y, y_value, y_value_np + + def trans_shape_to_bnsd( + self, + tensor, + shape, + layout, + headnums=None, + act_seq=None, + is_weights=False, + tensor_name=None, + ): + if layout in ["BSND"]: + B = shape[0] + S = shape[1] + N = shape[2] + D = 1 + if is_weights: + tensor = torch.unsqueeze(tensor, dim=-1) + else: + D = shape[3] + tensor = tensor.reshape(B, S, N, D).permute(0, 2, 1, 3) + return tensor, [B, N, S, D] + elif layout == "BSN": + print("shape", shape) + B = shape[0] + S = shape[1] + N = shape[2] + if is_weights: + D = 1 + tensor = torch.unsqueeze(tensor, dim=-1) # 补D轴 + tensor = tensor.reshape(B, S, N, D).permute(0, 2, 1, 3) + return tensor, [B, N, S, D] + else: + tensor = tensor.reshape(B, S, N).permute(0, 2, 1) + return tensor, [B, N, S] + elif layout in ["TND"]: + T = shape[0] + N = shape[1] + D = 1 + if is_weights: + tensor = torch.unsqueeze(tensor, dim=-1) + else: + D = shape[2] + B = len(act_seq) + S = max(act_seq) + new_tensor = torch.zeros((B, N, S, D), dtype=tensor.dtype) + t_start = 0 + for b_index in range(B): + act_s = act_seq[b_index] + t_end = t_start + act_s + if act_s == 0: + continue + for n_index in range(N): + new_tensor[b_index, n_index, 0:act_s, :] = tensor[t_start:t_end, n_index, :] + t_start += act_s + return new_tensor, [B, N, S, D] + elif layout == "TN": + T = shape[0] + N = shape[1] + D = 1 + B = len(act_seq) + S = max(act_seq) + new_tensor = torch.zeros((B, N, S), dtype=tensor.dtype) + t_start = 0 + for b_index in range(B): + act_s = act_seq[b_index] + t_end = t_start + act_s + if act_s == 0: + continue + for n_index in range(N): + new_tensor[b_index, n_index, 0:act_s] = tensor[t_start:t_end, n_index] + t_start += act_s + return new_tensor, [B, N, S] + else: + return tensor, shape + + def trans_tnd_actseq(self, list): + list_len = len(list) + if list_len == 0: + raise ValueError("TND情况下 act_seq需要必传") + list_new = [] + list_new.append(list[0]) + for i in range(list_len - 1): + new_item = list[i + 1] - list[i] + if new_item >= 0: + list_new.append(new_item) + else: + raise ValueError(f"TND情况下 act_seq_len 为非递减数列 act_seq_len={list}") + return list_new + + def cal_atten_per_batch_mx(self, b_idx): + cur_q = self.cur_q.to(dtype=torch.float32) + cur_k = self.cur_k.to(dtype=torch.float32) + cur_wt = self.cur_wt.to(dtype=torch.float32) + cur_q_scale = self.cur_q_scale.to(dtype=torch.float32) + cur_k_scale = self.cur_k_scale.to(dtype=torch.float32) + sparse_count = self.sparse_count + sparse_mode = self.sparse_mode + if is_mxfp4_quant_mode(self.quant_mode): + head_dim = cur_q.shape[-1] + if head_dim % MX_SCALE_GROUP_SIZE != 0: + raise ValueError(f"MXFP4 head_dim must be divisible by {MX_SCALE_GROUP_SIZE}, but got {head_dim}") + + cur_q = cur_q * cur_q_scale + cur_k = cur_k * cur_k_scale + qk_bmm_res = torch.bmm(cur_q.squeeze(0), cur_k.permute(0, 1, 3, 2).squeeze(0)).unsqueeze(0) + + qk_relu_out = qk_bmm_res.clamp_min(0.0).to(torch.bfloat16) + weight_matrix = cur_wt.to(torch.bfloat16).permute(0, 2, 3, 1).squeeze(0) + qk_matrix = qk_relu_out.permute(0, 2, 1, 3).squeeze(0) + brc_vmul_matrix = reduce_mxfp4_weighted_qk(weight_matrix, qk_matrix) + brc_vmul = brc_vmul_matrix.unsqueeze(0) + else: + cur_q = cur_q * cur_q_scale + cur_k = cur_k * cur_k_scale + qk_bmm_res = torch.bmm(cur_q.squeeze(0), cur_k.permute(0, 1, 3, 2).squeeze(0)).unsqueeze(0) + qk_relu_out = qk_bmm_res.to(dtype=torch.float32).clamp_min(0.0) + weight_matrix = cur_wt.permute(0, 2, 3, 1).squeeze(0) + qk_matrix = qk_relu_out.permute(0, 2, 1, 3).squeeze(0) + brc_vmul = reduce_mxfp8_weighted_qk(weight_matrix, qk_matrix).unsqueeze(0) + temp_b, temp_s1, temp_n1, temp_s2 = brc_vmul.shape + temp_n2 = self.k_head_num + actual_selected_count = min(temp_s2, sparse_count) + reduce_sum = brc_vmul.reshape(temp_b, temp_n2, temp_s1, temp_s2) + + if sparse_mode == 3: + cur_m = self.cur_m + cur_m_broadcasted = cur_m.reshape(1, 1, temp_s1, temp_s2) + cur_m_broadcasted = torch.broadcast_to(cur_m_broadcasted, (1, temp_n2, temp_s1, temp_s2)) + reduce_sum[cur_m_broadcasted.to(dtype=torch.bool)] = -torch.inf + to_be_sort_ele = reduce_sum.clone().to(torch.bfloat16) + b_sorted_indices = torch.full(to_be_sort_ele.shape, -1, dtype=torch.int32) + if sparse_mode == 3: + for i in range(temp_s1): + row_mask = cur_m_broadcasted[0, 0, i, :].to(dtype=torch.bool) + true_indices = torch.where(~row_mask)[0] + row_ele = to_be_sort_ele[0, 0, i, true_indices] + indices = torch.arange(len(row_ele), device=row_ele.device) + sorted_vals, sorted_idx = torch.sort(torch.stack([-row_ele, indices], dim=1), dim=0, stable=True) + b_sorted_indices[0, 0, i, true_indices] = true_indices[sorted_idx[:, 0]].to(torch.int32) + else: + for i in range(temp_s1): + row_ele = to_be_sort_ele[0, 0, i, :] + indices = torch.arange(len(row_ele), device=row_ele.device) + sorted_vals, sorted_idx = torch.sort(torch.stack([-row_ele, indices], dim=1), dim=0, stable=True) + b_sorted_indices[0, 0, i, :] = sorted_idx[:, 0] + topk_indices = b_sorted_indices[..., :actual_selected_count] + return topk_indices, to_be_sort_ele + + def cal_atten_per_batch_hifp8(self, b_idx): + cur_q = self.cur_q + cur_k = self.cur_k + cur_wt = self.cur_wt.to(dtype=torch.float32) + cur_q_scale = self.cur_q_scale.to(dtype=torch.float32) + cur_k_scale = self.cur_k_scale.to(dtype=torch.float32) + sparse_count = self.sparse_count + sparse_mode = self.sparse_mode + cmp_ratio = self.cmp_ratio + cur_q = trans_hifuint8_tensor_to_float(cur_q) + cur_k = trans_hifuint8_tensor_to_float(cur_k) + qk_bmm_res = torch.bmm(cur_q.squeeze(0), cur_k.permute(0, 1, 3, 2).squeeze(0)).unsqueeze(0) + cur_w = cur_wt * cur_q_scale + qk_relu_out = (qk_bmm_res.to(dtype=torch.float32)).clamp_min(0.0) + brc_vmul = torch.bmm( + cur_w.permute(0, 2, 3, 1).to(dtype=torch.float32).squeeze(0), + qk_relu_out.permute(0, 2, 1, 3).to(dtype=torch.float32).squeeze(0), + ).unsqueeze(0) + temp_b, temp_s1, temp_n1, temp_s2 = brc_vmul.shape + temp_g = self.group_size + temp_n2 = self.k_head_num + temp_b_idx = self.cur_b_idx + actual_selected_count = min(temp_s2, sparse_count) + reduce_sum = brc_vmul.reshape(temp_b, temp_n2, temp_s1, temp_s2) + reduce_sum[0, :, :, :] = reduce_sum[0, :, :, :] * cur_k_scale + + if sparse_mode == 3: + cur_m = self.cur_m + cur_m_broadcasted = cur_m.reshape(1, 1, temp_s1, temp_s2) + cur_m_broadcasted = torch.broadcast_to(cur_m_broadcasted, (1, temp_n2, temp_s1, temp_s2)) + # 根据布尔矩阵置-inf + reduce_sum[cur_m_broadcasted.to(dtype=torch.bool)] = -torch.inf + to_be_sort_ele = reduce_sum.clone() + to_be_sort_ele = to_be_sort_ele.to(torch.bfloat16) + # 稳定排序 + b_sorted_indices = torch.full(to_be_sort_ele.shape, -1, dtype=torch.int32) + if sparse_mode == 3: + for i in range(temp_s1): + row_mask = cur_m_broadcasted[0, 0, i, :].to(dtype=torch.bool) + true_indices = torch.where(~row_mask)[0] + row_ele = to_be_sort_ele[0, 0, i, true_indices] + indices = torch.arange(len(row_ele), device=row_ele.device) + + sorted_vals, sorted_idx = torch.sort(torch.stack([-row_ele, indices], dim=1), dim=0, stable=True) + b_sorted_indices[0, 0, i, true_indices] = true_indices[sorted_idx[:, 0]].to(torch.int32) + else: + for i in range(temp_s1): + row_ele = to_be_sort_ele[0, 0, i, :] + indices = torch.arange(len(row_ele), device=row_ele.device) + sorted_vals, sorted_idx = torch.sort(torch.stack([-row_ele, indices], dim=1), dim=0, stable=True) + b_sorted_indices[0, 0, i, :] = sorted_idx[:, 0] + topk_indices = b_sorted_indices[..., :actual_selected_count] + return topk_indices, to_be_sort_ele + + def cal_atten_per_batch_fp8(self, b_idx): + cur_q = self.cur_q + cur_k = self.cur_k + cur_wt = self.cur_wt.to(dtype=torch.float32) + cur_q_scale = self.cur_q_scale.to(dtype=torch.float32) + cur_k_scale = self.cur_k_scale.to(dtype=torch.float32) + sparse_count = self.sparse_count + sparse_mode = self.sparse_mode + cmp_ratio = self.cmp_ratio + qk_bmm_res = torch.bmm( + cur_q.to(dtype=torch.float32).squeeze(0), + cur_k.to(dtype=torch.float32).permute(0, 1, 3, 2).squeeze(0), + ).unsqueeze(0) + cur_w = cur_wt * cur_q_scale + qk_relu_out = (qk_bmm_res.to(dtype=torch.float32)).clamp_min(0.0) + brc_vmul = torch.bmm( + cur_w.permute(0, 2, 3, 1).to(dtype=torch.float32).squeeze(0), + qk_relu_out.permute(0, 2, 1, 3).to(dtype=torch.float32).squeeze(0), + ).unsqueeze(0) + temp_b, temp_s1, temp_n1, temp_s2 = brc_vmul.shape + temp_g = self.group_size + temp_n2 = self.k_head_num + temp_b_idx = self.cur_b_idx + actual_selected_count = min(temp_s2, sparse_count) + reduce_sum = brc_vmul.reshape(temp_b, temp_n2, temp_s1, temp_s2) + reduce_sum[0, :, :, :] = reduce_sum[0, :, :, :] * cur_k_scale + + if sparse_mode == 3: + cur_m = self.cur_m + cur_m_broadcasted = cur_m.reshape(1, 1, temp_s1, temp_s2) + cur_m_broadcasted = torch.broadcast_to(cur_m_broadcasted, (1, temp_n2, temp_s1, temp_s2)) + # 根据布尔矩阵置-inf + reduce_sum[cur_m_broadcasted.to(dtype=torch.bool)] = -torch.inf + to_be_sort_ele = reduce_sum.clone() + to_be_sort_ele = to_be_sort_ele.to(torch.bfloat16) + # 稳定排序 + b_sorted_indices = torch.full(to_be_sort_ele.shape, -1, dtype=torch.int32) + if sparse_mode == 3: + for i in range(temp_s1): + row_mask = cur_m_broadcasted[0, 0, i, :].to(dtype=torch.bool) + true_indices = torch.where(~row_mask)[0] + row_ele = to_be_sort_ele[0, 0, i, true_indices] + indices = torch.arange(len(row_ele), device=row_ele.device) + + sorted_vals, sorted_idx = torch.sort(torch.stack([-row_ele, indices], dim=1), dim=0, stable=True) + b_sorted_indices[0, 0, i, true_indices] = true_indices[sorted_idx[:, 0]].to(torch.int32) + else: + for i in range(temp_s1): + row_ele = to_be_sort_ele[0, 0, i, :] + indices = torch.arange(len(row_ele), device=row_ele.device) + sorted_vals, sorted_idx = torch.sort(torch.stack([-row_ele, indices], dim=1), dim=0, stable=True) + b_sorted_indices[0, 0, i, :] = sorted_idx[:, 0] + topk_indices = b_sorted_indices[..., :actual_selected_count] + return topk_indices, to_be_sort_ele + + def cal_atten_per_batch_int8(self, b_idx): + cur_q = self.cur_q + cur_k = self.cur_k + cur_wt = self.cur_wt.to(dtype=torch.float16) + cur_q_scale = self.cur_q_scale.to(dtype=torch.float16) + cur_k_scale = self.cur_k_scale.to(dtype=torch.float16) + sparse_count = self.sparse_count + sparse_mode = self.sparse_mode + cmp_ratio = self.cmp_ratio + qk_bmm_res = torch.bmm( + cur_q.to(dtype=torch.int32).squeeze(0), + cur_k.to(dtype=torch.int32).permute(0, 1, 3, 2).squeeze(0), + ).unsqueeze(0) + cur_w = cur_wt * cur_q_scale + qk_relu_out = (qk_bmm_res.to(dtype=torch.float32) / 1024.0).clamp_min(0.0).to(torch.float16) + brc_vmul = torch.bmm( + cur_w.permute(0, 2, 3, 1).to(dtype=torch.float32).squeeze(0), + qk_relu_out.permute(0, 2, 1, 3).to(dtype=torch.float32).squeeze(0), + ).unsqueeze(0) + temp_b, temp_s1, temp_n1, temp_s2 = brc_vmul.shape + temp_g = self.group_size + temp_n2 = self.k_head_num + temp_b_idx = self.cur_b_idx + actual_selected_count = min(temp_s2, sparse_count) + reduce_sum = brc_vmul.reshape(temp_b, temp_n2, temp_s1, temp_s2) + reduce_sum[0, :, :, :] = reduce_sum[0, :, :, :] * cur_k_scale + + if sparse_mode == 3: + cur_m = self.cur_m + cur_m_broadcasted = cur_m.reshape(1, 1, temp_s1, temp_s2) + cur_m_broadcasted = torch.broadcast_to(cur_m_broadcasted, (1, temp_n2, temp_s1, temp_s2)) + # 根据布尔矩阵置-inf + reduce_sum[cur_m_broadcasted.to(dtype=torch.bool)] = -torch.inf + + to_be_sort_ele = reduce_sum.clone() + # 稳定排序 + b_sorted_indices = torch.full(to_be_sort_ele.shape, -1, dtype=torch.int32) + if sparse_mode == 3: + for i in range(temp_s1): + row_mask = cur_m_broadcasted[0, 0, i, :].to(dtype=torch.bool) + true_indices = torch.where(~row_mask)[0] + row_ele = to_be_sort_ele[0, 0, i, true_indices] + indices = torch.arange(len(row_ele), device=row_ele.device) + + sorted_vals, sorted_idx = torch.sort(torch.stack([-row_ele, indices], dim=1), dim=0, stable=True) + b_sorted_indices[0, 0, i, true_indices] = true_indices[sorted_idx[:, 0]].to(torch.int32) + else: + for i in range(temp_s1): + row_ele = to_be_sort_ele[0, 0, i, :] + indices = torch.arange(len(row_ele), device=row_ele.device) + sorted_vals, sorted_idx = torch.sort(torch.stack([-row_ele, indices], dim=1), dim=0, stable=True) + b_sorted_indices[0, 0, i, :] = sorted_idx[:, 0] + topk_indices = b_sorted_indices[..., :actual_selected_count] + return topk_indices, to_be_sort_ele + + def trans_bnsd_to_layout(self, tensor, shape, layout, act_q=None): + # 此时的输出D轴是K轴 + if layout == "BSH": + output = tensor.permute(0, 2, 1, 3).contiguous().view(shape) + return output + elif layout == "BSND": + output = tensor.permute(0, 2, 1, 3).contiguous() + return output + elif layout in ["BSND_NBSD", "BNSD_NBSD", "BSH_NBSD"]: + output = tensor.permute(1, 0, 2, 3).contiguous() + return output + elif layout in ["TND", "TND_NTD"]: + T = sum(act_q) + B = tensor.shape[0] + N = tensor.shape[1] + D = tensor.shape[3] + output = torch.full(size=(T, N, D), fill_value=-1, dtype=tensor.dtype) + t_start = 0 + for b_index in range(B): + act_s = act_q[b_index] + t_end = t_start + act_s + if act_s == 0: + continue + for n_index in range(N): + output[t_start:t_end, n_index, :] = tensor[b_index, n_index, :act_s, :] + t_start += act_s + if layout == "TND_NTD": + output = output.permute(1, 0, 2).contiguous() + return output + else: + return tensor + + def broadcast_n_axis(self, n1, n2, temp_tensor, input_dtype): + g = n1 // n2 + temp_shape = temp_tensor.shape + B = temp_shape[0] + S = temp_shape[2] + D = temp_shape[3] + modify_tensor = torch.zeros([B, n1, S, D], dtype=temp_tensor.dtype) + for i in range(n1): + j = i // g + modify_tensor[:, i : i + 1, :, :] = temp_tensor[:, j : j + 1, :, :] + return modify_tensor, modify_tensor.shape + + def flatten_mx_scale_tail(self, tensor, shape): + scale_head_dim = shape[-2] * shape[-1] + flat_shape = list(shape[:-2]) + [scale_head_dim] + return tensor.reshape(flat_shape), flat_shape + + def broadcast_mx_scale_d_axis(self, tensor): + output = tensor.repeat_interleave(MX_SCALE_GROUP_SIZE, dim=-1) + return output, list(output.shape) + + def create_mask(self, m_shape, act_k, S1): + atten_masks = torch.zeros(tuple(m_shape), dtype=torch.uint8) + cmp_ratio = self.cmp_ratio + tmp_pos_orig = act_k - S1 + + for i in range(S1): + if ((tmp_pos_orig + i + 1) / cmp_ratio) < 0: + atten_masks[i, :] = 1 + else: + atten_masks[i, math.floor((tmp_pos_orig + i + 1) / cmp_ratio) :] = 1 + return atten_masks + + def create_mask_right_down(self, m_shape, actualSeqLengthsQ, actualSeqLengthsK, batch): + mask_s_q = m_shape[0] + mask_s_kv = m_shape[1] + cmp_ratio = self.cmp_ratio + cmp_residual_k = self.cmp_residual_k + next_tokens_list = [] + re_mask_batch = [] + pre_tokens = 214748647 + for i in range(batch): + if len(actualSeqLengthsQ) == 0: + S1 = mask_s_q + else: + S1 = actualSeqLengthsQ[i] + + if len(actualSeqLengthsK) == 0: + S2 = mask_s_kv + else: + S2 = math.floor(actualSeqLengthsK[i]) + next_tokens = S2 - S1 + next_tokens_list.append(next_tokens) + act_k = actualSeqLengthsK[i] * cmp_ratio + cmp_residual_k[i] + atten_masks = self.create_mask(m_shape, act_k, S1) + re_mask_batch.append(np.array(atten_masks, dtype=np.bool_)) + re_mask_np = np.array(re_mask_batch, dtype=np.bool_) + cpu_mask = torch.from_numpy(re_mask_np) + return cpu_mask, next_tokens_list + + def forward( + self, + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + block_table, + output_idx_offset, + ): + print("cpu执行中...") + + # 参数的初始化 + batch_size = self.batch_size + q_seq = self.q_seq + k_seq = self.k_seq + layout_query = self.layout_query + layout_key = self.layout_key + sparse_count = self.sparse_count + sparse_mode = self.sparse_mode + out_shape = self.out_shape + q_shape = self.q_shape + head_dim = self.head_dim + q_head_num = self.q_head_num + k_head_num = self.k_head_num + q_t_size = self.q_t_size + k_t_size = self.k_t_size + block_size = self.block_size + block_num = self.block_num + q_dtype = self.qk_dtype + k_dtype = self.qk_dtype + w_shape = self.w_shape + w_dtype = self.w_dtype + actual_seq_dtype = self.actual_seq_dtype + cmp_ratio = self.cmp_ratio + return_value = self.return_value + is_mx_mode = is_mx_quant_mode(self.quant_mode) + + if layout_query == "TND": + q_scale_shape = [q_t_size, q_head_num] + self.cu_seqlens_q = self.trans_tnd_actseq(cu_seqlens_q[1:]) + actualSeqLengths_q = self.cu_seqlens_q + if seqused_q is not None: + self.seqused_q = seqused_q + self.has_seqused_q = True + actualSeqLengths_q = self.seqused_q + elif layout_query == "BSND": + q_scale_shape = [batch_size, q_seq, q_head_num] + if seqused_q is not None: + self.seqused_q = seqused_q + actual_seq_lengths_query = seqused_q + self.has_seqused_q = True + else: + actual_seq_lengths_query = torch.tensor(np.random.uniform(q_seq, q_seq, batch_size)).to(torch.int32) + actualSeqLengths_q = actual_seq_lengths_query + + if layout_key == "TND": + layout_key_scale = "TN" + k_scale_shape = [k_t_size, k_head_num] + self.cu_seqlens_k = self.trans_tnd_actseq(cu_seqlens_k[1:]) + actualSeqLengths_k = self.cu_seqlens_k + k_shape = self.k_shape + if seqused_k is not None: + self.seqused_k = seqused_k + self.has_seqused_k = True + actualSeqLengths_k = self.seqused_k + elif layout_key == "BSND": + layout_key_scale = "BSN" + k_shape = self.k_shape + k_scale_shape = [batch_size, k_seq, k_head_num] + if seqused_k is not None: + self.seqused_k = seqused_k + actual_seq_lengths_key = seqused_k + self.has_seqused_k = True + else: + actual_seq_lengths_key = torch.tensor(np.random.uniform(k_seq, k_seq, batch_size)).to(torch.int32) + actualSeqLengths_k = actual_seq_lengths_key + + elif layout_key == "PA_BBND": + self.actual_seq_lengths_key = seqused_k + actualSeqLengths_k = self.actual_seq_lengths_key + layout_key_scale = layout_key + k_max_s2 = math.floor(max(actualSeqLengths_k)) + k_shape = [batch_size, k_head_num, k_max_s2, head_dim] + k_scale_shape = [batch_size, k_head_num, k_max_s2] + if is_mx_mode: + mx_scale_tail_shape = [head_dim // MX_SCALE_SHAPE_ALIGN, MX_SCALE_PACK_NUM] + q_scale_shape = q_scale_shape + mx_scale_tail_shape + k_scale_shape = k_scale_shape + mx_scale_tail_shape + if layout_key in ["BSND", "TND"]: + layout_key_scale = layout_key + query = query.cpu() + key = key.cpu() + weights = weights.cpu() + query_dequant_scale = query_dequant_scale.cpu() + key_dequant_scale = key_dequant_scale.cpu() + q_scale_is_weights = True + if is_mx_mode: + query_dequant_scale, q_scale_shape = self.flatten_mx_scale_tail(query_dequant_scale, q_scale_shape) + key_dequant_scale, k_scale_shape = self.flatten_mx_scale_tail(key_dequant_scale, k_scale_shape) + q_scale_is_weights = False + if output_idx_offset is not None: + output_idx_offset = output_idx_offset.cpu() + + # 将输入转化为BNSD + ## BSND / TND -> BNSD + if self.layout_query == "TND": + q_bnsd_tensor, q_bnsd_shape = self.trans_shape_to_bnsd( + query, q_shape, layout_query, q_head_num, self.cu_seqlens_q + ) + else: + q_bnsd_tensor, q_bnsd_shape = self.trans_shape_to_bnsd( + query, q_shape, layout_query, q_head_num, actualSeqLengths_q + ) + + ## BSND/TND/ -> BNSD + if self.layout_key == "TND": + k_bnsd_tensor, k_bnsd_shape = self.trans_shape_to_bnsd( + key, k_shape, layout_key, k_head_num, self.cu_seqlens_k + ) + k_scale_bnsd_tensor, k_scale_bnsd_shape = self.trans_shape_to_bnsd( + key_dequant_scale, + k_scale_shape, + layout_key_scale, + k_head_num, + self.cu_seqlens_k, + ) + else: + k_bnsd_tensor, k_bnsd_shape = self.trans_shape_to_bnsd( + key, + k_shape, + layout_key, + k_head_num, + torch.floor(actualSeqLengths_k).to(actual_seq_dtype), + ) + k_scale_bnsd_tensor, k_scale_bnsd_shape = self.trans_shape_to_bnsd( + key_dequant_scale, + k_scale_shape, + layout_key_scale, + k_head_num, + torch.floor(actualSeqLengths_k).to(actual_seq_dtype), + ) + + ## BSN1 -> BNS1 TN1 -> BNS1 + is_weights = True + if self.layout_query == "TND": + wt_bnsd_tensor, wt_bnsd_shape = self.trans_shape_to_bnsd( + weights, + w_shape, + layout_query, + q_head_num, + self.cu_seqlens_q, + is_weights, + ) + q_scale_bnsd_tensor, q_scale_bnsd_shape = self.trans_shape_to_bnsd( + query_dequant_scale, + q_scale_shape, + layout_query, + q_head_num, + self.cu_seqlens_q, + q_scale_is_weights, + ) + else: + wt_bnsd_tensor, wt_bnsd_shape = self.trans_shape_to_bnsd( + weights, + w_shape, + layout_query, + q_head_num, + actualSeqLengths_q, + is_weights, + ) + # BSN1 -> BNS1 + q_scale_bnsd_tensor, q_scale_bnsd_shape = self.trans_shape_to_bnsd( + query_dequant_scale, + q_scale_shape, + layout_query, + q_head_num, + actualSeqLengths_q, + q_scale_is_weights, + ) + # 将 k n2轴 广播为 n1 + if q_head_num != k_head_num: + k_bnsd_tensor, k_bnsd_shape = self.broadcast_n_axis(q_head_num, k_head_num, k_bnsd_tensor, k_dtype) + if is_mx_mode: + k_scale_bnsd_tensor, k_scale_bnsd_shape = self.broadcast_n_axis( + q_head_num, + k_head_num, + k_scale_bnsd_tensor, + k_scale_bnsd_tensor.dtype, + ) + if is_mx_mode: + q_scale_bnsd_tensor, q_scale_bnsd_shape = self.broadcast_mx_scale_d_axis(q_scale_bnsd_tensor) + k_scale_bnsd_tensor, k_scale_bnsd_shape = self.broadcast_mx_scale_d_axis(k_scale_bnsd_tensor) + self.q_bnsd_tensor = q_bnsd_tensor + self.q_bnsd_shape = q_bnsd_shape + self.k_bnsd_tensor = k_bnsd_tensor + self.k_bnsd_shape = k_bnsd_shape + self.wt_bnsd_tensor = wt_bnsd_tensor + self.wt_bnsd_shape = wt_bnsd_shape + self.q_scale_bnsd_tensor = q_scale_bnsd_tensor + self.q_scale_bnsd_shape = q_scale_bnsd_shape + self.k_scale_bnsd_tensor = k_scale_bnsd_tensor + self.k_scale_bnsd_shape = k_scale_bnsd_shape + # 生成mask, sparse_mode=3时使能 + m_shape_std = [q_bnsd_shape[2], k_bnsd_shape[2]] # m_shape应该是[s1,s2] + batch = q_bnsd_shape[0] + m_tensor = [] + if sparse_mode == 3: + m_tensor, next_tokens_list = self.create_mask_right_down( + m_shape_std, actualSeqLengths_q, actualSeqLengths_k, batch + ) + elif sparse_mode == 0: + pass + else: + raise ValueError("unsupported sparse_mode!") + self.m_tensor = m_tensor + y, y_value, y_value_np = self.cal_atten_bnsd(output_idx_offset) + sparse_value = torch.from_numpy(y_value_np) + + # TND & PA 需要传入out_shape为BNSD + out_shape_bnsd = copy.deepcopy(self.q_bnsd_shape) + out_shape_bnsd[1] = k_head_num + out_shape_bnsd[-1] = sparse_count + if self.layout_query == "TND": + y = self.trans_bnsd_to_layout(y, out_shape_bnsd, layout_query, self.cu_seqlens_q) # TODO + if return_value: + sparse_value = self.trans_bnsd_to_layout(sparse_value, out_shape_bnsd, layout_query, self.cu_seqlens_q) + else: + y = self.trans_bnsd_to_layout(y, out_shape_bnsd, layout_query, q_seq) + if return_value: + sparse_value = self.trans_bnsd_to_layout(sparse_value, out_shape_bnsd, layout_query, q_seq) + return y, y_value, sparse_value + + +def trans_prefix_actseq(self, list): + list_len = len(list) + if list_len == 0: + raise ValueError("PA场景下 act_seq需要必传") + list_new = [] + list_new.append(list[0]) + for i in range(list_len - 1): + new_item = list[i + 1] - list[i] + if new_item >= 0: + list_new.append(new_item) + else: + raise ValueError(f"PA场景下act seq 为非递减数列 act_seq ={list}") + return list_new + + +def qliv2_output_single( + params, + is_batch=False, + split_s1=DEFAULT_SPLIT_S1, + s1size=DEFAULT_S1SIZE, + generate_golden=True, +): + if is_batch: + params = normalize_qliv2_params(params) + ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + weight_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + output_idx_offset, + ) = params + + if is_batch: + dtype_map = { + "INT8": torch.int8, + "UINT8": torch.uint8, + "HIFLOAT8": torch.uint8, + "INT32": torch.int32, + "INT64": torch.int64, + "FP16": torch.float16, + "FLOAT16": torch.float16, + "FP32": torch.float32, + "FLOAT": torch.float32, + "FLOAT32": torch.float32, + "BF16": torch.bfloat16, + "FLOAT8_E4M3FN": torch.float8_e4m3fn, + "FLOAT8_E8M0": torch.float8_e8m0fnu, + "FLOAT8_E8M0FNU": torch.float8_e8m0fnu, + "FLOAT4_E2M1": MXFP4_TORCH_DTYPE, + "FLOAT4_E2M1FN_X2": MXFP4_TORCH_DTYPE, + } + qk_dtype = dtype_map.get(qk_dtype, qk_dtype) + weight_dtype = dtype_map.get(weight_dtype, weight_dtype) + dequant_dtype = dtype_map.get(dequant_dtype, dequant_dtype) + actual_seq_dtype = dtype_map.get(actual_seq_dtype, actual_seq_dtype) + + if cu_seqlens_q is not None and isinstance(cu_seqlens_q, str): + cu_seqlens_q = ast.literal_eval(cu_seqlens_q) + if cu_seqlens_k is not None and isinstance(cu_seqlens_k, str): + cu_seqlens_k = ast.literal_eval(cu_seqlens_k) + if seqused_q is not None and isinstance(seqused_q, str): + seqused_q = ast.literal_eval(seqused_q) + if seqused_k is not None and isinstance(seqused_k, str): + seqused_k = ast.literal_eval(seqused_k) + if cmp_residual_k is not None and isinstance(cmp_residual_k, str): + cmp_residual_k = ast.literal_eval(cmp_residual_k) + if query_datarange is not None and isinstance(query_datarange, str): + query_datarange = ast.literal_eval(query_datarange) + if key_datarange is not None and isinstance(key_datarange, str): + key_datarange = ast.literal_eval(key_datarange) + if weights_datarange is not None and isinstance(weights_datarange, str): + weights_datarange = ast.literal_eval(weights_datarange) + if output_idx_offset is not None and isinstance(output_idx_offset, str): + output_idx_offset = ast.literal_eval(output_idx_offset) + output_idx_offset = [int(x) for x in output_idx_offset] + if layout_query == "TND": + output_idx_offset_size = q_t_size * 1 + else: + output_idx_offset_size = batch_size * q_seq * 1 + output_idx_offset = [ + [random.randint(output_idx_offset[0], output_idx_offset[1]) for _ in range(output_idx_offset_size)] + for _ in range(1) + ] + if isinstance(q_scale_datarange, str): + q_scale_datarange = ast.literal_eval(q_scale_datarange) + if isinstance(k_scale_datarange, str): + k_scale_datarange = ast.literal_eval(k_scale_datarange) + + hifp8mode = 1 if quant_mode == QUANT_MODE_HIFLOAT8 else 0 + if is_mx_quant_mode(quant_mode): + validate_mx_scale_dtype(dequant_dtype) + + # ======================== 核心推导:从 cu_seqlens / seqused 推导个体长度 ======================== + # 辅助函数:从前缀和 cu_seqlens [B+1] 推导个体长度 [B] + def _cu_seqlens_to_lengths(cu_list): + return [cu_list[i + 1] - cu_list[i] for i in range(len(cu_list) - 1)] + + # Q 侧个体长度(CPU golden 用) + if layout_query == "TND": + # TND: 必传 cu_seqlens_q,从差分推导个体长度 + assert cu_seqlens_q is not None, "TND layout requires cu_seqlens_q" + lengths_q_list = _cu_seqlens_to_lengths(cu_seqlens_q) + else: + # BSND: 从 seqused_q 获取,若 None 则用 q_seq 填满 + if seqused_q is not None: + lengths_q_list = list(seqused_q) + else: + lengths_q_list = [q_seq] * batch_size + + # K 侧个体长度(CPU golden 用) + if layout_key == "TND": + # TND: 必传 cu_seqlens_k,从差分推导个体长度 + assert cu_seqlens_k is not None, "TND layout requires cu_seqlens_k" + lengths_k_list = _cu_seqlens_to_lengths(cu_seqlens_k) + elif layout_key == "PA_BBND": + # PA_BBND: 从 seqused_k 获取 + assert seqused_k is not None, f"{layout_key} layout requires seqused_k" + lengths_k_list = list(seqused_k) + else: + # BSND: 从 seqused_k 获取,若 None 则用 q_seq 填满 + if seqused_k is not None: + lengths_k_list = list(seqused_k) + else: + lengths_k_list = [k_seq] * batch_size + + # ======================== 构造 NPU 输入 tensor ======================== + # cu_seqlens tensor(仅 TND 传入) + if layout_query == "TND": + cu_seqlens_query = torch.tensor(cu_seqlens_q).to(actual_seq_dtype) + else: + cu_seqlens_query = None + + if layout_key == "TND": + cu_seqlens_key = torch.tensor(cu_seqlens_k).to(actual_seq_dtype) + else: + cu_seqlens_key = None + + # seqused tensor + if seqused_q is not None: + seqused_q_tensor = torch.tensor(seqused_q).to(actual_seq_dtype) + else: + seqused_q_tensor = None + if seqused_k is not None: + seqused_k_tensor = torch.tensor(seqused_k).to(actual_seq_dtype) + else: + seqused_k_tensor = None + + # ======================== CPU golden forward 用的 actual_seq ======================== + # TND: actual_seq 是前缀和格式,即 cu_seqlens[1:](去掉首位 0) + # golden.forward 内部会 trans_tnd_actseq 差分为个体长度 + # BSND/PA: actual_seq 是个体长度,即 seqused + # (actual_seq始终传入,CPU golden 也需要) + if layout_query == "TND": + actual_seq_lengths_query = torch.tensor(cu_seqlens_q[1:]).to(actual_seq_dtype) + else: + actual_seq_lengths_query = torch.tensor(lengths_q_list).to(actual_seq_dtype) + + if layout_key == "TND": + actual_seq_lengths_key = torch.tensor(cu_seqlens_k[1:]).to(actual_seq_dtype) + else: + actual_seq_lengths_key = torch.tensor(lengths_k_list).to(actual_seq_dtype) + + # PA_BBND key 构造用的 act_seq_k 列表 + act_seq_k = lengths_k_list + + # 检查 cmp_residual_k 参数 + if (sparse_mode == 0 or cmp_ratio == 1) and cmp_residual_k is not None: + print( + f"Warning: sparse_mode={sparse_mode} or cmp_ratio={cmp_ratio}, " + f"cmp_residual_k={cmp_residual_k}, should be None" + ) + print("Hint: set cmp_residual_k to None when sparse_mode==0 or cmp_ratio==1") + + # cmp_residual_k for CPU golden (always a list with zeros when cmp_ratio==1 or sparse_mode==0) + if cmp_ratio == 1 or sparse_mode == 0: + cmp_residual_k_for_cpu = [0] * batch_size + else: + cmp_residual_k_for_cpu = list(cmp_residual_k) + + # cmp_residual_k for NPU (None when cmp_ratio==1 or sparse_mode==0, tensor otherwise) + if cmp_ratio == 1 or sparse_mode == 0: + cmp_residual_k_for_npu = None + else: + cmp_residual_k_for_npu = torch.tensor(cmp_residual_k).to(actual_seq_dtype) + + if cu_seqlens_q is not None: + cu_seqlens_q = torch.tensor(cu_seqlens_q).to(torch.int32) + if cu_seqlens_k is not None: + cu_seqlens_k = torch.tensor(cu_seqlens_k).to(torch.int32) + if seqused_q is not None: + seqused_q = torch.tensor(seqused_q).to(torch.int32) + if seqused_k is not None: + seqused_k = torch.tensor(seqused_k).to(torch.int32) + # ======================== 构造 GeneralizedQLIV2 用于 CPU golden ======================== + # GeneralizedQLIV2 需要 act_seq 个体长度(用于 TND→BNSD 转换等) + test_qliv2 = GeneralizedQLIV2( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + weight_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + lengths_q_list, + lengths_k_list, + cmp_residual_k_for_cpu, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + split_s1=split_s1, + s1size=s1size, + ) + + qk_physical_head_dim = get_qk_physical_head_dim(head_dim, quant_mode) + mx_scale_tail_shape = (head_dim // MX_SCALE_SHAPE_ALIGN, MX_SCALE_PACK_NUM) + if layout_query == "BSND": + q_logical_shape = (batch_size, q_seq, q_head_num, head_dim) + q_physical_shape = (batch_size, q_seq, q_head_num, qk_physical_head_dim) + if is_mxfp4_quant_mode(quant_mode): + query, query_cpu_ref = make_mxfp4_tensor_pair(q_logical_shape, query_datarange, qk_dtype, to_npu=False) + else: + query_base = torch.tensor(np.random.uniform(query_datarange[0], query_datarange[1], q_logical_shape)).to( + torch.float + ) + if hifp8mode == 1: + query = trans_float_tensor_to_hifuint8(query_base, round_mode="hybrid", over_mode=True) + else: + query = query_base.to(qk_dtype) + query_cpu_ref = query + + q_scale = random.uniform(q_scale_datarange[0], q_scale_datarange[1]) + if is_mx_quant_mode(quant_mode): + query_dequant_scale, query_dequant_scale_cpu = make_mx_e8m0_scale_pair( + (batch_size, q_seq, q_head_num), + mx_scale_tail_shape, + q_scale_datarange, + dequant_dtype, + to_npu=False, + ) + elif quant_mode == QUANT_MODE_HIFLOAT8: + query_dequant_scale = torch.tensor([q_scale]).to(dequant_dtype) + query_dequant_scale_cpu = torch.tensor( + np.random.uniform(q_scale, q_scale, (batch_size, q_seq, q_head_num)) + ).to(dequant_dtype) + else: + query_dequant_scale = torch.tensor( + np.random.uniform( + q_scale_datarange[0], + q_scale_datarange[1], + (batch_size, q_seq, q_head_num), + ) + ).to(dequant_dtype) + query_dequant_scale_cpu = query_dequant_scale + + weights_cpu = torch.tensor( + np.random.uniform( + weights_datarange[0], + weights_datarange[1], + (batch_size, q_seq, q_head_num), + ) + ).to(weight_dtype) + weights = weights_cpu + if output_idx_offset is not None: + output_idx_offset = torch.tensor(output_idx_offset).reshape(batch_size, q_seq, 1).to(torch.int32) + elif layout_query == "TND": + q_logical_shape = (q_t_size, q_head_num, head_dim) + q_physical_shape = (q_t_size, q_head_num, qk_physical_head_dim) + if is_mxfp4_quant_mode(quant_mode): + query, query_cpu_ref = make_mxfp4_tensor_pair(q_logical_shape, query_datarange, qk_dtype, to_npu=False) + else: + query_base = torch.tensor(np.random.uniform(query_datarange[0], query_datarange[1], q_logical_shape)).to( + torch.float + ) + if hifp8mode == 1: + query = trans_float_tensor_to_hifuint8(query_base, round_mode="hybrid", over_mode=True) + else: + query = query_base.to(qk_dtype) + query_cpu_ref = query + + q_scale = random.uniform(q_scale_datarange[0], q_scale_datarange[1]) + if is_mx_quant_mode(quant_mode): + query_dequant_scale, query_dequant_scale_cpu = make_mx_e8m0_scale_pair( + (q_t_size, q_head_num), + mx_scale_tail_shape, + q_scale_datarange, + dequant_dtype, + to_npu=False, + ) + elif quant_mode == QUANT_MODE_HIFLOAT8: + query_dequant_scale = torch.tensor([q_scale]).to(dequant_dtype) + query_dequant_scale_cpu = torch.tensor(np.random.uniform(q_scale, q_scale, (q_t_size, q_head_num))).to( + dequant_dtype + ) + else: + query_dequant_scale = torch.tensor( + np.random.uniform(q_scale_datarange[0], q_scale_datarange[1], (q_t_size, q_head_num)) + ).to(dequant_dtype) + query_dequant_scale_cpu = query_dequant_scale + + weights_cpu = torch.tensor( + np.random.uniform(weights_datarange[0], weights_datarange[1], (q_t_size, q_head_num)) + ).to(weight_dtype) + weights = weights_cpu + if output_idx_offset is not None: + output_idx_offset = torch.tensor(output_idx_offset).reshape(q_t_size, 1).to(torch.int32) + + blockFusion = None + if layout_key == "BSND": + k_logical_shape = (batch_size, k_seq, k_head_num, head_dim) + k_physical_shape = (batch_size, k_seq, k_head_num, qk_physical_head_dim) + if is_mxfp4_quant_mode(quant_mode): + key, key_cpu_ref = make_mxfp4_tensor_pair(k_logical_shape, key_datarange, qk_dtype, to_npu=False) + else: + key_base = torch.tensor(np.random.uniform(key_datarange[0], key_datarange[1], k_logical_shape)).to( + torch.float + ) + if hifp8mode == 1: + key = trans_float_tensor_to_hifuint8(key_base, round_mode="hybrid", over_mode=True) + else: + key = key_base.to(qk_dtype) + key_cpu_ref = key + + k_scale = random.uniform(k_scale_datarange[0], k_scale_datarange[1]) + if is_mx_quant_mode(quant_mode): + key_dequant_scale, key_dequant_scale_cpu = make_mx_e8m0_scale_pair( + (batch_size, k_seq, k_head_num), + mx_scale_tail_shape, + k_scale_datarange, + dequant_dtype, + to_npu=False, + ) + elif quant_mode == QUANT_MODE_HIFLOAT8: + key_dequant_scale = torch.tensor([k_scale]).to(dequant_dtype) + key_dequant_scale_cpu = torch.tensor( + np.random.uniform(k_scale, k_scale, (batch_size, k_seq, k_head_num)) + ).to(dequant_dtype) + else: + key_dequant_scale = torch.tensor( + np.random.uniform( + k_scale_datarange[0], + k_scale_datarange[1], + (batch_size, k_seq, k_head_num), + ) + ).to(dequant_dtype) + key_dequant_scale_cpu = key_dequant_scale + + block_table = None + if generate_golden: + cpu_result, topk_value, cpu_topk_value = test_qliv2.forward( + query_cpu_ref, + key_cpu_ref, + weights, + query_dequant_scale_cpu, + key_dequant_scale_cpu, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + block_table, + output_idx_offset, + ) + else: + cpu_result, topk_value, cpu_topk_value = None, None, None + + elif layout_key == "TND": + k_logical_shape = (k_t_size, k_head_num, head_dim) + k_physical_shape = (k_t_size, k_head_num, qk_physical_head_dim) + if is_mxfp4_quant_mode(quant_mode): + key, key_cpu_ref = make_mxfp4_tensor_pair(k_logical_shape, key_datarange, qk_dtype, to_npu=False) + else: + key_base = torch.tensor(np.random.uniform(key_datarange[0], key_datarange[1], k_logical_shape)).to( + torch.float + ) + if hifp8mode == 1: + key = trans_float_tensor_to_hifuint8(key_base, round_mode="hybrid", over_mode=True) + else: + key = key_base.to(qk_dtype) + key_cpu_ref = key + + k_scale = random.uniform(k_scale_datarange[0], k_scale_datarange[1]) + if is_mx_quant_mode(quant_mode): + key_dequant_scale, key_dequant_scale_cpu = make_mx_e8m0_scale_pair( + (k_t_size, k_head_num), + mx_scale_tail_shape, + k_scale_datarange, + dequant_dtype, + to_npu=False, + ) + elif quant_mode == QUANT_MODE_HIFLOAT8: + key_dequant_scale = torch.tensor([k_scale]).to(dequant_dtype) + key_dequant_scale_cpu = torch.tensor(np.random.uniform(k_scale, k_scale, (k_t_size, k_head_num))).to( + dequant_dtype + ) + else: + key_dequant_scale = torch.tensor( + np.random.uniform(k_scale_datarange[0], k_scale_datarange[1], (k_t_size, k_head_num)) + ).to(dequant_dtype) + key_dequant_scale_cpu = key_dequant_scale + + block_table = None + if generate_golden: + cpu_result, topk_value, cpu_topk_value = test_qliv2.forward( + query_cpu_ref, + key_cpu_ref, + weights, + query_dequant_scale_cpu, + key_dequant_scale_cpu, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + block_table, + output_idx_offset, + ) + else: + cpu_result, topk_value, cpu_topk_value = None, None, None + + elif layout_key == "PA_BBND": + k_max_s2 = math.floor(max(act_seq_k)) + k_max_block_num_per_batch = math.ceil(k_max_s2 / block_size) + k_logical_shape = (batch_size, k_head_num, k_max_s2, head_dim) + k_physical_shape = (batch_size, k_head_num, k_max_s2, qk_physical_head_dim) + if is_mxfp4_quant_mode(quant_mode): + key_bnsd, key_bnsd_cpu_ref = make_mxfp4_tensor_pair(k_logical_shape, key_datarange, qk_dtype, to_npu=False) + else: + key_bnsd_base = torch.tensor(np.random.uniform(key_datarange[0], key_datarange[1], k_logical_shape)).to( + torch.float + ) + if hifp8mode == 1: + key_bnsd = trans_float_tensor_to_hifuint8(key_bnsd_base, round_mode="hybrid", over_mode=True) + else: + key_bnsd = key_bnsd_base.to(qk_dtype) + key_bnsd_cpu_ref = key_bnsd + + k_scale = random.uniform(k_scale_datarange[0], k_scale_datarange[1]) + if is_mx_quant_mode(quant_mode): + key_dequant_scale_bns_mx, key_dequant_scale_bns = make_mx_e8m0_scale_pair( + (batch_size, k_head_num, k_max_s2), + mx_scale_tail_shape, + k_scale_datarange, + dequant_dtype, + to_npu=False, + ) + elif quant_mode == QUANT_MODE_HIFLOAT8: + key_dequant_scale_bns = torch.tensor( + np.random.uniform(k_scale, k_scale, (batch_size, k_head_num, k_max_s2)) + ).to(dequant_dtype) + else: + key_dequant_scale_bns = torch.tensor( + np.random.uniform( + k_scale_datarange[0], + k_scale_datarange[1], + (batch_size, k_head_num, k_max_s2), + ) + ).to(dequant_dtype) + + key_block_num_per_batch = [] + key_block_num_sum = 0 + for cur_act_k in act_seq_k: + cur_cmp_act_k = math.floor(cur_act_k) + cur_key_block_num = math.ceil(cur_cmp_act_k / block_size) + key_block_num_per_batch.append(cur_key_block_num) + key_block_num_sum += cur_key_block_num + if block_num < key_block_num_sum: + raise ValueError("key actual block num < needed block num") + + block_id_list = np.arange(block_num) + block_id_list = np.random.permutation(block_id_list).astype(np.int32) + cur_block_id = 0 + block_table = np.full((batch_size, k_max_block_num_per_batch), fill_value=-1, dtype=np.int32) + batch_idx = 0 + for cur_block_id_threshold in key_block_num_per_batch: + for i_block_id in range(cur_block_id_threshold): + block_table[batch_idx][i_block_id] = block_id_list[cur_block_id] + cur_block_id += 1 + batch_idx += 1 + + if is_mxfp4_quant_mode(quant_mode): + # FP4 shell dtype仅用于接口语义;PA重排按其底层packed uint8字节完成。 + key_storage_dtype = torch.uint8 + key_bnsd_storage = key_bnsd.view(torch.uint8) + else: + key_storage_dtype = qk_dtype + key_bnsd_storage = key_bnsd + key_expand = torch.zeros( + ( + batch_size, + k_head_num, + k_max_block_num_per_batch * block_size, + qk_physical_head_dim, + ), + dtype=key_storage_dtype, + ) + key_expand[:, :, :k_max_s2, :] = key_bnsd_storage + key = torch.zeros( + (block_num, block_size, k_head_num, qk_physical_head_dim), + dtype=key_storage_dtype, + ) + for i_batch in range(batch_size): + for i_block, cur_block_id in enumerate(block_table[i_batch]): + block_start_pos = i_block * block_size + if cur_block_id == -1: + continue + else: + for i_n in range(k_head_num): + key[cur_block_id, :, i_n, :] = key_expand[ + i_batch, + i_n, + block_start_pos : block_start_pos + block_size, + :, + ] + + if is_mx_quant_mode(quant_mode): + key_dequant_scale_expand = make_e8m0_zero( + ( + batch_size, + k_head_num, + k_max_block_num_per_batch * block_size, + *mx_scale_tail_shape, + ), + dequant_dtype, + ) + key_dequant_scale_expand[:, :, :k_max_s2, :, :] = key_dequant_scale_bns_mx + key_dequant_scale_block = make_e8m0_zero( + (block_num, block_size, k_head_num, *mx_scale_tail_shape), dequant_dtype + ) + for i_batch in range(batch_size): + for i_block, cur_block_id in enumerate(block_table[i_batch]): + block_start_pos = i_block * block_size + if cur_block_id == -1: + continue + else: + for i_n in range(k_head_num): + key_dequant_scale_block[cur_block_id, :, i_n, :, :] = key_dequant_scale_expand[ + i_batch, + i_n, + block_start_pos : block_start_pos + block_size, + :, + :, + ] + else: + key_dequant_scale_expand = torch.zeros( + (batch_size, k_head_num, k_max_block_num_per_batch * block_size), + dtype=dequant_dtype, + ) + key_dequant_scale_expand[:, :, :k_max_s2] = key_dequant_scale_bns + key_dequant_scale_block = torch.zeros((block_num, block_size, k_head_num), dtype=dequant_dtype) + for i_batch in range(batch_size): + for i_block, cur_block_id in enumerate(block_table[i_batch]): + block_start_pos = i_block * block_size + if cur_block_id == -1: + continue + else: + for i_n in range(k_head_num): + key_dequant_scale_block[cur_block_id, :, i_n] = key_dequant_scale_expand[ + i_batch, + i_n, + block_start_pos : block_start_pos + block_size, + ] + + if not is_mx_quant_mode(quant_mode) and quant_mode != QUANT_MODE_HIFLOAT8 and DISCONTINUOUS_KEYS: + bytes_per_token = head_dim + key_dequant_scale_block.element_size() // key.element_size() + blockFusion = torch.zeros((block_num, block_size * k_head_num * bytes_per_token), dtype=qk_dtype) + key_flat = key.view(block_num, block_size * k_head_num * head_dim) + scale_flat = key_dequant_scale_block.view(block_num, block_size * k_head_num).view(qk_dtype) + blockFusion[:, : block_size * k_head_num * head_dim] = key_flat + blockFusion[:, block_size * k_head_num * head_dim :] = scale_flat + + if is_mxfp4_quant_mode(quant_mode): + key = key.view(qk_dtype) + if quant_mode == QUANT_MODE_HIFLOAT8: + key_dequant_scale = torch.tensor([k_scale]).to(dequant_dtype) + else: + key_dequant_scale = key_dequant_scale_block + if generate_golden: + cpu_result, topk_value, cpu_topk_value = test_qliv2.forward( + query_cpu_ref, + key_bnsd_cpu_ref, + weights, + query_dequant_scale_cpu, + key_dequant_scale_bns, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + block_table, + output_idx_offset, + ) + else: + cpu_result, topk_value, cpu_topk_value = None, None, None + block_table = torch.from_numpy(block_table).to(dtype=torch.int32) + # ======================== metadata 构造 ======================== + # max_seqlen 从个体长度中取 + max_seqlen_q_meta = actual_seq_lengths_query.max().item() + max_seqlen_k_meta = actual_seq_lengths_key.max().item() + + if is_batch: + if qk_dtype == torch.float8_e4m3fn: + query = query.to(dtype=torch.float16) + key = key.to(dtype=torch.float16) + if blockFusion is not None: + blockFusion = blockFusion.view(torch.uint8) + + golden_key = key_bnsd_cpu_ref if layout_key == "PA_BBND" else key_cpu_ref + golden_key_scale = key_dequant_scale_bns if layout_key == "PA_BBND" else key_dequant_scale_cpu + golden_block_table = block_table + if torch.is_tensor(golden_block_table): + golden_block_table = golden_block_table.detach().cpu() + golden_state = { + "model_args": ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + weight_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + lengths_q_list, + lengths_k_list, + cmp_residual_k_for_cpu, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + ), + "split_s1": split_s1, + "s1size": s1size, + "forward_inputs": { + "query": query_cpu_ref, + "key": golden_key, + "weights": weights, + "query_dequant_scale": query_dequant_scale_cpu, + "key_dequant_scale": golden_key_scale, + "cu_seqlens_q": cu_seqlens_q, + "cu_seqlens_k": cu_seqlens_k, + "seqused_q": seqused_q, + "seqused_k": seqused_k, + "block_table": golden_block_table, + "output_idx_offset": output_idx_offset, + }, + } + + output_tensors = { + "params": params, + "cpu_result": cpu_result, + "topk_value": topk_value, + "cpu_topk_value": cpu_topk_value, + "query": query, + "key": key, + "weights": weights, + "query_dequant_scale": query_dequant_scale, + "key_dequant_scale": key_dequant_scale, + "blockFusion": blockFusion, + "cu_seqlens_query": cu_seqlens_query, + "cu_seqlens_key": cu_seqlens_key, + "seqused_q": seqused_q_tensor, + "seqused_k": seqused_k_tensor, + "output_idx_offset": output_idx_offset, + "actual_seq_lengths_query": actual_seq_lengths_query, + "actual_seq_lengths_key": actual_seq_lengths_key, + "cmp_residual_k_for_npu": cmp_residual_k_for_npu, + "block_table": block_table, + "max_seqlen_q_meta": max_seqlen_q_meta, + "max_seqlen_k_meta": max_seqlen_k_meta, + "quant_mode": quant_mode, + "layout_query": layout_query, + "layout_key": layout_key, + "sparse_count": sparse_count, + "sparse_mode": sparse_mode, + "cmp_ratio": cmp_ratio, + "golden_state": golden_state, + } + return output_tensors + else: + metadata = torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata( + cu_seqlens_q=cu_seqlens_query.npu() if cu_seqlens_query is not None else None, + cu_seqlens_k=cu_seqlens_key.npu() if cu_seqlens_key is not None else None, + seqused_q=seqused_q_tensor.npu() if seqused_q_tensor is not None else None, + seqused_k=seqused_k_tensor.npu() if seqused_k_tensor is not None else None, + cmp_residual_k=cmp_residual_k_for_npu.npu() if cmp_residual_k_for_npu is not None else None, + batch_size=batch_size, + max_seqlen_q=max_seqlen_q_meta, + max_seqlen_k=max_seqlen_k_meta, + num_heads_q=q_head_num, + num_heads_k=k_head_num, + head_dim=head_dim, + topk=sparse_count, + quant_mode=quant_mode, + mask_mode=sparse_mode, + layout_q=layout_query, + layout_k=layout_key, + cmp_ratio=cmp_ratio, + ) + metadata = metadata.npu() + if blockFusion is not None: + blockFusion = blockFusion.npu() + key = blockFusion[:, : block_size * k_head_num * head_dim].view(block_num, block_size, k_head_num, head_dim) + key_dequant_scale_block = ( + blockFusion[:, block_size * k_head_num * head_dim :] + .view(dequant_dtype) + .view(block_num, block_size, k_head_num) + ) + else: + key = key.npu() + if quant_mode == QUANT_MODE_HIFLOAT8: + key_dequant_scale = torch.tensor([k_scale]).to(dequant_dtype).npu() + else: + if layout_key == "PA_BBND": + key_dequant_scale = key_dequant_scale_block.npu() + else: + key_dequant_scale = key_dequant_scale.npu() + + npu_result, npu_topk_value = torch.ops.cann_ops_transformer.quant_lightning_indexer( + query.npu(), + key, + weights.npu(), + query_dequant_scale.npu(), + key_dequant_scale, + cu_seqlens_q=cu_seqlens_query.npu() if cu_seqlens_query is not None else None, + cu_seqlens_k=cu_seqlens_key.npu() if cu_seqlens_key is not None else None, + seqused_q=seqused_q_tensor.npu() if seqused_q_tensor is not None else None, + seqused_k=seqused_k_tensor.npu() if seqused_k_tensor is not None else None, + cmp_residual_k=cmp_residual_k_for_npu.npu() if cmp_residual_k_for_npu is not None else None, + output_idx_offset=output_idx_offset.npu() if output_idx_offset is not None else None, + max_seqlen_q=max_seqlen_q, + block_table=block_table.npu() if block_table is not None else None, + metadata=metadata, + quant_mode=quant_mode, + layout_q=layout_query, + layout_k=layout_key, + topk=sparse_count, + mask_mode=sparse_mode, + cmp_ratio=cmp_ratio, + return_value=return_value, + ) + + torch.npu.synchronize() + if return_value: + if npu_topk_value.shape != npu_result.shape: + raise RuntimeError( + "sparse_values and sparse_indices must have the same shape when return_value=1, " + f"but got {tuple(npu_topk_value.shape)} and {tuple(npu_result.shape)}" + ) + npu_topk_value, npu_sort_order = npu_topk_value.sort(dim=-1, descending=True) + npu_result = torch.gather(npu_result, dim=-1, index=npu_sort_order) + return cpu_result, npu_result, topk_value, cpu_topk_value, npu_topk_value + + +def build_qliv2_metadata_input(params, tensor_dict): + """Return the canonical metadata arguments produced by the pytest case.""" + return { + "num_heads_q": int(params[5]), + "num_heads_k": int(params[6]), + "head_dim": int(params[7]), + "topk": int(tensor_dict["sparse_count"]), + "quant_mode": int(tensor_dict["quant_mode"]), + "cu_seqlens_q": tensor_dict["cu_seqlens_query"], + "cu_seqlens_k": tensor_dict["cu_seqlens_key"], + "seqused_q": tensor_dict["seqused_q"], + "seqused_k": tensor_dict["seqused_k"], + "cmp_residual_k": tensor_dict["cmp_residual_k_for_npu"], + "batch_size": int(params[0]), + "max_seqlen_q": int(tensor_dict["max_seqlen_q_meta"]), + "max_seqlen_k": int(tensor_dict["max_seqlen_k_meta"]), + "layout_q": tensor_dict["layout_query"], + "layout_k": tensor_dict["layout_key"], + "mask_mode": int(tensor_dict["sparse_mode"]), + "cmp_ratio": int(tensor_dict["cmp_ratio"]), + } + + +def generate_qliv2_test_data( + params, + split_s1=DEFAULT_SPLIT_S1, + s1size=DEFAULT_S1SIZE, + generate_golden=True, +): + """Generate QLI_V2 inputs and optionally materialize the CPU Golden.""" + data = qliv2_output_single( + params, + is_batch=True, + split_s1=split_s1, + s1size=s1size, + generate_golden=generate_golden, + ) + data["metadata_input"] = build_qliv2_metadata_input(data["params"], data) + return data + + +def generate_cpu_golden(input_data): + """Calculate CPU Golden from input-stage data without rerunning random generation.""" + state = input_data["golden_state"] + test_qliv2 = GeneralizedQLIV2( + *state["model_args"], + split_s1=state["split_s1"], + s1size=state["s1size"], + ) + values = state["forward_inputs"] + block_table = values["block_table"] + if torch.is_tensor(block_table): + # The CPU reference indexes block tables as an ndarray; keep the saved + # tensor intact for PT consumers and convert only this local argument. + block_table = block_table.detach().cpu().numpy() + cpu_result, topk_value, cpu_topk_value = test_qliv2.forward( + values["query"], + values["key"], + values["weights"], + values["query_dequant_scale"], + values["key_dequant_scale"], + values["cu_seqlens_q"], + values["cu_seqlens_k"], + values["seqused_q"], + values["seqused_k"], + block_table, + values["output_idx_offset"], + ) + state["cpu_result"] = cpu_result + state["topk_value"] = topk_value + state["cpu_topk_value"] = cpu_topk_value + input_data["cpu_result"] = cpu_result + input_data["topk_value"] = topk_value + input_data["cpu_topk_value"] = cpu_topk_value + return cpu_result, topk_value, cpu_topk_value + + +def fp32_ta_round_to_hif8(fraction32_int, hif8_bits_num, exponent): + if exponent == HIF8_EXP_ZERO_THRESHOLD: + return True, 0 + hif8_value_tmp = fraction32_int >> (FP32_FRACTION_BITS - (hif8_bits_num + 1)) + if hif8_value_tmp == pow(2, hif8_bits_num + 1) - 1: + return True, 0 + elif hif8_value_tmp == 0: + return False, 0 + elif hif8_value_tmp % 2 == 1: + hif8_value_tmp += 1 + return False, hif8_value_tmp >> 1 + else: + return False, hif8_value_tmp >> 1 + + +def fp32_ssr_round_to_hif8(fraction32_int, hif8_bits_num, exponent): + t14_mask = SSR_T14_MASK + if exponent == HIF8_EXP_ZERO_THRESHOLD: + f14_values = (fraction32_int >> SSR_DML_SHIFT) + SSR_F14_OFFSET + t14_values = fraction32_int & t14_mask + hif8_value = 0 + else: + hif8_value = fraction32_int >> (FP32_FRACTION_BITS - hif8_bits_num) + f14_t14 = fraction32_int - (hif8_value << (FP32_FRACTION_BITS - hif8_bits_num)) + f14_values = f14_t14 >> (FP32_FRACTION_BITS - hif8_bits_num - SSR_RESERVED_BITS) + t14_values = f14_t14 & t14_mask + if f14_values >= t14_values: + if hif8_value == pow(2, hif8_bits_num) - 1: + return True, 0 + else: + hif8_value += 1 + return False, hif8_value + else: + return False, hif8_value + + +def get_hif8_fraction_bits_number(exponent): + if exponent < HIF8_EXP_DML_MIN: + return HIF8_DOT_INVALID, HIF8_EXP_BITS_DML, HIF8_FRAC_BITS_DML + if HIF8_EXP_DML_MIN <= exponent < HIF8_EXP_DML_MAX: + return HIF8_DOT_DML, HIF8_EXP_BITS_DML, HIF8_FRAC_BITS_DML + if exponent == HIF8_EXP_D0: + return HIF8_DOT_D0, HIF8_EXP_BITS_D0, HIF8_FRAC_BITS_D0 + if abs(exponent) == HIF8_EXP_D1_BOUNDARY: + return HIF8_DOT_D1, HIF8_EXP_BITS_D1, HIF8_FRAC_BITS_D1 + if HIF8_EXP_D2_MIN <= abs(exponent) <= HIF8_EXP_D2_MAX: + return HIF8_DOT_D2, HIF8_EXP_BITS_D2, HIF8_FRAC_BITS_D2 + if HIF8_EXP_D3_MIN <= abs(exponent) <= HIF8_EXP_D3_MAX: + return HIF8_DOT_D3, HIF8_EXP_BITS_D3, HIF8_FRAC_BITS_D3 + if HIF8_EXP_D4_MIN <= abs(exponent) <= HIF8_EXP_D4_MAX: + return HIF8_DOT_D4, HIF8_EXP_BITS_D4, HIF8_FRAC_BITS_D4 + if exponent > HIF8_EXP_D4_MAX: + return HIF8_DOT_D4, HIF8_EXP_BITS_D4, HIF8_DOT_INVALID + + +def cvt_float32_to_hifuint8(x, round_mode="round", over_mode=True): + sign = False + sign_int_value = 0 + x_abs = math.fabs(x) + ec = 0 + over_value = HIF8_OVERFLOW_SCALE * pow(2.0, HIF8_EXP_D4_MAX + ec) + if x < 0.0: + sign = True + sign_int_value = HIF8_SIGN_MASK + if torch.isinf(x) or x_abs >= over_value: + if sign: + if over_mode: + return HIF8_NEG_INF + else: + return HIF8_NEG_MAX + else: + if over_mode: + return HIF8_POS_INF + else: + return HIF8_POS_MAX + if torch.isnan(x): + if over_mode: + return HIF8_NAN + else: + return 0 + if x_abs == 0.0: + return 0 + exponent = math.floor(math.log2(x_abs)) + if round_mode == "hybrid": + if abs(exponent) < HYBRID_ROUND_EXP_THRESHOLD: + cut_bit_type = "TA" + else: + cut_bit_type = "SSR" + elif round_mode == "round": + cut_bit_type = "TA" + elif round_mode == "storound": + cut_bit_type = "SSR" + else: + cut_bit_type = "TA" + fraction_int = int(x_abs * pow(2, FP32_FRACTION_BITS) * pow(2, -exponent) - pow(2, FP32_FRACTION_BITS)) + dot_hif8_value, exponent_hif8_bits, fraction_hif8_bits = get_hif8_fraction_bits_number(exponent) + if cut_bit_type == "TA": + carry_exp_status, hif8_frac_value = fp32_ta_round_to_hif8(fraction_int, fraction_hif8_bits, exponent) + elif cut_bit_type == "SSR": + carry_exp_status, hif8_frac_value = fp32_ssr_round_to_hif8(fraction_int, fraction_hif8_bits, exponent) + else: + print("unknown round type") + return 0 + + if carry_exp_status: + exponent += 1 + dot_hif8_value, exponent_hif8_bits, fraction_hif8_bits_new = get_hif8_fraction_bits_number(exponent) + fraction_hif8_bits = fraction_hif8_bits_new + if exponent < HIF8_EXP_ZERO_THRESHOLD: + return 0 + if exponent < 0: + sig_exp = 1 + else: + sig_exp = 0 + if dot_hif8_value <= 0: + if exponent <= HIF8_EXP_ZERO_THRESHOLD: + return 0 + else: + return sign_int_value + exponent + HIF8_DML_EXP_OFFSET + elif dot_hif8_value == 1: + dot_int_value = dot_hif8_value << HIF8_DOT_BIT_SHIFT + hif8_int_value = sign_int_value + dot_int_value + hif8_frac_value + else: + abs_exponent = abs(exponent) + abs_exponent = abs_exponent - pow(2, exponent_hif8_bits - 1) + exponent_int_value = abs_exponent << fraction_hif8_bits + sig_exp = sig_exp << (exponent_hif8_bits - 1 + fraction_hif8_bits) + dot_int_value = dot_hif8_value << HIF8_DOT_BIT_SHIFT + hif8_int_value = sign_int_value + dot_int_value + sig_exp + exponent_int_value + hif8_frac_value + return hif8_int_value + + +def trans_float_tensor_to_hifuint8(in_tensor, round_mode="round", over_mode=True): + """ + 通过向量操作,将 float32 Tensor 批量转换为 HiF8 编码的 uint8 Tensor + """ + + shape = in_tensor.shape + x = in_tensor.reshape(-1).to(torch.float32) + + # 先用int32作为输出类型,避免出现赋值错误 + out = torch.zeros_like(x, dtype=torch.int32) + + # 1. 符号位与绝对值提取 + sign_mask = x < 0.0 + sign_int_value = torch.where(sign_mask, HIF8_SIGN_MASK, 0) + x_abs = torch.abs(x) + + # 2. 溢出与边界条件判断 (Masks) + over_value = HIF8_OVERFLOW_SCALE * (2.0**HIF8_EXP_D4_MAX) + mask_inf_or_over = torch.isinf(x) | (x_abs >= over_value) + mask_nan = torch.isnan(x) + mask_zero = x_abs == 0.0 + + # 处理特殊边界填值 + if over_mode: + out = torch.where(mask_inf_or_over, torch.where(sign_mask, HIF8_NEG_INF, HIF8_POS_INF), out) + out = torch.where(mask_nan, HIF8_NAN, out) + else: + out = torch.where(mask_inf_or_over, torch.where(sign_mask, HIF8_NEG_MAX, HIF8_POS_MAX), out) + out = torch.where(mask_nan, 0, out) + out = torch.where(mask_zero, 0, out) + + # 提取正常数字的 Mask + mask_normal = ~(mask_inf_or_over | mask_nan | mask_zero) + if not mask_normal.any(): + return out.reshape(shape).to(torch.uint8) + + x_norm = x_abs[mask_normal] + sign_norm = sign_int_value[mask_normal] + + # 计算基本指数 + exponent = torch.floor(torch.log2(x_norm)).to(torch.int32) + + # 确定截断模式 (TA / SSR) + if round_mode == "hybrid": + cut_bit_is_ta = torch.abs(exponent) < HYBRID_ROUND_EXP_THRESHOLD + elif round_mode == "round": + cut_bit_is_ta = torch.ones_like(exponent, dtype=torch.bool) + elif round_mode == "storound": + cut_bit_is_ta = torch.zeros_like(exponent, dtype=torch.bool) + else: + cut_bit_is_ta = torch.ones_like(exponent, dtype=torch.bool) + + # 计算 fraction_int + fraction_int = ( + x_norm * (2.0**FP32_FRACTION_BITS) * torch.pow(2.0, -exponent.float()) - (2.0**FP32_FRACTION_BITS) + ).to(torch.int32) + + # 批量获取档位属性 (根据 exponent 映射) + abs_exp = torch.abs(exponent) + + dot = torch.full_like(exponent, HIF8_DOT_INVALID) + exp_bits = torch.zeros_like(exponent) + frac_bits = torch.zeros_like(exponent) + + # 条件区间映射 + m1 = exponent < HIF8_EXP_DML_MIN + dot = torch.where(m1, HIF8_DOT_INVALID, dot) + exp_bits = torch.where(m1, HIF8_EXP_BITS_DML, exp_bits) + frac_bits = torch.where(m1, HIF8_FRAC_BITS_DML, frac_bits) + + m2 = (~m1) & (exponent >= HIF8_EXP_DML_MIN) & (exponent < HIF8_EXP_DML_MAX) + dot = torch.where(m2, HIF8_DOT_DML, dot) + exp_bits = torch.where(m2, HIF8_EXP_BITS_DML, exp_bits) + frac_bits = torch.where(m2, HIF8_FRAC_BITS_DML, frac_bits) + + m3 = exponent == HIF8_EXP_D0 + dot = torch.where(m3, HIF8_DOT_D0, dot) + exp_bits = torch.where(m3, HIF8_EXP_BITS_D0, exp_bits) + frac_bits = torch.where(m3, HIF8_FRAC_BITS_D0, frac_bits) + + m4 = abs_exp == HIF8_EXP_D1_BOUNDARY + dot = torch.where(m4, HIF8_DOT_D1, dot) + exp_bits = torch.where(m4, HIF8_EXP_BITS_D1, exp_bits) + frac_bits = torch.where(m4, HIF8_FRAC_BITS_D1, frac_bits) + + m5 = (abs_exp >= HIF8_EXP_D2_MIN) & (abs_exp <= HIF8_EXP_D2_MAX) + dot = torch.where(m5, HIF8_DOT_D2, dot) + exp_bits = torch.where(m5, HIF8_EXP_BITS_D2, exp_bits) + frac_bits = torch.where(m5, HIF8_FRAC_BITS_D2, frac_bits) + + m6 = (abs_exp >= HIF8_EXP_D3_MIN) & (abs_exp <= HIF8_EXP_D3_MAX) + dot = torch.where(m6, HIF8_DOT_D3, dot) + exp_bits = torch.where(m6, HIF8_EXP_BITS_D3, exp_bits) + frac_bits = torch.where(m6, HIF8_FRAC_BITS_D3, frac_bits) + + m7 = (abs_exp >= HIF8_EXP_D4_MIN) & (abs_exp <= HIF8_EXP_D4_MAX) + dot = torch.where(m7, HIF8_DOT_D4, dot) + exp_bits = torch.where(m7, HIF8_EXP_BITS_D4, exp_bits) + frac_bits = torch.where(m7, HIF8_FRAC_BITS_D4, frac_bits) + + m8 = exponent > HIF8_EXP_D4_MAX + dot = torch.where(m8, HIF8_DOT_D4, dot) + exp_bits = torch.where(m8, HIF8_EXP_BITS_D4, exp_bits) + frac_bits = torch.where(m8, HIF8_DOT_INVALID, frac_bits) + + # ------------------ TA 舍入分支 ------------------ + carry_ta = torch.zeros_like(exponent, dtype=torch.bool) + frac_val_ta = torch.zeros_like(exponent) + + m_zero_thresh = exponent == HIF8_EXP_ZERO_THRESHOLD + carry_ta = torch.where(m_zero_thresh, True, carry_ta) + + m_ta_norm = ~m_zero_thresh + shift_bits = torch.clamp(FP32_FRACTION_BITS - (frac_bits + 1), min=0) + hif8_val_tmp = fraction_int >> shift_bits + + pow_frac = torch.pow(2, frac_bits + 1) - 1 + m_carry = m_ta_norm & (hif8_val_tmp == pow_frac) + carry_ta = torch.where(m_carry, True, carry_ta) + + m_odd = m_ta_norm & (~m_carry) & (hif8_val_tmp != 0) & (hif8_val_tmp % 2 == 1) + frac_val_ta = torch.where(m_odd, (hif8_val_tmp + 1) >> 1, frac_val_ta) + + m_even = m_ta_norm & (~m_carry) & (hif8_val_tmp != 0) & (hif8_val_tmp % 2 == 0) + frac_val_ta = torch.where(m_even, hif8_val_tmp >> 1, frac_val_ta) + + # ------------------ SSR 舍入分支 ------------------ + carry_ssr = torch.zeros_like(exponent, dtype=torch.bool) + frac_val_ssr = torch.zeros_like(exponent) + + f14_v1 = (fraction_int >> SSR_DML_SHIFT) + SSR_F14_OFFSET + t14_v1 = fraction_int & SSR_T14_MASK + hif8_v1 = torch.zeros_like(fraction_int) + + s_bits = torch.clamp(FP32_FRACTION_BITS - frac_bits, min=0) + hif8_v2 = fraction_int >> s_bits + f14_t14 = fraction_int - (hif8_v2 << s_bits) + s_bits_f14 = torch.clamp(FP32_FRACTION_BITS - frac_bits - SSR_RESERVED_BITS, min=0) + f14_v2 = f14_t14 >> s_bits_f14 + t14_v2 = f14_t14 & SSR_T14_MASK + + f14_values = torch.where(m_zero_thresh, f14_v1, f14_v2) + t14_values = torch.where(m_zero_thresh, t14_v1, t14_v2) + hif8_value = torch.where(m_zero_thresh, hif8_v1, hif8_v2) + + m_ge = f14_values >= t14_values + pow_frac_ssr = torch.pow(2, frac_bits) - 1 + m_ssr_carry = m_ge & (hif8_value == pow_frac_ssr) + carry_ssr = torch.where(m_ssr_carry, True, carry_ssr) + frac_val_ssr = torch.where(m_ge & (~m_ssr_carry), hif8_value + 1, frac_val_ssr) + frac_val_ssr = torch.where(~m_ge, hif8_value, frac_val_ssr) + + # ------------------ 合并舍入结果 ------------------ + carry_exp_status = torch.where(cut_bit_is_ta, carry_ta, carry_ssr) + hif8_frac_value = torch.where(cut_bit_is_ta, frac_val_ta, frac_val_ssr) + + exponent = torch.where(carry_exp_status, exponent + 1, exponent) + abs_exp = torch.abs(exponent) + + dot = torch.where( + carry_exp_status, + torch.where(exponent < HIF8_EXP_DML_MIN, HIF8_DOT_INVALID, dot), + dot, + ) + dot = torch.where( + carry_exp_status, + torch.where( + (exponent >= HIF8_EXP_DML_MIN) & (exponent < HIF8_EXP_DML_MAX), + HIF8_DOT_DML, + dot, + ), + dot, + ) + dot = torch.where(carry_exp_status, torch.where(exponent == HIF8_EXP_D0, HIF8_DOT_D0, dot), dot) + dot = torch.where( + carry_exp_status, + torch.where(abs_exp == HIF8_EXP_D1_BOUNDARY, HIF8_DOT_D1, dot), + dot, + ) + dot = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D2_MIN) & (abs_exp <= HIF8_EXP_D2_MAX), + HIF8_DOT_D2, + dot, + ), + dot, + ) + dot = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D3_MIN) & (abs_exp <= HIF8_EXP_D3_MAX), + HIF8_DOT_D3, + dot, + ), + dot, + ) + dot = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D4_MIN) & (abs_exp <= HIF8_EXP_D4_MAX), + HIF8_DOT_D4, + dot, + ), + dot, + ) + + frac_bits = torch.where( + carry_exp_status, + torch.where(exponent < HIF8_EXP_DML_MIN, HIF8_FRAC_BITS_DML, frac_bits), + frac_bits, + ) + frac_bits = torch.where( + carry_exp_status, + torch.where( + (exponent >= HIF8_EXP_DML_MIN) & (exponent < HIF8_EXP_DML_MAX), + HIF8_FRAC_BITS_DML, + frac_bits, + ), + frac_bits, + ) + frac_bits = torch.where( + carry_exp_status, + torch.where(exponent == HIF8_EXP_D0, HIF8_FRAC_BITS_D0, frac_bits), + frac_bits, + ) + frac_bits = torch.where( + carry_exp_status, + torch.where(abs_exp == HIF8_EXP_D1_BOUNDARY, HIF8_FRAC_BITS_D1, frac_bits), + frac_bits, + ) + frac_bits = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D2_MIN) & (abs_exp <= HIF8_EXP_D2_MAX), + HIF8_FRAC_BITS_D2, + frac_bits, + ), + frac_bits, + ) + frac_bits = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D3_MIN) & (abs_exp <= HIF8_EXP_D3_MAX), + HIF8_FRAC_BITS_D3, + frac_bits, + ), + frac_bits, + ) + frac_bits = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D4_MIN) & (abs_exp <= HIF8_EXP_D4_MAX), + HIF8_FRAC_BITS_D4, + frac_bits, + ), + frac_bits, + ) + + exp_bits = torch.where( + carry_exp_status, + torch.where(exponent < HIF8_EXP_DML_MIN, HIF8_EXP_BITS_DML, exp_bits), + exp_bits, + ) + exp_bits = torch.where( + carry_exp_status, + torch.where( + (exponent >= HIF8_EXP_DML_MIN) & (exponent < HIF8_EXP_DML_MAX), + HIF8_EXP_BITS_DML, + exp_bits, + ), + exp_bits, + ) + exp_bits = torch.where( + carry_exp_status, + torch.where(exponent == HIF8_EXP_D0, HIF8_EXP_BITS_D0, exp_bits), + exp_bits, + ) + exp_bits = torch.where( + carry_exp_status, + torch.where(abs_exp == HIF8_EXP_D1_BOUNDARY, HIF8_EXP_BITS_D1, exp_bits), + exp_bits, + ) + exp_bits = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D2_MIN) & (abs_exp <= HIF8_EXP_D2_MAX), + HIF8_EXP_BITS_D2, + exp_bits, + ), + exp_bits, + ) + exp_bits = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D3_MIN) & (abs_exp <= HIF8_EXP_D3_MAX), + HIF8_EXP_BITS_D3, + exp_bits, + ), + exp_bits, + ) + exp_bits = torch.where( + carry_exp_status, + torch.where( + (abs_exp >= HIF8_EXP_D4_MIN) & (abs_exp <= HIF8_EXP_D4_MAX), + HIF8_EXP_BITS_D4, + exp_bits, + ), + exp_bits, + ) + + # ------------------ 组合输出编码 ------------------ + hif8_int_value = torch.zeros_like(exponent) + sig_exp = torch.where(exponent < 0, 1, 0) + + # 分支 A: dot <= 0 + m_a = dot <= 0 + val_a = torch.where( + exponent <= HIF8_EXP_ZERO_THRESHOLD, + 0, + sign_norm + exponent + HIF8_DML_EXP_OFFSET, + ) + hif8_int_value = torch.where(m_a, val_a, hif8_int_value) + + # 分支 B: dot == 1 + m_b = dot == 1 + val_b = sign_norm + (dot << HIF8_DOT_BIT_SHIFT) + hif8_frac_value + hif8_int_value = torch.where(m_b, val_b, hif8_int_value) + + # 分支 C: dot > 1 + m_c = dot > 1 + abs_exponent = torch.abs(exponent) + abs_exponent = abs_exponent - torch.pow(2, exp_bits - 1) + exponent_int_value = abs_exponent << frac_bits + sig_exp_shifted = sig_exp << (exp_bits - 1 + frac_bits) + dot_int_value = dot << HIF8_DOT_BIT_SHIFT + val_c = sign_norm + dot_int_value + sig_exp_shifted + exponent_int_value + hif8_frac_value + hif8_int_value = torch.where(m_c, val_c, hif8_int_value) + + hif8_int_value = torch.where(exponent < HIF8_EXP_ZERO_THRESHOLD, 0, hif8_int_value) + + out[mask_normal] = hif8_int_value + + return out.reshape(shape).to(torch.uint8) + + +def cvt_hifuint8_to_float32(x, over_mode=True): + x = int(x) + if x == HIF8_ZERO: + return float(0) + elif x == HIF8_NAN: + if over_mode: + return float("nan") + else: + return float(0) + elif x == HIF8_NEG_INF: + if over_mode: + return -torch.inf + else: + return -HIF8_MAX_FINITE_VALUE + elif x == HIF8_POS_INF: + if over_mode: + return torch.inf + else: + return HIF8_MAX_FINITE_VALUE + else: + if x >= HIF8_NAN: + sign = -1.0 + else: + sign = 1.0 + dot_4_bits = x & HIF8_DOT_MASK + dot_4_value = dot_4_bits >> 3 + if dot_4_value >= HIF8_DOT_D4: + exponent = x & HIF8_EXP_MASK_D4 + exponent_int = exponent >> 1 + if exponent_int >= 8: + exponent_value = -exponent_int + else: + exponent_value = exponent_int + 8 + + fra_int = x & HIF8_FRAC_MASK_1BIT + m_value = 1.0 + fra_int * 0.5 + elif dot_4_value >= HIF8_DOT_D3: + exponent = x & HIF8_EXP_MASK_D3 + exponent_int = exponent >> 2 + if exponent_int >= 4: + exponent_value = -exponent_int + else: + exponent_value = exponent_int + 4 + + fra_int = x & HIF8_FRAC_MASK_2BIT + m_value = 1.0 + fra_int * 0.25 + elif dot_4_value >= HIF8_DOT_D2: + exponent = x & HIF8_EXP_MASK_D2 + exponent_int = exponent >> 3 + if exponent_int >= 2: + exponent_value = -exponent_int + else: + exponent_value = exponent_int + 2 + + fra_int = x & HIF8_FRAC_MASK_3BIT + m_value = 1.0 + fra_int * 0.125 + elif dot_4_value >= HIF8_DOT_D1: + exponent = x & HIF8_EXP_SIGN_MASK_D1 + exponent_sign = exponent >> 3 + if exponent_sign >= 1: + exponent_value = -1 + else: + exponent_value = 1 + + fra_int = x & HIF8_FRAC_MASK_3BIT + m_value = 1.0 + fra_int * 0.125 + elif dot_4_value == HIF8_DOT_D0: + exponent_value = 0 + fra_int = x & HIF8_FRAC_MASK_3BIT + m_value = 1.0 + fra_int * 0.125 + elif dot_4_value == HIF8_DOT_DML: + m_value = 1 + exponent_value = (x & HIF8_EXP_MASK_DML) - HIF8_DML_EXP_OFFSET + else: + print("error, dot error") + m_value = 0.0 + exponent_value = 0 + return sign * pow(2.0, exponent_value) * m_value + + +def trans_hifuint8_tensor_to_float(in_tensor, over_mode=True): + """ + 将 HiF8 编码的 uint8 Tensor 批量转换为 float32 Tensor (支持 CPU/GPU 矢量化) + """ + shape = in_tensor.shape + x = in_tensor.reshape(-1).to(torch.int32) + out = torch.zeros_like(x, dtype=torch.float32) + + # 1. 特殊值处理 (Masks) + mask_zero = x == HIF8_ZERO + mask_nan = x == HIF8_NAN + mask_ninf = x == HIF8_NEG_INF + mask_pinf = x == HIF8_POS_INF + + if over_mode: + out = torch.where(mask_nan, torch.tensor(float("nan"), device=x.device), out) + out = torch.where(mask_ninf, torch.tensor(-torch.inf, device=x.device), out) + out = torch.where(mask_pinf, torch.tensor(torch.inf, device=x.device), out) + else: + out = torch.where(mask_nan, 0.0, out) + out = torch.where(mask_ninf, float(-HIF8_MAX_FINITE_VALUE), out) + out = torch.where(mask_pinf, float(HIF8_MAX_FINITE_VALUE), out) + + # 正常数值的 Mask (排除特殊值) + mask_normal = ~(mask_zero | mask_nan | mask_ninf | mask_pinf) + if not mask_normal.any(): + return out.reshape(shape) + + # 提取正常数值子集进行计算 + x_norm = x[mask_normal] + + # 符号位计算 + sign = torch.where(x_norm >= HIF8_NAN, -1.0, 1.0) + + # 提取 dot 档位 + dot_4_value = (x_norm & HIF8_DOT_MASK) >> 3 + + # 初始化指数和尾数乘子 + exponent_value = torch.zeros_like(x_norm, dtype=torch.float32) + m_value = torch.zeros_like(x_norm, dtype=torch.float32) + + # --- 档位 D4 --- + m_d4 = dot_4_value >= HIF8_DOT_D4 + if m_d4.any(): + exp_int = (x_norm & HIF8_EXP_MASK_D4) >> 1 + exponent_value = torch.where( + m_d4, + torch.where(exp_int >= 8, -exp_int, exp_int + 8).float(), + exponent_value, + ) + m_value = torch.where(m_d4, 1.0 + (x_norm & HIF8_FRAC_MASK_1BIT) * 0.5, m_value) + + # --- 档位 D3 --- + m_d3 = (~m_d4) & (dot_4_value >= HIF8_DOT_D3) + if m_d3.any(): + exp_int = (x_norm & HIF8_EXP_MASK_D3) >> 2 + exponent_value = torch.where( + m_d3, + torch.where(exp_int >= 4, -exp_int, exp_int + 4).float(), + exponent_value, + ) + m_value = torch.where(m_d3, 1.0 + (x_norm & HIF8_FRAC_MASK_2BIT) * 0.25, m_value) + + # --- 档位 D2 --- + m_d2 = (~(m_d4 | m_d3)) & (dot_4_value >= HIF8_DOT_D2) + if m_d2.any(): + exp_int = (x_norm & HIF8_EXP_MASK_D2) >> 3 + exponent_value = torch.where( + m_d2, + torch.where(exp_int >= 2, -exp_int, exp_int + 2).float(), + exponent_value, + ) + m_value = torch.where(m_d2, 1.0 + (x_norm & HIF8_FRAC_MASK_3BIT) * 0.125, m_value) + + # --- 档位 D1 --- + m_d1 = (~(m_d4 | m_d3 | m_d2)) & (dot_4_value >= HIF8_DOT_D1) + if m_d1.any(): + exp_sign = (x_norm & HIF8_EXP_SIGN_MASK_D1) >> 3 + exponent_value = torch.where(m_d1, torch.where(exp_sign >= 1, -1.0, 1.0), exponent_value) + m_value = torch.where(m_d1, 1.0 + (x_norm & HIF8_FRAC_MASK_3BIT) * 0.125, m_value) + + # --- 档位 D0 --- + m_d0 = dot_4_value == HIF8_DOT_D0 + if m_d0.any(): + exponent_value = torch.where(m_d0, 0.0, exponent_value) + m_value = torch.where(m_d0, 1.0 + (x_norm & HIF8_FRAC_MASK_3BIT) * 0.125, m_value) + + # --- 档位 DML --- + m_dml = dot_4_value == HIF8_DOT_DML + if m_dml.any(): + exponent_value = torch.where( + m_dml, + ((x_norm & HIF8_EXP_MASK_DML) - HIF8_DML_EXP_OFFSET).float(), + exponent_value, + ) + m_value = torch.where(m_dml, 1.0, m_value) + + # 计算正常值结果并写回 + norm_res = sign * torch.pow(2.0, exponent_value) * m_value + out[mask_normal] = norm_res + + return out.reshape(shape) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/result_compare_method.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/result_compare_method.py new file mode 100644 index 000000000000..865b79390239 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/result_compare_method.py @@ -0,0 +1,837 @@ +#!/usr/bin/python +# ruff: noqa +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import ast +import datetime +import logging +import os +import sys +from time import time + +import numpy as np +import torch + +logging.basicConfig(level=logging.INFO, format="%(message)s", force=True) +logger = logging.getLogger(__name__) + + +def cal_relative_diff_np_isclose(real_data, expect_data, type_str="fp16"): + diff = abs(float(real_data) - float(expect_data)) + result = diff / (np.abs(expect_data) + 10e-10) + return result + + +def print_log(data=None, level="INFO"): + print( + "[%s] [%s]-%s:%s - %s" + % ( + datetime.datetime.now().strftime("%Y/%m/%d %H:%M:%S"), + level, + os.path.basename(sys._getframe().f_back.f_code.co_filename), + str(sys._getframe().f_back.f_lineno).zfill(4), + data, + ) + ) + + +def display_error_output(real_data, expect_data, err_idx, relative_diff): + print_log("Error Line-----------------------------------------------------------------------------") + print_log("Loop \t ExpectOut \t RealOut \t FpDiff \t RateDiff") + print_log("---------------------------------------------------------------------------------------") + count = 0 + len_err = len(err_idx) + for i in err_idx: + count += 1 + if count < 10 or (90 < count < 100): + print_log( + "%08d \t %.7f \t %.7f \t %.7f \t %.7f" + % ( + i, + expect_data[i], + real_data[i], + abs(np.float64(expect_data[i]) - np.float64(real_data[i])), + relative_diff[count - 1], + ) + ) + elif count == 10 or (count == 100 and len_err > 100): + dot_3 = "..." + print_log("%08s \t %07s \t %07s \t %07s \t %07s" % (dot_3, dot_3, dot_3, dot_3, dot_3)) + elif count > 100: + break + + print_log("Max-RE line:---------------------------------------------------------------------------") + max_error = max(relative_diff) + m_idx_list = err_idx[np.where(relative_diff == max_error)] + m_count = 0 + for m_idx in m_idx_list: + m_count += 1 + if m_count < 4: + print_log( + "%08d \t %.7f \t %.7f \t %.7f \t %.7f" + % ( + m_idx, + expect_data[m_idx], + real_data[m_idx], + abs(np.float64(expect_data[m_idx]) - np.float64(real_data[m_idx])), + max_error, + ) + ) + else: + break + print_log("---------------------------------------------------------------------------------------") + + +def display_output_np_isclose(real_data, expect_data, start, end, expect_fp32_data=None): + def display_inner(idx): + j = idx + start + diff_rate = cal_relative_diff_np_isclose(real_data[j], expect_data[j]) + if "inf" in str(expect_data[j]) or "nan" in str(expect_data[j]): + diff_abs = "inf" if "inf" in str(expect_data[j]) else "nan" + if expect_fp32_data is not None: + print_log( + "%08d \t %-7s \t %-7s \t %-7s \t %-7s \t %-7s" + % ( + start + idx + 1, + expect_fp32_data[j], + expect_data[j], + real_data[j], + diff_abs, + diff_rate, + ) + ) + else: + print_log( + "%08d \t %-7s \t %-7s \t %-7s \t %-7s" + % ( + start + idx + 1, + expect_data[j], + real_data[j], + diff_abs, + diff_rate, + ) + ) + else: + diff_abs = abs(np.float64(expect_data[j]) - np.float64(real_data[j])) + if expect_fp32_data is not None: + print_log( + "%08d \t %0.7f \t %0.7f \t %0.7f \t %0.7f \t %0.7f" + % ( + start + idx + 1, + expect_fp32_data[j], + expect_data[j], + real_data[j], + diff_abs, + diff_rate, + ) + ) + else: + print_log( + "%08d \t %0.7f \t %0.7f \t %0.7f \t %0.7f" + % ( + start + idx + 1, + expect_data[j], + real_data[j], + diff_abs, + diff_rate, + ) + ) + + print_log("---------------------------------------------------------------------------------------") + if expect_fp32_data is not None: + print_log("Loop \t ExpFP32Out \t ExpFP16Out \t NPUOut \tFpDiff(min) \t RateDiff") + else: + print_log("Loop \t ExpectOut \t RealOut \t FpDiff \t RateDiff") + print_log("---------------------------------------------------------------------------------------") + split_count = int(end - start) + if split_count <= 20: + for i in range(split_count + 1): + display_inner(i) + else: + for i in range(10): + display_inner(i) + print_log("... \t ... \t ... \t ... \t ...") + for i in range(split_count - 10 + 1, split_count + 1): + display_inner(i) + + +def find_batch_and_position(cu_seqlens, x): + """ + 判断x属于哪个batch以及在该batch中的位置 + + 参数: + cu_seqlens: 前缀和列表, cu_seqlens[b_idx]表示前(b_idx)个batch的总长度 + x: 需要判断的数值 + + 返回: + tuple: (batch_idx, position) + - batch_idx: 所属的batch索引(从0开始),超出范围则为-1 + - position: 在该batch中的位置(从0开始), 超出范围则为-1 + """ + if not cu_seqlens: + return (-1, -1) + # 遍历前缀和列表查找所属批次 + for batch_idx in range(len(cu_seqlens) - 1): + # 计算当前批次的起始位置 + start = cu_seqlens[batch_idx] + # 判断是否在当前批次范围内 + if start <= x < cu_seqlens[batch_idx + 1]: + # 计算在当前批次中的位置(偏移量) + position = x - start + return (batch_idx, position) + # 超出所有批次范围 + return (-1, -1) + + +def judge_value_by_isclose(real_data, data_compe, force_bf16=False): + atol = 2.5e-05 + rtol = 0.005 + pct_thd = 0.005 + diff_thd = 0.005 + # force_bf16: QLIV2 的 returnValue 固定为 bf16,但流程中已被 .float() 转成 float32, + # 无法通过 dtype 判断,需强制按 bf16 门限对比。 + is_bfloat16 = force_bf16 or (str(real_data.dtype) in ("bfloat16", "torch.bfloat16")) + if isinstance(real_data, torch.Tensor): + real_data = real_data.detach().cpu().float().numpy() + else: + real_data = np.asarray(real_data) + if isinstance(data_compe, torch.Tensor): + data_compe = data_compe.detach().cpu().float().numpy() + else: + data_compe = np.asarray(data_compe) + start = 0 + end = real_data.size - 1 + if end < start: + end = start + split_count = int(end - start + 1) if end != start else 1 + + if is_bfloat16: + # bf16 尾数位少、舍入误差大,误差门限放宽到 1/128(约 0.0078125) + atol = 0.0001 + rtol = 1.0 / 128 + diff_thd = 1.0 / 128 + diff_result = np.isclose( + real_data.astype(np.float32), + data_compe.astype(np.float32), + rtol=rtol, + atol=atol, + equal_nan=True, + ) + else: + diff_result = np.isclose(real_data, data_compe, rtol=rtol, atol=atol, equal_nan=True) + err_idx = np.where(diff_result != np.array((True,)))[0] + diff_abs = abs(data_compe - real_data) + b1 = np.maximum(np.abs(real_data), (np.abs(data_compe))) + b2 = float((1.0 / (1 << 14)) / diff_thd) + b = np.add(np.maximum(b1, b2), 10e-10) + eps = 10e-10 + err_diff = diff_abs / (b + eps) + err_diff = err_diff[err_idx] + fulfill_percent = float(split_count - err_idx.size) / float(split_count) * 100.0 + pct_thd = (1 - pct_thd) * 100.0 + result = True if (fulfill_percent >= pct_thd) else False + return result + + +def compare_topk_valid( + cur_cpu, + cur_npu, + topk_value, + bsn, + diff_npu, + diff_cpu, + cur_npu_output_value=None, + cur_cpu_output_value=None, + thres=0.001, + return_value_flag=False, + output_idx_offset=None, + layout_query=None, + cu_seqlens_q=None, + q_seq=0, +): + b_idx, s1_idx, n2_idx = bsn + max_re = 0.0 + npu_pass = True + cur_cpu = np.asarray(cur_cpu, dtype=np.int64) + cur_npu = np.asarray(cur_npu, dtype=np.int64) + + if output_idx_offset is not None: + # 统一转换后使用 + offset_data = ( + output_idx_offset.cpu().numpy() + if hasattr(output_idx_offset, "device") and output_idx_offset.device.type != "cpu" + else np.array(output_idx_offset) + ) + offset_flat = offset_data.flatten() + if layout_query == "TND": + cur_prefix = cu_seqlens_q[b_idx] + offset = offset_flat[cur_prefix + s1_idx] + else: + offset = offset_flat[b_idx * q_seq + s1_idx] + cpu_offset_mask = cur_cpu != -1 + npu_offset_mask = cur_npu != -1 + cur_cpu = np.where(cpu_offset_mask, cur_cpu - offset, cur_cpu) + cur_npu = np.where(npu_offset_mask, cur_npu - offset, cur_npu) + + element_list = topk_value[b_idx, n2_idx, s1_idx, :] + score_size = element_list.shape[-1] + invalid_cpu = (cur_cpu < 0) | (cur_cpu >= score_size) + invalid_npu = (cur_npu < 0) | (cur_npu >= score_size) + has_duplicate_cpu = np.unique(cur_cpu).size != cur_cpu.size + has_duplicate_npu = np.unique(cur_npu).size != cur_npu.size + if ( + cur_cpu.size != cur_npu.size + or np.any(invalid_cpu) + or np.any(invalid_npu) + or has_duplicate_cpu + or has_duplicate_npu + ): + diff_cpu.append(cur_cpu.tolist()) + diff_npu.append(cur_npu.tolist()) + return False, float("inf") + + npu_set = set(cur_npu) + cpu_set = set(cur_cpu) + if npu_set != cpu_set: + value_bm = topk_value[b_idx, n2_idx, s1_idx, cur_cpu[-1]] + only_in_npu = npu_set - cpu_set + only_in_cpu = cpu_set - npu_set + only_in_npu_list = list(only_in_npu) + only_in_cpu_list = list(only_in_cpu) + for diff_idx in range(len(only_in_npu_list)): + element_npu = element_list[only_in_npu_list[diff_idx]] + element_cpu = element_list[only_in_cpu_list[diff_idx]] + npu_ae = abs(element_npu - value_bm) + cpu_ae = abs(element_cpu - value_bm) + if value_bm == 0: + if npu_ae == 0: + npu_re = 0.0 + else: + npu_re = float("inf") + if cpu_ae == 0: + cpu_re = 0.0 + else: + cpu_re = float("inf") + else: + npu_re = abs(npu_ae / value_bm) + cpu_re = abs(cpu_ae / value_bm) + if npu_re > thres or cpu_re > thres: + if return_value_flag: + # 将 value 输出统一转为 numpy array,bfloat16 需先转 float32 再转 numpy + if torch.is_tensor(cur_npu_output_value): + npuValueArr = ( + cur_npu_output_value.float().cpu().numpy() + if cur_npu_output_value.dtype == torch.bfloat16 + else cur_npu_output_value.cpu().numpy() + ) + else: + npuValueArr = np.asarray(cur_npu_output_value) + if torch.is_tensor(cur_cpu_output_value): + cpuValueArr = ( + cur_cpu_output_value.float().cpu().numpy() + if cur_cpu_output_value.dtype == torch.bfloat16 + else cur_cpu_output_value.cpu().numpy() + ) + else: + cpuValueArr = np.asarray(cur_cpu_output_value) + if not judge_value_by_isclose(npuValueArr, cpuValueArr): + npu_pass = False + diff_npu.append(element_npu) + diff_cpu.append(element_cpu) + max_re = max(max_re, npu_re, cpu_re) + else: + npu_pass = False + diff_npu.append(element_npu) + diff_cpu.append(element_cpu) + max_re = max(max_re, npu_re) + return npu_pass, max_re + + +def compare_return_value(cur_npu_output_value=None, cur_cpu_output_value=None): + max_re = 0.0 + npu_pass = True + npu_pass = judge_value_by_isclose(cur_npu_output_value, cur_cpu_output_value) + return npu_pass, max_re + + +def trans_tnd_actseq(list): + list_len = len(list) + if list_len == 0: + raise ValueError("TND情况下 act_seq需要必传") + list_new = [] + list_new.append(list[0]) + for i in range(list_len - 1): + new_item = list[i + 1] - list[i] + if new_item >= 0: + list_new.append(new_item) + else: + raise ValueError(f"TND情况下 act_seq_len 为非递减数列 act_seq_len={list}") + return list_new + + +def _reshape_topk_value(topk_value, total_rows, sparse_count, params): + if isinstance(topk_value, torch.Tensor): + topk_value = topk_value.detach().cpu().float().numpy() + else: + topk_value = np.asarray(topk_value) + if topk_value.size == total_rows * sparse_count: + return topk_value.reshape(total_rows, sparse_count) + + batch_size, cu_seqlens_q, layout_query = params[0], params[14], params[21] + if layout_query != "TND" or topk_value.ndim != 4: + raise ValueError(f"topk value shape {topk_value.shape} cannot reshape to ({total_rows}, {sparse_count})") + cu_seqlens_q = _get_tnd_query_prefix(cu_seqlens_q, batch_size) + topk_value = np.concatenate( + [ + topk_value[batch_idx, :, : cu_seqlens_q[batch_idx + 1] - cu_seqlens_q[batch_idx], :] + .transpose(1, 0, 2) + .reshape(-1, sparse_count) + for batch_idx in range(batch_size) + ] + ) + return topk_value.reshape(total_rows, sparse_count) + + +def check_result( + expect, + result, + topk_value, + output_idx_offset, + params, + cpu_topk_value, + npu_topk_value, +): + ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + weight_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + _, + ) = params + + # Q 侧个体长度 + if layout_query == "TND": + # TND: 必传 cu_seqlens_q,从差分推导个体长度 + if isinstance(cu_seqlens_q, str): + lengths_q_list = ast.literal_eval(cu_seqlens_q) + else: + lengths_q_list = cu_seqlens_q[1:] + else: + # BSND: 从 seqused_q 获取,若 None 则用 q_seq 填满 + if seqused_q is not None: + if isinstance(seqused_q, str): + lengths_q_list = ast.literal_eval(seqused_q) + else: + lengths_q_list = list(seqused_q) + else: + lengths_q_list = [q_seq] * batch_size + + # K 侧个体长度 + if layout_key == "TND": + # TND: 必传 cu_seqlens_k,从差分推导个体长度 + if isinstance(cu_seqlens_k, str): + lengths_k_list = ast.literal_eval(cu_seqlens_k) + else: + lengths_k_list = cu_seqlens_k[1:] + elif layout_key == "PA_BBND": + # PA_BBND: 从 seqused_k 获取 + assert seqused_k is not None, f"{layout_key} layout requires seqused_k" + if isinstance(seqused_k, str): + lengths_k_list = ast.literal_eval(seqused_k) + else: + lengths_k_list = list(seqused_k) + else: + # BSND: 从 seqused_k 获取,若 None 则用 q_seq 填满 + if seqused_k is not None: + if isinstance(seqused_k, str): + lengths_k_list = ast.literal_eval(seqused_k) + else: + lengths_k_list = list(seqused_k) + else: + lengths_k_list = [k_seq] * batch_size + + act_seq_q = lengths_q_list + act_seq_k = lengths_k_list + + if isinstance(act_seq_q, int): + act_seq_q = [act_seq_q] + elif isinstance(act_seq_q, list): + act_seq_q = act_seq_q + else: + act_seq_q = ast.literal_eval(act_seq_q) + if isinstance(act_seq_k, int): + act_seq_k = [act_seq_k] + elif isinstance(act_seq_k, list): + act_seq_k = act_seq_k + else: + act_seq_k = ast.literal_eval(act_seq_k) + + if isinstance(cu_seqlens_q, int): + cu_seqlens_q = [cu_seqlens_q] + elif isinstance(cu_seqlens_q, list): + cu_seqlens_q = cu_seqlens_q + elif cu_seqlens_q is not None: + cu_seqlens_q = ast.literal_eval(cu_seqlens_q) + + if isinstance(cu_seqlens_k, int): + cu_seqlens_k = [cu_seqlens_k] + elif isinstance(cu_seqlens_k, list): + cu_seqlens_k = cu_seqlens_k + elif cu_seqlens_k is not None: + cu_seqlens_k = ast.literal_eval(cu_seqlens_k) + + if isinstance(seqused_q, int): + seqused_q = [seqused_q] + elif isinstance(seqused_q, list): + seqused_q = seqused_q + elif seqused_q is not None: + seqused_q = ast.literal_eval(seqused_q) + + if isinstance(seqused_k, int): + seqused_k = [seqused_k] + elif isinstance(seqused_k, list): + seqused_k = seqused_k + elif seqused_k is not None: + seqused_k = ast.literal_eval(seqused_k) + npu_pass = True + max_error = 0 + max_re = 0 + thres = 0.001 + diff_thd = 0.01 + pct_thd = 0.005 + max_diff_hd = 0.1 + rtol = 0.005 + atol = 0.000025 + max_error_idx = 10000000 + cpu_output = expect.cpu().numpy() + npu_output = result.cpu().numpy() + real_data = result.cpu().numpy() + data_compe = expect.cpu().numpy() + real_data = npu_output.flatten() + data_compe = cpu_output.flatten() + diff_cpu = [] + diff_npu = [] + + if layout_query in ["BSND"]: + sp = (batch_size, q_seq, k_head_num) + total_rows = batch_size * q_seq * k_head_num + elif layout_query in ["TND"]: + sp = (q_t_size, k_head_num) + total_rows = q_t_size * k_head_num + else: + total_rows = 0 + sp = (0, 0) + print(f"total_line is {total_rows}") + npu_reshape = npu_output.reshape([total_rows, sparse_count]) + cpu_reshape = cpu_output.reshape([total_rows, sparse_count]) + if return_value: + cpu_topk_value = _reshape_topk_value(cpu_topk_value, total_rows, sparse_count, params) + npu_topk_value = _reshape_topk_value(npu_topk_value, total_rows, sparse_count, params) + start_time = time() + invalid_data = cpu_reshape != -1 + valid_lens = invalid_data.sum(axis=-1) # (total_rows,) + # 判断有效值部分集合是否相同 + cpu_output_sorted = np.sort(cpu_reshape, axis=1) + npu_output_sorted = np.sort(npu_reshape, axis=1) + diff_rows = np.zeros(total_rows, dtype=bool) + diff_rows |= np.any(cpu_output_sorted != npu_output_sorted, axis=1) # 标记存在差异的行 + test_id = [] + rows = [] + if np.any(diff_rows): + rows = np.where(diff_rows)[0] + num_rows = len(rows) + if num_rows: + print(f"需要进行第二步比较的batch有{num_rows}") + else: + print("有效值集合相同,无需进行比较") + for t_id in rows: + bsn = np.unravel_index(t_id, sp) + npu_topk_output_value = None + cpu_topk_output_value = None + if layout_query == "TND": + b_idx, s1_idx = find_batch_and_position(cu_seqlens_q, bsn[0]) + bsn = (b_idx, s1_idx, bsn[-1]) + if return_value: + cpu_topk_output_value = cpu_topk_value[t_id, :] + npu_topk_output_value = npu_topk_value[t_id, :] + npu_pass_t = True + max_re_t = 0 + valid_len = valid_lens[t_id] + npu_pass_t, max_re_t = compare_topk_valid( + cpu_reshape[t_id, :valid_len], + npu_reshape[t_id, :valid_len], + topk_value, + bsn, + diff_npu, + diff_cpu, + npu_topk_output_value, + cpu_topk_output_value, + thres, + return_value, + output_idx_offset, + layout_query, + cu_seqlens_q, + q_seq, + ) + if not npu_pass_t: + npu_pass = False + end_time = time() + print(f"耗时:{end_time - start_time:.6f} 秒") + topk_precision = not diff_npu and not diff_cpu + if topk_precision: + print("[success]TopK精度通过, idx不同的地方的value误差在阈值之内") + else: + print("[fail]TopK精度失败") + print(f"npu_pass is {npu_pass}") + if real_data.size == 0 and real_data.size == data_compe.size: + print_log('The npu_output is [],and it is same as bm_output, the result of data_compare is "Pass"') + return "Pass", 100.0, 0 + start = 0 + end = real_data.size - 1 + if end < start: + end = start + diff_result = np.isclose(real_data, data_compe, rtol=rtol, atol=atol, equal_nan=True) + err_idx = np.where(diff_result != np.array((True,)))[0] + diff_abs = abs(data_compe - real_data) + b1 = np.maximum(np.abs(real_data), (np.abs(data_compe))) + b2 = float((1.0 / (1 << 14)) / diff_thd) + b = np.add(np.maximum(b1, b2), 10e-10) + eps = 10e-10 + err_diff = diff_abs / (b + eps) + err_diff = err_diff[err_idx] + split_count = int(end - start + 1) if end != start else 1 + print_log("split_count:%s; max_diff_hd:%s;" % (float(split_count), max_diff_hd)) + fulfill_percent = float(split_count - err_idx.size) / float(split_count) * 100.0 + display_output_np_isclose(real_data, data_compe, start, end) + pct_thd = (1 - pct_thd) * 100.0 + result = "Pass" if (npu_pass or topk_precision) else "Failed" + print_log("---------------------------------------------------------------------------------------") + print_log("Rtol \t Atol \t PctThd \t PctRlt \t Result") + print_log("---------------------------------------------------------------------------------------") + print_log("%.4f \t %.6f \t %.2f%% \t %.6f%% \t %s" % (rtol, atol, pct_thd, fulfill_percent, result)) + if len(err_diff) > 0: + print_log("Max-RelativeError is: %s. Threshold is: %s." % (max_error, max_diff_hd)) + if result == "Failed": + display_error_output(real_data, data_compe, err_idx, err_diff[0:max_error_idx]) + return result, fulfill_percent + + +def _to_flat_numpy(value, dtype=np.int64): + if value is None: + return None + if isinstance(value, str): + value = ast.literal_eval(value) + if isinstance(value, torch.Tensor): + value = value.detach().cpu().numpy() + return np.asarray(value, dtype=dtype).reshape(-1) + + +def _get_tnd_query_prefix(cu_seqlens_q, batch_size): + prefix = _to_flat_numpy(cu_seqlens_q) + if prefix is None: + raise ValueError("cu_seqlens_q is required") + if prefix.size == batch_size + 1 and prefix[0] == 0: + return prefix + if prefix.size == batch_size: + return np.concatenate((np.array([0], dtype=prefix.dtype), prefix)) + raise ValueError(f"invalid TND cu_seqlens_q length: {prefix.size}") + + +def _gather_return_values_by_index(topk_value, result_indices, params, output_idx_offset): + batch_size = params[0] + q_seq = params[1] + q_t_size = params[3] + k_head_num = params[6] + cu_seqlens_q = params[14] + layout_query = params[21] + sparse_count = params[23] + + full_score = topk_value.detach().cpu().float().numpy() + npu_indices = result_indices.detach().cpu().numpy().reshape(-1, sparse_count) + expected = np.full(npu_indices.shape, -np.inf, dtype=np.float32) + invalid_index = np.zeros(npu_indices.shape, dtype=bool) + offsets = _to_flat_numpy(output_idx_offset) + query_prefix = _get_tnd_query_prefix(cu_seqlens_q, batch_size) if layout_query == "TND" else None + + for row_idx in range(npu_indices.shape[0]): + if layout_query == "BSND": + b_idx, s1_idx, n2_idx = np.unravel_index(row_idx, (batch_size, q_seq, k_head_num)) + offset_pos = b_idx * q_seq + s1_idx + elif layout_query == "TND": + t_idx, n2_idx = np.unravel_index(row_idx, (q_t_size, k_head_num)) + b_idx = int(np.searchsorted(query_prefix[1:], t_idx, side="right")) + s1_idx = int(t_idx - query_prefix[b_idx]) + offset_pos = t_idx + else: + raise ValueError(f"unsupported query layout: {layout_query}") + + row_indices = npu_indices[row_idx] + logical_indices = row_indices.astype(np.int64, copy=True) + if offsets is not None: + logical_indices[row_indices >= 0] -= int(offsets[offset_pos]) + row_score = full_score[b_idx, n2_idx, s1_idx] + valid = (row_indices >= 0) & (logical_indices >= 0) & (logical_indices < row_score.shape[-1]) + expected[row_idx, valid] = row_score[logical_indices[valid]] + invalid_index[row_idx] = (row_indices < -1) | ((row_indices >= 0) & ~valid) + + return expected, invalid_index + + +def check_result_return_value( + expect, + result, + params, + expect_indices=None, + result_indices=None, + topk_value=None, + output_idx_offset=None, +): + ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + weight_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + _, + ) = params + + npu_pass = True + max_error = 0 + max_re = 0 + thres = 0.0001 + diff_thd = 0.01 + pct_thd = 0.005 + max_diff_hd = 0.1 + rtol = 0.005 + atol = 0.000025 + max_error_idx = 10000000 + npu_output = result.cpu().float().numpy() + if topk_value is not None and result_indices is not None: + if output_idx_offset is None: + output_idx_offset = params[32] + cpu_output, invalid_index = _gather_return_values_by_index( + topk_value, result_indices, params, output_idx_offset + ) + else: + cpu_output = expect.cpu().float().numpy() + invalid_index = np.zeros(cpu_output.shape, dtype=bool) + real_data = npu_output.flatten() + data_compe = cpu_output.flatten() + + if layout_query in ["BSND"]: + sp = (batch_size, q_seq, k_head_num) + total_rows = batch_size * q_seq * k_head_num + elif layout_query in ["TND"]: + sp = (q_t_size, k_head_num) + total_rows = q_t_size * k_head_num + else: + total_rows = 0 + sp = (0, 0) + print(f"total_line is {total_rows}") + npu_reshape = npu_output.reshape([total_rows, sparse_count]) + cpu_reshape = cpu_output.reshape([total_rows, sparse_count]) + invalid_index_reshape = invalid_index.reshape([total_rows, sparse_count]) + + start_time = time() + + # QLIV2 returnValue 为 bf16,强制使用 bf16 门限(误差阈值 1/128) + npu_pass = judge_value_by_isclose(npu_reshape, cpu_reshape, force_bf16=True) + if np.any(invalid_index_reshape): + npu_pass = False + end_time = time() + print(f"耗时:{end_time - start_time:.6f} 秒") + print(f"npu_pass is {npu_pass}") + if real_data.size == 0 and real_data.size == data_compe.size: + print_log('The npu_output is [],and it is same as bm_output, the result of data_compare is "Pass"') + return "Pass", 100.0, 0 + start = 0 + end = real_data.size - 1 + if end < start: + end = start + diff_result = np.isclose(real_data, data_compe, rtol=rtol, atol=atol, equal_nan=True) + err_idx = np.where(diff_result != np.array((True,)))[0] + diff_abs = abs(data_compe - real_data) + b1 = np.maximum(np.abs(real_data), (np.abs(data_compe))) + b2 = float((1.0 / (1 << 14)) / diff_thd) + b = np.add(np.maximum(b1, b2), 10e-10) + eps = 10e-10 + err_diff = diff_abs / (b + eps) + err_diff = err_diff[err_idx] + split_count = int(end - start + 1) if end != start else 1 + print_log("split_count:%s; max_diff_hd:%s;" % (float(split_count), max_diff_hd)) + fulfill_percent = float(split_count - err_idx.size) / float(split_count) * 100.0 + display_output_np_isclose(real_data, data_compe, start, end) + pct_thd = (1 - pct_thd) * 100.0 + result = "Pass" if npu_pass else "Failed" + print_log("---------------------------------------------------------------------------------------") + print_log("Rtol \t Atol \t PctThd \t PctRlt \t Result") + print_log("---------------------------------------------------------------------------------------") + print_log("%.4f \t %.6f \t %.2f%% \t %.6f%% \t %s" % (rtol, atol, pct_thd, fulfill_percent, result)) + if len(err_diff) > 0: + print_log("Max-RelativeError is: %s. Threshold is: %s." % (max_error, max_diff_hd)) + if result == "Failed": + display_error_output(real_data, data_compe, err_idx, err_diff[0:max_error_idx]) + return result, fulfill_percent diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_qliv2_test_utils.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_qliv2_test_utils.py new file mode 100644 index 000000000000..239b48aa021f --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_qliv2_test_utils.py @@ -0,0 +1,98 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +from multiprocessing.reduction import ForkingPickler +from pathlib import Path + +import pandas as pd +import pytest +from qliv2_parameter_normalization import normalize_qliv2_params +from qliv2_test_utils import ( + PARAM_NAMES, + QliV2CaseSelector, + QliV2ResultWriter, + ensure_comparison_passed, +) + + +def test_case_selector_preserves_requested_order(tmp_path): + for name in ("case_10.pt", "case_2.pt", "case_1.pt"): + (tmp_path / name).touch() + + natural = QliV2CaseSelector.resolve(tmp_path) + assert [Path(path).name for path in natural] == ["case_1.pt", "case_2.pt", "case_10.pt"] + + indexed = QliV2CaseSelector.resolve(tmp_path, case_indexes="3,1,2-2") + assert [Path(path).name for path in indexed] == ["case_10.pt", "case_1.pt", "case_2.pt"] + + named = QliV2CaseSelector.resolve(tmp_path, case_names="case_2,case_1.pt") + assert [Path(path).name for path in named] == ["case_2.pt", "case_1.pt"] + + +def test_case_selector_rejects_ambiguous_or_invalid_selection(tmp_path): + (tmp_path / "case_1.pt").touch() + with pytest.raises(ValueError, match="cannot be specified together"): + QliV2CaseSelector.resolve(tmp_path, case_names="case_1", case_indexes="1") + with pytest.raises(ValueError, match="out of range"): + QliV2CaseSelector.resolve(tmp_path, case_indexes="2") + + +def test_result_writer_uses_readable_name_and_migrates_legacy_result(tmp_path): + params = list(range(len(PARAM_NAMES))) + params[PARAM_NAMES.index("qk_dtype")] = "INT8" + params[PARAM_NAMES.index("layout_query")] = "BSND" + params[PARAM_NAMES.index("layout_key")] = "PA_BBND" + name = QliV2ResultWriter.case_name(params) + assert name == QliV2ResultWriter.case_name(params) + assert name == ("QLI_B0_S11_S22_N15_N26_D7_BSND_PA_BBND_INT8_QM20_SM24_CR30_K23_RV31") + assert ( + QliV2ResultWriter.case_name( + params, + explicit_name="quant li/default:a5 v2", + ) + == "quant_li_default_a5_v2" + ) + + row = QliV2ResultWriter.row(name, params, "Pass", 100.0) + output = tmp_path / "result.xlsx" + legacy = pd.DataFrame([{key: value for key, value in row.items() if key != "return_value"}]) + legacy.to_excel(output, index=False) + + QliV2ResultWriter.append(output, row) + result = pd.read_excel(output) + assert list(result.columns) == list(row.keys()) + assert len(result) == 2 + assert result.iloc[1]["return_value"] == params[PARAM_NAMES.index("return_value")] + + +def test_comparison_failure_raises_serializable_assertion(): + ensure_comparison_passed("case_pass", "Pass", 100.0) + + with pytest.raises(AssertionError, match="case_index_fail.*index result=Failed") as caught: + ensure_comparison_passed("case_index_fail", "Failed", 97.5) + restored = ForkingPickler.loads(ForkingPickler.dumps(caught.value)) + assert str(restored) == str(caught.value) + + with pytest.raises(AssertionError, match="case_value_fail.*value result=Failed"): + ensure_comparison_passed("case_value_fail", "Pass", 100.0, "Failed", 90.0) + + +def test_normalize_legacy_params_adds_weight_dtype(): + params = list(range(32)) + params[10] = "INT8" + params[11] = "FP16" + + normalized = normalize_qliv2_params(params) + + assert len(normalized) == 33 + assert normalized[:11] == tuple(params[:11]) + assert normalized[11] == params[11] + assert normalized[12:] == tuple(params[11:]) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_batch.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_batch.py new file mode 100644 index 000000000000..eda1681e0fab --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_batch.py @@ -0,0 +1,182 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import concurrent.futures +import os +from pathlib import Path + +import pytest +import result_compare_method +from batch import quant_lightning_indexer_v2_pt_loadprocess +from qliv2_test_utils import ( + QliV2CaseSelector, + QliV2ResultWriter, + ensure_comparison_passed, +) + +TEST_INPUT_PATH_ENV = os.environ.get("QLIV2_TESTCASE_DIR", "").strip() +TEST_INPUT_PATH = TEST_INPUT_PATH_ENV or "pt_path" +RESULT_PATH = Path(os.environ.get("QLIV2_RESULT_PATH", "result.xlsx").strip()) +DEVICE_ID = int(os.environ.get("QLIV2_DEVICE_ID", "0")) + +# 支持通过环境变量 QLIV2_TESTCASE_PATH 指定单条用例文件,实现进程级隔离执行: +# - 设置时:仅运行该条用例(配合 batch_isolated_run.sh 每条用例拉起独立进程) +# - 未设置:回退为原有行为,一次性加载目录下全部用例 +SINGLE_CASE_PATH = os.environ.get("QLIV2_TESTCASE_PATH", "").strip() +# flag:是否处于批量隔离模式(由 batch_isolated_run.sh 设置 QLIV2_TESTCASE_PATH 触发) +IS_ISOLATED_MODE = bool(SINGLE_CASE_PATH) +# flag:运行模式 eager / graph(通过环境变量 QLIV2_RUN_MODE 或命令行参数设置,默认 eager) +RUN_MODE = os.environ.get("QLIV2_RUN_MODE", "eager").strip().lower() +# 支持通过环境变量 QLIV2_PT_FILE_LIST 指定用例文件列表(逗号分隔),用于从 Excel 筛选的 batch_exec 模式 +PT_FILE_LIST = os.environ.get("QLIV2_PT_FILE_LIST", "").strip() +CASE_NAMES = os.environ.get("QLIV2_CASE_NAMES", "").strip() +CASE_INDEXES = os.environ.get("QLIV2_CASE_INDEXES", "").strip() + +try: + if SINGLE_CASE_PATH: + TESTCASE_FILES = QliV2CaseSelector.resolve(TEST_INPUT_PATH, explicit_files=SINGLE_CASE_PATH) + print(f"单用例隔离模式, 仅执行: {SINGLE_CASE_PATH}") + else: + TESTCASE_FILES = QliV2CaseSelector.resolve( + TEST_INPUT_PATH, + explicit_files=PT_FILE_LIST, + case_names=CASE_NAMES, + case_indexes=CASE_INDEXES, + ) + print(f"找到 {len(TESTCASE_FILES)} 个测试用例文件") +except ValueError as error: + has_explicit_selection = any( + ( + TEST_INPUT_PATH_ENV, + SINGLE_CASE_PATH, + PT_FILE_LIST, + CASE_NAMES, + CASE_INDEXES, + ) + ) + if has_explicit_selection: + raise + print(f"未配置 batch PT 用例,跳过收集: {error}") + TESTCASE_FILES = [] + + +def qliv2(testcase_file): + try: + if RUN_MODE == "graph": + ( + cpu_result, + npu_result, + topk_value, + cpu_topk_value, + npu_topk_value, + output_idx_offset, + params, + ) = quant_lightning_indexer_v2_pt_loadprocess.test_qliv2_process_graph(testcase_file, device_id=DEVICE_ID) + else: + ( + cpu_result, + npu_result, + topk_value, + cpu_topk_value, + npu_topk_value, + output_idx_offset, + params, + ) = quant_lightning_indexer_v2_pt_loadprocess.test_qliv2_process(testcase_file, device_id=DEVICE_ID) + if npu_result is not None: + result, fulfill_percent = result_compare_method.check_result( + cpu_result, + npu_result, + topk_value, + output_idx_offset, + params, + cpu_topk_value, + npu_topk_value, + ) + else: + result = "Failed" + fulfill_percent = 0 + return_value = params[31] + if return_value: + result_return_value, fulfill_precent_return_value = result_compare_method.check_result_return_value( + cpu_topk_value, + npu_topk_value, + params, + cpu_result, + npu_result, + topk_value, + output_idx_offset, + ) + print(f"result_return_value: {result_return_value}") + print(f"fulfill_precent_return_value: {fulfill_precent_return_value}") + else: + result_return_value = "N/A" + fulfill_precent_return_value = 0 + except Exception as error: + print("NPU ERROR:", error) + result = "NPU ERROR" + fulfill_percent = 0 + result_return_value = "N/A" + fulfill_precent_return_value = 0 + params = [None] * 33 + + row_data = QliV2ResultWriter.row( + Path(testcase_file).stem, + params, + result, + fulfill_percent, + result_return_value, + fulfill_precent_return_value, + ) + QliV2ResultWriter.append(RESULT_PATH, row_data) + + case_name = Path(testcase_file).stem + if result != "NPU ERROR": + try: + ensure_comparison_passed( + case_name, + result, + fulfill_percent, + result_return_value, + fulfill_precent_return_value, + ) + except AssertionError as error: + return str(error) + + if result == "NPU ERROR": + return f"用例执行失败:{Path(testcase_file).stem}" + return None + + +@pytest.mark.ci +@pytest.mark.parametrize("testcase_file", TESTCASE_FILES) +def test_qliv2(testcase_file): + if IS_ISOLATED_MODE: + # 批量隔离模式:shell 层已通过独立 pytest 进程提供进程隔离,内部使用线程池即可 + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + futures = executor.submit(qliv2, testcase_file) + for future in concurrent.futures.as_completed([futures]): + try: + result = future.result() + if result is not None: + pytest.fail(str(result)) + except Exception as e: + pytest.fail(f"当前用例线程执行失败:{e}") + else: + # 非隔离模式(直接 pytest):使用子进程隔离,防止单条用例崩溃影响整体 + with concurrent.futures.ProcessPoolExecutor(max_workers=1) as executor: + future1 = executor.submit(qliv2, testcase_file) + for future in concurrent.futures.as_completed([future1]): + try: + result = future.result() + if result is not None: + pytest.fail(str(result)) + except Exception as e: + pytest.fail(f"当前用例子进程执行失败:{e}") diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_paramset.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_paramset.py new file mode 100644 index 000000000000..b1f832867d53 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_paramset.py @@ -0,0 +1,933 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import random + +import torch + +# 定义测试参数组合 +# 参数规则: +# cu_seqlens_q/k: TND 时必传 [B+1] 前缀和(首元素=0),非 TND 时为 None +# seqused_q/k: 每个 batch 的实际有效元素数 [B] +# TND 时可选(golden 可从 cu_seqlens 推导) +# BSND 时可选(None 则用 q_seq/k_seq 填满) +# PA_BBND 时 seqused_k 必传 +TEST_PARAMS = { + # Ascend950 基础场景: BSND query + PA_BBND key + "quant_li_default_a5": { + "batch_size": [8], + "q_seq": [15], + "k_seq": [111], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [512], # 取16的整数倍,最多支持到1024 + "block_num": [8], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [[3, 3, 3, 3, 3, 3, 3, 3]], + "seqused_k": [[28, 24, 80, 96, 47, 76, 0, 111]], # PA场景每个batch的实际token数 + "cmp_residual_k": [None], + "max_seqlen_q": [-1], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-448, 448]], + "key_datarange": [[-20, 20]], + "weights_datarange": [[-123, 123]], + "q_scale_datarange": [[0, 255]], + "k_scale_datarange": [[0, 65504]], + "cmp_ratio": [1], # 1/2/4/8/16/32/64/128 + "return_value": [0], + "output_idx_offset": [None], + "run_mode": ["eager"], + }, + # Ascend950 基础场景v2: BSND query + PA_BBND key + "quant_li_default_a5_v2": { + "batch_size": [104], + "q_seq": [4], + "k_seq": [32768], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], # 取16的整数倍,最多支持到1024 + "block_num": [53256], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [None], + "seqused_k": [[14224] * 103 + [32768]], # PA场景每个batch的实际token数 + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [0], + "query_datarange": [[2, 10]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "return_value": [1], + "output_idx_offset": [[random.randint(2, 10) for _ in range(416)] for _ in range(1)], + }, + # Ascend950 基础场景v3: BSND query + PA_BBND key + "quant_li_default_a5_v3": { + "batch_size": [26], + "q_seq": [4], + "k_seq": [262144], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], # 取16的整数倍,最多支持到1024 + "block_num": [106528], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [None], + "seqused_k": [[18244] * 25 + [262144]], # PA场景每个batch的实际token数 + "cmp_residual_k": [ + [ + 0, + 2, + 1, + 2, + 1, + 0, + 0, + 1, + 0, + 3, + 3, + 3, + 3, + 3, + 2, + 1, + 3, + 0, + 3, + 0, + 3, + 0, + 1, + 0, + 2, + 2, + ] + ], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 0.5]], + "k_scale_datarange": [[0, 0.5]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "return_value": [1], + "output_idx_offset": [[random.randint(0, 1000) for _ in range(104)] for _ in range(1)], + }, + # Ascend950 基础场景v4: BSND query + PA_BBND key + "quant_li_default_a5_v4": { + "batch_size": [56], + "q_seq": [4], + "k_seq": [2048], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], # 取16的整数倍,最多支持到1024 + "block_num": [1792], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [None], + "seqused_k": [[1804] * 55 + [2048]], # PA场景每个batch的实际token数 + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [1024], + "sparse_mode": [0], + "query_datarange": [[-1, 1]], + "key_datarange": [[0, 0.001]], + "weights_datarange": [[-20, 20]], + "q_scale_datarange": [[0, 5]], + "k_scale_datarange": [[0, 5]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "return_value": [1], + "output_idx_offset": [[random.randint(0, 100) for _ in range(224)] for _ in range(1)], + }, + # Ascend950 基础场景v5: BSND query + PA_BBND key + "quant_li_default_a5_v5": { + "batch_size": [42], + "q_seq": [4], + "k_seq": [32768], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], # 取16的整数倍,最多支持到1024 + "block_num": [21523], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [None], + "seqused_k": [[20912] * 41 + [32768]], # PA场景每个batch的实际token数 + "cmp_residual_k": [ + [ + 2, + 1, + 3, + 0, + 3, + 0, + 0, + 0, + 3, + 1, + 1, + 2, + 2, + 2, + 3, + 2, + 1, + 2, + 1, + 2, + 1, + 2, + 2, + 0, + 2, + 1, + 2, + 3, + 1, + 0, + 0, + 0, + 3, + 2, + 0, + 3, + 0, + 2, + 2, + 1, + 2, + 0, + ] + ], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [1024], + "sparse_mode": [3], + "query_datarange": [[0, 0.001]], + "key_datarange": [[0.001, 0.01]], + "weights_datarange": [[-0.5, 0.5]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "return_value": [1], + "output_idx_offset": [[random.randint(2, 10) for _ in range(168)] for _ in range(1)], + }, + # Ascend950 基础场景v6: BSND query + PA_BBND key + "quant_li_default_a5_v6": { + "batch_size": [13], + "q_seq": [4], + "k_seq": [262144], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], # 取16的整数倍,最多支持到1024 + "block_num": [53290], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [None], + "seqused_k": [[201988] * 12 + [262144]], # PA场景每个batch的实际token数 + "cmp_residual_k": [[0, 2, 3, 3, 0, 1, 2, 0, 3, 2, 0, 0, 3]], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [1024], + "sparse_mode": [3], + "query_datarange": [[0.001, 0.01]], + "key_datarange": [[-5, 5]], + "weights_datarange": [[-2, -1]], + "q_scale_datarange": [[10, 255]], + "k_scale_datarange": [[10, 255]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "return_value": [1], + "output_idx_offset": [[random.randint(0, 1000000) for _ in range(52)] for _ in range(1)], + }, + # Ascend950 基础场景v7: BSND query + PA_BBND key + "quant_li_default_a5_v7": { + "batch_size": [4], + "q_seq": [4096], + "k_seq": [1024], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [512], # 取16的整数倍,最多支持到1024 + "block_num": [45], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [None], + "seqused_k": [[1000] * 3 + [1024]], # PA场景每个batch的实际token数 + "cmp_residual_k": [None], + "max_seqlen_q": [4096], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [0], + "query_datarange": [[-5, 5]], + "key_datarange": [[-100, 100]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "return_value": [1], + "output_idx_offset": [[random.randint(0, 1000000) for _ in range(16384)] for _ in range(1)], + }, + # Ascend950 基础场景v8: BSND query + PA_BBND key + "quant_li_default_a5_v8": { + "batch_size": [4], + "q_seq": [4096], + "k_seq": [1024], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [512], # 取16的整数倍,最多支持到1024 + "block_num": [41], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [[2736] * 3 + [4096]], + "seqused_k": [[92] * 3 + [1024]], # PA场景每个batch的实际token数 + "cmp_residual_k": [[3, 3, 2, 1]], + "max_seqlen_q": [4096], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [1024], + "sparse_mode": [3], + "query_datarange": [[-1, 1]], + "key_datarange": [[-0.5, 0.5]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "return_value": [0], + "output_idx_offset": [[random.randint(0, 1) for _ in range(16384)] for _ in range(1)], + }, + # Ascend950 hifp8 场景: BSND query + PA_BBND key + "quant_li_default_hifp8_a5": { + "batch_size": [3], + "q_seq": [13], + "k_seq": [111], + "q_t_size": [8], + "k_t_size": [15], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [1024], # 取16的整数倍,最多支持到1024 + "block_num": [100], + "qk_dtype": [torch.uint8], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], # BSND: cu_seqlens_q 不传 + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [[2, 5, 13]], + "seqused_k": [[2080, 2114, 1180]], + "max_seqlen_q": [-1], + "cmp_residual_k": [[3, 1, 3]], + "quant_mode": [4], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-448, 448]], + "key_datarange": [[-20, 20]], + "weights_datarange": [[-123, 123]], + "q_scale_datarange": [[0, 255]], + "k_scale_datarange": [[0, 65504]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "return_value": [0], + "output_idx_offset": [None], + }, + # Ascend910_93 场景: TND query + PA_BBND key + "quant_li_default_a3": { + "batch_size": [1], + "q_seq": [1], + "k_seq": [8192], + "q_t_size": [1], + "k_t_size": [8192], # 压缩后的值 + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [1024], # 取16的整数倍,最多支持到1024 + "block_num": [17], + "qk_dtype": [torch.int8], + "dequant_dtype": [torch.float16], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [[0, 1]], # TND: cu_seqlens_q 必传 [B+1] + "cu_seqlens_k": [None], # PA_BBND: cu_seqlens_k 不传 + "seqused_q": [None], # TND: seqused_q 可选,None 时从 cu_seqlens 推导 + "seqused_k": [[8196]], # PA_BBND: seqused_k 必传 + "cmp_residual_k": [[1]], # cmp_ratio=4 时需要 + "quant_mode": [2], # 910_93 tiling 要求 quant_mode=2 + "layout_query": ["TND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-100, 100]], + "key_datarange": [[-100, 100]], + "weights_datarange": [[-25, 25]], + "q_scale_datarange": [[0, 255]], + "k_scale_datarange": [[0, 65504]], + "cmp_ratio": [4], # 1/2/4/8/16/32/64/128 + "max_seqlen_q": [-1], + "return_value": [0], + "output_idx_offset": [None], + }, + # ==================== 白盒测试用例(针对LD+returnValue修改)==================== + # WB1: LD + return_value=0 — 验证合并分支仅 isNeedLD=true 时不输出 value + "wb_ld_rv0_bsnd_pa": { + "batch_size": [8], + "q_seq": [4], + "k_seq": [32768], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], + "block_num": [4096], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[4096] * 7 + [32768]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [0], + "output_idx_offset": [None], + "run_mode": ["eager"], + }, + # WB2: non-LD + return_value=1 — k_seq=128, block_size=128 → s2BlockNum=1, 不触发LD + # 验证合并分支仅 returnValueFlag=true 时正确输出 value + "wb_nold_rv1_bsnd_pa": { + "batch_size": [2], + "q_seq": [4], + "k_seq": [128], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [128], + "block_num": [2], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[128, 128]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [128], + "sparse_mode": [0], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 8], + "run_mode": ["eager"], + }, + # WB3: TND query + PA_BBND key + LD + return_value=1 + # 验证 infershape TND 分支 + ProcessLD value 输出 + "wb_ld_rv1_tnd_pa": { + "batch_size": [4], + "q_seq": [4], + "k_seq": [8192], + "q_t_size": [16], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], + "block_num": [512], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [[0, 4, 8, 12, 16]], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[2048, 2048, 2048, 8192]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["TND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 16], + "run_mode": ["eager"], + }, + # WB4: TND query + TND key + LD + return_value=1 + # 验证 TND+TND layout 路径 + ProcessLD + "wb_ld_rv1_tnd_tnd": { + "batch_size": [3], + "q_seq": [4], + "k_seq": [3072], + "q_t_size": [12], + "k_t_size": [3072], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], + "block_num": [48], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [[0, 4, 8, 12]], + "cu_seqlens_k": [[0, 1024, 2048, 3072]], + "seqused_q": [None], + "seqused_k": [None], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["TND"], + "layout_key": ["TND"], + "sparse_count": [256], + "sparse_mode": [0], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 12], + "run_mode": ["eager"], + }, + # WB5: BSND query + BSND key + LD + return_value=1 + # 验证非 PA key 路径 + ProcessLD + "wb_ld_rv1_bsnd_bsnd": { + "batch_size": [4], + "q_seq": [4], + "k_seq": [4096], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], + "block_num": [256], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[4096, 4096, 4096, 4096]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["BSND"], + "sparse_count": [512], + "sparse_mode": [0], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 16], + "run_mode": ["eager"], + }, + # WB6: quant_mode=4 (HIFLOAT8) + LD + return_value=1 + # 验证 HIFLOAT8 dtype + ProcessLD value 输出 + "wb_ld_rv1_hifp8": { + "batch_size": [4], + "q_seq": [4], + "k_seq": [8192], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], + "block_num": [512], + "qk_dtype": [torch.uint8], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[2048, 2048, 2048, 8192]], + "cmp_residual_k": [[0, 1, 2, 3]], + "max_seqlen_q": [4], + "quant_mode": [4], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-448, 448]], + "key_datarange": [[-20, 20]], + "weights_datarange": [[-123, 123]], + "q_scale_datarange": [[0, 255]], + "k_scale_datarange": [[0, 65504]], + "cmp_ratio": [4], + "return_value": [1], + "output_idx_offset": [[0] * 16], + "run_mode": ["eager"], + }, + # WB7: group_size=32 (q_head_num=32) + LD + return_value=1 + # 验证 gSize=32 路径 + ProcessLD + "wb_ld_rv1_g32": { + "batch_size": [4], + "q_seq": [4], + "k_seq": [8192], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [32], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], + "block_num": [512], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[2048, 2048, 2048, 8192]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 16], + "run_mode": ["eager"], + }, + # WB8: sparse_count=1 (最小边界) + LD + return_value=1 + # 验证 ProcessLD topk 边界 + "wb_ld_rv1_sparse1": { + "batch_size": [4], + "q_seq": [4], + "k_seq": [8192], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], + "block_num": [512], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[2048, 2048, 2048, 8192]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [1], + "sparse_mode": [0], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 16], + "run_mode": ["eager"], + }, + # WB9: block_size=16 (最小边界) + LD + return_value=1 + # 验证 ProcessLD 对齐边界 + "wb_ld_rv1_block16": { + "batch_size": [4], + "q_seq": [4], + "k_seq": [4096], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [16], + "block_num": [1024], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[1024, 1024, 1024, 4096]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [0], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 16], + "run_mode": ["eager"], + }, + # WB10: block_size=1024 (最大边界) + LD + return_value=1 + # 验证大 block_size 下 ProcessLD + "wb_ld_rv1_block1024": { + "batch_size": [4], + "q_seq": [4], + "k_seq": [32768], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [1024], + "block_num": [128], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[8192, 8192, 8192, 32768]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [512], + "sparse_mode": [3], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 16], + "run_mode": ["eager"], + }, + # WB11: sparse_count=2048 (最大边界) + LD + return_value=1 + # 验证 ProcessLD topk 最大值 + "wb_ld_rv1_sparse2048": { + "batch_size": [4], + "q_seq": [4], + "k_seq": [32768], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [64], + "block_num": [4096], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[8192, 8192, 8192, 32768]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [2048], + "sparse_mode": [0], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [1], + "output_idx_offset": [[0] * 16], + "run_mode": ["eager"], + }, + # WB12: non-LD + return_value=0 — 基线对照(两条件均为false) + "wb_nold_rv0_bsnd_pa": { + "batch_size": [2], + "q_seq": [4], + "k_seq": [128], + "q_t_size": [8], + "k_t_size": [15], + "q_head_num": [64], + "k_head_num": [1], + "head_dim": [128], + "block_size": [128], + "block_num": [2], + "qk_dtype": [torch.float8_e4m3fn], + "dequant_dtype": [torch.float32], + "actual_seq_dtype": [torch.int32], + "cu_seqlens_q": [None], + "cu_seqlens_k": [None], + "seqused_q": [None], + "seqused_k": [[128, 128]], + "cmp_residual_k": [None], + "max_seqlen_q": [4], + "quant_mode": [1], + "layout_query": ["BSND"], + "layout_key": ["PA_BBND"], + "sparse_count": [128], + "sparse_mode": [0], + "query_datarange": [[-1, 1]], + "key_datarange": [[-1, 1]], + "weights_datarange": [[-1, 1]], + "q_scale_datarange": [[0, 1]], + "k_scale_datarange": [[0, 1]], + "cmp_ratio": [1], + "return_value": [0], + "output_idx_offset": [None], + "run_mode": ["eager"], + }, +} + +# 按需选择要启用的测试参数(例如默认启用所有) +properties = torch.npu.get_device_properties() +if "Ascend910_93" in properties.name: + ENABLED_PARAMSETS = [ + ("quant_li_default_a3", TEST_PARAMS["quant_li_default_a3"]), + ] +elif "Ascend950" in properties.name: + ENABLED_PARAMSETS = [ + (name, TEST_PARAMS[name]) + for name in ( + "quant_li_default_a5_v2", + "quant_li_default_a5_v3", + "quant_li_default_a5_v4", + "quant_li_default_a5_v5", + "quant_li_default_a5_v6", + "quant_li_default_a5_v7", + "quant_li_default_a5_v8", + "quant_li_default_a5_mxfp8", + "quant_li_default_a5_mxfp4", + "quant_li_default_a5_mxfp8_bsnd", + "quant_li_default_a5_mxfp4_bsnd", + "quant_li_default_a5_mxfp8_tnd", + "quant_li_default_a5_mxfp4_tnd", + # 白盒测试用例 + "wb_ld_rv0_bsnd_pa", + "wb_nold_rv1_bsnd_pa", + "wb_ld_rv1_tnd_pa", + "wb_ld_rv1_tnd_tnd", + "wb_ld_rv1_bsnd_bsnd", + "wb_ld_rv1_hifp8", + "wb_ld_rv1_g32", + "wb_ld_rv1_sparse1", + "wb_ld_rv1_block16", + "wb_ld_rv1_block1024", + "wb_ld_rv1_sparse2048", + "wb_nold_rv0_bsnd_pa", + ) + ] + +ENABLED_PARAMS = [params for _, params in ENABLED_PARAMSETS] diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_single.py b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_single.py new file mode 100644 index 000000000000..fe76135153a2 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_quant_lightning_indexer_v2_single.py @@ -0,0 +1,242 @@ +#!/usr/bin/python +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import itertools +import os +from pathlib import Path + +import pytest +import quant_lightning_indexer_v2_golden +import result_compare_method +import torch +import torch_npu +from batch import quant_lightning_indexer_v2_pt_loadprocess +from qliv2_test_utils import QliV2ResultWriter, ensure_comparison_passed +from test_quant_lightning_indexer_v2_paramset import ENABLED_PARAMSETS + +SAVE_PT_DIR = os.environ.get("QLIV2_SINGLE_SAVE_PT_DIR", "").strip() +RESULT_PATH = os.environ.get("QLIV2_SINGLE_RESULT_PATH", "").strip() + +param_names = [ + "batch_size", + "q_seq", + "k_seq", + "q_t_size", + "k_t_size", + "q_head_num", + "k_head_num", + "head_dim", + "block_size", + "block_num", + "qk_dtype", + "weight_dtype", + "dequant_dtype", + "actual_seq_dtype", + "cu_seqlens_q", + "cu_seqlens_k", + "seqused_q", + "seqused_k", + "cmp_residual_k", + "max_seqlen_q", + "quant_mode", + "layout_query", + "layout_key", + "sparse_count", + "sparse_mode", + "query_datarange", + "key_datarange", + "weights_datarange", + "q_scale_datarange", + "k_scale_datarange", + "cmp_ratio", + "return_value", + "output_idx_offset", + "run_mode", +] + +param_combinations = [] +for paramset_name, params in ENABLED_PARAMSETS: + param_values = [ + params.get(name, params["dequant_dtype"] if name == "weight_dtype" else ["eager"]) for name in param_names + ] + combinations = list(itertools.product(*param_values)) + for combo_index, combo in enumerate(combinations, start=1): + param_dict = dict(zip(param_names, combo)) + param_dict["case_name"] = paramset_name if len(combinations) == 1 else f"{paramset_name}_{combo_index:03d}" + param_combinations.append(param_dict) + + +@pytest.mark.ci +@pytest.mark.parametrize("param_combinations", param_combinations) +def test_qliv2(param_combinations): # Init params and tensors + batch_size = param_combinations["batch_size"] + q_seq = param_combinations["q_seq"] + k_seq = param_combinations["k_seq"] + q_t_size = param_combinations["q_t_size"] + k_t_size = param_combinations["k_t_size"] + q_head_num = param_combinations["q_head_num"] + k_head_num = param_combinations["k_head_num"] + head_dim = param_combinations["head_dim"] + block_size = param_combinations["block_size"] + block_num = param_combinations["block_num"] + qk_dtype = param_combinations["qk_dtype"] + weight_dtype = param_combinations["weight_dtype"] + dequant_dtype = param_combinations["dequant_dtype"] + actual_seq_dtype = param_combinations["actual_seq_dtype"] + cu_seqlens_q = param_combinations["cu_seqlens_q"] + cu_seqlens_k = param_combinations["cu_seqlens_k"] + seqused_q = param_combinations["seqused_q"] + seqused_k = param_combinations["seqused_k"] + cmp_residual_k = param_combinations["cmp_residual_k"] + max_seqlen_q = param_combinations["max_seqlen_q"] + quant_mode = param_combinations["quant_mode"] + layout_query = param_combinations["layout_query"] + layout_key = param_combinations["layout_key"] + sparse_count = param_combinations["sparse_count"] + sparse_mode = param_combinations["sparse_mode"] + query_datarange = param_combinations["query_datarange"] + key_datarange = param_combinations["key_datarange"] + weights_datarange = param_combinations["weights_datarange"] + q_scale_datarange = param_combinations["q_scale_datarange"] + k_scale_datarange = param_combinations["k_scale_datarange"] + cmp_ratio = param_combinations["cmp_ratio"] + return_value = param_combinations["return_value"] + output_idx_offset = param_combinations["output_idx_offset"] + run_mode = os.environ.get("QLIV2_RUN_MODE", param_combinations["run_mode"]).strip().lower() + torch_npu.npu.set_device(0) + test_data = ( + batch_size, + q_seq, + k_seq, + q_t_size, + k_t_size, + q_head_num, + k_head_num, + head_dim, + block_size, + block_num, + qk_dtype, + weight_dtype, + dequant_dtype, + actual_seq_dtype, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + max_seqlen_q, + quant_mode, + layout_query, + layout_key, + sparse_count, + sparse_mode, + query_datarange, + key_datarange, + weights_datarange, + q_scale_datarange, + k_scale_datarange, + cmp_ratio, + return_value, + output_idx_offset, + ) + + case_name = QliV2ResultWriter.case_name( + test_data, + explicit_name=param_combinations["case_name"], + ) + if SAVE_PT_DIR: + case_data = quant_lightning_indexer_v2_golden.generate_qliv2_test_data(test_data) + case_path = Path(SAVE_PT_DIR) / f"{case_name}.pt" + case_path.parent.mkdir(parents=True, exist_ok=True) + torch.save(case_data, case_path) + print(f"当前用例 PT 已保存: {case_path}") + if run_mode == "eager": + ( + cpu_result, + npu_result, + topk_value, + cpu_topk_value, + npu_topk_value, + output_idx_offset, + _, + ) = quant_lightning_indexer_v2_pt_loadprocess.test_qliv2_process(case_path, device_id=0) + elif run_mode == "graph": + ( + cpu_result, + npu_result, + topk_value, + cpu_topk_value, + npu_topk_value, + output_idx_offset, + _, + ) = quant_lightning_indexer_v2_pt_loadprocess.test_qliv2_process_graph(case_path, device_id=0) + else: + raise ValueError(f"unsupported run_mode: {run_mode}") + elif run_mode == "eager": + cpu_result, npu_result, topk_value, cpu_topk_value, npu_topk_value = ( + quant_lightning_indexer_v2_golden.qliv2_output_single(test_data) + ) + elif run_mode == "graph": + import quant_lightning_indexer_v2_acl_graph + + cpu_result, npu_result, topk_value, cpu_topk_value, npu_topk_value = ( + quant_lightning_indexer_v2_acl_graph.qliv2_output_acl_graph(test_data) + ) + else: + raise ValueError(f"unsupported run_mode: {run_mode}") + # print("npu_result", npu_result) + # print("cpu_result:", cpu_result) + # Compare result accuracy + result, fulfill_percent = result_compare_method.check_result( + cpu_result, + npu_result, + topk_value, + output_idx_offset, + test_data, + cpu_topk_value, + npu_topk_value, + ) + print("result", result) + print("result", fulfill_percent) + result_return_value = "N/A" + fulfill_precent_return_value = 0 + if return_value: + result_return_value, fulfill_precent_return_value = result_compare_method.check_result_return_value( + cpu_topk_value, + npu_topk_value, + test_data, + cpu_result, + npu_result, + topk_value, + output_idx_offset, + ) + print(f"result_return_value: {result_return_value}") + print(f"result_return_value: {fulfill_precent_return_value}") + + if RESULT_PATH: + row = QliV2ResultWriter.row( + case_name, + test_data, + result, + fulfill_percent, + result_return_value, + fulfill_precent_return_value, + ) + QliV2ResultWriter.append(RESULT_PATH, row) + print(f"当前用例结果已写入: {RESULT_PATH}") + + ensure_comparison_passed( + case_name, + result, + fulfill_percent, + result_return_value, + fulfill_precent_return_value, + ) diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_run.sh b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_run.sh new file mode 100644 index 000000000000..ff7f047bba53 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/pytest/test_run.sh @@ -0,0 +1,226 @@ +#!/bin/bash +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +set -o pipefail + +SCRIPT_DIR=$(cd "$(dirname "$0")" && pwd) +DEFAULT_PT_PATH="$SCRIPT_DIR/pt_path" +PT_SAVE_SCRIPT="$SCRIPT_DIR/batch/quant_lightning_indexer_v2_pt_save.py" +LIST_PT_SCRIPT="$SCRIPT_DIR/batch/list_pt_from_excel.py" +BATCH_TEST_SCRIPT="$SCRIPT_DIR/test_quant_lightning_indexer_v2_batch.py" +SINGLE_TEST_SCRIPT="$SCRIPT_DIR/test_quant_lightning_indexer_v2_single.py" + +show_help() { + cat < [选项] + +命令: + single 执行 paramset 中的单用例,可保存本次实际输入 PT 和结果表 + batch 有 -E 时生成 PT 后执行;无 -E 时直接执行已有 PT + batch_exec 按 Excel 中的 Testcase_Name 筛选已有 PT,仅运行 NPU 和 compare + help 显示帮助 + +通用选项: + -M, --run-mode MODE eager|graph,默认 eager + -O, --output FILE 结果 Excel 路径;single 默认 single_result.xlsx,batch 默认 result.xlsx + +batch/batch_exec 选项: + -C, --cases NAMES 按给定顺序执行 case 名,逗号分隔,可省略 .pt + -I, --indexes INDEXES 按自然排序后的 1-based 序号执行,如 3,1,5-7 + -E, --excel FILE Excel 路径;batch 不传时跳过 PT 生成 + -S, --sheet NAME Sheet 名,默认 Sheet1 + -P, --pt-path DIR PT 生成和读取目录,默认 $DEFAULT_PT_PATH + +single 选项: + --save-pt DIR single 保存本次实际输入和 CPU golden 的目录 + +示例: + $0 single --save-pt ./single_pt -O ./result/single.xlsx + $0 batch -P ./pt_path + $0 batch -E ./excel/test_cases.xlsx -P ./pt_path -O ./result/batch.xlsx + $0 batch -P ./pt_path -I 3,1,5-7 + $0 batch_exec -E ./excel/test_cases.xlsx -P ./pt_path -M graph +EOF +} + +require_value() { + if [ -z "$2" ]; then + echo "错误: $1 缺少参数值" >&2 + exit 2 + fi +} + +validate_run_mode() { + if [ "$RUN_MODE" != "eager" ] && [ "$RUN_MODE" != "graph" ]; then + echo "错误: run mode 仅支持 eager/graph,当前值: $RUN_MODE" >&2 + exit 2 + fi +} + +run_batch_pytest() { + local explicit_files="$1" + QLIV2_TESTCASE_DIR="$PT_PATH" \ + QLIV2_PT_FILE_LIST="$explicit_files" \ + QLIV2_CASE_NAMES="$CASE_NAMES" \ + QLIV2_CASE_INDEXES="$CASE_INDEXES" \ + QLIV2_RESULT_PATH="$RESULT_PATH" \ + QLIV2_RUN_MODE="$RUN_MODE" \ + python3 -m pytest -rA -s "$BATCH_TEST_SCRIPT" -v -m ci \ + -W ignore::UserWarning -W ignore::DeprecationWarning +} + +run_single() { + echo "===== QLI_V2 single: mode=$RUN_MODE result=$RESULT_PATH =====" + QLIV2_SINGLE_SAVE_PT_DIR="$SAVE_PT_DIR" \ + QLIV2_SINGLE_RESULT_PATH="$RESULT_PATH" \ + QLIV2_RUN_MODE="$RUN_MODE" \ + python3 -m pytest -rA -s "$SINGLE_TEST_SCRIPT" -v -m ci \ + -W ignore::UserWarning -W ignore::DeprecationWarning +} + +run_batch() { + if [ -n "$EXCEL_PATH" ]; then + if [ ! -f "$EXCEL_PATH" ]; then + echo "错误: Excel 文件不存在: $EXCEL_PATH" >&2 + exit 1 + fi + echo "===== 生成 PT: excel=$EXCEL_PATH sheet=$EXCEL_SHEET output=$PT_PATH =====" + python3 "$PT_SAVE_SCRIPT" "$EXCEL_PATH" "$PT_PATH" --sheet "$EXCEL_SHEET" || exit 1 + elif [ ! -d "$PT_PATH" ]; then + echo "错误: PT 目录不存在: $PT_PATH" >&2 + exit 1 + fi + echo "===== 执行 PT: input=$PT_PATH mode=$RUN_MODE result=$RESULT_PATH =====" + run_batch_pytest "" +} + +run_batch_from_excel() { + if [ -z "$EXCEL_PATH" ]; then + echo "错误: batch_exec 必须指定 -E/--excel" >&2 + exit 2 + fi + if [ ! -f "$EXCEL_PATH" ]; then + echo "错误: Excel 文件不存在: $EXCEL_PATH" >&2 + exit 1 + fi + if [ ! -d "$PT_PATH" ]; then + echo "错误: PT 目录不存在: $PT_PATH" >&2 + exit 1 + fi + local file_list + file_list=$(python3 "$LIST_PT_SCRIPT" "$EXCEL_PATH" "$PT_PATH" --sheet "$EXCEL_SHEET") || exit 1 + echo "===== Excel 筛选后仅执行 NPU + compare: mode=$RUN_MODE result=$RESULT_PATH =====" + run_batch_pytest "$file_list" +} + +if [ $# -lt 1 ]; then + show_help + exit 2 +fi + +COMMAND="$1" +shift + +EXCEL_PATH="" +EXCEL_SHEET="Sheet1" +PT_PATH="$DEFAULT_PT_PATH" +RUN_MODE="eager" +RESULT_PATH="" +CASE_NAMES="" +CASE_INDEXES="" +SAVE_PT_DIR="" + +while [ $# -gt 0 ]; do + case "$1" in + -E|--excel) + require_value "$1" "$2" + EXCEL_PATH="$2" + shift 2 + ;; + -S|--sheet) + require_value "$1" "$2" + EXCEL_SHEET="$2" + shift 2 + ;; + -P|--pt-path) + require_value "$1" "$2" + PT_PATH="$2" + shift 2 + ;; + -M|--run-mode) + require_value "$1" "$2" + RUN_MODE="$2" + shift 2 + ;; + -O|--output) + require_value "$1" "$2" + RESULT_PATH="$2" + shift 2 + ;; + -C|--cases) + require_value "$1" "$2" + CASE_NAMES="$2" + shift 2 + ;; + -I|--indexes) + require_value "$1" "$2" + CASE_INDEXES="$2" + shift 2 + ;; + --save-pt) + require_value "$1" "$2" + SAVE_PT_DIR="$2" + shift 2 + ;; + -h|--help) + show_help + exit 0 + ;; + *) + echo "错误: 未知选项 $1" >&2 + show_help + exit 2 + ;; + esac +done + +if [ -n "$CASE_NAMES" ] && [ -n "$CASE_INDEXES" ]; then + echo "错误: --cases 和 --indexes 不能同时使用" >&2 + exit 2 +fi +if [ -z "$RESULT_PATH" ]; then + if [ "$COMMAND" = "single" ]; then + RESULT_PATH="$SCRIPT_DIR/single_result.xlsx" + else + RESULT_PATH="$SCRIPT_DIR/result.xlsx" + fi +fi +validate_run_mode + +case "$COMMAND" in + single) + run_single + ;; + batch) + run_batch + ;; + batch_exec) + run_batch_from_excel + ;; + help) + show_help + ;; + *) + echo "错误: 未知命令 $COMMAND" >&2 + show_help + exit 2 + ;; +esac diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/ut/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2/tests/ut/CMakeLists.txt new file mode 100644 index 000000000000..d5a84231c22b --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/ut/CMakeLists.txt @@ -0,0 +1,16 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/CMakeLists.txt new file mode 100644 index 000000000000..e611b728108a --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/CMakeLists.txt @@ -0,0 +1,20 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +if(UT_TEST_ALL OR OP_HOST_UT) + add_modules_ut_sources(UT_NAME ${OP_INFERSHAPE_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) +endif() + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch22/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch22/CMakeLists.txt new file mode 100644 index 000000000000..96095b3345f2 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch22/CMakeLists.txt @@ -0,0 +1,13 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +if(UT_TEST_ALL OR OP_HOST_UT) + add_modules_ut_sources(UT_NAME ${OP_TILING_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) +endif() diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch22/test_quant_lightning_indexer_v2_tiling.cpp b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch22/test_quant_lightning_indexer_v2_tiling.cpp new file mode 100644 index 000000000000..0a94620d1856 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch22/test_quant_lightning_indexer_v2_tiling.cpp @@ -0,0 +1,191 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ +#include +#include +#include "../test_quant_lightning_indexer_v2_utils.h" + +// DAV_2201 (Ascend910B) tiling cases for QuantLightningIndexerV2 +class QuantLightningIndexerV2TilingArch22 : public testing::Test { +protected: + static void SetUpTestCase() + { + std::cout << "QuantLightningIndexerV2TilingArch22 SetUp" << std::endl; + } + + static void TearDownTestCase() + { + std::cout << "QuantLightningIndexerV2TilingArch22 TearDown" << std::endl; + } +}; + +// BSND/PA_BBND int8 success on Ascend910B: quant_mode=2, topk=2048, mask_mode=0 +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_int8_pa_success) +{ + qliv2_ut::CaseParam p; + qliv2_ut::RunTilingCase(p, ge::GRAPH_SUCCESS); +} + +// BSND/PA_BBND int8 success with cmp_residual_k and output_idx_offset on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_cmp_residual_success) +{ + qliv2_ut::CaseParam p; + p.cmpRatio = 4; + p.maskMode = 3; + p.cmpResidual = {2}; + p.idxOffset = {2, 39, 64}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_SUCCESS); +} + +// layout_k only supports PA_BBND on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_layout_k_failed) +{ + qliv2_ut::CaseParam p; + p.layoutK = "BSND"; + p.blockTable = {}; + p.sequsedK = {}; + p.kShape = {2, 64, 1, 128}; + p.kScaleShape = {2, 64, 1}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// quant_mode only supports 2 on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_quant_mode_failed) +{ + qliv2_ut::CaseParam p; + p.quantMode = 1; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// return_value only supports false on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_return_value_failed) +{ + qliv2_ut::CaseParam p; + p.returnValue = 1; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// topk must > 0 and <= 2048 on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_topk_over_limit_failed) +{ + qliv2_ut::CaseParam p; + p.topk = 4096; + p.outShape = {2, 39, 1, 4096}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// cmp_ratio must be a power of 2 on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_cmp_ratio_not_pow2_failed) +{ + qliv2_ut::CaseParam p; + p.cmpRatio = 3; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of q and k must be int8 on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_q_dtype_failed) +{ + qliv2_ut::CaseParam p; + p.qType = ge::DT_FLOAT16; + p.kType = ge::DT_FLOAT16; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of q_descale and k_descale must be float16 on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_scale_dtype_failed) +{ + qliv2_ut::CaseParam p; + p.qScaleType = ge::DT_FLOAT; + p.kScaleType = ge::DT_FLOAT; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of w must be float16 on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_w_dtype_failed) +{ + qliv2_ut::CaseParam p; + p.wType = ge::DT_FLOAT; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// gSize (N1/N2) must equal 64 on Ascend910B +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_gsize_failed) +{ + qliv2_ut::CaseParam p; + p.qShape = {2, 39, 128, 128}; + p.wShape = {2, 39, 128}; + p.qScaleShape = {2, 39, 128}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// mask_mode only supports 0 or 3 +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_mask_mode_failed) +{ + qliv2_ut::CaseParam p; + p.maskMode = 1; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// metadata is required +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_metadata_missing_failed) +{ + qliv2_ut::CaseParam p; + p.metadata = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of q and k must be same: q=int8, k=fp8 should fail +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_qk_dtype_mismatch_failed) +{ + qliv2_ut::CaseParam p; + p.kType = ge::DT_FLOAT8_E4M3FN; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of sparse_values must be bfloat16 +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_values_dtype_failed) +{ + qliv2_ut::CaseParam p; + p.valuesType = ge::DT_FLOAT16; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// head dim of q only supports 128 +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_head_dim_failed) +{ + qliv2_ut::CaseParam p; + p.qShape = {2, 39, 64, 127}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// block_size of k must be a multiple of 16 +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_block_size_failed) +{ + qliv2_ut::CaseParam p; + p.kShape = {2, 17, 1, 128}; + p.kScaleShape = {2, 17, 1}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// head num of k only supports 1 +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_910b_tiling_k_headnum_failed) +{ + qliv2_ut::CaseParam p; + p.kShape = {2, 16, 2, 128}; + p.kScaleShape = {2, 16, 2}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// unsupported npu arch (Ascend310P) should fail +TEST_F(QuantLightningIndexerV2TilingArch22, QuantLightningIndexerV2_tiling_unsupported_arch_failed) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend310P"; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch35/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch35/CMakeLists.txt new file mode 100644 index 000000000000..96095b3345f2 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch35/CMakeLists.txt @@ -0,0 +1,13 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +if(UT_TEST_ALL OR OP_HOST_UT) + add_modules_ut_sources(UT_NAME ${OP_TILING_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) +endif() diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch35/test_quant_lightning_indexer_v2_tiling.cpp b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch35/test_quant_lightning_indexer_v2_tiling.cpp new file mode 100644 index 000000000000..813e0b49fb03 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/arch35/test_quant_lightning_indexer_v2_tiling.cpp @@ -0,0 +1,527 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ +#include +#include +#include "../test_quant_lightning_indexer_v2_utils.h" + +// DAV_3510 (Ascend950) tiling cases for QuantLightningIndexerV2 +class QuantLightningIndexerV2TilingArch35 : public testing::Test { +protected: + static void SetUpTestCase() + { + std::cout << "QuantLightningIndexerV2TilingArch35 SetUp" << std::endl; + } + + static void TearDownTestCase() + { + std::cout << "QuantLightningIndexerV2TilingArch35 TearDown" << std::endl; + } +}; + +namespace { +// Base of a valid Ascend950 TND/PA_BBND fp8 case +qliv2_ut::CaseParam Make950TndPaFp8() +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.qShape = {78, 64, 128}; + p.kShape = {2, 16, 1, 128}; + p.wShape = {78, 64}; + p.qScaleShape = {78, 64}; + p.kScaleShape = {2, 16, 1}; + p.outShape = {78, 1, 2048}; + p.qType = ge::DT_FLOAT8_E4M3FN; + p.kType = ge::DT_FLOAT8_E4M3FN; + p.wType = ge::DT_FLOAT; + p.qScaleType = ge::DT_FLOAT; + p.kScaleType = ge::DT_FLOAT; + p.cuSeqQ = {3}; + p.layoutQ = "TND"; + p.layoutK = "PA_BBND"; + p.quantMode = 1; + p.maxSeqlenQ = 64; + p.maskMode = 3; + return p; +} +} // namespace + +// TND/PA_BBND fp8 success on Ascend950: quant_mode=1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_fp8_tnd_pa_success) +{ + qliv2_ut::RunTilingCase(Make950TndPaFp8(), ge::GRAPH_SUCCESS); +} + +// TND/PA_BBND fp8 success with cmp_residual_k and output_idx_offset on Ascend950 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_fp8_cmp_residual_success) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.cmpRatio = 4; + p.cmpResidual = {2}; + p.idxOffset = {78, 64}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_SUCCESS); +} + +// BSND/BSND mxfp8 success on Ascend950: quant_mode=3, e8m0 scale +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_mxfp8_bsnd_bsnd_success) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.kShape = {2, 64, 1, 128}; + p.kScaleShape = {2, 64, 1, 2, 2}; + p.qScaleShape = {2, 39, 64, 2, 2}; + p.qType = ge::DT_FLOAT8_E4M3FN; + p.kType = ge::DT_FLOAT8_E4M3FN; + p.wType = ge::DT_FLOAT; + p.qScaleType = ge::DT_FLOAT8_E8M0; + p.kScaleType = ge::DT_FLOAT8_E8M0; + p.layoutK = "BSND"; + p.quantMode = 3; + p.sequsedK = {}; + p.blockTable = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_SUCCESS); +} + +// TND/TND hifloat8 success on Ascend950: quant_mode=4, return_value=1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_hif8_tnd_tnd_success) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.qShape = {78, 64, 128}; + p.kShape = {128, 1, 128}; + p.wShape = {78, 64}; + p.qScaleShape = {1}; + p.kScaleShape = {1}; + p.outShape = {78, 1, 2048}; + p.valuesShape = {78, 1, 2048}; + p.qType = ge::DT_HIFLOAT8; + p.kType = ge::DT_HIFLOAT8; + p.wType = ge::DT_FLOAT; + p.qScaleType = ge::DT_FLOAT; + p.kScaleType = ge::DT_FLOAT; + p.cuSeqQ = {3}; + p.cuSeqK = {3}; + p.layoutQ = "TND"; + p.layoutK = "TND"; + p.quantMode = 4; + p.maxSeqlenQ = 64; + p.returnValue = 1; + p.sequsedK = {}; + p.blockTable = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_SUCCESS); +} + +// BSND/PA_BBND int8 success on Ascend950: quant_mode=2, return_value=1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_int8_pa_rv_success) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.valuesShape = {2, 39, 1, 2048}; + p.returnValue = 1; + qliv2_ut::RunTilingCase(p, ge::GRAPH_SUCCESS); +} + +// TND/PA_BBND mxfp4 success on Ascend950: quant_mode=5 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_mxfp4_tnd_pa_success) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.qType = ge::DT_FLOAT4_E2M1; + p.kType = ge::DT_FLOAT4_E2M1; + p.qScaleShape = {78, 64, 2, 2}; + p.kScaleShape = {2, 16, 1, 2, 2}; + p.qScaleType = ge::DT_FLOAT8_E8M0; + p.kScaleType = ge::DT_FLOAT8_E8M0; + p.quantMode = 5; + p.maskMode = 0; + qliv2_ut::RunTilingCase(p, ge::GRAPH_SUCCESS); +} + +// layout_k only supports PA_BBND, BSND or TND on Ascend950 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_layout_k_invalid_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.layoutK = "XXX"; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// outside of PA, layout_q and layout_k must be the same +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_layout_mismatch_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.layoutK = "BSND"; + p.kShape = {2, 64, 1, 128}; + p.kScaleShape = {2, 64, 1}; + p.sequsedK = {}; + p.blockTable = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// topk must > 0 and <= 8192 on Ascend950 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_topk_over_limit_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.topk = 10000; + p.outShape = {78, 1, 10000}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// cmp_ratio must > 0 and <= 128 on Ascend950 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_cmp_ratio_over_limit_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.cmpRatio = 200; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// quant_mode only supports 1-5 on Ascend950 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_quant_mode_invalid_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.quantMode = 6; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// return_value only supports 0 or 1 on Ascend950 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_return_value_invalid_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.returnValue = 2; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// max_seqlen_q must >= -1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_max_seqlen_q_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.maxSeqlenQ = -2; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of q must match quant_mode: int8 with quant_mode=1 should fail +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_q_dtype_mismatch_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.qType = ge::DT_INT8; + p.kType = ge::DT_INT8; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of q_descale must match quant_mode: e8m0 with quant_mode=1 should fail +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_scale_dtype_mismatch_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.qScaleType = ge::DT_FLOAT8_E8M0; + p.kScaleType = ge::DT_FLOAT8_E8M0; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// when q is int8 (quant_mode=2), dtype of w must be float16 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_int8_w_dtype_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.quantMode = 2; + p.wType = ge::DT_FLOAT; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// when q is not int8 (quant_mode=1), dtype of w must be float +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_fp8_w_dtype_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.wType = ge::DT_FLOAT16; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of q and k must be same +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_qk_dtype_mismatch_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.kType = ge::DT_INT8; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of q_descale and k_descale must be same +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_scale_dtype_not_equal_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.kScaleType = ge::DT_FLOAT16; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of sparse_indices must be int32 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_out_dtype_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.outType = ge::DT_FLOAT; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of sparse_values must be bfloat16 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_values_dtype_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.valuesType = ge::DT_FLOAT16; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// PA_BBND requires block_table +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_pa_block_table_missing_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.blockTable = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// PA_BBND must not provide cu_seqlens_k +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_pa_cu_seqlens_k_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.cuSeqK = {2}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// TND k requires cu_seqlens_k +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_tnd_k_cu_seqlens_k_missing_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.layoutK = "TND"; + p.kShape = {128, 1, 128}; + p.kScaleShape = {128, 1}; + p.blockTable = {}; + p.sequsedK = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// BSND k must not provide cu_seqlens_k +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_bsnd_k_cu_seqlens_k_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.layoutK = "BSND"; + p.kShape = {2, 64, 1, 128}; + p.kScaleShape = {2, 64, 1}; + p.blockTable = {}; + p.sequsedK = {}; + p.cuSeqK = {2}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// non-PA layout must not provide block_table +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_bsnd_block_table_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.layoutK = "TND"; + p.kShape = {128, 1, 128}; + p.kScaleShape = {128, 1}; + p.sequsedK = {}; + p.cuSeqK = {3}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// cmp_ratio != 1 and mask_mode != 0 require cmp_residual_k +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_cmp_residual_missing_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.cmpRatio = 4; + p.cmpResidual = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// TND q requires cu_seqlens_q +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_tnd_cu_seqlens_q_missing_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.cuSeqQ = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// metadata shape size must be 1024 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_metadata_size_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.metadata = {512}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// mxfp8 scale dim num must be q dim num + 1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_mxfp8_scale_dim_failed) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.kShape = {2, 64, 1, 128}; + p.kScaleShape = {2, 64, 1}; + p.qScaleShape = {2, 39, 64}; + p.qType = ge::DT_FLOAT8_E4M3FN; + p.kType = ge::DT_FLOAT8_E4M3FN; + p.wType = ge::DT_FLOAT; + p.qScaleType = ge::DT_FLOAT8_E8M0; + p.kScaleType = ge::DT_FLOAT8_E8M0; + p.layoutK = "BSND"; + p.quantMode = 3; + p.sequsedK = {}; + p.blockTable = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// hifloat8 scale dim num must be 1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_hif8_scale_dim_failed) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.qShape = {78, 64, 128}; + p.kShape = {128, 1, 128}; + p.wShape = {78, 64}; + p.qScaleShape = {78, 64}; + p.kScaleShape = {1}; + p.outShape = {78, 1, 2048}; + p.qType = ge::DT_HIFLOAT8; + p.kType = ge::DT_HIFLOAT8; + p.wType = ge::DT_FLOAT; + p.qScaleType = ge::DT_FLOAT; + p.kScaleType = ge::DT_FLOAT; + p.cuSeqQ = {3}; + p.cuSeqK = {3}; + p.layoutQ = "TND"; + p.layoutK = "TND"; + p.quantMode = 4; + p.maxSeqlenQ = 64; + p.sequsedK = {}; + p.blockTable = {}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// fp8 scale dim num must be q dim num - 1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_fp8_scale_dim_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.qScaleShape = {78}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// head num of k only supports 1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_k_headnum_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.kShape = {2, 16, 2, 128}; + p.kScaleShape = {2, 16, 2}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// gSize must <= 64 on Ascend950 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_gsize_over_limit_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.qShape = {2, 39, 128, 128}; + p.wShape = {2, 39, 128}; + p.qScaleShape = {2, 39, 128}; + p.outShape = {2, 39, 1, 2048}; + p.layoutQ = "BSND"; + p.cuSeqQ = {}; + p.maskMode = 0; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// block_size of k must be a multiple of 16 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_block_size_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.kShape = {2, 17, 1, 128}; + p.kScaleShape = {2, 17, 1}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of cu_seqlens_q only supports int32 (TND/TND so that cu_seqlens_k desc is valid for error logging) +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_cu_seqlens_q_dtype_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.layoutK = "TND"; + p.kShape = {128, 1, 128}; + p.kScaleShape = {128, 1}; + p.blockTable = {}; + p.sequsedK = {}; + p.cuSeqK = {3}; + p.cuSeqQType = ge::DT_INT64; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of seqused_q only supports int32 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_seqused_q_dtype_failed) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.valuesShape = {2, 39, 1, 2048}; + p.sequsedQ = {2}; + p.sequsedQType = ge::DT_INT64; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of cmp_residual_k only supports int32 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_cmp_residual_dtype_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.cmpRatio = 4; + p.cmpResidual = {2}; + p.cmpResidualType = ge::DT_INT64; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of output_idx_offset only supports int32 (seqused_q provided so its desc is valid for error logging) +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_idx_offset_dtype_failed) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.idxOffset = {2, 39, 64}; + p.idxOffsetType = ge::DT_INT64; + p.sequsedQ = {2}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of block_table only supports int32 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_block_table_dtype_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.blockTableType = ge::DT_INT64; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// dtype of seqused_k only supports int32 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_seqused_k_dtype_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.sequsedKType = ge::DT_INT64; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// last dim of sparse_values must be same as topk when return_value=1 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_values_topk_mismatch_failed) +{ + qliv2_ut::CaseParam p; + p.soc = "Ascend950"; + p.coreNum = 56; + p.valuesShape = {2, 39, 1, 1024}; + p.returnValue = 1; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} + +// head dim of q only supports 128 +TEST_F(QuantLightningIndexerV2TilingArch35, QuantLightningIndexerV2_950_tiling_head_dim_failed) +{ + qliv2_ut::CaseParam p = Make950TndPaFp8(); + p.qShape = {78, 64, 127}; + p.wShape = {78, 64}; + p.qScaleShape = {78, 64}; + qliv2_ut::RunTilingCase(p, ge::GRAPH_FAILED); +} diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/test_quant_lightning_indexer_v2_infershape.cpp b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/test_quant_lightning_indexer_v2_infershape.cpp new file mode 100644 index 000000000000..e4ddbf0eeda9 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/test_quant_lightning_indexer_v2_infershape.cpp @@ -0,0 +1,216 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include +#include +#include "infer_shape_context_faker.h" +#include "infer_datatype_context_faker.h" +#include "infer_shape_case_executor.h" +#include "base/registry/op_impl_space_registry_v2.h" + +class QuantLightningIndexerV2Proto : public testing::Test { +protected: + static void SetUpTestCase() + { + std::cout << "QuantLightningIndexerV2Proto SetUp" << std::endl; + } + + static void TearDownTestCase() + { + std::cout << "QuantLightningIndexerV2Proto TearDown" << std::endl; + } +}; + +// BSND/BSND, return_value=1, topk=128 +TEST_F(QuantLightningIndexerV2Proto, QuantLightningIndexerV2_infershape_bsnd) +{ + gert::InfershapeContextPara infershapeContextPara( + "QuantLightningIndexerV2", + // 输入Tensor (13个) + { + {{{1, 8, 8, 128}, {1, 8, 8, 128}}, ge::DT_INT8, ge::FORMAT_ND}, // q input0 + {{{1, 64, 1, 128}, {1, 64, 1, 128}}, ge::DT_INT8, ge::FORMAT_ND}, // k input1 + {{{1, 8, 8}, {1, 8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND}, // w input2 + {{{1, 8, 8}, {1, 8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND}, // q_descale input3 + {{{1, 64, 1}, {1, 64, 1}}, ge::DT_FLOAT16, ge::FORMAT_ND}, // k_descale input4 + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cu_seqlens_q input5 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cu_seqlens_k input6 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // seqused_q input7 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // seqused_k input8 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cmp_residual_k input9 (optional) + {{{1, 4}, {1, 4}}, ge::DT_INT32, ge::FORMAT_ND}, // block_table input10 + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // output_idx_offset input11 (optional) + {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND} // metadata input12 + }, + // 输出Tensor + { + {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND}, // sparse_indices + {{{}, {}}, ge::DT_BF16, ge::FORMAT_ND} // sparse_values + }, + // 属性 + {{"topk", Ops::Transformer::AnyValue::CreateFrom(128)}, + {"quant_mode", Ops::Transformer::AnyValue::CreateFrom(2)}, + {"max_seqlen_q", Ops::Transformer::AnyValue::CreateFrom(-1)}, + {"layout_q", Ops::Transformer::AnyValue::CreateFrom("BSND")}, + {"layout_k", Ops::Transformer::AnyValue::CreateFrom("BSND")}, + {"mask_mode", Ops::Transformer::AnyValue::CreateFrom(0)}, + {"cmp_ratio", Ops::Transformer::AnyValue::CreateFrom(1)}, + {"return_value", Ops::Transformer::AnyValue::CreateFrom(1)}}); + + std::vector> expectOutputShape = {{1, 8, 1, 128}, {1, 8, 1, 128}}; + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); +} + +// TND/TND, return_value=0, topk=2048 +TEST_F(QuantLightningIndexerV2Proto, QuantLightningIndexerV2_infershape_tnd) +{ + gert::InfershapeContextPara infershapeContextPara( + "QuantLightningIndexerV2", + { + {{{64, 32, 128}, {64, 32, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND}, // q input0 + {{{64, 1, 128}, {64, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND}, // k input1 + {{{64, 32}, {64, 32}}, ge::DT_FLOAT, ge::FORMAT_ND}, // w input2 + {{{64, 32}, {64, 32}}, ge::DT_FLOAT, ge::FORMAT_ND}, // q_descale input3 + {{{64, 1}, {64, 1}}, ge::DT_FLOAT, ge::FORMAT_ND}, // k_descale input4 + {{{2}, {2}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cu_seqlens_q input5 + {{{2}, {2}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cu_seqlens_k input6 + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // seqused_q input7 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // seqused_k input8 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cmp_residual_k input9 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // block_table input10 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // output_idx_offset input11 (optional) + {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND} // metadata input12 + }, + { + {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND}, // sparse_indices + {{{}, {}}, ge::DT_BF16, ge::FORMAT_ND} // sparse_values + }, + {{"topk", Ops::Transformer::AnyValue::CreateFrom(2048)}, + {"quant_mode", Ops::Transformer::AnyValue::CreateFrom(1)}, + {"max_seqlen_q", Ops::Transformer::AnyValue::CreateFrom(64)}, + {"layout_q", Ops::Transformer::AnyValue::CreateFrom("TND")}, + {"layout_k", Ops::Transformer::AnyValue::CreateFrom("TND")}, + {"mask_mode", Ops::Transformer::AnyValue::CreateFrom(0)}, + {"cmp_ratio", Ops::Transformer::AnyValue::CreateFrom(1)}, + {"return_value", Ops::Transformer::AnyValue::CreateFrom(0)}}); + + std::vector> expectOutputShape = {{64, 1, 2048}, {0}}; + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); +} + +// TND/PA_BBND, return_value=1, topk=2048 +TEST_F(QuantLightningIndexerV2Proto, QuantLightningIndexerV2_infershape_tnd_pa) +{ + gert::InfershapeContextPara infershapeContextPara( + "QuantLightningIndexerV2", + { + {{{64, 32, 128}, {64, 32, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND}, // q input0 + {{{1, 16, 1, 128}, {1, 16, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND}, // k (PA) input1 + {{{64, 32}, {64, 32}}, ge::DT_FLOAT, ge::FORMAT_ND}, // w input2 + {{{64, 32}, {64, 32}}, ge::DT_FLOAT, ge::FORMAT_ND}, // q_descale input3 + {{{1, 16, 1}, {1, 16, 1}}, ge::DT_FLOAT, ge::FORMAT_ND}, // k_descale input4 + {{{2}, {2}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cu_seqlens_q input5 + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cu_seqlens_k input6 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // seqused_q input7 (optional) + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND, true}, // seqused_k input8 + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cmp_residual_k input9 (optional) + {{{1, 2}, {1, 2}}, ge::DT_INT32, ge::FORMAT_ND}, // block_table input10 + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // output_idx_offset input11 (optional) + {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND} // metadata input12 + }, + { + {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND}, // sparse_indices + {{{}, {}}, ge::DT_BF16, ge::FORMAT_ND} // sparse_values + }, + {{"topk", Ops::Transformer::AnyValue::CreateFrom(2048)}, + {"quant_mode", Ops::Transformer::AnyValue::CreateFrom(1)}, + {"max_seqlen_q", Ops::Transformer::AnyValue::CreateFrom(64)}, + {"layout_q", Ops::Transformer::AnyValue::CreateFrom("TND")}, + {"layout_k", Ops::Transformer::AnyValue::CreateFrom("PA_BBND")}, + {"mask_mode", Ops::Transformer::AnyValue::CreateFrom(0)}, + {"cmp_ratio", Ops::Transformer::AnyValue::CreateFrom(1)}, + {"return_value", Ops::Transformer::AnyValue::CreateFrom(1)}}); + + std::vector> expectOutputShape = {{64, 1, 2048}, {64, 1, 2048}}; + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); +} + +// invalid layout_q should fail +TEST_F(QuantLightningIndexerV2Proto, QuantLightningIndexerV2_infershape_layout_failed) +{ + gert::InfershapeContextPara infershapeContextPara( + "QuantLightningIndexerV2", + { + {{{1, 8, 8, 128}, {1, 8, 8, 128}}, ge::DT_INT8, ge::FORMAT_ND}, // q input0 + {{{1, 64, 1, 128}, {1, 64, 1, 128}}, ge::DT_INT8, ge::FORMAT_ND}, // k input1 + {{{1, 8, 8}, {1, 8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND}, // w input2 + {{{1, 8, 8}, {1, 8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND}, // q_descale input3 + {{{1, 64, 1}, {1, 64, 1}}, ge::DT_FLOAT16, ge::FORMAT_ND}, // k_descale input4 + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cu_seqlens_q input5 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cu_seqlens_k input6 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // seqused_q input7 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // seqused_k input8 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // cmp_residual_k input9 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // block_table input10 (optional) + {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true}, // output_idx_offset input11 (optional) + {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND} // metadata input12 + }, + { + {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND}, // sparse_indices + {{{}, {}}, ge::DT_BF16, ge::FORMAT_ND} // sparse_values + }, + {{"topk", Ops::Transformer::AnyValue::CreateFrom(128)}, + {"quant_mode", Ops::Transformer::AnyValue::CreateFrom(2)}, + {"max_seqlen_q", Ops::Transformer::AnyValue::CreateFrom(-1)}, + {"layout_q", Ops::Transformer::AnyValue::CreateFrom("SBND")}, + {"layout_k", Ops::Transformer::AnyValue::CreateFrom("BSND")}, + {"mask_mode", Ops::Transformer::AnyValue::CreateFrom(0)}, + {"cmp_ratio", Ops::Transformer::AnyValue::CreateFrom(1)}, + {"return_value", Ops::Transformer::AnyValue::CreateFrom(0)}}); + + ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED, {}); +} + +// infer dataType +TEST_F(QuantLightningIndexerV2Proto, QuantLightningIndexerV2_inferdtype) +{ + auto spaceRegistry = gert::DefaultOpImplSpaceRegistryV2::GetInstance().GetSpaceRegistry(); + ASSERT_NE(spaceRegistry, nullptr); + auto data_type_func = spaceRegistry->GetOpImpl("QuantLightningIndexerV2")->infer_datatype; + if (data_type_func != nullptr) { + ge::DataType inputQ = ge::DT_INT8; + ge::DataType inputK = ge::DT_INT8; + ge::DataType inputW = ge::DT_FLOAT16; + ge::DataType inputScale = ge::DT_FLOAT16; + ge::DataType inputI32 = ge::DT_INT32; + ge::DataType outputRef0 = ge::DT_INT32; + auto context_holder = + gert::InferDataTypeContextFaker() + .NodeIoNum(13, 2) + .NodeOutputTd(0, ge::FORMAT_ND, ge::FORMAT_ND) + .NodeOutputTd(1, ge::FORMAT_ND, ge::FORMAT_ND) + .InputDataTypes({&inputQ, &inputK, &inputW, &inputScale, &inputScale, &inputI32, &inputI32, &inputI32, + &inputI32, &inputI32, &inputI32, &inputI32, &inputI32}) + .NodeAttrs({{"topk", Ops::Transformer::AnyValue::CreateFrom(2048)}, + {"quant_mode", Ops::Transformer::AnyValue::CreateFrom(2)}, + {"max_seqlen_q", Ops::Transformer::AnyValue::CreateFrom(-1)}, + {"layout_q", Ops::Transformer::AnyValue::CreateFrom("BSND")}, + {"layout_k", Ops::Transformer::AnyValue::CreateFrom("PA_BBND")}, + {"mask_mode", Ops::Transformer::AnyValue::CreateFrom(0)}, + {"cmp_ratio", Ops::Transformer::AnyValue::CreateFrom(1)}, + {"return_value", Ops::Transformer::AnyValue::CreateFrom(0)}}) + .Build(); + auto context = context_holder.GetContext(); + EXPECT_EQ(data_type_func(context), ge::GRAPH_SUCCESS); + ASSERT_NE(context, nullptr); + + EXPECT_EQ(context->GetOutputDataType(0), outputRef0); + } +} diff --git a/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/test_quant_lightning_indexer_v2_utils.h b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/test_quant_lightning_indexer_v2_utils.h new file mode 100644 index 000000000000..d5b55db0c3a9 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/tests/ut/op_host/test_quant_lightning_indexer_v2_utils.h @@ -0,0 +1,114 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef TEST_QUANT_LIGHTNING_INDEXER_V2_UTILS_H +#define TEST_QUANT_LIGHTNING_INDEXER_V2_UTILS_H + +#include +#include +#include +#include "tiling_context_faker.h" +#include "tiling_case_executor.h" + +namespace qliv2_ut { + +constexpr uint64_t SKIP_TILING_KEY = UINT64_MAX; + +// Default case matches a valid Ascend910B PA_BBND int8 success case. +struct CaseParam { + std::string soc = "Ascend910B"; + uint64_t coreNum = 64; + uint64_t ubSize = 262144; + uint64_t l2Size = 16384; + std::vector qShape = {2, 39, 64, 128}; + std::vector kShape = {2, 16, 1, 128}; // PA_BBND: block_num, block_size, N2, D + std::vector wShape = {2, 39, 64}; + std::vector qScaleShape = {2, 39, 64}; + std::vector kScaleShape = {2, 16, 1}; + std::vector outShape = {2, 39, 1, 2048}; + std::vector valuesShape = {0}; + std::vector cuSeqQ; // empty means not provided + std::vector cuSeqK; + std::vector sequsedQ; + std::vector sequsedK = {2}; + std::vector cmpResidual; + std::vector blockTable = {2, 2}; + std::vector idxOffset; + std::vector metadata = {1024}; + ge::DataType qType = ge::DT_INT8; + ge::DataType kType = ge::DT_INT8; + ge::DataType wType = ge::DT_FLOAT16; + ge::DataType qScaleType = ge::DT_FLOAT16; + ge::DataType kScaleType = ge::DT_FLOAT16; + ge::DataType outType = ge::DT_INT32; + ge::DataType valuesType = ge::DT_BF16; + ge::DataType cuSeqQType = ge::DT_INT32; + ge::DataType cuSeqKType = ge::DT_INT32; + ge::DataType sequsedQType = ge::DT_INT32; + ge::DataType sequsedKType = ge::DT_INT32; + ge::DataType cmpResidualType = ge::DT_INT32; + ge::DataType blockTableType = ge::DT_INT32; + ge::DataType idxOffsetType = ge::DT_INT32; + std::string layoutQ = "BSND"; + std::string layoutK = "PA_BBND"; + int64_t topk = 2048; + int64_t quantMode = 2; + int64_t maxSeqlenQ = -1; + int64_t maskMode = 0; + int64_t cmpRatio = 1; + int64_t returnValue = 0; +}; + +inline gert::StorageShape ToStorageShape(const std::vector &dims) +{ + gert::StorageShape shape; + if (dims.empty()) { + return shape; + } + shape.MutableShape().SetDimNum(dims.size()); + shape.MutableStorageShape().SetDimNum(dims.size()); + for (size_t i = 0; i < dims.size(); i++) { + shape.MutableShape().SetDim(i, dims[i]); + shape.MutableStorageShape().SetDim(i, dims[i]); + } + return shape; +} + +inline gert::TilingContextPara::TensorDescription Desc(const std::vector &dims, ge::DataType dtype) +{ + return gert::TilingContextPara::TensorDescription(ToStorageShape(dims), dtype, ge::FORMAT_ND); +} + +inline void RunTilingCase(const CaseParam &p, ge::graphStatus expect) +{ + struct QLIV2CompileInfo { + } compileInfo; + gert::TilingContextPara para( + "QuantLightningIndexerV2", + {Desc(p.qShape, p.qType), Desc(p.kShape, p.kType), Desc(p.wShape, p.wType), Desc(p.qScaleShape, p.qScaleType), + Desc(p.kScaleShape, p.kScaleType), Desc(p.cuSeqQ, p.cuSeqQType), Desc(p.cuSeqK, p.cuSeqKType), + Desc(p.sequsedQ, p.sequsedQType), Desc(p.sequsedK, p.sequsedKType), Desc(p.cmpResidual, p.cmpResidualType), + Desc(p.blockTable, p.blockTableType), Desc(p.idxOffset, p.idxOffsetType), Desc(p.metadata, ge::DT_INT32)}, + {Desc(p.outShape, p.outType), Desc(p.valuesShape, p.valuesType)}, + {{"topk", Ops::Transformer::AnyValue::CreateFrom(p.topk)}, + {"quant_mode", Ops::Transformer::AnyValue::CreateFrom(p.quantMode)}, + {"max_seqlen_q", Ops::Transformer::AnyValue::CreateFrom(p.maxSeqlenQ)}, + {"layout_q", Ops::Transformer::AnyValue::CreateFrom(p.layoutQ)}, + {"layout_k", Ops::Transformer::AnyValue::CreateFrom(p.layoutK)}, + {"mask_mode", Ops::Transformer::AnyValue::CreateFrom(p.maskMode)}, + {"cmp_ratio", Ops::Transformer::AnyValue::CreateFrom(p.cmpRatio)}, + {"return_value", Ops::Transformer::AnyValue::CreateFrom(p.returnValue)}}, + &compileInfo, p.soc, p.coreNum, p.ubSize, p.l2Size); + ExecuteTestCase(para, expect, SKIP_TILING_KEY); +} + +} // namespace qliv2_ut + +#endif // TEST_QUANT_LIGHTNING_INDEXER_V2_UTILS_H diff --git a/csrc/attention/quant_lightning_indexer_v2/torch_extension/__init__.py b/csrc/attention/quant_lightning_indexer_v2/torch_extension/__init__.py new file mode 100644 index 000000000000..9c79a6acfa8e --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/torch_extension/__init__.py @@ -0,0 +1,16 @@ +__all__ = [ + "quant_lightning_indexer", + "quant_lightning_indexer_metadata", + "quant_lightning_indexer_candidate", + "quant_lightning_indexer_candidate_source", + "quant_lightning_indexer_candidate_consumer", +] + +from . import graph_convert_quant_lightning_indexer as graph_convert_quant_lightning_indexer +from .quant_lightning_indexer import ( + quant_lightning_indexer, + quant_lightning_indexer_candidate, + quant_lightning_indexer_candidate_consumer, + quant_lightning_indexer_candidate_source, + quant_lightning_indexer_metadata, +) diff --git a/csrc/attention/quant_lightning_indexer_v2/torch_extension/csrc/quant_lightning_indexer.cpp b/csrc/attention/quant_lightning_indexer_v2/torch_extension/csrc/quant_lightning_indexer.cpp new file mode 100644 index 000000000000..5c07ebce492c --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/torch_extension/csrc/quant_lightning_indexer.cpp @@ -0,0 +1,274 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file quant_lightning_indexer.cpp + * \brief + */ + +#include +#include "aclnn_common.h" + +namespace op_api { +using namespace at_npu::native; + +inline TensorWrapper MakeWrapper(const at::Tensor &tensor) +{ + return {tensor, ConvertToAclDataType(tensor.scalar_type())}; +} + +inline bool IsMxQuantMode(int64_t quantMode) { return quantMode == 3 || quantMode == 5; } + +inline bool IsE8M0Tensor(const at::Tensor &tensor) { return tensor.scalar_type() == at::kFloat8_e8m0fnu; } + +inline bool IsFp4CompatibleTensor(const at::Tensor &tensor) +{ + return tensor.scalar_type() == at::kFloat4_e2m1fn_x2 || tensor.scalar_type() == at::kByte; +} + +constexpr int64_t E8M0_SCALE_PACK_NUM = 2; + +inline void FixQLIV2AclDtypes(int64_t quantMode, TensorWrapper &queryWrapper, TensorWrapper &keyWrapper, + TensorWrapper &queryScaleWrapper, TensorWrapper &keyScaleWrapper) +{ + if (quantMode == 4) { + TORCH_CHECK(queryWrapper.tensor_.scalar_type() == at::kByte, "When quant_mode is 4, query must be hifp8 type"); + TORCH_CHECK(keyWrapper.tensor_.scalar_type() == at::kByte, "When quant_mode is 4, key must be hifp8 type"); + queryWrapper.dtype = ACL_HIFLOAT8; + keyWrapper.dtype = ACL_HIFLOAT8; + return; + } + + if (quantMode == 5) { + TORCH_CHECK(IsFp4CompatibleTensor(queryWrapper.tensor_), + "When quant_mode is 5, query must be torch.float4_e2m1fn_x2 or packed torch.uint8"); + TORCH_CHECK(IsFp4CompatibleTensor(keyWrapper.tensor_), + "When quant_mode is 5, key must be torch.float4_e2m1fn_x2 or packed torch.uint8"); + queryWrapper.dtype = ACL_FLOAT4_E2M1; + keyWrapper.dtype = ACL_FLOAT4_E2M1; + } + + if (IsMxQuantMode(quantMode)) { + TORCH_CHECK(IsE8M0Tensor(queryScaleWrapper.tensor_), + "When quant_mode is 3 or 5, query_dequant_scale must be torch.float8_e8m0fnu"); + TORCH_CHECK(IsE8M0Tensor(keyScaleWrapper.tensor_), + "When quant_mode is 3 or 5, key_dequant_scale must be torch.float8_e8m0fnu"); + // Cube loads E8M0 scales in two-byte groups through a bfloat16_t view. + TORCH_CHECK(queryScaleWrapper.tensor_.storage_offset() % E8M0_SCALE_PACK_NUM == 0, + "When quant_mode is 3 or 5, query_dequant_scale storage offset must satisfy 2-element E8M0 packing " + "alignment, but got ", + queryScaleWrapper.tensor_.storage_offset()); + TORCH_CHECK(keyScaleWrapper.tensor_.storage_offset() % E8M0_SCALE_PACK_NUM == 0, + "When quant_mode is 3 or 5, key_dequant_scale storage offset must satisfy 2-element E8M0 packing " + "alignment, but got ", + keyScaleWrapper.tensor_.storage_offset()); + queryScaleWrapper.dtype = ACL_FLOAT8_E8M0; + keyScaleWrapper.dtype = ACL_FLOAT8_E8M0; + } +} + +// npu tensor max size +const int SIZE = 8; +const int DIM_0 = 0; +const int DIM_1 = 1; +const int DIM_2 = 2; +const int DIM_3 = 3; + +constexpr int64_t QLI_V2_METADATA_SIZE = 1024; + +at::Tensor QuantLightningIndexerMetadata(int64_t numHeadsQ, int64_t numHeadsK, int64_t headDim, int64_t topk, + int64_t quantMode, const c10::optional &cuSeqlensQ, + const c10::optional &cuSeqlensK, + const c10::optional &sequsedQ, + const c10::optional &sequsedK, + const c10::optional &cmpResidualK, int64_t batchSize, + int64_t maxSeqlenQ, int64_t maxSeqlenK, c10::string_view layoutQ, + c10::string_view layoutK, int64_t maskMode, int64_t cmpRatio) +{ + at::Device outputDevice = at::Device(std::string("npu")); + if (cuSeqlensQ.has_value()) { + outputDevice = cuSeqlensQ.value().device(); + } else if (cuSeqlensK.has_value()) { + outputDevice = cuSeqlensK.value().device(); + } else if (sequsedQ.has_value()) { + outputDevice = sequsedQ.value().device(); + } else if (sequsedK.has_value()) { + outputDevice = sequsedK.value().device(); + } else if (cmpResidualK.has_value()) { + outputDevice = cmpResidualK.value().device(); + } + + at::Tensor output = torch::empty({QLI_V2_METADATA_SIZE}, torch::dtype(torch::kInt32).device(outputDevice)); + auto cuSeqlensQVal = get_valid_tensor(cuSeqlensQ, outputDevice); + auto cuSeqlensKVal = get_valid_tensor(cuSeqlensK, outputDevice); + auto sequsedQVal = get_valid_tensor(sequsedQ, outputDevice); + auto sequsedKVal = get_valid_tensor(sequsedK, outputDevice); + auto cmpResidualKVal = get_valid_tensor(cmpResidualK, outputDevice); + + std::string layoutQStr = std::string(layoutQ); + std::string layoutKStr = std::string(layoutK); + char *layoutQPtr = const_cast(layoutQStr.c_str()); + char *layoutKPtr = const_cast(layoutKStr.c_str()); + + ACLNN_CMD(aclnnQuantLightningIndexerV2Metadata, cuSeqlensQVal, cuSeqlensKVal, sequsedQVal, sequsedKVal, + cmpResidualKVal, numHeadsQ, numHeadsK, headDim, topk, quantMode, batchSize, maxSeqlenQ, maxSeqlenK, + layoutQPtr, layoutKPtr, maskMode, cmpRatio, output); + return output; +} + +// 工具函数,推导输出shape +std::tuple ConstructQuantLightningIndexerOutputTensor( + const at::Tensor &query, const at::Tensor &key, int64_t sparseCount, std::string queryLayoutStr, + std::string keyLayoutStr, int64_t returnValue) +{ + at::SmallVector outputSize; + for (size_t i = 0; i < query.sizes().size(); i++) { + TORCH_CHECK(query.size(i) > 0, + "All values within query's shape should be greater " + "than 0, but shape[", + i, "] is ", query.size(i)); + } + for (size_t i = 0; i < key.sizes().size(); i++) { + TORCH_CHECK(key.size(i) > 0, + "All values within key's shape should be greater " + "than 0, but shape[", + i, "] is ", key.size(i)); + } + TORCH_CHECK(sparseCount > 0, "sparse count should be greater than 0, but now is ", sparseCount); + int64_t keyHeadNum = (keyLayoutStr == "TND") ? key.size(DIM_1) : key.size(DIM_2); + if (queryLayoutStr == "BSND") { + outputSize = {query.size(DIM_0), query.size(DIM_1), keyHeadNum, sparseCount}; + } else { + int nDimIndex = 0; + nDimIndex = (keyLayoutStr == "TND") ? DIM_1 : DIM_2; + outputSize = {query.size(DIM_0), key.size(nDimIndex), sparseCount}; + } + at::Tensor sparseIndicesOut = at::empty(outputSize, query.options().dtype(at::kInt)); + at::Tensor sparseValuesOut; + if (returnValue) { + sparseValuesOut = at::empty(outputSize, query.options().dtype(at::kBFloat16)); + } else { + sparseValuesOut = at::empty({0}, query.options().dtype(at::kBFloat16)); + } + + return std::tuple(sparseIndicesOut, sparseValuesOut); +} + +std::tuple QuantLightningIndexer( + const at::Tensor &query, const at::Tensor &key, const at::Tensor &weights, const at::Tensor &queryDequantScale, + const at::Tensor &keyDequantScale, int64_t topk, int64_t quantMode, const c10::optional &cuSeqlensQ, + const c10::optional &cuSeqlensK, const c10::optional &sequsedQ, + const c10::optional &sequsedK, const c10::optional &cmpResidualK, + const c10::optional &blockTable, const c10::optional &outputIdxOffset, + const c10::optional &metadata, int64_t maxSeqlenQ, c10::string_view layoutQ, c10::string_view layoutK, + int64_t maskMode, int64_t cmpRatio, int64_t returnValue) +{ + TORCH_CHECK(query.numel() > 0, "Tensor query is empty.") + TORCH_CHECK(key.numel() > 0, "Tensor key is empty.") + + std::string queryLayoutStr = std::string(layoutQ); + std::string keyLayoutStr = std::string(layoutK); + + // construct the output tensor + std::tuple quantLightningIndexerOutput = + ConstructQuantLightningIndexerOutputTensor(query, key, topk, queryLayoutStr, keyLayoutStr, returnValue); + at::Tensor sparseIndicesOut = std::get<0>(quantLightningIndexerOutput); + at::Tensor sparseValuesOut = std::get<1>(quantLightningIndexerOutput); + // convert str + char *queryLayoutPtr = const_cast(queryLayoutStr.c_str()); + char *keyLayoutPtr = const_cast(keyLayoutStr.c_str()); + + auto queryWrapper = MakeWrapper(query); + auto keyWrapper = MakeWrapper(key); + auto queryScaleWrapper = MakeWrapper(queryDequantScale); + auto keyScaleWrapper = MakeWrapper(keyDequantScale); + FixQLIV2AclDtypes(quantMode, queryWrapper, keyWrapper, queryScaleWrapper, keyScaleWrapper); + + int64_t keyStride0Disabled = 0; // A11: 旧入口不启用 stride 显式属性 (保持现网行为) + int64_t keyScaleStride0Disabled = 0; + ACLNN_CMD(aclnnQuantLightningIndexerV2, queryWrapper, keyWrapper, weights, queryScaleWrapper, keyScaleWrapper, + cuSeqlensQ, cuSeqlensK, sequsedQ, sequsedK, cmpResidualK, blockTable, outputIdxOffset, metadata, topk, + quantMode, maxSeqlenQ, queryLayoutPtr, keyLayoutPtr, maskMode, cmpRatio, returnValue, + keyStride0Disabled, keyScaleStride0Disabled, sparseIndicesOut, + sparseValuesOut); + + return std::tuple(sparseIndicesOut, sparseValuesOut); +} + +// 两级TopK candidate 接口 (O1 方案b: 新增入口, 旧接口不变) +// candidate_mode: 1=source 输出 candidate_topk_index; 2=consumer 输入候选块; 3=关闭 +std::tuple QuantLightningIndexerCandidate( + const at::Tensor &query, const at::Tensor &key, const at::Tensor &weights, const at::Tensor &queryDequantScale, + const at::Tensor &keyDequantScale, int64_t topk, int64_t quantMode, + const c10::optional &candidateTopkIndexIn, const c10::optional &cuSeqlensQ, + const c10::optional &cuSeqlensK, const c10::optional &sequsedQ, + const c10::optional &sequsedK, const c10::optional &cmpResidualK, + const c10::optional &blockTable, const c10::optional &outputIdxOffset, + const c10::optional &metadata, int64_t maxSeqlenQ, c10::string_view layoutQ, + c10::string_view layoutK, int64_t maskMode, int64_t cmpRatio, int64_t candidateMode, + int64_t candidateTopkBlocks, int64_t candidateBlockSize) +{ + TORCH_CHECK(query.numel() > 0, "Tensor query is empty.") + TORCH_CHECK(key.numel() > 0, "Tensor key is empty.") + + std::string queryLayoutStr = std::string(layoutQ); + std::string keyLayoutStr = std::string(layoutK); + + std::tuple quantLightningIndexerOutput = + ConstructQuantLightningIndexerOutputTensor(query, key, topk, queryLayoutStr, keyLayoutStr, 0); + at::Tensor sparseIndicesOut = std::get<0>(quantLightningIndexerOutput); + at::Tensor sparseValuesOut = std::get<1>(quantLightningIndexerOutput); + + int64_t keyHeadNum = (keyLayoutStr == "TND") ? key.size(DIM_1) : key.size(DIM_2); + at::Tensor candidateTopkIndexOut; + if (candidateMode == 1) { + at::SmallVector candSize; + if (queryLayoutStr == "BSND") { + candSize = {query.size(DIM_0), query.size(DIM_1), keyHeadNum, candidateTopkBlocks}; + } else { + candSize = {query.size(DIM_0), keyHeadNum, candidateTopkBlocks}; + } + candidateTopkIndexOut = at::empty(candSize, query.options().dtype(at::kInt)); + } else { + candidateTopkIndexOut = at::empty({0}, query.options().dtype(at::kInt)); + } + + char *queryLayoutPtr = const_cast(queryLayoutStr.c_str()); + char *keyLayoutPtr = const_cast(keyLayoutStr.c_str()); + int64_t returnValue = 0; + + auto queryWrapper = MakeWrapper(query); + auto keyWrapper = MakeWrapper(key); + auto queryScaleWrapper = MakeWrapper(queryDequantScale); + auto keyScaleWrapper = MakeWrapper(keyDequantScale); + FixQLIV2AclDtypes(quantMode, queryWrapper, keyWrapper, queryScaleWrapper, keyScaleWrapper); + + // A11: key 0 轴非连续 — aclnn 动态调用下 tiling 拿不到 tensor stride (仅 TensorV2/图模式可见), + // 从 key/k_scale 的 torch stride(0) 自动显式传入 (紧凑存储时等于紧凑值, 走 kernel 兜底语义) + int64_t keyStride0 = key.stride(0); + int64_t keyScaleStride0 = keyDequantScale.stride(0); + + ACLNN_CMD(aclnnQuantLightningIndexerV2, queryWrapper, keyWrapper, weights, queryScaleWrapper, keyScaleWrapper, + cuSeqlensQ, cuSeqlensK, sequsedQ, sequsedK, cmpResidualK, blockTable, outputIdxOffset, + metadata, candidateTopkIndexIn, topk, quantMode, maxSeqlenQ, queryLayoutPtr, keyLayoutPtr, maskMode, + cmpRatio, returnValue, candidateMode, candidateTopkBlocks, candidateBlockSize, keyStride0, + keyScaleStride0, sparseIndicesOut, sparseValuesOut, candidateTopkIndexOut); + + return std::tuple(sparseIndicesOut, sparseValuesOut, candidateTopkIndexOut); +} +// Bind the C++ function to Python module +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) +{ + m.def("quant_lightning_indexer_metadata", &QuantLightningIndexerMetadata, "quant_lightning_indexer_metadata"); + m.def("quant_lightning_indexer", &QuantLightningIndexer, "quant_lightning_indexer"); + m.def("quant_lightning_indexer_candidate", &QuantLightningIndexerCandidate, + "quant_lightning_indexer_candidate"); +} +} // namespace op_api diff --git a/csrc/attention/quant_lightning_indexer_v2/torch_extension/graph_convert_quant_lightning_indexer.py b/csrc/attention/quant_lightning_indexer_v2/torch_extension/graph_convert_quant_lightning_indexer.py new file mode 100644 index 000000000000..e6dfda29acbb --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/torch_extension/graph_convert_quant_lightning_indexer.py @@ -0,0 +1,78 @@ +# ruff: noqa +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +# GE Converter for Graph Mode + +try: + from collections.abc import Callable + from typing import Any, Dict, List, Optional, Tuple, Union + + import torch + import torch_npu + import torchair + from torch.library import impl + from torchair._ge_concrete_graph import ge_apis as ge + from torchair._ge_concrete_graph.compat_ir import IrDef, ge_op + from torchair._ge_concrete_graph.fx2ge_converter import ( + declare_supported, + register_fx_node_ge_converter, + ) + from torchair._ge_concrete_graph.ge_ir_pb2 import ( + GraphDef, + OpDef, + TensorDef, + TensorDescriptor, + ) + from torchair._ge_concrete_graph.supported_declaration import Support + from torchair.ge import attr + from torchair.ge._ge_graph import ( + DataType, + Tensor, + TensorSpec, + TensorType, + auto_convert_to_tensor, + compat_as_bytes, + compat_as_bytes_list, + get_default_ge_graph, + get_invalid_desc, + next_unique_name, + trans_to_list_list_float, + trans_to_list_list_int, + ) + + _TORCHAIR_AVAILABLE = True +except ImportError: + _TORCHAIR_AVAILABLE = False + +if _TORCHAIR_AVAILABLE: + + @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata.default) + def convert_quant_lightning_indexer_metadata( + num_heads_q: int, + num_heads_kv: int, + head_dim: int, + topk: int, + quant_mode: int, + *, + cu_seqlens_q: Tensor | None = None, + cu_seqlens_k: Tensor | None = None, + seqused_q: Tensor | None = None, + seqused_k: Tensor | None = None, + cmp_residual_k: Tensor | None = None, + batch_size: int | None = None, + max_seqlen_q: int | None = None, + max_seqlen_k: int | None = None, + layout_q: str | None = None, + layout_k: str | None = None, + mask_mode: int | None = None, + cmp_ratio: int | None = None, + meta_outputs: TensorSpec = None, + ): + raise RuntimeError("GE converter doesn't support op: 'quant_lightning_indexer_metadata'") diff --git a/csrc/attention/quant_lightning_indexer_v2/torch_extension/quant_lightning_indexer.py b/csrc/attention/quant_lightning_indexer_v2/torch_extension/quant_lightning_indexer.py new file mode 100644 index 000000000000..2feb5c34ab2a --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2/torch_extension/quant_lightning_indexer.py @@ -0,0 +1,680 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import torch +from cann_ops_transformer.op_builder import OpBuilder, get_as_library +from torch.library import impl + +QLI_METADATA_SIZE = 1024 +QLI_METADATA_OP_NAME = "quant_lightning_indexer_metadata" + + +class QuantLightningIndexerOpBuilder(OpBuilder): + def __init__(self): + super().__init__("quant_lightning_indexer", category="attention") + + def sources(self): + """Path to C++ source code.""" + return ["csrc/attention/quant_lightning_indexer.cpp"] + + def schema(self) -> str: + """PyTorch operator signature.""" + return [ + "quant_lightning_indexer_metadata(int num_heads_q, int num_heads_k, int head_dim, int topk, " + "int quant_mode, *, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None, Tensor? seqused_q=None," + "Tensor? seqused_k=None, Tensor? cmp_residual_k=None, int? batch_size=None, int? max_seqlen_q=None," + "int? max_seqlen_k=None, str? layout_q=None, str? layout_k=None, int? mask_mode=None, " + "int? cmp_ratio=None) -> Tensor", + "quant_lightning_indexer(Tensor query, Tensor key, Tensor weights, Tensor query_dequant_scale, " + "Tensor key_dequant_scale, int topk, int quant_mode, *, Tensor? cu_seqlens_q=None, " + "Tensor? cu_seqlens_k=None, Tensor? seqused_q=None, Tensor? seqused_k=None, Tensor? " + "cmp_residual_k = None, Tensor? block_table=None, Tensor? output_idx_offset=None, Tensor? metadata=None, " + 'int max_seqlen_q=-1, str layout_q="BSND", str layout_k="BSND", int mask_mode=0, ' + "int cmp_ratio=1, int return_value=0) -> (Tensor, Tensor)", + # O1 方案b: candidate 两级TopK 新入口 (三元组), 旧 schema 保持不变 + "quant_lightning_indexer_candidate(Tensor query, Tensor key, Tensor weights, " + "Tensor query_dequant_scale, Tensor key_dequant_scale, int topk, int quant_mode, *, " + "Tensor? candidate_topk_index=None, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None, " + "Tensor? seqused_q=None, Tensor? seqused_k=None, Tensor? cmp_residual_k=None, " + "Tensor? block_table=None, Tensor? output_idx_offset=None, Tensor? metadata=None, " + 'int max_seqlen_q=-1, str layout_q="BSND", str layout_k="BSND", int mask_mode=0, ' + "int cmp_ratio=1, int candidate_mode=3, int candidate_topk_blocks=2048, " + "int candidate_block_size=8) -> (Tensor, Tensor, Tensor)", + ] + + def register_meta(self): + """ + Registers the Meta implementation (Shape/Dtype inference). + Essential for Autograd and FakeTensor support. + """ + + @torch.library.register_fake("cann_ops_transformer::" + QLI_METADATA_OP_NAME) + def quant_lightning_indexer_metadata_meta( + num_heads_q: int, + num_heads_k: int, + head_dim: int, + topk: int, + quant_mode: int, + cu_seqlens_q: torch.Tensor | None = None, + cu_seqlens_k: torch.Tensor | None = None, + seqused_q: torch.Tensor | None = None, + seqused_k: torch.Tensor | None = None, + cmp_residual_k: torch.Tensor | None = None, + batch_size: int | None = None, + max_seqlen_q: int | None = None, + max_seqlen_k: int | None = None, + layout_q: str | None = None, + layout_k: str | None = None, + mask_mode: int | None = None, + cmp_ratio: int | None = None, + ): + return torch.empty((QLI_METADATA_SIZE), dtype=torch.int32, device="npu") + + @impl(get_as_library(), self.name, "Meta") + def quant_lightning_indexer_meta( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + *, + cu_seqlens_q=None, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=None, + cmp_residual_k=None, + block_table=None, + output_idx_offset=None, + metadata=None, + max_seqlen_q=-1, + layout_q="BSND", + layout_k="BSND", + mask_mode=0, + cmp_ratio=1, + return_value=0, + ): + key_head_num = key.shape[1] if layout_k == "TND" else key.shape[2] + + if layout_q == "BSND": + sparse_indices_out = torch.empty( + [query.shape[0], query.shape[1], key_head_num, topk], + dtype=torch.int32, + device="meta", + ) + else: + sparse_indices_out = torch.empty( + [query.shape[0], key_head_num, topk], + dtype=torch.int32, + device="meta", + ) + if return_value: + if layout_q == "BSND": + sparse_values_out = torch.empty( + [query.shape[0], query.shape[1], key_head_num, topk], + dtype=torch.bfloat16, + device="meta", + ) + else: + sparse_values_out = torch.empty( + [query.shape[0], key_head_num, topk], + dtype=torch.bfloat16, + device="meta", + ) + else: + sparse_values_out = torch.empty([0], dtype=torch.bfloat16, device="meta") + return (sparse_indices_out, sparse_values_out) + + @torch.library.register_fake("cann_ops_transformer::quant_lightning_indexer_candidate") + def quant_lightning_indexer_candidate_meta( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + *, + candidate_topk_index=None, + cu_seqlens_q=None, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=None, + cmp_residual_k=None, + block_table=None, + output_idx_offset=None, + metadata=None, + max_seqlen_q=-1, + layout_q="BSND", + layout_k="BSND", + mask_mode=0, + cmp_ratio=1, + candidate_mode=3, + candidate_topk_blocks=2048, + candidate_block_size=8, + ): + key_head_num = key.shape[1] if layout_k == "TND" else key.shape[2] + if layout_q == "BSND": + sparse_indices_out = torch.empty( + [query.shape[0], query.shape[1], key_head_num, topk], + dtype=torch.int32, + device="meta", + ) + else: + sparse_indices_out = torch.empty( + [query.shape[0], key_head_num, topk], + dtype=torch.int32, + device="meta", + ) + sparse_values_out = torch.empty([0], dtype=torch.bfloat16, device="meta") + if candidate_mode == 1: + if layout_q == "BSND": + cand_out = torch.empty( + [query.shape[0], query.shape[1], key_head_num, candidate_topk_blocks], + dtype=torch.int32, + device="meta", + ) + else: + cand_out = torch.empty( + [query.shape[0], key_head_num, candidate_topk_blocks], + dtype=torch.int32, + device="meta", + ) + else: + cand_out = torch.empty([0], dtype=torch.int32, device="meta") + return (sparse_indices_out, sparse_values_out, cand_out) + + +# Instantiate the builder +quant_lightning_indexer_op_builder = QuantLightningIndexerOpBuilder() +quant_lightning_indexer_op_builder._ensure_initialized() + + +@impl(get_as_library(), QLI_METADATA_OP_NAME, "PrivateUse1") +def quant_lightning_indexer_metadata( + num_heads_q: int, + num_heads_k: int, + head_dim: int, + topk: int, + quant_mode: int, + cu_seqlens_q: torch.Tensor | None = None, + cu_seqlens_k: torch.Tensor | None = None, + seqused_q: torch.Tensor | None = None, + seqused_k: torch.Tensor | None = None, + cmp_residual_k: torch.Tensor | None = None, + batch_size: int | None = None, + max_seqlen_q: int | None = None, + max_seqlen_k: int | None = None, + layout_q: str | None = None, + layout_k: str | None = None, + mask_mode: int | None = None, + cmp_ratio: int | None = None, +): + """ + dispatcher implementation for NPU.zhe + 'PrivateUse1' is the combine key for custom NPU backends. + """ + batch_size = 0 if batch_size is None else batch_size + max_seqlen_q = -1 if max_seqlen_q is None else max_seqlen_q + max_seqlen_k = -1 if max_seqlen_k is None else max_seqlen_k + layout_q = "BSND" if layout_q is None else layout_q + layout_k = "BSND" if layout_k is None else layout_k + mask_mode = 0 if mask_mode is None else mask_mode + cmp_ratio = 1 if cmp_ratio is None else cmp_ratio + + op_module = quant_lightning_indexer_op_builder.load() + return op_module.quant_lightning_indexer_metadata( + num_heads_q, + num_heads_k, + head_dim, + topk, + quant_mode, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + batch_size, + max_seqlen_q, + max_seqlen_k, + layout_q, + layout_k, + mask_mode, + cmp_ratio, + ) + + +@torch.library.register_kernel("cann_ops_transformer::" + QLI_METADATA_OP_NAME, None) +def quant_lightning_indexer_metadata_fallback( + num_heads_q: int, + num_heads_k: int, + head_dim: int, + topk: int, + quant_mode: int, + cu_seqlens_q: torch.Tensor | None = None, + cu_seqlens_k: torch.Tensor | None = None, + seqused_q: torch.Tensor | None = None, + seqused_k: torch.Tensor | None = None, + cmp_residual_k: torch.Tensor | None = None, + batch_size: int | None = None, + max_seqlen_q: int | None = None, + max_seqlen_k: int | None = None, + layout_q: str | None = None, + layout_k: str | None = None, + mask_mode: int | None = None, + cmp_ratio: int | None = None, +): + # 处理所有 tensor 都为 None 的情况 + # 调用 NPU 实现 + return quant_lightning_indexer_metadata( + num_heads_q, + num_heads_k, + head_dim, + topk, + quant_mode, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + batch_size, + max_seqlen_q, + max_seqlen_k, + layout_q, + layout_k, + mask_mode, + cmp_ratio, + ) + + +torch.compiler.allow_in_graph(quant_lightning_indexer_metadata) + + +@impl(get_as_library(), "quant_lightning_indexer_candidate", "PrivateUse1") +def quant_lightning_indexer_candidate( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + *, + candidate_topk_index=None, + cu_seqlens_q=None, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=None, + cmp_residual_k=None, + block_table=None, + output_idx_offset=None, + metadata=None, + max_seqlen_q=-1, + layout_q="BSND", + layout_k="BSND", + mask_mode=0, + cmp_ratio=1, + candidate_mode=3, + candidate_topk_blocks=2048, + candidate_block_size=8, +): + """两级TopK candidate 入口: mode=1(source)/2(consumer)/3(关闭)""" + op_module = quant_lightning_indexer_op_builder.load() + return op_module.quant_lightning_indexer_candidate( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + candidate_topk_index, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + block_table, + output_idx_offset, + metadata, + max_seqlen_q, + layout_q, + layout_k, + mask_mode, + cmp_ratio, + candidate_mode, + candidate_topk_blocks, + candidate_block_size, + ) + + +torch.compiler.allow_in_graph(quant_lightning_indexer_candidate) + + +@impl(get_as_library(), quant_lightning_indexer_op_builder.name, "PrivateUse1") +def quant_lightning_indexer( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + *, + cu_seqlens_q=None, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=None, + cmp_residual_k=None, + block_table=None, + output_idx_offset=None, + metadata=None, + max_seqlen_q=-1, + layout_q="BSND", + layout_k="BSND", + mask_mode=0, + cmp_ratio=1, + return_value=0, +): + """ + dispatcher implementation for NPU.zhe + 'PrivateUse1' is the combine key for custom NPU backends. + """ + op_module = quant_lightning_indexer_op_builder.load() + return op_module.quant_lightning_indexer( + query, + key, + weights, + query_dequant_scale, + key_dequant_scale, + topk, + quant_mode, + cu_seqlens_q, + cu_seqlens_k, + seqused_q, + seqused_k, + cmp_residual_k, + block_table, + output_idx_offset, + metadata, + max_seqlen_q, + layout_q, + layout_k, + mask_mode, + cmp_ratio, + return_value, + ) + + +# ============================================================================================= +# 新接口规范适配层 (2026-09-09): py 分发封装, 按新芯片接口规范收参, 内部映射到已注册的 +# quant_lightning_indexer_candidate (分别固定 candidate_mode=1/2); 不改 csrc / 算子校验 / 数据类型。 +# 调用侧按芯片选择接口 (本后端走本适配层, 新芯片走其自带实现)。 +# +# 与内部接口的映射约定: +# q_descale / k_descale <-> query_dequant_scale / key_dequant_scale (仅改名) +# candidate_block_indices <-> candidate_topk_index (仅改名; 块级索引, 相对块号) +# seqused_q (新名, 每 batch key 有效长度, 即内部 seqused_k) <-> seqused_k +# candidate_block_length mode=1 输出: py 按行级有效长度公式生成 (mask 规则与 key 长度取小); +# mode=2 输入: 本后端忽略 (算子内部已按 mask 规则与 key 长度取小截断, 语义冗余) +# candidate_topk_blocks=-1 (规范默认, "无 candidate 机制") 在 source/consumer 场景下无意义, 取 2048 +# layout 参数消失 q 恒 3 维 (T1, N1, D): 有 cu_seqlens_q 视为 TND 变长拼接, 否则 B=1 BSND +# ============================================================================================= + + +def _qli_newapi_check(cond, msg): + if not cond: + raise RuntimeError("[quant_lightning_indexer new-api] " + msg) + + +def _qli_newapi_layout(cu_seqlens_q): + # 新规范无 layout 参数: 由 cu_seqlens_q 有无推断 (有 = TND 变长拼接; 无 = B=1 BSND) + return ("TND", True) if cu_seqlens_q is not None else ("BSND", False) + + +def _qli_newapi_build_metadata( + q, k, layout_q, layout_k, cu_seqlens_q, seqused_k, cmp_residual_k, topk, quant_mode, mask_mode, cmp_ratio +): + # 本后端主算子 metadata 必传; 新规范调用方不传时按 shape 自动推导。 + # 注: max_seqlen_q/k 需要标量, 此处 .item() 触发一次主机同步 (调用方可自行预生成 metadata 传入绕过)。 + num_heads_q = int(q.shape[1]) # N1 (q 恒 3 维) + num_heads_k = int(k.shape[2]) # N2 (k 为 PA 物理池布局 (block_num, block_size, N2, D)) + head_dim = int(q.shape[2]) + if layout_q == "TND": + batch_size = int(cu_seqlens_q.numel()) - 1 + max_seqlen_q = int((cu_seqlens_q[1:] - cu_seqlens_q[:-1]).max().item()) + else: + batch_size = 1 + max_seqlen_q = int(q.shape[0]) + max_seqlen_k = int(seqused_k.max().item()) + return quant_lightning_indexer_metadata( + num_heads_q, + num_heads_k, + head_dim, + topk, + quant_mode, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=seqused_k, + cmp_residual_k=cmp_residual_k, + batch_size=batch_size, + max_seqlen_q=max_seqlen_q, + max_seqlen_k=max_seqlen_k, + layout_q=layout_q, + layout_k=layout_k, + mask_mode=mask_mode, + cmp_ratio=cmp_ratio, + ) + + +def _qli_newapi_block_length(q, k, cu_seqlens_q, seqused_k, cmp_residual_k, cmp_ratio, mask_mode): + # source 第 4 输出 (T1, N2): 每行实际参与搜索的 key 长度 (mask 规则与 key 长度取小)。 + # mask_mode=3 (causal): vl(i) = clamp((act_k - S1 + i + 1) // cmp_ratio, 0, K), + # act_k = K * cmp_ratio + residual 为未压缩原始 key 长度, i 为 batch 内行号 (与 golden 同公式); + # mask_mode=0 (无 mask): 恒为 K。全程张量化, 不引入主机同步。 + device = q.device + t1 = int(q.shape[0]) + n2 = int(k.shape[2]) + s2_t = seqused_k.to(torch.int64) # (B,) 每 batch key 有效长度 + res_t = cmp_residual_k.to(torch.int64) if cmp_residual_k is not None else torch.zeros_like(s2_t) + act_k_t = s2_t * int(cmp_ratio) + res_t + if cu_seqlens_q is not None: # TND 变长拼接 + cu_q = cu_seqlens_q.to(torch.int64) + row_b = torch.repeat_interleave(torch.arange(s2_t.numel(), device=device), cu_q[1:] - cu_q[:-1]) + i_in_b = torch.arange(t1, device=device) - cu_q[row_b] # batch 内行号 + s1_row = (cu_q[1:] - cu_q[:-1])[row_b] # 所在 batch 的 query 行数 + else: # B=1 BSND + row_b = torch.zeros(t1, dtype=torch.int64, device=device) + i_in_b = torch.arange(t1, device=device) + s1_row = torch.full((t1,), t1, dtype=torch.int64, device=device) + if int(mask_mode) == 3: + vl = (act_k_t[row_b] - s1_row + i_in_b + 1) // int(cmp_ratio) + vl = torch.clamp(vl, min=0) # 全无效行 (vl < 0) 记 0 + vl = torch.minimum(vl, s2_t[row_b]) # 与 key 长度取小 + else: + vl = s2_t[row_b] + return vl.to(torch.int32).unsqueeze(-1).expand(t1, n2).contiguous() + + +def quant_lightning_indexer_candidate_source( + q, + k, + w, + q_descale, + k_descale, + topk, + quant_mode, + *, + cu_seqlens_q=None, + seqused_q=None, + cmp_residual_k=None, + block_table=None, + metadata=None, + max_seqlen_q=-1, + mask_mode=0, + cmp_ratio=1, + return_value=False, + candidate_topk_blocks=-1, + candidate_block_size=-1, +): + """新规范接口1 (source, 内部 candidate_mode=1): 输出候选块索引, 照常输出 sparse topk。 + + q: (T1, N1, D); k: PA 物理池 (block_num, block_size, N2, D), 经 block_table 寻址; + seqused_q: (B,) 每 batch key 有效长度 (按新规范注释, 语义为 key 截断, 即内部 seqused_k); + 返回 (sparse_indices (T1,N2,topk), sparse_values, candidate_block_indices (T1,N2,cb), + candidate_block_length (T1,N2))。 + 本后端限制: block_table 必传 (key 仅支持 PA_BBND 分页布局); return_value=True 不支持 + (内部入口恒 False); candidate_topk_blocks / candidate_block_size 传 -1 (默认) 时取 2048 / 8。 + """ + _qli_newapi_check( + block_table is not None, "block_table is required on this backend (key only supports PA_BBND paged layout)." + ) + _qli_newapi_check(not return_value, "return_value=True is not supported on this backend.") + _qli_newapi_check( + seqused_q is not None, + "seqused_q (per-batch key valid length) is required: candidate_block_length " + "computation and varlen dispatch both depend on it.", + ) + _qli_newapi_check( + cu_seqlens_q is not None or int(seqused_q.numel()) == 1, + "batch > 1 requires cu_seqlens_q (TND varlen layout); without it q is treated " + "as B=1 BSND and the batch dimension would be lost.", + ) + topk_blocks = 2048 if candidate_topk_blocks in (-1, None) else int(candidate_topk_blocks) + block_size = 8 if candidate_block_size in (-1, None) else int(candidate_block_size) + layout_q, is_tnd = _qli_newapi_layout(cu_seqlens_q) + if is_tnd: # TND: q/w/q_descale 均无 batch 维, 直传 + q_in, w_in, q_descale_in = q, w, q_descale + else: # B=1 BSND: 内部入口需要 batch 维 + _qli_newapi_check(q.dim() == 3, "q must be 3-D (T1, N1, D).") + q_in, w_in, q_descale_in = q.unsqueeze(0), w.unsqueeze(0), q_descale.unsqueeze(0) + if metadata is None: # 本后端主算子 metadata 必传, 自动推导 + metadata = _qli_newapi_build_metadata( + q, k, layout_q, "PA_BBND", cu_seqlens_q, seqused_q, cmp_residual_k, topk, quant_mode, mask_mode, cmp_ratio + ) + idx, vals, cand = quant_lightning_indexer_candidate( + q_in, + k, + w_in, + q_descale_in, + k_descale, + topk, + quant_mode, + candidate_topk_index=None, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=seqused_q, + cmp_residual_k=cmp_residual_k, + block_table=block_table, + output_idx_offset=None, + metadata=metadata, + max_seqlen_q=max_seqlen_q, + layout_q=layout_q, + layout_k="PA_BBND", + mask_mode=mask_mode, + cmp_ratio=cmp_ratio, + candidate_mode=1, + candidate_topk_blocks=topk_blocks, + candidate_block_size=block_size, + ) + if not is_tnd: + idx, cand = idx.squeeze(0), cand.squeeze(0) # (1,T1,N2,*) -> (T1,N2,*) + block_length = _qli_newapi_block_length(q, k, cu_seqlens_q, seqused_q, cmp_residual_k, cmp_ratio, mask_mode) + return idx, vals, cand, block_length + + +def quant_lightning_indexer_candidate_consumer( + q, + k, + w, + q_descale, + k_descale, + candidate_block_indices, + candidate_block_length, + topk, + quant_mode, + *, + cu_seqlens_q=None, + seqused_q=None, + cmp_residual_k=None, + block_table=None, + output_idx_offset=None, + metadata=None, + max_seqlen_q=-1, + mask_mode=0, + cmp_ratio=1, + return_value=False, + candidate_block_size=8, +): + """新规范接口2 (consumer, 内部 candidate_mode=2): 在候选块内选 topk。 + + candidate_block_indices: (T1, N2, cb) source 侧输出的候选块索引 (块级, 相对块号); + candidate_block_length: 本后端忽略 (算子内部已按 mask 规则与 key 长度取小截断, 该输入语义冗余); + candidate_topk_blocks 由 candidate_block_indices.shape[-1] 推导 (新规范无此属性)。 + 返回 (sparse_indices, sparse_values); 布局与内部一致: B=1 BSND (B,S1,N2,topk) / TND (T1,N2,topk)。 + """ + _qli_newapi_check( + block_table is not None, "block_table is required on this backend (key only supports PA_BBND paged layout)." + ) + _qli_newapi_check(not return_value, "return_value=True is not supported on this backend.") + _qli_newapi_check(candidate_block_indices is not None, "candidate_block_indices is required (consumer input).") + _qli_newapi_check( + candidate_block_indices.dim() == 3, "candidate_block_indices must be 3-D (T1, N2, candidate_topk_blocks)." + ) + _qli_newapi_check( + seqused_q is not None, "seqused_q (per-batch key valid length) is required: varlen dispatch depends on it." + ) + _qli_newapi_check( + cu_seqlens_q is not None or int(seqused_q.numel()) == 1, + "batch > 1 requires cu_seqlens_q (TND varlen layout); without it q is treated " + "as B=1 BSND and the batch dimension would be lost.", + ) + topk_blocks = int(candidate_block_indices.shape[-1]) + layout_q, is_tnd = _qli_newapi_layout(cu_seqlens_q) + if is_tnd: + q_in, w_in, q_descale_in, cand_in = q, w, q_descale, candidate_block_indices + else: + _qli_newapi_check(q.dim() == 3, "q must be 3-D (T1, N1, D).") + q_in, w_in, q_descale_in = q.unsqueeze(0), w.unsqueeze(0), q_descale.unsqueeze(0) + cand_in = candidate_block_indices.unsqueeze(0) # (1, T1, N2, cb) + if metadata is None: + metadata = _qli_newapi_build_metadata( + q, k, layout_q, "PA_BBND", cu_seqlens_q, seqused_q, cmp_residual_k, topk, quant_mode, mask_mode, cmp_ratio + ) + idx, vals, _ = quant_lightning_indexer_candidate( + q_in, + k, + w_in, + q_descale_in, + k_descale, + topk, + quant_mode, + candidate_topk_index=cand_in, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=seqused_q, + cmp_residual_k=cmp_residual_k, + block_table=block_table, + output_idx_offset=output_idx_offset, + metadata=metadata, + max_seqlen_q=max_seqlen_q, + layout_q=layout_q, + layout_k="PA_BBND", + mask_mode=mask_mode, + cmp_ratio=cmp_ratio, + candidate_mode=2, + candidate_topk_blocks=topk_blocks, + candidate_block_size=int(candidate_block_size), + ) + return idx, vals diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2_metadata/CMakeLists.txt index ed13e887f3cb..90d9ad1f2419 100644 --- a/csrc/attention/quant_lightning_indexer_v2_metadata/CMakeLists.txt +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/CMakeLists.txt @@ -8,4 +8,10 @@ # See LICENSE in the root of the software repository for the full text of the License. # ----------------------------------------------------------------------------------------------------------- -add_modules_sources_aicpu(OPTYPE quant_lightning_indexer_v2_metadata ACLNNTYPE aclnn) +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +list(REMOVE_ITEM CURRENT_DIRS tests) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/docs/aclnnQuantLightningIndexerV2Metadata.md b/csrc/attention/quant_lightning_indexer_v2_metadata/docs/aclnnQuantLightningIndexerV2Metadata.md new file mode 100644 index 000000000000..bd413e818f9e --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/docs/aclnnQuantLightningIndexerV2Metadata.md @@ -0,0 +1,743 @@ +# aclnnQuantLightningIndexerV2Metadata + +[📄 查看源码](https://gitcode.com/cann/ops-transformer/tree/master/attention/quant_lightning_indexer_v2_metadata) + +## 产品支持情况 + + +- Ascend 950PR/Ascend 950DT:支持 + + +- Atlas A3 训练系列产品/Atlas A3 推理系列产品:支持 + + +- Atlas A2 训练系列产品/Atlas A2 推理系列产品:支持 + + +- Atlas 200I/500 A2 推理产品:x + + +- Atlas 推理系列产品:x + + +- Atlas 训练系列产品:x + + +## 功能说明 + +- 算子功能:`aclnnQuantLightningIndexerV2Metadata`是`aclnnQuantLightningIndexerV2`算子的前置算子,用于生成负载均衡的任务划分方案。本算子不执行实际的LightningIndexer计算,而是根据输入参数在AI CPU计算出每个AI Core应处理的计算起止范围,从而最大化计算资源的利用率,避免各Core间负载不均衡的问题。 + + **该算子不建议单独使用,建议与aclnnQuantLightningIndexerV2算子配合使用,形成完整的工作流。** + 1. 接受aclnnQuantLightningIndexerV2算子接口输入数据shape信息,包含batchSize、qSeqlen、kSeqlen、mask。通过对输入分块并模拟计算耗时,均匀分配分块到可用核上,以降低aclnnQuantLightningIndexerV2算子的整体计算耗时,并提高硬件利用率。 + 2. 分配结果输出后,后续作为输入供aclnnQuantLightningIndexerV2算子使用。 + 3. 分配结果包含每个AIC核基本块的起始点和终止点,以及每个AIV核的FD任务信息。详细内容可以参考[调用示例](#调用示例)。 + +## 函数原型 + +每个算子分为[两段式接口](../../../docs/zh/context/two_phase_api.md),必须先调用"aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize"获取workspace大小,再调用"aclnnQuantLightningIndexerV2Metadata"执行计算 + +``` cpp +aclnnStatus aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize( + const aclTensor *cuSeqlensQOptional, + const aclTensor *cuSeqlensKOptional, + const aclTensor *sequsedQOptional, + const aclTensor *sequsedKOptional, + const aclTensor *cmpResidualKOptional, + int64_t numHeadsQ, + int64_t numHeadsK, + int64_t headDim, + int64_t topk, + int64_t quantMode, + int64_t batchSize, + int64_t maxSeqlenQ, + int64_t maxSeqlenK, + char *layoutQOptional, + char *layoutKOptional, + int64_t maskMode, + int64_t cmpRatio, + const aclTensor *metadata, + uint64_t *workspaceSize, + aclOpExecutor **executor) +``` + +``` cpp +aclnnStatus aclnnQuantLightningIndexerV2Metadata( + void *workspace, + uint64_t workspaceSize, + aclOpExecutor *executor, + aclrtStream stream) +``` + +## aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize + +- **参数说明** + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensor
cuSeqlensQOptional(aclTensor*)输入表示不同Batch中q的有效Sequence Length。
  • 支持空Tensor
  • shape固定为(B+1, )。
INT32ND1维√
cuSeqlensKOptional(aclTensor*)输入表示不同Batch中k的有效Sequence Length。
  • 支持空Tensor。
  • shape固定为(B+1, )。
INT32ND1维√
sequsedQOptional(aclTensor*)输入表示不同Batch中q实际参与运算的Sequence Length。
  • 支持空Tensor。
  • shape固定为(B, )。
INT32ND1维√
sequsedKOptional(aclTensor*)输入表示不同Batch中k实际参与运算的Sequence Length。
  • 支持空Tensor。
  • shape固定为(B, )。
INT32ND1维√
cmpResidualKOptional(aclTensor*)输入表示不同Batch中k压缩后Sequence Length的余数,配合cmpRatio实现mask和负载计算。
  • 支持空Tensor。
  • cmpRatio不为1,且mask为3场景下必传。
  • shape固定为(B, )。
INT32ND1维√
numHeadsQ(int64_t)输入表示q的head个数。当前支持[1, 64]。----
numHeadsK(int64_t)输入表示k的head个数。当前仅支持1。----
headDim(int64_t)输入注意力头的维度。当前仅支持128。----
topk(int64_t)输入表示从q中筛选出的关键稀疏token的个数。当前仅支持[1, 8192]。----
quantMode(int64_t)输入表示量化模式。
  • 当前支持1/2/3/4/5。
  • 1: qk: fp8(e4m3) per-token-head; scale: fp32。
  • 2: qk: int8 per-token-head; scale: fp16 w: fp16。
  • 3: qk: mxfp8(e4m3), scale: fp8(e8m0)。
  • 4: qk: hif8 per-tensor; scale: fp32。
  • 5: mxfp4(e2m1), scale: fp8(e8m0)。
----
batchSize(int64_t)输入表示Batch数量。
  • 支持非负数。
  • 建议值为0。
----
maxSeqlenQ(int64_t)输入表示q的最长Sequence Length。
  • 取值范围≥-1,-1表示任意可能长度。
  • 建议值为-1。
----
maxSeqlenK(int64_t)输入表示k的最长Sequence Length。
  • 取值范围≥-1,-1表示任意可能长度。
  • 建议值为-1。
----
layoutQOptional(char*)输入表示q的排列格式。
  • 支持 BSND、TND。
  • 建议值为BSND。
----
layoutKOptional(char*)输入表示k的排列格式。
  • 支持 BSND、TND、PA_BBND。
  • 建议值为BSND。
----
maskMode(int64_t)输入表示sparse模式。
  • 0: No mask。
  • 3: rightDownCausal模式的mask,对应以右顶点为划分的下三角场景。
  • 建议值为0。
----
cmpRatio(int64_t)输入表示k的压缩率。
  • 取值范围[1,128]。
  • 建议值1,表示无压缩。
----
metadata(aclTensor*)输出表示负载均衡结果输出。shape固定为(1024, )。INT32ND1维×
workspaceSize(uint64_t*)输出返回需要在Device侧申请的workspace大小。-----
executor(aclOpExecutor**)输出返回op执行器,包含了算子计算流程。-----
+ +
    + +
  • Atlas A3 训练系列产品/Atlas A3 推理系列产品 :numHeadsQ仅支持64,不支持quantMode = 1/3/4/5,topk仅支持[1, 2048],不支持layoutKOptional = BSND/TND,不支持cmpRatio在[1,128]任意取值,仅支持cmpRatio = 1/2/4/8/16/32/64/128。
  • + + +
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品 :numHeadsQ仅支持64,不支持quantMode = 1/3/4/5,topk仅支持[1, 2048],不支持layoutKOptional = BSND/TND,不支持cmpRatio在[1,128]任意取值,仅支持cmpRatio = 1/2/4/8/16/32/64/128。
  • + +
+ +- **返回值:** + + 返回aclnnStatus状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。 + + 第一段接口完成入参校验,出现以下场景时报错: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
返回值错误码描述
ACLNN_ERR_INNER_CREATE_EXECUTOR561101创建aclOpExecutor失败。
ACLNN_ERR_INNER_NULLPTR561103参数workspaceSize、executor是空指针,或参数cuSeqlensQOptional、cuSeqlensKOptional、sequsedQOptional、sequsedKOptional、cmpResidualKOptional进行Contiguous处理后为空指针。
ACLNN_ERR_PARAM_INVALID161002参数cuSeqlensQOptional、cuSeqlensKOptional、sequsedQOptional、sequsedKOptional、cmpResidualKOptional、numHeadsQ、numHeadsK、headDim、topk、quantMode、batchSize、maxSeqlenQ、maxSeqlenK、layoutQOptional、layoutKOptional、maskMode、cmpRatio的规格不在支持范围内。
+ +## aclnnQuantLightningIndexerV2Metadata + +- **参数说明:** + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出描述
workspace输入在Device侧申请的workspace内存地址
workspaceSize输入在Device侧申请的workspace大小,由第一段接口aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize获取
executor输入op执行器,包含了算子计算流程
stream输入指定执行任务的Stream
+ +- **返回值:** + + 返回aclnnStatus状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。 + +## 约束说明 + +- aclnnQuantLightningIndexerV2Metadata默认确定性实现。 +- B(Batch)表示输入样本批量大小,q、k为配套的aclnnQuantLightningIndexerV2算子的入参,S1表示layoutQOptional=BSND时,q shape中的S轴的大小,S2表示layoutKOptional=BSND时,k shape中的S轴的大小。 +- 参数cuSeqlensQOptional、cuSeqlensKOptional要求其值为当前Batch与前序Batch有效token数的累加值,第一个元素固定为0,后一个元素的值必须大于等于前一个元素的值。 +- 参数sequsedQOptional、sequsedKOptional要求其值表示每个Batch中的有效token数。 +- 非PA场景layoutQOptional、layoutKOptional须相同。 +- 参数cmpResidualKOptional需满足cmpResidualKOptional[i] < cmpRatio。 +- layoutQOptional=BSND场景 + - maxSeqlenQ必须传入S1的值。 +- layoutKOptional=BSND场景 + - maxSeqlenK必须传入S2的值。 +- layoutQOptional=TND场景 + - cuSeqlensQOptional必须传入。 +- layoutKOptional=TND场景 + - cuSeqlensKOptional必须传入。 +- layoutKOptional=PA_BBND场景 + - sequsedKOptional必须传入。 +- Batch取值规则 + - layoutQOptional为BSND时,优先通过sequsedQOptional的shape推导batch,sequsedQOptional未传入则通过batchSize获取batch数。 + - layoutQOptional为TND时,优先通过sequsedQOptional的shape推导batch,sequsedQOptional未传入则通过cuSeqlensQOptional的shape推导batch。 +- q Seqlen取值规则 + - layoutQOptional为BSND时,优先通过sequsedQOptional中的元素获取seqlen,sequsedQOptional未传入则通过maxSeqlenQ获取seqlen。 + - layoutQOptional为TND时,优先通过sequsedQOptional中的元素获取seqlen,sequsedQOptional未传入则通过cuSeqlensQOptional中的元素获取seqlen。 +- k Seqlen取值规则 + - layoutKOptional为BSND时,优先通过sequsedKOptional中的元素获取seqlen,sequsedKOptional未传入则通过maxSeqlenK获取seqlen。 + - layoutKOptional为TND时,优先通过sequsedKOptional中的元素获取seqlen,sequsedKOptional未传入则通过cuSeqlensKOptional中的元素获取seqlen。 + +## 调用示例 + +示例代码如下,仅供参考,具体编译和执行过程请参考[编译与运行样例](../../../docs/zh/context/compile_and_run_sample.md)。 + +``` cpp +#include +#include +#include +#include +#include +#include +#include +#include "acl/acl.h" +#include "aclnnop/aclnn_quant_lightning_indexer_v2_metadata.h" + +#define CHECK_LOG_RET(cond, ret_val, fmt, ...) \ + do { \ + if (!(cond)) { \ + printf(fmt "\n", ##__VA_ARGS__); \ + return (ret_val); \ + } \ + } while (0) + +// 参考 quant_lightning_indexer_v2_metadata.h +constexpr uint32_t AIC_CORE_MAX_NUM = 36; +constexpr uint32_t AIV_CORE_MAX_NUM = 72; +constexpr uint32_t QLI_V2_METADATA_TOTAL_SIZE = 1024; +constexpr uint32_t QLI_V2_METADATA_SIZE = 8; +constexpr uint32_t QLD_V2_METADATA_SIZE = 8; + +// QLI Metadata Index Definitions +constexpr uint32_t QLI_V2_CORE_ENABLE_INDEX = 0; +constexpr uint32_t QLI_V2_BN2_START_INDEX = 1; +constexpr uint32_t QLI_V2_M_START_INDEX = 2; +constexpr uint32_t QLI_V2_S2_START_INDEX = 3; +constexpr uint32_t QLI_V2_BN2_END_INDEX = 4; +constexpr uint32_t QLI_V2_M_END_INDEX = 5; +constexpr uint32_t QLI_V2_S2_END_INDEX = 6; +constexpr uint32_t QLI_V2_FIRST_QLD_V2_DATA_WORKSPACE_IDX_INDEX = 7; + +// QLD Metadata Index Definitions +constexpr uint32_t QLD_V2_CORE_ENABLE_INDEX = 0; +constexpr uint32_t QLD_V2_BN2_IDX_INDEX = 1; +constexpr uint32_t QLD_V2_M_IDX_INDEX = 2; +constexpr uint32_t QLD_V2_WORKSPACE_IDX_INDEX = 3; +constexpr uint32_t QLD_V2_WORKSPACE_NUM_INDEX = 4; +constexpr uint32_t QLD_V2_M_START_INDEX = 5; +constexpr uint32_t QLD_V2_M_NUM_INDEX = 6; + +struct QliV2Metadata { + uint32_t faData[AIC_CORE_MAX_NUM][QLI_V2_METADATA_SIZE]; + uint32_t fdData[AIV_CORE_MAX_NUM][QLD_V2_METADATA_SIZE]; +}; + +struct ScopeGuard +{ + explicit ScopeGuard(std::function onExitScope) : m_exitFunc(std::move(onExitScope)), + m_isDismissed(false) {} + // 禁止拷贝 + ScopeGuard(const ScopeGuard&) = delete; + ScopeGuard& operator=(const ScopeGuard&) = delete; + + ~ScopeGuard() + { + if (!m_isDismissed) { + m_exitFunc(); + } + } + + void Dismiss() + { + m_isDismissed = true; + } + + std::function m_exitFunc; + bool m_isDismissed; +}; + +struct Tensor { + void *hostAddr { nullptr }; + void *deviceAddr { nullptr }; + aclTensor *data { nullptr }; +}; + +struct ArgScenario { + bool hasCuSeq { false }; + bool hasSeqused { false }; +}; + +struct ArgContext { + // required input + int64_t numHeadsQ { 0 }; + int64_t numHeadsK { 0 }; + int64_t headDim { 0 }; + int64_t topk { 0 }; + int64_t quantMode { 2 }; + // optional input + Tensor cuSeqlensQOptional {}; + Tensor cuSeqlensKOptional {}; + Tensor sequsedQOptional {}; + Tensor sequsedKOptional {}; + Tensor cmpResidualKOptional {}; + int64_t batchSize { 0 }; + int64_t maxSeqlenQ { 0 }; + int64_t maxSeqlenK { 0 }; + char *layoutQOptional { nullptr }; + char *layoutKOptional { nullptr }; + int64_t maskMode { 0 }; + int64_t cmpRatio { 0 }; + // output + Tensor metadata {}; +}; + +int64_t GetShapeSize(const std::vector& shape) +{ + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} + +aclnnStatus Init(int32_t deviceId, aclrtStream* stream) +{ + // 固定写法,初始化 + auto ret = aclInit(nullptr); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclInit failed. ERROR: %d", ret); + ret = aclrtSetDevice(deviceId); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSetDevice failed. ERROR: %d", ret); + ret = aclrtCreateStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtCreateStream failed. ERROR: %d", ret); + return ACL_SUCCESS; +} + +void Finalize(int32_t deviceId, aclrtStream stream) +{ + aclrtDestroyStream(stream); + aclrtResetDevice(deviceId); + aclFinalize(); +} + +aclnnStatus CreateTensor(aclDataType dataType, const std::vector &shape, Tensor &tensor) +{ + auto size = GetShapeSize(shape) * aclDataTypeSize(dataType); + // 调用aclrtMallocHost申请host侧内存 + auto ret = aclrtMallocHost(&(tensor.hostAddr), size); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMallocHost failed. ERROR: %d", ret); + memset(tensor.hostAddr, 0, size); + // 调用aclrtMalloc申请device侧内存 + ret = aclrtMalloc(&(tensor.deviceAddr), size, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMalloc failed. ERROR: %d", ret); + // 调用aclCreateTensor接口创建aclTensor + tensor.data = aclCreateTensor(shape.data(), shape.size(), dataType, nullptr, 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), tensor.deviceAddr); + + // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 + ret = aclrtMemcpy(tensor.deviceAddr, size, tensor.hostAddr, size, ACL_MEMCPY_HOST_TO_DEVICE); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d", ret); + return ACL_SUCCESS; +} + +void DestroyTensor(Tensor &tensor) +{ + if (tensor.data != nullptr) { + aclDestroyTensor(tensor.data); + tensor.data = nullptr; + } + if (tensor.deviceAddr != nullptr) { + aclrtFree(tensor.deviceAddr); + tensor.deviceAddr = nullptr; + } + if (tensor.hostAddr != nullptr) { + aclrtFreeHost(tensor.hostAddr); + tensor.hostAddr = nullptr; + } +} + +void DestroyArgs(ArgContext &context) +{ + DestroyTensor(context.metadata); + DestroyTensor(context.cuSeqlensQOptional); + DestroyTensor(context.cuSeqlensKOptional); + DestroyTensor(context.sequsedQOptional); + DestroyTensor(context.sequsedKOptional); + DestroyTensor(context.cmpResidualKOptional); + + if (context.layoutQOptional != nullptr) { + free(context.layoutQOptional); + context.layoutQOptional = nullptr; + } + if (context.layoutKOptional != nullptr) { + free(context.layoutKOptional); + context.layoutKOptional = nullptr; + } +} + +aclnnStatus CreateArgs(const ArgScenario &scenario, ArgContext &context) +{ + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + aclnnStatus ret; + int64_t batchSize = 4; + + context.numHeadsQ = 1; + context.numHeadsK = 1; + context.headDim = 128; + context.topk = 0; + context.quantMode = 2; // 2: per-token-head / 3: group-scaling + ret = CreateTensor(aclDataType::ACL_INT32, { QLI_V2_METADATA_TOTAL_SIZE }, context.metadata); // 1024: Fix size + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create meta failed. Error: %d", ret); + + context.maskMode = 0; // 0: no mask, 3: causal + context.cmpRatio = 1; // [1, 128], 1: no compress + context.layoutQOptional = (char *)malloc(sizeof(char) * 16); + context.layoutKOptional = (char *)malloc(sizeof(char) * 16); + strcpy(context.layoutQOptional, "BSND"); // BSND,TND + strcpy(context.layoutKOptional, "BSND"); // BSND,TND,PA_BBND + + if (!scenario.hasCuSeq && !scenario.hasSeqused) { + context.batchSize = batchSize; + context.maxSeqlenK = 1024; + context.maxSeqlenQ = 1024; + argsGuard.Dismiss(); + return ACL_SUCCESS; + } + + if (scenario.hasCuSeq) { + // (B+1,), first element is always 0 + ret = CreateTensor(aclDataType::ACL_INT32, { batchSize + 1 }, context.cuSeqlensQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { batchSize + 1 }, context.cuSeqlensKOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensKOptional failed. Error: %d", ret); + } + + if (scenario.hasSeqused) { + // (B,) + ret = CreateTensor(aclDataType::ACL_INT32, { batchSize }, context.sequsedQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { batchSize }, context.sequsedKOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedKOptional failed. Error: %d", ret); + } + + argsGuard.Dismiss(); + return ACL_SUCCESS; +} + +int main() { + // 1.(固定写法)device/stream初始化,参考对外接口列表 + // 根据自己的实际device填写deviceId + int32_t deviceId = 0; + aclrtStream stream; + auto ret = Init(deviceId, &stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Init acl failed. ERROR: %d", ret); + ScopeGuard sysGuard([&] { Finalize(deviceId, stream); }); + + // 2. 构造输入与输出,需要根据API的接口定义构造 + ArgScenario scenario {}; + scenario.hasCuSeq = true; + scenario.hasSeqused = true; + ArgContext context {}; + ret = CreateArgs(scenario, context); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create input arguments failed. ERROR: %d", ret); + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + + // 3. 调用CANN算子库API,需要修改为具体的API + // 调用aclnnQuantLightningIndexerV2Metadata第一段接口 + uint64_t workspaceSize = 0; + aclOpExecutor *executor = nullptr; + void *workspaceAddr = nullptr; + ret = aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize( + context.cuSeqlensQOptional.data, context.cuSeqlensKOptional.data, context.sequsedQOptional.data, + context.sequsedKOptional.data, context.cmpResidualKOptional.data, + context.numHeadsQ, context.numHeadsK, context.headDim, context.topk, context.quantMode, + context.batchSize, context.maxSeqlenQ, context.maxSeqlenK, context.layoutQOptional, + context.layoutKOptional, context.maskMode, context.cmpRatio, + context.metadata.data, &workspaceSize, &executor); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, + "aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize failed. ERROR: %d\n", ret); + + if (workspaceSize > static_cast(0)) { + ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "allocate workspace failed. ERROR: %d\n", ret); + } + ScopeGuard workspaceGuard([&] { + if (workspaceAddr != nullptr) { + aclrtFree(workspaceAddr); + workspaceAddr = nullptr; + } + }); + + // 调用aclnnQuantLightningIndexerV2Metadata第二段接口 + ret = aclnnQuantLightningIndexerV2Metadata(workspaceAddr, workspaceSize, executor, stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclnnQuantLightningIndexerV2Metadata failed. ERROR: %d\n", ret); + + // 4.(固定写法)同步等待任务执行结束 + ret = aclrtSynchronizeStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSynchronizeStream failed. ERROR: %d\n", ret); + + // 5. 打印输出 + QliV2Metadata result {}; + ret = aclrtMemcpy(&result, sizeof(result), context.metadata.deviceAddr, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d\n", ret); + + for (uint32_t i = 0; i < AIC_CORE_MAX_NUM; ++i) { + printf("AIC Core%u\n", i); + printf(" Core Enable : %u\n", result.faData[i][QLI_V2_CORE_ENABLE_INDEX]); + printf(" Start BN2 : %u\n", result.faData[i][QLI_V2_BN2_START_INDEX]); + printf(" Start M : %u\n", result.faData[i][QLI_V2_M_START_INDEX]); + printf(" Start S2 : %u\n", result.faData[i][QLI_V2_S2_START_INDEX]); + printf(" End BN2 : %u\n", result.faData[i][QLI_V2_BN2_END_INDEX]); + printf(" End M : %u\n", result.faData[i][QLI_V2_M_END_INDEX]); + printf(" End S2 : %u\n", result.faData[i][QLI_V2_S2_END_INDEX]); + printf(" First Workspace Index : %u\n", result.faData[i][QLI_V2_FIRST_QLD_V2_DATA_WORKSPACE_IDX_INDEX]); + } + for (uint32_t i = 0; i < AIV_CORE_MAX_NUM; ++i) { + printf("AIV Core%u\n", i); + printf(" Core Enable : %u\n", result.fdData[i][QLD_V2_CORE_ENABLE_INDEX]); + printf(" FD Task BN2 Idx : %u\n", result.fdData[i][QLD_V2_BN2_IDX_INDEX]); + printf(" FD Task M Idx : %u\n", result.fdData[i][QLD_V2_M_IDX_INDEX]); + printf(" FD Task Workspace Idx : %u\n", result.fdData[i][QLD_V2_WORKSPACE_IDX_INDEX]); + printf(" FD Task Workspace Num : %u\n", result.fdData[i][QLD_V2_WORKSPACE_NUM_INDEX]); + printf(" FD Subtask M Start : %u\n", result.fdData[i][QLD_V2_M_START_INDEX]); + printf(" FD Subtask M Num : %u\n", result.fdData[i][QLD_V2_M_NUM_INDEX]); + } + + return 0; +} +``` diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/examples/test_aclnn_quant_lightning_indexer_v2_metadata.cpp b/csrc/attention/quant_lightning_indexer_v2_metadata/examples/test_aclnn_quant_lightning_indexer_v2_metadata.cpp new file mode 100644 index 000000000000..2d9ddb619fa8 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/examples/test_aclnn_quant_lightning_indexer_v2_metadata.cpp @@ -0,0 +1,335 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/** + * @file test_aclnn_quant_lightning_indexer_v2_metadata.cpp + */ +#include +#include +#include +#include +#include +#include +#include +#include "acl/acl.h" +#include "aclnnop/aclnn_quant_lightning_indexer_v2_metadata.h" + +#define CHECK_LOG_RET(cond, ret_val, fmt, ...) \ + do { \ + if (!(cond)) { \ + printf(fmt "\n", ##__VA_ARGS__); \ + return (ret_val); \ + } \ + } while (0) + +// 参考 quant_lightning_indexer_v2_metadata.h +constexpr uint32_t AIC_CORE_MAX_NUM = 36; +constexpr uint32_t AIV_CORE_MAX_NUM = 72; +constexpr uint32_t QLI_V2_METADATA_TOTAL_SIZE = 1024; +constexpr uint32_t QLI_V2_METADATA_SIZE = 8; +constexpr uint32_t QLD_V2_METADATA_SIZE = 8; + +// QLI Metadata Index Definitions +constexpr uint32_t QLI_V2_CORE_ENABLE_INDEX = 0; +constexpr uint32_t QLI_V2_BN2_START_INDEX = 1; +constexpr uint32_t QLI_V2_M_START_INDEX = 2; +constexpr uint32_t QLI_V2_S2_START_INDEX = 3; +constexpr uint32_t QLI_V2_BN2_END_INDEX = 4; +constexpr uint32_t QLI_V2_M_END_INDEX = 5; +constexpr uint32_t QLI_V2_S2_END_INDEX = 6; +constexpr uint32_t QLI_V2_FIRST_QLD_V2_DATA_WORKSPACE_IDX_INDEX = 7; + +// QLD Metadata Index Definitions +constexpr uint32_t QLD_V2_CORE_ENABLE_INDEX = 0; +constexpr uint32_t QLD_V2_BN2_IDX_INDEX = 1; +constexpr uint32_t QLD_V2_M_IDX_INDEX = 2; +constexpr uint32_t QLD_V2_WORKSPACE_IDX_INDEX = 3; +constexpr uint32_t QLD_V2_WORKSPACE_NUM_INDEX = 4; +constexpr uint32_t QLD_V2_M_START_INDEX = 5; +constexpr uint32_t QLD_V2_M_NUM_INDEX = 6; + +struct QliV2Metadata { + uint32_t faData[AIC_CORE_MAX_NUM][QLI_V2_METADATA_SIZE]; + uint32_t fdData[AIV_CORE_MAX_NUM][QLD_V2_METADATA_SIZE]; +}; + +struct ScopeGuard { + explicit ScopeGuard(std::function onExitScope) + : m_exitFunc(std::move(onExitScope)), + m_isDismissed(false) + {} + // 禁止拷贝 + ScopeGuard(const ScopeGuard &) = delete; + ScopeGuard &operator=(const ScopeGuard &) = delete; + + ~ScopeGuard() + { + if (!m_isDismissed) { + m_exitFunc(); + } + } + + void Dismiss() + { + m_isDismissed = true; + } + + std::function m_exitFunc; + bool m_isDismissed; +}; + +struct Tensor { + void *hostAddr{nullptr}; + void *deviceAddr{nullptr}; + aclTensor *data{nullptr}; +}; + +struct ArgScenario { + bool hasCuSeq{false}; + bool hasSeqused{false}; +}; + +struct ArgContext { + // required input + int64_t numHeadsQ{0}; + int64_t numHeadsK{0}; + int64_t headDim{0}; + int64_t topk{0}; + int64_t quantMode{2}; + // optional input + Tensor cuSeqlensQOptional{}; + Tensor cuSeqlensKOptional{}; + Tensor sequsedQOptional{}; + Tensor sequsedKOptional{}; + Tensor cmpResidualKOptional{}; + int64_t batchSize{0}; + int64_t maxSeqlenQ{0}; + int64_t maxSeqlenK{0}; + char *layoutQOptional{nullptr}; + char *layoutKOptional{nullptr}; + int64_t maskMode{0}; + int64_t cmpRatio{0}; + // output + Tensor metadata{}; +}; + +int64_t GetShapeSize(const std::vector &shape) +{ + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} + +aclnnStatus Init(int32_t deviceId, aclrtStream *stream) +{ + // 固定写法,初始化 + auto ret = aclInit(nullptr); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclInit failed. ERROR: %d", ret); + ret = aclrtSetDevice(deviceId); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSetDevice failed. ERROR: %d", ret); + ret = aclrtCreateStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtCreateStream failed. ERROR: %d", ret); + return ACL_SUCCESS; +} + +void Finalize(int32_t deviceId, aclrtStream stream) +{ + aclrtDestroyStream(stream); + aclrtResetDevice(deviceId); + aclFinalize(); +} + +aclnnStatus CreateTensor(aclDataType dataType, const std::vector &shape, Tensor &tensor) +{ + auto size = GetShapeSize(shape) * aclDataTypeSize(dataType); + // 调用aclrtMallocHost申请host侧内存 + auto ret = aclrtMallocHost(&(tensor.hostAddr), size); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMallocHost failed. ERROR: %d", ret); + memset(tensor.hostAddr, 0, size); + // 调用aclrtMalloc申请device侧内存 + ret = aclrtMalloc(&(tensor.deviceAddr), size, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMalloc failed. ERROR: %d", ret); + // 调用aclCreateTensor接口创建aclTensor + tensor.data = aclCreateTensor(shape.data(), shape.size(), dataType, nullptr, 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), tensor.deviceAddr); + + // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 + ret = aclrtMemcpy(tensor.deviceAddr, size, tensor.hostAddr, size, ACL_MEMCPY_HOST_TO_DEVICE); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d", ret); + return ACL_SUCCESS; +} + +void DestroyTensor(Tensor &tensor) +{ + if (tensor.data != nullptr) { + aclDestroyTensor(tensor.data); + tensor.data = nullptr; + } + if (tensor.deviceAddr != nullptr) { + aclrtFree(tensor.deviceAddr); + tensor.deviceAddr = nullptr; + } + if (tensor.hostAddr != nullptr) { + aclrtFreeHost(tensor.hostAddr); + tensor.hostAddr = nullptr; + } +} + +void DestroyArgs(ArgContext &context) +{ + DestroyTensor(context.metadata); + DestroyTensor(context.cuSeqlensQOptional); + DestroyTensor(context.cuSeqlensKOptional); + DestroyTensor(context.sequsedQOptional); + DestroyTensor(context.sequsedKOptional); + DestroyTensor(context.cmpResidualKOptional); + + if (context.layoutQOptional != nullptr) { + free(context.layoutQOptional); + context.layoutQOptional = nullptr; + } + if (context.layoutKOptional != nullptr) { + free(context.layoutKOptional); + context.layoutKOptional = nullptr; + } +} + +aclnnStatus CreateArgs(const ArgScenario &scenario, ArgContext &context) +{ + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + aclnnStatus ret; + int64_t batchSize = 4; + + context.numHeadsQ = 1; + context.numHeadsK = 1; + context.headDim = 128; + context.topk = 0; + context.quantMode = 2; // 2: per-token-head / 3: group-scaling + ret = CreateTensor(aclDataType::ACL_INT32, {QLI_V2_METADATA_TOTAL_SIZE}, context.metadata); // 1024: Fix size + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create meta failed. Error: %d", ret); + + context.maskMode = 0; // 0: no mask, 3: causal + context.cmpRatio = 1; // [1, 128], 1: no compress + context.layoutQOptional = (char *)malloc(sizeof(char) * 16); + context.layoutKOptional = (char *)malloc(sizeof(char) * 16); + strcpy(context.layoutQOptional, "BSND"); // BSND,TND + strcpy(context.layoutKOptional, "BSND"); // BSND,TND,PA_BBND + + if (!scenario.hasCuSeq && !scenario.hasSeqused) { + context.batchSize = batchSize; + context.maxSeqlenK = 1024; + context.maxSeqlenQ = 1024; + argsGuard.Dismiss(); + return ACL_SUCCESS; + } + + if (scenario.hasCuSeq) { + // (B+1,), first element is always 0 + ret = CreateTensor(aclDataType::ACL_INT32, {batchSize + 1}, context.cuSeqlensQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, {batchSize + 1}, context.cuSeqlensKOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensKOptional failed. Error: %d", ret); + } + + if (scenario.hasSeqused) { + // (B,) + ret = CreateTensor(aclDataType::ACL_INT32, {batchSize}, context.sequsedQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, {batchSize}, context.sequsedKOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedKOptional failed. Error: %d", ret); + } + + argsGuard.Dismiss(); + return ACL_SUCCESS; +} + +int main() +{ + // 1. (固定写法)device/stream初始化,参考对外接口列表 + // 根据自己的实际device填写deviceId + int32_t deviceId = 0; + aclrtStream stream; + auto ret = Init(deviceId, &stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Init acl failed. ERROR: %d", ret); + ScopeGuard sysGuard([&] { Finalize(deviceId, stream); }); + + // 2. 构造输入与输出,需要根据API的接口定义构造 + ArgScenario scenario{}; + scenario.hasCuSeq = true; + scenario.hasSeqused = true; + ArgContext context{}; + ret = CreateArgs(scenario, context); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create input arguments failed. ERROR: %d", ret); + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + + // 3. 调用CANN算子库API,需要修改为具体的API + // 调用aclnnLightningIndexerV2Metadata第一段接口 + uint64_t workspaceSize = 0; + aclOpExecutor *executor = nullptr; + void *workspaceAddr = nullptr; + ret = aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize( + context.cuSeqlensQOptional.data, context.cuSeqlensKOptional.data, context.sequsedQOptional.data, + context.sequsedKOptional.data, context.cmpResidualKOptional.data, context.numHeadsQ, context.numHeadsK, + context.headDim, context.topk, context.quantMode, context.batchSize, context.maxSeqlenQ, context.maxSeqlenK, + context.layoutQOptional, context.layoutKOptional, context.maskMode, context.cmpRatio, context.metadata.data, + &workspaceSize, &executor); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize failed. ERROR: %d\n", + ret); + + if (workspaceSize > static_cast(0)) { + ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "allocate workspace failed. ERROR: %d\n", ret); + } + ScopeGuard workspaceGuard([&] { + if (workspaceAddr != nullptr) { + aclrtFree(workspaceAddr); + workspaceAddr = nullptr; + } + }); + + // 调用aclnnLightningIndexerV2Metadata第二段接口 + ret = aclnnQuantLightningIndexerV2Metadata(workspaceAddr, workspaceSize, executor, stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclnnQuantLightningIndexerV2Metadata failed. ERROR: %d\n", ret); + + // 4. (固定写法)同步等待任务执行结束 + ret = aclrtSynchronizeStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSynchronizeStream failed. ERROR: %d\n", ret); + + // 5. 打印输出 + QliV2Metadata result{}; + ret = aclrtMemcpy(&result, sizeof(result), context.metadata.deviceAddr, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d\n", ret); + + for (uint32_t i = 0; i < AIC_CORE_MAX_NUM; ++i) { + printf("AIC Core%u\n", i); + printf(" Core Enable : %u\n", result.faData[i][QLI_V2_CORE_ENABLE_INDEX]); + printf(" Start BN2 : %u\n", result.faData[i][QLI_V2_BN2_START_INDEX]); + printf(" Start M : %u\n", result.faData[i][QLI_V2_M_START_INDEX]); + printf(" Start S2 : %u\n", result.faData[i][QLI_V2_S2_START_INDEX]); + printf(" End BN2 : %u\n", result.faData[i][QLI_V2_BN2_END_INDEX]); + printf(" End M : %u\n", result.faData[i][QLI_V2_M_END_INDEX]); + printf(" End S2 : %u\n", result.faData[i][QLI_V2_S2_END_INDEX]); + printf(" First Workspace Index : %u\n", result.faData[i][QLI_V2_FIRST_QLD_V2_DATA_WORKSPACE_IDX_INDEX]); + } + for (uint32_t i = 0; i < AIV_CORE_MAX_NUM; ++i) { + printf("AIV Core%u\n", i); + printf(" Core Enable : %u\n", result.fdData[i][QLD_V2_CORE_ENABLE_INDEX]); + printf(" FD Task BN2 Idx : %u\n", result.fdData[i][QLD_V2_BN2_IDX_INDEX]); + printf(" FD Task M Idx : %u\n", result.fdData[i][QLD_V2_M_IDX_INDEX]); + printf(" FD Task Workspace Idx : %u\n", result.fdData[i][QLD_V2_WORKSPACE_IDX_INDEX]); + printf(" FD Task Workspace Num : %u\n", result.fdData[i][QLD_V2_WORKSPACE_NUM_INDEX]); + printf(" FD Subtask M Start : %u\n", result.fdData[i][QLD_V2_M_START_INDEX]); + printf(" FD Subtask M Num : %u\n", result.fdData[i][QLD_V2_M_NUM_INDEX]); + } + + return 0; +} diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/examples/test_torch_quant_lightning_indexer_v2_metadata.py b/csrc/attention/quant_lightning_indexer_v2_metadata/examples/test_torch_quant_lightning_indexer_v2_metadata.py new file mode 100644 index 000000000000..fdc5fb0821fa --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/examples/test_torch_quant_lightning_indexer_v2_metadata.py @@ -0,0 +1,31 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2025 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import torch + +metadata = torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata( + cu_seqlens_q=None, + cu_seqlens_k=None, + seqused_q=None, + seqused_k=None, + cmp_residual_k=torch.tensor([3, 3, 3, 3, 3, 3, 3, 3], dtype=torch.int32).npu(), + batch_size=8, + max_seqlen_q=10, + max_seqlen_k=10, + num_heads_q=64, + num_heads_k=1, + head_dim=128, + topk=2048, + quant_mode=1, + mask_mode=3, + layout_q="BSND", + layout_k="BSND", + cmp_ratio=128, +) diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/CMakeLists.txt new file mode 100644 index 000000000000..a842512c95eb --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/CMakeLists.txt @@ -0,0 +1,12 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +add_op_to_compiled_list() +add_modules_sources(OPTYPE quant_lightning_indexer_v2_metadata ACLNNTYPE aclnn) diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/aclnn_quant_lightning_indexer_v2_metadata.cpp b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/aclnn_quant_lightning_indexer_v2_metadata.cpp new file mode 100644 index 000000000000..399d3c1ea8cc --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/aclnn_quant_lightning_indexer_v2_metadata.cpp @@ -0,0 +1,134 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file aclnn_quant_lightning_indexer_v2_metadata.cpp + * \brief + */ + +#include "aclnn_quant_lightning_indexer_v2_metadata.h" +#include "../quant_lightning_indexer_v2_metadata_check.h" +#include "aclnn/aclnn_base.h" +#include "aclnn_kernels/common/op_error_check.h" +#include "aclnn_kernels/contiguous.h" +#include "aclnn_kernels/reshape.h" +#include "quant_lightning_indexer_v2_metadata.h" +#include "opdev/common_types.h" +#include "opdev/data_type_utils.h" +#include "opdev/format_utils.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/tensor_view_utils.h" +#include "opdev/platform.h" + +#ifdef __cplusplus +extern "C" { +#endif + +aclnnStatus aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize( + const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensKOptional, const aclTensor *sequsedQOptional, + const aclTensor *sequsedKOptional, const aclTensor *cmpResidualKOptional, int64_t numHeadsQ, int64_t numHeadsK, + int64_t headDim, int64_t topk, int64_t quantMode, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenK, + char *layoutQOptional, char *layoutKOptional, int64_t maskMode, int64_t cmpRatio, const aclTensor *metadata, + uint64_t *workspaceSize, aclOpExecutor **executor) +{ + if (workspaceSize == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "workspaceSize is nullptr"); + return ACLNN_ERR_INNER_NULLPTR; + } + if (executor == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "executor is nullptr"); + return ACLNN_ERR_INNER_NULLPTR; + } + L2_DFX_PHASE_1(aclnnQuantLightningIndexerV2Metadata, + DFX_IN(cuSeqlensQOptional, cuSeqlensKOptional, sequsedQOptional, sequsedKOptional, + cmpResidualKOptional, numHeadsQ, numHeadsK, headDim, topk, quantMode, batchSize, maxSeqlenQ, + maxSeqlenK, layoutQOptional, layoutKOptional, maskMode, cmpRatio), + DFX_OUT(metadata)); + + auto uniqueExecutor = CREATE_EXECUTOR(); + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); + + const op::PlatformInfo &npuInfo = op::GetCurrentPlatformInfo(); + uint32_t aicCoreNum = npuInfo.GetCubeCoreNum(); + uint32_t aivCoreNum = npuInfo.GetVectorCoreNum(); + const std::string socVersion = npuInfo.GetSocLongVersion(); + + auto ret = ParamsCheckQliV2(cuSeqlensQOptional, cuSeqlensKOptional, sequsedQOptional, sequsedKOptional, + cmpResidualKOptional, numHeadsQ, numHeadsK, headDim, topk, quantMode, batchSize, + maxSeqlenQ, maxSeqlenK, layoutQOptional, layoutKOptional, maskMode, cmpRatio, metadata, + aicCoreNum, aivCoreNum, socVersion); + CHECK_RET(ret == ACLNN_SUCCESS, ret); + + const aclTensor *cuSeqlensQOptionalContiguous = nullptr; + if (cuSeqlensQOptional != nullptr) { + cuSeqlensQOptionalContiguous = l0op::Contiguous(cuSeqlensQOptional, uniqueExecutor.get()); + if (cuSeqlensQOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "cu_seqlens_q contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *cuSeqlensKOptionalContiguous = nullptr; + if (cuSeqlensKOptional != nullptr) { + cuSeqlensKOptionalContiguous = l0op::Contiguous(cuSeqlensKOptional, uniqueExecutor.get()); + if (cuSeqlensKOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "cu_seqlens_k contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *sequsedQOptionalContiguous = nullptr; + if (sequsedQOptional != nullptr) { + sequsedQOptionalContiguous = l0op::Contiguous(sequsedQOptional, uniqueExecutor.get()); + if (sequsedQOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "seqused_q contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *sequsedKOptionalContiguous = nullptr; + if (sequsedKOptional != nullptr) { + sequsedKOptionalContiguous = l0op::Contiguous(sequsedKOptional, uniqueExecutor.get()); + if (sequsedKOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "seqused_k contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *cmpResidualLKOptionalContiguous = nullptr; + if (cmpResidualKOptional != nullptr) { + cmpResidualLKOptionalContiguous = l0op::Contiguous(cmpResidualKOptional, uniqueExecutor.get()); + if (cmpResidualLKOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "cmp_residual_k contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + + auto output = l0op::QuantLightningIndexerV2Metadata( + cuSeqlensQOptionalContiguous, cuSeqlensKOptionalContiguous, sequsedQOptionalContiguous, + sequsedKOptionalContiguous, cmpResidualLKOptionalContiguous, numHeadsQ, numHeadsK, headDim, topk, quantMode, + batchSize, maxSeqlenQ, maxSeqlenK, layoutQOptional, layoutKOptional, maskMode, cmpRatio, aicCoreNum, aivCoreNum, + socVersion.c_str(), metadata, uniqueExecutor.get()); + CHECK_RET(output != nullptr, ACLNN_ERR_INNER_NULLPTR); + + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); + uniqueExecutor.ReleaseTo(executor); + return ACLNN_SUCCESS; +} + +aclnnStatus aclnnQuantLightningIndexerV2Metadata(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, + aclrtStream stream) +{ + L2_DFX_PHASE_2(aclnnQuantLightningIndexerV2Metadata); + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); +} + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/aclnn_quant_lightning_indexer_v2_metadata.h b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/aclnn_quant_lightning_indexer_v2_metadata.h new file mode 100644 index 000000000000..766c77c5adba --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/aclnn_quant_lightning_indexer_v2_metadata.h @@ -0,0 +1,37 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef ACLNN_QUANT_LIGHTNING_INDEXER_V2_METADATA_H +#define ACLNN_QUANT_LIGHTNING_INDEXER_V2_METADATA_H + +#include +#include "aclnn/aclnn_base.h" + +#ifdef __cplusplus +extern "C" { +#endif + +__attribute__((visibility("default"))) +aclnnStatus aclnnQuantLightningIndexerV2MetadataGetWorkspaceSize( + const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensKOptional, const aclTensor *sequsedQOptional, + const aclTensor *sequsedKOptional, const aclTensor *cmpResidualKOptional, int64_t numHeadsQ, int64_t numHeadsK, + int64_t headDim, int64_t topk, int64_t quantMode, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenK, + char *layoutQOptional, char *layoutKOptional, int64_t maskMode, int64_t cmpRatio, const aclTensor *metadata, + uint64_t *workspaceSize, aclOpExecutor **executor); + +__attribute__((visibility("default"))) +aclnnStatus aclnnQuantLightningIndexerV2Metadata(void* workspace, uint64_t workspaceSize, aclOpExecutor *executor, + aclrtStream stream); + +#ifdef __cplusplus +} +#endif + +#endif // ACLNN_QUANT_LIGHTNING_INDEXER_V2_METADATA_AICPU_H diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/quant_lightning_indexer_v2_metadata.cpp b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/quant_lightning_indexer_v2_metadata.cpp new file mode 100644 index 000000000000..dc4d283578df --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/quant_lightning_indexer_v2_metadata.cpp @@ -0,0 +1,57 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file quant_lightning_indexer_v2_metadata.cpp + * \brief + */ + +#include "quant_lightning_indexer_v2_metadata.h" +#include "opdev/aicpu/aicpu_task.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_def.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/shape_utils.h" + +using namespace op; +namespace l0op { +OP_TYPE_REGISTER(QuantLightningIndexerV2Metadata); + +const aclTensor *QuantLightningIndexerV2Metadata( + const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensKOptional, const aclTensor *sequsedQOptional, + const aclTensor *sequsedKOptional, const aclTensor *cmpResidualKOptional, int64_t numHeadsQ, int64_t numHeadsK, + int64_t headDim, int64_t topk, int64_t quantMode, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenK, + char *layoutQOptional, char *layoutKOptional, int64_t maskMode, int64_t cmpRatio, int64_t aicCoreNum, + int64_t aivCoreNum, const char *socVersion, const aclTensor *metadata, aclOpExecutor *executor) +{ + L0_DFX(QuantLightningIndexerV2Metadata, cuSeqlensQOptional, cuSeqlensKOptional, sequsedQOptional, sequsedKOptional, + cmpResidualKOptional, numHeadsQ, numHeadsK, headDim, topk, quantMode, batchSize, maxSeqlenQ, maxSeqlenK, + layoutQOptional, layoutKOptional, maskMode, cmpRatio, aicCoreNum, aivCoreNum, socVersion, metadata); + + static internal::AicpuTaskSpace space("QuantLightningIndexerV2Metadata"); + + auto ret = ADD_TO_LAUNCHER_LIST_AICPU( + QuantLightningIndexerV2Metadata, + OP_ATTR_NAMES({ "num_heads_q", "num_heads_k", "head_dim", "topk", "quant_mode", "batch_size", "max_seqlen_q", + "max_seqlen_k", "layout_q", "layout_k", "mask_mode", "cmp_ratio", "aic_core_num", + "aiv_core_num", "soc_version" }), + OP_INPUT(cuSeqlensQOptional, cuSeqlensKOptional, sequsedQOptional, sequsedKOptional, cmpResidualKOptional), + OP_OUTPUT(metadata), + OP_ATTR(numHeadsQ, numHeadsK, headDim, topk, quantMode, batchSize, maxSeqlenQ, maxSeqlenK, layoutQOptional, + layoutKOptional, maskMode, cmpRatio, aicCoreNum, aivCoreNum, socVersion)); + + OP_CHECK(ret == ACL_SUCCESS, + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "QuantLightningIndexerV2Metadata ADD_TO_LAUNCHER_LIST_AICPU failed."), + return nullptr); + return metadata; +} +} // namespace l0op diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/quant_lightning_indexer_v2_metadata.h b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/quant_lightning_indexer_v2_metadata.h new file mode 100644 index 000000000000..35aae86d203a --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/op_api/quant_lightning_indexer_v2_metadata.h @@ -0,0 +1,25 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef L0_QUANT_LIGHTNING_INDEXER_V2_METADATA_H +#define L0_QUANT_LIGHTNING_INDEXER_V2_METADATA_H + +#include "opdev/op_executor.h" + +namespace l0op { +const aclTensor *QuantLightningIndexerV2Metadata( + const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensKOptional, const aclTensor *sequsedQOptional, + const aclTensor *sequsedKOptional, const aclTensor *cmpResidualKOptional, int64_t numHeadsQ, int64_t numHeadsK, + int64_t headDim, int64_t topk, int64_t quantMode, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenK, + char *layoutQOptional, char *layoutKOptional, int64_t maskMode, int64_t cmpRatio, int64_t aicCoreNum, + int64_t aivCoreNum, const char *socVersion, const aclTensor *metadata, aclOpExecutor *executor); +} // namespace l0op + +#endif diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/quant_lightning_indexer_v2_metadata_check.h b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/quant_lightning_indexer_v2_metadata_check.h index f54a25987271..a8e7f2558577 100644 --- a/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/quant_lightning_indexer_v2_metadata_check.h +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_host/quant_lightning_indexer_v2_metadata_check.h @@ -39,6 +39,7 @@ inline constexpr int64_t QLI_V2_CMP_RATIO_LOWER_BOUND = 1; inline constexpr int64_t QLI_V2_CMP_RATIO_UPPER_BOUND = 128; inline constexpr int64_t QLI_V2_NUM_HEADS_Q_LOWER_BOUND = 1; inline constexpr int64_t QLI_V2_NUM_HEADS_Q_UPPER_BOUND = 64; +inline constexpr int64_t QLI_V2_NUM_HEADS_Q_G32 = 32; // g=32 支持 (A1 对齐 v1 推导) inline constexpr int64_t QLI_V2_TOPK_LOWER_BOUND = 1; inline constexpr int64_t QLI_V2_A5_TOPK_UPPER_BOUND = 8192; inline constexpr int64_t QLI_V2_A3_TOPK_UPPER_BOUND = 2048; @@ -66,7 +67,10 @@ aclDataType GetDataTypeQliV2(const aclTensor *tensor) return dataType; } -inline bool IsTensorSourceQLiV2(const std::string &source) { return source != "batch_size"; } +inline bool IsTensorSourceQLiV2(const std::string &source) +{ + return source != "batch_size"; +} inline int64_t GetRawShapeSizeQLiV2(const std::string &source, int64_t batchValue) { @@ -185,9 +189,11 @@ aclnnStatus CheckSingleParamQliV2(int64_t numHeadsQ, int64_t numHeadsK, int64_t } // 校验 A2/A3 参数 if (socVersion.find("Ascend950") == std::string::npos) { - // num_heads_q 校验 - CHECK_COND(numHeadsQ == QLI_V2_NUM_HEADS_Q_UPPER_BOUND, ACLNN_ERR_PARAM_INVALID, - "num_heads_q should be %lld, but got %lld", QLI_V2_NUM_HEADS_Q_UPPER_BOUND, numHeadsQ); + // num_heads_q 校验 (g=64/32, 32 参照 v1 对齐 mBaseSize=4*g 推导) + CHECK_COND((numHeadsQ == QLI_V2_NUM_HEADS_Q_UPPER_BOUND) || (numHeadsQ == QLI_V2_NUM_HEADS_Q_G32), + ACLNN_ERR_PARAM_INVALID, + "num_heads_q should be %lld or %lld, but got %lld", QLI_V2_NUM_HEADS_Q_UPPER_BOUND, + QLI_V2_NUM_HEADS_Q_G32, numHeadsQ); // topk 校验 CHECK_COND(topk >= QLI_V2_TOPK_LOWER_BOUND && topk <= QLI_V2_A3_TOPK_UPPER_BOUND, ACLNN_ERR_PARAM_INVALID, "topk should be [%lld, %lld], but got %lld", QLI_V2_TOPK_LOWER_BOUND, QLI_V2_A3_TOPK_UPPER_BOUND, diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/CMakeLists.txt b/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/CMakeLists.txt new file mode 100644 index 000000000000..e509bc15e9c3 --- /dev/null +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/CMakeLists.txt @@ -0,0 +1,29 @@ +# --------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# --------------------------------------------------------------------------------------------------------- + +if (BUILD_WITH_INSTALLED_DEPENDENCY_CANN_PKG) + if (NOT (UT_TEST_ALL OR OP_KERNEL_AICPU_UT)) + add_definitions(-D_GLIBCXX_USE_CXX11_ABI=1) + set(CMAKE_CXX_COMPILER ${ASCEND_DIR}/toolkit/toolchain/hcc/bin/aarch64-target-linux-gnu-g++) + endif() + + # aicpu json + file(GLOB_RECURSE JSON_FILE ${CMAKE_CURRENT_SOURCE_DIR}/*.json) + + # aicpu cust kernel + file(GLOB AICPU_SRC ${CMAKE_CURRENT_SOURCE_DIR}/*_aicpu*.cpp) + message(STATUS "[quant_lightning_indexer_v2_metadata] Found aicpu sources: ${AICPU_SRC}, ascend dir: ${ASCEND_DIR}, ophsot name: ${OPHOST_NAME}") + + add_aicpu_cust_kernel_modules(quant_lightning_indexer_v2_metadata ${AICPU_SRC} ${JSON_FILE}) +endif() + +if(UT_TEST_ALL OR OP_KERNEL_AICPU_UT) + AddAicpuOpTestCase(quant_lightning_indexer_v2_metadata) +endif() diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.cpp b/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.cpp index d6d563871566..7ea9f61f163c 100644 --- a/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.cpp +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.cpp @@ -61,7 +61,6 @@ bool QuantLightningIndexerV2MetadataCpuKernel::Prepare(CpuKernelContext &ctx) bool QuantLightningIndexerV2MetadataCpuKernel::ParamsCheck() { - // 校验输出 metadata 是否为空 if (metadata_ == nullptr) { KERNEL_LOG_ERROR("Output metadata is nullptr"); return false; @@ -69,113 +68,131 @@ bool QuantLightningIndexerV2MetadataCpuKernel::ParamsCheck() KERNEL_LOG_ERROR("Output metadata data is nullptr"); return false; } - int32_t batchSize = GetQueryBatchSize(); - // 校验 cu_seqlens_q 元素 - if (layoutQ_ == "TND") { - if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) { - const int32_t *cuSeqlensQPtr = static_cast(cuSeqlensQ_->GetData()); - // 校验 cu_seqlens_q 首元素为 0 - if (cuSeqlensQPtr[0] != 0) { - KERNEL_LOG_ERROR("The first element of cu_seqlens_q should be 0, but got %d", cuSeqlensQPtr[0]); - return false; - } - for (int i = 0; i < batchSize + 1; i++) { - // 校验 cu_seqlens_q 元素递增 - if (i > 0 && cuSeqlensQPtr[i - 1] > cuSeqlensQPtr[i]) { - KERNEL_LOG_ERROR("The elements in cu_seqlens_q must be in ascending order, " - "but got cu_seqlens_q[%d] = %d, cu_seqlens_q[%d] = %d", - i - 1, cuSeqlensQPtr[i - 1], i, cuSeqlensQPtr[i]); - return false; - } - } + const int32_t batchSize = GetQueryBatchSize(); + return CheckCuSeqlensQ(batchSize) && CheckCuSeqlensK(batchSize) && CheckSequsedQ(batchSize) && + CheckSequsedK(batchSize) && CheckCmpResidualK(batchSize); +} + +bool QuantLightningIndexerV2MetadataCpuKernel::CheckCuSeqlensQ(int32_t batchSize) const +{ + if (layoutQ_ != "TND" || cuSeqlensQ_ == nullptr || cuSeqlensQ_->GetData() == nullptr) { + return true; + } + const int32_t *cuSeqlensQPtr = static_cast(cuSeqlensQ_->GetData()); + if (cuSeqlensQPtr[0] != 0) { + KERNEL_LOG_ERROR("The first element of cu_seqlens_q should be 0, but got %d", cuSeqlensQPtr[0]); + return false; + } + for (int32_t i = 1; i < batchSize + 1; ++i) { + if (cuSeqlensQPtr[i - 1] > cuSeqlensQPtr[i]) { + KERNEL_LOG_ERROR("The elements in cu_seqlens_q must be in ascending order, " + "but got cu_seqlens_q[%d] = %d, cu_seqlens_q[%d] = %d", + i - 1, cuSeqlensQPtr[i - 1], i, cuSeqlensQPtr[i]); + return false; } } - // 校验 cu_seqlens_k 元素 - if (layoutK_ == "TND") { - if (cuSeqlensK_ != nullptr && cuSeqlensK_->GetData() != nullptr) { - const int32_t *cuSeqlensKPtr = static_cast(cuSeqlensK_->GetData()); - // 校验 cu_seqlens_k 首元素为 0 - if (cuSeqlensKPtr[0] != 0) { - KERNEL_LOG_ERROR("The first element of cu_seqlens_k should be 0, but got %d", cuSeqlensKPtr[0]); - return false; - } - for (int i = 0; i < batchSize + 1; i++) { - // 校验 cu_seqlens_k 元素递增 - if (i > 0 && cuSeqlensKPtr[i - 1] > cuSeqlensKPtr[i]) { - KERNEL_LOG_ERROR("The elements in cu_seqlens_k must be in ascending order, " - "but got cu_seqlens_k[%d] = %d, cu_seqlens_k[%d] = %d", - i - 1, cuSeqlensKPtr[i - 1], i, cuSeqlensKPtr[i]); - return false; - } - } + return true; +} + +bool QuantLightningIndexerV2MetadataCpuKernel::CheckCuSeqlensK(int32_t batchSize) const +{ + if (layoutK_ != "TND" || cuSeqlensK_ == nullptr || cuSeqlensK_->GetData() == nullptr) { + return true; + } + const int32_t *cuSeqlensKPtr = static_cast(cuSeqlensK_->GetData()); + if (cuSeqlensKPtr[0] != 0) { + KERNEL_LOG_ERROR("The first element of cu_seqlens_k should be 0, but got %d", cuSeqlensKPtr[0]); + return false; + } + for (int32_t i = 1; i < batchSize + 1; ++i) { + if (cuSeqlensKPtr[i - 1] > cuSeqlensKPtr[i]) { + KERNEL_LOG_ERROR("The elements in cu_seqlens_k must be in ascending order, " + "but got cu_seqlens_k[%d] = %d, cu_seqlens_k[%d] = %d", + i - 1, cuSeqlensKPtr[i - 1], i, cuSeqlensKPtr[i]); + return false; } } - // 校验 seqused_q 元素非负 - if (sequsedQ_ != nullptr && sequsedQ_->GetData() != nullptr) { - const int32_t *sequsedQPtr = static_cast(sequsedQ_->GetData()); - const int32_t *cuSeqlensQPtr = (layoutQ_ == "TND" && cuSeqlensQ_ != nullptr && - cuSeqlensQ_->GetData() != nullptr) ? - static_cast(cuSeqlensQ_->GetData()) : nullptr; - for (int i = 0; i < batchSize; i++) { - if (sequsedQPtr[i] < 0) { - KERNEL_LOG_ERROR("The elements in seqused_q should be >= 0, but got seqused_q[%d] = %d", i, - sequsedQPtr[i]); - return false; - } - // 校验 seqused_q 元素不大于 max_seqlen_q (BSND) 或 cu_seqlens_q 序列长度 (TND) - if (layoutQ_ == "BSND" && sequsedQPtr[i] > maxSeqlenQ_) { - KERNEL_LOG_ERROR("The elements in seqused_q should not be greater than max_seqlen_q %d, " - "but got seqused_q[%d] = %d", maxSeqlenQ_, i, sequsedQPtr[i]); + return true; +} + +bool QuantLightningIndexerV2MetadataCpuKernel::CheckSequsedQ(int32_t batchSize) const +{ + if (sequsedQ_ == nullptr || sequsedQ_->GetData() == nullptr) { + return true; + } + const int32_t *sequsedQPtr = static_cast(sequsedQ_->GetData()); + const int32_t *cuSeqlensQPtr = (layoutQ_ == "TND" && cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) ? + static_cast(cuSeqlensQ_->GetData()) : + nullptr; + for (int32_t i = 0; i < batchSize; ++i) { + if (sequsedQPtr[i] < 0) { + KERNEL_LOG_ERROR("The elements in seqused_q should be >= 0, but got seqused_q[%d] = %d", i, sequsedQPtr[i]); + return false; + } + if (layoutQ_ == "BSND" && sequsedQPtr[i] > maxSeqlenQ_) { + KERNEL_LOG_ERROR("The elements in seqused_q should not be greater than max_seqlen_q %d, " + "but got seqused_q[%d] = %d", + maxSeqlenQ_, i, sequsedQPtr[i]); + return false; + } + if (cuSeqlensQPtr != nullptr) { + const int32_t seqLen = cuSeqlensQPtr[i + 1] - cuSeqlensQPtr[i]; + if (sequsedQPtr[i] > seqLen) { + KERNEL_LOG_ERROR("The elements in seqused_q should not be greater than the sequence length " + "from cu_seqlens_q %d, but got seqused_q[%d] = %d", + seqLen, i, sequsedQPtr[i]); return false; } - if (cuSeqlensQPtr != nullptr) { - int32_t seqLen = cuSeqlensQPtr[i + 1] - cuSeqlensQPtr[i]; - if (sequsedQPtr[i] > seqLen) { - KERNEL_LOG_ERROR("The elements in seqused_q should not be greater than the sequence length " - "from cu_seqlens_q %d, but got seqused_q[%d] = %d", seqLen, i, sequsedQPtr[i]); - return false; - } - } } } - // 校验 seqused_k 元素非负 - if (sequsedK_ != nullptr && sequsedK_->GetData() != nullptr) { - const int32_t *sequsedKPtr = static_cast(sequsedK_->GetData()); - const int32_t *cuSeqlensKPtr = (layoutK_ == "TND" && cuSeqlensK_ != nullptr && - cuSeqlensK_->GetData() != nullptr) ? - static_cast(cuSeqlensK_->GetData()) : nullptr; - for (int i = 0; i < batchSize; i++) { - if (sequsedKPtr[i] < 0) { - KERNEL_LOG_ERROR("The elements in seqused_k should be >= 0, but got seqused_k[%d] = %d", i, - sequsedKPtr[i]); - return false; - } - // 校验 seqused_k 元素不大于 max_seqlen_k (BSND) 或 cu_seqlens_k 序列长度 (TND) - if (layoutK_ == "BSND" && sequsedKPtr[i] > maxSeqlenK_) { - KERNEL_LOG_ERROR("The elements in seqused_k should not be greater than max_seqlen_k %d, " - "but got seqused_k[%d] = %d", maxSeqlenK_, i, sequsedKPtr[i]); + return true; +} + +bool QuantLightningIndexerV2MetadataCpuKernel::CheckSequsedK(int32_t batchSize) const +{ + if (sequsedK_ == nullptr || sequsedK_->GetData() == nullptr) { + return true; + } + const int32_t *sequsedKPtr = static_cast(sequsedK_->GetData()); + const int32_t *cuSeqlensKPtr = (layoutK_ == "TND" && cuSeqlensK_ != nullptr && cuSeqlensK_->GetData() != nullptr) ? + static_cast(cuSeqlensK_->GetData()) : + nullptr; + for (int32_t i = 0; i < batchSize; ++i) { + if (sequsedKPtr[i] < 0) { + KERNEL_LOG_ERROR("The elements in seqused_k should be >= 0, but got seqused_k[%d] = %d", i, sequsedKPtr[i]); + return false; + } + if (layoutK_ == "BSND" && sequsedKPtr[i] > maxSeqlenK_) { + KERNEL_LOG_ERROR("The elements in seqused_k should not be greater than max_seqlen_k %d, " + "but got seqused_k[%d] = %d", + maxSeqlenK_, i, sequsedKPtr[i]); + return false; + } + if (cuSeqlensKPtr != nullptr) { + const int32_t seqLen = cuSeqlensKPtr[i + 1] - cuSeqlensKPtr[i]; + if (sequsedKPtr[i] > seqLen) { + KERNEL_LOG_ERROR("The elements in seqused_k should not be greater than the sequence length " + "from cu_seqlens_k %d, but got seqused_k[%d] = %d", + seqLen, i, sequsedKPtr[i]); return false; } - if (cuSeqlensKPtr != nullptr) { - int32_t seqLen = cuSeqlensKPtr[i + 1] - cuSeqlensKPtr[i]; - if (sequsedKPtr[i] > seqLen) { - KERNEL_LOG_ERROR("The elements in seqused_k should not be greater than the sequence length " - "from cu_seqlens_k %d, but got seqused_k[%d] = %d", seqLen, i, sequsedKPtr[i]); - return false; - } - } } } - // 校验 cmp_residual_k 元素 - if (cmpResidualK_ != nullptr && cmpResidualK_->GetData() != nullptr) { - const int32_t *cmpResidualKPtr = static_cast(cmpResidualK_->GetData()); - for (int i = 0; i < batchSize; i++) { - if (cmpResidualKPtr[i] < 0 || cmpResidualKPtr[i] >= cmpRatio_) { - KERNEL_LOG_ERROR("The elements in cmp_residual_k should be in [0, cmpRatio_(%d)), but got " - "cmp_residual_k[%d] = %d", cmpRatio_, - i, cmpResidualKPtr[i]); - return false; - } + return true; +} + +bool QuantLightningIndexerV2MetadataCpuKernel::CheckCmpResidualK(int32_t batchSize) const +{ + if (cmpResidualK_ == nullptr || cmpResidualK_->GetData() == nullptr) { + return true; + } + const int32_t *cmpResidualKPtr = static_cast(cmpResidualK_->GetData()); + for (int32_t i = 0; i < batchSize; ++i) { + if (cmpResidualKPtr[i] < 0 || cmpResidualKPtr[i] >= cmpRatio_) { + KERNEL_LOG_ERROR("The elements in cmp_residual_k should be in [0, cmpRatio_(%d)), but got " + "cmp_residual_k[%d] = %d", + cmpRatio_, i, cmpResidualKPtr[i]); + return false; } } return true; @@ -224,13 +241,13 @@ bool QuantLightningIndexerV2MetadataCpuKernel::ParamsInit() } else if (mode == SparseMode::BAND) { attentionMode_ = 1; } - groupSize_ = numHeadsQ_ / numHeadsK_; + groupSize_ = static_cast(numHeadsQ_) / static_cast(numHeadsK_); ValidSocVersion validSocVersion = ProcessSocVersion(); if (validSocVersion == ValidSocVersion::ASCEND910B) { mBaseSize_ = s1BaseSize_ * groupSize_; s2BaseSize_ = 2048U; } else if (validSocVersion == ValidSocVersion::ASCEND950) { - if (topk_ > TOPK_6K) { + if (static_cast(topk_) > TOPK_6K) { s1BaseSize_ = S1_BASE_SIZE_SMALL; } mBaseSize_ = s1BaseSize_ * groupSize_; @@ -285,7 +302,8 @@ uint64_t QuantLightningIndexerV2MetadataCpuKernel::GetRevertS2Size(uint32_t bIdx uint32_t cmpS2Size = GetS2SeqSize(bIdx); if (cmpResidualK_ != nullptr && cmpResidualK_->GetData() != nullptr) { const int32_t *residualPtr = static_cast(cmpResidualK_->GetData()); - return static_cast(cmpS2Size) * static_cast(cmpRatio_) + residualPtr[bIdx]; + return static_cast(cmpS2Size) * static_cast(cmpRatio_) + + static_cast(residualPtr[bIdx]); } else { return static_cast(cmpS2Size) * static_cast(cmpRatio_); } @@ -295,7 +313,7 @@ void QuantLightningIndexerV2MetadataCpuKernel::CalcSplitInfo(SplitContext &split { // 计算每个batch的切分,统计是否为空batch,记录最后有效batch(每个batch的每个N2切分是一样的) SplitInfo &splitInfo = splitContext.splitInfo; - for (uint32_t bIdx = 0; bIdx < batchSize_; bIdx++) { + for (uint32_t bIdx = 0; bIdx < static_cast(batchSize_); bIdx++) { uint32_t s1Size = GetS1SeqSize(bIdx); uint32_t s2Size = GetS2SeqSize(bIdx); maxS2Size_ = std::max(maxS2Size_, s2Size); @@ -309,7 +327,7 @@ void QuantLightningIndexerV2MetadataCpuKernel::CalcSplitInfo(SplitContext &split } ValidSocVersion validSocVersion = ProcessSocVersion(); if (validSocVersion == ValidSocVersion::ASCEND950) { - if (maxS2Size_ < topk_) { + if (maxS2Size_ < static_cast(topk_)) { supportFd_ = false; return; } @@ -320,7 +338,7 @@ void QuantLightningIndexerV2MetadataCpuKernel::CalcSplitInfo(SplitContext &split } } -int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcPreTokenLeftUp(uint32_t s1Size, uint64_t s2Size) +int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcPreTokenLeftUp(uint32_t s1Size, uint64_t s2Size) const { auto mode = static_cast(maskMode_); if (mode == SparseMode::BAND) { @@ -329,7 +347,7 @@ int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcPreTokenLeftUp(uint32_t s1 return preToken_; } -int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcNextTokenLeftUp(uint32_t s1Size, uint64_t s2Size) +int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcNextTokenLeftUp(uint32_t s1Size, uint64_t s2Size) const { auto mode = static_cast(maskMode_); switch (mode) { @@ -346,7 +364,7 @@ int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcNextTokenLeftUp(uint32_t s } } -int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcCost(uint32_t basicM, uint32_t basicS2) +int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcCost(uint32_t basicM, uint32_t basicS2) const { uint32_t alignCoefM = 16U; uint32_t alignCoefS2 = 64U; @@ -356,7 +374,8 @@ int64_t QuantLightningIndexerV2MetadataCpuKernel::CalcCost(uint32_t basicM, uint } BlockCost QuantLightningIndexerV2MetadataCpuKernel::CalcCostTable(uint32_t s1NormalSize, uint32_t s2NormalSize, - uint32_t s1GTailSize, uint32_t s2TailSize) + uint32_t s1GTailSize, + uint32_t s2TailSize) const { BlockCost typeCost{}; typeCost[NORMAL_BLOCK][NORMAL_BLOCK] = CalcCost(s1NormalSize, s2NormalSize); @@ -523,10 +542,10 @@ void QuantLightningIndexerV2MetadataCpuKernel::CalcCostInfo(SplitContext &splitC } // 计算batch的负载并记录,用于按batch分配,需要按行计算起止点,统计块数、负载 - for (uint32_t bIdx = 0; bIdx < batchSize_; bIdx++) { + for (uint32_t bIdx = 0; bIdx < static_cast(batchSize_); bIdx++) { CalcBatchCost(bIdx, splitContext, costInfo); costInfo.totalCost += costInfo.bN2CostOfEachBatch[bIdx] * numHeadsK_; - costInfo.totalBlockNum += costInfo.bN2BlockOfEachBatch[bIdx] * numHeadsK_; + costInfo.totalBlockNum += costInfo.bN2BlockOfEachBatch[bIdx] * static_cast(numHeadsK_); } } @@ -553,15 +572,16 @@ void QuantLightningIndexerV2MetadataCpuKernel::UpdateCursor(const SplitContext & } // Update Batch - if (assignContext.curBN2Idx == batchSize_ * numHeadsK_) { // 所有负载全部分配完,设置最后一个核的右开区间,返回 + if (assignContext.curBN2Idx == + static_cast(batchSize_) * static_cast(numHeadsK_)) { // 所有负载全部分配完 assignContext.curS1GIdx = 0U; assignContext.curS2Idx = 0U; assignContext.isFinished = true; return; } - if (assignContext.curBN2Idx / numHeadsK_ != assignContext.curBIdx) { - assignContext.curBIdx = assignContext.curBN2Idx / numHeadsK_; + if (assignContext.curBN2Idx / static_cast(numHeadsK_) != assignContext.curBIdx) { + assignContext.curBIdx = assignContext.curBN2Idx / static_cast(numHeadsK_); assignContext.curS1GIdx = 0U; UpdateBatch = true; UpdateS1G = true; @@ -595,7 +615,7 @@ void QuantLightningIndexerV2MetadataCpuKernel::AssignByBatch(const SplitContext assignContext.curBN2Idx++; // to the end - if (assignContext.curBN2Idx == batchSize_ * numHeadsK_) { + if (assignContext.curBN2Idx == static_cast(batchSize_) * static_cast(numHeadsK_)) { assignContext.curS1GIdx = 0U; assignContext.curS2Idx = 0U; assignContext.isFinished = true; @@ -603,8 +623,8 @@ void QuantLightningIndexerV2MetadataCpuKernel::AssignByBatch(const SplitContext } // next batch - if (assignContext.curBN2Idx / numHeadsK_ != assignContext.curBIdx) { - assignContext.curBIdx = assignContext.curBN2Idx / numHeadsK_; + if (assignContext.curBN2Idx / static_cast(numHeadsK_) != assignContext.curBIdx) { + assignContext.curBIdx = assignContext.curBN2Idx / static_cast(numHeadsK_); CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache); } @@ -721,7 +741,7 @@ void QuantLightningIndexerV2MetadataCpuKernel::RecordFDInfo(const SplitContext & { const SplitInfo &splitInfo = splitContext.splitInfo; // 需要规约的行是上一个核的切分点所在位置 - uint32_t splitBIdx = result.bN2End[assignContext.curCoreIdx - 1U] / numHeadsK_; + uint32_t splitBIdx = result.bN2End[assignContext.curCoreIdx - 1U] / static_cast(numHeadsK_); uint32_t splitS1GIdx = result.gS1End[assignContext.curCoreIdx - 1U]; uint32_t s1Size = GetS1SeqSize(splitBIdx); @@ -811,6 +831,13 @@ void QuantLightningIndexerV2MetadataCpuKernel::CalcSplitPlan(int64_t costLimit, } assignContext.curCoreIdx = i; AssignBlockToCore(splitContext, assignContext, result); + // Keep trailing zero-cost batches in the last AIC range so their outputs are initialized. + if (!assignContext.isFinished && assignContext.unassignedCost <= 0) { + result.bN2End[i] = static_cast(batchSize_) * static_cast(numHeadsK_); + result.gS1End[i] = 0U; + result.s2End[i] = 0U; + assignContext.isFinished = true; + } } result.usedCoreNum = assignContext.curCoreIdx + 1; } @@ -867,7 +894,7 @@ bool QuantLightningIndexerV2MetadataCpuKernel::BalanceSchedule(SplitResult &spli // 全空case if (splitContext.splitInfo.isKvSeqAllZero) { splitRes.usedCoreNum = 1U; - splitRes.bN2End[0] = batchSize_ * numHeadsK_; + splitRes.bN2End[0] = static_cast(batchSize_) * static_cast(numHeadsK_); splitRes.gS1End[0] = 0U; splitRes.s2End[0] = 0U; return true; diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.h b/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.h index 6cbc92d8d130..9552fca1c164 100644 --- a/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.h +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.h @@ -33,7 +33,11 @@ constexpr uint32_t COST_WEIGHT_S2 = 10U; constexpr uint32_t S1_BASE_SIZE_SMALL = 2; constexpr uint32_t TOPK_6K = 6144; -enum BlockType : uint32_t { NORMAL_BLOCK = 0, TAIL_BLOCK, BLOCK_MAX_TYPE }; +enum BlockType : uint32_t { + NORMAL_BLOCK = 0, + TAIL_BLOCK, + BLOCK_MAX_TYPE +}; enum class SparseMode : uint8_t { DEFAULT_MASK = 0, @@ -44,7 +48,11 @@ enum class SparseMode : uint8_t { SPARSE_BUTT, }; -enum class ValidSocVersion { ASCEND910B = 0, ASCEND950, RESERVED_VERSION = 99999 }; +enum class ValidSocVersion { + ASCEND910B = 0, + ASCEND950, + RESERVED_VERSION = 99999 +}; template using Range = std::pair; @@ -109,7 +117,11 @@ struct SplitResult { FlashDecodeResult fdRes{0U, 0U}; // FD信息 SplitResult(uint32_t aicNum, uint32_t aivNum) - : bN2End(aicNum), gS1End(aicNum), s2End(aicNum), firstFdDataWorkspaceIdx(aicNum), fdRes(aicNum, aivNum) {}; + : bN2End(aicNum), + gS1End(aicNum), + s2End(aicNum), + firstFdDataWorkspaceIdx(aicNum), + fdRes(aicNum, aivNum) {}; }; // 分核功能模块内部使用:记录切分信息 @@ -121,7 +133,10 @@ struct SplitInfo { bool isKvSeqAllZero{true}; explicit SplitInfo(uint32_t batchSize) - : s1GBaseNum(batchSize), s2BaseNum(batchSize), s1GTailSize(batchSize), s2TailSize(batchSize) + : s1GBaseNum(batchSize), + s2BaseNum(batchSize), + s1GTailSize(batchSize), + s2TailSize(batchSize) {} }; @@ -135,7 +150,9 @@ struct CostInfo { uint64_t maxS1GCost{0}; // 新增 explicit CostInfo(uint32_t batchSize) - : bN2CostOfEachBatch(batchSize), bN2BlockOfEachBatch(batchSize), bN2LastBlockCostOfEachBatch(batchSize) + : bN2CostOfEachBatch(batchSize), + bN2BlockOfEachBatch(batchSize), + bN2LastBlockCostOfEachBatch(batchSize) {} }; @@ -144,7 +161,10 @@ struct SplitContext { SplitInfo splitInfo{0U}; CostInfo costInfo{0U}; - explicit SplitContext(uint32_t batchSize) : splitInfo(batchSize), costInfo(batchSize) {} + explicit SplitContext(uint32_t batchSize) + : splitInfo(batchSize), + costInfo(batchSize) + {} }; // 分核功能模块内部使用:记录batch相关的临时信息 @@ -203,6 +223,11 @@ class QuantLightningIndexerV2MetadataCpuKernel : public CpuKernel { private: bool Prepare(CpuKernelContext &ctx); bool ParamsCheck(); + bool CheckCuSeqlensQ(int32_t batchSize) const; + bool CheckCuSeqlensK(int32_t batchSize) const; + bool CheckSequsedQ(int32_t batchSize) const; + bool CheckSequsedK(int32_t batchSize) const; + bool CheckCmpResidualK(int32_t batchSize) const; int32_t GetQueryBatchSize(); ValidSocVersion ProcessSocVersion(); bool ParamsInit(); @@ -213,12 +238,12 @@ class QuantLightningIndexerV2MetadataCpuKernel : public CpuKernel { uint32_t GetS1SeqSize(uint32_t bIdx); uint32_t GetS2SeqSize(uint32_t bIdx); uint64_t GetRevertS2Size(uint32_t bIdx); - int64_t CalcPreTokenLeftUp(uint32_t s1Size, uint64_t s2Size); - int64_t CalcNextTokenLeftUp(uint32_t s1Size, uint64_t s2Size); + int64_t CalcPreTokenLeftUp(uint32_t s1Size, uint64_t s2Size) const; + int64_t CalcNextTokenLeftUp(uint32_t s1Size, uint64_t s2Size) const; Range CalcS2TokenRange(uint32_t s1GIdx, const BatchCache &batchCache); - int64_t CalcCost(uint32_t basicM, uint32_t basicS2); + int64_t CalcCost(uint32_t basicM, uint32_t basicS2) const; BlockCost CalcCostTable(uint32_t s1NormalSize, uint32_t s2NormalSize, uint32_t s1GTailSize, - uint32_t s2TailSize); + uint32_t s2TailSize) const; // cache calculation void CalcBatchCache(uint32_t bIdx, const SplitContext &splitContext, BatchCache &batchCache); @@ -247,7 +272,6 @@ class QuantLightningIndexerV2MetadataCpuKernel : public CpuKernel { void CalcSplitPlan(int64_t costLimit, const SplitContext &splitContext, SplitResult &result); private: - CpuKernelContext *context_ = nullptr; // input Tensor *cuSeqlensQ_ = nullptr; Tensor *cuSeqlensK_ = nullptr; diff --git a/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.json b/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.json index a56f16c5316f..6a91c60b5cd2 100644 --- a/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.json +++ b/csrc/attention/quant_lightning_indexer_v2_metadata/op_kernel_aicpu/quant_lightning_indexer_v2_metadata_aicpu.json @@ -12,4 +12,4 @@ "workspaceSize":"100" } } -} \ No newline at end of file +} diff --git a/csrc/attention/sparse_flash_mla/CMakeLists.txt b/csrc/attention/sparse_flash_mla/CMakeLists.txt new file mode 100644 index 000000000000..a3ab34a930dc --- /dev/null +++ b/csrc/attention/sparse_flash_mla/CMakeLists.txt @@ -0,0 +1,19 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +if(NOT ENABLE_TEST AND NOT BENCHMARK) + list(REMOVE_ITEM CURRENT_DIRS tests) +endif() +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() \ No newline at end of file diff --git a/csrc/attention/sparse_flash_mla/README.md b/csrc/attention/sparse_flash_mla/README.md new file mode 100644 index 000000000000..de8339a573d3 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/README.md @@ -0,0 +1,324 @@ +# SparseFlashMla + +开发者学习入口:[算子设计与开发指南](./docs/design.md),包含 tiling、Metadata 分核、内存分配、计算流水、架构差异与验证建议。 + +倍率扩展:[A2/A3 cmp_ratio=1/2 适配说明](./docs/ratio2_a2a3.md)。 + +## 产品支持情况 + +| 产品 | 是否支持 | +| :------------------------------------------------------------ | :------: | +|Ascend 950PR/Ascend 950DT | √ | +|Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ | +|Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ | +|Atlas 200I/500 A2推理系列产品 | × | +|Atlas 推理系列产品 | × | +|Atlas 训练系列产品 | × | + +## 功能说明 + +- 算子功能: + + `SparseFlashMla`算子旨在完成以下公式描述的Attention计算,支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)三类Attention计算场景。调用时需要使用`SparseFlashMlaMetadata`生成的任务列表`metadata`。 + + 典型调用流程如下: + + 1. 准备`q`、`ori_kv`、`cmp_kv`、序列长度、`block table`、`sinks`等输入。 + 2. 调用`SparseFlashMlaMetadata`生成`metadata`。 + 3. 调用`SparseFlashMla`,将上一步得到的`metadata`传入主算子。 + +- 计算公式: + + $$ + O = \text{softmax}(Q@\tilde{K}^T \cdot \text{softmax\_scale})@\tilde{V} + $$ + + 其中$\tilde{K}=\tilde{V}$为基于ori_kv、cmp_kv以及cmp_ratio等入参控制的实际参与计算的 $KV$。 + +## 参数说明 +> +> **说明:**
+> 参数维度含义:B表示Batch Size,Q_S、ORI\_KV\_S和CMP\_KV\_S分别表示query、oriKv和cmpKv的Sequence Length,Q_N和KV_N分别表示query和key/value的Head Num,Q_T、ORI_KV_T和CMP_KV_T分别表示query、oriKv和cmpKv的Total Tokens。 + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出/属性描述数据类型数据格式
q输入对应公式中的Q。BFLOAT16、FLOAT16ND
ori_kv可选输入对应公式中K和V的一部分,表示原始不经压缩的KV。BFLOAT16、FLOAT16ND
cmp_kv可选输入对应公式中K和V的一部分,表示经过压缩的KV。BFLOAT16、FLOAT16ND
ori_sparse_indices可选输入SWA稀疏ori_kv场景表示从ori_kv中离散取数的逻辑索引,-1表示无效或填充slot,shape为(T1, N2, K)或(B, S1, N2, K)。INT32ND
cmp_sparse_indices可选输入表示从cmp_kv中离散取数的索引。INT32ND
ori_block_table可选输入表示PageAttention中ori_kv使用的block映射表。INT32ND
cmp_block_table可选输入表示PageAttention中cmp_kv使用的block映射表。INT32ND
cu_seqlens_q可选输入表示TND布局下不同batch中q的累积序列长度。INT32ND
cu_seqlens_ori_kv可选输入表示TND布局下不同batch中ori_kv的累积序列长度。INT32ND
cu_seqlens_cmp_kv可选输入表示TND布局下不同batch中cmp_kv的累积序列长度。INT32ND
seqused_q可选输入表示不同batch中q实际参与计算的token数。INT32ND
seqused_ori_kv可选输入表示不同batch中ori_kv实际参与计算的token数。INT32ND
seqused_cmp_kv可选输入表示不同batch中cmp_kv实际参与计算的token数。INT32ND
cmp_residual_kv可选输入表示压缩KV余数,用于恢复cmp侧mask使用的压缩前KV长度。INT32ND
ori_topk_length可选输入SWA稀疏ori_kv场景表示不同q token对应的ori_kv关键稀疏token的实际个数,shape为(T1, N2)或(B, S1, N2)。INT32ND
cmp_topk_length可选输入预留输入,当前版本不支持传入非空Tensor。INT32ND
sinks可选输入表示attention sinks输入。FLOATND
metadata可选输入配套metadata前置接口生成的任务切分结果。INT32ND
softmax_scale可选属性对应公式中的softmax_scale。FLOAT-
cmp_ratio可选属性表示cmp_kv相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度;仅传入ori_kv时不参与压缩KV计算。支持1到128。INT-
ori_mask_mode可选属性表示q和ori_kv计算的mask模式。
0: No Mask。
3: RightDownCausal模式。
4: Band模式。
INT-
cmp_mask_mode可选属性表示q和cmp_kv计算的mask模式。
0: No Mask。
3: RightDownCausal模式。
INT-
ori_win_left可选属性表示q和ori_kv计算中q对过去token计算的数量,支持-1或非负数,其中-1表示窗口不受限。INT-
ori_win_right可选属性表示q和ori_kv计算中q对未来token计算的数量,支持-1或非负数,其中-1表示窗口不受限。INT-
layout_q可选属性表示输入q的数据排布格式,支持"BSND"和"TND"。STRING-
layout_kv可选属性表示输入ori_kv和cmp_kv的数据排布格式,支持"BSND"、"TND"和"PA_BBND"。STRING-
topk_value_mode可选属性表示TopK索引取值模式。INT-
return_softmax_lse可选属性表示是否返回softmax_lse。BOOL-
attn_out输出对应公式中的输出O。BFLOAT16、FLOAT16ND
softmax_lse可选输出返回softmax的log-sum-exp结果。FLOATND
+ +## 约束说明 + +- 该接口支持训练、推理场景下使用。 +- 该接口支持aclgraph模式。 +- 该接口支持batch一致性。 +- 该接口当前支持四种计算场景:SWA(Sliding Window Attention)场景仅传入`ori_kv`;SWA稀疏ori_kv场景传入`ori_kv`、`ori_sparse_indices`及`ori_topk_length`;CSA(Compressed Sparse Attention)场景传入`ori_kv`、`cmp_kv`及`cmp_sparse_indices`;HCA(Heavily Compressed Attention)场景传入`ori_kv`及`cmp_kv`。 +- 通用规格约束如下: + - KV\_N仅支持1,D仅支持512。其中,`ori_kv`和`cmp_kv`的D_kv由nope(448)和rope(64)拼接而成。 + - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时,`cmp_ratio`不参与压缩KV计算,取值为0;传入`cmp_kv`时支持1到128。 + - `ori_mask_mode`、`cmp_mask_mode`、`ori_win_left`和`ori_win_right`的取值随产品型号而变化,详见“产品型号约束”。 + - PageAttention的block_size支持1到1024。 + - `layout_q`和`layout_kv`组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下`layout_q`和`layout_kv`必须一致。 + - SWA稀疏ori_kv场景下,`ori_topk_length`必须传入,配套Metadata接口的`ori_topk`为`ori_sparse_indices`最后一维K,且`ori_topk_length`的元素取值应在[0, K]范围内;其他场景`ori_topk_length`传入nullptr或空Tensor。 +- 产品型号约束如下: + - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:Q\_N支持1、2、4、8、16、32、64、128,KV\_N只支持1;cmp_ratio在SWA场景传入0,CSA支持传入1、2或4,HCA支持传入128;block_size取值为16的倍数,最大支持1024;SWA稀疏ori_kv场景支持`ori_sparse_indices`和`ori_topk_length`,`ori_mask_mode`为0,`ori_win_left`和`ori_win_right`为非负数;非SWA稀疏ori_kv场景的`ori_mask_mode`为4、`ori_win_left`为127、`ori_win_right`为0,`cmp_sparse_indices`的最后一维K2当前支持512或1024,`cmp_mask_mode`仅支持3。 + - Ascend 950PR/Ascend 950DT:Q\_N支持1-128,KV\_N只支持1。`ori_mask_mode`支持0、3、4,`cmp_mask_mode`支持0、3;`ori_win_left`和`ori_win_right`支持-1或非负数,-1表示对应方向不受限。只有`ori_mask_mode`为4时,`ori_win_left`和`ori_win_right`可以>=0。 + +- 当`layout_q`为TND时,功能使用限制如下: + - `q`的shape需要为[Q\_T, Q\_N, D]。 + - SWA稀疏ori_kv场景下,`ori_sparse_indices`的shape为[Q\_T, KV\_N, K],`ori_topk_length`的shape为[Q\_T, KV\_N]。 + - `cmp_sparse_indices`的shape需要为[Q\_T, KV\_N, K2],其中K2为对`cmp_kv`一次离散选取的token数。 + - `cu_seqlens_q`必须传入,shape为[B+1,],第一个数固定为0,即前缀0。后面的每个元素表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。 + +- 当`layout_q`为BSND时,功能使用限制如下: + - `q`的shape需要为[B, Q\_S, Q\_N, D]。 + - SWA稀疏ori_kv场景下,`ori_sparse_indices`的shape为[B, Q\_S, KV\_N, K],`ori_topk_length`的shape为[B, Q\_S, KV\_N]。 + - `cmp_sparse_indices`的shape需要为[B, Q\_S, KV\_N, K2],其中K2为对`cmp_kv`一次离散选取的token数。 + +- PageAttention场景下,功能使用限制如下: + - `ori_kv`和`cmp_kv`的shape分别为[ori\_block\_num, ori\_block\_size, KV\_N, D]和[cmp\_block\_num, cmp\_block\_size, KV\_N, D],其中ori\_block\_num和cmp\_block\_num为PageAttention时block总数,ori\_block\_size和cmp\_block\_size为一个block的token数,ori\_block\_size和cmp\_block\_size取值支持1到1024,KV_N仅支持1。 + - `ori_block_table`和`cmp_block_table`的shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的S2和S3对应的block数量,即ORI\_KV\_S\_max / block\_size和CMP\_KV\_S\_max / block\_size向上取整。 +- `metadata`为算子实际需要使用的分核结果,目前该参数必传,shape大小固定为[1024]。 +- `layout_kv`支持输入"BSND"、"TND"和"PA_BBND",需满足上述`layout_q`和`layout_kv`组合约束。 + - 当输入为PA_BBND时,`seqused_ori_kv`和`ori_block_table`必须传入;当输入为BSND时,`seqused_ori_kv`可用于表达每个batch的`ori_kv`有效长度,且有效长度不超过key/value中的KV_S的大小且不小于0;当输入为TND时,`ori_kv`总长度由`cu_seqlens_ori_kv`表达,若`seqused_ori_kv`被传入,则`ori_kv`的有效长度由`seqused_ori_kv`表达,否则有效长度与总长度相同。 + - 当输入为BSND时,`ori_kv`和`cmp_kv`的layout都必须为BSND,ori_kv的shape为[B, ORI\_KV\_S, KV\_N, D],cmp_kv的shape为[B, CMP\_KV\_S, KV\_N, D]。 + - 当输入为TND时,`cu_seqlens_ori_kv`必须传入;若存在`cmp_kv`,`cu_seqlens_cmp_kv`也必须传入。 +- `return_softmax_lse`为False时返回占位Tensor;为True时返回softmax的log-sum-exp结果。 +- SWA稀疏ori_kv场景仅支持SWA模板,仅传入`ori_kv`,必须同时传入`ori_sparse_indices`和`ori_topk_length`,`ori_win_left`和`ori_win_right`仅在`ori_mask_mode`为4时取非负数;配套Metadata接口的`ori_topk`为K,该场景不传入`cmp_kv`。 +- `ori_topk_length`表示每个q token和KV head的实际有效索引条目数,取值应在[0, K]范围内。`ori_sparse_indices`的[0, ori_topk_length)区间为左对齐的有效索引条目,[ori_topk_length, K)区间为无效或填充条目,建议填-1。 +- 除`cmp_topk_length`等预留输入可不传或传入空Tensor外,其余已传入Tensor不支持为空。 +- `seqused_cmp_kv`为所有`layout_kv`下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度;未传时由`cmp_kv` shape、`cu_seqlens_cmp_kv`或PA block table相关语义推导。 +- `cmp_residual_kv`为主接口和metadata前置接口的可选入参;传入后用于按`cmp_len * cmp_ratio + residual`恢复cmp侧mask使用的压缩前KV长度,其中`cmp_len`优先来自显式传入的`seqused_cmp_kv`。 +- `q`、`ori_kv`、`cmp_kv`数据排布格式支持从多种维度解读,B(Batch)表示输入样本批量大小、S(Seq-Length)表示输入样本序列长度、H(Hidden-Size)表示隐藏层的大小、N(Head-Num)表示多头数、D(Head-Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。 +- Q\_S表示q shape中的S,ORI\_KV\_S表示ori_kv shape中的S,CMP\_KV\_S表示cmp_kv shape中的S;Q\_N表示num\_q\_heads,KV\_N表示num\_ori_kv\_heads和num\_cmp_kv\_heads;Q\_T表示q shape中的输入样本序列长度的累加和。 + +## 调用说明 + +| 调用方式 | 样例代码 | 说明 | +| --------- | ------------------------------------------------------------ | ------------------------------------------------------------ | +| aclnn API | [test_aclnnSparseFlashMla](./examples/test_aclnn_sparse_flash_mla.cpp) | 通过[aclnnSparseFlashMla](./docs/aclnnSparseFlashMla.md)调用SparseFlashMla算子 | +| PyTorch API | [sparse_flash_mla](../../torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md) | 通过`cann_ops_transformer.sparse_flash_mla`调用SparseFlashMla算子 | diff --git a/csrc/attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md b/csrc/attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md new file mode 100644 index 000000000000..b8d94b965d11 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md @@ -0,0 +1,1142 @@ +# aclnnSparseFlashMla + +## 产品支持情况 + + +- Ascend 950PR/Ascend 950DT:支持 + + +- Atlas A3 训练系列产品/Atlas A3 推理系列产品:支持 + + +- Atlas A2 训练系列产品/Atlas A2 推理系列产品:支持 + + +- Atlas 200I/500 A2 推理产品:不支持 + + +- Atlas 推理系列产品:不支持 + + +- Atlas 训练系列产品:不支持 + + +## 功能说明 + +- 接口功能: + + `aclnnSparseFlashMla`算子实现基于共享KV(Key=Value)的稀疏注意力计算,支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)三类Attention计算场景。该算子适用于大语言模型训练、推理场景,通过滑动窗口和KV压缩机制大幅降低长序列注意力计算的开销。调用时需要使用`aclnnSparseFlashMlaMetadata`生成的任务列表`metadata`。 + + **该算子不建议单独使用,建议与aclnnSparseFlashMla算子配合使用,形成完整的工作流。** + + 典型调用流程如下: + + 1. 准备`q`、`ori_kv`、`cmp_kv`、序列长度、`block table`、`sinks`等输入。 + 2. 调用`aclnnSparseFlashMlaMetadata`生成`metadata`。 + 3. 调用`aclnnSparseFlashMla`,将上一步得到的`metadata`传入主算子。 + +- 计算公式: + + $$ + O = \text{softmax}(Q \cdot \tilde{K}^T \cdot \text{softmax\_scale}) \cdot \tilde{V} + $$ + + 其中$\tilde{K} = \tilde{V}$(共享KV),$\tilde{K}$由滑动窗口内的原始KV和因果边界内的压缩KV拼接而成,具体参与计算的KV范围由模板模式和mask参数决定: + + - 滑动窗口部分(oriKv):对第$i_{S1}$个Query token,其因果对角线位置为$\text{ori\_threshold} = S2_{act} - S1_{act} + i_{S1} + 1$,窗口范围为$[\max(\text{ori\_threshold} - \text{ori\_win\_left} - 1, 0), \text{ori\_threshold} + \text{ori\_win\_right})$。 + + - 压缩KV部分(cmpKv):因果边界阈值为$\text{cmp\_threshold} = \lfloor \frac{\text{ori\_threshold}}{\text{cmp\_ratio}} \rfloor$。HCA场景取$[0, \text{cmp\_threshold})$内的连续压缩KV;CSA场景通过TopK索引从压缩KV中按需收集,仅保留$\text{begin\_idx} < \text{cmp\_threshold}$的块。 + + 注意力计算采用Online Softmax(Flash Attention V2),S2方向按512分块循环,sinks作为每行softmax的初始最大值: + + $$ + \text{row\_max}^{(0)} = \text{sinks}[g], \quad \text{row\_sum}^{(0)} = 1.0, \quad O^{(0)} = 0 + $$ + + $$ + S^{(t)} = Q \cdot K_{tile}^{(t)T} \cdot \text{softmax\_scale} + $$ + + $$ + \text{row\_max}^{(t+1)} = \max(\text{row\_max}^{(t)}, \max(S^{(t)}, \text{dim}=-1)) + $$ + + $$ + \text{row\_sum}^{(t+1)} = \exp(\text{row\_max}^{(t)} - \text{row\_max}^{(t+1)}) \cdot \text{row\_sum}^{(t)} + \sum \exp(S^{(t)} - \text{row\_max}^{(t+1)}) + $$ + + $$ + O^{(t+1)} = \exp(\text{row\_max}^{(t)} - \text{row\_max}^{(t+1)}) \cdot O^{(t)} + \exp(S^{(t)} - \text{row\_max}^{(t+1)}) \cdot V_{tile}^{(t)} + $$ + + $$ + O_{final} = O^{(T_{s2})} / \text{row\_sum}^{(T_{s2})} + $$ + +- 符号说明 + + | 符号 | 含义 | + | ------------------- | --------------------------------------------------------- | + | Q | Query输入,形状为[G, D](单行) | + | K_tile_t | 第t个S2分块的KV数据,K=V(共享KV) | + | S_t | 第t个分块的QK缩放注意力分数 | + | P_t | 第t个分块的softmax概率 | + | O_t | 第t个分块后的累加输出 | + | softmax_scale | 缩放系数,通常取每个注意力头维度的倒数平方根 | + | B | Batch Size | + | S1/S1_act | Query序列长度/实际有效长度 | + | S2/S2_act | 原始KV序列长度/实际有效长度 | + | N1 | Query头数 | + | N2 | KV头数 | + | G | GQA分组比,G=N1/N2 | + | D | 每个注意力头的维度 | + | sinks | 注意力汇点,形状为[N1] | + | cmp_ratio | cmpKv的压缩倍率,用于换算cmp侧mask的压缩前KV长度 | + +## 函数原型 + +每个算子分为[两段式接口](../../../docs/zh/context/two_phase_api.md),必须先调用`aclnnSparseFlashMlaGetWorkspaceSize`接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用`aclnnSparseFlashMla`执行实际计算。 + +```c++ +aclnnStatus aclnnSparseFlashMlaGetWorkspaceSize( + const aclTensor *q, + const aclTensor *oriKvOptional, + const aclTensor *cmpKvOptional, + const aclTensor *oriSparseIndicesOptional, + const aclTensor *cmpSparseIndicesOptional, + const aclTensor *oriBlockTableOptional, + const aclTensor *cmpBlockTableOptional, + const aclTensor *cuSeqlensQOptional, + const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, + const aclTensor *sequsedQOptional, + const aclTensor *sequsedOriKvOptional, + const aclTensor *sequsedCmpKvOptional, + const aclTensor *cmpResidualKvOptional, + const aclTensor *oriTopkLengthOptional, + const aclTensor *cmpTopkLengthOptional, + const aclTensor *sinksOptional, + const aclTensor *metadataOptional, + double softmaxScale, + int64_t cmpRatio, + int64_t oriMaskMode, + int64_t cmpMaskMode, + int64_t oriWinLeft, + int64_t oriWinRight, + char *layoutQOptional, + char *layoutKvOptional, + int64_t topkValueMode, + bool returnSoftmaxLse, + const aclTensor *attnOutOut, + const aclTensor *softmaxLseOutOptional, + uint64_t *workspaceSize, + aclOpExecutor **executor) +``` + +```c++ +aclnnStatus aclnnSparseFlashMla( + void *workspace, + uint64_t workspaceSize, + aclOpExecutor *executor, + aclrtStream stream) +``` + +## aclnnSparseFlashMlaGetWorkspaceSize + +- **参数说明** + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensor
q(aclTensor*)输入Query输入张量。不支持空Tensor。N1支持1-128;D仅支持512。N1/N2的平台差异见表格后说明。BFLOAT16、FLOAT16ND +
    +
  • layoutQ为BSND时:(B, S1, N1, D)
  • +
  • layoutQ为TND时:(T1, N1, D)
  • +
+
√
oriKvOptional(aclTensor*)输入原始KV输入张量,Key与Value共享同一份数据。SWA、CSA、HCA场景必须传入。BFLOAT16、FLOAT16ND +
    +
  • layoutKv为PA_BBND时:(ori_block_num, ori_block_size, N2, D),ori_block_size支持1到1024
  • +
  • layoutKv为BSND时:(B, S2, N2, D)
  • +
  • layoutKv为TND时:(T2, N2, D)
  • +
+ N2仅支持1,D仅支持512,由nope(448)和rope(64)拼接而成。 +
√
cmpKvOptional(aclTensor*)输入压缩KV输入张量,Key与Value共享同一份数据。- BFLOAT16、FLOAT16ND +
    +
  • layoutKv为PA_BBND时:(cmp_block_num, cmp_block_size, N2, D),cmp_block_size支持1到1024
  • +
  • layoutKv为BSND时:(B, S3, N2, D)
  • +
  • layoutKv为TND时:(T3, N2, D)
  • +
+ N2仅支持1,D仅支持512,由nope(448)和rope(64)拼接而成。 +
√
oriSparseIndicesOptional(aclTensor*)输入代表离散取oriKvCache的逻辑索引,-1表示无效或填充slot。ori_kv稀疏场景必须传入,其他场景不传入。INT32ND +
    +
  • layoutQ为BSND时:(B, S1, N2, K)
  • +
  • layoutQ为TND时:(T1, N2, K)
  • +
+ 其中K为oriKv的TopK稀疏选择数。 +
√
cmpSparseIndicesOptional(aclTensor*)输入代表离散取cmpKvCache的TopK索引。cmp_kv稀疏场景必须传入,其他场景不传入。INT32ND +
    +
  • layoutQ为BSND时:(B, S1, N2, K2)
  • +
  • layoutQ为TND时:(T1, N2, K2)
  • +
+ 其中K2为cmpKv的TopK稀疏选择数。 +
√
oriBlockTableOptional(aclTensor*)输入PageAttention中oriKvCache存储使用的block映射表。layoutKv为PA_BBND时必须传入。第二维长度不小于所有batch中最大的S2对应的block数量。INT32ND(B, ori_max_block_num_per_batch)√
cmpBlockTableOptional(aclTensor*)输入PageAttention中cmpKvCache存储使用的block映射表。cmpKv传入且layoutKv为PA_BBND时必须传入。INT32ND(B, cmp_max_block_num_per_batch)√
cuSeqlensQOptional(aclTensor*)输入表示不同Batch中q的有效token数(前缀和形式)。layoutQOptional为TND时必须传入。每个元素表示当前batch与之前所有batch的token数总和。INT32ND(B+1,)√
cuSeqlensOriKvOptional(aclTensor*)输入表示不同Batch中oriKv的有效token数(前缀和形式)。layoutKvOptional为TND时必须传入。INT32ND(B+1,)√
cuSeqlensCmpKvOptional(aclTensor*)输入表示不同Batch中cmpKv的有效token数(前缀和形式)。layoutKvOptional为TND且存在cmpKvOptional时必须传入。INT32ND(B+1,)√
sequsedQOptional(aclTensor*)输入表示不同Batch中q实际参与运算的token数。当前暂不支持指定该参数。INT32ND(B,)√
sequsedOriKvOptional(aclTensor*)输入表示不同Batch中oriKv实际参与运算的token数。layoutKvOptional为PA_BBND时必须传入;layoutKvOptional为BSND时可选传入,用于指定每个batch的oriKv有效长度;layoutKvOptional为TND时使用cuSeqlensOriKvOptional表达序列边界。INT32ND(B,)√
sequsedCmpKvOptional(aclTensor*)输入表示不同Batch中cmpKv实际参与运算的token数。可选输入。传入时shape必须为(B,),作为每个batch的cmp逻辑有效长度,优先于cmpKvOptional shape、cuSeqlensCmpKvOptional或PA block table推导。INT32ND(B,)√
cmpResidualKvOptional(aclTensor*)输入压缩KV余数,用于恢复cmp侧mask使用的压缩前KV长度。可选输入。传入时shape必须为(B,),第b个batch按cmp_len * cmpRatio + cmpResidualKvOptional[b]恢复压缩前KV长度;在cmpRatio不等于1且cmpMaskMode为3场景必传。INT32ND(B,)√
oriTopkLengthOptional(aclTensor*)输入表示不同q token对应的oriKv关键稀疏token的实际个数。ori_kv稀疏的场景必须传入,其他场景传入nullptr或空Tensor。INT32NDlayoutQ为BSND时:(B, S1, N2);layoutQ为TND时:(T1, N2)。shape必须与oriSparseIndicesOptional去掉最后一维K后保持一致。√
cmpTopkLengthOptional(aclTensor*)输入表示不同q token对应的cmpKv关键稀疏token的实际个数。必须传入nullptr或空Tensor;传入非空Tensor会返回参数错误。INT32ND-√
sinksOptional(aclTensor*)输入注意力汇点tensor,作为每行softmax的初始最大值。必须传入。FLOAT32ND(N1,)√
metadataOptional(aclTensor*)输入AICPU算子aclnnSparseFlashMlaMetadata的分核结果。必须传入。由aclnnSparseFlashMlaMetadata算子生成。INT32ND(1024,)√
softmaxScale(double)输入缩放系数,对应公式中的softmaxScale。建议值为1/√D,其中D为每个注意力头的维度。----
cmpRatio(int64_t)输入cmpKv相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度。支持1到128;仅传入oriKv时不参与压缩KV计算。----
oriMaskMode(int64_t)输入q和oriKv计算的mask模式。0: No Mask。
3: RightDownCausal模式。
4: Band模式。
----
cmpMaskMode(int64_t)输入q和cmpKv计算的mask模式。0: No Mask。
3: RightDownCausal模式。cmpKv未传入时该参数取默认值0。
----
oriWinLeft(int64_t)输入q和oriKv计算中,在因果边界基础上向左多看的token数。支持-1或非负数,其中-1表示窗口不受限。----
oriWinRight(int64_t)输入q和oriKv计算中,在因果边界基础上向右多看的token数。支持-1或非负数,其中-1表示窗口不受限。----
layoutQOptional(char*)输入标识输入q的数据排布格式。支持"BSND"和"TND"。----
layoutKvOptional(char*)输入标识输入oriKvOptional和cmpKvOptional的数据排布格式。支持"PA_BBND"、"BSND"和"TND"。----
topkValueMode(int64_t)输入topk索引取值模式。当前支持1。----
returnSoftmaxLse(bool)输入是否返回softmaxLse。支持true或false。----
attnOutOut(aclTensor*)输出注意力计算输出。-BFLOAT16、FLOAT16ND与q的shape一致×
softmaxLseOutOptional(aclTensor*)输出softmax的log-sum-exp结果。returnSoftmaxLse为false时返回占位Tensor;returnSoftmaxLse为true时返回softmax的log-sum-exp结果。FLOAT32ND +
    +
  • layoutQ为BSND时:(B, N2, S1, N1/N2)
  • +
  • layoutQ为TND时:(N2, T1, N1/N2)
  • +
  • returnSoftmaxLse为false时:占位Tensor
  • +
+
×
workspaceSize(uint64_t*)输出返回需要在Device侧申请的workspace大小。-----
executor(aclOpExecutor**)输出返回op执行器,包含了算子计算流程。-----
+ + + - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:N1/N2支持1、2、4、8、16、32、64、128;cmp_ratio在SWA场景传入0,CSA支持传入1、2或4,HCA支持传入128;block_size取值为16的倍数,最大支持1024;SWA稀疏ori_kv场景支持ori_sparse_indices及ori_topk_length,oriWinLeft和oriWinRight支持非负数,cmp_sparse_indices的最后一维K2当前支持512或1024。 + + + - Ascend 950PR/Ascend 950DT:N1支持1-128,N2只支持1。 + + + +- **返回值** + + aclnnStatus:返回状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。 + + 第一段接口完成入参校验,出现以下场景时报错: + + + - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
返回值错误码描述
ACLNN_ERR_PARAM_NULLPTR161001传入参数是必选输入、输出或者必选属性,且是空指针。
ACLNN_ERR_PARAM_INVALID161002输入变量的数据类型和数据格式不在支持的范围内。
N1不在[1,128]范围内,或N2不为1,或N1不能被N2整除,或N1/N2不是[1,128]范围内的2的幂。
D不为512。
非SWA稀疏ori_kv场景oriMaskMode不为4,SWA稀疏ori_kv场景oriMaskMode不为0,或cmpMaskMode不为3。
SWA场景cmpRatio不为0,或cmpRatio与CSA、HCA场景不匹配。
非SWA稀疏ori_kv场景oriWinLeft不为127,或oriWinRight不为0;SWA稀疏ori_kv场景oriWinLeft或oriWinRight为负数。
layoutQOptional、layoutKvOptional、topkValueMode、cmpSparseIndicesOptional、metadataOptional、sinksOptional、cuSeqlens或seqused相关参数规格不在支持范围内。
+ + + + - Ascend 950PR/Ascend 950DT: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
返回值错误码描述
ACLNN_ERR_PARAM_NULLPTR161001传入参数是必选输入、输出或者必选属性,且是空指针。
ACLNN_ERR_PARAM_INVALID161002输入变量的数据类型和数据格式不在支持的范围内。
N1不在[1,128]范围内,或N2不为1,或N1不能被N2整除,或N1/N2不在[1,128]范围内。
D不为512。
oriMaskMode不为0、3、4,或cmpMaskMode不为0、3。
hasCmpKv为true时,cmpRatio不在[1,128]范围内。
oriWinLeft或oriWinRight小于-1。
layoutQOptional、layoutKvOptional、topkValueMode、cmpSparseIndicesOptional、metadataOptional、sinksOptional、cuSeqlens或seqused相关参数规格不在支持范围内。
+ + + +## aclnnSparseFlashMla + +- **参数说明** + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出描述
workspace输入在Device侧申请的workspace内存地址。
workspaceSize输入在Device侧申请的workspace大小,由第一段接口aclnnSparseFlashMlaGetWorkspaceSize获取。
executor输入op执行器,包含了算子计算流程。
stream输入指定执行任务的Stream。
+ +- **返回值** + + 返回aclnnStatus状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。 + +## 约束说明 + +- 确定性计算 + + - aclnnSparseFlashMla默认采用确定性实现,相同输入多次调用结果一致。 + +- 使用约束 + + - SWA稀疏ori_kv场景仅支持SWA模板,仅传入`oriKvOptional`,同时传入`oriSparseIndicesOptional`和`oriTopkLengthOptional`;配套Metadata接口的`oriTopk`为oriSparseIndicesOptional最后一维K,`oriMaskMode`为0/3/4,`oriWinLeft`和`oriWinRight`仅在`oriMaskMode`为4时取非负数,且`cmpKvOptional`不传入。 + - SWA稀疏ori_kv场景下,`oriTopkLengthOptional`的元素表示实际有效索引条目数,取值应在[0, K]范围内。对每个q token和KV head,`oriSparseIndicesOptional`的[0, oriTopkLengthOptional)区间为左对齐的有效索引条目,[oriTopkLengthOptional, K)区间为无效或填充条目,建议填-1;其他场景`oriTopkLengthOptional`传入nullptr或空Tensor。 + - 除`cmpTopkLengthOptional`等预留输入可传入nullptr或空Tensor外,其余已传入Tensor不支持为空。 + - `metadataOptional`参数必须传入,由`aclnnSparseFlashMlaMetadata`算子生成,shape固定为(1024,)。 + - `cmpResidualKvOptional`为主算子和`aclnnSparseFlashMlaMetadata`的可选入参;传入后用于按`cmp_len * cmpRatio + residual`恢复cmp侧mask使用的压缩前长度。 + +- 三种Attention场景输入要求 + + | 场景 | oriKvOptional | cmpKvOptional | cmpSparseIndicesOptional | 说明 | + | :--- | :----- | :----- | :----------------- | :--- | + | SWA | 必须传入 | 不传入 | 不传入 | 仅滑动窗口注意力 | + | CSA | 必须传入 | 必须传入 | 必须传入 | 滑动窗口 + TopK稀疏压缩KV | + | HCA | 必须传入 | 必须传入 | 不传入 | 滑动窗口 + 稠密压缩KV | + +- Layout约束 + + - `layoutQOptional`和`layoutKvOptional`组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下`layoutQOptional`和`layoutKvOptional`必须一致。 + - 当`layoutQOptional`为TND时,`cuSeqlensQOptional`必须传入。 + - 当`layoutKvOptional`为PA_BBND时,`sequsedOriKvOptional`必须传入,`oriBlockTableOptional`必须传入。BSND场景可选传入`sequsedOriKvOptional`覆盖每个batch的oriKv有效长度;TND场景使用`cuSeqlensOriKvOptional`表达oriKv序列边界。 + - 当`layoutKvOptional`为TND时,`cuSeqlensOriKvOptional`必须传入。 + - 当`layoutKvOptional`为TND且存在`cmpKvOptional`时,`cuSeqlensCmpKvOptional`必须传入。 + - `sequsedCmpKvOptional`为所有layoutKvOptional下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度。 + +## 调用示例 + +调用示例代码如下,仅供参考,具体编译和执行过程请参考[编译与运行样例](../../../docs/zh/context/compile_and_run_sample.md)。 + +```c++ +/*! + * \file test_aclnn_sparse_flash_mla.cpp + * \brief SparseFlashMla + SparseFlashMlaMetadata 算子调用示例(CSA) + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include "acl/acl.h" +#include "aclnnop/aclnn_sparse_flash_mla.h" +#include "aclnnop/aclnn_sparse_flash_mla_metadata.h" + +#define CHECK_RET(cond, return_expr) \ + do { \ + if (!(cond)) { \ + return_expr; \ + } \ + } while (0) + +#define LOG_PRINT(message, ...) \ + do { \ + printf(message, ##__VA_ARGS__); \ + } while (0) + +namespace { + +int64_t GetShapeSize(const std::vector& shape) +{ + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} + +uint16_t FloatToFp16(float f) +{ + uint32_t bits; + std::memcpy(&bits, &f, sizeof(bits)); + uint32_t sign = (bits >> 31) & 0x1u; + int32_t exp = static_cast((bits >> 23) & 0xffu) - 127 + 15; + uint32_t mant = (bits >> 13) & 0x3ffu; + if (exp <= 0) { + return static_cast(sign << 15); + } + if (exp >= 31) { + return static_cast((sign << 15) | 0x7c00u); + } + return static_cast((sign << 15) | (static_cast(exp) << 10) | mant); +} + +float Fp16ToFloat(uint16_t h) +{ + uint32_t sign = (h >> 15) & 0x1u; + uint32_t exp = (h >> 10) & 0x1fu; + uint32_t mant = h & 0x3ffu; + uint32_t f; + if (exp == 0) { + f = (sign << 31) | (mant << 13); + } else if (exp == 31) { + f = (sign << 31) | 0x7f800000u | (mant << 13); + } else { + f = (sign << 31) | ((exp + 127u - 15u) << 23) | (mant << 13); + } + float result; + std::memcpy(&result, &f, sizeof(result)); + return result; +} + +void PrintOutResult(const std::vector& shape, void** deviceAddr) +{ + auto size = GetShapeSize(shape); + std::vector resultData(size, 0); + auto ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), + *deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); return); + for (int64_t i = 0; i < size && i < 10; i++) { + LOG_PRINT("result[%ld] is: %f\n", i, Fp16ToFloat(resultData[i])); + } +} + +int Init(int32_t deviceId, aclrtContext* context, aclrtStream* stream) +{ + auto ret = aclInit(nullptr); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret); + ret = aclrtSetDevice(deviceId); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret); + ret = aclrtCreateContext(context, deviceId); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateContext failed. ERROR: %d\n", ret); return ret); + ret = aclrtSetCurrentContext(*context); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetCurrentContext failed. ERROR: %d\n", ret); return ret); + ret = aclrtCreateStream(stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret); + return 0; +} + +template +int CreateAclTensor(const std::vector& hostData, const std::vector& shape, void** deviceAddr, + aclDataType dataType, aclTensor** tensor) +{ + auto size = GetShapeSize(shape) * sizeof(T); + if (size > 0) { + auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret); + ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret); + } else { + *deviceAddr = nullptr; + } + + std::vector strides(shape.size(), 1); + for (int64_t i = static_cast(shape.size()) - 2; i >= 0; i--) { + strides[i] = shape[i + 1] * strides[i + 1]; + } + + *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), *deviceAddr); + return 0; +} + +std::vector MakeFp16Data(int64_t size, float value) +{ + std::vector data(static_cast(size), FloatToFp16(value)); + return data; +} + +} // namespace + +int main() +{ + // 1. (固定写法)device/stream初始化,参考acl API手册 + // 根据自己的实际device填写deviceId + int32_t deviceId = 0; + aclrtContext context = nullptr; + aclrtStream stream = nullptr; + auto ret = Init(deviceId, &context, &stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); + + int64_t B = 4; + int64_t S1 = 128; + int64_t S2 = 8192; + int64_t N1 = 64; + int64_t N2 = 1; + int64_t D = 512; + int64_t K = 512; + int64_t oriBlockSize = 128; + int64_t cmpBlockSize = 128; + int64_t s2Act = 4096; + int64_t cmpRatio = 4; + int64_t oriWinLeft = 127; + int64_t oriWinRight = 0; + int64_t oriMaskMode = 4; + int64_t cmpMaskMode = 3; + double softmaxScale = 1.0 / sqrt(static_cast(D)); + + int64_t T1 = B * S1; + int64_t cmpKvLen = s2Act / cmpRatio; + int64_t oriBlockNum = ((s2Act + oriBlockSize - 1) / oriBlockSize) * B; + int64_t cmpBlockNum = ((cmpKvLen + cmpBlockSize - 1) / cmpBlockSize) * B; + + // 2. 构造输入与输出,需要根据API的接口自定义构造 + std::vector qShape = {T1, N1, D}; + std::vector oriKvShape = {oriBlockNum, oriBlockSize, N2, D}; + std::vector cmpKvShape = {cmpBlockNum, cmpBlockSize, N2, D}; + std::vector cmpSparseIndicesShape = {T1, N2, K}; + std::vector oriBlockTableShape = {B, (s2Act + oriBlockSize - 1) / oriBlockSize}; + std::vector cmpBlockTableShape = {B, (cmpKvLen + cmpBlockSize - 1) / cmpBlockSize}; + std::vector cuSeqLensQShape = {B + 1}; + std::vector seqUsedOriKvShape = {B}; + std::vector seqUsedCmpKvShape = {B}; + std::vector cmpResidualKvShape = {B}; + std::vector sinksShape = {N1}; + std::vector metadataShape = {1024}; + std::vector attnOutShape = {T1, N1, D}; + std::vector softmaxLseShape = {T1, N1, 1}; + // 对全部 5 个输入调用 Contiguous,optional 输入传 shape 为 {0} 的空 tensor。 + std::vector emptyShape = {0}; + + void* qDeviceAddr = nullptr; + void* oriKvDeviceAddr = nullptr; + void* cmpKvDeviceAddr = nullptr; + void* cmpSparseIndicesDeviceAddr = nullptr; + void* oriBlockTableDeviceAddr = nullptr; + void* cmpBlockTableDeviceAddr = nullptr; + void* cuSeqLensQDeviceAddr = nullptr; + void* cuSeqLensOriKvDeviceAddr = nullptr; + void* cuSeqLensCmpKvDeviceAddr = nullptr; + void* seqUsedQDeviceAddr = nullptr; + void* seqUsedOriKvDeviceAddr = nullptr; + void* seqUsedCmpKvDeviceAddr = nullptr; + void* cmpResidualKvDeviceAddr = nullptr; + void* sinksDeviceAddr = nullptr; + void* metadataDeviceAddr = nullptr; + void* attnOutDeviceAddr = nullptr; + void* softmaxLseDeviceAddr = nullptr; + + aclTensor* q = nullptr; + aclTensor* oriKv = nullptr; + aclTensor* cmpKv = nullptr; + aclTensor* cmpSparseIndices = nullptr; + aclTensor* oriBlockTable = nullptr; + aclTensor* cmpBlockTable = nullptr; + aclTensor* cuSeqLensQ = nullptr; + aclTensor* cuSeqLensOriKv = nullptr; + aclTensor* cuSeqLensCmpKv = nullptr; + aclTensor* seqUsedQ = nullptr; + aclTensor* seqUsedOriKv = nullptr; + aclTensor* seqUsedCmpKv = nullptr; + aclTensor* cmpResidualKv = nullptr; + aclTensor* sinks = nullptr; + aclTensor* metadata = nullptr; + aclTensor* attnOut = nullptr; + aclTensor* softmaxLse = nullptr; + + int64_t qSize = GetShapeSize(qShape); + int64_t oriKvSize = GetShapeSize(oriKvShape); + int64_t cmpKvSize = GetShapeSize(cmpKvShape); + int64_t cmpSparseIndicesSize = GetShapeSize(cmpSparseIndicesShape); + int64_t oriBlockTableSize = GetShapeSize(oriBlockTableShape); + int64_t cmpBlockTableSize = GetShapeSize(cmpBlockTableShape); + int64_t attnOutSize = GetShapeSize(attnOutShape); + int64_t softmaxLseSize = GetShapeSize(softmaxLseShape); + + std::vector qHostData = MakeFp16Data(qSize, 1.0f); + std::vector oriKvHostData = MakeFp16Data(oriKvSize, 1.0f); + std::vector cmpKvHostData = MakeFp16Data(cmpKvSize, 1.0f); + std::vector cmpSparseIndicesHostData(cmpSparseIndicesSize); + std::vector oriBlockTableHostData(oriBlockTableSize); + std::iota(oriBlockTableHostData.begin(), oriBlockTableHostData.end(), 0); + std::vector cmpBlockTableHostData(cmpBlockTableSize); + std::iota(cmpBlockTableHostData.begin(), cmpBlockTableHostData.end(), 0); + std::vector cuSeqLensQHostData(B + 1); + for (int64_t i = 0; i <= B; i++) { + cuSeqLensQHostData[i] = static_cast(i * S1); + } + std::vector emptyHostData; + std::vector seqUsedOriKvHostData(B, static_cast(s2Act)); + std::vector seqUsedCmpKvHostData(B, static_cast(cmpKvLen)); + std::vector cmpResidualKvHostData(B, static_cast(s2Act % cmpRatio)); + std::vector sinksHostData(N1, 1.0f); + std::vector metadataHostData(1024, 0); + std::vector attnOutHostData = MakeFp16Data(attnOutSize, 0.0f); + std::vector softmaxLseHostData(softmaxLseSize, 0.0f); + + std::mt19937 gen(42); + for (int64_t t = 0; t < T1; t++) { + for (int64_t n = 0; n < N2; n++) { + for (int64_t k = 0; k < K; k++) { + cmpSparseIndicesHostData[t * N2 * K + n * K + k] = static_cast(gen() % cmpKvLen); + } + } + } + + ret = CreateAclTensor(qHostData, qShape, &qDeviceAddr, aclDataType::ACL_FLOAT16, &q); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(oriKvHostData, oriKvShape, &oriKvDeviceAddr, aclDataType::ACL_FLOAT16, &oriKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpKvHostData, cmpKvShape, &cmpKvDeviceAddr, aclDataType::ACL_FLOAT16, &cmpKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpSparseIndicesHostData, cmpSparseIndicesShape, &cmpSparseIndicesDeviceAddr, + aclDataType::ACL_INT32, &cmpSparseIndices); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(oriBlockTableHostData, oriBlockTableShape, &oriBlockTableDeviceAddr, aclDataType::ACL_INT32, + &oriBlockTable); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpBlockTableHostData, cmpBlockTableShape, &cmpBlockTableDeviceAddr, aclDataType::ACL_INT32, + &cmpBlockTable); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cuSeqLensQHostData, cuSeqLensQShape, &cuSeqLensQDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensQ); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(emptyHostData, emptyShape, &cuSeqLensOriKvDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensOriKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(emptyHostData, emptyShape, &cuSeqLensCmpKvDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensCmpKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(emptyHostData, emptyShape, &seqUsedQDeviceAddr, aclDataType::ACL_INT32, &seqUsedQ); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(seqUsedOriKvHostData, seqUsedOriKvShape, &seqUsedOriKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedOriKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(seqUsedCmpKvHostData, seqUsedCmpKvShape, &seqUsedCmpKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedCmpKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpResidualKvHostData, cmpResidualKvShape, &cmpResidualKvDeviceAddr, aclDataType::ACL_INT32, &cmpResidualKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(sinksHostData, sinksShape, &sinksDeviceAddr, aclDataType::ACL_FLOAT, &sinks); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(metadataHostData, metadataShape, &metadataDeviceAddr, aclDataType::ACL_INT32, &metadata); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(attnOutHostData, attnOutShape, &attnOutDeviceAddr, aclDataType::ACL_FLOAT16, &attnOut); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(softmaxLseHostData, softmaxLseShape, &softmaxLseDeviceAddr, aclDataType::ACL_FLOAT, &softmaxLse); + CHECK_RET(ret == ACL_SUCCESS, return ret); + + char layoutQ[] = "TND"; + char layoutKv[] = "PA_BBND"; + + uint64_t metadataWorkspaceSize = 0; + aclOpExecutor* metadataExecutor = nullptr; + + // 3. 调用CANN算子库API,需要修改为具体的Api名称 + ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( + cuSeqLensQ, cuSeqLensOriKv, cuSeqLensCmpKv, + seqUsedQ, seqUsedOriKv, + seqUsedCmpKv, cmpResidualKv, nullptr, nullptr, + N1, N2, D, B, S1, S2, cmpKvLen, + 0, K, cmpRatio, + oriMaskMode, cmpMaskMode, + oriWinLeft, oriWinRight, + layoutQ, layoutKv, + true, true, + metadata, + &metadataWorkspaceSize, &metadataExecutor); + CHECK_RET(ret == ACL_SUCCESS, + LOG_PRINT("aclnnSparseFlashMlaMetadataGetWorkspaceSize failed. ERROR: %d\n", ret); return ret); + + void* metadataWorkspaceAddr = nullptr; + if (metadataWorkspaceSize > 0) { + ret = aclrtMalloc(&metadataWorkspaceAddr, metadataWorkspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate metadata workspace failed. ERROR: %d\n", ret); return ret); + } + + ret = aclnnSparseFlashMlaMetadata(metadataWorkspaceAddr, metadataWorkspaceSize, metadataExecutor, stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseFlashMlaMetadata failed. ERROR: %d\n", ret); return ret); + + // 4. (固定写法)同步等待任务执行结束 + ret = aclrtSynchronizeStream(stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream after metadata failed. ERROR: %d\n", ret); return ret); + + uint64_t workspaceSize = 0; + aclOpExecutor* executor = nullptr; + + ret = aclnnSparseFlashMlaGetWorkspaceSize( + q, oriKv, cmpKv, + nullptr, cmpSparseIndices, + oriBlockTable, cmpBlockTable, + cuSeqLensQ, nullptr, nullptr, + nullptr, seqUsedOriKv, + seqUsedCmpKv, cmpResidualKv, nullptr, nullptr, + sinks, metadata, + softmaxScale, cmpRatio, + oriMaskMode, cmpMaskMode, + oriWinLeft, oriWinRight, + layoutQ, layoutKv, + 1, + false, + attnOut, softmaxLse, + &workspaceSize, &executor); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseFlashMlaGetWorkspaceSize failed. ERROR: %d\n", ret); return ret); + + void* workspaceAddr = nullptr; + if (workspaceSize > 0) { + ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret); + } + + ret = aclnnSparseFlashMla(workspaceAddr, workspaceSize, executor, stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseFlashMla failed. ERROR: %d\n", ret); return ret); + + ret = aclrtSynchronizeStream(stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret); + + // 5.获取输出的值,将device侧内存上的结果拷贝至host侧,需要根据具体API的接口定义修改 + PrintOutResult(attnOutShape, &attnOutDeviceAddr); + + // 6. 释放aclTensor和aclScalar,需要根据具体API的接口定义修改 + aclDestroyTensor(q); + aclDestroyTensor(oriKv); + aclDestroyTensor(cmpKv); + aclDestroyTensor(cmpSparseIndices); + aclDestroyTensor(oriBlockTable); + aclDestroyTensor(cmpBlockTable); + aclDestroyTensor(cuSeqLensQ); + aclDestroyTensor(cuSeqLensOriKv); + aclDestroyTensor(cuSeqLensCmpKv); + aclDestroyTensor(seqUsedQ); + aclDestroyTensor(seqUsedOriKv); + aclDestroyTensor(seqUsedCmpKv); + aclDestroyTensor(cmpResidualKv); + aclDestroyTensor(sinks); + aclDestroyTensor(metadata); + aclDestroyTensor(attnOut); + aclDestroyTensor(softmaxLse); + + // 7. 释放device资源 + aclrtFree(qDeviceAddr); + aclrtFree(oriKvDeviceAddr); + aclrtFree(cmpKvDeviceAddr); + aclrtFree(cmpSparseIndicesDeviceAddr); + aclrtFree(oriBlockTableDeviceAddr); + aclrtFree(cmpBlockTableDeviceAddr); + if (cuSeqLensQDeviceAddr != nullptr) { + aclrtFree(cuSeqLensQDeviceAddr); + } + if (seqUsedOriKvDeviceAddr != nullptr) { + aclrtFree(seqUsedOriKvDeviceAddr); + } + if (seqUsedCmpKvDeviceAddr != nullptr) { + aclrtFree(seqUsedCmpKvDeviceAddr); + } + if (cmpResidualKvDeviceAddr != nullptr) { + aclrtFree(cmpResidualKvDeviceAddr); + } + aclrtFree(sinksDeviceAddr); + aclrtFree(metadataDeviceAddr); + aclrtFree(attnOutDeviceAddr); + aclrtFree(softmaxLseDeviceAddr); + if (metadataWorkspaceSize > 0) { + aclrtFree(metadataWorkspaceAddr); + } + if (workspaceSize > 0) { + aclrtFree(workspaceAddr); + } + aclrtDestroyStream(stream); + aclrtDestroyContext(context); + aclrtResetDevice(deviceId); + aclFinalize(); + + return 0; +} +``` diff --git a/csrc/attention/sparse_flash_mla/docs/design.md b/csrc/attention/sparse_flash_mla/docs/design.md new file mode 100644 index 000000000000..17cdc8e67dc8 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/docs/design.md @@ -0,0 +1,400 @@ +# SparseFlashMla 算子设计与开发指南 + +本文面向第一次接触该算子的开发者,按“计算什么 → 怎样分工 → 数据放在哪里 → 怎样执行 → 怎样验证”的顺序介绍当前仓库实现。范围为 `attention/sparse_flash_mla`,并追踪配套 `sparse_flash_mla_metadata` 的任务切分。文中数值来自当前源码,不是适用于所有 MLA 算子的通用参数。 + +建议先读第 1~3 节建立概念,再结合第 4~6 节读 Kernel。接口完整约束见 [README](../README.md) 和 [aclnn 接口说明](aclnnSparseFlashMla.md)。本文区分公开接口约束与内部模板分支;存在模板不代表任意输入组合都能通过校验。 + +A2/A3 倍率 1/2 的具体改动与验证范围见 [cmp_ratio=1/2 适配说明](ratio2_a2a3.md)。当前 A2/A3 的 SWA 接受 0,CSA 接受 1/2/4,HCA 接受 128;cmp causal mask 下倍率不为 1 时必须提供 residual(包括余数为 0 的情况)。 + +## 1. 算子计算什么 + +### 1.1 MLA、稀疏选择和融合计算 + +本实现的 Q 和 KV 的最后一维都是 `D=512`,KV 由 448 维 nope 与 64 维 rope 拼接组成;KV head 数 `N2=1`。每个 query head 对选出的 KV 计算注意力,**同一份 KV 同时作为 K 和 V,输出也是 512 维**。不要把其他 MLA 实现中“仅部分维度作为 V”的规则套到这里。 + +设 query head 数为 `N1`,分组比 `G=N1/N2`。给定 batch `b`、query token `i`、query head `h`,令 `J(b,i)` 为实际参与计算的 KV 条目序列,条目可以来自原始缓存 `ori_kv` 或压缩缓存 `cmp_kv`。对有效条目: + +\[ +x_j=\text{softmax\_scale}\sum_{d=0}^{511}Q_{b,i,h,d}KV_{j,d}. +\] + +本实现还包含每个 head 的 `sinks[h]`。记它为 \(s_h\),完整计算为: + +\[ +Z=e^{s_h}+\sum_{j\in J(b,i)}e^{x_j},\qquad +O_d=\frac{\sum_{j\in J(b,i)}e^{x_j}KV_{j,d}}{Z},\qquad +LSE=\log Z. +\] + +`sinks` 相当于一个 value 为零的额外 softmax 项:影响分母,不增加输出分子。它不是 KV cache 中的一个 token,也不是窗口左端保留的 token。当前 Host 的 `GetSinks` 会检查 sinks 是否存在,不能因为接口表将其标为可选,就直接省略。 + +算子不负责生成压缩 KV,也不负责从全量 KV 计算 TopK 排名;调用方提供压缩结果及稀疏索引。融合的核心是按小块完成 `QKᵀ → mask/softmax → PV → 累加`,避免存储完整注意力矩阵。 + +公式与精度参考:[Golden](../tests/pytest/sparse_flash_mla_golden.py) 的 `calculate_by_bnsd`、`sinks_softmax`;输出 shape 参考 [InferShape](../op_host/sparse_flash_mla_infershape.cpp)。 + +### 1.2 场景与模板路由 + +| 输入形态 | 算法含义 | arch22 路由 | arch35 路由 | +| --- | --- | --- | --- | +| ori KV,无稀疏索引 | SWA,按窗口或 mask 访问 ori | SWA Kernel | SWA Kernel | +| ori KV + ori 索引 + ori 有效 TopK 长度 | 稀疏 ori | SWA 模板 + `hasOriSparseIndices` | `ORI_SPARSE`,CSA Kernel | +| ori KV + cmp KV + cmp 索引 | CSA,ori 与稀疏 cmp 共同归一化 | CSA Kernel | CSA Kernel | +| ori KV + cmp KV,无 cmp 索引 | HCA,ori 与连续压缩 KV 共同归一化 | SWA Kernel | SWA Kernel | +| ori/cmp 均带索引 | 内部 `ORI_CMP_SPARSE` 分支 | 不应据此推断支持 | CSA Kernel,仍须通过 checker | + +“CSA Kernel”是复用的执行框架名,不意味着只处理传统 CSA。HCA 也没有独立的 `hca_kernel.h`。真实选择过程见 [Host](../op_host/sparse_flash_mla_tiling.cpp) 的 `GetSMLATemplateMode` 和 [Kernel 入口](../op_kernel/sparse_flash_mla.cpp)。 + +平台对应:A2/A3 使用 `DAV_2201` 的 Host 逻辑和 `arch22`;950 路径使用 `DAV_3510` 与 `arch35`,Kernel 编译入口以 `__CCE_AICORE__ == 310` 区分。文件夹名、Host 架构枚举和编译宏不是同一个编号体系。 + +## 2. 输入、布局与寻址 + +### 2.1 维度约定 + +| 符号 | 含义 | +| --- | --- | +| B | batch 数 | +| S1 | 单 batch 的 Q 序列长度;Host 中 TND 的 `s1Size` 可表示总 Q token 数 | +| S2 / cmpS2 | ori / cmp 序列长度,计算时再按有效长度、mask、索引收缩 | +| T | TND 的所有 batch token 总数 | +| N1 / N2 / G | Q head 数 / KV head 数 / 每个 KV head 对应的 Q head 数 | +| M | 矩阵乘的行,来自 query token 与 G 的合轴,不一定等于 query token 数 | +| Ktop | 稀疏索引张量末维容量,与矩阵乘的归约维 D 无关 | + +Q 支持 `[B,S1,N1,D]`(BSND)、`[T,N1,D]`(TND);KV 支持 BSND、TND、`[block_num,block_size,N2,D]`(PA_BBND)。非分页时 Q/KV 布局必须匹配;分页时 Q 可为 BSND 或 TND。 + +`attn_out` shape 和 dtype 与 Q 相同。开启 LSE 时,BSND 输出 LSE 为 `[B,N2,S1,G]`,TND 为 `[N2,T,G]`,dtype 为 FP32;关闭时为 shape `[0]` 的占位输出。LSE 不能简单按 Q 去掉最后一维来解释。 + +### 2.2 分清存储长度、有效长度和坐标 + +`cu_seqlens_*` 是前缀和,TND 中 batch `b` 的存储起点为 `cu[b]`,存储长度为 `cu[b+1]-cu[b]`。`seqused_*` 描述实际参与运算的长度,存在时用于收缩有效区间,不应拿它代替 TND 的存储起点。BSND 则仍按固定 shape/stride 找 batch 起点。 + +压缩侧 mask 需要原始时间轴。显式 `cmp_residual_kv` 用于恢复: + +\[ +L_{\text{cmp,original}}=L_{\text{cmp,valid}}\times\text{cmp\_ratio}+\text{residual}. +\] + +未显式提供长度或 residual 时,各布局有自己的推导分支,应追踪 `ComputeParamBatch` 和 Golden 的长度解析,不能总用 ori 长度替代 cmp 的时间轴。 + +### 2.3 稀疏索引与分页是两层映射 + +稀疏索引决定“取哪个逻辑 token”;block table 决定“逻辑 token 存在哪个物理页”。例如 PA block size 为 16,某索引为 35,则逻辑页号为 2、页内偏移为 3。若 `block_table[b,2]=7`,实际读取物理页 7 的第 3 个 token。 + +对于连续存储、N2=1 的 PA KV,元素偏移可以写成: + +```text +page = logical_token / block_size +offset_in_page = logical_token % block_size +physical_page = block_table[b, page] +element_offset = (physical_page * block_size + offset_in_page) * D + d +``` + +上式是连续存储示例。源码还传递 `oriKvStride0`、`cmpKvStride0`、`oriKeyStride0` 等 stride,扩展非连续输入时必须按实际 stride 计算。ori/cmp 的 block size、block table 和长度各自独立。 + +稀疏 ori 的索引 shape 为 `[B,S1,N2,Ktop]` 或 `[T,N2,Ktop]`,`ori_topk_length` 给每个 query/KV head 的有效条目数,范围 `[0,Ktop]`;有效条目左对齐,尾部建议填 `-1`。索引有效性检查、长度裁剪和 mask 都不能因 gather 已经完成而省略。`cmp_topk_length` 在公开接口中仍是预留输入,不可仅根据 Kernel 的参数名认定可以传入。 + +### 2.4 mask 如何影响选中条目 + +原始侧右下对齐的 query 位置为 `p=Lori-Lq+i`。mode 3 保留 `j<=p`;mode 4 的有限窗口保留 `p-win_left<=j<=p+win_right`,再与 `[0,Lori)` 相交;mode 0 不施加窗口 mask。`-1` 无界窗口仅适用于支持它的平台/模式。 + +压缩侧 mode 3 按压缩比例与恢复的原始长度判断可见范围,不能直接拿压缩 token 编号和 Q 编号比较。边界、索引裁剪的完整实现应对照 [KV 参数工具](../op_kernel/arch35/sparse_flash_mla_kvcache.h) 和 Golden。A2/A3 与 950 的 mask、head 数、block size、压缩比支持范围不同,以 README 和相应 checker 为准。 + +## 3. Tiling:三层分工 + +### 3.1 Host tiling 决定静态配置 + +`TilingForSparseFlashMla` 的顺序是: + +1. `SMLAInfoParser::Parse` 解析平台、shape、dtype、layout、stride、可选输入及模式。 +2. arch22 走 `SMLATilingCheck`;其他分支走独立的 `SparseFlashMlaChecker`。 +3. `DoOpTiling` 设置 blockDim,计算基本块和 workspace,写入 tiling data 和 tiling key。 + +`blockDim` 通过平台的 `CalcTschBlockDim` 计算,不是写死 36。Host 的 `usedCoreNum` 记录可用 AIC 数;实际哪些核有任务由 metadata 的 enable 字段决定。 + +主要 tiling 字段: + +| 字段组 | 作用 | +| --- | --- | +| `batchSize/qSeqSize/kvSeqSize/nNumOfQInOneGroup` | batch、序列和 head 分组参数 | +| `mBaseSize/s2BaseSize/mmResUbSize/bmm2ResUbSize` | Host 基本块和中间结果容量,后两项单位是元素数 | +| layout、stride、各侧 block size/table 宽度 | 将逻辑坐标转换为地址 | +| mask、window、scale、returnSoftmaxLse | 数值语义及输出控制 | +| ori/cmp sparse count、index width、有效长度维度 | 稀疏和变长输入解析 | + +### 3.2 基本块与架构差异 + +令 `align(x,a)=ceil(x/a)*a`。Host `SplitBalanced` 的计算如下: + +| 分支 | Host M 基本块 | Host S2 基本块 | +| --- | --- | --- | +| arch22 CSA | G | 512 | +| arch22 连续 SWA/HCA | `floor(256/G)*G` | 512 | +| arch22 稀疏 ori | G,即一个 query token 的 heads | 512 | +| arch35 | 默认 64 | 默认 512 | + +Host 还计算: + +```text +Mcap = min(G * host_s1Size, mBaseSize) +R1 = align(s2BaseSize,32) * align(Mcap,16) # mmResUbSize,元素数 +R2 = align(D,32) * align(Mcap,16) # bmm2ResUbSize,元素数 +``` + +**arch35 Kernel 在 Init 中明确设置 `s1BaseSize=64`、`s2BaseSize=128`。** 其本地矩阵缓冲区按这个实际块配置建立,不能用 Host 的 512 来计算 UB/L1 大小。配套 Metadata 在 arch35 以 `mBaseSize=G`、S2=128 组织逻辑任务,表示每个 query token 的一组 heads;64 是单 AIC 的矩阵行容量。 + +arch35 在 `G>64` 时开启 `SPLIT_G`:一对 AIC 处理相同 query 的两部分 heads,第一个处理 `ceil(G/2)`,第二个处理剩余 heads,共享 KV gather 缓存。因而逻辑任务槽数由 AIC 数 C 降为 `C/2`。这是切 head;沿 S2 切分再归约是另一件事。 + +### 3.3 Metadata 决定实际任务区间 + +调用链:`SparseFlashMlaMetadata → metadata Tensor → SparseFlashMla`。前置算子使用实际长度、窗口、稀疏容量等构建任务,主算子不能只凭 Q/KV shape 还原这些任务。 + +[Metadata 实现](../../sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.cpp) 的 `BalanceSchedule` 依次执行:`CalcSplitInfo → CalcCostInfo → CalcSplitPlan → SplitFD → GenMetadata`。分配内部包含按 batch、按行、按 S2 block 分配的路径;使用 ori/cmp 成本及尾块信息,而非简单平均分 batch。必要时一行的 S2 工作跨多个核,产生 FD(Flash Decode)归约任务。 + +metadata 为固定 `[1024]` 的 INT32 Tensor,即 4096 字节。当前协议容纳 36 组 FA 记录和 72 组 FD 记录,实际占用 `36*9+72*8=900` 个元素,剩余空间不应自行定义用途。 + +| 区域 | 单条元素数 | 内容 | +| --- | --- | --- | +| FA,按 AIC | 9 | enable;起止 BN2/M/S2 游标;首个 FD workspace 索引;S2 最大轮次 | +| FD,按 AIV | 8 | enable;BN2/M;workspace 起点与份数;归约 M 起点与行数;预留位置 | + +FA 元素地址为 `core*9+field`,FD 为 `36*9+core*8+field`,见 [协议头](../op_kernel/sparse_flash_mla_kernel_metadata.h)。M/S2 是协议游标,S2 经 `ConvertS2MetadataBlockToToken`、`ApplyS2MetadataRange` 转换和裁剪,不能把原始字段直接当 GM token 偏移。 + +### 3.4 Tiling key 和一致性模式 + +key 模板维度为 `FLASH_DECODE, Q_LAYOUT, KV_LAYOUT, TEMPLATE_MODE, SPLIT_G, HEAD_RATIO_ONE, BATCH_CONSISTENCY, IS_VEC_S2PHYADDR`。当前 Host 第一个参数传 0;这不代表执行过程中不存在 FD,arch35 的 FD 任务由 metadata 驱动。`HEAD_RATIO_ONE` 为 arch22 CSA 且 G=1 的特化。 + +batch consistency 来自执行上下文的 deterministic level。arch35 按每行实际负载形成稳定 reduction block:以 `floor(totalLoad/32)` 向上对齐到 S2 基本块,至少一个基本块;先在核内归约,再按协议做跨核归约。其目的在于稳定 batch 变化时的浮点累加顺序,不能理解成只设置一个随机种子。它需要专门的 workspace,且不能据此承诺跨架构、跨 dtype 逐位一致。 + +## 4. 内存分配与生命周期 + +### 4.1 内存层级 + +| 位置 | 保存内容 | 主要使用者 | +| --- | --- | --- | +| GM | 输入/输出、metadata、跨阶段或跨核 workspace | 所有核 | +| L1 | Q、K/V、softmax 权重 P 的矩阵输入块 | Cube,部分数据由 Vector 写入 | +| L0A/L0B | 当前矩阵乘操作数 | Cube | +| L0C | FP32 矩阵乘累加结果 | Cube/Fixpipe | +| UB | gather、mask、softmax 状态、输出累加 | Vector | + +FP16/BF16 输入和 P 占 2 字节;矩阵结果、max/sum 和累加输出通常采用 FP32,占 4 字节。后面的 KiB 均为 1024 字节。 + +workspace 总量由 Host 申请;Kernel 入口用 `GetUserWorkspace` 去掉框架工作区前缀,再按内部偏移切片。不能将 Host 的库工作区大小再重复加到 user 指针上。 + +### 4.2 arch22 的 GM workspace + +令 C 为 AIC 数,R1/R2 为第 3.2 节的元素数。`DoOpTiling` 为每核两套流水槽分配: + +| 区域 | 字节数 | 数据流 | +| --- | --- | --- | +| MM1 结果 | `2*R1*4*C` | Cube → Vector,QK 分数 | +| Vec1 结果 | `2*R1*2*C` | Vector → Cube,softmax 权重 | +| MM2 结果 | `2*R2*4*C` | Cube → Vector,PV | +| Vec2 结果 | `2*R2*4*C` | 分块累加结果暂存 | +| KV merge,仅 CSA/稀疏 ori | `3*512*512*2*C` | 三槽 gather 缓存 | + +所以此路径用户工作区为 `C*(12*R1+16*R2)`,再按需加 KV merge;框架库工作区另加。源码里的容量字段虽然包含 `Ub`,这里是用于计算 GM 中间结果预留量,不能当成单个 UB 分配。 + +例如 arch22 CSA、G=32、D=512,R1=R2=16384,四类结果每核共 448 KiB,KV merge 每核 1536 KiB,合计 1984 KiB;再乘实际 C。此例只计算该分支的 workspace,不包含 Q/KV/输出和库工作区。 + +### 4.3 arch22 的片上缓冲区(CSA 路径) + +[Cube `InitBuffers`](../op_kernel/arch22/sparse_flash_mla_csa_block_cube.h) 中 `L1_BLOCK_SIZE=64*512*2=64 KiB`:Q/P L1 分配 4 块即 256 KiB,KV L1 分配 3 块即 192 KiB;L0A/L0B 各双缓冲共 64 KiB,L0C 双缓冲共 128 KiB。 + +[Vector `InitBuffers`](../op_kernel/arch22/sparse_flash_mla_csa_block_vector.h) 包含:inputBuff1 64 KiB、inputBuff2 32 KiB、outputBuff1 32 KiB、tmpBuff1 32 KiB、有效长度缓冲 8 KiB;另有 max/exp/sum 流水状态、默认状态、sinks 与广播缓冲,LSE 开启时加输出缓冲。部分逻辑 Tensor 是对已有缓冲的视图,不应对每个 Tensor 名重复计费。SWA 有独立的 Cube/Vector 类,应分别检查,不能将 CSA 预算无条件套用。 + +### 4.4 arch35 的 GM workspace + +定义逻辑槽数 `L=C`,split-G 时 `L=C/2`;AIV 数为 V。Host 申请量分三部分: + +1. 当 split-G 或 CSA/ORI_SPARSE/ORI_CMP_SPARSE 时,KV gather 三槽为 `3*128*512*2*L` 字节;有效索引辅助区预留 `3*128*4*V` 字节。 +2. 若启用物理地址向量化,增加选中侧的地址表:`totalQ*align(Ktop,128)*8` 字节/侧,`totalQ` 为 BSND 的 `B*S1` 或 TND 的 T。 +3. FD staging:每槽 `F=G*4*(align(D,32)+2*8)` 字节,保存部分输出和广播布局的 max/sum。普通模式预留 `2*L*F`;batch consistency 模式预留 `(2*L+C*33)*F`。 + +物理地址向量化并非所有 PA case 都开启:Host 估算索引、INT64 地址和 block table 所需 UB,要求不超过 184 KiB;PA 还要求涉及的 block size 满足二次幂条件,且模板属于支持的稀疏路径。失败则回落到非向量化路径。 + +这里列的是 **Host 的申请公式**。Kernel 的 `InitMMResBuf`、`GetKVPhyAddr` 使用 `GetBlockNum()` 等计算各区域偏移,不应仅凭 Host 的预留总量反推每个内部区域起点;特别是辅助索引区的 Host 预留可能大于设备实际布局。修改时要逐区验证设备最大访问地址不超过申请量。 + +### 4.5 arch35 的片上静态布局 + +[CSA Kernel `InitMMResBuf`](../op_kernel/arch35/sparse_flash_mla_csa_kernel_arch35.h) 中: + +```text +每 AIC 的 L1: [P槽0 16KiB][P槽1 16KiB][Cube侧Q/KV区...] +每 AIV 的 UB: [BMM2 64KiB][BMM1槽0 16KiB][BMM1槽1 16KiB][Vector私有区...] +``` + +P 每槽 `64*128*2=16 KiB`;每 AIV 负责半个 M 块,所以 BMM1 每槽 `32*128*4=16 KiB`,BMM2 为 `32*512*4=64 KiB`。P 必须位于 L1 前部且 Vector/Cube 地址一致。 + +[Cube 缓冲](../op_kernel/arch35/sparse_flash_mla_csa_block_cube_arch35.h) 使用三槽 Q、三槽 KV,以及双槽 L0A/B/C。单槽常量见 [arch35 common](../op_kernel/arch35/sparse_flash_mla_common_arch35.h):Q 32 KiB、KV 128 KiB、L0A 16 KiB、L0B 32 KiB、L0C 128 KiB。L1 中连同 P 共 `32+3*32+3*128=512 KiB`。 + +[Vector `InitLocalBuffer`](../op_kernel/arch35/sparse_flash_mla_csa_block_vector_arch35.h) 在上述 96 KiB UB 后继续排布: + +| 对象 | 大小 | +| --- | --- | +| softmax sum/max/exp,各双槽 | 合计 `6*256=1536` 字节 | +| common / sinks | 各 512 字节 | +| 稀疏 gather stage0,双槽 | `2*16*512*2=32 KiB` | +| LSE,按需双槽 | 512 字节 | +| stage1 P 输出,双槽 | `2*33*128*2=16896` 字节,stride 为 33 | +| stage2 FP32 输出 | `32*512*4=64 KiB` | +| batch consistency 附加状态 | `4*256+768*4=4096` 字节 | + +因此稀疏路径、开启 LSE、关闭 batch consistency 时,该静态主流程布局合计 216576 字节;开启一致性时为 220672 字节。地址向量化和 FD 有各自的执行相位/缓冲布局,不能将不同相位占用机械相加。无稀疏 gather 的 SWA 路径也不分配 stage0 双槽。 + +### 4.6 为什么需要双缓冲和三缓冲 + +双缓冲让一个槽被消费时,另一个槽可以写入;三槽 KV 则覆盖 gather、加载和矩阵消费之间更长的流水距离。复用的前提是“旧消费者已结束”,并不只是 `taskId%2` 或 `%3` 算对。 + +arch35 使用 `CROSSCORE_V0RES`、`CROSSCORE_BMM1`、`CROSSCORE_L1P`、`CROSSCORE_BMM2` 等跨核 flag,以及 `INNERCORE_*` 的 MTE/Vector/Cube 事件。flag 绑定资源就绪/释放关系,修改循环、提前 return、跳过空任务时必须保持生产与消费配对。split-G 中无计算任务的核也可能仍需参加同步,不能直接删除其补齐循环。 + +## 5. 算子计算流 + +### 5.1 总体流程 + +```mermaid +flowchart TD + A[准备输入、长度、索引与属性] --> B[Metadata 生成 FA 与 FD 任务] + B --> C[Host 校验、tiling key、workspace] + C --> D[Kernel 初始化与解析 metadata] + D --> E[按 batch / query组 / S2块遍历] + E --> F[寻址与稀疏 gather] + F --> G[Cube BMM1: Q乘K转置] + G --> H[Vector Vec1: scale、mask、在线softmax] + H --> I[Cube BMM2: P乘V] + I --> J[Vector Vec2: 重缩放并累加] + J --> E + J --> K[完整行输出或写部分结果] + K --> L[有跨核任务时执行 FD 归约] + L --> M[attn_out 与可选 LSE] +``` + +图中循环是逻辑依赖;真实实现将不同 S2 块的阶段重叠执行。 + +### 5.2 一个计算块内部的职责 + +| 阶段 | 执行单元 | 输入 → 输出 | +| --- | --- | --- | +| 参数计算 | 标量逻辑 | metadata、长度、mask → batch/query/S2 范围与尾块 | +| Vec0 | Vector,稀疏路径 | 索引、block table、KV → 连续 KV 小块与有效性信息 | +| LoadQK | Cube 搬运流水 | Q、gather/连续 KV → L1/L0 | +| BMM1 | Cube | `[M,D] * [D,N] → [M,N]` FP32 分数 | +| Vec1 | Vector | 分数 → 缩放、mask、max/sum、低精度 P | +| BMM2 | Cube | `[M,N] * [N,D] → [M,D]` FP32 块结果 | +| Vec2 | Vector | 新块结果、旧累加值、指数修正 → 更新后的输出状态 | +| FD | Vector,按 metadata | 部分输出/max/sum → 最终归一化输出 | + +连续 SWA/HCA 可直接读取连续或分页 KV,少掉通用稀疏 gather 阶段。ori 与 cmp 可以分多轮读取,但必须共享同一行的 softmax 状态;分别 softmax 后直接相加不等价。 + +### 5.3 在线 softmax 的正确递推 + +以一行解释。保存行最大值 m、指数和 l、未归一化输出向量 a。若在此处一次性计入有限 sink,可初始化 `m=sink, l=1, a=0`。对当前 KV 块的有效 logits x: + +\[ +m'=\max(m,\max_j x_j),\quad \alpha=e^{m-m'},\quad p_j=e^{x_j-m'}, +\] +\[ +l'=\alpha l+\sum_jp_j,\qquad a'=\alpha a+\sum_jp_jV_j. +\] + +所有块结束后 `O=a/l`、`LSE=m+log(l)`。m 变大时旧累加量必须乘 alpha,否则结果错误。实现用 FP32 管理状态,但 P 在送入第二次矩阵乘时会转换为输入精度;Golden 也显式模拟这一转换,所以不能期望与全 FP32 dense attention 完全一致。 + +无效索引或 mask 掉的条目必须对指数和贡献零。空行、全 mask、全空 batch 的输出有专门初始化/跳过分支;特别是 LSE 的空行写出规则要按实际分支与 Golden 核对,不可直接执行可能产生 NaN 的 `-inf-(-inf)`。 + +### 5.4 arch35 CSA 的流水时序 + +`ProcessMainLoop` 使用 `RunInfo[4]` 保存不同在途任务。稳定阶段,以当前计数 t 表示: + +| 逻辑任务 | 本轮推进的阶段 | +| --- | --- | +| t | Vec0 gather | +| t-1 | Cube LoadQK | +| t-2 | Cube BMM1,随后 Vector Vec1 | +| t-3 | Cube BMM2,随后 Vector Vec2 | + +同一任务的 BMM1/Vec1、BMM2/Vec2 通过 flag 保证先后,表格不表示它们同时读写同一块数据。Q/KV 三槽、BMM1/P 双槽与 RunInfo 四槽承担不同生命周期,不能统一改成一个取模值。 + +末尾还要执行排空轮次,让最后几块走完 Vec1/BMM2/Vec2。`notLastThreeLoop`、`notLastTwoLoop`、`notLast` 正是控制预热和排空。只遍历“真实 S2 块数”而删除排空,会漏写尾部结果。 + +### 5.5 跨核 S2 归约 + +若一行被切到多个核,每份结果带有局部 \((m_r,l_r,a_r)\)。合并公式为: + +\[ +m=\max_r m_r,\quad l=\sum_r e^{m_r-m}l_r,\quad +a=\sum_r e^{m_r-m}a_r,\quad O=a/l. +\] + +若 staging 保存的是局部已归一化输出 \(O_r\),则输出权重必须为 \(e^{m_r-m}l_r\),不能平均各核输出。sink 只能在整行中计入一次;修改分块或 FD 时必须检查其初始化和合并位置。 + +arch35 主循环完成后,AIV 经同步再按 FD metadata 执行 `ProcessFlashDecode`;相关 staging、归约辅助代码见 [Vector](../op_kernel/arch35/sparse_flash_mla_csa_block_vector_arch35.h) 和 [flash_decode](../op_kernel/arch35/common/flash_decode.h)。 + +## 6. 用一个例子串起来 + +选择 arch35 CSA,BSND,`B=1,S1=1,G=32,D=512`。假设经 mask/有效长度裁剪后有 128 个 ori 条目、512 个有效 cmp 索引条目,且这个 query 未跨核拆分。 + +1. Metadata 组织一个 query 组,ori 为一个 S2 块,cmp 为四个 S2 块;共五块,每块 128 条。该例的单核分配是说明条件,不是任意设备上的调度保证。 +2. 此 query 的 Q 矩阵逻辑 shape 为 `[32,512]`。基本容量仍按 64 行分配,尾行由运行参数控制。 +3. 每个 S2 块形成 `[128,512]` KV;cmp 块先按索引 gather。BMM1 有效结果为 `[32,128]`。 +4. Vec1 对五块维护同一份 softmax 状态,Vec2 对五份 `[32,512]` 结果做指数修正累加。 +5. 最后归一化输出 `[1,1,32,512]`;开启 LSE 时输出 `[1,1,1,32]`。 + +如果 Metadata 把同一行 S2 划给两核,两核改为写局部状态,最后走 FD。若 G 改为 128,arch35 还会启用 split-G,一对 AIC 各处理 64 个 heads;head 切分本身不需要把两个不同 head 的输出相加。 + +## 7. 代码阅读顺序与修改入口 + +| 阅读顺序 | 文件/符号 | 需要回答的问题 | +| --- | --- | --- | +| 1 | [README](../README.md)、[InferShape](../op_host/sparse_flash_mla_infershape.cpp) | 支持哪些输入,输出如何排列? | +| 2 | [Golden](../tests/pytest/sparse_flash_mla_golden.py) | 选择、mask、sink、LSE 的数学语义是什么? | +| 3 | [Host tiling](../op_host/sparse_flash_mla_tiling.cpp)、[tiling 结构](../op_host/sparse_flash_mla_tiling.h) | 哪个分支,哪些字段和 workspace? | +| 4 | [Metadata](../../sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.cpp) | 任务游标、负载和 FD 如何生成? | +| 5 | [Kernel 入口](../op_kernel/sparse_flash_mla.cpp) | 实际实例化哪个模板? | +| 6 | [arch35 CSA](../op_kernel/arch35/sparse_flash_mla_csa_kernel_arch35.h)、[SWA](../op_kernel/arch35/sparse_flash_mla_swa_kernel_arch35.h) | Init/ProcessMainLoop 如何调度? | +| 7 | [arch22 CSA](../op_kernel/arch22/sparse_flash_mla_csa_kernel.h)、[SWA](../op_kernel/arch22/sparse_flash_mla_swa_kernel.h) | GM 中转与旧架构流水有何差异? | +| 8 | 对应 block_cube、block_vector、kvcache | 每次搬运、矩阵乘、mask 和同步具体做什么? | + +增加布局要同时检查 shape/stride 校验、Metadata 长度解析、寻址和输出布局。增加 mask 要同时修改 Metadata 负载估计与 Kernel 可见范围,避免任务遗漏。改变基本块要同时检查 Host workspace、Metadata block 单位、本地静态布局、尾块和同步轮次。增加可选输入还要同步算子注册、Torch 绑定、参数索引及测试;不能只在 Kernel 函数中追加参数。 + +## 8. 验证与故障定位 + +### 8.1 建议的验证范围 + +| 类别 | 最小覆盖重点 | +| --- | --- | +| Host UT | dtype/shape/layout、缺失输入、无效属性、tiling key、workspace;分别覆盖 arch22/35 | +| 数值 | SWA、稀疏 ori、CSA、HCA;FP16/BF16;sink 与 LSE 开关 | +| 尾块 | S2 在基本块边界前后;小 G、G=64、G>64;950 的非偶数 G | +| 变长 | TND 前缀和,seqused 小于存储长度,空有效行,不同 batch 长度 | +| 稀疏/分页 | TopK 长度为 0/容量值,尾部 -1,跨物理页,不同 ori/cmp block size | +| 压缩 mask | cmp ratio、显式 cmp 长度与 residual,右下因果边界 | +| 调度 | 行内 S2 跨核、split-G、空闲核、流水排空 | +| 一致性 | 单样本与拼 batch/重排 batch;普通模式与 batch consistency;aclgraph | + +已有 [Host UT](../tests/ut/op_host/test_sparse_flash_mla_tiling.cpp)、[arch35 UT](../tests/ut/op_host/arch35/test_sparse_flash_mla_tiling.cpp)、[单跑测试](../tests/pytest/test_sparse_flash_mla_single.py)、[batch consistency 测试](../tests/pytest/test_sparse_flash_mla_batch_consistency.py) 可作为起点。Golden 比较应同时检查输出与 LSE,并沿用仓库比较方法,避免随意放宽阈值掩盖边界错误。 + +在安装了匹配 CANN、torch_npu 和本仓库扩展的 Ascend 环境中,可从测试目录使用已有脚本: + +```bash +cd attention/sparse_flash_mla/tests/pytest +bash test_run.sh --help +bash test_run.sh single +bash test_run.sh single --batch-consistency on +``` + +用例参数来自相应 paramset/测试输入,批量和 aclgraph 模式按脚本帮助配置。本文编写只进行了源码与文档静态核对,未执行 Ascend 编译或上板测试。 + +### 8.2 从现象找位置 + +| 现象 | 优先检查 | +| --- | --- | +| 所有输出有系统性缩放偏差 | sink 是否进入分母,scale 是否重复应用,ori/cmp 是否被分别归一化 | +| 只有长序列不对 | 在线 alpha 修正、FD 部分结果、workspace 槽位 | +| 只有尾部 query/S2 不对 | metadata 游标、有效长度、尾块 mask、流水排空 | +| PA 错而连续正确 | 逻辑 token → 页号 → 物理页,ori/cmp block size 和 stride | +| 单 batch 正确、拼 batch 错 | TND 存储起点与有效长度混用、metadata 未匹配、归约顺序 | +| 卡住或偶现错误 | flag 配对、槽位提前复用、空核同步、split-G 补齐轮次 | +| 输出正常但 LSE 错 | sink、max/sum 更新、LSE 布局、空行初始化 | +| workspace 越界 | 元素/字节混用、逻辑槽/物理核混用、Host 512 与 arch35 128 混用 | + +性能分析首先分辨瓶颈:稀疏 gather 是否受不连续 GM 访问限制,Cube 是否因小 G/尾块利用率不足,Vector softmax 是否拖慢流水,FD 是否占比过高,核间实际负载是否均衡。基本块、缓存数量与归约策略相互影响,应在正确性覆盖后用对应平台 profiling 数据决定优化方向。 diff --git a/csrc/attention/sparse_flash_mla/docs/ratio2_a2a3.md b/csrc/attention/sparse_flash_mla/docs/ratio2_a2a3.md new file mode 100644 index 000000000000..89d2d4968b9e --- /dev/null +++ b/csrc/attention/sparse_flash_mla/docs/ratio2_a2a3.md @@ -0,0 +1,98 @@ +# A2/A3 cmp_ratio=1/2 适配说明 + +## 1. 范围与结果 + +本次为前向 SparseFlashMla 的 CSA 增加压缩倍率 1 和 2,配套修改 SparseFlashMlaMetadata。A2/A3 HCA 保持仅支持 128 的原有逻辑。950 的倍率范围不变;不涉及梯度算子、压缩算子或上游索引器的实现。 + +| A2/A3 模式 | 修改前 | 修改后 | +| --- | --- | --- | +| SWA,无 cmp KV | 1 | 0 | +| CSA,有 cmp KV 和 cmp 索引 | 4 | 1、2、4 | +| HCA,有 cmp KV、无 cmp 索引 | 128 | 128 | + +其余约束保持现有实现,例如 cmp causal mask 为 3、CSA TopK 容量为 512/1024、ori 窗口为左 127/右 0。ratio=1/2 不意味着扩大其他输入规格。 + +## 2. 倍率在代码中的传递 + +```text +调用方 cmp_ratio=1/2 + Metadata Host: IsCmpRatioSupportSmla → ParamsCheck + AICPU: cmpRatio_ → GetRevertS2Size → CalcCmpBlockRange + → block/cost → 分核 → metadata + 主算子 Host: CheckSingleParaCmpRatio → cmpParams.cmpRatio + arch22 Kernel: constInfo.cmpRatio + → 压缩有效范围 → gather / 逐行 mask → attention +``` + +倍率是运行时 tiling 参数,不是 tiling key 的模板维度。两次调用必须使用同样的倍率、有效长度、residual,改变这些参数后应重新生成 metadata。 + +## 3. 适配点与实现决策 + +| 层次 | 文件 / 符号 | 本次处理 | +| --- | --- | --- | +| 主算子 Host | [sparse_flash_mla_tiling.cpp](../op_host/sparse_flash_mla_tiling.cpp),`CheckSingleParaCmpRatio` | CSA 增加 1/2;保持 HCA=128,SWA 改用 0,更新报错 | +| Metadata Host | [metadata_check.h](../../sparse_flash_mla_metadata/op_host/sparse_flash_mla_metadata_check.h),`IsCmpRatioSupportSmla` | 与主算子允许集合一致,更新错误信息 | +| Metadata AICPU | [metadata_aicpu.cpp](../../sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.cpp) | 保留已有运行时乘除公式、residual 范围校验及 block/cost 逻辑 | +| CSA Kernel | [csa_kernel.h](../op_kernel/arch22/sparse_flash_mla_csa_kernel.h) | 已通过 cmpRatio 计算长度和 cmpS2IdLimit,无需新增 ratio 模板 | +| CSA gather | [csa_block_vector.h](../op_kernel/arch22/sparse_flash_mla_csa_block_vector.h) | 沿用 cmpS2IdLimit 检查索引,重点验证 causal 边界 | +| HCA Kernel / mask | [swa_kernel.h](../op_kernel/arch22/sparse_flash_mla_swa_kernel.h) | 保持原有实现,不增加倍率 1/2 支持 | +| tiling / 内存 | Host SplitBalanced、DoOpTiling | 保持 S2=512 和原缓冲区配置;增加序列长度会增加循环次数,不直接扩大片上基本块 | +| 接口 / 绑定 | 现有整数属性透传 | 不修改 ABI、输出 shape、dtype 或 tiling key | + +Kernel 与 AICPU 的相关公式已经参数化,所以本次不为“有 Kernel 变更”而改写等价计算。真正的生产代码变更是两处 Host 倍率白名单。 + +## 4. 长度、residual 与 causal 边界 + +令 Lc 为压缩后有效长度,r 为倍率,residual 为余数: + +```text +L = Lc * r + residual +p = L - Lq + query_index +visible_cmp = clamp((p + 1) / r, 0, Lc) +``` + +非负坐标下除法向下取整,负范围由具体分支裁剪。r=2 时 residual 只能为 0 或 1。cmp_mask_mode=3 且 r!=1 时 residual 必须同时传给 Metadata 和主算子,即使余数为 0。 + +例:原始各 batch 长度 [3,3],压缩有效长度为 [1,1],residual 为 [1,1],压缩 TND 前缀和为 [0,1,2]。不能将 ori 前缀和 [0,3,6] 逐项除以 2 得到 [0,1,3] 后当作压缩前缀和。 + +CSA 调用方必须实际生成对应倍率的 KV、索引和分页表。倍率 1 的压缩长度等于原长度且不需要 residual;倍率 2 的 residual 取 0 或 1。相同原始长度下,从倍率 4 改为 1/2 会增加压缩 KV 条目数量,应重新预算输入 cache;不能只改变算子属性。A2/A3 的 HCA 仍传 128。 + +## 5. 测试与验收 + +算子源库随附的验证用例: + +- Host tiling UT:覆盖 A2/A3 的 CSA=1/2 成功、HCA 非法倍率拒绝、缺少 residual 和 SWA 非零倍率拒绝。 +- Metadata API UT:通过 ParamsCheck 检查主接口与前置接口的倍率规则一致。 +- ratio2 数值回归:覆盖 FP16/BF16、BSND/TND/PA_BBND、residual=0/1、边界压缩长度、双 batch 和多 query 行,并比较 attn_out 与 LSE。 + +vLLM Ascend 侧另有组网路由单测,验证 ratio 0/1/2 均进入原生 SparseFlashMla 路径。算子源库的 CANN UT/pytest 未复制到 `csrc` 发布目录。 + +上板前须重新编译/安装修改后的主算子与 Metadata,并在算子源库执行 Host UT、Metadata API UT 和真实算子数值用例。A2、A3 都要覆盖 ratio=1/2,已有 ratio=4/128 用例也需回归;Host UT 不执行 AICPU 任务切分,必须由真实 Metadata + 主算子调用补齐。 + +## 6. 已知边界与后续验证 + +arch22 的 TND 压缩长度读取当前直接使用 cu_seqlens_cmp_kv 相邻差值,Metadata 则优先使用 seqused_cmp_kv。对“有效长度小于存储长度”的 TND 输入,这两个口径需要另行统一。本次不改变原有长度接口语义,新增 TND 用例使用二者一致的长度,不将这种带 padding 的有效长度覆盖场景声明为已解决。 + +CSA 多核调度、G=1/128、非均匀 batch、aclgraph、极端空范围还应在目标平台验收时扩展覆盖。普通模式本次保持现有流水和内存设计,性能结论必须由 profiling 给出。 + +## 7. 当前仓库的模型编译范围 + +当前仓库按 Aurora 调用范围裁剪编译模板;前面列出的源库通用能力及其回归矩阵不等于本仓库裁剪后的支持范围。主算子只编译 BF16、TND Q、PA_BBND KV、SWA/CSA。BSND、非分页 KV、HCA 和两个独立 ori sparse 模板不再编译;host 会拒绝这些未编译的布局、模式或 dtype。 + +| 构建目标 | 裁剪前 key 数 | 当前 key 数 | 保留的硬件特化 | +| --- | --- | --- | --- | +| A2/A3 | 320 | 6 | CSA 的 `HEAD_RATIO_ONE=0/1`,`SPLIT_G=IS_VEC_S2PHYADDR=0` | +| A5 | 320 | 12 | `HEAD_RATIO_ONE=0`,保留 split-G 和 CSA 物理地址向量化 | +| Host | 320 | 14 | 两个设备集合的并集,仅用于 key 编码及校验 | + +两个设备集合共有 4 个 key,共用部分分别编译。A2/A3 不编译 A5 专用标志组合;A5 不编译 A2/A3 的单头专用模板。确定性级别仍进入 host key,因此两边都保留 `BATCH_CONSISTENCY=0/1`。`FLASH_DECODE` 固定为 0,decode 调度由 metadata 驱动。模板声明的参数顺序、位宽和取值不变,保留的 key 编码不变。 + +`DeepseekV41EagerAttentionImpl._native_attention` 的 C0 走 SWA,C1/C2 走 CSA;cache spec 固定 BF16,query 与 KV 需要同 dtype。压缩比例、TopK 和请求长度都是运行时参数。编译保留全部本地 head 数边界,包括 TP 后只剩一个 query head 的 CSA。 + +无需 CANN 的矩阵回归: + +```bash +python3 -m unittest discover -s tests/ut/ops -p test_aurora_tiling_keys.py -v +``` + +该检查枚举实际头文件的预处理结果,覆盖两种架构、模型调用布局、C0/C1/C2、单头和确定性边界,并检查 key 声明不变。它不执行 CANN 编译或 NPU kernel。key 数不包含 dtype 对编译任务数的影响,也不能直接换算整包编译耗时;耗时及数值结果需在配套 CANN/NPU 环境中重编译后验证。 diff --git a/csrc/attention/sparse_flash_mla/examples/test_aclnn_sparse_flash_mla.cpp b/csrc/attention/sparse_flash_mla/examples/test_aclnn_sparse_flash_mla.cpp new file mode 100644 index 000000000000..b64c1c565a46 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/examples/test_aclnn_sparse_flash_mla.cpp @@ -0,0 +1,432 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This file is a part of the CANN Open Software. + * Licensed under CANN Open Software License Agreement Version 1.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file test_aclnn_sparse_flash_mla.cpp + * \brief SparseFlashMla + SparseFlashMlaMetadata 算子调用示例(CSA) + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include "acl/acl.h" +#include "aclnnop/aclnn_sparse_flash_mla.h" +#include "aclnnop/aclnn_sparse_flash_mla_metadata.h" + +#define CHECK_RET(cond, return_expr) \ + do { \ + if (!(cond)) { \ + return_expr; \ + } \ + } while (0) + +#define LOG_PRINT(message, ...) \ + do { \ + printf(message, ##__VA_ARGS__); \ + } while (0) + +namespace { + +int64_t GetShapeSize(const std::vector& shape) +{ + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} + +uint16_t FloatToFp16(float f) +{ + uint32_t bits; + std::memcpy(&bits, &f, sizeof(bits)); + uint32_t sign = (bits >> 31) & 0x1u; + int32_t exp = static_cast((bits >> 23) & 0xffu) - 127 + 15; + uint32_t mant = (bits >> 13) & 0x3ffu; + if (exp <= 0) { + return static_cast(sign << 15); + } + if (exp >= 31) { + return static_cast((sign << 15) | 0x7c00u); + } + return static_cast((sign << 15) | (static_cast(exp) << 10) | mant); +} + +float Fp16ToFloat(uint16_t h) +{ + uint32_t sign = (h >> 15) & 0x1u; + uint32_t exp = (h >> 10) & 0x1fu; + uint32_t mant = h & 0x3ffu; + uint32_t f; + if (exp == 0) { + f = (sign << 31) | (mant << 13); + } else if (exp == 31) { + f = (sign << 31) | 0x7f800000u | (mant << 13); + } else { + f = (sign << 31) | ((exp + 127u - 15u) << 23) | (mant << 13); + } + float result; + std::memcpy(&result, &f, sizeof(result)); + return result; +} + +void PrintOutResult(const std::vector& shape, void** deviceAddr) +{ + auto size = GetShapeSize(shape); + std::vector resultData(size, 0); + auto ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), + *deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); return); + for (int64_t i = 0; i < size && i < 10; i++) { + LOG_PRINT("result[%ld] is: %f\n", i, Fp16ToFloat(resultData[i])); + } +} + +int Init(int32_t deviceId, aclrtContext* context, aclrtStream* stream) +{ + auto ret = aclInit(nullptr); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret); + ret = aclrtSetDevice(deviceId); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret); + ret = aclrtCreateContext(context, deviceId); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateContext failed. ERROR: %d\n", ret); return ret); + ret = aclrtSetCurrentContext(*context); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetCurrentContext failed. ERROR: %d\n", ret); return ret); + ret = aclrtCreateStream(stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret); + return 0; +} + +template +int CreateAclTensor(const std::vector& hostData, const std::vector& shape, void** deviceAddr, + aclDataType dataType, aclTensor** tensor) +{ + auto size = GetShapeSize(shape) * sizeof(T); + if (size > 0) { + auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret); + ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret); + } else { + *deviceAddr = nullptr; + } + + std::vector strides(shape.size(), 1); + for (int64_t i = static_cast(shape.size()) - 2; i >= 0; i--) { + strides[i] = shape[i + 1] * strides[i + 1]; + } + + *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), *deviceAddr); + return 0; +} + +std::vector MakeFp16Data(int64_t size, float value) +{ + std::vector data(static_cast(size), FloatToFp16(value)); + return data; +} + +} // namespace + +int main() +{ + // 1. (固定写法)device/stream初始化,参考acl API手册 + // 根据自己的实际device填写deviceId + int32_t deviceId = 0; + aclrtContext context = nullptr; + aclrtStream stream = nullptr; + auto ret = Init(deviceId, &context, &stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); + + int64_t B = 4; + int64_t S1 = 128; + int64_t S2 = 8192; + int64_t N1 = 64; + int64_t N2 = 1; + int64_t D = 512; + int64_t K = 512; + int64_t oriBlockSize = 128; + int64_t cmpBlockSize = 128; + int64_t s2Act = 4096; + int64_t cmpRatio = 4; + int64_t oriWinLeft = 127; + int64_t oriWinRight = 0; + int64_t oriMaskMode = 4; + int64_t cmpMaskMode = 3; + double softmaxScale = 1.0 / sqrt(static_cast(D)); + + int64_t T1 = B * S1; + int64_t cmpKvLen = s2Act / cmpRatio; + int64_t oriBlockNum = ((s2Act + oriBlockSize - 1) / oriBlockSize) * B; + int64_t cmpBlockNum = ((cmpKvLen + cmpBlockSize - 1) / cmpBlockSize) * B; + + // 2. 构造输入与输出,需要根据API的接口自定义构造 + std::vector qShape = {T1, N1, D}; + std::vector oriKvShape = {oriBlockNum, oriBlockSize, N2, D}; + std::vector cmpKvShape = {cmpBlockNum, cmpBlockSize, N2, D}; + std::vector cmpSparseIndicesShape = {T1, N2, K}; + std::vector oriBlockTableShape = {B, (s2Act + oriBlockSize - 1) / oriBlockSize}; + std::vector cmpBlockTableShape = {B, (cmpKvLen + cmpBlockSize - 1) / cmpBlockSize}; + std::vector cuSeqLensQShape = {B + 1}; + std::vector seqUsedOriKvShape = {B}; + std::vector seqUsedCmpKvShape = {B}; + std::vector cmpResidualKvShape = {B}; + std::vector sinksShape = {N1}; + std::vector metadataShape = {1024}; + std::vector attnOutShape = {T1, N1, D}; + std::vector softmaxLseShape = {T1, N1, 1}; + // 对全部 5 个输入调用 Contiguous,optional 输入传 shape 为 {0} 的空 tensor。 + std::vector emptyShape = {0}; + + void* qDeviceAddr = nullptr; + void* oriKvDeviceAddr = nullptr; + void* cmpKvDeviceAddr = nullptr; + void* cmpSparseIndicesDeviceAddr = nullptr; + void* oriBlockTableDeviceAddr = nullptr; + void* cmpBlockTableDeviceAddr = nullptr; + void* cuSeqLensQDeviceAddr = nullptr; + void* cuSeqLensOriKvDeviceAddr = nullptr; + void* cuSeqLensCmpKvDeviceAddr = nullptr; + void* seqUsedQDeviceAddr = nullptr; + void* seqUsedOriKvDeviceAddr = nullptr; + void* seqUsedCmpKvDeviceAddr = nullptr; + void* cmpResidualKvDeviceAddr = nullptr; + void* sinksDeviceAddr = nullptr; + void* metadataDeviceAddr = nullptr; + void* attnOutDeviceAddr = nullptr; + void* softmaxLseDeviceAddr = nullptr; + + aclTensor* q = nullptr; + aclTensor* oriKv = nullptr; + aclTensor* cmpKv = nullptr; + aclTensor* cmpSparseIndices = nullptr; + aclTensor* oriBlockTable = nullptr; + aclTensor* cmpBlockTable = nullptr; + aclTensor* cuSeqLensQ = nullptr; + aclTensor* cuSeqLensOriKv = nullptr; + aclTensor* cuSeqLensCmpKv = nullptr; + aclTensor* seqUsedQ = nullptr; + aclTensor* seqUsedOriKv = nullptr; + aclTensor* seqUsedCmpKv = nullptr; + aclTensor* cmpResidualKv = nullptr; + aclTensor* sinks = nullptr; + aclTensor* metadata = nullptr; + aclTensor* attnOut = nullptr; + aclTensor* softmaxLse = nullptr; + + int64_t qSize = GetShapeSize(qShape); + int64_t oriKvSize = GetShapeSize(oriKvShape); + int64_t cmpKvSize = GetShapeSize(cmpKvShape); + int64_t cmpSparseIndicesSize = GetShapeSize(cmpSparseIndicesShape); + int64_t oriBlockTableSize = GetShapeSize(oriBlockTableShape); + int64_t cmpBlockTableSize = GetShapeSize(cmpBlockTableShape); + int64_t attnOutSize = GetShapeSize(attnOutShape); + int64_t softmaxLseSize = GetShapeSize(softmaxLseShape); + + std::vector qHostData = MakeFp16Data(qSize, 1.0f); + std::vector oriKvHostData = MakeFp16Data(oriKvSize, 1.0f); + std::vector cmpKvHostData = MakeFp16Data(cmpKvSize, 1.0f); + std::vector cmpSparseIndicesHostData(cmpSparseIndicesSize); + std::vector oriBlockTableHostData(oriBlockTableSize); + std::iota(oriBlockTableHostData.begin(), oriBlockTableHostData.end(), 0); + std::vector cmpBlockTableHostData(cmpBlockTableSize); + std::iota(cmpBlockTableHostData.begin(), cmpBlockTableHostData.end(), 0); + std::vector cuSeqLensQHostData(B + 1); + for (int64_t i = 0; i <= B; i++) { + cuSeqLensQHostData[i] = static_cast(i * S1); + } + std::vector emptyHostData; + std::vector seqUsedOriKvHostData(B, static_cast(s2Act)); + std::vector seqUsedCmpKvHostData(B, static_cast(cmpKvLen)); + std::vector cmpResidualKvHostData(B, static_cast(s2Act % cmpRatio)); + std::vector sinksHostData(N1, 1.0f); + std::vector metadataHostData(1024, 0); + std::vector attnOutHostData = MakeFp16Data(attnOutSize, 0.0f); + std::vector softmaxLseHostData(softmaxLseSize, 0.0f); + + std::mt19937 gen(42); + for (int64_t t = 0; t < T1; t++) { + for (int64_t n = 0; n < N2; n++) { + for (int64_t k = 0; k < K; k++) { + cmpSparseIndicesHostData[t * N2 * K + n * K + k] = static_cast(gen() % cmpKvLen); + } + } + } + + ret = CreateAclTensor(qHostData, qShape, &qDeviceAddr, aclDataType::ACL_FLOAT16, &q); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(oriKvHostData, oriKvShape, &oriKvDeviceAddr, aclDataType::ACL_FLOAT16, &oriKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpKvHostData, cmpKvShape, &cmpKvDeviceAddr, aclDataType::ACL_FLOAT16, &cmpKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpSparseIndicesHostData, cmpSparseIndicesShape, &cmpSparseIndicesDeviceAddr, + aclDataType::ACL_INT32, &cmpSparseIndices); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(oriBlockTableHostData, oriBlockTableShape, &oriBlockTableDeviceAddr, aclDataType::ACL_INT32, + &oriBlockTable); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpBlockTableHostData, cmpBlockTableShape, &cmpBlockTableDeviceAddr, aclDataType::ACL_INT32, + &cmpBlockTable); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cuSeqLensQHostData, cuSeqLensQShape, &cuSeqLensQDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensQ); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(emptyHostData, emptyShape, &cuSeqLensOriKvDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensOriKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(emptyHostData, emptyShape, &cuSeqLensCmpKvDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensCmpKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(emptyHostData, emptyShape, &seqUsedQDeviceAddr, aclDataType::ACL_INT32, &seqUsedQ); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(seqUsedOriKvHostData, seqUsedOriKvShape, &seqUsedOriKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedOriKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(seqUsedCmpKvHostData, seqUsedCmpKvShape, &seqUsedCmpKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedCmpKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpResidualKvHostData, cmpResidualKvShape, &cmpResidualKvDeviceAddr, aclDataType::ACL_INT32, &cmpResidualKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(sinksHostData, sinksShape, &sinksDeviceAddr, aclDataType::ACL_FLOAT, &sinks); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(metadataHostData, metadataShape, &metadataDeviceAddr, aclDataType::ACL_INT32, &metadata); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(attnOutHostData, attnOutShape, &attnOutDeviceAddr, aclDataType::ACL_FLOAT16, &attnOut); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(softmaxLseHostData, softmaxLseShape, &softmaxLseDeviceAddr, aclDataType::ACL_FLOAT, &softmaxLse); + CHECK_RET(ret == ACL_SUCCESS, return ret); + + char layoutQ[] = "TND"; + char layoutKv[] = "PA_BBND"; + + uint64_t metadataWorkspaceSize = 0; + aclOpExecutor* metadataExecutor = nullptr; + + // 3. 调用CANN算子库API,需要修改为具体的Api名称 + ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( + cuSeqLensQ, cuSeqLensOriKv, cuSeqLensCmpKv, + seqUsedQ, seqUsedOriKv, + seqUsedCmpKv, cmpResidualKv, nullptr, nullptr, + N1, N2, D, B, S1, S2, cmpKvLen, + 0, K, cmpRatio, + oriMaskMode, cmpMaskMode, + oriWinLeft, oriWinRight, + layoutQ, layoutKv, + true, true, + metadata, + &metadataWorkspaceSize, &metadataExecutor); + CHECK_RET(ret == ACL_SUCCESS, + LOG_PRINT("aclnnSparseFlashMlaMetadataGetWorkspaceSize failed. ERROR: %d\n", ret); return ret); + + void* metadataWorkspaceAddr = nullptr; + if (metadataWorkspaceSize > 0) { + ret = aclrtMalloc(&metadataWorkspaceAddr, metadataWorkspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate metadata workspace failed. ERROR: %d\n", ret); return ret); + } + + ret = aclnnSparseFlashMlaMetadata(metadataWorkspaceAddr, metadataWorkspaceSize, metadataExecutor, stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseFlashMlaMetadata failed. ERROR: %d\n", ret); return ret); + + // 4. (固定写法)同步等待任务执行结束 + ret = aclrtSynchronizeStream(stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream after metadata failed. ERROR: %d\n", ret); return ret); + + uint64_t workspaceSize = 0; + aclOpExecutor* executor = nullptr; + + ret = aclnnSparseFlashMlaGetWorkspaceSize( + q, oriKv, cmpKv, + nullptr, cmpSparseIndices, + oriBlockTable, cmpBlockTable, + cuSeqLensQ, nullptr, nullptr, + nullptr, seqUsedOriKv, + seqUsedCmpKv, cmpResidualKv, nullptr, nullptr, + sinks, metadata, + softmaxScale, cmpRatio, + oriMaskMode, cmpMaskMode, + oriWinLeft, oriWinRight, + layoutQ, layoutKv, + 1, + false, + attnOut, softmaxLse, + &workspaceSize, &executor); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseFlashMlaGetWorkspaceSize failed. ERROR: %d\n", ret); return ret); + + void* workspaceAddr = nullptr; + if (workspaceSize > 0) { + ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret); + } + + ret = aclnnSparseFlashMla(workspaceAddr, workspaceSize, executor, stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseFlashMla failed. ERROR: %d\n", ret); return ret); + + ret = aclrtSynchronizeStream(stream); + CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret); + + // 5.获取输出的值,将device侧内存上的结果拷贝至host侧,需要根据具体API的接口定义修改 + PrintOutResult(attnOutShape, &attnOutDeviceAddr); + + // 6. 释放aclTensor和aclScalar,需要根据具体API的接口定义修改 + aclDestroyTensor(q); + aclDestroyTensor(oriKv); + aclDestroyTensor(cmpKv); + aclDestroyTensor(cmpSparseIndices); + aclDestroyTensor(oriBlockTable); + aclDestroyTensor(cmpBlockTable); + aclDestroyTensor(cuSeqLensQ); + aclDestroyTensor(cuSeqLensOriKv); + aclDestroyTensor(cuSeqLensCmpKv); + aclDestroyTensor(seqUsedQ); + aclDestroyTensor(seqUsedOriKv); + aclDestroyTensor(seqUsedCmpKv); + aclDestroyTensor(cmpResidualKv); + aclDestroyTensor(sinks); + aclDestroyTensor(metadata); + aclDestroyTensor(attnOut); + aclDestroyTensor(softmaxLse); + + // 7. 释放device资源 + aclrtFree(qDeviceAddr); + aclrtFree(oriKvDeviceAddr); + aclrtFree(cmpKvDeviceAddr); + aclrtFree(cmpSparseIndicesDeviceAddr); + aclrtFree(oriBlockTableDeviceAddr); + aclrtFree(cmpBlockTableDeviceAddr); + if (cuSeqLensQDeviceAddr != nullptr) { + aclrtFree(cuSeqLensQDeviceAddr); + } + if (seqUsedOriKvDeviceAddr != nullptr) { + aclrtFree(seqUsedOriKvDeviceAddr); + } + if (seqUsedCmpKvDeviceAddr != nullptr) { + aclrtFree(seqUsedCmpKvDeviceAddr); + } + if (cmpResidualKvDeviceAddr != nullptr) { + aclrtFree(cmpResidualKvDeviceAddr); + } + aclrtFree(sinksDeviceAddr); + aclrtFree(metadataDeviceAddr); + aclrtFree(attnOutDeviceAddr); + aclrtFree(softmaxLseDeviceAddr); + if (metadataWorkspaceSize > 0) { + aclrtFree(metadataWorkspaceAddr); + } + if (workspaceSize > 0) { + aclrtFree(workspaceAddr); + } + aclrtDestroyStream(stream); + aclrtDestroyContext(context); + aclrtResetDevice(deviceId); + aclFinalize(); + + return 0; +} diff --git a/csrc/attention/sparse_flash_mla/op_host/CMakeLists.txt b/csrc/attention/sparse_flash_mla/op_host/CMakeLists.txt new file mode 100644 index 000000000000..7d756eccd601 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/CMakeLists.txt @@ -0,0 +1,39 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +add_op_to_compiled_list() + +if (BUILD_OPEN_PROJECT) + set(sparse_flash_mla_depends attention/common CACHE INTERNAL "Dependencies for sparse_flash_mla") + target_sources(op_host_aclnn PRIVATE + sparse_flash_mla_def.cpp + ) +endif() + +add_ops_compile_options( + OP_NAME SparseFlashMla + OPTIONS --cce-auto-sync=off + -Wno-deprecated-declarations + -mllvm -cce-vf-remove-membar=false + -mllvm -cce-aicore-hoist-movemask=false +) + +add_modules_sources(OPTYPE sparse_flash_mla ACLNNTYPE aclnn) + +include(${CMAKE_CURRENT_SOURCE_DIR}/checkers/checker_sources.cmake) +set(SPARSE_MLA_CHECKER_SRC_FILES + ${CMAKE_CURRENT_SOURCE_DIR}/checkers/sparse_flash_mla_checker.cpp +) + +add_tiling_modules() +add_sparse_mla_common_checker_sources(${OPHOST_NAME}_tiling_obj) +target_sources(${OPHOST_NAME}_tiling_obj PRIVATE + ${SPARSE_MLA_CHECKER_SRC_FILES} + sparse_flash_mla_tiling.cpp +) diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/base_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/base_checker.cpp new file mode 100644 index 000000000000..91faa86c42aa --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/base_checker.cpp @@ -0,0 +1,197 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "base_checker_sparse_flash_mla.h" +#include +#include +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +namespace { +const char *GetOpName(const CheckContext &context) { return context.opName == nullptr ? "SparseMla" : context.opName; } + +std::string ShapeToString(const gert::Shape *shape) +{ + if (shape == nullptr) { + return "nullptr"; + } + std::ostringstream oss; + oss << "("; + for (size_t i = 0; i < shape->GetDimNum(); ++i) { + if (i > 0) { + oss << ", "; + } + oss << shape->GetDim(i); + } + oss << ")"; + return oss.str(); +} + +template +std::string ValuesToString(std::initializer_list values) +{ + std::ostringstream oss; + bool first = true; + for (const auto value : values) { + if (!first) { + oss << ", "; + } + oss << static_cast(value); + first = false; + } + return oss.str(); +} +} // namespace + +ge::graphStatus BaseChecker::CheckSinglePara(const CheckContext &context) const +{ + (void)context; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus BaseChecker::CheckParaExistence(const CheckContext &context) const +{ + (void)context; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus BaseChecker::CheckFeature(const CheckContext &context) const +{ + (void)context; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus BaseChecker::CheckMultiPara(const CheckContext &context) const +{ + (void)context; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus BaseChecker::CheckTensorDesc(const CheckContext &context, const TensorParam ¶m, const char *name, + std::initializer_list dtypes) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(param.desc == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(GetOpName(context), name, "Tensor desc cannot be null"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + param.desc->GetOriginFormat() != ge::FORMAT_ND, + OP_LOGE_FOR_INVALID_FORMAT(GetOpName(context), name, + std::to_string(static_cast(param.desc->GetOriginFormat())).c_str(), "ND"), + return ge::GRAPH_FAILED); + const ge::DataType actual = param.desc->GetDataType(); + const std::string expectedDtypes = ValuesToString(dtypes); + OP_CHECK_IF(std::find(dtypes.begin(), dtypes.end(), actual) == dtypes.end(), + OP_LOGE_FOR_INVALID_DTYPE(GetOpName(context), name, + std::to_string(static_cast(actual)).c_str(), expectedDtypes.c_str()), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus BaseChecker::CheckDimNum(const CheckContext &context, const TensorParam ¶m, const char *name, + std::initializer_list dimNums) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(param.shape == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(GetOpName(context), name, "Tensor shape cannot be null"), + return ge::GRAPH_FAILED); + const size_t actual = param.shape->GetDimNum(); + const std::string expectedDimNums = ValuesToString(dimNums); + OP_CHECK_IF( + std::find(dimNums.begin(), dimNums.end(), actual) == dimNums.end(), + OP_LOGE_FOR_INVALID_SHAPEDIM(GetOpName(context), name, std::to_string(actual).c_str(), expectedDimNums.c_str()), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus BaseChecker::CheckNoEmptyDim(const CheckContext &context, const TensorParam ¶m, + const char *name) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(param.shape == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(GetOpName(context), name, "Tensor shape cannot be null"), + return ge::GRAPH_FAILED); + for (size_t i = 0; i < param.shape->GetDimNum(); ++i) { + OP_CHECK_IF(param.shape->GetDim(i) <= 0, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(GetOpName(context), name, ShapeToString(param.shape).c_str(), + "Each dimension must be greater than 0"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus BaseChecker::CheckShape(const CheckContext &context, const TensorParam ¶m, const char *name, + std::initializer_list expected) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(param.shape == nullptr || param.shape->GetDimNum() != expected.size(), + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(GetOpName(context), name, ShapeToString(param.shape).c_str(), + "Shape dim number does not match the documented shape"), + return ge::GRAPH_FAILED); + size_t index = 0; + for (const int64_t value : expected) { + OP_CHECK_IF(param.shape->GetDim(index) != value, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + GetOpName(context), name, ShapeToString(param.shape).c_str(), + ("Dimension " + std::to_string(index) + " must be " + std::to_string(value)).c_str()), + return ge::GRAPH_FAILED); + ++index; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus BaseChecker::CheckSameShape(const CheckContext &context, const TensorParam &left, const char *leftName, + const TensorParam &right, const char *rightName) const +{ + OP_CHECK_IF( + left.shape == nullptr || right.shape == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + GetOpName(context), (std::string(leftName) + " and " + rightName).c_str(), "Tensor shape cannot be null"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(left.shape->GetDimNum() != right.shape->GetDimNum(), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + GetOpName(context), (std::string(leftName) + " and " + rightName).c_str(), + (ShapeToString(left.shape) + " and " + ShapeToString(right.shape)).c_str(), + "Tensor dim numbers must be the same"), + return ge::GRAPH_FAILED); + for (size_t i = 0; i < left.shape->GetDimNum(); ++i) { + OP_CHECK_IF(left.shape->GetDim(i) != right.shape->GetDim(i), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + GetOpName(context), (std::string(leftName) + " and " + rightName).c_str(), + (ShapeToString(left.shape) + " and " + ShapeToString(right.shape)).c_str(), + ("Dimension " + std::to_string(i) + " must be the same").c_str()), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +int64_t BaseChecker::GetDim(const TensorParam ¶m, size_t index) const +{ + if (param.shape == nullptr || index >= param.shape->GetDimNum()) { + return -1; + } + return param.shape->GetDim(index); +} + +bool BaseChecker::CanOmitSequsedOriKv(const CheckContext &context) const { return context.oriTopkLength.present; } + +bool BaseChecker::CanOmitSequsedCmpKv(const CheckContext &context) const { return context.cmpTopkLength.present; } + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/base_checker_sparse_flash_mla.h b/csrc/attention/sparse_flash_mla/op_host/checkers/base_checker_sparse_flash_mla.h new file mode 100644 index 000000000000..f595b9fd4b43 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/base_checker_sparse_flash_mla.h @@ -0,0 +1,49 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_BASE_CHECKER_SPARSE_FLASH_MLA_H +#define SPARSE_MLA_BASE_CHECKER_SPARSE_FLASH_MLA_H + +#include +#include +#include "checker_context.h" + +namespace optiling { +namespace sparse_mla_checker { + +class BaseChecker { +public: + BaseChecker() = default; + virtual ~BaseChecker() = default; + + virtual ge::graphStatus CheckSinglePara(const CheckContext &context) const; + virtual ge::graphStatus CheckParaExistence(const CheckContext &context) const; + virtual ge::graphStatus CheckFeature(const CheckContext &context) const; + virtual ge::graphStatus CheckMultiPara(const CheckContext &context) const; + +protected: + ge::graphStatus CheckTensorDesc(const CheckContext &context, const TensorParam ¶m, const char *name, + std::initializer_list dtypes) const; + ge::graphStatus CheckDimNum(const CheckContext &context, const TensorParam ¶m, const char *name, + std::initializer_list dimNums) const; + ge::graphStatus CheckNoEmptyDim(const CheckContext &context, const TensorParam ¶m, const char *name) const; + ge::graphStatus CheckShape(const CheckContext &context, const TensorParam ¶m, const char *name, + std::initializer_list expected) const; + ge::graphStatus CheckSameShape(const CheckContext &context, const TensorParam &left, const char *leftName, + const TensorParam &right, const char *rightName) const; + int64_t GetDim(const TensorParam ¶m, size_t index) const; + bool CanOmitSequsedOriKv(const CheckContext &context) const; + bool CanOmitSequsedCmpKv(const CheckContext &context) const; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_BASE_CHECKER_SPARSE_FLASH_MLA_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/checker_adapter.h b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_adapter.h new file mode 100644 index 000000000000..b76401852176 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_adapter.h @@ -0,0 +1,100 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_CHECKER_ADAPTER_H +#define SPARSE_MLA_CHECKER_ADAPTER_H + +#include "checker_context.h" + +namespace optiling { +namespace sparse_mla_checker { + +template +void PopulateOptionalTensorParam(gert::TilingContext *context, uint32_t index, OptionalParam ¶m) +{ + param.desc = context->GetOptionalInputDesc(index); + param.shape = context->GetOptionalInputShape(index); + param.tensor = context->GetOptionalInputTensor(index); +} + +template +TensorParam MakeRequiredTensor(const RequiredParam ¶m) +{ + return {param.desc, param.shape == nullptr ? nullptr : ¶m.shape->GetStorageShape(), + param.desc != nullptr || param.shape != nullptr}; +} + +template +TensorParam MakeOptionalTensor(const OptionalParam ¶m) +{ + const gert::Shape *shape = param.shape == nullptr ? nullptr : ¶m.shape->GetStorageShape(); + if (shape == nullptr && param.tensor != nullptr) { + shape = ¶m.tensor->GetStorageShape(); + } + return {param.desc, shape, shape != nullptr}; +} + +inline int64_t GetLastDim(const TensorParam ¶m) +{ + return param.shape == nullptr || param.shape->GetDimNum() == 0 ? 0 : + param.shape->GetDim(param.shape->GetDimNum() - 1); +} + +template +void PopulateCommonContext(CheckContext &context, const TilingInfo &info) +{ + const auto ¶m = info.opParamInfo; + context.opName = info.opName; + context.q = MakeRequiredTensor(param.q); + context.oriKv = MakeOptionalTensor(param.oriKv); + context.cmpKv = MakeOptionalTensor(param.cmpKv); + context.oriSparseIndices = MakeOptionalTensor(param.oriSparseIndices); + context.cmpSparseIndices = MakeOptionalTensor(param.cmpSparseIndices); + context.oriBlockTable = MakeOptionalTensor(param.oriBlockTable); + context.cmpBlockTable = MakeOptionalTensor(param.cmpBlockTable); + context.cuSeqlensQ = MakeOptionalTensor(param.cuSeqLensQ); + context.cuSeqlensOriKv = MakeOptionalTensor(param.cuSeqLensOriKv); + context.cuSeqlensCmpKv = MakeOptionalTensor(param.cuSeqLensCmpKv); + context.sequsedQ = MakeOptionalTensor(param.seqUsedQ); + context.sequsedOriKv = MakeOptionalTensor(param.sequsedOriKv); + context.sequsedCmpKv = MakeOptionalTensor(param.sequsedCmpKv); + context.cmpResidualKv = MakeOptionalTensor(param.cmpResidualKv); + context.oriTopkLength = MakeOptionalTensor(param.oriTopkLength); + context.cmpTopkLength = MakeOptionalTensor(param.cmpTopkLength); + context.sinks = MakeOptionalTensor(param.sinks); + context.metadata = MakeOptionalTensor(param.metadata); + context.attentionOut = MakeRequiredTensor(param.attnOut); + context.softmaxLse = MakeRequiredTensor(param.softmaxLse); + + context.qLayout = static_cast(static_cast(info.qLayout)); + context.kvLayout = static_cast(static_cast(info.kvLayout)); + context.bSize = info.bSize; + context.qNumHeads = info.n1Size; + context.kvNumHeads = info.n2Size; + context.qSeqSize = info.s1Size; + context.qTotalSize = info.qTSize; + context.oriBlockSize = info.oriBlockSize; + context.cmpBlockSize = info.cmpBlockSize; + context.oriTopk = GetLastDim(context.oriSparseIndices); + context.cmpTopk = GetLastDim(context.cmpSparseIndices); + context.softmaxScale = info.softmaxScale; + context.cmpRatio = info.cmpRatio; + context.oriMaskMode = static_cast(info.oriMaskMode); + context.cmpMaskMode = static_cast(info.cmpMaskMode); + context.oriWinLeft = info.oriWinLeft; + context.oriWinRight = info.oriWinRight; + context.topkValueMode = info.topkValueMode; + context.returnSoftmaxLse = info.returnSoftmaxLse; +} + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_CHECKER_ADAPTER_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/checker_context.h b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_context.h new file mode 100644 index 000000000000..47e0a0b0d5a2 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_context.h @@ -0,0 +1,101 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_CHECKER_CONTEXT_H +#define SPARSE_MLA_CHECKER_CONTEXT_H + +#include +#include +#include "tiling/tiling_api.h" +#include "exe_graph/runtime/tiling_context.h" + +namespace optiling { +namespace sparse_mla_checker { + +enum class OperatorVariant : uint32_t { + SPARSE = 0, + MIXED_QUANT = 1, + QUANT = 2, +}; + +enum class Layout : uint32_t { + BSND = 0, + TND = 1, + PA_BBND = 2, +}; + +struct TensorParam { + const gert::CompileTimeTensorDesc *desc = nullptr; + const gert::Shape *shape = nullptr; + bool present = false; +}; + +struct CheckContext { + const char *opName = nullptr; + OperatorVariant variant = OperatorVariant::SPARSE; + + TensorParam q; + TensorParam oriKv; + TensorParam cmpKv; + TensorParam qDescale; + TensorParam oriKvDescale; + TensorParam cmpKvDescale; + TensorParam oriSparseIndices; + TensorParam cmpSparseIndices; + TensorParam oriBlockTable; + TensorParam cmpBlockTable; + TensorParam cuSeqlensQ; + TensorParam cuSeqlensOriKv; + TensorParam cuSeqlensCmpKv; + TensorParam sequsedQ; + TensorParam sequsedOriKv; + TensorParam sequsedCmpKv; + TensorParam cmpResidualKv; + TensorParam oriTopkLength; + TensorParam cmpTopkLength; + TensorParam sinks; + TensorParam metadata; + TensorParam attentionOut; + TensorParam softmaxLse; + + Layout qLayout = Layout::BSND; + Layout kvLayout = Layout::BSND; + int64_t bSize = 0; + int64_t qNumHeads = 0; + int64_t kvNumHeads = 0; + int64_t qSeqSize = 0; + int64_t qTotalSize = 0; + int64_t qHeadDim = 0; + int64_t oriKvHeadDim = 0; + int64_t cmpKvHeadDim = 0; + int64_t oriBlockSize = 0; + int64_t cmpBlockSize = 0; + int64_t oriTopk = 0; + int64_t cmpTopk = 0; + + int64_t quantMode = 0; + int64_t ropeHeadDim = 0; + float softmaxScale = 1.0F; + int64_t cmpRatio = 0; + int64_t oriMaskMode = 0; + int64_t cmpMaskMode = 0; + int64_t oriWinLeft = -1; + int64_t oriWinRight = -1; + int64_t topkValueMode = 1; + bool returnSoftmaxLse = false; + + std::vector oriKvStrides; + std::vector cmpKvStrides; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_CHECKER_CONTEXT_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/checker_runner.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_runner.cpp new file mode 100644 index 000000000000..ee053694148f --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_runner.cpp @@ -0,0 +1,63 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "checker_runner.h" +#include "common_checker_sparse_flash_mla.h" +#include "mask_checker_sparse_flash_mla.h" +#include "metadata_checker_sparse_flash_mla.h" +#include "paged_attention_checker_sparse_flash_mla.h" +#include "seq_len_checker_sparse_flash_mla.h" +#include "sinks_checker_sparse_flash_mla.h" +#include "softmax_lse_checker_sparse_flash_mla.h" +#include "sparse_compression_checker.h" + +namespace optiling { +namespace sparse_mla_checker { + +void CheckerRunner::Add(std::unique_ptr checker) +{ + checkers_.push_back(std::move(checker)); +} + +ge::graphStatus CheckerRunner::Run(CheckMethod method, const CheckContext &context) const +{ + for (const auto &checker : checkers_) { + if ((checker.get()->*method)(context) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus CheckerRunner::Process(const CheckContext &context) const +{ + if (Run(&BaseChecker::CheckSinglePara, context) != ge::GRAPH_SUCCESS || + Run(&BaseChecker::CheckParaExistence, context) != ge::GRAPH_SUCCESS || + Run(&BaseChecker::CheckFeature, context) != ge::GRAPH_SUCCESS || + Run(&BaseChecker::CheckMultiPara, context) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +void RegisterCommonCheckers(CheckerRunner &runner) +{ + runner.Add(std::make_unique()); + runner.Add(std::make_unique()); + runner.Add(std::make_unique()); + runner.Add(std::make_unique()); + runner.Add(std::make_unique()); + runner.Add(std::make_unique()); + runner.Add(std::make_unique()); + runner.Add(std::make_unique()); +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/checker_runner.h b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_runner.h new file mode 100644 index 000000000000..3eb759ec09ed --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_runner.h @@ -0,0 +1,40 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_CHECKER_RUNNER_H +#define SPARSE_MLA_CHECKER_RUNNER_H + +#include +#include +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class CheckerRunner { +public: + CheckerRunner() = default; + ~CheckerRunner() = default; + + void Add(std::unique_ptr checker); + ge::graphStatus Process(const CheckContext &context) const; + +private: + using CheckMethod = ge::graphStatus (BaseChecker::*)(const CheckContext &) const; + ge::graphStatus Run(CheckMethod method, const CheckContext &context) const; + std::vector> checkers_; +}; + +void RegisterCommonCheckers(CheckerRunner &runner); + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_CHECKER_RUNNER_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/checker_sources.cmake b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_sources.cmake new file mode 100644 index 000000000000..ecea0b8f15ca --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/checker_sources.cmake @@ -0,0 +1,30 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +set(SPARSE_MLA_COMMON_CHECKER_SRC_FILES + ${CMAKE_CURRENT_LIST_DIR}/base_checker.cpp + ${CMAKE_CURRENT_LIST_DIR}/checker_runner.cpp + ${CMAKE_CURRENT_LIST_DIR}/common_checker.cpp + ${CMAKE_CURRENT_LIST_DIR}/mask_checker.cpp + ${CMAKE_CURRENT_LIST_DIR}/metadata_checker.cpp + ${CMAKE_CURRENT_LIST_DIR}/paged_attention_checker.cpp + ${CMAKE_CURRENT_LIST_DIR}/seq_len_checker.cpp + ${CMAKE_CURRENT_LIST_DIR}/sinks_checker.cpp + ${CMAKE_CURRENT_LIST_DIR}/softmax_lse_checker.cpp + ${CMAKE_CURRENT_LIST_DIR}/sparse_compression_checker.cpp +) + +function(add_sparse_mla_common_checker_sources target_name) + get_property(common_checker_added TARGET ${target_name} PROPERTY SPARSE_MLA_COMMON_CHECKER_ADDED) + if(NOT common_checker_added) + target_sources(${target_name} PRIVATE ${SPARSE_MLA_COMMON_CHECKER_SRC_FILES}) + set_property(TARGET ${target_name} PROPERTY SPARSE_MLA_COMMON_CHECKER_ADDED TRUE) + endif() +endfunction() diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/common_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/common_checker.cpp new file mode 100644 index 000000000000..b9db1461d8c6 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/common_checker.cpp @@ -0,0 +1,259 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "common_checker_sparse_flash_mla.h" +#include +#include +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +constexpr uint32_t DIM_IDX_TWO = 2; +constexpr uint32_t DIM_IDX_THREE = 3; + +namespace { +const char *Op(const CheckContext &context) +{ + return context.opName == nullptr ? "SparseMla" : context.opName; +} +} // namespace + +ge::graphStatus CommonChecker::CheckQuery(const CheckContext &context) const +{ + if (!context.q.present) { + return ge::GRAPH_SUCCESS; + } + ge::graphStatus status = ge::GRAPH_FAILED; + if (context.variant == OperatorVariant::SPARSE) { + status = CheckTensorDesc(context, context.q, "q", {ge::DT_FLOAT16, ge::DT_BF16}); + } else if (context.variant == OperatorVariant::MIXED_QUANT) { + status = CheckTensorDesc(context, context.q, "q", {ge::DT_BF16}); + } else { + status = CheckTensorDesc(context, context.q, "q", {ge::DT_HIFLOAT8}); + } + if (status != ge::GRAPH_SUCCESS || CheckNoEmptyDim(context, context.q, "q") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + const size_t expectedDimNum = context.qLayout == Layout::BSND ? 4U : 3U; + return CheckDimNum(context, context.q, "q", {expectedDimNum}); +} + +ge::graphStatus CommonChecker::CheckKv(const CheckContext &context, const TensorParam &kv, const char *name) const +{ + if (!kv.present) { + return ge::GRAPH_SUCCESS; + } + ge::graphStatus status = ge::GRAPH_FAILED; + if (context.variant == OperatorVariant::SPARSE) { + status = CheckTensorDesc(context, kv, name, {ge::DT_FLOAT16, ge::DT_BF16}); + } else if (context.variant == OperatorVariant::MIXED_QUANT) { + status = CheckTensorDesc(context, kv, name, {ge::DT_FLOAT8_E4M3FN}); + } else { + status = CheckTensorDesc(context, kv, name, {ge::DT_HIFLOAT8}); + } + if (status != ge::GRAPH_SUCCESS || CheckNoEmptyDim(context, kv, name) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + const size_t expectedDimNum = context.kvLayout == Layout::TND ? 3U : 4U; + return CheckDimNum(context, kv, name, {expectedDimNum}); +} + +ge::graphStatus CommonChecker::CheckOutput(const CheckContext &context) const +{ + if (!context.attentionOut.present) { + return ge::GRAPH_SUCCESS; + } + ge::graphStatus status = ge::GRAPH_FAILED; + if (context.variant == OperatorVariant::SPARSE) { + status = CheckTensorDesc(context, context.attentionOut, "attention_out", {ge::DT_FLOAT16, ge::DT_BF16}); + } else { + status = CheckTensorDesc(context, context.attentionOut, "attention_out", {ge::DT_BF16}); + } + if (status != ge::GRAPH_SUCCESS || + CheckNoEmptyDim(context, context.attentionOut, "attention_out") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + const size_t expectedDimNum = context.qLayout == Layout::BSND ? 4U : 3U; + return CheckDimNum(context, context.attentionOut, "attention_out", {expectedDimNum}); +} + +ge::graphStatus CommonChecker::CheckSinglePara(const CheckContext &context) const +{ + OP_CHECK_IF( + context.qLayout != Layout::BSND && context.qLayout != Layout::TND, + OP_LOGE_FOR_INVALID_VALUE(Op(context), "layout_q", + std::to_string(static_cast(context.qLayout)).c_str(), "BSND or TND"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + context.kvLayout != Layout::BSND && context.kvLayout != Layout::TND && context.kvLayout != Layout::PA_BBND, + OP_LOGE_FOR_INVALID_VALUE(Op(context), "layout_kv", + std::to_string(static_cast(context.kvLayout)).c_str(), + "BSND, TND or PA_BBND"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + !std::isfinite(context.softmaxScale), + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + Op(context), "softmax_scale", std::to_string(context.softmaxScale).c_str(), "Softmax_scale must be finite"), + return ge::GRAPH_FAILED); + if (CheckQuery(context) != ge::GRAPH_SUCCESS || CheckKv(context, context.oriKv, "ori_kv") != ge::GRAPH_SUCCESS || + CheckKv(context, context.cmpKv, "cmp_kv") != ge::GRAPH_SUCCESS || CheckOutput(context) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus CommonChecker::CheckParaExistence(const CheckContext &context) const +{ + OP_CHECK_IF(!context.q.present, OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "q", "Q is required"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + !context.oriKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "ori_kv", "Ori_kv is required in all supported modes"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(!context.attentionOut.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "attention_out", "Attention_out is required"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus CommonChecker::CheckFeature(const CheckContext &context) const +{ + OP_CHECK_IF( + context.kvLayout != Layout::PA_BBND && context.qLayout != context.kvLayout, + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON(Op(context), "layout_q and layout_kv", + (std::to_string(static_cast(context.qLayout)) + " and " + + std::to_string(static_cast(context.kvLayout))) + .c_str(), + "Non-PA layout_q and layout_kv must be the same"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus CommonChecker::CheckQueryAxes(const CheckContext &context) const +{ + const int64_t qSeq = GetDim(context.q, 0); + int64_t qHeads = -1; + int64_t qDim = -1; + if (context.qLayout == Layout::BSND) { + OP_CHECK_IF(GetDim(context.q, 0) <= 0 || GetDim(context.q, 1) <= 0, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(Op(context), "q", + ("Batch=" + std::to_string(GetDim(context.q, 0)) + + ", sequence=" + std::to_string(GetDim(context.q, 1))) + .c_str(), + "Batch and sequence dimensions must be greater than 0"), + return ge::GRAPH_FAILED); + qHeads = GetDim(context.q, DIM_IDX_TWO); + qDim = GetDim(context.q, DIM_IDX_THREE); + } else { + OP_CHECK_IF(qSeq <= 0, OP_LOGE_FOR_INVALID_SHAPESIZE(Op(context), "q_t", std::to_string(qSeq).c_str(), "> 0"), + return ge::GRAPH_FAILED); + qHeads = GetDim(context.q, 1); + qDim = GetDim(context.q, DIM_IDX_TWO); + } + OP_CHECK_IF(qHeads <= 0 || qHeads > 128, // 128:最大支持的注意力头数 + OP_LOGE_FOR_INVALID_VALUE(Op(context), "q_n", std::to_string(qHeads).c_str(), "[1, 128]"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(qDim != 512, // 512:每个注意力头的固定维度大小 + OP_LOGE_FOR_INVALID_VALUE(Op(context), "q head dimension", std::to_string(qDim).c_str(), "512"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus CommonChecker::CheckKvAxes(const CheckContext &context, const TensorParam &kv, const char *name) const +{ + if (!kv.present) { + return ge::GRAPH_SUCCESS; + } + int64_t numHeads = -1; + int64_t headDim = -1; + if (context.kvLayout == Layout::TND) { + numHeads = GetDim(kv, 1); + headDim = GetDim(kv, DIM_IDX_TWO); + } else { + numHeads = GetDim(kv, DIM_IDX_TWO); + headDim = GetDim(kv, DIM_IDX_THREE); + } + OP_CHECK_IF(numHeads != 1, + OP_LOGE_FOR_INVALID_VALUE(Op(context), (std::string(name) + " kv_n").c_str(), + std::to_string(numHeads).c_str(), "1"), + return ge::GRAPH_FAILED); + int64_t expectedDim = 512; // 512:预期维度大小 + if (context.variant == OperatorVariant::MIXED_QUANT) { + expectedDim = context.quantMode == 1 ? 608 : 584; // 608,584:混合量化模式下根据量化模式选择不同维度 + } + OP_CHECK_IF(headDim != expectedDim, + OP_LOGE_FOR_INVALID_VALUE(Op(context), (std::string(name) + " head dimension").c_str(), + std::to_string(headDim).c_str(), std::to_string(expectedDim).c_str()), + return ge::GRAPH_FAILED); + if (context.kvLayout == Layout::PA_BBND) { + const int64_t blockSize = GetDim(kv, 1); + OP_CHECK_IF(blockSize <= 0 || blockSize > 1024, + OP_LOGE_FOR_INVALID_VALUE(Op(context), (std::string(name) + " block size").c_str(), + std::to_string(blockSize).c_str(), "[1, 1024]"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus CommonChecker::CheckMultiPara(const CheckContext &context) const +{ + if (CheckSameShape(context, context.q, "q", context.attentionOut, "attention_out") != ge::GRAPH_SUCCESS || + CheckQueryAxes(context) != ge::GRAPH_SUCCESS || + CheckKvAxes(context, context.oriKv, "ori_kv") != ge::GRAPH_SUCCESS || + CheckKvAxes(context, context.cmpKv, "cmp_kv") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(context.variant == OperatorVariant::SPARSE && + context.q.desc->GetDataType() != context.attentionOut.desc->GetDataType(), + OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON( + Op(context), "q and attention_out", + (std::to_string(static_cast(context.q.desc->GetDataType())) + " and " + + std::to_string(static_cast(context.attentionOut.desc->GetDataType()))) + .c_str(), + "Q and attention_out dtype must be the same"), + return ge::GRAPH_FAILED); + if (context.variant == OperatorVariant::SPARSE) { + OP_CHECK_IF(context.q.desc->GetDataType() != context.oriKv.desc->GetDataType(), + OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON( + Op(context), "q and ori_kv", + (std::to_string(static_cast(context.q.desc->GetDataType())) + " and " + + std::to_string(static_cast(context.oriKv.desc->GetDataType()))) + .c_str(), + "Q and ori_kv dtype must be the same"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.cmpKv.present && context.q.desc->GetDataType() != context.cmpKv.desc->GetDataType(), + OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON( + Op(context), "q and cmp_kv", + (std::to_string(static_cast(context.q.desc->GetDataType())) + " and " + + std::to_string(static_cast(context.cmpKv.desc->GetDataType()))) + .c_str(), + "Q and cmp_kv dtype must be the same"), + return ge::GRAPH_FAILED); + } + if (context.kvLayout == Layout::BSND) { + const int64_t qBatch = GetDim(context.q, 0); + OP_CHECK_IF(GetDim(context.oriKv, 0) != qBatch, + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + Op(context), "q and ori_kv batch", + (std::to_string(qBatch) + " and " + std::to_string(GetDim(context.oriKv, 0))).c_str(), + "Ori_kv batch must equal q batch"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.cmpKv.present && GetDim(context.cmpKv, 0) != qBatch, + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + Op(context), "q and cmp_kv batch", + (std::to_string(qBatch) + " and " + std::to_string(GetDim(context.cmpKv, 0))).c_str(), + "Cmp_kv batch must equal q batch"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/common_checker_sparse_flash_mla.h b/csrc/attention/sparse_flash_mla/op_host/checkers/common_checker_sparse_flash_mla.h new file mode 100644 index 000000000000..33adb09a1850 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/common_checker_sparse_flash_mla.h @@ -0,0 +1,37 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_COMMON_CHECKER_SPARSE_FLASH_MLA_H +#define SPARSE_MLA_COMMON_CHECKER_SPARSE_FLASH_MLA_H + +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class CommonChecker : public BaseChecker { +public: + ge::graphStatus CheckSinglePara(const CheckContext &context) const override; + ge::graphStatus CheckParaExistence(const CheckContext &context) const override; + ge::graphStatus CheckFeature(const CheckContext &context) const override; + ge::graphStatus CheckMultiPara(const CheckContext &context) const override; + +private: + ge::graphStatus CheckQuery(const CheckContext &context) const; + ge::graphStatus CheckKv(const CheckContext &context, const TensorParam &kv, const char *name) const; + ge::graphStatus CheckOutput(const CheckContext &context) const; + ge::graphStatus CheckQueryAxes(const CheckContext &context) const; + ge::graphStatus CheckKvAxes(const CheckContext &context, const TensorParam &kv, const char *name) const; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_COMMON_CHECKER_SPARSE_FLASH_MLA_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/mask_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/mask_checker.cpp new file mode 100644 index 000000000000..6b46fded5d5e --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/mask_checker.cpp @@ -0,0 +1,65 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "mask_checker_sparse_flash_mla.h" +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +namespace { +const char *Op(const CheckContext &context) +{ + return context.opName == nullptr ? "SparseMla" : context.opName; +} +} // namespace + +ge::graphStatus MaskChecker::CheckSinglePara(const CheckContext &context) const +{ + OP_CHECK_IF(context.oriMaskMode != 0 && context.oriMaskMode != 3 && + context.oriMaskMode != 4, // 3: RightDownCausal模式 4: Band模式 + OP_LOGE_FOR_INVALID_VALUE(Op(context), "ori_mask_mode", std::to_string(context.oriMaskMode).c_str(), + "0, 3 or 4"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + context.cmpMaskMode != 0 && context.cmpMaskMode != 3, // 3: RightDownCausal模式 + OP_LOGE_FOR_INVALID_VALUE(Op(context), "cmp_mask_mode", std::to_string(context.cmpMaskMode).c_str(), "0 or 3"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.oriWinLeft < -1 || context.oriWinRight < -1, + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + Op(context), "ori_win_left and ori_win_right", + (std::to_string(context.oriWinLeft) + ", " + std::to_string(context.oriWinRight)).c_str(), + "Both values must be -1 or non-negative"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MaskChecker::CheckFeature(const CheckContext &context) const +{ + if (context.oriMaskMode != 4) { // 4: Band模式 + OP_CHECK_IF(context.oriWinLeft != -1 || context.oriWinRight != -1, + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + Op(context), "ori_win_left and ori_win_right", + (std::to_string(context.oriWinLeft) + ", " + std::to_string(context.oriWinRight)).c_str(), + "Ori_win_left and ori_win_right must be -1 when ori_mask_mode is not 4"), + return ge::GRAPH_FAILED); + } + + if (!context.cmpKv.present) { + OP_CHECK_IF(context.cmpMaskMode != 0, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(Op(context), "cmp_mask_mode", + std::to_string(context.cmpMaskMode).c_str(), + "Cmp_mask_mode must be 0 when cmp_kv is absent"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/mask_checker_sparse_flash_mla.h b/csrc/attention/sparse_flash_mla/op_host/checkers/mask_checker_sparse_flash_mla.h new file mode 100644 index 000000000000..6eb6d5bfdc0a --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/mask_checker_sparse_flash_mla.h @@ -0,0 +1,28 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_MASK_CHECKER_SPARSE_FLASH_MLA_H +#define SPARSE_MLA_MASK_CHECKER_SPARSE_FLASH_MLA_H + +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class MaskChecker : public BaseChecker { +public: + ge::graphStatus CheckSinglePara(const CheckContext &context) const override; + ge::graphStatus CheckFeature(const CheckContext &context) const override; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_MASK_CHECKER_SPARSE_FLASH_MLA_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/metadata_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/metadata_checker.cpp new file mode 100644 index 000000000000..7d3a49efaed2 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/metadata_checker.cpp @@ -0,0 +1,45 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "metadata_checker_sparse_flash_mla.h" +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +namespace { +const char *Op(const CheckContext &context) +{ + return context.opName == nullptr ? "SparseMla" : context.opName; +} +} // namespace + +ge::graphStatus MetadataChecker::CheckSinglePara(const CheckContext &context) const +{ + if (!context.metadata.present) { + return ge::GRAPH_SUCCESS; + } + if (CheckTensorDesc(context, context.metadata, "metadata", {ge::DT_INT32}) != ge::GRAPH_SUCCESS || + CheckShape(context, context.metadata, "metadata", {1024}) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MetadataChecker::CheckParaExistence(const CheckContext &context) const +{ + OP_CHECK_IF(!context.metadata.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "metadata", "Metadata is required in the current version"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/metadata_checker_sparse_flash_mla.h b/csrc/attention/sparse_flash_mla/op_host/checkers/metadata_checker_sparse_flash_mla.h new file mode 100644 index 000000000000..a53d88adf1c2 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/metadata_checker_sparse_flash_mla.h @@ -0,0 +1,28 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_METADATA_CHECKER_SPARSE_FLASH_MLA_H +#define SPARSE_MLA_METADATA_CHECKER_SPARSE_FLASH_MLA_H + +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class MetadataChecker : public BaseChecker { +public: + ge::graphStatus CheckSinglePara(const CheckContext &context) const override; + ge::graphStatus CheckParaExistence(const CheckContext &context) const override; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_METADATA_CHECKER_SPARSE_FLASH_MLA_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/paged_attention_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/paged_attention_checker.cpp new file mode 100644 index 000000000000..9225d331a059 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/paged_attention_checker.cpp @@ -0,0 +1,101 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "paged_attention_checker_sparse_flash_mla.h" +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +namespace { +const char *Op(const CheckContext &context) +{ + return context.opName == nullptr ? "SparseMla" : context.opName; +} +} // namespace + +ge::graphStatus PagedAttentionChecker::CheckBlockTable(const CheckContext &context, const TensorParam ¶m, + const char *name) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + if (CheckTensorDesc(context, param, name, {ge::DT_INT32}) != ge::GRAPH_SUCCESS || + CheckDimNum(context, param, name, {2U}) != ge::GRAPH_SUCCESS || + CheckNoEmptyDim(context, param, name) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus PagedAttentionChecker::CheckSinglePara(const CheckContext &context) const +{ + if (CheckBlockTable(context, context.oriBlockTable, "ori_block_table") != ge::GRAPH_SUCCESS || + CheckBlockTable(context, context.cmpBlockTable, "cmp_block_table") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus PagedAttentionChecker::CheckParaExistence(const CheckContext &context) const +{ + if (context.kvLayout != Layout::PA_BBND) { + OP_CHECK_IF(context.oriBlockTable.present || context.cmpBlockTable.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "ori_block_table and cmp_block_table", + "Block tables are only supported when layout_kv is PA_BBND"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; + } + + OP_CHECK_IF(!context.oriBlockTable.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "ori_block_table", "Ori_block_table is required when layout_kv is PA_BBND"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.cmpKv.present != context.cmpBlockTable.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "cmp_block_table", + "Cmp_block_table must be present exactly when PA cmp_kv is present"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.oriBlockTable.present && !context.sequsedOriKv.present && !CanOmitSequsedOriKv(context), + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "seqused_ori_kv and ori_topk_length", + "Seqused_ori_kv is required with ori_block_table unless ORI_SPARSE or ORI_CMP_SPARSE uses " + "mask_mode=0 with ori_topk_length"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.cmpBlockTable.present && !context.sequsedCmpKv.present && !CanOmitSequsedCmpKv(context), + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "seqused_cmp_kv and cmp_topk_length", + "Seqused_cmp_kv is required with cmp_block_table unless ORI_CMP_SPARSE uses mask_mode=0 " + "with cmp_topk_length"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus PagedAttentionChecker::CheckMultiPara(const CheckContext &context) const +{ + if (context.oriBlockTable.present) { + OP_CHECK_IF(GetDim(context.oriBlockTable, 0) != context.bSize, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + Op(context), "ori_block_table", std::to_string(GetDim(context.oriBlockTable, 0)).c_str(), + ("The first dimension must equal batch size " + std::to_string(context.bSize)).c_str()), + return ge::GRAPH_FAILED); + } + if (context.cmpBlockTable.present) { + OP_CHECK_IF(GetDim(context.cmpBlockTable, 0) != context.bSize, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + Op(context), "cmp_block_table", std::to_string(GetDim(context.cmpBlockTable, 0)).c_str(), + ("The first dimension must equal batch size " + std::to_string(context.bSize)).c_str()), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/paged_attention_checker_sparse_flash_mla.h b/csrc/attention/sparse_flash_mla/op_host/checkers/paged_attention_checker_sparse_flash_mla.h new file mode 100644 index 000000000000..fa4b6a6c04db --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/paged_attention_checker_sparse_flash_mla.h @@ -0,0 +1,32 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_PAGED_ATTENTION_CHECKER_SPARSE_FLASH_MLA_H +#define SPARSE_MLA_PAGED_ATTENTION_CHECKER_SPARSE_FLASH_MLA_H + +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class PagedAttentionChecker : public BaseChecker { +public: + ge::graphStatus CheckSinglePara(const CheckContext &context) const override; + ge::graphStatus CheckParaExistence(const CheckContext &context) const override; + ge::graphStatus CheckMultiPara(const CheckContext &context) const override; + +private: + ge::graphStatus CheckBlockTable(const CheckContext &context, const TensorParam ¶m, const char *name) const; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_PAGED_ATTENTION_CHECKER_SPARSE_FLASH_MLA_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/seq_len_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/seq_len_checker.cpp new file mode 100644 index 000000000000..4b2ec0a2c902 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/seq_len_checker.cpp @@ -0,0 +1,135 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "seq_len_checker_sparse_flash_mla.h" +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +namespace { +const char *Op(const CheckContext &context) +{ + return context.opName == nullptr ? "SparseMla" : context.opName; +} +} // namespace + +ge::graphStatus SeqLenChecker::CheckLengthTensor(const CheckContext &context, const TensorParam ¶m, + const char *name) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + if (CheckTensorDesc(context, param, name, {ge::DT_INT32}) != ge::GRAPH_SUCCESS || + CheckDimNum(context, param, name, {1U}) != ge::GRAPH_SUCCESS || + CheckNoEmptyDim(context, param, name) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SeqLenChecker::CheckSinglePara(const CheckContext &context) const +{ + if (CheckLengthTensor(context, context.cuSeqlensQ, "cu_seqlens_q") != ge::GRAPH_SUCCESS || + CheckLengthTensor(context, context.cuSeqlensOriKv, "cu_seqlens_ori_kv") != ge::GRAPH_SUCCESS || + CheckLengthTensor(context, context.cuSeqlensCmpKv, "cu_seqlens_cmp_kv") != ge::GRAPH_SUCCESS || + CheckLengthTensor(context, context.sequsedQ, "seqused_q") != ge::GRAPH_SUCCESS || + CheckLengthTensor(context, context.sequsedOriKv, "seqused_ori_kv") != ge::GRAPH_SUCCESS || + CheckLengthTensor(context, context.sequsedCmpKv, "seqused_cmp_kv") != ge::GRAPH_SUCCESS || + CheckLengthTensor(context, context.cmpResidualKv, "cmp_residual_kv") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SeqLenChecker::CheckParaExistence(const CheckContext &context) const +{ + if (context.qLayout == Layout::TND) { + OP_CHECK_IF(!context.cuSeqlensQ.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "cu_seqlens_q", "Cu_seqlens_q is required when layout_q is TND"), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(context.cuSeqlensQ.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "cu_seqlens_q", "Cu_seqlens_q is only supported when layout_q is TND"), + return ge::GRAPH_FAILED); + } + + if (context.kvLayout == Layout::TND) { + OP_CHECK_IF(!context.cuSeqlensOriKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "cu_seqlens_ori_kv", + "Cu_seqlens_ori_kv is required when layout_kv is TND"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.cmpKv.present && !context.cuSeqlensCmpKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "cu_seqlens_cmp_kv", + "Cu_seqlens_cmp_kv is required when TND cmp_kv is present"), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(context.cuSeqlensOriKv.present || context.cuSeqlensCmpKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "cu_seqlens_ori_kv and cu_seqlens_cmp_kv", + "KV cu_seqlens inputs are only supported when layout_kv is TND"), + return ge::GRAPH_FAILED); + } + + if (context.kvLayout == Layout::PA_BBND) { + OP_CHECK_IF(!context.sequsedOriKv.present && !CanOmitSequsedOriKv(context), + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "seqused_ori_kv and ori_topk_length", + "Seqused_ori_kv is required for Paged Attention unless ORI_SPARSE or ORI_CMP_SPARSE " + "uses mask_mode=0 with ori_topk_length"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.cmpKv.present && !context.sequsedCmpKv.present && !CanOmitSequsedCmpKv(context), + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "seqused_cmp_kv and cmp_topk_length", + "Seqused_cmp_kv is required for Paged Attention cmp_kv unless ORI_CMP_SPARSE uses " + "mask_mode=0 with cmp_topk_length"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SeqLenChecker::CheckLength(const CheckContext &context, const TensorParam ¶m, const char *name, + int64_t expected) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(GetDim(param, 0) != expected, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + Op(context), name, std::to_string(GetDim(param, 0)).c_str(), + ("The length must be " + std::to_string(expected)).c_str()), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SeqLenChecker::CheckMultiPara(const CheckContext &context) const +{ + const int64_t batch = context.bSize; + OP_CHECK_IF(batch <= 0, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(Op(context), "batch size", std::to_string(batch).c_str(), + "Batch size must be greater than 0"), + return ge::GRAPH_FAILED); + if (CheckLength(context, context.cuSeqlensQ, "cu_seqlens_q", batch + 1) != ge::GRAPH_SUCCESS || + CheckLength(context, context.cuSeqlensOriKv, "cu_seqlens_ori_kv", batch + 1) != ge::GRAPH_SUCCESS || + CheckLength(context, context.cuSeqlensCmpKv, "cu_seqlens_cmp_kv", batch + 1) != ge::GRAPH_SUCCESS || + CheckLength(context, context.sequsedQ, "seqused_q", batch) != ge::GRAPH_SUCCESS || + CheckLength(context, context.sequsedOriKv, "seqused_ori_kv", batch) != ge::GRAPH_SUCCESS || + CheckLength(context, context.sequsedCmpKv, "seqused_cmp_kv", batch) != ge::GRAPH_SUCCESS || + CheckLength(context, context.cmpResidualKv, "cmp_residual_kv", batch) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/seq_len_checker_sparse_flash_mla.h b/csrc/attention/sparse_flash_mla/op_host/checkers/seq_len_checker_sparse_flash_mla.h new file mode 100644 index 000000000000..0a0f0f1303de --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/seq_len_checker_sparse_flash_mla.h @@ -0,0 +1,34 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_SEQ_LEN_CHECKER_SPARSE_FLASH_MLA_H +#define SPARSE_MLA_SEQ_LEN_CHECKER_SPARSE_FLASH_MLA_H + +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class SeqLenChecker : public BaseChecker { +public: + ge::graphStatus CheckSinglePara(const CheckContext &context) const override; + ge::graphStatus CheckParaExistence(const CheckContext &context) const override; + ge::graphStatus CheckMultiPara(const CheckContext &context) const override; + +private: + ge::graphStatus CheckLengthTensor(const CheckContext &context, const TensorParam ¶m, const char *name) const; + ge::graphStatus CheckLength(const CheckContext &context, const TensorParam ¶m, const char *name, + int64_t expected) const; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_SEQ_LEN_CHECKER_SPARSE_FLASH_MLA_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/sinks_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/sinks_checker.cpp new file mode 100644 index 000000000000..234776fdb1e7 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/sinks_checker.cpp @@ -0,0 +1,56 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "sinks_checker_sparse_flash_mla.h" +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +namespace { +const char *Op(const CheckContext &context) +{ + return context.opName == nullptr ? "SparseMla" : context.opName; +} +} // namespace + +ge::graphStatus SinksChecker::CheckSinglePara(const CheckContext &context) const +{ + if (!context.sinks.present) { + return ge::GRAPH_SUCCESS; + } + if (CheckTensorDesc(context, context.sinks, "sinks", {ge::DT_FLOAT}) != ge::GRAPH_SUCCESS || + CheckDimNum(context, context.sinks, "sinks", {1U}) != ge::GRAPH_SUCCESS || + CheckNoEmptyDim(context, context.sinks, "sinks") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SinksChecker::CheckParaExistence(const CheckContext &context) const +{ + OP_CHECK_IF(!context.sinks.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "sinks", "Sinks is required in the current version"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SinksChecker::CheckMultiPara(const CheckContext &context) const +{ + OP_CHECK_IF(GetDim(context.sinks, 0) != context.qNumHeads, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + Op(context), "sinks", std::to_string(GetDim(context.sinks, 0)).c_str(), + ("The length must equal q_n " + std::to_string(context.qNumHeads)).c_str()), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/sinks_checker_sparse_flash_mla.h b/csrc/attention/sparse_flash_mla/op_host/checkers/sinks_checker_sparse_flash_mla.h new file mode 100644 index 000000000000..e6ac8a102cab --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/sinks_checker_sparse_flash_mla.h @@ -0,0 +1,29 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_SINKS_CHECKER_SPARSE_FLASH_MLA_H +#define SPARSE_MLA_SINKS_CHECKER_SPARSE_FLASH_MLA_H + +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class SinksChecker : public BaseChecker { +public: + ge::graphStatus CheckSinglePara(const CheckContext &context) const override; + ge::graphStatus CheckParaExistence(const CheckContext &context) const override; + ge::graphStatus CheckMultiPara(const CheckContext &context) const override; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_SINKS_CHECKER_SPARSE_FLASH_MLA_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/softmax_lse_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/softmax_lse_checker.cpp new file mode 100644 index 000000000000..b173e4022917 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/softmax_lse_checker.cpp @@ -0,0 +1,53 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "softmax_lse_checker_sparse_flash_mla.h" +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +namespace { +const char *Op(const CheckContext &context) +{ + return context.opName == nullptr ? "SparseMla" : context.opName; +} +} // namespace + +ge::graphStatus SoftmaxLseChecker::CheckSinglePara(const CheckContext &context) const +{ + if (!context.returnSoftmaxLse || !context.softmaxLse.present) { + return ge::GRAPH_SUCCESS; + } + return CheckTensorDesc(context, context.softmaxLse, "softmax_lse", {ge::DT_FLOAT}); +} + +ge::graphStatus SoftmaxLseChecker::CheckMultiPara(const CheckContext &context) const +{ + if (!context.returnSoftmaxLse || !context.softmaxLse.present) { + return ge::GRAPH_SUCCESS; + } + + OP_CHECK_IF(context.qNumHeads % context.kvNumHeads != 0, + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + Op(context), "qNumHeads and kvNumHeads", + (std::to_string(context.qNumHeads) + ", " + std::to_string(context.kvNumHeads)).c_str(), + "QNumHeads must be divisible by kvNumHeads for softmax_lse"), + return ge::GRAPH_FAILED); + const int64_t group = context.qNumHeads / context.kvNumHeads; + if (context.qLayout == Layout::BSND) { + return CheckShape(context, context.softmaxLse, "softmax_lse", + {context.bSize, context.kvNumHeads, context.qSeqSize, group}); + } + return CheckShape(context, context.softmaxLse, "softmax_lse", + {context.kvNumHeads, context.qTotalSize, group}); +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/softmax_lse_checker_sparse_flash_mla.h b/csrc/attention/sparse_flash_mla/op_host/checkers/softmax_lse_checker_sparse_flash_mla.h new file mode 100644 index 000000000000..88a951eea514 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/softmax_lse_checker_sparse_flash_mla.h @@ -0,0 +1,28 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_SOFTMAX_LSE_CHECKER_SPARSE_FLASH_MLA_H +#define SPARSE_MLA_SOFTMAX_LSE_CHECKER_SPARSE_FLASH_MLA_H + +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class SoftmaxLseChecker : public BaseChecker { +public: + ge::graphStatus CheckSinglePara(const CheckContext &context) const override; + ge::graphStatus CheckMultiPara(const CheckContext &context) const override; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_SOFTMAX_LSE_CHECKER_SPARSE_FLASH_MLA_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_compression_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_compression_checker.cpp new file mode 100644 index 000000000000..74789e1631ec --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_compression_checker.cpp @@ -0,0 +1,201 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "sparse_compression_checker.h" +#include "log/log.h" + +namespace optiling { +namespace sparse_mla_checker { +namespace { +const char *Op(const CheckContext &context) { return context.opName == nullptr ? "SparseMla" : context.opName; } +} // namespace + +ge::graphStatus SparseCompressionChecker::CheckIndex(const CheckContext &context, const TensorParam ¶m, + const char *name) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + const size_t dimNum = context.qLayout == Layout::BSND ? 4U : 3U; + if (CheckTensorDesc(context, param, name, {ge::DT_INT32}) != ge::GRAPH_SUCCESS || + CheckDimNum(context, param, name, {dimNum}) != ge::GRAPH_SUCCESS || + CheckNoEmptyDim(context, param, name) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SparseCompressionChecker::CheckTopkLength(const CheckContext &context, const TensorParam ¶m, + const char *name) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + const size_t dimNum = context.qLayout == Layout::BSND ? 3U : 2U; + if (CheckTensorDesc(context, param, name, {ge::DT_INT32}) != ge::GRAPH_SUCCESS || + CheckDimNum(context, param, name, {dimNum}) != ge::GRAPH_SUCCESS || + CheckNoEmptyDim(context, param, name) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SparseCompressionChecker::CheckSinglePara(const CheckContext &context) const +{ + OP_CHECK_IF( + context.cmpRatio < 0 || context.cmpRatio > 128 || (context.cmpKv.present && context.cmpRatio == 0), + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(Op(context), "cmp_ratio", std::to_string(context.cmpRatio).c_str(), + "Cmp_ratio must be 0 or 1 when cmp_kv is absent, or in range [1, 128] " + "when cmp_kv is present"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + context.topkValueMode != 1, + OP_LOGE_FOR_INVALID_VALUE(Op(context), "topk_value_mode", std::to_string(context.topkValueMode).c_str(), "1"), + return ge::GRAPH_FAILED); + if (CheckIndex(context, context.oriSparseIndices, "ori_sparse_indices") != ge::GRAPH_SUCCESS || + CheckIndex(context, context.cmpSparseIndices, "cmp_sparse_indices") != ge::GRAPH_SUCCESS || + CheckTopkLength(context, context.oriTopkLength, "ori_topk_length") != ge::GRAPH_SUCCESS || + CheckTopkLength(context, context.cmpTopkLength, "cmp_topk_length") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SparseCompressionChecker::CheckParaExistence(const CheckContext &context) const +{ + if (!context.cmpKv.present) { + OP_CHECK_IF(context.cmpSparseIndices.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "cmp_sparse_indices", + "Cmp_sparse_indices requires cmp_kv"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + context.cmpTopkLength.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "cmp_topk_length", "Cmp_topk_length requires cmp_kv"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + context.cmpResidualKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "cmp_residual_kv", "Cmp_residual_kv requires cmp_kv"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + context.cmpBlockTable.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "cmp_block_table", "Cmp_block_table requires cmp_kv"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.cuSeqlensCmpKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "cu_seqlens_cmp_kv", + "Cu_seqlens_cmp_kv requires cmp_kv"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + context.sequsedCmpKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "seqused_cmp_kv", "Seqused_cmp_kv requires cmp_kv"), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + context.cmpRatio != 0 && context.cmpRatio != 1, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(Op(context), "cmp_ratio", std::to_string(context.cmpRatio).c_str(), + "Cmp_ratio must be 0 or 1 when cmp_kv is absent"), + return ge::GRAPH_FAILED); + } + + const bool needResidual = context.cmpKv.present && context.cmpMaskMode == 3 && context.cmpRatio != 1; + OP_CHECK_IF(needResidual && !context.cmpResidualKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "cmp_residual_kv", + "Cmp_residual_kv is required when cmp_mask_mode is 3 and cmp_ratio is not 1"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.oriSparseIndices.present && context.oriMaskMode == 0 && !context.oriTopkLength.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "ori_topk_length", + "Ori_topk_length is required when ori_sparse_indices is present and ori_mask_mode is 0"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context.cmpSparseIndices.present && context.cmpMaskMode == 0 && !context.cmpTopkLength.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + Op(context), "cmp_topk_length", + "Cmp_topk_length is required when cmp_sparse_indices is present and cmp_mask_mode is 0"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SparseCompressionChecker::CheckFeature(const CheckContext &context) const +{ + OP_CHECK_IF(context.cmpSparseIndices.present && !context.cmpKv.present, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(Op(context), "cmp_sparse_indices", + "Cmp_sparse_indices requires cmp_kv"), + return ge::GRAPH_FAILED); + if (context.cmpSparseIndices.present) { + OP_CHECK_IF(context.cmpTopk <= 0, + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + Op(context), "cmp_sparse_indices", std::to_string(context.cmpTopk).c_str(), + "The last dimension must be greater than 0 when cmp_sparse_indices is present"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SparseCompressionChecker::CheckIndexShape(const CheckContext &context, const TensorParam ¶m, + const char *name) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + if (context.qLayout == Layout::BSND) { + const std::string actualShape = "(" + std::to_string(GetDim(param, 0)) + ", " + + std::to_string(GetDim(param, 1)) + ", " + std::to_string(GetDim(param, 2)) + + ", " + std::to_string(GetDim(param, 3)) + ")"; + OP_CHECK_IF(GetDim(param, 0) != GetDim(context.q, 0) || GetDim(param, 1) != GetDim(context.q, 1) || + GetDim(param, 2) != 1, + OP_LOGE_FOR_INVALID_SHAPE(Op(context), name, actualShape.c_str(), "(b, q_s, kv_n, topk)"), + return ge::GRAPH_FAILED); + } else { + const std::string actualShape = "(" + std::to_string(GetDim(param, 0)) + ", " + + std::to_string(GetDim(param, 1)) + ", " + std::to_string(GetDim(param, 2)) + + ")"; + OP_CHECK_IF(GetDim(param, 0) != GetDim(context.q, 0) || GetDim(param, 1) != 1, + OP_LOGE_FOR_INVALID_SHAPE(Op(context), name, actualShape.c_str(), "(q_t, kv_n, topk)"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SparseCompressionChecker::CheckTopkLengthShape(const CheckContext &context, const TensorParam ¶m, + const char *name) const +{ + if (!param.present) { + return ge::GRAPH_SUCCESS; + } + if (context.qLayout == Layout::BSND) { + const std::string actualShape = "(" + std::to_string(GetDim(param, 0)) + ", " + + std::to_string(GetDim(param, 1)) + ", " + std::to_string(GetDim(param, 2)) + + ")"; + OP_CHECK_IF(GetDim(param, 0) != GetDim(context.q, 0) || GetDim(param, 1) != GetDim(context.q, 1) || + GetDim(param, 2) != 1, + OP_LOGE_FOR_INVALID_SHAPE(Op(context), name, actualShape.c_str(), "(b, q_s, kv_n)"), + return ge::GRAPH_FAILED); + } else { + const std::string actualShape = + "(" + std::to_string(GetDim(param, 0)) + ", " + std::to_string(GetDim(param, 1)) + ")"; + OP_CHECK_IF(GetDim(param, 0) != GetDim(context.q, 0) || GetDim(param, 1) != 1, + OP_LOGE_FOR_INVALID_SHAPE(Op(context), name, actualShape.c_str(), "(q_t, kv_n)"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SparseCompressionChecker::CheckMultiPara(const CheckContext &context) const +{ + if (CheckIndexShape(context, context.oriSparseIndices, "ori_sparse_indices") != ge::GRAPH_SUCCESS || + CheckIndexShape(context, context.cmpSparseIndices, "cmp_sparse_indices") != ge::GRAPH_SUCCESS || + CheckTopkLengthShape(context, context.oriTopkLength, "ori_topk_length") != ge::GRAPH_SUCCESS || + CheckTopkLengthShape(context, context.cmpTopkLength, "cmp_topk_length") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +} // namespace sparse_mla_checker +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_compression_checker.h b/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_compression_checker.h new file mode 100644 index 000000000000..19cd50f57ace --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_compression_checker.h @@ -0,0 +1,36 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_MLA_SPARSE_COMPRESSION_CHECKER_H +#define SPARSE_MLA_SPARSE_COMPRESSION_CHECKER_H + +#include "base_checker_sparse_flash_mla.h" + +namespace optiling { +namespace sparse_mla_checker { + +class SparseCompressionChecker : public BaseChecker { +public: + ge::graphStatus CheckSinglePara(const CheckContext &context) const override; + ge::graphStatus CheckParaExistence(const CheckContext &context) const override; + ge::graphStatus CheckFeature(const CheckContext &context) const override; + ge::graphStatus CheckMultiPara(const CheckContext &context) const override; + +private: + ge::graphStatus CheckIndex(const CheckContext &context, const TensorParam ¶m, const char *name) const; + ge::graphStatus CheckTopkLength(const CheckContext &context, const TensorParam ¶m, const char *name) const; + ge::graphStatus CheckIndexShape(const CheckContext &context, const TensorParam ¶m, const char *name) const; + ge::graphStatus CheckTopkLengthShape(const CheckContext &context, const TensorParam ¶m, const char *name) const; +}; + +} // namespace sparse_mla_checker +} // namespace optiling + +#endif // SPARSE_MLA_SPARSE_COMPRESSION_CHECKER_H diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_flash_mla_checker.cpp b/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_flash_mla_checker.cpp new file mode 100644 index 000000000000..44983846e55f --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_flash_mla_checker.cpp @@ -0,0 +1,37 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "sparse_flash_mla_checker.h" +#include "checker_adapter.h" +#include "checker_runner.h" + +namespace optiling { +namespace { +using sparse_mla_checker::CheckContext; +CheckContext BuildContext(const SMLATilingInfo &info) +{ + CheckContext context; + sparse_mla_checker::PopulateCommonContext(context, info); + context.variant = sparse_mla_checker::OperatorVariant::SPARSE; + context.qHeadDim = info.qHeadDim; + context.oriKvHeadDim = info.oriKvHeadDim; + context.cmpKvHeadDim = info.cmpKvHeadDim; + return context; +} +} // namespace + +ge::graphStatus SparseFlashMlaChecker::Process() const +{ + sparse_mla_checker::CheckerRunner runner; + sparse_mla_checker::RegisterCommonCheckers(runner); + return runner.Process(BuildContext(info_)); +} + +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_flash_mla_checker.h b/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_flash_mla_checker.h new file mode 100644 index 000000000000..e6b46041cbcb --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/checkers/sparse_flash_mla_checker.h @@ -0,0 +1,32 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef SPARSE_FLASH_MLA_CHECKER_H +#define SPARSE_FLASH_MLA_CHECKER_H + +#include "../sparse_flash_mla_tiling.h" +#include "log/error_code.h" + +namespace optiling { + +class SparseFlashMlaChecker { +public: + explicit SparseFlashMlaChecker(const SMLATilingInfo &info) + : info_(info) + {} + ge::graphStatus Process() const; + +private: + const SMLATilingInfo &info_; +}; + +} // namespace optiling + +#endif // SPARSE_FLASH_MLA_CHECKER_H diff --git a/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_def.cpp b/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_def.cpp new file mode 100644 index 000000000000..fed7e73e83d8 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_def.cpp @@ -0,0 +1,147 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! +* \file sparse_flash_mla_def.cpp +* \brief +*/ + +#include "register/op_def_registry.h" + +namespace ops { +class SparseFlashMla : public OpDef { +public: + explicit SparseFlashMla(const char *name) : OpDef(name) + { + // Aurora stores both SWA and compressed KV in BF16. + this->Input("q") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("ori_kv") + .ParamType(OPTIONAL) + .DataType({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .IgnoreContiguous(); + this->Input("cmp_kv") + .ParamType(OPTIONAL) + .DataType({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .IgnoreContiguous(); + this->Input("ori_sparse_indices") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("cmp_sparse_indices") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("ori_block_table") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("cmp_block_table") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("cu_seqlens_q") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("cu_seqlens_ori_kv") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("cu_seqlens_cmp_kv") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("seqused_q") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("seqused_ori_kv") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("seqused_cmp_kv") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("cmp_residual_kv") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("ori_topk_length") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("cmp_topk_length") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("sinks") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("metadata") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Output("attn_out") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}); + this->Output("softmax_lse") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}); + this->Attr("softmax_scale").AttrType(OPTIONAL).Float(1.0); + this->Attr("cmp_ratio").AttrType(OPTIONAL).Int(0); + this->Attr("ori_mask_mode").AttrType(OPTIONAL).Int(0); // ori_mask_mode默认值0 + this->Attr("cmp_mask_mode").AttrType(OPTIONAL).Int(0); // cmp_mask_mode默认值0 + this->Attr("ori_win_left").AttrType(OPTIONAL).Int(-1); // ori_win_left默认值-1 + this->Attr("ori_win_right").AttrType(OPTIONAL).Int(-1); // ori_win_right默认值-1 + this->Attr("layout_q").AttrType(OPTIONAL).String("BSND"); + this->Attr("layout_kv").AttrType(OPTIONAL).String("BSND"); + this->Attr("topk_value_mode").AttrType(OPTIONAL).Int(1); + this->Attr("return_softmax_lse").AttrType(OPTIONAL).Bool(false); + + OpAICoreConfig aicore_config; + aicore_config.DynamicCompileStaticFlag(true) + .DynamicFormatFlag(true) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .NeedCheckSupportFlag(false) + .PrecisionReduceFlag(true) + .ExtendCfgInfo("aclnnSupport.value", "support_aclnn"); + this->AICore().AddConfig("ascend910b", aicore_config); + this->AICore().AddConfig("ascend910_93", aicore_config); + this->AICore().AddConfig("ascend950", aicore_config); + } +}; +OP_ADD(SparseFlashMla); +} // namespace ops diff --git a/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_infershape.cpp b/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_infershape.cpp new file mode 100644 index 000000000000..4c4914344d47 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_infershape.cpp @@ -0,0 +1,134 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_infershape.cpp + * \brief + */ + +#include +#include +#include "err/ops_err.h" + +using namespace ge; + +namespace ops { +constexpr uint32_t DIM_NUM_1 = 1; +constexpr uint32_t DIM_NUM_3 = 3; +constexpr uint32_t DIM_NUM_4 = 4; +constexpr uint32_t DIM_INDEX_0 = 0; +constexpr uint32_t DIM_INDEX_1 = 1; +constexpr uint32_t DIM_INDEX_2 = 2; +constexpr uint32_t DIM_INDEX_3 = 3; +constexpr uint32_t QUERY_INPUT_INDEX = 0; +constexpr uint32_t ORI_KV_INPUT_INDEX = 1; +constexpr uint32_t CMP_KV_INPUT_INDEX = 2; +constexpr uint32_t RETURN_SOFTMAX_INDEX = 9; +constexpr uint32_t LAYOUT_KV_ATTR_INDEX = 7; + +static std::vector ToVectorFunc(const gert::Shape *shape) +{ + size_t shapeSize = shape->GetDimNum(); + std::vector shapeVec(shapeSize, 0); + + for (size_t i = 0; i < shapeSize; i++) { + shapeVec[i] = shape->GetDim(i); + } + return shapeVec; +} + +static std::string ToStringFunc(const gert::Shape *shape) +{ + std::ostringstream oss; + auto v = ToVectorFunc(shape); + if (v.size() > 0) { + for (size_t i = 0; i < v.size() - 1; ++i) { + oss << v[i] << ", "; + } + oss << v[v.size() - 1]; + } + return oss.str(); +} + +static int64_t GetKvHeadNum(const gert::Shape *kvShape, const std::string &layoutKv) +{ + if (layoutKv == "TND") { + return kvShape->GetDim(DIM_INDEX_1); + } + return kvShape->GetDim(DIM_INDEX_2); +} + +const gert::Shape *GetOptionalStorageShape(const gert::InferShapeContext *context, uint32_t inputIndex) +{ + return context->GetOptionalInputShape(inputIndex); +} + +ge::graphStatus InferShapeSparseFlashMla(gert::InferShapeContext *context) +{ + OP_CHECK_IF(context == nullptr, OP_LOGE("SparseFlashMla", "InferShapeContext is nullptr"), return ge::GRAPH_FAILED); + const gert::Shape *queryShape = context->GetInputShape(QUERY_INPUT_INDEX); + OP_CHECK_NULL_WITH_CONTEXT(context, queryShape); + const gert::Shape *oriKvShape = GetOptionalStorageShape(context, ORI_KV_INPUT_INDEX); + const gert::Shape *cmpKvShape = GetOptionalStorageShape(context, CMP_KV_INPUT_INDEX); + const gert::Shape *kvShape = (oriKvShape != nullptr) ? oriKvShape : cmpKvShape; + OP_CHECK_NULL_WITH_CONTEXT(context, kvShape); + + gert::Shape *attentionOutShape = context->GetOutputShape(0); + OP_CHECK_NULL_WITH_CONTEXT(context, attentionOutShape); + *attentionOutShape = *queryShape; + + gert::Shape *softmaxLseShape = context->GetOutputShape(1); + OP_CHECK_NULL_WITH_CONTEXT(context, softmaxLseShape); + auto attrs = context->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context, attrs); + const bool *returnSoftmaxLsePtr = attrs->GetAttrPointer(RETURN_SOFTMAX_INDEX); + const char *layoutKvPtr = attrs->GetAttrPointer(LAYOUT_KV_ATTR_INDEX); + OP_CHECK_NULL_WITH_CONTEXT(context, layoutKvPtr); + std::string layoutKv = std::string(layoutKvPtr); + bool returnSoftmaxLse = (returnSoftmaxLsePtr != nullptr) ? *returnSoftmaxLsePtr : false; + int64_t kvHeadNum = GetKvHeadNum(kvShape, layoutKv); + + OP_CHECK_IF(kvHeadNum <= 0, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + "SparseFlashMla", "ori_kv or cmp_kv", ToStringFunc(kvShape).c_str(), + "The head num of ori_kv or cmp_kv should be greater than 0 but got " + std::to_string(kvHeadNum)), + return ge::GRAPH_FAILED); + + if (returnSoftmaxLse) { + if (queryShape->GetDimNum() == DIM_NUM_3) { + softmaxLseShape->SetDimNum(DIM_NUM_3); + softmaxLseShape->SetDim(DIM_INDEX_0, kvHeadNum); + softmaxLseShape->SetDim(DIM_INDEX_1, queryShape->GetDim(DIM_INDEX_0)); + softmaxLseShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_1) / kvHeadNum); + } else { + softmaxLseShape->SetDimNum(DIM_NUM_4); + softmaxLseShape->SetDim(DIM_INDEX_0, queryShape->GetDim(DIM_INDEX_0)); + softmaxLseShape->SetDim(DIM_INDEX_1, kvHeadNum); + softmaxLseShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_1)); + softmaxLseShape->SetDim(DIM_INDEX_3, queryShape->GetDim(DIM_INDEX_2) / kvHeadNum); + } + } else { + softmaxLseShape->SetDimNum(DIM_NUM_1); + softmaxLseShape->SetDim(DIM_INDEX_0, 0); + } + return GRAPH_SUCCESS; +} + +ge::graphStatus InferDataTypeSparseFlashMla(gert::InferDataTypeContext *context) +{ + OP_CHECK_IF(context == nullptr, OP_LOGE("SparseFlashMla", "InferShapeContext is nullptr"), return ge::GRAPH_FAILED); + const auto inputDataType = context->GetInputDataType(QUERY_INPUT_INDEX); + context->SetOutputDataType(0, inputDataType); + context->SetOutputDataType(1, ge::DT_FLOAT); + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(SparseFlashMla).InferShape(InferShapeSparseFlashMla).InferDataType(InferDataTypeSparseFlashMla); +} // namespace ops diff --git a/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.cpp b/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.cpp new file mode 100644 index 000000000000..773d51c0cb44 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.cpp @@ -0,0 +1,2539 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_tiling.cpp + * \brief + */ + +#include "sparse_flash_mla_tiling.h" +#include "checkers/checker_adapter.h" +#include "checkers/sparse_flash_mla_checker.h" +#include "../op_kernel/sparse_flash_mla_template_tiling_key.h" +#include "register/op_def_registry.h" + +using namespace ge; +using namespace AscendC; +using std::map; +using std::pair; +using std::string; + +namespace optiling { + +static const std::string QUERY_NAME = "query"; +static const std::string ORI_KV_NAME = "ori_kv"; +static const std::string CMP_KV_NAME = "cmp_kv"; +static const std::string CU_SEQLENS_ORI_KV_NAME = "cu_seqlens_ori_kv"; +static const std::string CU_SEQLENS_CMP_KV_NAME = "cu_seqlens_cmp_kv"; +static const std::string ORI_SPARSE_INDICES = "ori_sparse_indices"; +static const std::string CMP_SPARSE_INDICES = "cmp_sparse_indices"; +static const std::string ORI_BLOCK_TABLE_NAME = "ori_block_table"; +static const std::string CMP_BLOCK_TABLE_NAME = "cmp_block_table"; +static const std::string SINKS_NAME = "sinks"; +static const std::string METADATA_NAME = "metadata"; +static const std::string ATTEN_OUT_NAME = "attn_out"; +static const std::string CU_SEQLENS_Q_NAME = "cu_seqlens_q"; +static const std::string SEQUSED_Q_NAME = "seqused_q"; +static const std::string SEQUSED_ORI_KV_NAME = "seqused_ori_kv"; +static const std::string SEQUSED_CMP_KV_NAME = "seqused_cmp_kv"; +static const std::string CMP_RESIDUAL_KV_NAME = "cmp_residual_kv"; +static const std::string ORI_TOPK_LENGTH_NAME = "ori_topk_length"; +static const std::string CMP_TOPK_LENGTH_NAME = "cmp_topk_length"; +static const std::string A2_A3_PLATFORM_LOG = "A2/A3"; +static const std::string A5_PLATFORM_LOG = "A5"; +constexpr uint32_t FD_MAX_S2_SPLIT_NUM = 2U; +constexpr uint32_t BATCH_CONSISTENCY_MAX_REDUCE_BLOCK_NUM = 33U; +constexpr uint32_t FD_BROADCAST_ELEMS = 8U; +constexpr int64_t BATCH_CONSISTENCY_LEVEL = 3; + +static std::vector ToVector(const gert::Shape &shape) +{ + size_t shapeSize = shape.GetDimNum(); + std::vector shapeVec(shapeSize, 0); + + for (size_t i = 0; i < shapeSize; i++) { + shapeVec[i] = shape.GetDim(i); + } + return shapeVec; +} + +static std::string ToStringRaw(const gert::Shape &shape) +{ + std::ostringstream oss; + auto v = ToVector(shape); + if (v.size() > 0) { + for (size_t i = 0; i < v.size() - 1; ++i) { + oss << v[i] << ", "; + } + oss << v[v.size() - 1]; + } + return oss.str(); +} + +static bool IsNonEmptyOptionalTensor(const gert::Tensor *tensor) +{ + return tensor != nullptr && tensor->GetShapeSize() > 0; +} + +static bool IsPowerOfTwoInRange(uint32_t value, uint32_t minValue, uint32_t maxValue) +{ + return value >= minValue && value <= maxValue && (value & (value - 1U)) == 0U; +} + +static bool IsA5Arch(NpuArch npuArch) +{ + return npuArch == NpuArch::DAV_3510; +} + +static bool IsPaBlockSizeSupport(NpuArch npuArch, int32_t blockSize) +{ + if (IsA5Arch(npuArch)) { + return blockSize >= 1 && blockSize <= static_cast(BLOCK_SIZE_LIMIT); + } + return blockSize >= 16U && blockSize <= static_cast(BLOCK_SIZE_LIMIT) && blockSize % 16U == 0; +} + +static const std::map> DTYPE_SUPPORT_MAP = { + {QUERY_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, + {ORI_KV_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, + {CMP_KV_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, + {CU_SEQLENS_ORI_KV_NAME, {ge::DT_INT32}}, + {CU_SEQLENS_CMP_KV_NAME, {ge::DT_INT32}}, + {ORI_SPARSE_INDICES, {ge::DT_INT32}}, + {CMP_SPARSE_INDICES, {ge::DT_INT32}}, + {ATTEN_OUT_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, + {ORI_BLOCK_TABLE_NAME, {ge::DT_INT32}}, + {CMP_BLOCK_TABLE_NAME, {ge::DT_INT32}}, + {SINKS_NAME, {ge::DT_FLOAT}}, + {METADATA_NAME, {ge::DT_INT32}}, + {CU_SEQLENS_Q_NAME, {ge::DT_INT32}}, + {SEQUSED_Q_NAME, {ge::DT_INT32}}, + {SEQUSED_ORI_KV_NAME, {ge::DT_INT32}}, + {SEQUSED_CMP_KV_NAME, {ge::DT_INT32}}, + {CMP_RESIDUAL_KV_NAME, {ge::DT_INT32}}, + {ORI_TOPK_LENGTH_NAME, {ge::DT_INT32}}, + {CMP_TOPK_LENGTH_NAME, {ge::DT_INT32}}}; + +static const std::map> LAYOUT_SUPPORT_MAP = { + {QUERY_NAME, {SMLALayout::BSND, SMLALayout::TND}}, + {ORI_KV_NAME, {SMLALayout::PA_BBND, SMLALayout::TND, SMLALayout::BSND}}, + {CMP_KV_NAME, {SMLALayout::PA_BBND, SMLALayout::TND, SMLALayout::BSND}}, + {ATTEN_OUT_NAME, {SMLALayout::BSND, SMLALayout::TND}}, + {ORI_SPARSE_INDICES, {SMLALayout::BSND, SMLALayout::TND}}, + {CMP_SPARSE_INDICES, {SMLALayout::BSND, SMLALayout::TND}}, +}; + +static const std::map DATATYPE_TO_STRING_MAP = { + {ge::DT_UNDEFINED, "DT_UNDEFINED"}, // Used to indicate a DataType field has not been set. + {ge::DT_FLOAT, "DT_FLOAT"}, // float type + {ge::DT_FLOAT16, "DT_FLOAT16"}, // fp16 type + {ge::DT_INT8, "DT_INT8"}, // int8 type + {ge::DT_INT16, "DT_INT16"}, // int16 type + {ge::DT_UINT16, "DT_UINT16"}, // uint16 type + {ge::DT_UINT8, "DT_UINT8"}, // uint8 type + {ge::DT_INT32, "DT_INT32"}, // uint32 type + {ge::DT_INT64, "DT_INT64"}, // int64 type + {ge::DT_UINT32, "DT_UINT32"}, // unsigned int32 + {ge::DT_UINT64, "DT_UINT64"}, // unsigned int64 + {ge::DT_BOOL, "DT_BOOL"}, // bool type + {ge::DT_DOUBLE, "DT_DOUBLE"}, // double type + {ge::DT_DUAL, "DT_DUAL"}, // dual output type + {ge::DT_DUAL_SUB_INT8, "DT_DUAL_SUB_INT8"}, // dual output int8 type + {ge::DT_DUAL_SUB_UINT8, "DT_DUAL_SUB_UINT8"}, // dual output uint8 type + {ge::DT_COMPLEX32, "DT_COMPLEX32"}, // complex32 type + {ge::DT_COMPLEX64, "DT_COMPLEX64"}, // complex64 type + {ge::DT_COMPLEX128, "DT_COMPLEX128"}, // complex128 type + {ge::DT_QINT8, "DT_QINT8"}, // qint8 type + {ge::DT_QINT16, "DT_QINT16"}, // qint16 type + {ge::DT_QINT32, "DT_QINT32"}, // qint32 type + {ge::DT_QUINT8, "DT_QUINT8"}, // quint8 type + {ge::DT_QUINT16, "DT_QUINT16"}, // quint16 type + {ge::DT_RESOURCE, "DT_RESOURCE"}, // resource type + {ge::DT_STRING_REF, "DT_STRING_REF"}, // string ref type + {ge::DT_STRING, "DT_STRING"}, // string type + {ge::DT_VARIANT, "DT_VARIANT"}, // dt_variant type + {ge::DT_BF16, "DT_BFLOAT16"}, // dt_bfloat16 type + {ge::DT_INT4, "DT_INT4"}, // dt_variant type + {ge::DT_UINT1, "DT_UINT1"}, // dt_variant type + {ge::DT_INT2, "DT_INT2"}, // dt_variant type + {ge::DT_UINT2, "DT_UINT2"} // dt_variant type +}; + +static uint64_t GetStorageShapeStride0(const gert::Shape &storageShape) +{ + if (storageShape.GetDimNum() <= DIM_NUM_ONE) { + return 0ULL; + } + + uint64_t stride0 = 1ULL; + for (size_t i = 1; i < storageShape.GetDimNum(); ++i) { + int64_t dim = storageShape.GetDim(i); + if (dim <= 0) { + return 0ULL; + } + stride0 *= static_cast(dim); + } + return stride0; +} + +template +static auto GetStride0FromStrideObject(const StrideT &stride, int) -> decltype(stride.GetDimNum(), stride.GetStride(0), + uint64_t()) +{ + if (stride.GetDimNum() <= 0) { + return 0ULL; + } + int64_t stride0 = stride.GetStride(0); + return stride0 > 0 ? static_cast(stride0) : 0ULL; +} + +template +static uint64_t GetStride0FromStrideObject(const StrideT &, ...) +{ + return 0ULL; +} + +template +static auto GetStride0FromStrideScalar(const StrideT &stride, int) -> decltype(stride > 0, + static_cast(stride)) +{ + return stride > 0 ? static_cast(stride) : 0ULL; +} + +template +static uint64_t GetStride0FromStrideScalar(const StrideT &, ...) +{ + return 0ULL; +} + +template +static uint64_t GetStride0FromStrideElement(const StrideT &stride) +{ + // CANN stride APIs return a dimension-wise stride array. In newer headers, stride[0] is scalar stride0. + // In compatibility headers it may be a stride object. Non-positive stride is treated as unavailable and + // falls back to the storage-shape contiguous calculation. + uint64_t stride0 = GetStride0FromStrideScalar(stride, 0); + if (stride0 > 0) { + return stride0; + } + return GetStride0FromStrideObject(stride, 0); +} + +template +static uint64_t GetStride0FromStrideArray(const StrideT *stride) +{ + if (stride == nullptr) { + return 0ULL; + } + return GetStride0FromStrideElement(stride[0]); +} + +template +static auto TryGetOptionalInputStride0(ContextT *context, uint32_t inputIndex, + int) -> decltype(context->GetOptionalInputStride(inputIndex), uint64_t()) +{ + return GetStride0FromStrideArray(context->GetOptionalInputStride(inputIndex)); +} + +template +static uint64_t TryGetOptionalInputStride0(ContextT *, uint32_t, ...) +{ + return 0ULL; +} + +// Compatibility path for CANN headers that do not expose GetOptionalInputStride. +// Some tiling contexts only provide real stride for view inputs through InputIsView/GetInputStride. +// Returning 0 means the stride is unavailable; the caller then falls back to storage-shape contiguous stride. +template +static auto TryGetInputViewStride0(ContextT *context, uint32_t inputIndex, + int) -> decltype(context->InputIsView(inputIndex), + context->GetInputStride(inputIndex), uint64_t()) +{ + if (!context->InputIsView(inputIndex)) { + return 0ULL; + } + return GetStride0FromStrideArray(context->GetInputStride(inputIndex)); +} + +template +static uint64_t TryGetInputViewStride0(ContextT *, uint32_t, ...) +{ + return 0ULL; +} + +std::string SMLALayoutToSerialString(SMLALayout layout) +{ + switch (layout) { + case SMLALayout::BSND: + return "BSND"; + case SMLALayout::TND: + return "TND"; + case SMLALayout::PA_BBND: + return "PA_BBND"; + default: + return "UNKNOWN"; + } +} + +struct SMLACompileInfo { + int64_t core_num; +}; + +static const std::map> SMLA_LAYOUT_AXIS_MAP = { + {SMLALayout::BSND, {SMLAAxis::B, SMLAAxis::S, SMLAAxis::N, SMLAAxis::D}}, + {SMLALayout::TND, {SMLAAxis::T, SMLAAxis::N, SMLAAxis::D}}, + {SMLALayout::PA_BBND, {SMLAAxis::Bn, SMLAAxis::Bs, SMLAAxis::N, SMLAAxis::D}}, +}; + +static const std::map SMLA_LAYOUT_DIM_MAP = { + {SMLALayout::BSND, DIM_NUM_FOUR}, + {SMLALayout::TND, DIM_NUM_THREE}, + {SMLALayout::PA_BBND, DIM_NUM_FOUR}, +}; + +static std::string SMLADataTypeToSerialString(ge::DataType type) +{ + const auto it = DATATYPE_TO_STRING_MAP.find(type); + if (it != DATATYPE_TO_STRING_MAP.end()) { + return it->second; + } else { + OP_LOGE("sparseFlashMla", "datatype %d not support", type); + return "UNDEFINED"; + } +} + +// --------------------------SMLAInfoParser类成员函数定义------------------------------------ +ge::graphStatus SMLAInfoParser::CheckRequiredInOutExistence() const +{ + OP_CHECK_IF(opParamInfo_.q.shape == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "query", "The shape of query is nullptr"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(opParamInfo_.q.desc == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "query", "The desc of query is nullptr"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(opParamInfo_.oriKv.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "ori_kv", "The tensor of ori_kv is nullptr"), + return ge::GRAPH_FAILED); + if (std::string(opParamInfo_.layoutKv) == "PA_BBND") { + OP_CHECK_IF( + opParamInfo_.oriBlockTable.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "ori_block_table", "The tensor of ori_block_table is nullptr when layoutKv is PA_BBND"), + return ge::GRAPH_FAILED); + } + if (perfMode_ == SMLATemplateMode::HCA_TEMPLATE_MODE) { + OP_CHECK_IF(opParamInfo_.cmpKv.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "cmp_kv", "The tensor of cmp_kv is nullptr"), + return ge::GRAPH_FAILED); + } + if (perfMode_ == SMLATemplateMode::CSA_TEMPLATE_MODE) { + OP_CHECK_IF(opParamInfo_.cmpKv.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "cmp_kv", "The tensor of cmp_kv is nullptr"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(opParamInfo_.cmpSparseIndices.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "cmp_sparse_indices", + "The tensor of cmp_sparse_indices is nullptr"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::CheckRequiredAttrExistence() const +{ + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::CheckRequiredParaExistence() const +{ + if (CheckRequiredInOutExistence() != ge::GRAPH_SUCCESS || CheckRequiredAttrExistence() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::CheckUnrequiredParaExistence() const +{ + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetOpName() +{ + if (context_->GetNodeName() == nullptr) { + OP_LOGE("SparseFlashMla", "opName got from TilingContext is nullptr"); + return ge::GRAPH_FAILED; + } + opName_ = context_->GetNodeName(); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetNpuInfo() +{ + platformInfo_ = context_->GetPlatformInfo(); + OP_CHECK_IF(platformInfo_ == nullptr, OP_LOGE(opName_, "GetPlatformInfo is nullptr."), return ge::GRAPH_FAILED); + + auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo_); + aivNum_ = ascendcPlatform.GetCoreNumAiv(); + aicNum_ = ascendcPlatform.GetCoreNumAic(); + OP_CHECK_IF(aicNum_ == 0 || aivNum_ == 0, OP_LOGE(opName_, "num of core obtained is 0."), return ge::GRAPH_FAILED); + + npuArch_ = ascendcPlatform.GetCurNpuArch(); + if (npuArch_ != NpuArch::DAV_2201 && npuArch_ != NpuArch::DAV_3510) { + OP_LOGE(opName_, "Npu Arch Version[%d] is not support.", static_cast(npuArch_)); + return ge::GRAPH_FAILED; + } + batchConsistency_ = (context_->GetDeterministicLevel() == BATCH_CONSISTENCY_LEVEL); + OP_LOGD(opName_, "deterministic_level=%d", context_->GetDeterministicLevel()); + + return ge::GRAPH_SUCCESS; +} + +void SMLAInfoParser::GetOptionalInputParaInfo() +{ + sparse_mla_checker::PopulateOptionalTensorParam(context_, ORI_KV_INDEX, opParamInfo_.oriKv); + sparse_mla_checker::PopulateOptionalTensorParam(context_, CMP_KV_INDEX, opParamInfo_.cmpKv); + sparse_mla_checker::PopulateOptionalTensorParam(context_, ORI_SPARSE_INDICES_INDEX, opParamInfo_.oriSparseIndices); + sparse_mla_checker::PopulateOptionalTensorParam(context_, CMP_SPARSE_INDICES_INDEX, opParamInfo_.cmpSparseIndices); + sparse_mla_checker::PopulateOptionalTensorParam(context_, ORI_BLOCK_TABLE_INDEX, opParamInfo_.oriBlockTable); + sparse_mla_checker::PopulateOptionalTensorParam(context_, CMP_BLOCK_TABLE_INDEX, opParamInfo_.cmpBlockTable); + sparse_mla_checker::PopulateOptionalTensorParam(context_, SINKS_INDEX, opParamInfo_.sinks); + sparse_mla_checker::PopulateOptionalTensorParam(context_, CU_SEQLENS_Q_INDEX, opParamInfo_.cuSeqLensQ); + sparse_mla_checker::PopulateOptionalTensorParam(context_, CU_SEQLENS_ORI_KV_INDEX, opParamInfo_.cuSeqLensOriKv); + sparse_mla_checker::PopulateOptionalTensorParam(context_, CU_SEQLENS_CMP_KV_INDEX, opParamInfo_.cuSeqLensCmpKv); + sparse_mla_checker::PopulateOptionalTensorParam(context_, SEQUSED_Q_INDEX, opParamInfo_.seqUsedQ); + sparse_mla_checker::PopulateOptionalTensorParam(context_, SEQUSED_ORI_KV_INDEX, opParamInfo_.sequsedOriKv); + sparse_mla_checker::PopulateOptionalTensorParam(context_, SEQUSED_CMP_KV_INDEX, opParamInfo_.sequsedCmpKv); + sparse_mla_checker::PopulateOptionalTensorParam(context_, CMP_RESIDUAL_KV_INDEX, opParamInfo_.cmpResidualKv); + sparse_mla_checker::PopulateOptionalTensorParam(context_, ORI_TOPK_LENGTH_INDEX, opParamInfo_.oriTopkLength); + sparse_mla_checker::PopulateOptionalTensorParam(context_, CMP_TOPK_LENGTH_INDEX, opParamInfo_.cmpTopkLength); + sparse_mla_checker::PopulateOptionalTensorParam(context_, METADATA_INDEX, opParamInfo_.metadata); +} + +void SMLAInfoParser::GetInputParaInfo() +{ + opParamInfo_.q.desc = context_->GetInputDesc(Q_INDEX); + opParamInfo_.q.shape = context_->GetInputShape(Q_INDEX); + GetOptionalInputParaInfo(); +} + +void SMLAInfoParser::GetOutputParaInfo() +{ + opParamInfo_.attnOut.desc = context_->GetOutputDesc(ATTN_OUT_INDEX); + opParamInfo_.attnOut.shape = context_->GetOutputShape(ATTN_OUT_INDEX); + opParamInfo_.softmaxLse.desc = context_->GetOutputDesc(SOFTMAX_LSE_INDEX); + opParamInfo_.softmaxLse.shape = context_->GetOutputShape(SOFTMAX_LSE_INDEX); +} + +ge::graphStatus SMLAInfoParser::GetAttrParaInfo() +{ + auto attrs = context_->GetAttrs(); + OP_CHECK_IF(attrs == nullptr, OPS_REPORT_VECTOR_INNER_ERR(context_->GetNodeName(), "attrs got from ge is nullptr"), + return ge::GRAPH_FAILED); + OP_LOGI(context_->GetNodeName(), "GetAttrParaInfo start"); + opParamInfo_.softmaxScale = attrs->GetAttrPointer(ATTR_SOFTMAX_SCALE_INDEX); + opParamInfo_.cmpRatio = attrs->GetAttrPointer(ATTR_CMP_RATIO_INDEX); + opParamInfo_.oriMaskMode = attrs->GetAttrPointer(ATTR_ORI_MASK_MODE_INDEX); + opParamInfo_.cmpMaskMode = attrs->GetAttrPointer(ATTR_CMP_MASK_MODE_INDEX); + opParamInfo_.oriWinLeft = attrs->GetAttrPointer(ATTR_ORI_WIN_LEFT_INDEX); + opParamInfo_.oriWinRight = attrs->GetAttrPointer(ATTR_ORI_WIN_RIGHT_INDEX); + opParamInfo_.layoutQ = attrs->GetStr(ATTR_LAYOUT_Q_INDEX); + opParamInfo_.layoutKv = attrs->GetStr(ATTR_LAYOUT_KV_INDEX); + opParamInfo_.topkValueMode = attrs->GetAttrPointer(ATTR_TOPK_VALUE_MODE_INDEX); + opParamInfo_.returnSoftmaxLse = attrs->GetAttrPointer(ATTR_RETURN_SOFTMAX_LSE_INDEX); + + auto oriKeyStrides = context_->GetDynamicInputStride(ORI_KV_INDEX, 0); + if (oriKeyStrides != nullptr && oriKeyStrides->GetDimNum() > 0) { + for (size_t i = 0; i < oriKeyStrides->GetDimNum(); i++) { + oriKeyStridesVec_.push_back(oriKeyStrides->GetStride(i)); + } + } + auto cmpKeyStrides = context_->GetDynamicInputStride(CMP_KV_INDEX, 0); + if (cmpKeyStrides != nullptr && cmpKeyStrides->GetDimNum() > 0) { + for (size_t i = 0; i < cmpKeyStrides->GetDimNum(); i++) { + cmpKeyStridesVec_.push_back(cmpKeyStrides->GetStride(i)); + } + } + + OP_LOGI(context_->GetNodeName(), "GetAttrParaInfo end"); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetOpParaInfo() +{ + GetInputParaInfo(); + GetOutputParaInfo(); + if (ge::GRAPH_SUCCESS != GetAttrParaInfo()) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +uint64_t SMLAInfoParser::GetOptionalInputStride0(uint32_t inputIndex) const +{ + const gert::Tensor *inputTensor = nullptr; + if (inputIndex == ORI_KV_INDEX) { + inputTensor = opParamInfo_.oriKv.tensor; + } else if (inputIndex == CMP_KV_INDEX) { + inputTensor = opParamInfo_.cmpKv.tensor; + } + if (inputTensor == nullptr) { + return 0ULL; + } + + uint64_t stride0 = TryGetOptionalInputStride0(context_, inputIndex, 0); + if (stride0 > 0) { + return stride0; + } + + // Compatible with CANN packages that only expose view stride by normal input index. + stride0 = TryGetInputViewStride0(context_, inputIndex, 0); + if (stride0 > 0) { + return stride0; + } + + const gert::Shape &storageShape = inputTensor->GetStorageShape(); + stride0 = GetStorageShapeStride0(storageShape); + const char *inputName = inputIndex == ORI_KV_INDEX ? "ori_kv" : "cmp_kv"; + OP_LOGW(context_->GetNodeName(), + "Cannot get %s stride0 from tiling context stride APIs. Use storage shape to infer contiguous " + "stride0(%lu). Non-contiguous %s requires GetOptionalInputStride or GetInputStride support.", + inputName, stride0, inputName); + return stride0; +} +ge::graphStatus SMLAInfoParser::GetInOutDataType() +{ + qType_ = opParamInfo_.q.desc->GetDataType(); + outputType_ = opParamInfo_.attnOut.desc->GetDataType(); + if (opParamInfo_.oriKv.desc != nullptr) { + oriKvType_ = opParamInfo_.oriKv.desc->GetDataType(); + } + if (opParamInfo_.cmpKv.desc != nullptr) { + cmpKvType_ = opParamInfo_.cmpKv.desc->GetDataType(); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetSMLATemplateMode() +{ + if (opParamInfo_.oriKv.desc != nullptr) { + if (opParamInfo_.cmpKv.desc != nullptr && opParamInfo_.cmpSparseIndices.tensor != nullptr) { + if (opParamInfo_.oriSparseIndices.tensor != nullptr) { + perfMode_ = SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE; + } else { + perfMode_ = SMLATemplateMode::CSA_TEMPLATE_MODE; + } + } else if (opParamInfo_.cmpKv.desc != nullptr && opParamInfo_.cmpSparseIndices.tensor == nullptr) { + perfMode_ = SMLATemplateMode::HCA_TEMPLATE_MODE; + } else if (opParamInfo_.cmpKv.desc == nullptr && opParamInfo_.cmpSparseIndices.tensor == nullptr) { + if (opParamInfo_.oriSparseIndices.tensor != nullptr) { + // A2A3此处dspark用 SWA_TEMPLATE_MODE+hasOri判断;给A5留ORI_SPARSE_TEMPLATE_MODE分支 + if (IsA5Arch(npuArch_)) { + perfMode_ = SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE; + } else { + // DSpark: oriMaskMode=0 + ori_sparse_indices on SWA kernel path. + if (opParamInfo_.oriMaskMode == nullptr || *opParamInfo_.oriMaskMode != 0U) { + OP_LOGE(opName_, "SWA ori sparse (DSpark) requires oriMaskMode 0, but got %u.", + opParamInfo_.oriMaskMode != nullptr ? *opParamInfo_.oriMaskMode : UINT32_MAX); + return ge::GRAPH_FAILED; + } + perfMode_ = SMLATemplateMode::SWA_TEMPLATE_MODE; + } + } else { + perfMode_ = SMLATemplateMode::SWA_TEMPLATE_MODE; + } + } else { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "cmp_kv", "When cmp_sparse_indices is not nullptr, cmp_kv cannot be nullptr"); + return ge::GRAPH_FAILED; + } + if (perfMode_ == SMLATemplateMode::HCA_TEMPLATE_MODE || perfMode_ == SMLATemplateMode::CSA_TEMPLATE_MODE) { + if (kvLayout_ == SMLALayout::TND && opParamInfo_.cuSeqLensCmpKv.tensor == nullptr) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "cu_seqlens_cmp_kv", + "The layout_kv is" + SMLALayoutToSerialString(kvLayout_) + ", seqlens_cmp_kv must be provided"); + return ge::GRAPH_FAILED; + } + } + } else { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "ori_kv", "ori_kv is nullptr"); + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetQueryAndOutLayout() +{ + const map> layoutMap = { + {"BSND", {SMLALayout::BSND, SMLALayout::BSND}}, + {"TND", {SMLALayout::TND, SMLALayout::TND}}, + }; + std::string layout(opParamInfo_.layoutQ); + auto it = layoutMap.find(layout); + if (it != layoutMap.end()) { + qLayout_ = it->second.first; + outLayout_ = it->second.second; + oriSparseIndicesLayout_ = qLayout_; + cmpSparseIndicesLayout_ = qLayout_; + } else { + OP_LOGE_FOR_INVALID_VALUE(opName_, "layout_q", layout.c_str(), "BSND or TND"); + return ge::GRAPH_FAILED; + } + if (qLayout_ == SMLALayout::BSND) { + OP_CHECK_IF(opParamInfo_.cuSeqLensQ.tensor != nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "cu_seqlens_q", + "When layout_q is BSND, cu_seqlens_q should be null"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetKvLayout() +{ + const map layoutKVMap = { + {"PA_BBND", SMLALayout::PA_BBND}, + {"TND", SMLALayout::TND}, + {"BSND", SMLALayout::BSND}, + }; + std::string layout(opParamInfo_.layoutKv); + auto it = layoutKVMap.find(layout); + if (it != layoutKVMap.end()) { + kvLayout_ = it->second; + } else { + OP_LOGE_FOR_INVALID_VALUE(opName_, "layout_kv", layout.c_str(), "BSND, PA_BBND or TND"); + return ge::GRAPH_FAILED; + } + if (kvLayout_ != SMLALayout::PA_BBND && qLayout_ != kvLayout_) { + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + opName_, "layout_q and layout_kv", + SMLALayoutToSerialString(qLayout_) + " and " + SMLALayoutToSerialString(kvLayout_), + "Layout_q and layout_kv only support BSND/BSND, TND/TND, BSND/PA_BBND or TND/PA_BBND"); + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +// =============Parser function==================== +bool SMLAInfoParser::HasAxis(const SMLAAxis &axis, const SMLALayout &layout, const gert::Shape &shape) const +{ + const auto &layoutIt = SMLA_LAYOUT_AXIS_MAP.find(layout); + if (layoutIt == SMLA_LAYOUT_AXIS_MAP.end()) { + return false; + } + + const std::vector &axes = layoutIt->second; + const auto &axisIt = std::find(axes.begin(), axes.end(), axis); + if (axisIt == axes.end()) { + return false; + } + const auto &dimIt = SMLA_LAYOUT_DIM_MAP.find(layout); + if (dimIt == SMLA_LAYOUT_DIM_MAP.end() || dimIt->second != shape.GetDimNum()) { + return false; + } + return true; +} + +size_t SMLAInfoParser::GetAxisIdx(const SMLAAxis &axis, const SMLALayout &layout) const +{ + const std::vector &axes = SMLA_LAYOUT_AXIS_MAP.find(layout)->second; + const auto &axisIt = std::find(axes.begin(), axes.end(), axis); + return std::distance(axes.begin(), axisIt); +} + +uint32_t SMLAInfoParser::GetAxisNum(const gert::Shape &shape, const SMLAAxis &axis, const SMLALayout &layout) const +{ + return HasAxis(axis, layout, shape) ? shape.GetDim(GetAxisIdx(axis, layout)) : invalidDimValue_; +} + +void SMLAInfoParser::SetSMLAShape() +{ + qShape_ = opParamInfo_.q.shape->GetStorageShape(); + if (opParamInfo_.oriKv.tensor != nullptr) { + oriKvShape_ = opParamInfo_.oriKv.tensor->GetStorageShape(); + } else { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "q", "Q is nullptr, please check input parameters"); + } + if (opParamInfo_.cmpKv.tensor != nullptr) { + cmpKvShape_ = opParamInfo_.cmpKv.tensor->GetStorageShape(); + } + if (opParamInfo_.oriSparseIndices.tensor != nullptr) { + oriSparseIndicesShape_ = opParamInfo_.oriSparseIndices.tensor->GetStorageShape(); + hasOriSparseIndices_ = true; + oriSparseIndexWidth_ = GetAxisNum(oriSparseIndicesShape_, SMLAAxis::K, oriSparseIndicesLayout_); + } + if (perfMode_ == SMLATemplateMode::CSA_TEMPLATE_MODE || + perfMode_ == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (opParamInfo_.cmpSparseIndices.tensor != nullptr) { + cmpSparseIndicesShape_ = opParamInfo_.cmpSparseIndices.tensor->GetStorageShape(); + } else { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "cmp_sparse_indices", + "Cmp_sparse_indices is nullptr, please check input parameters"); + } + } + + if (perfMode_ == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + perfMode_ == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (opParamInfo_.oriSparseIndices.tensor != nullptr) { + oriSparseIndicesShape_ = opParamInfo_.oriSparseIndices.tensor->GetStorageShape(); + } + } +} + +// 根据layout计算期望的连续stride +std::vector SMLAInfoParser::GetKvstride(const gert::Shape &shape, const SMLALayout &layout) const +{ + std::vector expectedStrides; + if (layout == SMLALayout::BSND || layout == SMLALayout::PA_BBND) { + uint64_t dim1 = static_cast(shape.GetDim(1)); + uint64_t dim2 = static_cast(shape.GetDim(2)); + uint64_t dim3 = static_cast(shape.GetDim(3)); + expectedStrides = {dim1 * dim2 * dim3, dim2 * dim3, dim3, 1}; + } else if (layout == SMLALayout::TND) { + uint64_t dim1 = static_cast(shape.GetDim(1)); + uint64_t dim2 = static_cast(shape.GetDim(2)); + expectedStrides = {dim1 * dim2, dim2, 1}; + } + return expectedStrides; +} + +// 非连续校验:通过shape计算expected stride进行校验 +// PA_BBND时,只允许0轴非连续,其余轴必须连续 +// 非PA_BBND时,所有轴都必须连续 +ge::graphStatus SMLAInfoParser::CheckContiguous() const +{ + bool oriKeyNonContiguous = false; + bool cmpKeyNonContiguous = false; + size_t checkStartIdx = (kvLayout_ == SMLALayout::PA_BBND) ? 1 : 0; + if (opParamInfo_.oriKv.tensor != nullptr && !oriKeyStridesVec_.empty() && + opParamInfo_.oriKv.tensor->GetShapeSize() > 0) { + std::vector oriExpectedStrides = GetKvstride(oriKvShape_, kvLayout_); + OP_CHECK_IF(oriKeyStridesVec_.size() != oriExpectedStrides.size(), + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "ori_kv", + "Ori_kv strideVec size[" + std::to_string(oriKeyStridesVec_.size()) + + "] not match layout_kv expect len[" + std::to_string(oriExpectedStrides.size()) + "]"), + return ge::GRAPH_FAILED); + oriKeyNonContiguous = + static_cast(oriKeyStridesVec_[checkStartIdx]) != oriExpectedStrides[checkStartIdx]; + } + if (opParamInfo_.cmpKv.tensor != nullptr && !cmpKeyStridesVec_.empty() && + opParamInfo_.cmpKv.tensor->GetShapeSize() > 0) { + std::vector cmpExpectedStrides = GetKvstride(cmpKvShape_, kvLayout_); + OP_CHECK_IF(cmpKeyStridesVec_.size() != cmpExpectedStrides.size(), + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "cmp_kv", + "Cmp_kv strideVec size[" + std::to_string(cmpKeyStridesVec_.size()) + + "] not match kvLayout expect len[" + std::to_string(cmpExpectedStrides.size()) + "]"), + return ge::GRAPH_FAILED); + cmpKeyNonContiguous = + static_cast(cmpKeyStridesVec_[checkStartIdx]) != cmpExpectedStrides[checkStartIdx]; + } + + OP_CHECK_IF(oriKeyNonContiguous, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "ori_kv", + "Ori_kv only support non-continuous keying on the 0-axis"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(cmpKeyNonContiguous, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "cmp_kv", + "Cmp_kv only support non-continuous keying on the 0-axis"), + return ge::GRAPH_FAILED); + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetN1Size() +{ + n1Size_ = GetAxisNum(qShape_, SMLAAxis::N, qLayout_); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetN2Size() +{ + if (opParamInfo_.oriKv.tensor != nullptr) { + n2Size_ = GetAxisNum(oriKvShape_, SMLAAxis::N, kvLayout_); + } + if (opParamInfo_.cmpKv.tensor != nullptr) { + n2Size_ = GetAxisNum(cmpKvShape_, SMLAAxis::N, kvLayout_); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetGSize() +{ + if (n2Size_ != 0) { + gSize_ = n1Size_ / n2Size_; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor, SMLALayout &layout, + const std::string &name) const +{ + if ((tensor == nullptr)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, name.c_str(), + "When layout_q is " + SMLALayoutToSerialString(layout) + ", " + name + " must be provided"); + return ge::GRAPH_FAILED; + } + int64_t shapeSize = tensor->GetShapeSize(); + if (shapeSize <= 0) { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, name.c_str(), std::to_string(shapeSize).c_str(), + "The shape size of " + name + " should be greater than 0"); + return ge::GRAPH_FAILED; + } + size = static_cast(shapeSize) - 1; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetActualSeqLenQSize(uint32_t &size) +{ + return GetActualSeqLenSize(size, opParamInfo_.cuSeqLensQ.tensor, qLayout_, "cu_seqlens_q"); +} + +ge::graphStatus SMLAInfoParser::GetBatchSize() +{ + // 获取B基准 // 1、非TND: 以query的batch_size维度为基 + // 2、TND: actual_seq_lens_q必须传入, 以actual_seq_lens_q数组的长度为B轴大小 + if (qLayout_ == SMLALayout::TND) { + return GetActualSeqLenQSize(bSize_); + } else { // BSND + bSize_ = GetAxisNum(qShape_, SMLAAxis::B, qLayout_); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetQTSize() +{ + // 获取query的T基准 // 1、非TND: 以query的batch_size维度为基准 + // 2、TND: actual_seq_lens_q必须传入, 以actual_seq_lens_q数组的长度为B轴大小 + qTSize_ = (qLayout_ == SMLALayout::TND) ? GetAxisNum(qShape_, SMLAAxis::T, qLayout_) : 0; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetS1Size() +{ + // 获取S1基准 // 1、非TND: 以query的S维度为基准 + // 2、TND: actual_seq_lens_q必须传入, 以actual_seq_lens_q数组中的最大值为基准 + if (qLayout_ == SMLALayout::TND) { + s1Size_ = GetAxisNum(qShape_, SMLAAxis::T, qLayout_); + } else { // BSND + s1Size_ = GetAxisNum(qShape_, SMLAAxis::S, qLayout_); + } + if (perfMode_ == SMLATemplateMode::CSA_TEMPLATE_MODE || + perfMode_ == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (cmpSparseIndicesLayout_ == SMLALayout::TND) { + uint32_t cmpSparseIndicesT = GetAxisNum(cmpSparseIndicesShape_, SMLAAxis::T, cmpSparseIndicesLayout_); + OP_CHECK_IF( + cmpSparseIndicesT != s1Size_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_sparse_indices", Ops::Base::ToString(cmpSparseIndicesShape_), "T size check failed"), + return ge::GRAPH_FAILED); + } else { + uint32_t cmpSparseIndicesS1 = GetAxisNum(cmpSparseIndicesShape_, SMLAAxis::S, cmpSparseIndicesLayout_); + OP_CHECK_IF( + cmpSparseIndicesS1 != s1Size_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_sparse_indices", Ops::Base::ToString(cmpSparseIndicesShape_), "S1 size check failed"), + return ge::GRAPH_FAILED); + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetMaxBlockNumPerBatch() +{ + if (opParamInfo_.oriBlockTable.tensor == nullptr) { + return ge::GRAPH_SUCCESS; + } + uint32_t oriDimNum = opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDimNum(); + if (oriDimNum != DIM_NUM_TWO) { + OP_LOGE_FOR_INVALID_SHAPEDIM(opName_, "ori_block_table", std::to_string(oriDimNum).c_str(), + std::to_string(DIM_NUM_TWO).c_str()); + return ge::GRAPH_FAILED; + } + if (opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1) < 0) { + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "ori_block_table", ToStringRaw(opParamInfo_.oriBlockTable.tensor->GetStorageShape()).c_str(), + "Ori_block_table's second dimension(" + + std::to_string(opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1)) + + ") should be non-negative number"); + return ge::GRAPH_FAILED; + } + oriMaxBlockNumPerBatch_ = opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1); + + if (opParamInfo_.cmpBlockTable.tensor != nullptr) { + uint32_t cmpDimNum = opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDimNum(); + if (cmpDimNum != DIM_NUM_TWO) { + OP_LOGE_FOR_INVALID_SHAPEDIM(opName_, "cmp_block_table", std::to_string(cmpDimNum).c_str(), + std::to_string(DIM_NUM_TWO).c_str()); + return ge::GRAPH_FAILED; + } + if (qLayout_ == SMLALayout::TND || qLayout_ == SMLALayout::BSND) { + if (opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0) != bSize_) { + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_block_table", + ToStringRaw(opParamInfo_.cmpBlockTable.tensor->GetStorageShape()).c_str(), + "Cmp_block_table's first dimension(" + + std::to_string(opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0)) + + ") should be equal to query's B(" + std::to_string(bSize_) + ")"); + return ge::GRAPH_FAILED; + } + } + if (opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(1) <= 0) { + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_block_table", ToStringRaw(opParamInfo_.cmpBlockTable.tensor->GetStorageShape()).c_str(), + "cmp_block_table's second dimension(" + + std::to_string(opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(1)) + + ") should be greater than 0"); + return ge::GRAPH_FAILED; + } + cmpMaxBlockNumPerBatch_ = opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(1); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetBlockSize() +{ + oriBlockSize_ = GetAxisNum(oriKvShape_, SMLAAxis::Bs, kvLayout_); + cmpBlockSize_ = GetAxisNum(cmpKvShape_, SMLAAxis::Bs, kvLayout_); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetS2SizeForPageAttention() +{ + if (GetMaxBlockNumPerBatch() != ge::GRAPH_SUCCESS || GetBlockSize() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + s2Size_ = oriMaxBlockNumPerBatch_ * oriBlockSize_; + cmpS2Size_ = cmpMaxBlockNumPerBatch_ * cmpBlockSize_; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetS2Size() +{ + if (kvLayout_ == SMLALayout::TND) { + s2Size_ = GetAxisNum(oriKvShape_, SMLAAxis::T, kvLayout_); + cmpS2Size_ = GetAxisNum(cmpKvShape_, SMLAAxis::T, kvLayout_); + return ge::GRAPH_SUCCESS; + } else if (kvLayout_ == SMLALayout::BSND) { + s2Size_ = GetAxisNum(oriKvShape_, SMLAAxis::S, kvLayout_); + cmpS2Size_ = GetAxisNum(cmpKvShape_, SMLAAxis::S, kvLayout_); + return ge::GRAPH_SUCCESS; + } else if (kvLayout_ == SMLALayout::PA_BBND) { + // 获取S2基准PAGE_ATTENTION S2 = block_table.dim1 * block_size + return GetS2SizeForPageAttention(); + } + return ge::GRAPH_FAILED; +} + +ge::graphStatus SMLAInfoParser::GetQHeadDim() +{ + qHeadDim_ = GetAxisNum(qShape_, SMLAAxis::D, qLayout_); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetValueHeadDim() +{ + if (opParamInfo_.oriKv.tensor != nullptr) { + oriKvHeadDim_ = GetAxisNum(oriKvShape_, SMLAAxis::D, kvLayout_); + } + if (opParamInfo_.cmpKv.tensor != nullptr) { + cmpKvHeadDim_ = GetAxisNum(cmpKvShape_, SMLAAxis::D, kvLayout_); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetSparseBlockCount() +{ + if (opParamInfo_.oriSparseIndices.tensor != nullptr) { + oriSparseBlockCount_ = GetAxisNum(oriSparseIndicesShape_, SMLAAxis::K, oriSparseIndicesLayout_); + } + if (opParamInfo_.cmpSparseIndices.tensor != nullptr) { + cmpSparseBlockCount_ = GetAxisNum(cmpSparseIndicesShape_, SMLAAxis::K, cmpSparseIndicesLayout_); + } + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetSinks() +{ + if (opParamInfo_.sinks.tensor == nullptr) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "sinks", "sinks must be provided"); + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLAInfoParser::GetActualseqInfo() +{ + if (qLayout_ == SMLALayout::TND) { + if (opParamInfo_.cuSeqLensQ.tensor != nullptr) { + if (opParamInfo_.cuSeqLensQ.tensor->GetShapeSize() != bSize_ + 1) { + OP_LOGE_FOR_INVALID_SHAPESIZE(opName_, "cu_seqlens_q", + std::to_string(opParamInfo_.cuSeqLensQ.tensor->GetShapeSize()).c_str(), + std::to_string(bSize_ + 1)); + return ge::GRAPH_FAILED; + } + actualLenDimsQ_ = opParamInfo_.cuSeqLensQ.tensor->GetShapeSize() - 1; // cuSeqLensQ shape is B+1 + } else { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "cu_seqlens_q", + "When layout_q is TND, cu_seqlens_q must be provided"); + return ge::GRAPH_FAILED; + } + } else { + if (opParamInfo_.seqUsedQ.tensor != nullptr) { + actualLenDimsQ_ = opParamInfo_.seqUsedQ.tensor->GetShapeSize(); + } + } + if (kvLayout_ != SMLALayout::PA_BBND && kvLayout_ != SMLALayout::BSND && kvLayout_ != SMLALayout::TND) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "layout_kv", SMLALayoutToSerialString(kvLayout_), + "Ori_kv and cmp_kv only support PA_BBND, BSND and TND"); + return ge::GRAPH_FAILED; + } + if (opParamInfo_.sequsedOriKv.tensor != nullptr) { + actualLenDimsOriKV_ = opParamInfo_.sequsedOriKv.tensor->GetShapeSize(); + } + if (opParamInfo_.sequsedCmpKv.tensor != nullptr) { + actualLenDimsCmpKV_ = opParamInfo_.sequsedCmpKv.tensor->GetShapeSize(); + if (opParamInfo_.sequsedCmpKv.tensor->GetShapeSize() != bSize_) { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + opName_, "seqused_cmp_kv", std::to_string(opParamInfo_.sequsedCmpKv.tensor->GetShapeSize()), + "Seqused_cmp_kv's dimension should be equal to " + std::to_string(bSize_)); + return ge::GRAPH_FAILED; + } + } + if (opParamInfo_.cmpResidualKv.tensor != nullptr) { + cmpResidualKVSize_ = opParamInfo_.cmpResidualKv.tensor->GetShapeSize(); + if (opParamInfo_.cmpResidualKv.tensor->GetShapeSize() != bSize_) { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + opName_, "cmp_residual_kv", std::to_string(opParamInfo_.cmpResidualKv.tensor->GetShapeSize()), + "Cmp_residual_kv's dimension should be equal to " + std::to_string(bSize_)); + return ge::GRAPH_FAILED; + } + } + if (!IsA5Arch(npuArch_) && IsNonEmptyOptionalTensor(opParamInfo_.cmpTopkLength.tensor)) { + OP_LOGE(opName_, "cmp_topk_length is reserved and does not support non-empty tensor on %s.", + A2_A3_PLATFORM_LOG.c_str()); + return ge::GRAPH_FAILED; + } + if (kvLayout_ == SMLALayout::PA_BBND) { + if (opParamInfo_.sequsedOriKv.tensor != nullptr) { + if (qLayout_ == SMLALayout::BSND) { + if (opParamInfo_.sequsedOriKv.tensor->GetShapeSize() != bSize_) { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + opName_, "seqused_ori_kv", std::to_string(opParamInfo_.sequsedOriKv.tensor->GetShapeSize()), + "Seqused_ori_kv's dimension should be equal to " + std::to_string(bSize_)); + return ge::GRAPH_FAILED; + } + } else { + if (opParamInfo_.sequsedOriKv.tensor->GetShapeSize() != bSize_) { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + opName_, "seqused_ori_kv", std::to_string(opParamInfo_.sequsedOriKv.tensor->GetShapeSize()), + "Seqused_ori_kv's dimension should be equal to " + std::to_string(bSize_)); + return ge::GRAPH_FAILED; + } + } + actualLenDimsKV_ = opParamInfo_.sequsedOriKv.tensor->GetShapeSize(); + } else if (perfMode_ != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE && + perfMode_ != SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "seqused_ori_kv", + "Seqused_ori_kv must be provided when layout_kv is PA_BBND"); + return ge::GRAPH_FAILED; + } + } else if (kvLayout_ == SMLALayout::TND) { + } else if (kvLayout_ == SMLALayout::BSND) { + actualLenDimsKV_ = actualLenDimsOriKV_; + } else { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "layout_kv", SMLALayoutToSerialString(kvLayout_), + "Ori_kv and cmp_kv only support PA_BBND, TND and BSND"); + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +void SMLAInfoParser::GenerateInfo(SMLATilingInfo &smlaInfo) +{ + smlaInfo.opName = opName_; + smlaInfo.platformInfo = platformInfo_; + smlaInfo.opParamInfo = opParamInfo_; + smlaInfo.npuArch = npuArch_; + + smlaInfo.bSize = bSize_; + smlaInfo.n1Size = n1Size_; + smlaInfo.n2Size = n2Size_; + smlaInfo.s1Size = s1Size_; + smlaInfo.s2Size = s2Size_; + smlaInfo.cmpS2Size = cmpS2Size_; + smlaInfo.gSize = gSize_; + smlaInfo.qHeadDim = qHeadDim_; + smlaInfo.oriKvHeadDim = oriKvHeadDim_; + smlaInfo.cmpKvHeadDim = cmpKvHeadDim_; + smlaInfo.qTSize = qTSize_; + smlaInfo.oriSparseBlockCount = oriSparseBlockCount_; + smlaInfo.cmpSparseBlockCount = cmpSparseBlockCount_; + smlaInfo.sparseBlockCount = cmpSparseBlockCount_; + smlaInfo.hasOriSparseIndices = hasOriSparseIndices_; + smlaInfo.oriSparseIndexWidth = oriSparseIndexWidth_; + smlaInfo.oriWinLeft = oriWinLeft_; + smlaInfo.oriWinRight = oriWinRight_; + smlaInfo.qType = qType_; + smlaInfo.oriKvType = oriKvType_; + smlaInfo.cmpKvType = cmpKvType_; + smlaInfo.outputType = outputType_; + smlaInfo.perfMode = perfMode_; + + smlaInfo.sparseBlockSize = 1; + smlaInfo.oriBlockSize = oriBlockSize_; + smlaInfo.cmpBlockSize = cmpBlockSize_; + smlaInfo.oriMaxBlockNumPerBatch = oriMaxBlockNumPerBatch_; + smlaInfo.cmpMaxBlockNumPerBatch = cmpMaxBlockNumPerBatch_; + + smlaInfo.actualLenDimsQ = actualLenDimsQ_; + smlaInfo.actualLenDimsKV = actualLenDimsKV_; + + smlaInfo.softmaxScale = *opParamInfo_.softmaxScale; + smlaInfo.cmpRatio = *opParamInfo_.cmpRatio; + smlaInfo.oriMaskMode = *opParamInfo_.oriMaskMode; + smlaInfo.cmpMaskMode = *opParamInfo_.cmpMaskMode; + smlaInfo.oriKvStride0 = GetOptionalInputStride0(ORI_KV_INDEX); + smlaInfo.cmpKvStride0 = GetOptionalInputStride0(CMP_KV_INDEX); + smlaInfo.oriWinLeft = *opParamInfo_.oriWinLeft; + smlaInfo.oriWinRight = *opParamInfo_.oriWinRight; + + smlaInfo.topkValueMode = *opParamInfo_.topkValueMode; + + smlaInfo.qLayout = qLayout_; + smlaInfo.oriSparseIndicesLayout = oriSparseIndicesLayout_; + smlaInfo.cmpSparseIndicesLayout = cmpSparseIndicesLayout_; + smlaInfo.kvLayout = kvLayout_; + smlaInfo.outLayout = outLayout_; + smlaInfo.returnSoftmaxLse = *opParamInfo_.returnSoftmaxLse; + smlaInfo.batchConsistency = batchConsistency_; + + smlaInfo.actualLenDimsOriKV = actualLenDimsOriKV_; + smlaInfo.actualLenDimsCmpKV = actualLenDimsCmpKV_; + smlaInfo.cmpResidualKVSize = cmpResidualKVSize_; + + if (!IsA5Arch(npuArch_)) { + smlaInfo.oriKeyStride0 = 0; + smlaInfo.cmpKeyStride0 = 0; + } else { + if (!oriKeyStridesVec_.empty()) { + smlaInfo.oriKeyStride0 = static_cast(oriKeyStridesVec_[0]); + } else { + smlaInfo.oriKeyStride0 = GetKvstride(oriKvShape_, kvLayout_)[0]; + } + + if (!cmpKeyStridesVec_.empty()) { + smlaInfo.cmpKeyStride0 = static_cast(cmpKeyStridesVec_[0]); + } else { + smlaInfo.cmpKeyStride0 = GetKvstride(cmpKvShape_, kvLayout_)[0]; + } + } +} + +ge::graphStatus SMLAInfoParser::Parse(SMLATilingInfo &smlaInfo) +{ + if (context_ == nullptr) { + OP_LOGE("SparseFlashAttention", "tiling context is nullptr!"); + return ge::GRAPH_FAILED; + } + + if (ge::GRAPH_SUCCESS != GetOpName() || ge::GRAPH_SUCCESS != GetNpuInfo() || ge::GRAPH_SUCCESS != GetOpParaInfo() || + ge::GRAPH_SUCCESS != CheckRequiredParaExistence() || ge::GRAPH_SUCCESS != CheckUnrequiredParaExistence()) { + return ge::GRAPH_FAILED; + } + + if (ge::GRAPH_SUCCESS != GetInOutDataType() || ge::GRAPH_SUCCESS != GetQueryAndOutLayout() || + ge::GRAPH_SUCCESS != GetKvLayout() || ge::GRAPH_SUCCESS != GetSMLATemplateMode()) { + return ge::GRAPH_FAILED; + } + + // Match the model-facing template selection on both A2/A3 and A5. + OP_CHECK_IF(qLayout_ != SMLALayout::TND || kvLayout_ != SMLALayout::PA_BBND, + OP_LOGE(opName_, "Aurora SparseFlashMla only compiles TND Q with PA_BBND KV."), + return ge::GRAPH_FAILED); + OP_CHECK_IF(perfMode_ != SMLATemplateMode::SWA_TEMPLATE_MODE && + perfMode_ != SMLATemplateMode::CSA_TEMPLATE_MODE, + OP_LOGE(opName_, "Aurora SparseFlashMla only compiles SWA and CSA templates."), + return ge::GRAPH_FAILED); + OP_CHECK_IF(qType_ != ge::DT_BF16, + OP_LOGE(opName_, "Aurora SparseFlashMla only compiles BF16 Q/KV."), return ge::GRAPH_FAILED); + + SetSMLAShape(); + if (ge::GRAPH_SUCCESS != GetN1Size() || ge::GRAPH_SUCCESS != GetN2Size() || ge::GRAPH_SUCCESS != GetGSize() || + ge::GRAPH_SUCCESS != GetBatchSize() || ge::GRAPH_SUCCESS != GetQTSize() || ge::GRAPH_SUCCESS != GetS1Size() || + ge::GRAPH_SUCCESS != GetS2Size() || ge::GRAPH_SUCCESS != GetQHeadDim() || + ge::GRAPH_SUCCESS != GetValueHeadDim() || ge::GRAPH_SUCCESS != GetSparseBlockCount() || + ge::GRAPH_SUCCESS != GetSinks()) { + return ge::GRAPH_FAILED; + } + if (ge::GRAPH_SUCCESS != GetActualseqInfo()) { + return ge::GRAPH_FAILED; + } + if (ge::GRAPH_SUCCESS != CheckContiguous()) { + return ge::GRAPH_FAILED; + } + GenerateInfo(smlaInfo); + return ge::GRAPH_SUCCESS; +} + +void SMLATilingCheck::Init() +{ + opName_ = smlaInfo_.opName; + platformInfo_ = smlaInfo_.platformInfo; + opParamInfo_ = smlaInfo_.opParamInfo; + npuArch_ = smlaInfo_.npuArch; + bSize_ = smlaInfo_.bSize; + n1Size_ = smlaInfo_.n1Size; + n2Size_ = smlaInfo_.n2Size; + s1Size_ = smlaInfo_.s1Size; + s2Size_ = smlaInfo_.s2Size; + cmpS2Size_ = smlaInfo_.cmpS2Size; + gSize_ = smlaInfo_.gSize; + qHeadDim_ = smlaInfo_.qHeadDim; + oriKvHeadDim_ = smlaInfo_.oriKvHeadDim; + cmpKvHeadDim_ = smlaInfo_.cmpKvHeadDim; + oriBlockSize_ = smlaInfo_.oriBlockSize; + cmpBlockSize_ = smlaInfo_.cmpBlockSize; + qTSize_ = smlaInfo_.qTSize; + qType_ = smlaInfo_.qType; + oriKvType_ = smlaInfo_.oriKvType; + cmpKvType_ = smlaInfo_.cmpKvType; + outputType_ = smlaInfo_.outputType; + cmpRatio_ = smlaInfo_.cmpRatio; + qLayout_ = smlaInfo_.qLayout; + oriSparseIndicesLayout_ = smlaInfo_.oriSparseIndicesLayout; + cmpSparseIndicesLayout_ = smlaInfo_.cmpSparseIndicesLayout; + oriWinLeft_ = smlaInfo_.oriWinLeft; + oriWinRight_ = smlaInfo_.oriWinRight; + hasOriSparseIndices_ = smlaInfo_.hasOriSparseIndices; + oriSparseIndexWidth_ = smlaInfo_.oriSparseIndexWidth; + kvLayout_ = smlaInfo_.kvLayout; + outLayout_ = smlaInfo_.outLayout; + topkValueMode_ = smlaInfo_.topkValueMode; +} + +void SMLATilingCheck::LogErrorDtypeSupport(const std::vector &expectDtypeList, + const ge::DataType &actualDtype, const std::string &name) const +{ + std::ostringstream oss; + for (size_t i = 0; i < expectDtypeList.size(); ++i) { + oss << SMLADataTypeToSerialString(expectDtypeList[i]); + if (i < expectDtypeList.size() - 1) { + oss << ", "; + } + } + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, name.c_str(), SMLADataTypeToSerialString(actualDtype).c_str(), + "The dtype of " + name + " only supports " + oss.str()); +} + +ge::graphStatus SMLATilingCheck::CheckDtypeSupport(const gert::CompileTimeTensorDesc *desc, + const std::string &name) const +{ + if (desc != nullptr) { + const auto &it = DTYPE_SUPPORT_MAP.find(name); + OP_CHECK_IF(it == DTYPE_SUPPORT_MAP.end(), + OP_LOGE(opName_, "%s datatype support list should be specify in DTYPE_SUPPORT_MAP", name.c_str()), + return ge::GRAPH_FAILED); + auto &expectDtypeList = it->second; + OP_CHECK_IF( + std::find(expectDtypeList.begin(), expectDtypeList.end(), desc->GetDataType()) == expectDtypeList.end(), + LogErrorDtypeSupport(expectDtypeList, desc->GetDataType(), name), return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +void SMLATilingCheck::LogErrorLayoutSupport(const std::vector &expectLayoutList, + const SMLALayout &actualLayout, const std::string &name) const +{ + std::ostringstream oss; + for (size_t i = 0; i < expectLayoutList.size(); ++i) { + oss << SMLALayoutToSerialString(expectLayoutList[i]); + if (i < expectLayoutList.size() - 1) { + oss << ", "; + } + } + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, name.c_str(), SMLALayoutToSerialString(actualLayout).c_str(), + "Tensor " + name + " only supports layout " + oss.str()); +} + +ge::graphStatus SMLATilingCheck::CheckLayoutSupport(const SMLALayout &actualLayout, const std::string &name) const +{ + const auto &it = LAYOUT_SUPPORT_MAP.find(name); + OP_CHECK_IF(it == LAYOUT_SUPPORT_MAP.end(), + OP_LOGE(opName_, "%s layout support list should be specify in LAYOUT_SUPPORT_MAP", name.c_str()), + return ge::GRAPH_FAILED); + auto &expectLayoutList = it->second; + OP_CHECK_IF(std::find(expectLayoutList.begin(), expectLayoutList.end(), actualLayout) == expectLayoutList.end(), + LogErrorLayoutSupport(expectLayoutList, actualLayout, name), return ge::GRAPH_FAILED); + + return ge::GRAPH_SUCCESS; +} + +template +void SMLATilingCheck::LogErrorNumberSupport(const std::vector &expectNumberList, const T &actualValue, + const std::string &name, const std::string subName) const +{ + std::ostringstream oss; + for (size_t i = 0; i < expectNumberList.size(); ++i) { + oss << std::to_string(expectNumberList[i]); + if (i < expectNumberList.size() - 1) { + oss << ", "; + } + } + OP_LOGE(opName_, "%s %s only supports %s, but got %s", name.c_str(), subName.c_str(), oss.str().c_str(), + std::to_string(actualValue).c_str()); +} + +template +void SMLATilingCheck::LogErrorDimNumSupport(const std::vector &expectNumberList, const T &actualValue, + const std::string &name) const +{ + LogErrorNumberSupport(expectNumberList, actualValue, name, "dimension"); +} + +ge::graphStatus SMLATilingCheck::CheckDimNumSupport(const gert::StorageShape *shape, + const std::vector &expectDimNumList, + const std::string &name) const +{ + if (shape == nullptr) { + return ge::GRAPH_SUCCESS; + } + + if (std::find(expectDimNumList.begin(), expectDimNumList.end(), shape->GetStorageShape().GetDimNum()) == + expectDimNumList.end()) { + LogErrorDimNumSupport(expectDimNumList, shape->GetStorageShape().GetDimNum(), name); + return ge::GRAPH_FAILED; + } + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckDimNumInLayoutSupport(const SMLALayout &layout, const gert::StorageShape *shape, + const std::string &name) const +{ + const auto &dimIt = SMLA_LAYOUT_DIM_MAP.find(layout); + OP_CHECK_IF(shape->GetStorageShape().GetDimNum() != dimIt->second, + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON( + opName_, name.c_str(), std::to_string(shape->GetStorageShape().GetDimNum()).c_str(), + "When layout is " + SMLALayoutToSerialString(layout) + ", the shape dim of " + name + + " should be " + std::to_string(dimIt->second)), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaQuery() const +{ + if (opParamInfo_.q.desc == nullptr || opParamInfo_.q.shape == nullptr) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "q", "Q must be provided"); + return ge::GRAPH_FAILED; + } + const std::vector queryDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.q.desc, QUERY_NAME) || + ge::GRAPH_SUCCESS != CheckLayoutSupport(qLayout_, QUERY_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(opParamInfo_.q.shape, queryDimNumList, QUERY_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(qLayout_, opParamInfo_.q.shape, QUERY_NAME)) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaOriKv() const +{ + const std::vector oriKvDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriKv.desc, ORI_KV_NAME) || + ge::GRAPH_SUCCESS != CheckLayoutSupport(kvLayout_, ORI_KV_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriKv.tensor->GetShape(), oriKvDimNumList, ORI_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumInLayoutSupport(kvLayout_, &opParamInfo_.oriKv.tensor->GetShape(), ORI_KV_NAME)) { + return ge::GRAPH_FAILED; + } + if (kvLayout_ == SMLALayout::BSND) { + OP_CHECK_IF( + opParamInfo_.oriKv.tensor->GetStorageShape().GetDim(0) != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "ori_kv", ToStringRaw(opParamInfo_.oriKv.tensor->GetStorageShape()), + "Ori_kv's batch dimension(" + std::to_string(opParamInfo_.oriKv.tensor->GetStorageShape().GetDim(0)) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCmpKv() const +{ + if (smlaInfo_.perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::HCA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + const std::vector cmpKvDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpKv.desc, CMP_KV_NAME) || + ge::GRAPH_SUCCESS != CheckLayoutSupport(kvLayout_, CMP_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.cmpKv.tensor->GetShape(), cmpKvDimNumList, CMP_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumInLayoutSupport(kvLayout_, &opParamInfo_.cmpKv.tensor->GetShape(), CMP_KV_NAME)) { + return ge::GRAPH_FAILED; + } + if (kvLayout_ == SMLALayout::BSND) { + OP_CHECK_IF(opParamInfo_.cmpKv.tensor->GetStorageShape().GetDim(0) != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_kv", ToStringRaw(opParamInfo_.cmpKv.tensor->GetStorageShape()), + "Cmp_kv's batch dimension(" + + std::to_string(opParamInfo_.cmpKv.tensor->GetStorageShape().GetDim(0)) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCuSeqLensQ() const +{ + if (opParamInfo_.cuSeqLensQ.tensor == nullptr) { + return ge::GRAPH_SUCCESS; + } + const std::vector dimNumList = {DIM_NUM_ONE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cuSeqLensQ.desc, CU_SEQLENS_Q_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.cuSeqLensQ.tensor->GetShape(), dimNumList, CU_SEQLENS_Q_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(opParamInfo_.cuSeqLensQ.tensor->GetShapeSize() != bSize_ + 1, + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + opName_, "cu_seqlens_q", std::to_string(opParamInfo_.cuSeqLensQ.tensor->GetShapeSize()).c_str(), + "The shape size of cu_seqlens_q is not equal to B + 1:" + std::to_string(bSize_ + 1)), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCuSeqLensOriKv() const +{ + if (opParamInfo_.cuSeqLensOriKv.tensor == nullptr) { + return ge::GRAPH_SUCCESS; + } + const std::vector dimNumList = {DIM_NUM_ONE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cuSeqLensOriKv.desc, CU_SEQLENS_ORI_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.cuSeqLensOriKv.tensor->GetShape(), dimNumList, CU_SEQLENS_ORI_KV_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF( + opParamInfo_.cuSeqLensOriKv.tensor->GetShapeSize() != bSize_ + 1, + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + opName_, "cu_seqlens_ori_kv", std::to_string(opParamInfo_.cuSeqLensOriKv.tensor->GetShapeSize()).c_str(), + "The shape size of cu_seqlens_ori_kv is not equal to B + 1:" + std::to_string(bSize_ + 1)), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCuSeqLensCmpKv() const +{ + if (opParamInfo_.cuSeqLensCmpKv.tensor == nullptr) { + return ge::GRAPH_SUCCESS; + } + const std::vector dimNumList = {DIM_NUM_ONE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cuSeqLensCmpKv.desc, CU_SEQLENS_CMP_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.cuSeqLensCmpKv.tensor->GetShape(), dimNumList, CU_SEQLENS_CMP_KV_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF( + opParamInfo_.cuSeqLensCmpKv.tensor->GetShapeSize() != bSize_ + 1, + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + opName_, "cu_seqlens_cmp_kv", std::to_string(opParamInfo_.cuSeqLensCmpKv.tensor->GetShapeSize()).c_str(), + "The shape size of cu_seqlens_cmp_kv is not equal to B + 1:" + std::to_string(bSize_ + 1)), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaNumHeads() const +{ + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaKvHeadNums() const +{ + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaOriSparseIndices() const +{ + if (opParamInfo_.oriSparseIndices.tensor == nullptr) { + return ge::GRAPH_SUCCESS; + } + if (!IsA5Arch(npuArch_)) { + OP_CHECK_IF(smlaInfo_.perfMode != SMLATemplateMode::SWA_TEMPLATE_MODE, + OP_LOGE(opName_, "ori_sparse_indices is only supported in SWA mode."), return ge::GRAPH_FAILED); + OP_CHECK_IF(opParamInfo_.oriMaskMode == nullptr || *opParamInfo_.oriMaskMode != 0U, + OP_LOGE(opName_, "ori_sparse_indices SWA (DSpark) requires oriMaskMode 0."), + return ge::GRAPH_FAILED); + } + OP_CHECK_IF(opParamInfo_.oriSparseIndices.tensor->GetStorageShape().GetShapeSize() == 0, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "ori_sparse_indices", + "Ori_sparse_indices cannot be empty tensor"), + return ge::GRAPH_FAILED); + const std::vector oriSparseIndicesDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriSparseIndices.desc, ORI_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckLayoutSupport(oriSparseIndicesLayout_, ORI_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriSparseIndices.tensor->GetShape(), + oriSparseIndicesDimNumList, ORI_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(oriSparseIndicesLayout_, + &opParamInfo_.oriSparseIndices.tensor->GetShape(), + ORI_SPARSE_INDICES)) { + return ge::GRAPH_FAILED; + } + if (oriSparseIndicesLayout_ == SMLALayout::BSND) { + OP_CHECK_IF( + opParamInfo_.oriSparseIndices.tensor->GetStorageShape().GetDim(0) != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "ori_sparse_indices", ToStringRaw(opParamInfo_.oriSparseIndices.tensor->GetStorageShape()), + "Ori_sparse_indices's batch dimension(" + + std::to_string(opParamInfo_.oriSparseIndices.tensor->GetStorageShape().GetDim(0)) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCmpSparseIndices() const +{ + if (smlaInfo_.perfMode == optiling::SMLATemplateMode::CSA_TEMPLATE_MODE || + smlaInfo_.perfMode == optiling::SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + OP_CHECK_IF(opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetShapeSize() == 0, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "cmp_sparse_indices", + "When cmp_sparse_indices is not nullptr(CSA), cmp_sparse_indices cannot be empty tensor"), + return ge::GRAPH_FAILED); + const std::vector cmpSparseIndicesDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpSparseIndices.desc, CMP_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckLayoutSupport(cmpSparseIndicesLayout_, CMP_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpSparseIndices.tensor->GetShape(), + cmpSparseIndicesDimNumList, CMP_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(cmpSparseIndicesLayout_, + &opParamInfo_.cmpSparseIndices.tensor->GetShape(), + CMP_SPARSE_INDICES)) { + return ge::GRAPH_FAILED; + } + if (cmpSparseIndicesLayout_ == SMLALayout::BSND) { + OP_CHECK_IF( + opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(0) != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_sparse_indices", ToStringRaw(opParamInfo_.cmpSparseIndices.tensor->GetStorageShape()), + "Cmp_sparse_indices's batch dimension(" + + std::to_string(opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(0)) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaOriBlockTable() const +{ + if (kvLayout_ != SMLALayout::PA_BBND) { + return ge::GRAPH_SUCCESS; + } + const std::vector oriBlockTableDimNumList = {DIM_NUM_TWO}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriBlockTable.desc, ORI_BLOCK_TABLE_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriBlockTable.tensor->GetShape(), oriBlockTableDimNumList, + ORI_BLOCK_TABLE_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(0) != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "ori_block_table", ToStringRaw(opParamInfo_.oriBlockTable.tensor->GetStorageShape()), + "Ori_block_table's first dimension(" + + std::to_string(opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(0)) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(!IsPaBlockSizeSupport(npuArch_, oriBlockSize_), + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "ori_kv", ToStringRaw(opParamInfo_.oriKv.tensor->GetStorageShape()).c_str(), + "OriBlockSize_ should be in [1, 1024] on " + A5_PLATFORM_LOG + " or 16-aligned [16, 1024] on " + + A2_A3_PLATFORM_LOG + ", but got: " + std::to_string(oriBlockSize_)), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCmpBlockTable() const +{ + if (kvLayout_ != SMLALayout::PA_BBND) { + return ge::GRAPH_SUCCESS; + } + if (smlaInfo_.perfMode == optiling::SMLATemplateMode::CSA_TEMPLATE_MODE || + smlaInfo_.perfMode == optiling::SMLATemplateMode::HCA_TEMPLATE_MODE || + smlaInfo_.perfMode == optiling::SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + OP_CHECK_IF(opParamInfo_.cmpBlockTable.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "cmp_block_table", + "Cmp_block_table must be provided when layout_kv is PA_BBND in CSA/HCA/ORI_CMP_SPARSE mode"), + return ge::GRAPH_FAILED); + const std::vector cmpBlockTableDimNumList = {DIM_NUM_TWO}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpBlockTable.desc, CMP_BLOCK_TABLE_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpBlockTable.tensor->GetShape(), + cmpBlockTableDimNumList, CMP_BLOCK_TABLE_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0) != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_block_table", ToStringRaw(opParamInfo_.cmpBlockTable.tensor->GetStorageShape()), + "cmp_block_table's first dimension(" + + std::to_string(opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0)) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(!IsPaBlockSizeSupport(npuArch_, cmpBlockSize_), + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_kv", ToStringRaw(opParamInfo_.oriKv.tensor->GetStorageShape()).c_str(), + "CmpBlockSize should be in [1, 1024] on " + A5_PLATFORM_LOG + " or 16-aligned [16, 1024] on " + + A2_A3_PLATFORM_LOG + ", but got: " + std::to_string(cmpBlockSize_)), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaSinks() const +{ + OP_CHECK_IF(opParamInfo_.sinks.tensor->GetStorageShape().GetShapeSize() == 0, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "sinks", "Sinks cannot be empty tensor"), + return ge::GRAPH_FAILED); + if (opParamInfo_.sinks.tensor->GetStorageShape().GetDimNum() != DIM_NUM_ONE) { + OP_LOGE_FOR_INVALID_SHAPEDIM(opName_, SINKS_NAME.c_str(), + std::to_string(opParamInfo_.sinks.tensor->GetStorageShape().GetDimNum()).c_str(), + std::to_string(DIM_NUM_ONE).c_str()); + return ge::GRAPH_FAILED; + } + if (opParamInfo_.sinks.tensor->GetStorageShape().GetDim(0) != n1Size_) { + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "sinks", ToStringRaw(opParamInfo_.sinks.tensor->GetStorageShape()), + "Sinks's dimension(" + std::to_string(opParamInfo_.sinks.tensor->GetStorageShape().GetDim(0)) + + ") should be equal to the head num of query(" + std::to_string(n1Size_) + ")"); + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(opParamInfo_.sinks.desc->GetDataType() != ge::DT_FLOAT, + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON( + opName_, "sinks", SMLADataTypeToSerialString(opParamInfo_.sinks.desc->GetDataType()).c_str(), + "The dtype of sinks must be DT_FLOAT"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaMetadata() const +{ + if (opParamInfo_.metadata.tensor == nullptr) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "metadata", "Metadata must be provided"); + return ge::GRAPH_FAILED; + } + const std::vector dimNumList = {DIM_NUM_ONE}; + if (ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.metadata.tensor->GetShape(), dimNumList, METADATA_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF((opParamInfo_.metadata.tensor->GetShapeSize() != METADATA_LIMIT), + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(opName_, "metadata", + ToStringRaw(opParamInfo_.metadata.tensor->GetStorageShape()), + "metadata dim 0 must be" + std::to_string(METADATA_LIMIT)), + return ge::GRAPH_FAILED); + OP_CHECK_IF(opParamInfo_.metadata.desc->GetDataType() != ge::DT_INT32, + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON( + opName_, "metadata", SMLADataTypeToSerialString(opParamInfo_.metadata.desc->GetDataType()).c_str(), + "The dtype of metadata must be DT_INT32"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCmpRatio() const +{ + if (IsA5Arch(npuArch_)) { + if (opParamInfo_.cmpKv.tensor != nullptr) { + OP_CHECK_IF( + cmpRatio_ < 1 || cmpRatio_ > 128, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "cmp_ratio", std::to_string(cmpRatio_).c_str(), + "Cmp_ratio should be in range [1, 128] on " + A5_PLATFORM_LOG), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; + } + + const auto checkRatio = [this](bool isSupported, const char *expectedRatios, const char *modeName, + const char *modeReason) { + OP_CHECK_IF(!isSupported, + OP_LOGE(opName_, "cmpRatio should be %s in %s on %s %s, but got %ld.", expectedRatios, modeName, + A2_A3_PLATFORM_LOG.c_str(), modeReason, cmpRatio_), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; + }; + + switch (smlaInfo_.perfMode) { + case SMLATemplateMode::CSA_TEMPLATE_MODE: + return checkRatio(cmpRatio_ == 1 || cmpRatio_ == 2 || cmpRatio_ == 4, "1, 2 or 4", "CSA", + "when cmp_sparse_indices is provided"); + case SMLATemplateMode::HCA_TEMPLATE_MODE: + return checkRatio(cmpRatio_ == 128, "128", "HCA", "when cmp_sparse_indices is not provided"); + default: + return checkRatio(cmpRatio_ == 0, "0", "SWA", "when cmp_kv is not provided"); + } +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaOriMaskMode() const +{ + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCmpMaskMode() const +{ + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaOriWinLeft() const +{ + OP_CHECK_IF(oriWinLeft_ < -1, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "ori_win_left", std::to_string(oriWinLeft_).c_str(), + "Ori_win_left should be -1(unlimited) or non-negative"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaOriWinRight() const +{ + OP_CHECK_IF(oriWinRight_ < -1, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "ori_win_right", std::to_string(oriWinRight_).c_str(), + "Ori_win_right should be -1(unlimited) or non-negative"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCmpResidualKv() const +{ + bool isCmpTemplate = smlaInfo_.perfMode == SMLATemplateMode::HCA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE; + if (isCmpTemplate && *opParamInfo_.cmpMaskMode == 3 && cmpRatio_ != 1) { // 3: RightDownCausal模式 + OP_CHECK_IF( + opParamInfo_.cmpResidualKv.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "cmp_residual_kv", "Cmp_residual_kv is required when cmp_mask_mode=3 and cmp_ratio != 1"), + return ge::GRAPH_FAILED); + } + if (opParamInfo_.cmpResidualKv.tensor == nullptr) { + return ge::GRAPH_SUCCESS; + } + const std::vector dimNumList = {DIM_NUM_ONE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpResidualKv.desc, CMP_RESIDUAL_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.cmpResidualKv.tensor->GetShape(), dimNumList, CMP_RESIDUAL_KV_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF( + opParamInfo_.cmpResidualKv.tensor->GetShapeSize() != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_residual_kv", ToStringRaw(opParamInfo_.cmpResidualKv.tensor->GetStorageShape()), + "Cmp_residual_kv's first dimension(" + std::to_string(opParamInfo_.cmpResidualKv.tensor->GetShapeSize()) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaTopkLength() const +{ + if (IsA5Arch(npuArch_)) { + if (opParamInfo_.oriTopkLength.tensor != nullptr) { + const std::vector dimNumList = {DIM_NUM_TWO, DIM_NUM_THREE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriTopkLength.desc, ORI_TOPK_LENGTH_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriTopkLength.tensor->GetShape(), dimNumList, + ORI_TOPK_LENGTH_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(opParamInfo_.oriTopkLength.desc->GetDataType() != ge::DT_INT32, + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON( + opName_, "ori_topk_length", + SMLADataTypeToSerialString(opParamInfo_.oriTopkLength.desc->GetDataType()).c_str(), + "The dtype of ori_topk_length must be DT_INT32"), + return ge::GRAPH_FAILED); + } + if (opParamInfo_.cmpTopkLength.tensor != nullptr) { + const std::vector dimNumList = {DIM_NUM_TWO, DIM_NUM_THREE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpTopkLength.desc, CMP_TOPK_LENGTH_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpTopkLength.tensor->GetShape(), dimNumList, + CMP_TOPK_LENGTH_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(opParamInfo_.cmpTopkLength.desc->GetDataType() != ge::DT_INT32, + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON( + opName_, "cmp_topk_length", + SMLADataTypeToSerialString(opParamInfo_.cmpTopkLength.desc->GetDataType()).c_str(), + "The dtype of cmp_topk_length must be DT_INT32"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; + } + if (IsNonEmptyOptionalTensor(opParamInfo_.cmpTopkLength.tensor)) { + OP_CHECK_IF(smlaInfo_.perfMode != SMLATemplateMode::CSA_TEMPLATE_MODE, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "cmp_topk_length", + "Cmp_topk_length is only supported in CSA mode"), + return ge::GRAPH_FAILED); + } + if (!IsNonEmptyOptionalTensor(opParamInfo_.oriTopkLength.tensor)) { + OP_CHECK_IF(hasOriSparseIndices_ && smlaInfo_.perfMode == SMLATemplateMode::SWA_TEMPLATE_MODE, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "ori_topk_length", + "Ori_topk_length is required for SWA ori sparse"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(opParamInfo_.oriSparseIndices.tensor == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "ori_sparse_indices", + "Ori_topk_length requires ori_sparse_indices"), + return ge::GRAPH_FAILED); + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriTopkLength.desc, ORI_TOPK_LENGTH_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(opParamInfo_.oriTopkLength.desc->GetDataType() != ge::DT_INT32, + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON( + opName_, "ori_topk_length", + SMLADataTypeToSerialString(opParamInfo_.oriTopkLength.desc->GetDataType()).c_str(), + "The dtype of ori_topk_length must be DT_INT32"), + return ge::GRAPH_FAILED); + const gert::Shape &topkLenShape = opParamInfo_.oriTopkLength.tensor->GetStorageShape(); + const gert::Shape &sparseShape = opParamInfo_.oriSparseIndices.tensor->GetStorageShape(); + if (oriSparseIndicesLayout_ == SMLALayout::TND) { + OP_CHECK_IF(topkLenShape.GetDimNum() != 2 || sparseShape.GetDimNum() != 3, + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + opName_, "ori_topk_length and ori_sparse_indices", + Ops::Base::ToString(topkLenShape) + " and " + Ops::Base::ToString(sparseShape), + "TND ori_topk_length shape must be [T, N2], ori_sparse_indices [T, N2, K]"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(topkLenShape.GetDim(0) != sparseShape.GetDim(0) || topkLenShape.GetDim(1) != sparseShape.GetDim(1), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + opName_, "ori_topk_length and ori_sparse_indices", + Ops::Base::ToString(topkLenShape) + " and " + Ops::Base::ToString(sparseShape), + "Ori_topk_length shape must match ori_sparse_indices without K dim"), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF( + topkLenShape.GetDimNum() != 3 || sparseShape.GetDimNum() != 4, + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(opName_, "ori_topk_length", Ops::Base::ToString(topkLenShape), + "BSND ori_topk_length shape must be [B, S1, N2]"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(topkLenShape.GetDim(0) != sparseShape.GetDim(0) || + topkLenShape.GetDim(1) != sparseShape.GetDim(1) || + topkLenShape.GetDim(2) != sparseShape.GetDim(2), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + opName_, "ori_topk_length and ori_sparse_indices", + Ops::Base::ToString(topkLenShape) + " and " + Ops::Base::ToString(sparseShape), + "Ori_topk_length shape must match ori_sparse_indices without K dim"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSinglePara() const +{ + if (ge::GRAPH_SUCCESS != CheckSingleParaQuery() || ge::GRAPH_SUCCESS != CheckSingleParaOriKv() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpKv() || ge::GRAPH_SUCCESS != CheckSingleParaCuSeqLensQ() || + ge::GRAPH_SUCCESS != CheckSingleParaCuSeqLensOriKv() || ge::GRAPH_SUCCESS != CheckSingleParaCuSeqLensCmpKv() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpRatio() || ge::GRAPH_SUCCESS != CheckSingleParaCmpResidualKv() || + ge::GRAPH_SUCCESS != CheckSingleParaTopkLength() || ge::GRAPH_SUCCESS != CheckSingleParaNumHeads() || + ge::GRAPH_SUCCESS != CheckSingleParaKvHeadNums() || ge::GRAPH_SUCCESS != CheckSingleParaOriSparseIndices() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpSparseIndices() || ge::GRAPH_SUCCESS != CheckSingleParaOriBlockTable() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpBlockTable() || ge::GRAPH_SUCCESS != CheckSingleParaSinks() || + ge::GRAPH_SUCCESS != CheckSingleParaMetadata() || ge::GRAPH_SUCCESS != CheckSingleParaOriMaskMode() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpMaskMode() || ge::GRAPH_SUCCESS != CheckSingleParaOriWinLeft() || + ge::GRAPH_SUCCESS != CheckSingleParaOriWinRight()) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckExists(const void *pointer, const std::string &name) const +{ + OP_CHECK_IF(pointer == nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, name.c_str(), name + " should not be null"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckNotExists(const void *pointer, const std::string &name) const +{ + OP_CHECK_IF(pointer != nullptr, + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, name.c_str(), name + " should not be null"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckExistsByMap(const std::map ¶mMap) const +{ + for (const auto &kv : paramMap) { + if (CheckExists(kv.second, kv.first) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckNotExistsByMap(const std::map ¶mMap) const +{ + for (const auto &kv : paramMap) { + if (CheckNotExists(kv.second, kv.first) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckExistenceByMap(std::map &existMap, + std::map ¬ExistMap) const +{ + if (CheckExistsByMap(existMap) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + if (CheckNotExistsByMap(notExistMap) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckParaExistence() const +{ + if (npuArch_ == NpuArch::DAV_3510) { + OP_CHECK_IF((kvLayout_ == SMLALayout::TND && opParamInfo_.cuSeqLensOriKv.tensor == nullptr), + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON( + opName_, "cu_seqlens_ori_kv", "cu_seqlens_ori_kv must be provided when kv layout is TND"), + return ge::GRAPH_FAILED); + } else { + if (kvLayout_ == SMLALayout::PA_BBND) { + std::map ParamExistMap = { + {"actualSeqLengths", opParamInfo_.sequsedOriKv.tensor}, + {"oriBlockTable", opParamInfo_.oriBlockTable.tensor}, + }; + std::map ParamNotExistMap = {}; + if (CheckExistenceByMap(ParamExistMap, ParamNotExistMap) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckFeatureShape() const +{ + if (qLayout_ == SMLALayout::TND) { + OP_CHECK_IF(bSize_ <= 0, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cu_seqlens_q", ToStringRaw(opParamInfo_.cuSeqLensQ.tensor->GetStorageShape()).c_str(), + "Batch_size should be greater than 0, but got " + std::to_string(bSize_)), + return ge::GRAPH_FAILED); + } else { // BSND + OP_CHECK_IF(bSize_ <= 0, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "q", ToStringRaw(opParamInfo_.q.shape->GetStorageShape()).c_str(), + "Batch_size should be greater than 0, but got " + std::to_string(bSize_)), + return ge::GRAPH_FAILED); + } + + OP_CHECK_IF(qTSize_ <= 0 && (qLayout_ == SMLALayout::TND), + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "q", ToStringRaw(opParamInfo_.q.shape->GetStorageShape()).c_str(), + "T_size of q should be greater than 0, but got " + std::to_string(qTSize_)), + return ge::GRAPH_FAILED); + + if (IsA5Arch(npuArch_)) { + OP_CHECK_IF(n1Size_ < 1 || n1Size_ > 128, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(opName_, "q", + ToStringRaw(opParamInfo_.q.shape->GetStorageShape()).c_str(), + "The head num of q should be in [1, 128] on" + + A5_PLATFORM_LOG + ", but got " + std::to_string(n1Size_)), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(!IsPowerOfTwoInRange(n1Size_, 1, 128), + OP_LOGE(opName_, "q_head_num should be power of two in [1, 128] on %s, but got %u", + A2_A3_PLATFORM_LOG.c_str(), n1Size_), + return ge::GRAPH_FAILED); + } + + if (opParamInfo_.oriKv.tensor != nullptr) { + OP_CHECK_IF(n2Size_ != 1, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "ori_kv", ToStringRaw(opParamInfo_.oriKv.tensor->GetStorageShape()).c_str(), + "The head num of ori_kv should be 1, but got " + std::to_string(n2Size_)), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(n1Size_ % n2Size_ != 0, + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + opName_, "q and ori_kv", + Ops::Base::ToString(opParamInfo_.q.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.oriKv.tensor->GetStorageShape()), + "The head num of q(" + std::to_string(n1Size_) + + ") must be divisible by the head num of ori_kv(" + std::to_string(n2Size_) + ")"), + return ge::GRAPH_FAILED); + } + if (opParamInfo_.cmpKv.tensor != nullptr) { + OP_CHECK_IF(n2Size_ != 1, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_kv", ToStringRaw(opParamInfo_.cmpKv.tensor->GetStorageShape()).c_str(), + "The head num of cmp_kv should be 1, but got " + std::to_string(n2Size_)), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(n1Size_ % n2Size_ != 0, + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + opName_, "q and cmp_kv", + Ops::Base::ToString(opParamInfo_.q.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.cmpKv.tensor->GetStorageShape()), + "The head num of q(" + std::to_string(n1Size_) + + ") must be divisible by the head num of cmp_kv(" + std::to_string(n2Size_) + ")"), + return ge::GRAPH_FAILED); + } + + if (IsA5Arch(npuArch_)) { + if (opParamInfo_.oriKv.tensor != nullptr) { + OP_CHECK_IF(gSize_ < 1 || gSize_ > 128, + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + opName_, "q and ori_kv", + Ops::Base::ToString(opParamInfo_.q.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.oriKv.tensor->GetStorageShape()), + "The value of (the head num of q ceildivided by the head num of ori_kv) " + "should be in [1, 128] on " + + A5_PLATFORM_LOG + ", but got " + std::to_string(gSize_)), + return ge::GRAPH_FAILED); + } + if (opParamInfo_.cmpKv.tensor != nullptr) { + OP_CHECK_IF(gSize_ < 1 || gSize_ > 128, + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + opName_, "q and cmp_kv", + Ops::Base::ToString(opParamInfo_.q.shape->GetStorageShape()) + " and " + + Ops::Base::ToString(opParamInfo_.cmpKv.tensor->GetStorageShape()), + "The value of (the head num of q ceildivided by the head num of cmp_kv) " + "should be in [1, 128] on " + + A5_PLATFORM_LOG + ", but got " + std::to_string(gSize_)), + return ge::GRAPH_FAILED); + } + } else { + OP_CHECK_IF(!IsPowerOfTwoInRange(gSize_, 1, 128), + OP_LOGE(opName_, "group num should be power of two in [1, 128] on %s, but got %u", + A2_A3_PLATFORM_LOG.c_str(), gSize_), + return ge::GRAPH_FAILED); + } + + OP_CHECK_IF( + qHeadDim_ != DIM_LIMIT, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "q", ToStringRaw(opParamInfo_.q.shape->GetStorageShape()).c_str(), + "The head num of q only support " + std::to_string(DIM_LIMIT) + ", but got " + std::to_string(qHeadDim_)), + return ge::GRAPH_FAILED); + + if (opParamInfo_.oriKv.tensor != nullptr) { + OP_CHECK_IF(oriKvHeadDim_ != DIM_LIMIT, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "ori_kv", ToStringRaw(opParamInfo_.oriKv.tensor->GetStorageShape()).c_str(), + "The head num of ori_kv only support " + std::to_string(DIM_LIMIT) + ", but got " + + std::to_string(oriKvHeadDim_)), + return ge::GRAPH_FAILED); + } + if (!(smlaInfo_.perfMode == SMLATemplateMode::SWA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) && + (opParamInfo_.cmpKv.tensor != nullptr)) { + OP_CHECK_IF(cmpKvHeadDim_ != DIM_LIMIT, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "cmp_kv", ToStringRaw(opParamInfo_.cmpKv.tensor->GetStorageShape()).c_str(), + "The head num of cmp_kv only support " + std::to_string(DIM_LIMIT) + ", but got " + + std::to_string(cmpKvHeadDim_)), + return ge::GRAPH_FAILED); + } + + OP_CHECK_IF(!(qType_ == oriKvType_), + OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON( + opName_, "q and ori_kv", + SMLADataTypeToSerialString(qType_) + " and " + SMLADataTypeToSerialString(oriKvType_), + "The dtype of q and ori_kv must be same"), + return ge::GRAPH_FAILED); + + if (IsA5Arch(npuArch_)) { + OP_CHECK_IF(*opParamInfo_.oriMaskMode != 0 && *opParamInfo_.oriMaskMode != 3 && *opParamInfo_.oriMaskMode != 4, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "ori_mask_mode", + std::to_string(*opParamInfo_.oriMaskMode).c_str(), + "Ori_mask_mode should be {0, 3, 4} on " + A5_PLATFORM_LOG), + return ge::GRAPH_FAILED); + OP_CHECK_IF(*opParamInfo_.cmpMaskMode != 0 && *opParamInfo_.cmpMaskMode != 3, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "cmp_mask_mode", + std::to_string(*opParamInfo_.cmpMaskMode).c_str(), + "Cmp_mask_mode should be {0, 3} on " + A5_PLATFORM_LOG), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + topkValueMode_ != 1, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "topk_value_mode", std::to_string(topkValueMode_).c_str(), + "Topk_value_mode should be 1"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(oriWinLeft_ < -1, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + opName_, "ori_win_left", std::to_string(oriWinLeft_).c_str(), + "Ori_win_left should be -1(unlimited) or non-negative on " + A5_PLATFORM_LOG), + return ge::GRAPH_FAILED); + OP_CHECK_IF(oriWinRight_ < -1, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + opName_, "ori_win_right", std::to_string(oriWinRight_).c_str(), + "Ori_win_right should be -1(unlimited) or non-negative on " + A5_PLATFORM_LOG), + return ge::GRAPH_FAILED); + } else { + if (smlaInfo_.perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::HCA_TEMPLATE_MODE) { + OP_CHECK_IF(*opParamInfo_.cmpMaskMode != 3, + OP_LOGE(opName_, "cmpMaskMode should be 3 on %s, but got %u", A2_A3_PLATFORM_LOG.c_str(), + *opParamInfo_.cmpMaskMode), + return ge::GRAPH_FAILED); + } + if (hasOriSparseIndices_) { + OP_CHECK_IF(*opParamInfo_.oriMaskMode != 0U, + OP_LOGE(opName_, "oriMaskMode must be 0 for SWA ori sparse (DSpark) on %s, but got %u.", + A2_A3_PLATFORM_LOG.c_str(), *opParamInfo_.oriMaskMode), + return ge::GRAPH_FAILED); + OP_CHECK_IF(oriWinLeft_ < 0 || oriWinRight_ < 0, + OP_LOGE(opName_, "ori_win_left/right should be non-negative for ori sparse SWA on %s.", + A2_A3_PLATFORM_LOG.c_str()), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(*opParamInfo_.oriMaskMode != 4U, + OP_LOGE(opName_, "oriMaskMode should be 4 on %s when ori_sparse_indices is empty, but got %u", + A2_A3_PLATFORM_LOG.c_str(), *opParamInfo_.oriMaskMode), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + oriWinLeft_ != 127, + OP_LOGE(opName_, "oriWinLeft_ should be 127 on %s when ori_sparse_indices is empty, but got %ld", + A2_A3_PLATFORM_LOG.c_str(), oriWinLeft_), + return ge::GRAPH_FAILED); + OP_CHECK_IF(oriWinRight_ != 0, + OP_LOGE(opName_, "oriWinRight_ should be 0 on %s when ori_sparse_indices is empty, but got %ld", + A2_A3_PLATFORM_LOG.c_str(), oriWinRight_), + return ge::GRAPH_FAILED); + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckFeatureLayout() const +{ + const std::vector layoutQuerySupportList = {"BSND", "TND"}; + std::string layoutQuery = opParamInfo_.layoutQ; + OP_CHECK_IF(std::find(layoutQuerySupportList.begin(), layoutQuerySupportList.end(), layoutQuery) == + layoutQuerySupportList.end(), + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "layout_q", layoutQuery.c_str(), + "Layout_q only supports BSND or TND"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckFeatureDtype() const +{ + OP_CHECK_IF(qType_ != ge::DT_BF16 && qType_ != ge::DT_FLOAT16, + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "q", SMLADataTypeToSerialString(qType_).c_str(), + "The dtype of q only supports " + + SMLADataTypeToSerialString(ge::DT_BF16) + " and " + + SMLADataTypeToSerialString(ge::DT_FLOAT16)), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckFeaturePa() const +{ + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckFeature() const +{ + if (ge::GRAPH_SUCCESS != CheckFeatureShape() || ge::GRAPH_SUCCESS != CheckFeatureLayout() || + ge::GRAPH_SUCCESS != CheckFeatureDtype() || ge::GRAPH_SUCCESS != CheckFeaturePa()) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckDTypeConsistency(const ge::DataType &actualDtype, const ge::DataType &expectDtype, + const std::string &name) const +{ + if (actualDtype != expectDtype) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON( + opName_, name.c_str(), SMLADataTypeToSerialString(actualDtype).c_str(), + "The dtype of " + name + " should be " + SMLADataTypeToSerialString(expectDtype)); + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckOriAndCmpKv() const +{ + if (smlaInfo_.perfMode == SMLATemplateMode::HCA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (ge::GRAPH_SUCCESS != CheckDTypeConsistency(cmpKvType_, oriKvType_, CMP_KV_NAME)) { + return ge::GRAPH_FAILED; + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckAttenOut() const +{ + if (opParamInfo_.attnOut.desc == nullptr) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(opName_, "attn_out", "Attn_out must be provided"); + return ge::GRAPH_FAILED; + } + const std::vector dimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.attnOut.desc, ATTEN_OUT_NAME) || + ge::GRAPH_SUCCESS != CheckLayoutSupport(outLayout_, ATTEN_OUT_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(opParamInfo_.attnOut.shape, dimNumList, ATTEN_OUT_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(outLayout_, opParamInfo_.attnOut.shape, ATTEN_OUT_NAME)) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckActualSeqLensQ() const +{ + if (opParamInfo_.seqUsedQ.tensor == nullptr) { + return ge::GRAPH_SUCCESS; + } + const std::vector dimNumList = {DIM_NUM_ONE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.seqUsedQ.desc, SEQUSED_Q_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.seqUsedQ.tensor->GetShape(), dimNumList, SEQUSED_Q_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(opParamInfo_.seqUsedQ.tensor->GetShapeSize() != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "seqused_q", ToStringRaw(opParamInfo_.seqUsedQ.tensor->GetStorageShape()), + "Seqused_q's first dimension(" + std::to_string(opParamInfo_.seqUsedQ.tensor->GetShapeSize()) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckActualSeqLens() const +{ + if (opParamInfo_.sequsedOriKv.tensor != nullptr) { + const std::vector dimNumList = {DIM_NUM_ONE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.sequsedOriKv.desc, SEQUSED_ORI_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.sequsedOriKv.tensor->GetShape(), dimNumList, SEQUSED_ORI_KV_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF( + opParamInfo_.sequsedOriKv.tensor->GetShapeSize() != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "seqused_ori_kv", ToStringRaw(opParamInfo_.sequsedOriKv.tensor->GetStorageShape()), + "Seqused_ori_kv's first dimension(" + std::to_string(opParamInfo_.sequsedOriKv.tensor->GetShapeSize()) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + } + if (opParamInfo_.sequsedCmpKv.tensor != nullptr) { + const std::vector dimNumList = {DIM_NUM_ONE}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.sequsedCmpKv.desc, SEQUSED_CMP_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.sequsedCmpKv.tensor->GetShape(), dimNumList, SEQUSED_CMP_KV_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF( + opParamInfo_.sequsedCmpKv.tensor->GetShapeSize() != bSize_, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + opName_, "seqused_cmp_kv", ToStringRaw(opParamInfo_.sequsedCmpKv.tensor->GetStorageShape()), + "Seqused_cmp_kv's first dimension(" + std::to_string(opParamInfo_.sequsedCmpKv.tensor->GetShapeSize()) + + ") should be equal to B(" + std::to_string(bSize_) + ")"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} +ge::graphStatus SMLATilingCheck::CheckMultiParaConsistency() +{ + if (ge::GRAPH_SUCCESS != CheckOriAndCmpKv() || ge::GRAPH_SUCCESS != CheckAttenOut() || + ge::GRAPH_SUCCESS != CheckActualSeqLensQ() || ge::GRAPH_SUCCESS != CheckActualSeqLens()) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::Process() +{ + Init(); + if (CheckSinglePara() != ge::GRAPH_SUCCESS || CheckParaExistence() != ge::GRAPH_SUCCESS || + CheckFeature() != ge::GRAPH_SUCCESS || CheckMultiParaConsistency() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +void SparseFlashMlaTiling::CalcUbBmm(const SMLATilingInfo *tilingInfo) +{ + uint32_t cubeMSize = tilingInfo->gSize * tilingInfo->s1Size; + uint32_t maxMSize = mBaseSize_; + if (cubeMSize > maxMSize) { + cubeMSize = maxMSize; + } + mmResUbSize_ = sInnerSizeAlign_ * Align(cubeMSize, 16U); // kernel按照16对齐写出,tiling按照这个原则分配内存 + bmm2ResUbSize_ = headDimAlign_ * Align(cubeMSize, 16U); // kernel按照16对齐写出,tiling按照这个原则分配内存 +} + +void SparseFlashMlaTiling::SplitBalanced(SMLATilingInfo *tilingInfo) +{ + sInnerSizeAlign_ = Align(sInnerSize_, BYTE_BLOCK); + if (tilingInfo->npuArch == NpuArch::DAV_2201) { + mBaseSize_ = tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE ? + tilingInfo->gSize : + (256 / tilingInfo->gSize) * tilingInfo->gSize; // 256:DAV_2201架构下的基础计算块大小 + // DSpark ori_sparse_indices are per query token; keep one S1 row per M block. + if (tilingInfo->hasOriSparseIndices && tilingInfo->perfMode == SMLATemplateMode::SWA_TEMPLATE_MODE) { + mBaseSize_ = tilingInfo->gSize; + } + } + headDimAlign_ = Align(tilingInfo->qHeadDim, BYTE_BLOCK); + CalcUbBmm(tilingInfo); + + tilingData_.baseParams.set_mBaseSize(mBaseSize_); + tilingData_.baseParams.set_s2BaseSize(sInnerSize_); + tilingData_.baseParams.set_mmResUbSize(mmResUbSize_); + tilingData_.baseParams.set_bmm2ResUbSize(bmm2ResUbSize_); +} + +uint32_t SparseFlashMlaTiling::CalcFdLogicalSlotCount(const SMLATilingInfo *tilingInfo, uint32_t aicNum) const +{ + bool isSplitG = IsA5Arch(tilingInfo->npuArch) && tilingInfo->gSize > 64; + if (isSplitG) { + return aicNum >> 1U; + } + return aicNum; +} + +uint64_t SparseFlashMlaTiling::CalcFdStagingWorkspaceSize(const SMLATilingInfo *tilingInfo, uint32_t aicNum) const +{ + if (!IsA5Arch(tilingInfo->npuArch) || aicNum == 0U) { + return 0ULL; + } + + const uint64_t logicalCoreSlots = CalcFdLogicalSlotCount(tilingInfo, aicNum); + const uint64_t bytesPerSlot = tilingInfo->gSize * sizeof(float) * (headDimAlign_ + 2ULL * FD_BROADCAST_ELEMS); + if (tilingInfo->batchConsistency) { + // Intra-core reduction ping-pongs by multiCoreIdxMod2. Cross-core staging + // keeps one reduce-block range for every physical AIC. + const uint64_t intraCoreSlots = 2ULL * logicalCoreSlots; + const uint64_t crossCoreSlots = static_cast(aicNum) * BATCH_CONSISTENCY_MAX_REDUCE_BLOCK_NUM; + return (intraCoreSlots + crossCoreSlots) * bytesPerSlot; + } + return logicalCoreSlots * FD_MAX_S2_SPLIT_NUM * bytesPerSlot; +} + +uint64_t SparseFlashMlaTiling::CalcVectorizeKvPhyAddrWorkspaceSize(const SMLATilingInfo *tilingInfo, + uint32_t &vectorizeFlag) const +{ + vectorizeFlag = 0U; + if (tilingInfo->npuArch != NpuArch::DAV_3510) { + return 0ULL; + } + constexpr uint32_t SPARSE_BLOCK_ALIGN_NUM = 128; + constexpr uint32_t UB_SIZE = 184 * 1024; + uint32_t alignedOriSparseBlockCount = (tilingInfo->oriSparseBlockCount + SPARSE_BLOCK_ALIGN_NUM - 1) / + SPARSE_BLOCK_ALIGN_NUM * SPARSE_BLOCK_ALIGN_NUM; + uint32_t alignedCmpSparseBlockCount = (tilingInfo->cmpSparseBlockCount + SPARSE_BLOCK_ALIGN_NUM - 1) / + SPARSE_BLOCK_ALIGN_NUM * SPARSE_BLOCK_ALIGN_NUM; + bool isPa = (tilingInfo->kvLayout == SMLALayout::PA_BBND); + uint32_t oriBlocksizeFlag = + static_cast(tilingInfo->oriBlockSize & static_cast(tilingInfo->oriBlockSize - 1)) == 0; + uint32_t cmpBlocksizeFlag = + static_cast(tilingInfo->cmpBlockSize & static_cast(tilingInfo->cmpBlockSize - 1)) == 0; + uint32_t blocksizeFlag = isPa ? ((tilingInfo->perfMode == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) ? + oriBlocksizeFlag : + (oriBlocksizeFlag != 0 && cmpBlocksizeFlag != 0)) : + 1U; + uint64_t paExtraUb = std::max(static_cast(tilingInfo->oriMaxBlockNumPerBatch) * sizeof(int32_t), + static_cast(tilingInfo->cmpMaxBlockNumPerBatch) * sizeof(int32_t)); + uint64_t cmpUbSize = (isPa ? paExtraUb : 0U) + static_cast(alignedCmpSparseBlockCount) * sizeof(int32_t) + + static_cast(alignedCmpSparseBlockCount) * sizeof(int64_t); + uint64_t oriUbSize = (isPa ? paExtraUb : 0U) + static_cast(alignedOriSparseBlockCount) * sizeof(int32_t) + + static_cast(alignedOriSparseBlockCount) * sizeof(int64_t); + uint64_t vectorizeUbSize = std::max(cmpUbSize, oriUbSize); + vectorizeFlag = static_cast((tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE || + tilingInfo->perfMode == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + tilingInfo->perfMode == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) && + (vectorizeUbSize <= UB_SIZE) && (blocksizeFlag != 0)); + + if (vectorizeFlag == 0U) { + return 0ULL; + } + uint32_t totalBS1 = + (tilingInfo->qLayout == SMLALayout::TND) ? tilingInfo->s1Size : (tilingInfo->bSize * tilingInfo->s1Size); + uint64_t oriPhyAddrSize = 0; + if (tilingInfo->perfMode == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + tilingInfo->perfMode == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + oriPhyAddrSize = static_cast(totalBS1) * alignedOriSparseBlockCount * sizeof(int64_t); + } + uint64_t cmpPhyAddrSize = 0; + if (tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE || + tilingInfo->perfMode == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + cmpPhyAddrSize = static_cast(totalBS1) * alignedCmpSparseBlockCount * sizeof(int64_t); + } + return oriPhyAddrSize + cmpPhyAddrSize; +} + +// --------------------------SparseFlashMlaTiling类成员函数定义---------------------- +ge::graphStatus SparseFlashMlaTiling::DoOpTiling(SMLATilingInfo *tilingInfo) +{ + auto ascendcPlatform = platform_ascendc::PlatformAscendC(tilingInfo->platformInfo); + uint32_t aivNum = ascendcPlatform.GetCoreNumAiv(); + uint32_t aicNum = ascendcPlatform.GetCoreNumAic(); + uint32_t blockDim = ascendcPlatform.CalcTschBlockDim(aivNum, aicNum, aivNum); + context_->SetBlockDim(blockDim); + OP_LOGI(tilingInfo->opName, "SMLA block dim: %u aiv Num: %u aic Num: %u.", blockDim, aivNum, aicNum); + + SplitBalanced(tilingInfo); + + constexpr uint32_t MM1_RES_ELEM_SIZE = 4; + constexpr uint32_t VEC1_RES_ELEM_SIZE = 2; + constexpr uint32_t MM2_RES_ELEM_SIZE = 4; + constexpr uint32_t VEC2_RES_ELEM_SIZE = 4; + constexpr uint32_t PRELOAD_NUM = 2; + + size_t workspaceSize = ascendcPlatform.GetLibApiWorkSpaceSize(); + if (tilingInfo->npuArch == NpuArch::DAV_3510) { + bool isSplitG = tilingInfo->gSize > 64; + if (isSplitG || tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE || + tilingInfo->perfMode == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + tilingInfo->perfMode == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + constexpr uint32_t TRIPLE_BUFFER_NUM = 3; + constexpr uint32_t S2_BASE_SIZE = 128; + constexpr uint32_t D_SIZE = 512; + constexpr uint32_t VEC_RES_ELEM_SIZE = 2; + uint32_t gmSize = S2_BASE_SIZE * D_SIZE * VEC_RES_ELEM_SIZE * TRIPLE_BUFFER_NUM; + if (isSplitG) { + gmSize *= (aicNum >> 1U); + } else { + gmSize *= aicNum; + } + workspaceSize += gmSize; + constexpr uint32_t S2_REAL_BUF_LEN = 128; + workspaceSize += TRIPLE_BUFFER_NUM * S2_REAL_BUF_LEN * sizeof(int32_t) * aivNum; + } + } else { + workspaceSize += PRELOAD_NUM * mmResUbSize_ * MM1_RES_ELEM_SIZE * aicNum; + workspaceSize += PRELOAD_NUM * mmResUbSize_ * VEC1_RES_ELEM_SIZE * aicNum; + workspaceSize += PRELOAD_NUM * bmm2ResUbSize_ * MM2_RES_ELEM_SIZE * aicNum; + workspaceSize += PRELOAD_NUM * bmm2ResUbSize_ * VEC2_RES_ELEM_SIZE * aicNum; + if (tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE || + (tilingInfo->perfMode == SMLATemplateMode::SWA_TEMPLATE_MODE && tilingInfo->hasOriSparseIndices)) { + constexpr uint32_t MERGE_CACHE_GM_BUF_NUM = 3; // 3:缓存缓冲区数量 + workspaceSize += MERGE_CACHE_GM_BUF_NUM * 512 * 512 * 2 * aicNum; // 缓冲区的尺寸为512x512字节,2表示双缓冲 + } + } + + // 计算vectorizeFlag (稀疏KV物理地址向量化) + uint32_t vectorizeFlag = 0U; + workspaceSize += CalcVectorizeKvPhyAddrWorkspaceSize(tilingInfo, vectorizeFlag); + + workspaceSize += CalcFdStagingWorkspaceSize(tilingInfo, aicNum); + size_t *workSpaces = context_->GetWorkspaceSizes(1); + workSpaces[0] = workspaceSize; + + tilingData_.baseParams.set_batchSize(tilingInfo->bSize); + tilingData_.baseParams.set_kvSeqSize(tilingInfo->s2Size); + tilingData_.baseParams.set_qSeqSize(tilingInfo->s1Size); + tilingData_.baseParams.set_nNumOfQInOneGroup(tilingInfo->gSize); + tilingData_.baseParams.set_paBlockSize(tilingInfo->blockSize); + tilingData_.baseParams.set_oriBlockSize(tilingInfo->oriBlockSize); + tilingData_.baseParams.set_cmpBlockSize(tilingInfo->cmpBlockSize); + tilingData_.baseParams.set_oriMaxBlockNumPerBatch(tilingInfo->oriMaxBlockNumPerBatch); + tilingData_.baseParams.set_actualLenDimsQ(tilingInfo->actualLenDimsQ); + tilingData_.baseParams.set_actualLenDimsKV(tilingInfo->actualLenDimsKV); + + tilingData_.baseParams.set_softmaxScale(tilingInfo->softmaxScale); + tilingData_.baseParams.set_outputLayout(static_cast(tilingInfo->outLayout)); + tilingData_.baseParams.set_oriMaskMode(tilingInfo->oriMaskMode); + tilingData_.baseParams.set_oriKvStride0(tilingInfo->oriKvStride0); + tilingData_.baseParams.set_oriWinLeft(tilingInfo->oriWinLeft); + tilingData_.baseParams.set_oriWinRight(tilingInfo->oriWinRight); + tilingData_.baseParams.set_sparseBlockSize(tilingInfo->sparseBlockSize); + tilingData_.baseParams.set_returnSoftmaxLse(tilingInfo->returnSoftmaxLse); + + tilingData_.cmpParams.set_cmpMaxBlockNumPerBatch(tilingInfo->cmpMaxBlockNumPerBatch); + tilingData_.cmpParams.set_cmpRatio(tilingInfo->cmpRatio); + tilingData_.cmpParams.set_cmpMaskMode(tilingInfo->cmpMaskMode); + tilingData_.cmpParams.set_cmpKvStride0(tilingInfo->cmpKvStride0); + tilingData_.cmpParams.set_cmpKvSeqSize(tilingInfo->cmpS2Size); + tilingData_.baseParams.set_actualLenDimsOriKV(tilingInfo->actualLenDimsOriKV); + tilingData_.baseParams.set_actualLenDimsCmpKV(tilingInfo->actualLenDimsCmpKV); + tilingData_.baseParams.set_cmpResidualKVSize(tilingInfo->cmpResidualKVSize); + tilingData_.baseParams.set_oriKeyStride0(tilingInfo->oriKeyStride0); + tilingData_.baseParams.set_hasOriSparseIndices(tilingInfo->hasOriSparseIndices ? 1U : 0U); + tilingData_.baseParams.set_oriSparseIndexWidth(tilingInfo->oriSparseIndexWidth); + tilingData_.cmpParams.set_cmpKeyStride0(tilingInfo->cmpKeyStride0); + + if (tilingInfo->npuArch == NpuArch::DAV_3510) { + tilingData_.baseParams.set_oriSparseBlockCount(tilingInfo->oriSparseBlockCount); + tilingData_.baseParams.set_topkValueMode(tilingInfo->topkValueMode); + tilingData_.cmpParams.set_cmpSparseBlockCount(tilingInfo->cmpSparseBlockCount); + } else { + tilingData_.cmpParams.set_sparseBlockCount(tilingInfo->sparseBlockCount); + } + + usedCoreNum_ = aicNum; + tilingData_.baseParams.set_usedCoreNum(usedCoreNum_); + tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity()); + context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize()); + + uint32_t qLayout = static_cast(tilingInfo->qLayout); + uint32_t inputKvLayout = static_cast(tilingInfo->kvLayout); + + uint64_t tilingKey; + uint32_t splitG = 0U; + uint32_t headRatioOne = + static_cast(tilingInfo->npuArch == NpuArch::DAV_2201 && + tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE && tilingInfo->gSize == 1U); + if (tilingInfo->npuArch == NpuArch::DAV_3510) { + splitG = static_cast(tilingInfo->gSize > 64); // 64:分组拆分阈值 + } + tilingKey = GET_TPL_TILING_KEY(0U, qLayout, inputKvLayout, static_cast(tilingInfo->perfMode), splitG, + headRatioOne, static_cast(tilingInfo->batchConsistency), vectorizeFlag); + context_->SetScheduleMode(1); + context_->SetTilingKey(tilingKey); + + return ge::GRAPH_SUCCESS; +} + +} // namespace optiling + +namespace optiling { +// --------------------------TilingPrepare函数定义------------------------------------- +static ge::graphStatus TilingPrepareForSparseFlashMla(gert::TilingParseContext * /* context */) +{ + return ge::GRAPH_SUCCESS; +} + +// --------------------------Tiling函数定义--------------------------- +ge::graphStatus TilingForSparseFlashMla(gert::TilingContext *context) +{ + OP_CHECK_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("SparseFlashMla", "Tiling context is null."), + return ge::GRAPH_FAILED); + + SMLATilingInfo smlaInfo; + SMLAInfoParser smlaInfoParser(context); + if (smlaInfoParser.Parse(smlaInfo) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + + if (smlaInfo.npuArch == NpuArch::DAV_2201) { + SMLATilingCheck smlaTilingChecker(smlaInfo); + if (smlaTilingChecker.Process() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + } else { + SparseFlashMlaChecker smlaTilingChecker(smlaInfo); + if (smlaTilingChecker.Process() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + } + + SparseFlashMlaTiling tiling(context); + return tiling.DoOpTiling(&smlaInfo); +} +// --------------------------Tiling函数及TilingPrepare函数注册-------- +IMPL_OP_OPTILING(SparseFlashMla) + .Tiling(TilingForSparseFlashMla) + .TilingParse(TilingPrepareForSparseFlashMla); +} // namespace optiling diff --git a/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.h b/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.h new file mode 100644 index 000000000000..fd35bdd21845 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.h @@ -0,0 +1,575 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_tiling.h + * \brief + */ +#ifndef SPARSE_FLASH_MLA_TILING_H +#define SPARSE_FLASH_MLA_TILING_H + +#include +#include +#include "register/tilingdata_base.h" +#include "tiling/tiling_api.h" +#include "err/ops_err.h" +#include "platform/soc_spec.h" + +namespace optiling { +// ------------------公共定义-------------------------- +struct SMLATilingRequiredParaInfo { + const gert::CompileTimeTensorDesc *desc; + const gert::StorageShape *shape; +}; + +struct SMLATilingOptionalParaInfo { + const gert::CompileTimeTensorDesc *desc; + const gert::Tensor *tensor; + const gert::StorageShape *shape; +}; + +enum class SMLALayout : uint32_t { + BSND = 0, + TND = 1, + PA_BBND = 2 +}; + +enum class SMLAAxis : uint32_t { + B = 0, + S = 1, + N = 2, + D = 3, + K = 3, // sparse_indices的K和key的D枚举值相同,表达相同位置, 最后一维 + T = 5, + Bn = 6, // block number + Bs = 7 // block size +}; + +enum class SMLATemplateMode : uint32_t { + SWA_TEMPLATE_MODE = 0, + HCA_TEMPLATE_MODE = 1, + CSA_TEMPLATE_MODE = 2, + ORI_SPARSE_TEMPLATE_MODE = 3, + ORI_CMP_SPARSE_TEMPLATE_MODE = 4 +}; + +// ------------------算子原型索引常量定义---------------- +// Inputs Index (0-10, common) +constexpr uint32_t Q_INDEX = 0; +constexpr uint32_t ORI_KV_INDEX = 1; +constexpr uint32_t CMP_KV_INDEX = 2; +constexpr uint32_t ORI_SPARSE_INDICES_INDEX = 3; +constexpr uint32_t CMP_SPARSE_INDICES_INDEX = 4; +constexpr uint32_t ORI_BLOCK_TABLE_INDEX = 5; +constexpr uint32_t CMP_BLOCK_TABLE_INDEX = 6; +constexpr uint32_t CU_SEQLENS_Q_INDEX = 7; +constexpr uint32_t CU_SEQLENS_ORI_KV_INDEX = 8; +constexpr uint32_t CU_SEQLENS_CMP_KV_INDEX = 9; +constexpr uint32_t SEQUSED_Q_INDEX = 10; +constexpr uint32_t SEQUSED_ORI_KV_INDEX = 11; +constexpr uint32_t SEQUSED_CMP_KV_INDEX = 12; +constexpr uint32_t CMP_RESIDUAL_KV_INDEX = 13; +constexpr uint32_t ORI_TOPK_LENGTH_INDEX = 14; +constexpr uint32_t CMP_TOPK_LENGTH_INDEX = 15; +constexpr uint32_t SINKS_INDEX = 16; +constexpr uint32_t METADATA_INDEX = 17; +// Outputs Index +constexpr uint32_t ATTN_OUT_INDEX = 0; +constexpr uint32_t SOFTMAX_LSE_INDEX = 1; + +// Attributes Index +constexpr uint32_t ATTR_SOFTMAX_SCALE_INDEX = 0; +constexpr uint32_t ATTR_CMP_RATIO_INDEX = 1; +constexpr uint32_t ATTR_ORI_MASK_MODE_INDEX = 2; +constexpr uint32_t ATTR_CMP_MASK_MODE_INDEX = 3; +constexpr uint32_t ATTR_ORI_WIN_LEFT_INDEX = 4; +constexpr uint32_t ATTR_ORI_WIN_RIGHT_INDEX = 5; +constexpr uint32_t ATTR_LAYOUT_Q_INDEX = 6; +constexpr uint32_t ATTR_LAYOUT_KV_INDEX = 7; +constexpr uint32_t ATTR_TOPK_VALUE_MODE_INDEX = 8; // A2/A3 +constexpr uint32_t ATTR_RETURN_SOFTMAX_LSE_INDEX = 9; + +// Dim Index +constexpr uint32_t DIM_IDX_ONE = 1; +constexpr uint32_t DIM_IDX_TWO = 2; +constexpr uint32_t DIM_IDX_THREE = 3; +constexpr uint32_t DIM_IDX_FOUR = 4; + +// Dim Num +constexpr uint32_t DIM_NUM_ONE = 1; +constexpr uint32_t DIM_NUM_TWO = 2; +constexpr uint32_t DIM_NUM_THREE = 3; +constexpr uint32_t DIM_NUM_FOUR = 4; + +// 常量 +constexpr uint32_t BYTE_BLOCK = 32; + +// 入参限制常量 +constexpr uint32_t METADATA_LIMIT = 1024; +constexpr uint32_t DIM_LIMIT = 512; +constexpr uint32_t BLOCK_SIZE_LIMIT = 1024; + +// -----------算子TilingData定义(A2/A3字段顺序 + A5追加字段)--------------- +BEGIN_TILING_DATA_DEF(SparseFlashMlaSwaParams) +TILING_DATA_FIELD_DEF(uint32_t, batchSize) +TILING_DATA_FIELD_DEF(uint32_t, qSeqSize) +TILING_DATA_FIELD_DEF(uint32_t, kvSeqSize) +TILING_DATA_FIELD_DEF(int64_t, paBlockSize) +TILING_DATA_FIELD_DEF(int64_t, oriBlockSize) +TILING_DATA_FIELD_DEF(int64_t, cmpBlockSize) +TILING_DATA_FIELD_DEF(uint32_t, oriMaxBlockNumPerBatch) +TILING_DATA_FIELD_DEF(uint32_t, nNumOfQInOneGroup) +TILING_DATA_FIELD_DEF(uint32_t, actualLenDimsQ) +TILING_DATA_FIELD_DEF(uint32_t, actualLenDimsKV) +TILING_DATA_FIELD_DEF(float, softmaxScale) // 即 scaleValue +TILING_DATA_FIELD_DEF(uint32_t, outputLayout) +TILING_DATA_FIELD_DEF(uint64_t, oriMaskMode) +TILING_DATA_FIELD_DEF(int64_t, oriKvStride0) // A2/A3 +TILING_DATA_FIELD_DEF(int64_t, oriWinLeft) +TILING_DATA_FIELD_DEF(int64_t, oriWinRight) +TILING_DATA_FIELD_DEF(int64_t, sparseBlockSize) +TILING_DATA_FIELD_DEF(uint32_t, oriSparseBlockCount) // A5 + +TILING_DATA_FIELD_DEF(int64_t, topkValueMode) + +TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum); + +TILING_DATA_FIELD_DEF(uint32_t, mmResUbSize); +TILING_DATA_FIELD_DEF(uint32_t, bmm2ResUbSize); +TILING_DATA_FIELD_DEF(uint32_t, returnSoftmaxLse) +TILING_DATA_FIELD_DEF(uint32_t, mBaseSize) +TILING_DATA_FIELD_DEF(uint32_t, s2BaseSize) + +TILING_DATA_FIELD_DEF(uint32_t, actualLenDimsOriKV) +TILING_DATA_FIELD_DEF(uint32_t, actualLenDimsCmpKV) +TILING_DATA_FIELD_DEF(uint32_t, cmpResidualKVSize) +TILING_DATA_FIELD_DEF(uint32_t, kvHeadNum) +TILING_DATA_FIELD_DEF(uint32_t, oriKeyStride0) // A5 +TILING_DATA_FIELD_DEF(uint32_t, hasOriSparseIndices) +TILING_DATA_FIELD_DEF(uint32_t, oriSparseIndexWidth) +END_TILING_DATA_DEF +REGISTER_TILING_DATA_CLASS(SparseFlashMlaSwaParamsOp, SparseFlashMlaSwaParams) + +BEGIN_TILING_DATA_DEF(SparseFlashMlaCmpParams) +TILING_DATA_FIELD_DEF(uint32_t, cmpMaxBlockNumPerBatch) +TILING_DATA_FIELD_DEF(uint32_t, sparseBlockCount) // A2/A3 +TILING_DATA_FIELD_DEF(int64_t, cmpRatio) +TILING_DATA_FIELD_DEF(uint64_t, cmpMaskMode) +TILING_DATA_FIELD_DEF(int64_t, cmpKvStride0) // A2/A3 +TILING_DATA_FIELD_DEF(uint32_t, cmpSparseBlockCount) // A5 +TILING_DATA_FIELD_DEF(uint32_t, cmpKvSeqSize) // A5 +TILING_DATA_FIELD_DEF(uint32_t, cmpKeyStride0) // A5 +END_TILING_DATA_DEF +REGISTER_TILING_DATA_CLASS(SparseFlashMlaCmpParamsOp, SparseFlashMlaCmpParams) + +BEGIN_TILING_DATA_DEF(SparseFlashMlaTilingData) +TILING_DATA_FIELD_DEF_STRUCT(SparseFlashMlaSwaParams, baseParams); +TILING_DATA_FIELD_DEF_STRUCT(SparseFlashMlaCmpParams, cmpParams); +END_TILING_DATA_DEF +REGISTER_TILING_DATA_CLASS(SparseFlashMla, SparseFlashMlaTilingData) + +struct SMLAParaInfo { + SMLATilingRequiredParaInfo q = {nullptr, nullptr}; + SMLATilingOptionalParaInfo oriKv = {nullptr, nullptr}; + SMLATilingOptionalParaInfo cmpKv = {nullptr, nullptr}; + SMLATilingOptionalParaInfo oriSparseIndices = {nullptr, nullptr}; + SMLATilingOptionalParaInfo cmpSparseIndices = {nullptr, nullptr}; + SMLATilingOptionalParaInfo oriBlockTable = {nullptr, nullptr}; + SMLATilingOptionalParaInfo cmpBlockTable = {nullptr, nullptr}; + SMLATilingOptionalParaInfo cuSeqLensQ = {nullptr, nullptr}; + SMLATilingOptionalParaInfo cuSeqLensOriKv = {nullptr, nullptr}; // A5 + SMLATilingOptionalParaInfo cuSeqLensCmpKv = {nullptr, nullptr}; // A5 + SMLATilingOptionalParaInfo cuSeqLensKv = {nullptr, nullptr}; // A2/A3 + SMLATilingOptionalParaInfo seqUsedQ = {nullptr, nullptr}; + SMLATilingOptionalParaInfo sequsedOriKv = {nullptr, nullptr}; + SMLATilingOptionalParaInfo sequsedCmpKv = {nullptr, nullptr}; + SMLATilingOptionalParaInfo cmpResidualKv = {nullptr, nullptr}; + SMLATilingOptionalParaInfo oriTopkLength = {nullptr, nullptr}; + SMLATilingOptionalParaInfo cmpTopkLength = {nullptr, nullptr}; + SMLATilingOptionalParaInfo sinks = {nullptr, nullptr}; + SMLATilingOptionalParaInfo metadata = {nullptr, nullptr}; + SMLATilingRequiredParaInfo attnOut = {nullptr, nullptr}; + SMLATilingRequiredParaInfo softmaxLse = {nullptr, nullptr}; + + const float *softmaxScale = nullptr; + const uint32_t *cmpRatio = nullptr; + const uint32_t *oriMaskMode = nullptr; + const uint32_t *cmpMaskMode = nullptr; + const int32_t *oriWinLeft = nullptr; + const int32_t *oriWinRight = nullptr; + const char *layoutQ = nullptr; + const char *layoutKv = nullptr; + const uint32_t *topkValueMode = nullptr; + const bool *returnSoftmaxLse = nullptr; +}; + +static std::string SMLADataTypeToSerialString(ge::DataType type); +std::string SMLALayoutToSerialString(SMLALayout layout); + +// -----------算子Tiling入参信息类--------------- +class SMLATilingInfo { +public: + const char *opName = nullptr; + fe::PlatFormInfos *platformInfo = nullptr; + SMLAParaInfo opParamInfo; + + // Base Param + NpuArch npuArch = NpuArch::DAV_2201; + uint32_t bSize = 0; + uint32_t n1Size = 0; + uint32_t n2Size = 0; + uint32_t s1Size = 0; + int64_t s2Size = 0; + int64_t cmpS2Size = 0; // A5 + uint32_t gSize = 0; + uint32_t qHeadDim = 0; + uint32_t oriKvHeadDim = 0; + uint32_t cmpKvHeadDim = 0; + uint32_t qTSize = 0; // 仅TND时生效 + + uint32_t actualLenDimsQ = 0; + uint32_t actualLenDimsKV = 0; + + uint32_t actualLenDimsOriKV = 0; + uint32_t actualLenDimsCmpKV = 0; + uint32_t cmpResidualKVSize = 0; + + uint32_t oriKeyStride0 = 0; // A5 + uint32_t cmpKeyStride0 = 0; // A5 + + float softmaxScale = 0; + int64_t cmpRatio = 0; + uint64_t oriMaskMode = 0; + uint64_t cmpMaskMode = 0; + uint64_t oriKvStride0 = 0; // A2/A3 + uint64_t cmpKvStride0 = 0; // A2/A3 + int64_t oriWinLeft = 0; + int64_t oriWinRight = 0; + int64_t sparseBlockSize = 0; + int64_t oriSparseBlockCount = 0; + int64_t cmpSparseBlockCount = 0; + int64_t sparseBlockCount = 0; // A2/A3 + bool hasOriSparseIndices = false; + uint32_t oriSparseIndexWidth = 0; + + int64_t topkValueMode = 0; + // Others Flag + bool returnSoftmaxLse = false; + bool batchConsistency = false; + + // PageAttention + uint32_t oriMaxBlockNumPerBatch = 0; + int32_t blockSize = 0; + int32_t oriBlockSize = 0; + int32_t cmpBlockSize = 0; + uint32_t cmpMaxBlockNumPerBatch = 0; + + // DType + ge::DataType qType = ge::DT_FLOAT16; + ge::DataType oriKvType = ge::DT_FLOAT16; + ge::DataType cmpKvType = ge::DT_FLOAT16; + ge::DataType outputType = ge::DT_FLOAT16; + + // Layout + SMLALayout qLayout = SMLALayout::TND; + SMLALayout cmpSparseIndicesLayout = SMLALayout::TND; + SMLALayout oriSparseIndicesLayout = SMLALayout::TND; + SMLALayout kvLayout = SMLALayout::PA_BBND; + SMLALayout outLayout = SMLALayout::BSND; + + // template mode + SMLATemplateMode perfMode = SMLATemplateMode::SWA_TEMPLATE_MODE; +}; + +// -----------算子Tiling入参信息解析及Check类--------------- +class SMLATilingCheck { +public: + explicit SMLATilingCheck(const SMLATilingInfo &smlaInfo) + : smlaInfo_(smlaInfo) {}; + ~SMLATilingCheck() = default; + virtual ge::graphStatus Process(); + +private: + void Init(); + + void LogErrorDtypeSupport(const std::vector &expectDtypeList, const ge::DataType &actualDtype, + const std::string &name) const; + ge::graphStatus CheckLayoutSupport(const SMLALayout &actualLayout, const std::string &name) const; + template + void LogErrorDimNumSupport(const std::vector &expectNumberList, const T &actualValue, + const std::string &name) const; + template + void LogErrorNumberSupport(const std::vector &expectNumberList, const T &actualValue, const std::string &name, + const std::string subName) const; + ge::graphStatus CheckDimNumSupport(const gert::StorageShape *shape, const std::vector &expectDimNumList, + const std::string &name) const; + void LogErrorLayoutSupport(const std::vector &expectLayoutList, const SMLALayout &actualLayout, + const std::string &name) const; + ge::graphStatus CheckDimNumInLayoutSupport(const SMLALayout &layout, const gert::StorageShape *shape, + const std::string &name) const; + ge::graphStatus CheckDtypeSupport(const gert::CompileTimeTensorDesc *desc, const std::string &name) const; + ge::graphStatus CheckSinglePara() const; + ge::graphStatus CheckSingleParaQuery() const; + ge::graphStatus CheckSingleParaOriKv() const; + ge::graphStatus CheckSingleParaCmpKv() const; + ge::graphStatus CheckSingleParaCuSeqLensQ() const; + ge::graphStatus CheckSingleParaCuSeqLensOriKv() const; + ge::graphStatus CheckSingleParaCuSeqLensCmpKv() const; + ge::graphStatus CheckSingleParaCmpResidualKv() const; + ge::graphStatus CheckSingleParaTopkLength() const; + ge::graphStatus CheckSingleParaNumHeads() const; + ge::graphStatus CheckSingleParaKvHeadNums() const; + ge::graphStatus CheckSingleParaOriSparseIndices() const; + ge::graphStatus CheckSingleParaCmpSparseIndices() const; + ge::graphStatus CheckSingleParaSinks() const; + ge::graphStatus CheckSingleParaMetadata() const; + ge::graphStatus CheckSingleParaCmpRatio() const; + ge::graphStatus CheckSingleParaOriMaskMode() const; + ge::graphStatus CheckSingleParaCmpMaskMode() const; + ge::graphStatus CheckSingleParaOriKvStride0() const; // A2/A3 + ge::graphStatus CheckSingleParaCmpKvStride0() const; // A2/A3 + ge::graphStatus CheckSingleParaOriWinLeft() const; + ge::graphStatus CheckSingleParaOriWinRight() const; + ge::graphStatus CheckSingleParaOriBlockTable() const; + ge::graphStatus CheckSingleParaCmpBlockTable() const; + + ge::graphStatus CheckParaExistence() const; + ge::graphStatus CheckExists(const void *pointer, const std::string &name) const; + ge::graphStatus CheckNotExists(const void *pointer, const std::string &name) const; + ge::graphStatus CheckExistsByMap(const std::map ¶mMap) const; + ge::graphStatus CheckNotExistsByMap(const std::map ¶mMap) const; + ge::graphStatus CheckExistenceByMap(std::map &existMap, + std::map ¬ExistMap) const; + + ge::graphStatus CheckFeature() const; + ge::graphStatus CheckFeatureShape() const; + ge::graphStatus CheckFeatureLayout() const; + ge::graphStatus CheckFeatureDtype() const; + ge::graphStatus CheckFeaturePa() const; + + ge::graphStatus CheckMultiParaConsistency(); + ge::graphStatus CheckDTypeConsistency(const ge::DataType &actualDtype, const ge::DataType &expectDtype, + const std::string &name) const; + ge::graphStatus CheckOriAndCmpKv() const; + ge::graphStatus CheckAttenOut() const; + ge::graphStatus CheckActualSeqLensQ() const; + ge::graphStatus CheckActualSeqLens() const; + + const char *opName_; + fe::PlatFormInfos *platformInfo_; + SMLAParaInfo opParamInfo_; + const SMLATilingInfo &smlaInfo_; + + uint32_t bSize_ = 0; + uint32_t n1Size_ = 0; + uint32_t n2Size_ = 0; + uint32_t gSize_ = 0; + uint32_t s1Size_ = 0; + int64_t s2Size_ = 0; + int64_t cmpS2Size_ = 0; // A5 + uint32_t qHeadDim_ = 0; + uint32_t oriKvHeadDim_ = 0; + uint32_t cmpKvHeadDim_ = 0; + + uint32_t qTSize_ = 0; // 仅TND时生效 + int64_t cmpRatio_ = 0; + int64_t oriWinLeft_ = 0; + int64_t oriWinRight_ = 0; + bool hasOriSparseIndices_ = false; + uint32_t oriSparseIndexWidth_ = 0; + + int64_t topkValueMode_; + + SMLALayout qLayout_ = SMLALayout::TND; + SMLALayout cmpSparseIndicesLayout_ = SMLALayout::TND; + SMLALayout oriSparseIndicesLayout_ = SMLALayout::TND; + SMLALayout outLayout_ = SMLALayout::TND; + SMLALayout kvLayout_ = SMLALayout::PA_BBND; + + int32_t oriBlockSize_ = 0; + int32_t cmpBlockSize_ = 0; + + NpuArch npuArch_ = NpuArch::DAV_2201; + + ge::DataType qType_ = ge::DT_FLOAT16; + ge::DataType oriKvType_ = ge::DT_FLOAT16; + ge::DataType cmpKvType_ = ge::DT_FLOAT16; + ge::DataType outputType_ = ge::DT_FLOAT16; +}; + +template +inline T Align(T num, T rnd) +{ + return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd) * (rnd))); +} + +class SMLAInfoParser { +public: + explicit SMLAInfoParser(gert::TilingContext *context) + : context_(context) + {} + ~SMLAInfoParser() = default; + + ge::graphStatus CheckRequiredInOutExistence() const; + ge::graphStatus CheckRequiredAttrExistence() const; + ge::graphStatus CheckRequiredParaExistence() const; + ge::graphStatus CheckUnrequiredParaExistence() const; + + ge::graphStatus GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor, SMLALayout &layout, + const std::string &name) const; + ge::graphStatus GetActualSeqLenQSize(uint32_t &size); + ge::graphStatus GetOpName(); + ge::graphStatus GetNpuInfo(); + void GetOptionalInputParaInfo(); + void GetInputParaInfo(); + void GetOutputParaInfo(); + ge::graphStatus GetAttrParaInfo(); + ge::graphStatus GetOpParaInfo(); + + ge::graphStatus GetInOutDataType(); + ge::graphStatus GetQueryAndOutLayout(); + ge::graphStatus GetKvLayout(); + ge::graphStatus GetSMLATemplateMode(); + void SetSMLAShape(); + ge::graphStatus GetN1Size(); + ge::graphStatus GetN2Size(); + ge::graphStatus GetGSize(); + ge::graphStatus GetBatchSize(); + ge::graphStatus GetQTSize(); + ge::graphStatus GetS1Size(); + ge::graphStatus GetS2SizeForPageAttention(); + ge::graphStatus GetS2SizeForTND(); // A2/A3 + ge::graphStatus GetS2Size(); + ge::graphStatus GetMaxBlockNumPerBatch(); + ge::graphStatus GetBlockSize(); + ge::graphStatus GetQHeadDim(); + ge::graphStatus GetValueHeadDim(); + ge::graphStatus GetSparseBlockCount(); + ge::graphStatus GetActualseqInfo(); + ge::graphStatus GetSinks(); + uint64_t GetOptionalInputStride0(uint32_t inputIndex) const; + void GenerateInfo(SMLATilingInfo &smlaInfo); + ge::graphStatus Parse(SMLATilingInfo &smlaInfo); + std::vector GetKvstride(const gert::Shape &shape, const SMLALayout &layout) const; + ge::graphStatus CheckContiguous() const; // A5 + + gert::TilingContext *context_ = nullptr; + const char *opName_; + fe::PlatFormInfos *platformInfo_; + SMLAParaInfo opParamInfo_; + + bool HasAxis(const SMLAAxis &axis, const SMLALayout &layout, const gert::Shape &shape) const; + size_t GetAxisIdx(const SMLAAxis &axis, const SMLALayout &layout) const; + uint32_t GetAxisNum(const gert::Shape &shape, const SMLAAxis &axis, const SMLALayout &layout) const; + static constexpr uint32_t invalidDimValue_ = std::numeric_limits::min(); + + // BaseParams + uint32_t bSize_ = 0; + uint32_t n1Size_ = 0; + uint32_t n2Size_ = 0; + uint32_t gSize_ = 0; + uint32_t s1Size_ = 0; + int64_t s2Size_ = 0; + int64_t cmpS2Size_ = 0; // A5 + uint32_t qTSize_ = 0; + uint32_t qHeadDim_ = 0; + uint32_t oriKvHeadDim_ = 0; + uint32_t cmpKvHeadDim_ = 0; + int64_t sparseBlockSize_ = 0; + int64_t oriSparseBlockCount_ = 0; // A5 + int64_t cmpSparseBlockCount_ = 0; // A5 + int64_t oriWinLeft_ = 0; + int64_t oriWinRight_ = 0; + bool hasOriSparseIndices_ = false; + uint32_t oriSparseIndexWidth_ = 0; + int64_t topkValueMode_ = 0; + uint32_t actualLenDimsKV_ = 0; + uint32_t actualLenDimsQ_ = 0; + bool batchConsistency_ = false; + + uint32_t actualLenDimsOriKV_ = 0; + uint32_t actualLenDimsCmpKV_ = 0; + uint32_t cmpResidualKVSize_ = 0; + + std::vector oriKeyStridesVec_; // A5 + std::vector cmpKeyStridesVec_; // A5 + + uint32_t aicNum_ = 0; + uint32_t aivNum_ = 0; + // Layout + SMLALayout qLayout_ = SMLALayout::TND; + SMLALayout cmpSparseIndicesLayout_ = SMLALayout::TND; + SMLALayout oriSparseIndicesLayout_ = SMLALayout::TND; + SMLALayout outLayout_ = SMLALayout::BSND; + SMLALayout kvLayout_ = SMLALayout::PA_BBND; + // PageAttention + uint32_t oriMaxBlockNumPerBatch_ = 0; + uint32_t cmpMaxBlockNumPerBatch_ = 0; + int32_t oriBlockSize_ = 0; + int32_t cmpBlockSize_ = 0; + + // template mode + SMLATemplateMode perfMode_ = SMLATemplateMode::SWA_TEMPLATE_MODE; + + NpuArch npuArch_ = NpuArch::DAV_2201; + ge::DataType qType_ = ge::DT_FLOAT16; + ge::DataType oriKvType_ = ge::DT_FLOAT16; + ge::DataType cmpKvType_ = ge::DT_FLOAT16; + ge::DataType cmpSparseIndicesType_ = ge::DT_INT32; + ge::DataType oriBlockTableType_ = ge::DT_INT32; + ge::DataType cmpBlockTableType_ = ge::DT_INT32; + ge::DataType cuSeqLensQType_ = ge::DT_INT32; + ge::DataType seqsedKvType_ = ge::DT_INT32; + ge::DataType sinksType_ = ge::DT_INT32; + ge::DataType metadataType_ = ge::DT_INT32; + ge::DataType outputType_ = ge::DT_FLOAT16; + + gert::Shape qShape_{}; + gert::Shape oriKvShape_{}; + gert::Shape cmpKvShape_{}; + gert::Shape oriSparseIndicesShape_{}; + gert::Shape cmpSparseIndicesShape_{}; +}; + +// ---------------算子Tiling类--------------- +class SparseFlashMlaTiling { +public: + explicit SparseFlashMlaTiling(gert::TilingContext *context) + : context_(context) {}; + ge::graphStatus DoOpTiling(SMLATilingInfo *tilingInfo); + +private: + void SplitBalanced(SMLATilingInfo *tilingInfo); + void CalcUbBmm(const SMLATilingInfo *tilingInfo); + uint32_t CalcFdLogicalSlotCount(const SMLATilingInfo *tilingInfo, uint32_t aicNum) const; + uint64_t CalcFdStagingWorkspaceSize(const SMLATilingInfo *tilingInfo, uint32_t aicNum) const; + uint64_t CalcVectorizeKvPhyAddrWorkspaceSize(const SMLATilingInfo *tilingInfo, uint32_t &vectorizeFlag) const; + gert::TilingContext *context_ = nullptr; + SMLATemplateMode perfMode_ = SMLATemplateMode::SWA_TEMPLATE_MODE; + SparseFlashMlaTilingData tilingData_; + uint32_t blockDim_{0}; + uint64_t workspaceSize_{0}; + uint64_t tilingKey_{0}; + + SMLATilingInfo *smlaInfo_ = nullptr; + + size_t mmResUbSize_ = 0; + size_t bmm2ResUbSize_ = 0; + uint32_t sInnerLoopTimes_ = 0; + uint32_t sInnerSize_ = 512; // s2固定切分512 + uint32_t sInnerSizeAlign_ = 0; + uint32_t usedCoreNum_ = 0; + + uint32_t headDimAlign_ = 0; + uint32_t mBaseSize_ = 64; +}; + +} // namespace optiling +#endif diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_arch22_metadata.h b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_arch22_metadata.h new file mode 100644 index 000000000000..c24cb7559375 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_arch22_metadata.h @@ -0,0 +1,80 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_arch22_metadata.h + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_ARCH22_METADATA_H +#define SPARSE_FLASH_MLA_ARCH22_METADATA_H + +#include + +namespace optiling { + +// Constants +constexpr uint32_t AIC_CORE_NUM = 36; +constexpr uint32_t AIV_CORE_NUM = 72; +constexpr uint32_t SMLA_META_SIZE = 1024; +using SMLA_METADATA_T = int32_t; + +constexpr uint32_t FA_METADATA_SIZE = 9; +constexpr uint32_t FD_METADATA_SIZE = 8; + +// FA Metadata Index Definitions +constexpr uint32_t FA_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FA_BN2_START_INDEX = 1; +constexpr uint32_t FA_M_START_INDEX = 2; +constexpr uint32_t FA_S2_START_INDEX = 3; +constexpr uint32_t FA_BN2_END_INDEX = 4; +constexpr uint32_t FA_M_END_INDEX = 5; +constexpr uint32_t FA_S2_END_INDEX = 6; +constexpr uint32_t FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX = 7; +constexpr uint32_t FA_S2_MAX_NUM = 8; + +// FD Metadata Index Definitions +constexpr uint32_t FD_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FD_BN2_IDX_INDEX = 1; +constexpr uint32_t FD_M_IDX_INDEX = 2; +constexpr uint32_t FD_WORKSPACE_IDX_INDEX = 3; +constexpr uint32_t FD_WORKSPACE_NUM_INDEX = 4; +constexpr uint32_t FD_M_START_INDEX = 5; +constexpr uint32_t FD_M_NUM_INDEX = 6; + +/** + * @brief 获取属性的绝对索引 + * @param coreIdx 核索引 + * @param metaIdx 元数据索引 + * @param isAIV 是否为AIV数据,默认为false + * @return 返回属性的绝对索引 + */ +#ifdef __CCE_AICORE__ +__aicore__ inline uint32_t GetAttrAbsIndex(uint32_t coreIdx, uint32_t metaIdx, bool isAIV = false) +{ + if (isAIV) { + return FA_METADATA_SIZE * AIC_CORE_NUM + FD_METADATA_SIZE * coreIdx + metaIdx; + } else { + return FA_METADATA_SIZE * coreIdx + metaIdx; + } +} +#endif + +namespace detail { +struct SasMetadata { + uint32_t faMetadata[AIC_CORE_NUM][FA_METADATA_SIZE]; + uint32_t fdMetadata[AIV_CORE_NUM][FD_METADATA_SIZE]; +}; +} // namespace detail + +static_assert(SMLA_META_SIZE * sizeof(SMLA_METADATA_T) >= sizeof(detail::SasMetadata)); +} // namespace optiling + +#endif // SPARSE_FLASH_MLA_ARCH22_METADATA_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_common_arch22.h b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_common_arch22.h new file mode 100644 index 000000000000..35cc6e72798b --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_common_arch22.h @@ -0,0 +1,376 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_common_arch22.h + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_COMMON_ARCH22_H +#define SPARSE_FLASH_MLA_COMMON_ARCH22_H + +#include "kernel_operator.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" + +namespace SMLAKernel { +using namespace AscendC; +// 将isCheckTiling设置为false, 输入输出的max&sum&exp的shape为(m, 1) +constexpr SoftmaxConfig SMLA_SOFTMAX_FLASHV2_CFG_WITHOUT_BRC = {false, 0, 0, SoftmaxMode::SOFTMAX_OUTPUT_WITHOUT_BRC}; + +enum class SMLA_RUN_MODE { + SWA_MODE = 0, + CSA_MODE = 1, + HCA_MODE = 2, +}; + +enum class SMLA_LAYOUT { + BSND = 0, + TND = 1, + PA_BBND = 2 +}; + +template +struct SMLAType { + using queryType = Q_T; + using kvType = KV_T; + using outputType = OUT_T; + static constexpr bool flashDecode = FLASH_DECODE; + static constexpr SMLA_LAYOUT layout = LAYOUT_T; + static constexpr SMLA_LAYOUT kvLayout = KV_LAYOUT_T; + static constexpr bool pageAttention = (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND); + static constexpr int templateMode = TEMPLATE_MODE; + static constexpr bool headRatioOne = HEAD_RATIO_ONE; +}; + +// ================================Util functions================================== +template +__aicore__ inline T1 SMLAAlign(T1 num, T2 rnd) +{ + return (rnd == 0) ? 0 : ((num + rnd - 1) / rnd * rnd); +} + +template +__aicore__ inline T1 CeilDiv(T1 num, T2 rnd) +{ + return (rnd == 0) ? 0 : ((num + rnd - 1) / rnd); +} + +template +__aicore__ inline T1 Min(T1 a, T2 b) +{ + return (a > b) ? b : a; +} + +template +__aicore__ inline T1 Max(T1 a, T2 b) +{ + return (a > b) ? a : b; +} + +template +__aicore__ inline size_t BlockAlign(size_t s) +{ + if constexpr (IsSameType::value) { + return (s + 63) / 64 * 64; + } + size_t n = (32 / sizeof(T)); + return (s + n - 1) / n * n; +} + +struct PAShape { + uint32_t blockSize; + uint32_t headNum; // 一般为kv的head num,对应n2 + uint32_t headDim; // 512 对应d + uint32_t kvStride; + uint32_t maxblockNumPerBatch; // block table 每一行的最大个数 + uint32_t actHeadDim; // 实际拷贝col大小,考虑到N切块 s*d, 对应d + uint32_t copyRowNum; // 总共要拷贝的行数 + uint32_t copyRowNumAlign; +}; + +struct Position { + uint32_t bIdx; + uint32_t n2Idx; + uint32_t s2Idx; + uint32_t dIdx; +}; + +// 场景:query、key、value GM to L1 +// GM按ND格式存储 +// L1按NZ格式存储 +// GM的行、列、列的stride +template +__aicore__ inline void DataCopyGmNDToL1(LocalTensor &l1Tensor, GlobalTensor &gmTensor, uint32_t rowAct, + uint32_t rowAlign, + uint32_t col, // D + uint32_t colStride) // D or N*D +{ + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = rowAct; // nd矩阵的行数 + // T为int4场景下,dValue = col / 2,srcDValue = colStride / 2 + nd2nzPara.dValue = col; // nd矩阵的列数 + nd2nzPara.srcDValue = colStride; // 同一nd矩阵相邻行起始地址间的偏移 + nd2nzPara.dstNzC0Stride = rowAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(l1Tensor, gmTensor, nd2nzPara); +} + +/* + 适用PA数据从GM拷贝到L1,支持ND、NZ数据; + PA的layout分 BBND(blockNum,N,blockSize,D) BBH(blockNum,blockSize,N*D + BSH\BSND\TND 为BBH + shape.copyRowNumAlign 需要16字节对齐,如拷贝k矩阵,一次拷贝128*512,遇到尾块 10*512 需对齐到16*512 +*/ +template +__aicore__ inline void DataCopyPA(LocalTensor &dstTensor, // l1 + GlobalTensor &srcTensor, // gm + GlobalTensor &blockTableGm, + const PAShape &shape, // blockSize, headNum, headDim + const Position &startPos) // bacthIdx nIdx curSeqIdx +{ + uint32_t copyFinishRowCnt = 0; + uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; + uint32_t curS2Idx = startPos.s2Idx; + uint32_t blockElementCnt = 32 / sizeof(T); + while (copyFinishRowCnt < shape.copyRowNum) { + uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引 + uint64_t reaminRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数 + uint64_t idInBlockTable = + blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上的获取编号 + uint32_t copyRowCnt = shape.blockSize - reaminRowCnt; // 一次只能处理一个Block + if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) { + copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满 + } + uint64_t offset = idInBlockTable * shape.kvStride; // PA的偏移 + uint64_t dStride = shape.headDim; + if constexpr (SRC_LAYOUT == SMLA_LAYOUT::BSND || SRC_LAYOUT == SMLA_LAYOUT::TND) { + offset += (uint64_t)(startPos.n2Idx * shape.headDim) + reaminRowCnt * shape.headDim * shape.headNum + + startPos.dIdx; + dStride = shape.headDim * shape.headNum; + } else { + offset += (uint64_t)(startPos.n2Idx * shape.headDim * shape.blockSize) + reaminRowCnt * shape.headDim + + startPos.dIdx; + } + + uint32_t dValue = shape.actHeadDim; + uint32_t srcDValue = dStride; + LocalTensor tmpDstTensor = dstTensor[copyFinishRowCnt * blockElementCnt]; + GlobalTensor tmpSrcTensor = srcTensor[offset]; + + DataCopyGmNDToL1(tmpDstTensor, tmpSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue); + copyFinishRowCnt += copyRowCnt; + curS2Idx += copyRowCnt; + } +} + +template +__aicore__ inline void DataCopyPABySlots(LocalTensor &dstTensor, GlobalTensor &srcTensor, + GlobalTensor &blockTableGm, GlobalTensor &sparseIndicesGm, + const PAShape &shape, const Position &startPos, uint64_t sparseIndexBaseOffset, + uint32_t sparseIndexStart) +{ + uint32_t blockElementCnt = 32 / sizeof(T); + uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; + for (uint32_t row = 0; row < shape.copyRowNum; ++row) { + int32_t logicalIdx = sparseIndicesGm.GetValue(sparseIndexBaseOffset + sparseIndexStart + row); + if (logicalIdx < 0) { + continue; + } + uint64_t blockTableIdx = static_cast(logicalIdx) / shape.blockSize; + uint64_t inBlockIdx = static_cast(logicalIdx) % shape.blockSize; + uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockTableIdx); + uint64_t offset = idInBlockTable * shape.kvStride; + offset += static_cast(startPos.n2Idx * shape.headDim * shape.blockSize) + inBlockIdx * shape.headDim + + startPos.dIdx; + + LocalTensor tmpDstTensor = dstTensor[row * blockElementCnt]; + GlobalTensor tmpSrcTensor = srcTensor[offset]; + DataCopyGmNDToL1(tmpDstTensor, tmpSrcTensor, 1, shape.copyRowNumAlign, shape.actHeadDim, shape.headDim); + } +} + +struct RunInfo { + uint32_t loop = 0; + uint32_t cmpLoop = 0; // 用于判断取 用于merge的4块GM 中的哪一块 + uint32_t bIdx = 0; + uint32_t gIdx = 0; + uint32_t s1Idx = 0; + uint32_t s2Idx = 0; + uint32_t n2IdxReal = 0; + uint32_t relativeS2Idx = 0; + uint32_t bn2IdxInCurCore = 0; + uint32_t curSInnerLoopTimes = 0; + uint64_t tndBIdxOffsetForQ = 0; + uint64_t tndBIdxOffsetForKV = 0; + uint64_t tensorCmpBOffset = 0; + uint64_t tensorAOffset = 0; + uint64_t tensorBOffset = 0; + uint64_t attenOutOffset = 0; + uint64_t qTokenOffset = 0; + uint64_t attenMaskOffset = 0; + uint64_t topKBaseOffset = 0; + uint32_t actualSingleProcessSInnerSize = 0; + uint32_t actualSingleProcessSInnerSizeAlign = 0; + uint32_t actualSingleProcessSInnerOriSize = 0; + uint32_t actualSingleProcessSInnerOriAlignSize = 0; + uint32_t actualSingleProcessSInnerCmpSize = 0; + uint32_t actualSingleProcessSInnerCmpAlignSize = 0; + uint32_t s2BatchOffset = 0; + uint32_t gSize = 0; + uint32_t s1Size = 0; + uint32_t s2Size = 0; + uint32_t mSize = 0; + uint32_t mSizeV = 0; + uint32_t mSizeVStart = 0; + uint32_t tndIsS2SplitCore = 0; + uint32_t tndCoreStartKVSplitPos = 0; + bool isFirstSInnerLoop = false; + bool isBmm2Output = false; + bool isValid = false; + bool isLastS2Loop = 0; + int64_t inValidRowCount = 0; + + uint64_t actS1Size = 1; + uint64_t actS2SizeOri = 0ULL; + static constexpr uint32_t n2Idx = 0; + uint32_t gS1Idx = 0; + uint64_t actS2Size = 1; + uint64_t actOriS2Size = 1; + uint32_t actMBaseSize = 0; + int32_t nextTokensPerBatch = 0; + int64_t threshold = 0; + uint64_t curOffsetInSparseBlock = 0; + uint32_t curTopKIdx = 0; + bool isOriOnly = true; // 判断当前块是在Ori部分还是Cmp部分 + bool isOriCmpMix = false; + uint8_t resv[2]; + uint64_t s2StartPoint = 0; + int64_t cmpS2IdStart = 0; + int64_t cmpS2IdLimit = 0; + int32_t v0S2DealSize = 0; + int32_t v0S2Start = 0; + uint32_t oriDealSize = 0; + int32_t cmpMaskRight = 0; +}; + +struct ConstInfo { + // CUBE与VEC核间同步的模式 + static constexpr uint32_t SMLA_SYNC_MODE2 = 2; + // BUFFER的字节数 + static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32; + static constexpr uint32_t BUFFER_SIZE_BYTE_64B = 64; + static constexpr uint32_t BUFFER_SIZE_BYTE_256B = 256; + static constexpr uint32_t BUFFER_SIZE_BYTE_512B = 512; + static constexpr uint32_t BUFFER_SIZE_BYTE_1K = 1024; + static constexpr uint32_t BUFFER_SIZE_BYTE_2K = 2048; + static constexpr uint32_t BUFFER_SIZE_BYTE_4K = 4096; + static constexpr uint32_t BUFFER_SIZE_BYTE_8K = 8192; + static constexpr uint32_t BUFFER_SIZE_BYTE_16K = 16384; + static constexpr uint32_t BUFFER_SIZE_BYTE_32K = 32768; + // FP32的0值和极大值 + static constexpr float FLOAT_ZERO = 0; + static constexpr float FLOAT_MAX = 3.402823466e+38F; + + // preLoad的总次数 + uint32_t preLoadNum = 0U; + uint32_t nBufferMBaseSize = 0U; + // CUBE和VEC的核间同步EventID + uint32_t syncV0C1 = 0U; + uint32_t syncC1V1 = 0U; + uint32_t syncV1C2 = 0U; + uint32_t syncC2V2 = 0U; + + uint32_t mmResUbSize = 0U; // Matmul1输出结果GM上的大小 + uint32_t vec1ResUbSize = 0U; // Vector1输出结果GM上的大小 + uint32_t bmm2ResUbSize = 0U; // Matmul2输出结果GM上的大小 + uint32_t usedCoreNum = 0U; + uint64_t batchSize = 0ULL; + uint64_t gSize = 0ULL; + uint64_t qHeadNum = 0ULL; + uint64_t kvHeadNum = 0; + uint64_t headDim = 0; + uint64_t kvSeqSize = 0ULL; // kv最大S长度 + uint64_t qSeqSize = 1ULL; // q最大S长度 + int64_t kvCacheBlockSize = 0; // PA场景的block size + uint64_t paCmpBlockSize = 0; + uint64_t paOriBlockSize = 0; + int64_t orikvCacheBlockSize = 0; + int64_t cmpkvCacheBlockSize = 0; + uint32_t oriMaxBlockNumPerBatch = 0; // PA场景的最大单batch block number + uint32_t cmpMaxBlockNumPerBatch = 0; + uint32_t splitKVNum = 0U; // S2核间切分的切分份数 + SMLA_LAYOUT outputLayout; // 输出的Transpose格式 + uint32_t oriMaskMode = 0; + uint32_t cmpMaskMode = 0; + uint64_t oriKvStride0 = 0; + uint64_t cmpKvStride0 = 0; + bool needInit = false; + uint32_t templateMode = 0; + + // FlashDecoding + uint32_t actualCombineLoopSize = 0U; // FlashDecoding场景, S2在核间切分的最大份数 + uint64_t combineLseOffset = 0ULL; + uint64_t combineAccumOutOffset = 0ULL; + + uint32_t actualLenDimsQ = 0U; // query的actualSeqLength 的维度 + uint32_t actualLenDimsKV = 0U; // KV 的actualSeqLength 的维度 + + uint32_t actualLenDimsCmpKV = 0U; + uint32_t cmpResidualKVSize = 0U; + + // TND + uint32_t s2Start = 0U; // TND场景下,S2的起始位置 + uint32_t s2End = 0U; // 单核TND场景下S2循环index上限 + + uint32_t bN2Start = 0U; + uint32_t bN2End = 0U; + uint32_t gS1Start = 0U; + uint32_t gS1End = 0U; + + uint32_t tndFDCoreArrLen = 0U; // TNDFlashDecoding相关分核信息array的长度 + uint32_t coreStartKVSplitPos = 0U; // TNDFlashDecoding kv起始位置 + + uint32_t mBaseSize = 1ULL; + uint32_t s2BaseSize = 1ULL; + + // sparse attr + int64_t sparseBlockSize = 0; + uint32_t sparseBlockCount = 0; + bool hasOriSparseIndices = false; + bool hasOriTopkLength = false; + uint32_t oriSparseIndexWidth = 0; + + // cmp attr + int64_t cmpRatio = 0; + + uint64_t cmpSeqSize = 0ULL; + + // win + int32_t oriWinRight = 0; + int32_t oriWinLeft = 128; + + bool returnSoftmaxLse = false; +}; + +struct MSplitInfo { + uint32_t nBufferIdx = 0U; + uint32_t nBufferStartM = 0U; + uint32_t nBufferDealM = 0U; + uint32_t vecStartM = 0U; + uint32_t vecDealM = 0U; +}; +} // namespace SMLAKernel +#endif // SPARSE_FLASH_MLA_COMMON_ARCH22_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_block_cube.h b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_block_cube.h new file mode 100644 index 000000000000..4ddc85b9c796 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_block_cube.h @@ -0,0 +1,864 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_csa_block_cube.h + * \brief use 7 buffer for matmul l1, better pipeline + */ +#ifndef SPARSE_FLASH_MLA_CSA_BLOCK_CUBE_H +#define SPARSE_FLASH_MLA_CSA_BLOCK_CUBE_H + +#include "kernel_operator.h" +#include "kernel_operator_list_tensor_intf.h" +#include "kernel_tiling/kernel_tiling.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" +#include "sparse_flash_mla_common_arch22.h" + +namespace SMLAKernel { +template +class SMLACubeBlock { +public: + // 中间计算数据类型为float, 高精度模式 + using T = float; + using Q_T = typename SMLAT::queryType; + using KV_T = typename SMLAT::kvType; + using OUT_T = typename SMLAT::outputType; + using MM_OUT_T = T; + + __aicore__ inline SMLACubeBlock(){}; + __aicore__ inline void InitParams(const ConstInfo &constInfo); + __aicore__ inline void InitMm1GlobalTensor(GlobalTensor queryGm, GlobalTensor oriKvGm, + GlobalTensor cmpKV, GlobalTensor mm1ResGm); + __aicore__ inline void InitMm2GlobalTensor(GlobalTensor vec1ResGm, GlobalTensor mm2ResGm, + GlobalTensor attentionOutGm); + __aicore__ inline void InitPageAttentionInfo(GlobalTensor oriKvGm, const GlobalTensor &kvMergeGm, + GlobalTensor oriBlockTableGm, + GlobalTensor cmpBlockTableGm); + __aicore__ inline void InitBuffers(TPipe *pipe); + + __aicore__ inline void AllocEventID(); + __aicore__ inline void FreeEventID(); + __aicore__ inline void ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo); + __aicore__ inline void ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo); + +private: + static constexpr bool PAGE_ATTENTION = SMLAT::pageAttention; + static constexpr int TEMPLATE_MODE = SMLAT::templateMode; + static constexpr bool FLASH_DECODE = SMLAT::flashDecode; + static constexpr SMLA_LAYOUT LAYOUT_T = SMLAT::layout; + static constexpr SMLA_LAYOUT KV_LAYOUT_T = SMLAT::kvLayout; + + static constexpr uint32_t MERGE_CACHE_GM_BUF_NUM = 3; + static constexpr uint32_t M_SPLIT_SIZE = 128; // m方向切分 + static constexpr uint32_t N_SPLIT_SIZE = 128; // n方向切分 + static constexpr uint32_t K_L0_SPLIT_SIZE = 128; // k方向L0切分 + static constexpr uint32_t K_L1_SPLIT_SIZE = 256; // k方向L1切分 + static constexpr uint32_t N_WORKSPACE_SIZE = 512; // n方向切分 + static constexpr uint32_t D_SPLIT_SIZE = 256; // d轴切分 + + static constexpr uint32_t L1_BLOCK_SIZE = (64 * 512 * sizeof(Q_T)); + static constexpr uint32_t L1_BLOCK_OFFSET = 64 * 512; + + static constexpr uint32_t L0A_PP_SIZE = (32 * 1024); + static constexpr uint32_t L0B_PP_SIZE = (32 * 1024); + static constexpr uint32_t L0C_PP_SIZE = (64 * 1024); + + // mte2 <> mte1 EventID + // L1 3buf, 使用3个eventId + static constexpr uint32_t L1_EVENT0 = EVENT_ID2; + static constexpr uint32_t L1_EVENT1 = EVENT_ID3; + static constexpr uint32_t L1_EVENT2 = EVENT_ID4; + static constexpr uint32_t L1_EVENT3 = EVENT_ID5; + static constexpr uint32_t L1_EVENT4 = EVENT_ID6; + static constexpr uint32_t L1_EVENT5 = EVENT_ID7; + static constexpr uint32_t L1_EVENT6 = EVENT_ID1; + + // m <> mte1 EventID + static constexpr uint32_t L0AB_EVENT0 = EVENT_ID3; + static constexpr uint32_t L0AB_EVENT1 = EVENT_ID4; + + static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding + static constexpr uint32_t mte21QPIds[4] = {L1_EVENT0, L1_EVENT1, L1_EVENT2, L1_EVENT3}; // mte12复用 + static constexpr uint32_t mte21KVIds[3] = {L1_EVENT4, L1_EVENT5, L1_EVENT6}; + + ConstInfo constInfo{}; + + // L1分成3块buf, 用于记录 + uint32_t qpL1BufIter = 0; + uint32_t kvL1BufIter = -1; + uint32_t abL0BufIter = 0; + uint32_t cL0BufIter = 0; + + // mm1 + GlobalTensor queryGm; + GlobalTensor keyGm; + GlobalTensor mm1ResGm; + GlobalTensor oriKvGm; + GlobalTensor kvMergeGm_; + GlobalTensor cmpKvGm; + + // mm2 + GlobalTensor vec1ResGm; + GlobalTensor valueGm; + GlobalTensor mm2ResGm; + GlobalTensor attentionOutGm; + + // block_table + GlobalTensor oriBlockTableGm; + GlobalTensor cmpBlockTableGm; + + TBuf bufQPL1; + TBuf bufKVL1; + TBuf tmpBufL0A; + TBuf tmpBufL0B; + TBuf tmpBufL0C; + + LocalTensor l1QPTensor; + LocalTensor l1KVTensor; + LocalTensor aL0TensorPingPong; + LocalTensor bL0TensorPingPong; + LocalTensor cL0TensorPingPong; + + // L0AB m <> mte1 EventID + __aicore__ inline uint32_t Mte1MmABEventId(uint32_t idx) + { + return (L0AB_EVENT0 + idx); + } + + __aicore__ inline uint32_t GetQPL1RealIdx(uint32_t mIdx, uint32_t k1Idx) + { + uint32_t idxMap[] = {0, 2}; // 确保0块和1块连在一起, 2和3块连在一起, 来保证同一m块的地址相连 + return idxMap[mIdx % 2] + k1Idx; + } + + __aicore__ inline void CopyGmToL1(LocalTensor &l1Tensor, GlobalTensor &gmSrcTensor, uint32_t srcN, + uint32_t srcD, uint32_t srcDstride); + __aicore__ inline void CopyInMm1AToL1(LocalTensor &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx, + uint32_t mSizeAct, uint32_t headSize, uint32_t headOffset); + + __aicore__ inline void CopyInMm2AToL1(LocalTensor &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx, + uint32_t subMSizeAct, uint32_t nSize, uint32_t nOffset); + __aicore__ inline void LoadDataMm1A(LocalTensor &aL0Tensor, LocalTensor &aL1Tensor, uint32_t idx, + uint32_t kSplitSize, uint32_t mSize, uint32_t kSize); + __aicore__ inline void LoadDataMm1B(LocalTensor &bL0Tensor, LocalTensor &bL1Tensor, uint32_t idx, + uint32_t kSplitSize, uint32_t kSize, uint32_t nSize); +}; + +template +__aicore__ inline void SMLACubeBlock::InitParams(const ConstInfo &constInfo) +{ + this->constInfo = constInfo; +} + +template +__aicore__ inline void SMLACubeBlock::InitMm1GlobalTensor(GlobalTensor queryGm, GlobalTensor oriKvGm, + GlobalTensor cmpKvGm, + GlobalTensor mm1ResGm) +{ + // mm1 + this->queryGm = queryGm; + this->oriKvGm = oriKvGm; + this->cmpKvGm = cmpKvGm; + this->mm1ResGm = mm1ResGm; +} + +template +__aicore__ inline void SMLACubeBlock::InitMm2GlobalTensor(GlobalTensor vec1ResGm, + GlobalTensor mm2ResGm, + GlobalTensor attentionOutGm) +{ + // mm2 + this->vec1ResGm = vec1ResGm; + this->mm2ResGm = mm2ResGm; + this->attentionOutGm = attentionOutGm; +} + +template +__aicore__ inline void SMLACubeBlock::InitPageAttentionInfo(GlobalTensor oriKvGm, + const GlobalTensor &kvMergeGm, + GlobalTensor oriBlockTableGm, + GlobalTensor cmpBlockTableGm) +{ + this->oriKvGm = oriKvGm; + this->kvMergeGm_ = kvMergeGm; + this->oriBlockTableGm = oriBlockTableGm; + this->cmpBlockTableGm = cmpBlockTableGm; +} + +template +__aicore__ inline void SMLACubeBlock::InitBuffers(TPipe *pipe) +{ + pipe->InitBuffer(bufQPL1, L1_BLOCK_SIZE * 4); + l1QPTensor = bufQPL1.Get(); + pipe->InitBuffer(bufKVL1, L1_BLOCK_SIZE * 3); + l1KVTensor = bufKVL1.Get(); + + // L0A + pipe->InitBuffer(tmpBufL0A, L0A_PP_SIZE * 2); // 64K + aL0TensorPingPong = tmpBufL0A.Get(); + // L0B + pipe->InitBuffer(tmpBufL0B, L0B_PP_SIZE * 2); // 64K + bL0TensorPingPong = tmpBufL0B.Get(); + // L0C + pipe->InitBuffer(tmpBufL0C, L0C_PP_SIZE * 2); // 128K + cL0TensorPingPong = tmpBufL0C.Get(); +} + +template +__aicore__ inline void SMLACubeBlock::AllocEventID() +{ + SetFlag(L1_EVENT0); + SetFlag(L1_EVENT1); + SetFlag(L1_EVENT2); + SetFlag(L1_EVENT3); + SetFlag(L1_EVENT4); + SetFlag(L1_EVENT5); + SetFlag(L1_EVENT6); + SetFlag(L0AB_EVENT0); + SetFlag(L0AB_EVENT1); +} + +template +__aicore__ inline void SMLACubeBlock::FreeEventID() +{ + WaitFlag(L1_EVENT0); + WaitFlag(L1_EVENT1); + WaitFlag(L1_EVENT2); + WaitFlag(L1_EVENT3); + WaitFlag(L1_EVENT4); + WaitFlag(L1_EVENT5); + WaitFlag(L1_EVENT6); + WaitFlag(L0AB_EVENT0); + WaitFlag(L0AB_EVENT1); +} + +template +__aicore__ inline void SMLACubeBlock::CopyGmToL1(LocalTensor &l1Tensor, GlobalTensor &gmSrcTensor, + uint32_t srcN, uint32_t srcD, uint32_t srcDstride) +{ + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = srcN; // 行数 + nd2nzPara.dValue = srcD; + nd2nzPara.srcDValue = srcDstride; + nd2nzPara.dstNzC0Stride = (srcN + 15) / 16 * 16; // 对齐到16 单位block + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(l1Tensor, gmSrcTensor, nd2nzPara); +} + +template +__aicore__ inline void SMLACubeBlock::CopyInMm1AToL1(LocalTensor &l1Tensor, const RunInfo &info, + uint32_t mSeqIdx, uint32_t mSizeAct, uint32_t headSize, + uint32_t headOffset) +{ + auto srcGm = queryGm[info.tensorAOffset + mSeqIdx * constInfo.headDim + headOffset]; + CopyGmToL1(l1Tensor, srcGm, mSizeAct, headSize, constInfo.headDim); +} + +template +__aicore__ inline void SMLACubeBlock::LoadDataMm1A(LocalTensor &aL0Tensor, LocalTensor &aL1Tensor, + uint32_t idx, uint32_t kSplitSize, uint32_t mSize, + uint32_t kSize) +{ + LocalTensor srcTensor = aL1Tensor[mSize * kSplitSize * idx]; + LoadData3DParamsV2 loadData3DParams; + // SetFmatrixParams + loadData3DParams.l1H = mSize / 16; // Hin=M1=8 + loadData3DParams.l1W = 16; // Win=M0 + loadData3DParams.padList[0] = 0; + loadData3DParams.padList[1] = 0; + loadData3DParams.padList[2] = 0; + loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果 + + // SetLoadToA0Params + loadData3DParams.mExtension = mSize; // M + loadData3DParams.kExtension = kSize; // K + loadData3DParams.mStartPt = 0; + loadData3DParams.kStartPt = 0; + loadData3DParams.strideW = 1; + loadData3DParams.strideH = 1; + loadData3DParams.filterW = 1; + loadData3DParams.filterSizeW = (1 >> 8) & 255; + loadData3DParams.filterH = 1; + loadData3DParams.filterSizeH = (1 >> 8) & 255; + loadData3DParams.dilationFilterW = 1; + loadData3DParams.dilationFilterH = 1; + loadData3DParams.enTranspose = 0; + loadData3DParams.fMatrixCtrl = 0; + loadData3DParams.channelSize = kSize; // Cin=K + LoadData(aL0Tensor, srcTensor, loadData3DParams); +} + +template +__aicore__ inline void SMLACubeBlock::LoadDataMm1B(LocalTensor &l0Tensor, LocalTensor &l1Tensor, + uint32_t idx, uint32_t kSplitSize, uint32_t kSize, + uint32_t nSize) +{ + // N 方向全载 + LocalTensor srcTensor = l1Tensor[nSize * kSplitSize * idx]; + + LoadData2DParams loadData2DParams; + loadData2DParams.startIndex = 0; + loadData2DParams.repeatTimes = (nSize + 15) / 16 * kSize / (32 / sizeof(KV_T)); + loadData2DParams.srcStride = 1; + loadData2DParams.dstGap = 0; + loadData2DParams.ifTranspose = false; + LoadData(l0Tensor, srcTensor, loadData2DParams); +} + +template +__aicore__ inline void SMLACubeBlock::CopyInMm2AToL1(LocalTensor &aL1Tensor, const RunInfo &info, + uint32_t mSeqIdx, uint32_t subMSizeAct, uint32_t nSize, + uint32_t nOffset) +{ + auto srcGm = vec1ResGm[(info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize + + mSeqIdx * info.actualSingleProcessSInnerSizeAlign + nOffset]; + CopyGmToL1(aL1Tensor, srcGm, subMSizeAct, nSize, info.actualSingleProcessSInnerSizeAlign); +} + +template +__aicore__ inline void SMLACubeBlock::ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo) +{ + uint32_t mSize = mSplitInfo.nBufferDealM; + uint32_t mL1Size = M_SPLIT_SIZE; + uint32_t mL1SizeAlign = SMLAAlign(M_SPLIT_SIZE, 16); + uint32_t mL1Loops = CeilDiv(mSize, M_SPLIT_SIZE); + + uint32_t nSize = info.actualSingleProcessSInnerSize; + uint32_t nL1Size = N_SPLIT_SIZE; + uint32_t nL1SizeAlign = SMLAAlign(N_SPLIT_SIZE, 16); + uint32_t nL1Loops = CeilDiv(nSize, N_SPLIT_SIZE); + + uint32_t kSize = 512; + uint32_t kL1Size = 256; + uint32_t kL1Loops = 2; + uint32_t kL0Size = 128; + uint32_t kL0Loops = CeilDiv(kL1Size, kL0Size); + + LocalTensor bL1Tensor; + uint32_t ka = 0, kb = 0; + + // L1 切n切k + for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { + if (nL1 == (nL1Loops - 1)) { + // 尾块重新计算size + nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE; + nL1SizeAlign = SMLAAlign(nL1Size, 16); + } + + for (uint32_t kL1 = 0; kL1 < kL1Loops; kL1++) { + kvL1BufIter++; + uint32_t kb = kvL1BufIter % 3; + WaitFlag(mte21KVIds[kb]); + // 从k当中取当前的块 + bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET]; + uint32_t curSeqIdx = info.s2BatchOffset + nL1 * N_SPLIT_SIZE; + if (info.isOriOnly) { + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint; + uint32_t copyFinishRowCnt = 0; + LocalTensor kTensor; + uint32_t copyRowCnt = 0; + + while (copyFinishRowCnt < nL1Size) { + // 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位 + copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize; + if (copyFinishRowCnt + copyRowCnt > nL1Size) { + copyRowCnt = nL1Size - copyFinishRowCnt; + } + PAShape shape; + shape.blockSize = constInfo.paOriBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.oriKvStride0; + shape.actHeadDim = D_SPLIT_SIZE; + shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = nL1SizeAlign; + kTensor = bL1Tensor[copyFinishRowCnt * 16]; + + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + startPos.dIdx = + kL1 * D_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + DataCopyPA(kTensor, oriKvGm, oriBlockTableGm, shape, startPos); + + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; // 行数 + nd2nzPara.dValue = D_SPLIT_SIZE; // 256 + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + uint32_t headStride = constInfo.headDim; + uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim; + uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? + static_cast(constInfo.kvSeqSize) * seqStride : + constInfo.oriKvStride0; + + uint32_t curS2 = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE; + DataCopy(bL1Tensor, oriKvGm[offset], nd2nzPara); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint; + if (kL1 == 0) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; + nd2nzPara.dValue = constInfo.headDim >> 1; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor, + oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + + nL1 * N_SPLIT_SIZE * constInfo.headDim], + nd2nzPara); + } else { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; + nd2nzPara.dValue = constInfo.headDim >> 1; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor, + oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + + (constInfo.headDim >> 1) + nL1 * N_SPLIT_SIZE * constInfo.headDim], + nd2nzPara); + } + } + } else { + if (kL1 == 0) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; + nd2nzPara.dValue = constInfo.headDim >> 1; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor, + kvMergeGm_[info.cmpLoop % MERGE_CACHE_GM_BUF_NUM * N_WORKSPACE_SIZE * kSize + + nL1 * N_SPLIT_SIZE * constInfo.headDim], + nd2nzPara); + } else { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; + nd2nzPara.dValue = constInfo.headDim >> 1; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor, + kvMergeGm_[info.cmpLoop % MERGE_CACHE_GM_BUF_NUM * N_WORKSPACE_SIZE * kSize + + (constInfo.headDim >> 1) + nL1 * N_SPLIT_SIZE * constInfo.headDim], + nd2nzPara); + } + } + SetFlag(mte21KVIds[kb]); + WaitFlag(mte21KVIds[kb]); + mL1Size = M_SPLIT_SIZE; + mL1SizeAlign = SMLAAlign(M_SPLIT_SIZE, 16U); + for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) { + uint32_t aL1PaddingSize = 0; // 用于使左矩阵对齐到尾部, 以保证两块32K内存连续 + if (mL1 == (mL1Loops - 1)) { + mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE; + mL1SizeAlign = SMLAAlign(mL1Size, 16U); + aL1PaddingSize = (M_SPLIT_SIZE - mL1SizeAlign) * 256; + } + uint32_t mIdx = qpL1BufIter + mL1; + ka = GetQPL1RealIdx(mIdx, kL1); + LocalTensor aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET + (1 - kL1) * aL1PaddingSize]; + if (nL1 == 0) { + if (kL1 == 0) { + WaitFlag(mte21QPIds[ka]); + WaitFlag(mte21QPIds[ka + 1]); + CopyInMm1AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 256, 0); + } else { + LocalTensor qTmpTensor = aL1Tensor; + CopyInMm1AToL1(qTmpTensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 256, + 256); + } + SetFlag(mte21QPIds[ka]); + WaitFlag(mte21QPIds[ka]); + } + // 使用unitflag同步 + LocalTensor cL0Tensor = + cL0TensorPingPong[(cL0BufIter % 2) * + (L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致 + for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) { + WaitFlag(Mte1MmABEventId(abL0BufIter % 2)); + LocalTensor aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE / sizeof(KV_T))]; + LoadDataMm1A(aL0Tensor, aL1Tensor, kL0, kL0Size, mL1SizeAlign, kL0Size); + LocalTensor bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE / sizeof(KV_T))]; + LoadDataMm1B(bL0Tensor, bL1Tensor, kL0, kL0Size, kL0Size, nL1SizeAlign); + SetFlag(Mte1MmABEventId(abL0BufIter % 2)); + WaitFlag(Mte1MmABEventId(abL0BufIter % 2)); + + MmadParams mmadParams; + mmadParams.m = mL1SizeAlign; + mmadParams.n = nL1SizeAlign; + mmadParams.k = kL0Size; + mmadParams.cmatrixInitVal = (kL1 == 0 && kL0 == 0); + mmadParams.cmatrixSource = false; + mmadParams.unitFlag = + (kL1 == 1 && kL0 == (kL0Loops - 1)) ? 0b11 : 0b10; // 累加最后一次翻转flag, 表示可以搬出 + Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams); + if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) { + PipeBarrier(); + } + SetFlag(Mte1MmABEventId(abL0BufIter % 2)); + abL0BufIter++; + } + + if (nL1 == (nL1Loops - 1)) { + SetFlag(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完 + } + + if (kL1 == 1) { // 最后一轮kL1循环 + FixpipeParamsV220 fixParams; + fixParams.nSize = nL1SizeAlign; + fixParams.mSize = mL1SizeAlign; + fixParams.srcStride = mL1SizeAlign; + // 改成nSizeAlign + fixParams.dstStride = info.actualSingleProcessSInnerSizeAlign; // mm1ResGm两行之间的间隔 + fixParams.unitFlag = 0b11; + fixParams.ndNum = 1; // 输出ND + + Fixpipe(mm1ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.mmResUbSize + nL1 * N_SPLIT_SIZE + + (mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) * + info.actualSingleProcessSInnerSizeAlign], + cL0Tensor, fixParams); + } + if (mL1Loops == 2) { + cL0BufIter++; + } + } + + SetFlag(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完 + } + if (mL1Loops == 1) { + cL0BufIter++; + } + } + qpL1BufIter += mL1Loops; +} + +template +__aicore__ inline void SMLACubeBlock::ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo) +{ + uint32_t mSize = mSplitInfo.nBufferDealM; + uint32_t mSizeAlign = (mSize + 16 - 1) / 16; + uint32_t mL1Loops = (mSize + M_SPLIT_SIZE - 1) / M_SPLIT_SIZE; + uint32_t mL1SizeAlign = M_SPLIT_SIZE; // 16对齐 + uint32_t mL1Size = M_SPLIT_SIZE; // m的实际大小 + + uint32_t nSize = BlockAlign(constInfo.headDim); + uint32_t nL1Loops = (nSize + N_SPLIT_SIZE - 1) / N_SPLIT_SIZE; + uint32_t nL1SizeAlign = N_SPLIT_SIZE; // 16对齐 + uint32_t nL1Size = N_SPLIT_SIZE; // n的实际大小 + + uint32_t kSize = info.actualSingleProcessSInnerSize; + uint32_t kL1Size = 256; + uint32_t kL1SizeAlign = SMLAAlign(kL1Size, 16U); + uint32_t kL1Loops = (kSize + kL1Size - 1) / kL1Size; + uint32_t kL0Size = 128; + uint32_t kL0Loops = (kL1Size + kL0Size - 1) / kL0Size; + uint32_t kL0SizeAlign = kL0Size; + LocalTensor bL1Tensor; + LocalTensor subvTensor; + + // ka表示左矩阵4buf选择哪一块buf, kb表示右矩阵3buf选择哪一块buf + uint32_t ka = 0, kb = 0; + uint32_t mBaseIdx = qpL1BufIter; + for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // n切L1 -> D + if (nL1 == (nL1Loops - 1)) { + // 尾块 + nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE; + nL1SizeAlign = SMLAAlign(nL1Size, 16U); + } + // k l1写成一个循环, 和mm1保持一致 + kL1Size = 256; + kL1SizeAlign = SMLAAlign(kL1Size, 16U); + uint32_t copyRowCnt = 0; + + for (uint32_t k1 = 0; k1 < kL1Loops; k1++) { // k切L1, 这里套了一层l0来操作 -> S2,每次256 + if (k1 == (kL1Loops - 1)) { + // 尾块 + kL1Size = kSize - (kL1Loops - 1) * 256; + kL1SizeAlign = SMLAAlign(kL1Size, 16U); + } + kvL1BufIter++; + uint32_t kb = kvL1BufIter % 3; + WaitFlag(mte21KVIds[kb]); + bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET]; + uint32_t kOffset = k1 * kL0Loops; + kL0Size = 128; + // 此处必须先初始化kL0Size, 再求kL0Loops, 否则由于循环会改变kL0Size大小, 导致kL0Loops错误 + kL0Loops = (kL1Size + kL0Size - 1) / kL0Size; + kL0SizeAlign = kL0Size; + for (uint32_t kL1 = kOffset; kL1 < kL0Loops + kOffset; kL1++) { // 128 循环搬pa,每次128 + if (kL1 == kOffset + kL0Loops - 1) { + // 尾块 + kL0Size = kL1Size - (kL0Loops - 1) * kL0Size; + kL0SizeAlign = SMLAAlign(kL0Size, 16U); + } + + uint32_t curSeqIdx = info.s2BatchOffset + (kL1 - kOffset) * 128 + k1 * 256; + if (info.isOriOnly) { + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + uint32_t copyFinishRowCnt = 0; + uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint; + while (copyFinishRowCnt < kL0Size) { + copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize; + if (copyFinishRowCnt + copyRowCnt > kL0Size) { + copyRowCnt = kL0Size - copyFinishRowCnt; + } + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + startPos.dIdx = + nL1 * N_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + PAShape shape; + shape.blockSize = constInfo.paOriBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.oriKvStride0; + shape.actHeadDim = nL1Size; + shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = kL0SizeAlign; + subvTensor = bL1Tensor[(kL1 - kOffset) * 128 * N_SPLIT_SIZE + copyFinishRowCnt * 16]; + + DataCopyPA(subvTensor, oriKvGm, oriBlockTableGm, shape, startPos); + + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + uint32_t headStride = constInfo.headDim; + uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim; + uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? + static_cast(constInfo.kvSeqSize) * seqStride : + constInfo.oriKvStride0; + + uint32_t curS2 = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE; + subvTensor = bL1Tensor[(kL1 - kOffset) * 128 * N_SPLIT_SIZE]; + DataCopy(subvTensor, oriKvGm[offset], nd2nzPara); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint; + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor[(kL1 - kOffset) * 128 * N_SPLIT_SIZE], + oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + + kL1 * 128 * constInfo.headDim + nL1 * N_SPLIT_SIZE], + nd2nzPara); + } + } else { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor[(kL1 - kOffset) * 128 * N_SPLIT_SIZE], + kvMergeGm_[info.cmpLoop % MERGE_CACHE_GM_BUF_NUM * N_WORKSPACE_SIZE * 512 + + kL1 * 128 * constInfo.headDim + nL1 * N_SPLIT_SIZE], + nd2nzPara); + } + } + SetFlag(mte21KVIds[kb]); + WaitFlag(mte21KVIds[kb]); + mL1SizeAlign = M_SPLIT_SIZE; + mL1Size = M_SPLIT_SIZE; // m的实际大小 + for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) { + if (mL1 == (mL1Loops - 1)) { + // 尾块 + mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE; + mL1SizeAlign = SMLAAlign(mL1Size, 16U); + } + + uint32_t mIdx = mBaseIdx + mL1; + ka = GetQPL1RealIdx(mIdx, k1); + LocalTensor aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET]; + if (nL1 == 0) { + WaitFlag(mte21QPIds[ka]); + CopyInMm2AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, kL1Size, + 256 * k1); + SetFlag(mte21QPIds[ka]); + WaitFlag(mte21QPIds[ka]); + } + + LocalTensor cL0Tensor = + cL0TensorPingPong[(cL0BufIter % 2) * + (L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致 + uint32_t baseK = 128; + uint32_t baseN = 128; + kL0Size = 128; + kL0SizeAlign = kL0Size; + for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) { + if (kL0 + 1 == kL0Loops) { + kL0Size = kL1Size - (kL0Loops - 1) * kL0Size; + kL0SizeAlign = SMLAAlign(kL0Size, 16U); + } + WaitFlag(Mte1MmABEventId(abL0BufIter % 2)); + LocalTensor bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE / sizeof(KV_T))]; + LoadData3DParamsV2 loadData3DParamsForB; + loadData3DParamsForB.l1H = kL0SizeAlign / 16; // 源操作数height + loadData3DParamsForB.l1W = 16; // 源操作数weight=16,目的height=l1H*L1W + loadData3DParamsForB.padList[0] = 0; + loadData3DParamsForB.padList[1] = 0; + loadData3DParamsForB.padList[2] = 0; + loadData3DParamsForB.padList[3] = 255; // 尾部数据不影响滑窗的结果 + + loadData3DParamsForB.mExtension = kL0SizeAlign; // 在目的操作数height维度的传输长度 + loadData3DParamsForB.kExtension = nL1SizeAlign; // 在目的操作数width维度的传输长度 + loadData3DParamsForB.mStartPt = 0; // 卷积核在目的操作数width维度的起点 + loadData3DParamsForB.kStartPt = 0; // 卷积核在目的操作数height维度的起点 + loadData3DParamsForB.strideW = 1; + loadData3DParamsForB.strideH = 1; + loadData3DParamsForB.filterW = 1; + loadData3DParamsForB.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素 + loadData3DParamsForB.filterH = 1; + loadData3DParamsForB.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素 + loadData3DParamsForB.dilationFilterW = 1; // 卷积核width膨胀系数 + loadData3DParamsForB.dilationFilterH = 1; // 卷积核height膨胀系数 + loadData3DParamsForB.enTranspose = 1; // 是否启用转置功能 + loadData3DParamsForB.fMatrixCtrl = + 0; // 使用FMATRIX_LEFT还是使用FMATRIX_RIGHT,=0使用FMATRIX_LEFT,=1使用FMATRIX_RIGHT 1 + loadData3DParamsForB.channelSize = + nL1SizeAlign; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize + LoadData(bL0Tensor, bL1Tensor[kL0 * baseK * baseN], loadData3DParamsForB); + + LocalTensor aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE / sizeof(KV_T))]; + LoadData3DParamsV2 loadData3DParamsForA; + loadData3DParamsForA.l1H = mL1SizeAlign / 16; // 源操作数height + loadData3DParamsForA.l1W = 16; // 源操作数weight + loadData3DParamsForA.padList[0] = 0; + loadData3DParamsForA.padList[1] = 0; + loadData3DParamsForA.padList[2] = 0; + loadData3DParamsForA.padList[3] = 255; // 尾部数据不影响滑窗的结果 + + loadData3DParamsForA.mExtension = mL1SizeAlign; // 在目的操作数height维度的传输长度 + loadData3DParamsForA.kExtension = kL0SizeAlign; // 在目的操作数width维度的传输长度 + loadData3DParamsForA.mStartPt = 0; // 卷积核在目的操作数width维度的起点 + loadData3DParamsForA.kStartPt = 0; // 卷积核在目的操作数height维度的起点 + loadData3DParamsForA.strideW = 1; // 卷积核在源操作数width维度滑动的步长 + loadData3DParamsForA.strideH = 1; // 卷积核在源操作数height维度滑动的步长 + loadData3DParamsForA.filterW = 1; // 卷积核width + loadData3DParamsForA.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素 + loadData3DParamsForA.filterH = 1; // 卷积核height + loadData3DParamsForA.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素 + loadData3DParamsForA.dilationFilterW = 1; // 卷积核width膨胀系数 + loadData3DParamsForA.dilationFilterH = 1; // 卷积核height膨胀系数 + loadData3DParamsForA.enTranspose = 0; // 是否启用转置功能,对整个目标矩阵进行转置 + loadData3DParamsForA.fMatrixCtrl = 0; + loadData3DParamsForA.channelSize = + kL0SizeAlign; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize + LoadData(aL0Tensor, aL1Tensor[kL0 * baseK * mL1SizeAlign], + loadData3DParamsForA); + SetFlag(Mte1MmABEventId(abL0BufIter % 2)); + WaitFlag(Mte1MmABEventId(abL0BufIter % 2)); + + MmadParams mmadParams; + mmadParams.m = mL1SizeAlign; + mmadParams.n = nL1SizeAlign; + mmadParams.k = kL0Size; + mmadParams.cmatrixInitVal = (kL0 == 0 && k1 == 0); + mmadParams.cmatrixSource = false; + mmadParams.unitFlag = ((k1 == (kL1Loops - 1)) && (kL0 == (kL0Loops - 1))) ? 0b11 : 0b10; + + Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams); + if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) { + PipeBarrier(); + } + SetFlag(Mte1MmABEventId(abL0BufIter % 2)); + abL0BufIter++; + } + + if (nL1 == (nL1Loops - 1)) { // nL1最后一轮, 需要将B驻留在L1中, 用于下一轮的计算? + SetFlag(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完 + } + + if (k1 == (kL1Loops - 1)) { + // ND + FixpipeParamsV220 fixParams; + fixParams.nSize = nL1SizeAlign; + fixParams.mSize = mL1SizeAlign; + fixParams.srcStride = mL1SizeAlign; + fixParams.dstStride = nSize; // mm2ResGm两行之间的间隔 + fixParams.ndNum = 1; // 输出ND + fixParams.unitFlag = 0b11; + + uint64_t mm2Offset = (mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) * nSize + nL1 * N_SPLIT_SIZE; + Fixpipe(mm2ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.bmm2ResUbSize + mm2Offset], + cL0Tensor, fixParams); + } + + if (mL1Loops == 2) { + cL0BufIter++; + } + } + SetFlag(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完 + } + // cL0BufIter已经不在使用 + if (mL1Loops == 1) { + cL0BufIter++; + } + } + qpL1BufIter += mL1Loops; +} +} // namespace SMLAKernel +#endif // SPARSE_FLASH_MLA_CSA_BLOCK_CUBE_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_block_vector.h b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_block_vector.h new file mode 100644 index 000000000000..092628c852fc --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_block_vector.h @@ -0,0 +1,1141 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_csa_block_vector.h + * \brief + */ +#ifndef SPARSE_FLASH_MLA_CSA_BLOCK_VECTOR_H +#define SPARSE_FLASH_MLA_CSA_BLOCK_VECTOR_H + +#include "kernel_operator.h" +#include "kernel_operator_list_tensor_intf.h" +#include "kernel_tiling/kernel_tiling.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" +#include "sparse_flash_mla_common_arch22.h" + +namespace SMLAKernel { +using AscendC::CrossCoreSetFlag; +using AscendC::CrossCoreWaitFlag; + +template +class SMLAVectorBlock { +public: + // 中间计算数据类型为float,高精度模式 + using T = float; + using KV_T = typename SMLAT::kvType; + using OUT_T = typename SMLAT::outputType; + using UPDATE_T = T; + using SINKS_T = T; + using MM1_OUT_T = float; + using MM2_OUT_T = float; + + __aicore__ inline SMLAVectorBlock(){}; + __aicore__ inline void ProcessVec0L(const RunInfo &runInfo); + __aicore__ inline void ProcessVec1L(const RunInfo &info); + __aicore__ inline void ProcessVec2L(const RunInfo &info); + __aicore__ inline void InitBuffers(TPipe *pipe); + __aicore__ inline void InitParams(const struct ConstInfo &constInfo, + const SparseFlashMlaTilingData *__restrict tilingData); + __aicore__ inline void InitVec0GlobalTensor(const GlobalTensor &kvMergeGm, const GlobalTensor &oriKvGm, + const GlobalTensor &cmpKvGm, + const GlobalTensor &oriBlockTableGm, + const GlobalTensor &cmpBlockTableGm); + __aicore__ inline void InitVec1GlobalTensor(GlobalTensor mm1ResGm, GlobalTensor vec1ResGm, + GlobalTensor actualSeqLengthsQGm, + GlobalTensor actualSeqLengthsKVGm, + GlobalTensor topKGm, GlobalTensor sinksGm, + GlobalTensor softmaxLseGm); + __aicore__ inline void InitVec2GlobalTensor(GlobalTensor accumOutGm, GlobalTensor vec2ResGm, + GlobalTensor mm2ResGm, GlobalTensor attentionOutGm); + __aicore__ inline void AllocEventID(); + __aicore__ inline void FreeEventID(); + __aicore__ inline void CopySinksIn(); + __aicore__ inline void SliceAndContactSinksValue(uint32_t nIdx, uint32_t dealRowCount); + __aicore__ inline void InitSoftmaxDefaultBuffer(); + // ================================Base Vector========================================== + __aicore__ inline void RowDivs(LocalTensor dstUb, LocalTensor src0Ub, LocalTensor src1Ub, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + __aicore__ inline void RowMuls(LocalTensor dstUb, LocalTensor src0Ub, LocalTensor src1Ub, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + // ================================Vector0========================================== + __aicore__ inline int64_t GetKeyGmOffset(int64_t realS2Idx, const RunInfo &runInfo, int64_t s2IdLimit); + __aicore__ inline void GetRealS2Idx(int64_t s2GmOffset, int64_t &realS2Idx, int64_t topkGmBaseOffset, + const RunInfo &runInfo); + __aicore__ inline void CopyInKv(int64_t &mte2Size, int64_t mte3Size, int64_t mergeMte3Idx, int64_t realS2Idx1, + int64_t realS2Idx2, const RunInfo &runInfo); + __aicore__ inline void CopyOutMrgeResult(int64_t mte2Size, int64_t mte3Size, int64_t s2StartGmOffset, + int64_t mergeMte3Idx, const RunInfo &runInfo); + __aicore__ inline void CopyInSingleKv(int64_t &mte2Size, int64_t mte3Size, int64_t mergeMte3Idx, int64_t realS2Idx, + int64_t keyBNBOffset, int64_t s2IdLimit, const RunInfo &runInfo); + // ================================Vector1========================================== + __aicore__ inline void ProcessVec1SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo); + __aicore__ inline void DealBmm1ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t startRow, + uint32_t dealRowCount, uint32_t columnCount, uint32_t loopId); + __aicore__ inline void SoftmaxFlashV2Compute(const RunInfo &info, const MSplitInfo &mSplitInfo, + LocalTensor &mmResUb, LocalTensor &softmaxTmpUb, + uint32_t startRow, uint32_t dealRowCount, uint32_t columnCount, + uint32_t actualColumnCount); + + __aicore__ inline void ElewiseCompute(const RunInfo &info, const LocalTensor &mmResUb, uint32_t dealRowCount, + uint32_t columnCount); + __aicore__ inline void ProcessLse(const RunInfo &info, const MSplitInfo &mSplitInfo); + // ================================Vecotr2========================================== + __aicore__ inline void ProcessVec2SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo); + __aicore__ inline void DealBmm2ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t startRow, + uint32_t dealRowCount, uint32_t columnCount, + uint32_t actualColumnCount); + __aicore__ inline void ProcessVec2Inner(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t mStartRow, + uint32_t mDealSize); + __aicore__ inline void Bmm2DataCopyOutTrans(const RunInfo &info, LocalTensor &attenOutUb, uint32_t wsMStart, + uint32_t dealRowCount, uint32_t columnCount, + uint32_t actualColumnCount); + __aicore__ inline void Bmm2ResCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, uint32_t wsMStart, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + __aicore__ inline void Bmm2CastAndCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, uint32_t wsMStart, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + __aicore__ inline void Bmm2FDDataCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, uint32_t wsMStart, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + __aicore__ inline uint64_t CalcAccumOffset(uint32_t bN2Idx, uint32_t gS1Idx); + + // BLOCK和REPEAT的字节数 + static constexpr uint64_t BYTE_BLOCK = 32UL; + static constexpr uint32_t REPEAT_BLOCK_BYTE = 256U; + // BLOCK和REPEAT的FP32元素数 + static constexpr uint32_t FP32_BLOCK_ELEMENT_NUM = BYTE_BLOCK / sizeof(float); + static constexpr uint32_t FP32_REPEAT_ELEMENT_NUM = REPEAT_BLOCK_BYTE / sizeof(float); + // repeat stride不能超过256 + static constexpr uint32_t REPEATE_STRIDE_UP_BOUND = 256; + +private: + static constexpr bool PAGE_ATTENTION = SMLAT::pageAttention; + static constexpr int TEMPLATE_MODE = SMLAT::templateMode; + static constexpr bool FLASH_DECODE = SMLAT::flashDecode; + static constexpr SMLA_LAYOUT LAYOUT_T = SMLAT::layout; + static constexpr SMLA_LAYOUT KV_LAYOUT_T = SMLAT::kvLayout; + static constexpr bool HEAD_RATIO_ONE = SMLAT::headRatioOne; + + static constexpr uint64_t MERGE_CACHE_GM_BUF_NUM = 3; + static constexpr uint64_t SYNC_INPUT_BUF1_FLAG = 2; + static constexpr uint64_t SYNC_INPUT_BUF1_PONG_FLAG = 3; + static constexpr uint64_t SYNC_INPUT_BUF2_FLAG = 4; + static constexpr uint64_t SYNC_INPUT_BUF2_PONG_FLAG = 5; + static constexpr uint64_t SYNC_OUTPUT_BUF1_FLAG = 4; + static constexpr uint64_t SYNC_OUTPUT_BUF2_FLAG = 5; + static constexpr uint64_t SYNC_SINKS_BUF_FLAG = 6; + static constexpr uint64_t SYNC_INPUT_V0BUF_FLAG = 7; + static constexpr uint32_t INPUT1_BUFFER_OFFSET = ConstInfo::BUFFER_SIZE_BYTE_32K; + static constexpr uint32_t INPUT2_BUFFER_OFFSET = ConstInfo::BUFFER_SIZE_BYTE_16K; + static constexpr uint32_t SOFTMAX_TMP_BUFFER_OFFSET = ConstInfo::BUFFER_SIZE_BYTE_1K; + static constexpr uint32_t BASE_BLOCK_MAX_ELEMENT_NUM = ConstInfo::BUFFER_SIZE_BYTE_32K / sizeof(T); // 32768/4=8096 + static constexpr uint32_t BLOCK_ELEMENT_NUM = BYTE_BLOCK / sizeof(T); // 32/4=8 + static constexpr uint32_t MAX_N1_SIZE = 128U; + static constexpr T SOFTMAX_MIN_NUM = -2e38; + static constexpr SINKS_T R0 = 1.0f; + + const SparseFlashMlaTilingData *__restrict tilingData; + + uint32_t pingpongFlag = 0U; + ConstInfo constInfo = {}; + + GlobalTensor mm1ResGm; + GlobalTensor vec1ResGm; + GlobalTensor softmaxMaxGm; + GlobalTensor softmaxSumGm; + GlobalTensor sinksGm; + + GlobalTensor actualSeqLengthsQGm; + GlobalTensor actualSeqLengthsKVGm; + GlobalTensor vec2ResGm; + GlobalTensor mm2ResGm; + GlobalTensor accumOutGm; + GlobalTensor attentionOutGm; + GlobalTensor softmaxLseGm; + + GlobalTensor blkTableGm_; + GlobalTensor kvMergeGm_; + GlobalTensor keyGm_; + GlobalTensor topkGm_; + GlobalTensor oriKvGm_; + GlobalTensor cmpKvGm_; + GlobalTensor oriBlockTableGm_; + GlobalTensor cmpBlockTableGm_; + + // ================================Local Buffer区==================================== + TBuf<> inputBuff1; // 32K + TBuf<> inputBuff2; // 16K + TBuf<> outputBuff1; // 32K + TBuf<> outputBuff2; // 32K + + TBuf<> tmpBuff1; // 32K + TBuf<> v0ValidSizeBuff; // 8K + + TBuf<> sinksBuff; // 1K + TBuf<> sinksBrcbBuff; // 12K + + TBuf<> softmaxMaxBuff; // PRE_LOAD_NUM * 2K + TBuf<> softmaxExpBuff; // PRE_LOAD_NUM * 2K + TBuf<> softmaxSumBuff; // PRE_LOAD_NUM * 2K + TBuf<> softmaxMaxDefaultBuff; // 2K + TBuf<> softmaxSumDefaultBuff; // 2K + + LocalTensor softmaxMaxDefaultUb; + LocalTensor softmaxSumDefaultUb; + + LocalTensor softmaxMaxUb; + LocalTensor softmaxSumUb; + LocalTensor softmaxExpUb; + LocalTensor kvMergUb_; + LocalTensor v0ValidSizeUb_; + LocalTensor sinksUb; + LocalTensor sinksBrcbUb; + + uint32_t mergeMte3Idx = 0; +}; + +template +__aicore__ inline void SMLAVectorBlock::InitBuffers(TPipe *pipe) +{ + pipe->InitBuffer(inputBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K * 2); // 2:pingpong + pipe->InitBuffer(inputBuff2, ConstInfo::BUFFER_SIZE_BYTE_16K * 2); // 2:pingpong + pipe->InitBuffer(outputBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K); + if (constInfo.returnSoftmaxLse) { + pipe->InitBuffer(outputBuff2, ConstInfo::BUFFER_SIZE_BYTE_1K); + } + + pipe->InitBuffer(tmpBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K); + pipe->InitBuffer(v0ValidSizeBuff, ConstInfo::BUFFER_SIZE_BYTE_8K); + + // M_MAX = 512/2vector = 256, 256 * sizeof(T) * N_Buffer + + pipe->InitBuffer(softmaxMaxBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum); + pipe->InitBuffer(softmaxExpBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum); + pipe->InitBuffer(softmaxSumBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum); + + pipe->InitBuffer(softmaxMaxDefaultBuff, ConstInfo::BUFFER_SIZE_BYTE_1K); + pipe->InitBuffer(softmaxSumDefaultBuff, ConstInfo::BUFFER_SIZE_BYTE_1K); + + pipe->InitBuffer(sinksBuff, MAX_N1_SIZE * sizeof(SINKS_T)); + // 分配256+N1大小内存,其中256是m轴VEC最大切块 + pipe->InitBuffer(sinksBrcbBuff, MAX_N1_SIZE * sizeof(SINKS_T) * BLOCK_ELEMENT_NUM * 3U); + + softmaxMaxUb = softmaxMaxBuff.Get(); + softmaxSumUb = softmaxSumBuff.Get(); + softmaxExpUb = softmaxExpBuff.Get(); + + softmaxMaxDefaultUb = softmaxMaxDefaultBuff.Get(); + softmaxSumDefaultUb = softmaxSumDefaultBuff.Get(); + + kvMergUb_ = inputBuff2.Get(); + + v0ValidSizeUb_ = v0ValidSizeBuff.Get(); + + sinksUb = sinksBuff.Get(); + sinksBrcbUb = sinksBrcbBuff.Get(); +} + +template +__aicore__ inline void SMLAVectorBlock::InitParams(const struct ConstInfo &constInfo, + const SparseFlashMlaTilingData *__restrict tilingData) +{ + this->constInfo = constInfo; + this->tilingData = tilingData; +} + +template +__aicore__ inline void SMLAVectorBlock::InitVec0GlobalTensor(const GlobalTensor &kvMergeGm, + const GlobalTensor &oriKvGm, + const GlobalTensor &cmpKvGm, + const GlobalTensor &oriBlockTableGm, + const GlobalTensor &cmpBlockTableGm) +{ + this->kvMergeGm_ = kvMergeGm; + this->oriKvGm_ = oriKvGm; + this->cmpKvGm_ = cmpKvGm; + this->oriBlockTableGm_ = oriBlockTableGm; + this->cmpBlockTableGm_ = cmpBlockTableGm; +} + +template +__aicore__ inline void SMLAVectorBlock::InitVec1GlobalTensor( + GlobalTensor mm1ResGm, GlobalTensor vec1ResGm, GlobalTensor actualSeqLengthsQGm, + GlobalTensor actualSeqLengthsKVGm, GlobalTensor topKGm, GlobalTensor sinksGm, + GlobalTensor softmaxLseGm) +{ + this->mm1ResGm = mm1ResGm; + this->vec1ResGm = vec1ResGm; + this->actualSeqLengthsQGm = actualSeqLengthsQGm; + this->actualSeqLengthsKVGm = actualSeqLengthsKVGm; + this->topkGm_ = topKGm; + this->sinksGm = sinksGm; + this->softmaxLseGm = softmaxLseGm; +} + +template +__aicore__ inline void SMLAVectorBlock::InitVec2GlobalTensor(GlobalTensor accumOutGm, + GlobalTensor vec2ResGm, + GlobalTensor mm2ResGm, + GlobalTensor attentionOutGm) +{ + this->accumOutGm = accumOutGm; + this->vec2ResGm = vec2ResGm; + this->mm2ResGm = mm2ResGm; + this->attentionOutGm = attentionOutGm; +} + +template +__aicore__ inline void SMLAVectorBlock::AllocEventID() +{ + SetFlag(SYNC_INPUT_BUF1_FLAG); + SetFlag(SYNC_INPUT_BUF1_PONG_FLAG); + SetFlag(SYNC_INPUT_BUF2_FLAG); + SetFlag(SYNC_INPUT_BUF2_PONG_FLAG); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + SetFlag(SYNC_OUTPUT_BUF2_FLAG); +} + +template +__aicore__ inline void SMLAVectorBlock::FreeEventID() +{ + WaitFlag(SYNC_INPUT_BUF1_FLAG); + WaitFlag(SYNC_INPUT_BUF1_PONG_FLAG); + WaitFlag(SYNC_INPUT_BUF2_FLAG); + WaitFlag(SYNC_INPUT_BUF2_PONG_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF2_FLAG); +} + +template +__aicore__ inline void SMLAVectorBlock::CopySinksIn() +{ + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1U; + dataCopyParams.blockLen = constInfo.qHeadNum * sizeof(T); + dataCopyParams.srcStride = 0U; + dataCopyParams.dstStride = 0U; + DataCopyPadExtParams padParams; + DataCopyPad(sinksUb, sinksGm, dataCopyParams, padParams); + SetFlag(SYNC_SINKS_BUF_FLAG); + WaitFlag(SYNC_SINKS_BUF_FLAG); + uint32_t repeatTimes = (constInfo.qHeadNum + BLOCK_ELEMENT_NUM - 1U) / BLOCK_ELEMENT_NUM; // 每次处理 8 datablocks + Brcb(sinksBrcbUb, sinksUb, repeatTimes, {1, BLOCK_ELEMENT_NUM}); + PipeBarrier(); + + DataCopyParams repeatParams; + repeatParams.blockCount = 1; // 搬到有一个块超过单个vec核减分核M轴大小即可,核间切分每个vec256 + repeatParams.blockLen = constInfo.qHeadNum; + repeatParams.srcStride = 0U; + repeatParams.dstStride = 0U; + for (uint32_t i = 1U; i <= 256U / constInfo.qHeadNum; i++) { + DataCopy(sinksBrcbUb[constInfo.qHeadNum * BLOCK_ELEMENT_NUM * i], sinksBrcbUb, repeatParams); + } + PipeBarrier(); +} + +template +__aicore__ inline void SMLAVectorBlock::SliceAndContactSinksValue(uint32_t nIdx, uint32_t dealRowCount) +{ + // WholeReduceMax接口中repeatTimes支持范围(0,255),因此需要分多次调用WholeReduceMax,每次repeatTime=128 + uint32_t repeatTimesOnce = 128; + uint32_t loopTimes = (dealRowCount + repeatTimesOnce - 1) / repeatTimesOnce; + uint32_t repeatTimes = repeatTimesOnce; + + for (uint32_t loop = 0; loop < loopTimes; ++loop) { + if (loop == loopTimes - 1) { + repeatTimes = dealRowCount - loop * repeatTimesOnce; + } + WholeReduceMax(softmaxMaxDefaultUb[loop * repeatTimesOnce], + sinksBrcbUb[(nIdx + loop * repeatTimesOnce) * BLOCK_ELEMENT_NUM], + BLOCK_ELEMENT_NUM * BLOCK_ELEMENT_NUM, repeatTimes, 1, 0, 1, ReduceOrder::ORDER_ONLY_VALUE); + PipeBarrier(); + } +} + +template +__aicore__ inline void SMLAVectorBlock::InitSoftmaxDefaultBuffer() +{ + CopySinksIn(); + Duplicate(softmaxMaxDefaultUb, SOFTMAX_MIN_NUM, SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T)); + Duplicate(softmaxSumDefaultUb, R0, SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T)); +} + +template +__aicore__ inline void SMLAVectorBlock::ElewiseCompute(const RunInfo &info, const LocalTensor &mmResUb, + uint32_t dealRowCount, uint32_t columnCount) +{ + Muls(mmResUb, mmResUb, static_cast(tilingData->baseParams.softmaxScale), dealRowCount * columnCount); +} + +template +__aicore__ inline void SMLAVectorBlock::ProcessLse(const RunInfo &info, const MSplitInfo &mSplitInfo) +{ + if (mSplitInfo.vecDealM == 0) { + return; + } + uint64_t lseOffset; + if (constInfo.outputLayout == SMLA_LAYOUT::TND) { + uint32_t tBase = actualSeqLengthsQGm.GetValue(info.bIdx); + lseOffset = (tBase + info.s1Idx) * constInfo.gSize + // T轴、s1轴偏移 + info.n2IdxReal * constInfo.qSeqSize * constInfo.gSize; // N2轴偏移 + } else if (constInfo.outputLayout == SMLA_LAYOUT::BSND) { + lseOffset = info.bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize + // B轴偏移 + info.n2IdxReal * constInfo.qSeqSize * constInfo.gSize + // N2轴偏移 + info.s1Idx * constInfo.gSize; // S1轴偏移 + } + lseOffset = lseOffset + mSplitInfo.nBufferStartM + mSplitInfo.vecStartM; + uint32_t baseOffset = mSplitInfo.nBufferStartM / 2; + uint32_t outIdx = info.loop % (constInfo.preLoadNum); + uint32_t softmaxOffset = outIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset; + auto sumTensor = softmaxSumUb[softmaxOffset]; + auto maxTensor = softmaxMaxUb[softmaxOffset]; + auto outLSETensor = outputBuff2.Get(); + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1; + dataCopyParams.blockLen = mSplitInfo.vecDealM * sizeof(T); + dataCopyParams.srcStride = 0; + dataCopyParams.dstStride = 0; + + WaitFlag(SYNC_OUTPUT_BUF2_FLAG); + PipeBarrier(); + Log(outLSETensor, sumTensor, mSplitInfo.vecDealM); + PipeBarrier(); + Add(outLSETensor, outLSETensor, maxTensor, mSplitInfo.vecDealM); + SetFlag(SYNC_OUTPUT_BUF2_FLAG); + WaitFlag(SYNC_OUTPUT_BUF2_FLAG); + + DataCopyPad(softmaxLseGm[lseOffset], outLSETensor, dataCopyParams); + SetFlag(SYNC_OUTPUT_BUF2_FLAG); +} + +template +__aicore__ inline void SMLAVectorBlock::SoftmaxFlashV2Compute(const RunInfo &info, const MSplitInfo &mSplitInfo, + LocalTensor &mmResUb, + LocalTensor &softmaxTmpUb, + uint32_t startRow, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + LocalTensor inSumTensor; + LocalTensor inMaxTensor; + uint32_t baseOffset = mSplitInfo.nBufferStartM / 2 + startRow; + uint32_t outIdx = info.loop % (constInfo.preLoadNum); + uint32_t softmaxOutOffset = outIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset; + if (info.isFirstSInnerLoop) { + inMaxTensor = softmaxMaxDefaultUb[startRow]; + inSumTensor = softmaxSumDefaultUb; + } else { + uint32_t inIdx = (info.loop - 1) % (constInfo.preLoadNum); + inMaxTensor = softmaxMaxUb[inIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset]; + inSumTensor = softmaxSumUb[inIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset]; + } + if (actualColumnCount != 0) { + SoftMaxShapeInfo srcShape{dealRowCount, columnCount, dealRowCount, actualColumnCount}; + SoftMaxTiling newTiling = + SoftMaxFlashV2TilingFunc(srcShape, sizeof(T), sizeof(T), softmaxTmpUb.GetSize(), true, false); + SoftmaxFlashV2( + mmResUb, softmaxSumUb[softmaxOutOffset], softmaxMaxUb[softmaxOutOffset], mmResUb, + softmaxExpUb[softmaxOutOffset], inSumTensor, inMaxTensor, softmaxTmpUb, newTiling, srcShape); + } else { + uint32_t dealRowCountAlign = SMLAAlign(dealRowCount, FP32_BLOCK_ELEMENT_NUM); + DataCopy(softmaxSumUb[softmaxOutOffset], inSumTensor, dealRowCountAlign); + PipeBarrier(); + DataCopy(softmaxMaxUb[softmaxOutOffset], inMaxTensor, dealRowCountAlign); + } +} + +template +__aicore__ inline void SMLAVectorBlock::DealBmm1ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, + uint32_t startRow, uint32_t dealRowCount, + uint32_t columnCount, uint32_t loopId) +{ + uint32_t computeSize = dealRowCount * columnCount; + uint64_t inOutGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize + + (mSplitInfo.nBufferStartM + mSplitInfo.vecStartM + startRow) * columnCount; + LocalTensor mmResUb = inputBuff1.Get(); + mmResUb = mmResUb[pingpongFlag * INPUT1_BUFFER_OFFSET / sizeof(MM1_OUT_T)]; + WaitFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + + DataCopy(mmResUb, mm1ResGm[inOutGmOffset], computeSize); + SetFlag(SYNC_INPUT_BUF1_FLAG); + WaitFlag(SYNC_INPUT_BUF1_FLAG); + + ElewiseCompute(info, mmResUb, dealRowCount, columnCount); + + PipeBarrier(); + LocalTensor tmpAFloorUb = tmpBuff1.Get(); + LocalTensor softmaxTmpUb = tmpAFloorUb.template ReinterpretCast(); + + SoftmaxFlashV2Compute(info, mSplitInfo, mmResUb, softmaxTmpUb, startRow, dealRowCount, columnCount, + info.actualSingleProcessSInnerSize); + + PipeBarrier(); + LocalTensor tmpMMResCastTensor = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + + Cast(tmpMMResCastTensor, mmResUb, AscendC::RoundMode::CAST_ROUND, computeSize); + SetFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + pingpongFlag ^= 1; // pingpong 0 1 切换 + + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + DataCopy(vec1ResGm[inOutGmOffset], tmpMMResCastTensor, computeSize); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); +} + +template +__aicore__ inline void SMLAVectorBlock::ProcessVec1SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo) +{ + if (mSplitInfo.vecDealM == 0) { + return; + } + uint32_t mSplitSize = info.actualSingleProcessSInnerSize == 0 ? + 16 : + BASE_BLOCK_MAX_ELEMENT_NUM / info.actualSingleProcessSInnerSizeAlign; + // 1. 向下8对齐是因为UB操作至少32B + // 2. info.actualSingleProcessSInnerSizeAlign最大512, mSplitSize可以确保最小为16 + mSplitSize = mSplitSize / 8 * 8; + + if (mSplitSize > mSplitInfo.vecDealM) { + mSplitSize = mSplitInfo.vecDealM; + } + uint32_t loopCount = (mSplitInfo.vecDealM + mSplitSize - 1) / mSplitSize; + uint32_t tailSplitSize = mSplitInfo.vecDealM - (loopCount - 1) * mSplitSize; + + uint32_t sinkHeadIdx = + (info.n2IdxReal * constInfo.gSize + mSplitInfo.nBufferStartM + mSplitInfo.vecStartM) % constInfo.qHeadNum; + SliceAndContactSinksValue(sinkHeadIdx, mSplitInfo.vecDealM); + + for (uint32_t i = 0, dealSize = mSplitSize; i < loopCount; i++) { + if (i == (loopCount - 1)) { + dealSize = tailSplitSize; + } + DealBmm1ResBaseBlock(info, mSplitInfo, i * mSplitSize, dealSize, info.actualSingleProcessSInnerSizeAlign, i); + } +} + +template +__aicore__ inline void SMLAVectorBlock::GetRealS2Idx(int64_t s2GmOffset, int64_t &realS2Idx, + int64_t topkGmBaseOffset, const RunInfo &runInfo) +{ + int64_t cmpS2Offset = s2GmOffset; + int64_t topkGmIdx = cmpS2Offset / constInfo.sparseBlockSize; + if (unlikely(topkGmIdx >= constInfo.sparseBlockCount || s2GmOffset >= runInfo.v0S2DealSize)) { + realS2Idx = -1; + return; + } + realS2Idx = topkGm_.GetValue(topkGmBaseOffset + topkGmIdx) * static_cast(constInfo.sparseBlockSize) + + static_cast(cmpS2Offset % constInfo.sparseBlockSize); +} + +template +__aicore__ inline int64_t SMLAVectorBlock::GetKeyGmOffset(int64_t realS2Idx, const RunInfo &runInfo, + int64_t s2IdLimit) +{ + if (realS2Idx < 0 || realS2Idx >= s2IdLimit) { + return -1; + } + int64_t realKeyGmOffset = 0; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + int64_t blkTableIdx = realS2Idx / constInfo.paCmpBlockSize; + int64_t blkTableOffset = realS2Idx % constInfo.paCmpBlockSize; + realKeyGmOffset = + cmpBlockTableGm_.GetValue(runInfo.bIdx * constInfo.cmpMaxBlockNumPerBatch + blkTableIdx) * + static_cast(constInfo.cmpKvStride0) + + blkTableOffset * static_cast(constInfo.kvHeadNum) * static_cast(constInfo.headDim); + + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + int64_t batchStride = (constInfo.cmpKvStride0 == 0) ? static_cast(constInfo.cmpSeqSize) * + static_cast(constInfo.kvHeadNum) * + static_cast(constInfo.headDim) : + static_cast(constInfo.cmpKvStride0); + realKeyGmOffset = + static_cast(runInfo.bIdx) * batchStride + + realS2Idx * static_cast(constInfo.kvHeadNum) * static_cast(constInfo.headDim) + + static_cast(runInfo.n2Idx) * static_cast(constInfo.headDim); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + realKeyGmOffset = + runInfo.tensorCmpBOffset + + realS2Idx * static_cast(constInfo.kvHeadNum) * static_cast(constInfo.headDim) + + static_cast(runInfo.n2Idx) * static_cast(constInfo.headDim); + } + return realKeyGmOffset; +} + +template +__aicore__ inline void SMLAVectorBlock::CopyInSingleKv(int64_t &mte2Size, int64_t mte3Size, int64_t mergeMte3Idx, + int64_t realS2Idx, int64_t keyBNBOffset, + int64_t s2IdLimit, const RunInfo &runInfo) +{ + if (keyBNBOffset < 0) { + return; + } + int64_t validS2Count = + (realS2Idx + constInfo.sparseBlockSize > s2IdLimit ? s2IdLimit - realS2Idx : constInfo.sparseBlockSize); + DataCopyExtParams intriParams; + intriParams.blockLen = validS2Count * constInfo.headDim * sizeof(KV_T); + intriParams.blockCount = 1; + intriParams.dstStride = 0; + intriParams.srcStride = 0; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + DataCopyPad( + kvMergUb_[mergeMte3Idx % 2 * INPUT2_BUFFER_OFFSET / sizeof(KV_T) + (mte2Size - mte3Size) * constInfo.headDim], + cmpKvGm_[keyBNBOffset], intriParams, padParams); + mte2Size += validS2Count; +} + +template +__aicore__ inline void SMLAVectorBlock::CopyInKv(int64_t &mte2Size, int64_t mte3Size, int64_t mergeMte3Idx, + int64_t realS2Idx1, int64_t realS2Idx2, const RunInfo &runInfo) +{ + int64_t s2IdLimit = runInfo.cmpS2IdLimit; + + int64_t keyOffset1 = GetKeyGmOffset(realS2Idx1, runInfo, s2IdLimit); + int64_t keyOffset2 = GetKeyGmOffset(realS2Idx2, runInfo, s2IdLimit); + if (unlikely(keyOffset1 < 0 && keyOffset2 < 0)) { + return; + } + + int64_t keySrcStride = 0; + keySrcStride = ((keyOffset1 > keyOffset2 ? (keyOffset1 - keyOffset2) : (keyOffset2 - keyOffset1)) - + constInfo.sparseBlockSize * constInfo.headDim) * + sizeof(KV_T); + if (unlikely(keySrcStride >= INT32_MAX || keySrcStride < 0 || realS2Idx1 + constInfo.sparseBlockSize >= s2IdLimit || + realS2Idx2 + constInfo.sparseBlockSize >= s2IdLimit)) { + // stride溢出、stride为负数、s2超长等异常场景,还原成2条搬运指令 + // 因为需要拷贝两块 + CopyInSingleKv(mte2Size, mte3Size, mergeMte3Idx, realS2Idx1, keyOffset1, s2IdLimit, runInfo); + CopyInSingleKv(mte2Size, mte3Size, mergeMte3Idx, realS2Idx2, keyOffset2, s2IdLimit, runInfo); + } else { + DataCopyExtParams intriParams; + intriParams.blockLen = constInfo.sparseBlockSize * constInfo.headDim * sizeof(KV_T); + intriParams.blockCount = (keyOffset1 >= 0) + (keyOffset2 >= 0); + intriParams.dstStride = 0; + intriParams.srcStride = keySrcStride; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + + int64_t startGmOffset = keyOffset1 > -1 ? keyOffset1 : keyOffset2; + if (keyOffset2 > -1 && keyOffset2 < keyOffset1) { + startGmOffset = keyOffset2; + } + DataCopyPad(kvMergUb_[mergeMte3Idx % 2 * INPUT2_BUFFER_OFFSET / sizeof(KV_T) + + (mte2Size - mte3Size) * constInfo.headDim], + cmpKvGm_[startGmOffset], intriParams, padParams); + mte2Size += ((keyOffset1 > -1) + (keyOffset2 > -1)) * constInfo.sparseBlockSize; + } +} + +template +__aicore__ inline void SMLAVectorBlock::CopyOutMrgeResult(int64_t mte2Size, int64_t mte3Size, + int64_t s2GmStartOffset, int64_t mergeMte3Idx, + const RunInfo &runInfo) +{ + if (mte2Size <= mte3Size) { + return; + } + SetFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + WaitFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = mte2Size - mte3Size; + dataCopyParams.blockLen = constInfo.headDim * sizeof(KV_T); + dataCopyParams.srcStride = 0; + dataCopyParams.dstStride = 0; + + DataCopyPad(kvMergeGm_[runInfo.cmpLoop % MERGE_CACHE_GM_BUF_NUM * 512 * 512 + + (s2GmStartOffset + mte3Size) * constInfo.headDim], + kvMergUb_[mergeMte3Idx % 2 * INPUT2_BUFFER_OFFSET / sizeof(KV_T)], dataCopyParams); +} + +// b s1 k +template +__aicore__ inline void SMLAVectorBlock::ProcessVec0L(const RunInfo &runInfo) +{ + int64_t s2ProcessSize = runInfo.v0S2DealSize; + int64_t s2Pair = CeilDiv(s2ProcessSize, 2 * constInfo.sparseBlockSize); + int64_t topkGmBaseOffset = runInfo.topKBaseOffset + runInfo.v0S2Start; + int64_t mte2Size = 0; + int64_t mte3Size = 0; + int64_t s2IdxArray0 = -1; + int64_t s2IdxArray1 = -1; + bool needWaitMte3ToMte2 = true; + int64_t s2SplitPoint = SMLAAlign(s2Pair, 2) * constInfo.sparseBlockSize; + int64_t s2GmStartOffset = GetSubBlockIdx() == 0 ? 0 : s2SplitPoint; + int64_t s2GmLimit = GetSubBlockIdx() == 0 ? s2SplitPoint : s2ProcessSize; + if (s2GmLimit > s2ProcessSize) { + s2GmLimit = s2ProcessSize; + } + // 处理两个基本块 + for (int64_t s2GmOffsetArray = s2GmStartOffset; s2GmOffsetArray < s2GmLimit; + s2GmOffsetArray += 2 * constInfo.sparseBlockSize) { + if (needWaitMte3ToMte2) { + WaitFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + needWaitMte3ToMte2 = false; + } + GetRealS2Idx(s2GmOffsetArray, s2IdxArray0, topkGmBaseOffset, runInfo); + if (unlikely(s2IdxArray0 < 0)) { + CopyOutMrgeResult(mte2Size, mte3Size, s2GmStartOffset, mergeMte3Idx, runInfo); + SetFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + mergeMte3Idx++; + break; + } + GetRealS2Idx(s2GmOffsetArray + constInfo.sparseBlockSize, s2IdxArray1, topkGmBaseOffset, runInfo); + CopyInKv(mte2Size, mte3Size, mergeMte3Idx, s2IdxArray0, s2IdxArray1, runInfo); + if ((mte2Size - mte3Size + 2 * constInfo.sparseBlockSize > 16) || + s2GmOffsetArray + 2 * constInfo.sparseBlockSize >= s2GmLimit) { + CopyOutMrgeResult(mte2Size, mte3Size, s2GmStartOffset, mergeMte3Idx, runInfo); + mte3Size = mte2Size; + SetFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + mergeMte3Idx++; + needWaitMte3ToMte2 = true; + } + } + return; +} + +template +__aicore__ inline void SMLAVectorBlock::ProcessVec1L(const RunInfo &info) +{ + uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize; + uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize; + for (uint32_t i = 0; i < nBufferLoopTimes; i++) { + MSplitInfo mSplitInfo; + mSplitInfo.nBufferIdx = i; + mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize; + mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail; + + mSplitInfo.vecDealM = (mSplitInfo.nBufferDealM <= 16) ? mSplitInfo.nBufferDealM : + (((mSplitInfo.nBufferDealM + 15) / 16 + 1) / 2 * 16); + mSplitInfo.vecStartM = 0; + if (GetBlockIdx() % 2 == 1) { + mSplitInfo.vecStartM = mSplitInfo.vecDealM; + mSplitInfo.vecDealM = mSplitInfo.nBufferDealM - mSplitInfo.vecDealM; + } + + if constexpr (HEAD_RATIO_ONE) { + CrossCoreWaitFlag(constInfo.syncC1V1); + } else { + CrossCoreWaitFlag(constInfo.syncC1V1); + } + // vec1 compute + ProcessVec1SingleBuf(info, mSplitInfo); + CrossCoreSetFlag(constInfo.syncV1C2); + + // move lse for flash decode or FA + if (constInfo.returnSoftmaxLse && info.s2Idx == info.curSInnerLoopTimes - 1) { + ProcessLse(info, mSplitInfo); + } + } +} + +template +__aicore__ inline uint64_t SMLAVectorBlock::CalcAccumOffset(uint32_t bN2Idx, uint32_t gS1Idx) +{ + return 0; +} + +template +__aicore__ inline void SMLAVectorBlock::ProcessVec2SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo) +{ + if (mSplitInfo.vecDealM == 0) { + return; + } + + ProcessVec2Inner(info, mSplitInfo, 0, mSplitInfo.vecDealM); +} + +template +__aicore__ inline void SMLAVectorBlock::ProcessVec2L(const RunInfo &info) +{ + uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize; + uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize; + for (uint32_t i = 0; i < nBufferLoopTimes; i++) { + MSplitInfo mSplitInfo; + mSplitInfo.nBufferIdx = i; + mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize; + mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail; + + mSplitInfo.vecDealM = (mSplitInfo.nBufferDealM <= 16) ? mSplitInfo.nBufferDealM : + (((mSplitInfo.nBufferDealM + 15) / 16 + 1) / 2 * 16); + mSplitInfo.vecStartM = 0; + if (GetBlockIdx() % 2 == 1) { + mSplitInfo.vecStartM = mSplitInfo.vecDealM; + mSplitInfo.vecDealM = mSplitInfo.nBufferDealM - mSplitInfo.vecDealM; + } + if constexpr (HEAD_RATIO_ONE) { + CrossCoreWaitFlag(constInfo.syncC2V2); + } else { + CrossCoreWaitFlag(constInfo.syncC2V2); + } + ProcessVec2SingleBuf(info, mSplitInfo); + } +} + +template +__aicore__ inline void SMLAVectorBlock::ProcessVec2Inner(const RunInfo &info, const MSplitInfo &mSplitInfo, + uint32_t mStartRow, uint32_t mDealSize) +{ + uint32_t mSplitSize = BASE_BLOCK_MAX_ELEMENT_NUM / constInfo.headDim; + if (mSplitSize > mDealSize) { + mSplitSize = mDealSize; + } + + uint32_t loopCount = (mDealSize + mSplitSize - 1) / mSplitSize; + uint32_t tailSplitSize = mDealSize - (loopCount - 1) * mSplitSize; + for (uint32_t i = 0, dealSize = mSplitSize; i < loopCount; i++) { + if (i == (loopCount - 1)) { + dealSize = tailSplitSize; + } + DealBmm2ResBaseBlock(info, mSplitInfo, i * mSplitSize + mStartRow, dealSize, constInfo.headDim, + constInfo.headDim); + } +} + +template +__aicore__ inline void SMLAVectorBlock::Bmm2FDDataCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, + uint32_t wsMStart, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + LocalTensor tmp = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + DataCopy(tmp, bmm2ResUb, columnCount * dealRowCount); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + uint64_t accumTmpOutNum = CalcAccumOffset(info.bIdx, info.gS1Idx); + uint64_t offset = + accumTmpOutNum * constInfo.kvHeadNum * constInfo.mBaseSize * constInfo.headDim + // taskoffset + info.tndCoreStartKVSplitPos * constInfo.kvHeadNum * constInfo.mBaseSize * constInfo.headDim + // 份数offset + wsMStart * actualColumnCount; // m轴offset + GlobalTensor dst = accumOutGm[offset]; + if (info.actualSingleProcessSInnerSize == 0) { + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = dealRowCount; + dataCopyParams.blockLen = actualColumnCount * sizeof(T); + dataCopyParams.srcStride = (columnCount - actualColumnCount) / (BYTE_BLOCK / sizeof(T)); + dataCopyParams.dstStride = 0; + DataCopyPad(dst, tmp, dataCopyParams); + } else { + matmul::InitOutput(dst, dealRowCount * actualColumnCount, ConstInfo::FLOAT_ZERO); + } + SetFlag(SYNC_OUTPUT_BUF1_FLAG); +} + +template +__aicore__ inline void SMLAVectorBlock::Bmm2DataCopyOutTrans(const RunInfo &info, LocalTensor &attenOutUb, + uint32_t wsMStart, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = dealRowCount; + dataCopyParams.blockLen = actualColumnCount * sizeof(OUT_T); + dataCopyParams.srcStride = (columnCount - actualColumnCount) / (BYTE_BLOCK / sizeof(OUT_T)); + dataCopyParams.dstStride = 0; + DataCopyPad(attentionOutGm[info.attenOutOffset + wsMStart * actualColumnCount], attenOutUb, dataCopyParams); + return; +} + +template +__aicore__ inline void SMLAVectorBlock::Bmm2CastAndCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, + uint32_t wsMStart, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + LocalTensor tmpBmm2ResCastTensor = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + if constexpr (IsSameType::value) { // bf16 采取四舍六入五成双模式 + Cast(tmpBmm2ResCastTensor, bmm2ResUb, AscendC::RoundMode::CAST_RINT, dealRowCount * columnCount); + } else { + Cast(tmpBmm2ResCastTensor, bmm2ResUb, AscendC::RoundMode::CAST_ROUND, dealRowCount * columnCount); + } + + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + Bmm2DataCopyOutTrans(info, tmpBmm2ResCastTensor, wsMStart, dealRowCount, columnCount, actualColumnCount); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); +} + +template +__aicore__ inline void SMLAVectorBlock::Bmm2ResCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, + uint32_t wsMStart, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + if constexpr (FLASH_DECODE) { + if (info.tndIsS2SplitCore) { + Bmm2FDDataCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount); + } else { + Bmm2CastAndCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount); + } + } else { + Bmm2CastAndCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount); + } +} + +template +__aicore__ inline void SMLAVectorBlock::DealBmm2ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, + uint32_t startRow, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + uint32_t vec2ComputeSize = dealRowCount * columnCount; + uint32_t mStart = mSplitInfo.nBufferStartM + mSplitInfo.vecStartM + startRow; + uint64_t srcGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + mStart * columnCount; + LocalTensor tmpBmm2ResUb = inputBuff1.Get(); + tmpBmm2ResUb = tmpBmm2ResUb[pingpongFlag * INPUT1_BUFFER_OFFSET / sizeof(MM2_OUT_T)]; + WaitFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + DataCopy(tmpBmm2ResUb, mm2ResGm[srcGmOffset], vec2ComputeSize); + + SetFlag(SYNC_INPUT_BUF1_FLAG); + WaitFlag(SYNC_INPUT_BUF1_FLAG); + + LocalTensor bmm2ResUb = tmpBuff1.Get(); + bmm2ResUb.SetSize(vec2ComputeSize); + DataCopy(bmm2ResUb, tmpBmm2ResUb, vec2ComputeSize); + SetFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + pingpongFlag ^= 1; // pingpong 0 1切换 + + uint32_t inOutBaseOffset = mStart * columnCount; + uint32_t baseOffset = mSplitInfo.nBufferStartM / 2 + startRow; + + if constexpr (!HEAD_RATIO_ONE) { + if (!info.isFirstSInnerLoop) { + event_t eventIdMte2WaitMte3 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2)); + SetFlag(eventIdMte2WaitMte3); + WaitFlag(eventIdMte2WaitMte3); + + LocalTensor bmm2ResPreUb = inputBuff1.Get(); + bmm2ResPreUb = bmm2ResPreUb[pingpongFlag * INPUT1_BUFFER_OFFSET / sizeof(MM2_OUT_T)]; + WaitFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + + uint64_t vec2ResGmOffset = + ((info.loop - 1) % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + inOutBaseOffset; + DataCopy(bmm2ResPreUb, vec2ResGm[vec2ResGmOffset], vec2ComputeSize); + + SetFlag(SYNC_INPUT_BUF1_FLAG); + WaitFlag(SYNC_INPUT_BUF1_FLAG); + + uint32_t idx = info.loop % (constInfo.preLoadNum); + LocalTensor expUb = v0ValidSizeBuff.Get()[384]; + Brcb(expUb, softmaxExpUb[idx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset], (dealRowCount + 7) / 8, + {1, 8}); + PipeBarrier(); + + RowMuls(bmm2ResPreUb, bmm2ResPreUb, expUb, dealRowCount, columnCount, actualColumnCount); + AscendC::PipeBarrier(); + Add(bmm2ResUb, bmm2ResUb, bmm2ResPreUb, vec2ComputeSize); + AscendC::PipeBarrier(); + + SetFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + pingpongFlag ^= 1; + } + + if (info.isLastS2Loop) { + uint32_t idx = info.loop % (constInfo.preLoadNum); + LocalTensor tmpSumUb = v0ValidSizeBuff.Get()[384]; + Brcb(tmpSumUb, softmaxSumUb[idx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset], + (dealRowCount + 7) / 8, {1, 8}); + PipeBarrier(); + RowDivs(bmm2ResUb, bmm2ResUb, tmpSumUb, dealRowCount, columnCount, actualColumnCount); + PipeBarrier(); + Bmm2ResCopyOut(info, bmm2ResUb, mStart, dealRowCount, columnCount, actualColumnCount); + } else { + LocalTensor outUb = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + DataCopy(outUb, bmm2ResUb, dealRowCount * columnCount); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + uint64_t vec2ResGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + inOutBaseOffset; + DataCopy(vec2ResGm[vec2ResGmOffset], outUb, vec2ComputeSize); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + } + } else { + // 除第一个循环外,均需要更新中间计算结果 + if (!info.isFirstSInnerLoop) { + event_t eventIdMte2WaitMte3 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2)); + SetFlag(eventIdMte2WaitMte3); + WaitFlag(eventIdMte2WaitMte3); + + LocalTensor bmm2ResPreUb = inputBuff1.Get(); + bmm2ResPreUb = bmm2ResPreUb[pingpongFlag * INPUT1_BUFFER_OFFSET / sizeof(MM2_OUT_T)]; + WaitFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + + uint64_t accumGmOffset = + ((info.loop - 1) % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + inOutBaseOffset; + DataCopy(bmm2ResPreUb, mm2ResGm[accumGmOffset], vec2ComputeSize); + + SetFlag(SYNC_INPUT_BUF1_FLAG); + WaitFlag(SYNC_INPUT_BUF1_FLAG); + + uint32_t idx = info.loop % (constInfo.preLoadNum); + LocalTensor expUb = v0ValidSizeBuff.Get()[384]; // sumUb用临时内存 16 * 32B = 512B + Brcb(expUb, softmaxExpUb[idx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset], (dealRowCount + 7) / 8, + {1, 8}); + PipeBarrier(); + + RowMuls(bmm2ResPreUb, bmm2ResPreUb, expUb, dealRowCount, columnCount, actualColumnCount); + AscendC::PipeBarrier(); + Add(bmm2ResUb, bmm2ResUb, bmm2ResPreUb, vec2ComputeSize); + AscendC::PipeBarrier(); + + SetFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + pingpongFlag ^= 1; // pingpong 0 1 切换 + } + + // 最后一次输出计算结果,否则将中间结果暂存至workspace + if (info.isLastS2Loop) { + uint32_t idx = info.loop % (constInfo.preLoadNum); + LocalTensor tmpSumUb = v0ValidSizeBuff.Get()[384]; // sumUb用临时内存 16 * 32B = 512B + Brcb(tmpSumUb, softmaxSumUb[idx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset], + (dealRowCount + 7) / 8, {1, 8}); + PipeBarrier(); + RowDivs(bmm2ResUb, bmm2ResUb, tmpSumUb, dealRowCount, columnCount, actualColumnCount); + PipeBarrier(); + Bmm2ResCopyOut(info, bmm2ResUb, mStart, dealRowCount, columnCount, actualColumnCount); + } else if (!info.isFirstSInnerLoop) { + LocalTensor outUb = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + DataCopy(outUb, bmm2ResUb, dealRowCount * columnCount); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + uint64_t accumGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + inOutBaseOffset; + DataCopy(mm2ResGm[accumGmOffset], outUb, vec2ComputeSize); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + } + } +} + +template +__aicore__ inline void SMLAVectorBlock::RowDivs(LocalTensor dstUb, LocalTensor src0Ub, + LocalTensor src1Ub, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + // divs by row, 每行的元素除以相同的元素 + // dstUb[i, (j * 8) : (j * 8 + 7)] = src0Ub[i, (j * 8) : (j * 8 + 7)] / src1Ub[i, 0 : 7] + // src0Ub:[dealRowCount, columnCount], src1Ub:[dealRowCount, FP32_BLOCK_ELEMENT_NUM] dstUb:[dealRowCount, + // columnCount] + uint32_t dtypeMask = FP32_REPEAT_ELEMENT_NUM; + uint32_t dLoop = actualColumnCount / dtypeMask; + uint32_t dRemain = actualColumnCount % dtypeMask; + + BinaryRepeatParams repeatParamsDiv; + repeatParamsDiv.src0BlkStride = 1; + repeatParamsDiv.src1BlkStride = 0; + repeatParamsDiv.dstBlkStride = 1; + repeatParamsDiv.src0RepStride = columnCount / FP32_BLOCK_ELEMENT_NUM; + repeatParamsDiv.src1RepStride = 1; + repeatParamsDiv.dstRepStride = columnCount / FP32_BLOCK_ELEMENT_NUM; + uint32_t columnRepeatCount = dLoop; + if (columnRepeatCount <= dealRowCount) { + uint32_t offset = 0; + for (uint32_t i = 0; i < dLoop; i++) { + Div(dstUb[offset], src0Ub[offset], src1Ub, dtypeMask, dealRowCount, repeatParamsDiv); + offset += dtypeMask; + } + } else { + BinaryRepeatParams columnRepeatParams; + columnRepeatParams.src0BlkStride = 1; + columnRepeatParams.src1BlkStride = 0; + columnRepeatParams.dstBlkStride = 1; + columnRepeatParams.src0RepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + columnRepeatParams.src1RepStride = 0; + columnRepeatParams.dstRepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + uint32_t offset = 0; + for (uint32_t i = 0; i < dealRowCount; i++) { + Div(dstUb[offset], src0Ub[offset], src1Ub[i * FP32_BLOCK_ELEMENT_NUM], dtypeMask, columnRepeatCount, + columnRepeatParams); + offset += columnCount; + } + } + if (dRemain > 0) { + Div(dstUb[dLoop * dtypeMask], src0Ub[dLoop * dtypeMask], src1Ub, dRemain, dealRowCount, repeatParamsDiv); + } +} + +template +__aicore__ inline void SMLAVectorBlock::RowMuls(LocalTensor dstUb, LocalTensor src0Ub, + LocalTensor src1Ub, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + // muls by row, 每行的元素乘以相同的元素 + // dstUb[i, (j * 8) : (j * 8 + 7)] = src0Ub[i, (j * 8) : (j * 8 + 7)] * src1Ub[i, 0 : 7] + // src0Ub:[dealRowCount, columnCount] src1Ub:[dealRowCount, FP32_BLOCK_ELEMENT_NUM] dstUb:[dealRowCount, + // columnCount] + // dealRowCount is repeat times, must be less 256 + uint32_t repeatElementNum = FP32_REPEAT_ELEMENT_NUM; + uint32_t blockElementNum = FP32_BLOCK_ELEMENT_NUM; + + if constexpr (std::is_same::value) { + // 此限制由于每个repeat至多连续读取256B数据 + repeatElementNum = FP32_REPEAT_ELEMENT_NUM * 2; // 256/4 * 2=128 + blockElementNum = FP32_BLOCK_ELEMENT_NUM * 2; // 32/4 * 2 = 16 + } + + // 每次只能连续读取256B的数据进行计算,故每次只能处理256B/sizeof(dType)= + // 列方向分dLoop次,每次处理8列数据 + uint32_t dLoop = actualColumnCount / repeatElementNum; + uint32_t dRemain = actualColumnCount % repeatElementNum; + // REPEATE_STRIDE_UP_BOUND=256, 此限制由于src0RepStride数据类型为uint8之多256个datablock间距 + if (columnCount < REPEATE_STRIDE_UP_BOUND * blockElementNum) { + BinaryRepeatParams repeatParams; + repeatParams.src0BlkStride = 1; + repeatParams.src1BlkStride = 0; + repeatParams.dstBlkStride = 1; + repeatParams.src0RepStride = columnCount / blockElementNum; + repeatParams.src1RepStride = 1; + repeatParams.dstRepStride = columnCount / blockElementNum; + + // 如果以列为repeat所处理的次数小于行处理次数,则以列方式处理。反之则以行进行repeat处理 + if (dLoop <= dealRowCount) { + uint32_t offset = 0; + for (uint32_t i = 0; i < dLoop; i++) { + Mul(dstUb[offset], src0Ub[offset], src1Ub, repeatElementNum, dealRowCount, repeatParams); + offset += repeatElementNum; + } + } else { + BinaryRepeatParams columnRepeatParams; + columnRepeatParams.src0BlkStride = 1; + columnRepeatParams.src1BlkStride = 0; + columnRepeatParams.dstBlkStride = 1; + columnRepeatParams.src0RepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + columnRepeatParams.src1RepStride = 0; + columnRepeatParams.dstRepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + for (uint32_t i = 0; i < dealRowCount; i++) { + Mul(dstUb[i * columnCount], src0Ub[i * columnCount], src1Ub[i * blockElementNum], repeatElementNum, + dLoop, columnRepeatParams); + } + } + + // 最后一次完成[dealRowCount, dRemain] * [dealRowCount, blockElementNum] 只计算有效部分 + if (dRemain > 0) { + Mul(dstUb[dLoop * repeatElementNum], src0Ub[dLoop * repeatElementNum], src1Ub, dRemain, dealRowCount, + repeatParams); + } + } else { + BinaryRepeatParams repeatParams; + repeatParams.src0RepStride = 8; // 每个repeat为256B数据,正好8个datablock + repeatParams.src0BlkStride = 1; + repeatParams.src1RepStride = 0; + repeatParams.src1BlkStride = 0; + repeatParams.dstRepStride = 8; + repeatParams.dstBlkStride = 1; + // 每次计算一行,共计算dealRowCount行 + for (uint32_t i = 0; i < dealRowCount; i++) { + // 计算一行中的dLoop个repeat, 每个repeat计算256/block_size 个data_block + Mul(dstUb[i * columnCount], src0Ub[i * columnCount], src1Ub[i * blockElementNum], repeatElementNum, dLoop, + repeatParams); + // 计算一行中的尾块 + if (dRemain > 0) { + Mul(dstUb[i * columnCount + dLoop * repeatElementNum], + src0Ub[i * columnCount + dLoop * repeatElementNum], src1Ub[i * blockElementNum], dRemain, 1, + repeatParams); + } + } + } +} +} // namespace SMLAKernel +#endif // SPARSE_FLASH_MLA_CSA_BLOCK_VECTOR_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_kernel.h b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_kernel.h new file mode 100644 index 000000000000..e0c85cca7d40 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_csa_kernel.h @@ -0,0 +1,979 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_csa_kernel.h + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_CSA_KERNEL_H +#define SPARSE_FLASH_MLA_CSA_KERNEL_H + +#include "kernel_operator.h" +#include "kernel_operator_list_tensor_intf.h" +#include "kernel_tiling/kernel_tiling.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" +#include "sparse_flash_mla_common_arch22.h" +#include "sparse_flash_mla_csa_block_cube.h" +#include "sparse_flash_mla_csa_block_vector.h" +#include "sparse_flash_mla_arch22_metadata.h" + +namespace SMLAKernel { +using namespace matmul; +using namespace optiling; +using AscendC::CrossCoreSetFlag; +using AscendC::CrossCoreWaitFlag; + +// 由于S2循环前,RunInfo还没有赋值,使用Bngs1Param临时存放B、N、S1轴相关的信息;同时减少重复计算 +struct TempLoopInfo { + uint32_t bn2IdxInCurCore = 0; + uint32_t bIdx = 0U; + uint32_t n2Idx = 0U; + uint64_t s2BasicSizeTail = 0U; // S2方向循环的尾基本块大小 + uint32_t s2LoopTimes = 0U; // S2方向循环的总次数,无论TND还是BXXD都是等于实际次数,不用减1 + + int32_t actS1Size = 0; // TND场景下当前Batch循环处理的S1轴的大小 + int32_t actOriS2Size = 0; + int32_t actCmpS2Size = 0; + + bool curActSeqLenIsZero = false; + + uint32_t tndCoreStartKVSplitPos = 0; + bool tndIsS2SplitCore = false; + uint32_t gS1Idx = 0U; + uint32_t s1StartIdx = 0; + uint32_t s1EndIdx = 0; + uint64_t mBasicSizeTail = 0U; // gS1方向循环的尾基本块大小 + uint32_t cmpLoopTimes = 0; + uint32_t oriLoopTimes = 0; + uint32_t v0OriSize = 0; + uint32_t v0CmpSize = 0; + + // sparsemode = 4 + int32_t oriMaskRight = 0; + int32_t oriMaskLeft = 0; + + // sparsemode = 3 + int32_t cmpMaskRight = 0; + + uint64_t actualSeqQPrefixSum = 0; + uint64_t actualSeqKVPrefixSum = 0; + uint64_t actualSeqCmpKVPrefixSum = 0; +}; + +template +class SparseFlashMlaCsa { +public: + // 中间计算数据类型为float,高精度模式 + using T = float; + using Q_T = typename SMLAT::queryType; + using KV_T = typename SMLAT::kvType; + using OUT_T = typename SMLAT::outputType; + using SINKS_T = float; + using UPDATE_T = T; + using MM1_OUT_T = T; + using MM2_OUT_T = T; + + __aicore__ inline SparseFlashMlaCsa(){}; + __aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, + __gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensKV, __gm__ uint8_t *cuSeqlensCmpKV, + __gm__ uint8_t *seqUsedQ, __gm__ uint8_t *seqUsedKV, __gm__ uint8_t *seqUsedCmpKV, + __gm__ uint8_t *cmpResidualKV, __gm__ uint8_t *oriTopkLength, + __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata, + __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, __gm__ uint8_t *workspace, + const SparseFlashMlaTilingData *__restrict tiling, __gm__ uint8_t *gmTiling, + TPipe *tPipe); + + __aicore__ inline void Process(); + +private: + static constexpr bool PAGE_ATTENTION = SMLAT::pageAttention; + static constexpr bool FLASH_DECODE = SMLAT::flashDecode; + static constexpr SMLA_LAYOUT LAYOUT_T = SMLAT::layout; + static constexpr SMLA_LAYOUT KV_LAYOUT_T = SMLAT::kvLayout; + static constexpr bool HEAD_RATIO_ONE = SMLAT::headRatioOne; + + static constexpr uint32_t PRELOAD_NUM = 2; + static constexpr uint32_t N_BUFFER_M_BASIC_SIZE = 256; + static constexpr uint32_t SMLA_PRELOAD_TASK_CACHE_SIZE = 3; + static constexpr uint32_t MERGE_CACHE_GM_BUF_NUM = 3; + + static constexpr uint32_t SYNC_V0_C1_FLAG = 6; + static constexpr uint32_t SYNC_C1_V1_FLAG = 7; + static constexpr uint32_t SYNC_V1_C2_FLAG = 8; + static constexpr uint32_t SYNC_C2_V2_FLAG = 9; + + static constexpr uint64_t headDim = 512ULL; + + static constexpr uint32_t dbWorkspaceRatio = PRELOAD_NUM; + + const SparseFlashMlaTilingData *__restrict tilingData = nullptr; + + TPipe *pipe = nullptr; + GlobalTensor metadataGm; + uint64_t mSizeVStart = 0ULL; + uint64_t topKBaseOffset = 0ULL; + uint64_t tensorACoreOffset = 0ULL; + uint64_t tensorBCoreOffset = 0ULL; + uint64_t tensorCmpBCoreOffset = 0ULL; + + uint32_t tmpBlockIdx = 0U; + uint32_t aiCoreIdx = 0U; + + ConstInfo constInfo{}; + TempLoopInfo tempLoopInfo{}; + + SMLACubeBlock cubeBlock; + SMLAVectorBlock vectorBlock; + + GlobalTensor queryGm; + GlobalTensor oriKvGm; + GlobalTensor cmpKvGm; + GlobalTensor sinksGm; + + GlobalTensor attentionOutGm; + GlobalTensor softmaxLseGm; + + GlobalTensor oriBlockTableGm; + GlobalTensor cmpBlockTableGm; + GlobalTensor topKGm; + + GlobalTensor actualSeqLengthsQGm; + GlobalTensor actualSeqLengthsKVGm; + GlobalTensor actualSeqLengthsCmpKVGm; + GlobalTensor cmpResidualKVGm; + + // workspace + GlobalTensor mm1ResGm; + GlobalTensor vec1ResGm; + GlobalTensor mm2ResGm; + GlobalTensor kvMergeGm_; + + GlobalTensor vec2ResGm; + + GlobalTensor accumOutGm; + + // ================================Init functions================================== + __aicore__ inline void InitTilingData(); + __aicore__ inline void InitCalcParamsEach(); + __aicore__ inline void InitBuffers(); + __aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKV); + __aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKV, + __gm__ uint8_t *actualSeqLengthsCmpKV); + __aicore__ inline void InitOutputSingleCore(); + // ================================Process functions================================ + __aicore__ inline void ProcessBalance(); + __aicore__ inline void PreloadPipeline(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, uint64_t s2LoopIdx, + RunInfo extraInfo[SMLA_PRELOAD_TASK_CACHE_SIZE]); + // ================================Offset Calc===================================== + __aicore__ inline void GetSparseActualSeqLen(); + __aicore__ inline int32_t CountValidCmpSparseLen(int32_t maxLen); + __aicore__ inline void UpdateInnerLoopCond(); + __aicore__ inline void CalcParams(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, uint32_t s2LoopIdx, + RunInfo &info); + __aicore__ inline int32_t GetActualSeqLenQ(uint32_t bIdx); + __aicore__ inline int32_t GetActualSeqLenKV(uint32_t bIdx); + __aicore__ inline int32_t GetActualSeqLenCmpKV(uint32_t bIdx, int32_t actualOriS2Size); + __aicore__ inline int32_t GetCmpMaskS2Size(uint32_t bIdx, int32_t actualOriS2Size, int32_t actualCmpS2Size); + __aicore__ inline void GetBN2Idx(uint32_t bN2Idx, uint32_t &bIdx, uint32_t &n2Idx); + // ================================Mm1============================================== + __aicore__ inline void ComputeMm1(const RunInfo &info); + // ================================Mm2============================================== + __aicore__ inline void ComputeMm2(const RunInfo &info); + __aicore__ inline void InitAllZeroOutput(uint32_t bIdx, uint32_t s1Idx, uint32_t n2Idx); +}; + +template +__aicore__ inline void SparseFlashMlaCsa::InitTilingData() +{ + // singleCoreParams + // singleCoreTensorSize + constInfo.mmResUbSize = tilingData->baseParams.mmResUbSize; + constInfo.bmm2ResUbSize = tilingData->baseParams.bmm2ResUbSize; + + // baseParams + constInfo.batchSize = tilingData->baseParams.batchSize; + constInfo.gSize = tilingData->baseParams.nNumOfQInOneGroup; + constInfo.kvHeadNum = (tilingData->baseParams.kvHeadNum == 0) ? 1 : tilingData->baseParams.kvHeadNum; + constInfo.qHeadNum = constInfo.gSize * constInfo.kvHeadNum; + constInfo.kvSeqSize = tilingData->baseParams.kvSeqSize; + constInfo.qSeqSize = tilingData->baseParams.qSeqSize; + constInfo.oriMaxBlockNumPerBatch = tilingData->baseParams.oriMaxBlockNumPerBatch; + constInfo.cmpMaxBlockNumPerBatch = tilingData->cmpParams.cmpMaxBlockNumPerBatch; + constInfo.kvCacheBlockSize = tilingData->baseParams.paBlockSize; + constInfo.paOriBlockSize = tilingData->baseParams.oriBlockSize; + constInfo.paCmpBlockSize = tilingData->baseParams.cmpBlockSize; + constInfo.outputLayout = static_cast(tilingData->baseParams.outputLayout); + constInfo.headDim = headDim; + constInfo.oriMaskMode = tilingData->baseParams.oriMaskMode; + constInfo.oriKvStride0 = tilingData->baseParams.oriKvStride0; + constInfo.oriWinLeft = tilingData->baseParams.oriWinLeft; + constInfo.oriWinRight = tilingData->baseParams.oriWinRight; + + constInfo.actualLenDimsQ = tilingData->baseParams.actualLenDimsQ; + constInfo.actualLenDimsKV = tilingData->baseParams.actualLenDimsKV; + constInfo.actualLenDimsCmpKV = tilingData->baseParams.actualLenDimsCmpKV; + constInfo.cmpResidualKVSize = tilingData->baseParams.cmpResidualKVSize; + constInfo.returnSoftmaxLse = tilingData->baseParams.returnSoftmaxLse; + // innerSplitParams + constInfo.mBaseSize = constInfo.gSize; + constInfo.s2BaseSize = tilingData->baseParams.s2BaseSize; + + constInfo.preLoadNum = PRELOAD_NUM; + constInfo.nBufferMBaseSize = N_BUFFER_M_BASIC_SIZE; + constInfo.syncV0C1 = SYNC_V0_C1_FLAG; + constInfo.syncC1V1 = SYNC_C1_V1_FLAG; + constInfo.syncV1C2 = SYNC_V1_C2_FLAG; + constInfo.syncC2V2 = SYNC_C2_V2_FLAG; + + // cmp + constInfo.cmpRatio = tilingData->cmpParams.cmpRatio; + constInfo.sparseBlockCount = tilingData->cmpParams.sparseBlockCount; + constInfo.sparseBlockSize = 1; // sparseBlockSize 固定为1 + constInfo.cmpMaskMode = tilingData->cmpParams.cmpMaskMode; + constInfo.cmpKvStride0 = tilingData->cmpParams.cmpKvStride0; + constInfo.cmpSeqSize = tilingData->cmpParams.cmpKvSeqSize; +} + +template +__aicore__ inline void SparseFlashMlaCsa::InitBuffers() +{ + if ASCEND_IS_AIV { + vectorBlock.InitBuffers(pipe); + } else { + cubeBlock.InitBuffers(pipe); + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, + __gm__ uint8_t *actualSeqLengthsKV) +{ + if (constInfo.actualLenDimsKV != 0) { + actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsKV, constInfo.actualLenDimsKV); + } + if (constInfo.actualLenDimsQ != 0) { + actualSeqLengthsQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ, constInfo.actualLenDimsQ); + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, + __gm__ uint8_t *actualSeqLengthsKV, + __gm__ uint8_t *actualSeqLengthsCmpKV) +{ + if (constInfo.actualLenDimsKV != 0) { + actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsKV, constInfo.actualLenDimsKV); + } + if (constInfo.actualLenDimsCmpKV != 0) { + actualSeqLengthsCmpKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsCmpKV, constInfo.actualLenDimsCmpKV); + } + if (constInfo.actualLenDimsQ != 0) { + actualSeqLengthsQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ, constInfo.actualLenDimsQ); + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::InitAllZeroOutput(uint32_t bIdx, uint32_t s1Idx, uint32_t n2Idx) +{ + if (constInfo.outputLayout == SMLA_LAYOUT::TND) { + if (tempLoopInfo.actS1Size == 0) { + return; + } + uint32_t tBase = actualSeqLengthsQGm.GetValue(bIdx); + uint32_t s1Count = tempLoopInfo.actS1Size; + + uint64_t attenOutOffset = (tBase + s1Idx) * constInfo.kvHeadNum * constInfo.gSize * constInfo.headDim + + n2Idx * constInfo.gSize * constInfo.headDim; + uint64_t lseOffset = (tBase + s1Idx) * constInfo.gSize + // T轴、s1轴偏移 + n2Idx * constInfo.qSeqSize * constInfo.gSize; // N2轴偏移 + matmul::InitOutput(attentionOutGm[attenOutOffset], constInfo.gSize * constInfo.headDim, 0); + if (constInfo.returnSoftmaxLse) { + matmul::InitOutput(softmaxLseGm[lseOffset], constInfo.gSize, 0); + } + } else if (constInfo.outputLayout == SMLA_LAYOUT::BSND) { + uint64_t attenOutOffset = + bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize * constInfo.headDim + + s1Idx * constInfo.kvHeadNum * constInfo.gSize * constInfo.headDim + + n2Idx * constInfo.gSize * constInfo.headDim; + uint64_t lseOffset = bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize + // B轴偏移 + n2Idx * constInfo.qSeqSize * constInfo.gSize + // N2轴偏移 + s1Idx * constInfo.gSize; // S1轴偏移 + matmul::InitOutput(attentionOutGm[attenOutOffset], constInfo.gSize * constInfo.headDim, 0); + if (constInfo.returnSoftmaxLse) { + matmul::InitOutput(softmaxLseGm[lseOffset], constInfo.gSize, 0); + } + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::InitOutputSingleCore() +{ + uint32_t coreNum = GetBlockNum(); + if (coreNum != 0) { + uint64_t totalOutputSize = constInfo.batchSize * constInfo.qHeadNum * constInfo.qSeqSize * constInfo.headDim; + uint64_t singleCoreSize = (totalOutputSize + (2 * coreNum) - 1) / (2 * coreNum); // 2 means c:v = 1:2 + uint64_t tailSize = totalOutputSize - tmpBlockIdx * singleCoreSize; + uint64_t singleInitOutputSize = tailSize < singleCoreSize ? tailSize : singleCoreSize; + if (singleInitOutputSize > 0) { + matmul::InitOutput(attentionOutGm[tmpBlockIdx * singleCoreSize], singleInitOutputSize, 0); + } + SyncAll(); + } +} + +template +__aicore__ inline int32_t SparseFlashMlaCsa::GetActualSeqLenQ(uint32_t bIdx) +{ + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + int32_t actualSeqQPrefixSum = actualSeqLengthsQGm.GetValue(bIdx); + int32_t actualSeqQNextSum = actualSeqLengthsQGm.GetValue(bIdx + 1); + tempLoopInfo.actualSeqQPrefixSum = static_cast(actualSeqQPrefixSum); + return actualSeqQNextSum - actualSeqQPrefixSum; + } else { + tempLoopInfo.actualSeqQPrefixSum = static_cast(bIdx * constInfo.qSeqSize); + if (constInfo.actualLenDimsQ == 0) { + return static_cast(constInfo.qSeqSize); + } else { + return actualSeqLengthsQGm.GetValue(bIdx); + } + } +} + +template +__aicore__ inline int32_t SparseFlashMlaCsa::GetActualSeqLenKV(uint32_t bIdx) +{ + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + tempLoopInfo.actualSeqKVPrefixSum = static_cast(bIdx * constInfo.kvSeqSize); + if (constInfo.actualLenDimsKV == 0) { + return static_cast(constInfo.kvSeqSize); + } + return actualSeqLengthsKVGm.GetValue(bIdx); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + tempLoopInfo.actualSeqKVPrefixSum = static_cast(bIdx * constInfo.kvSeqSize); + if (constInfo.actualLenDimsKV != 0) { + return actualSeqLengthsKVGm.GetValue(bIdx); + } + return static_cast(constInfo.kvSeqSize); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + int32_t actualSeqKVPrefixSum = actualSeqLengthsKVGm.GetValue(bIdx); + int32_t actualSeqKVNextSum = actualSeqLengthsKVGm.GetValue(bIdx + 1); + tempLoopInfo.actualSeqKVPrefixSum = actualSeqKVPrefixSum; + return actualSeqKVNextSum - actualSeqKVPrefixSum; + } +} + +template +__aicore__ inline int32_t SparseFlashMlaCsa::GetActualSeqLenCmpKV(uint32_t bIdx, int32_t actualOriS2Size) +{ + (void)actualOriS2Size; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + int32_t actualSeqCmpKVPrefixSum = actualSeqLengthsCmpKVGm.GetValue(bIdx); + int32_t actualSeqCmpKVNextSum = actualSeqLengthsCmpKVGm.GetValue(bIdx + 1); + tempLoopInfo.actualSeqCmpKVPrefixSum = actualSeqCmpKVPrefixSum; + return actualSeqCmpKVNextSum - actualSeqCmpKVPrefixSum; + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + tempLoopInfo.actualSeqCmpKVPrefixSum = static_cast(bIdx * constInfo.cmpSeqSize); + if (constInfo.actualLenDimsCmpKV != 0) { + return actualSeqLengthsCmpKVGm.GetValue(bIdx); + } + return static_cast(constInfo.cmpSeqSize); + } else { + tempLoopInfo.actualSeqCmpKVPrefixSum = static_cast(bIdx * constInfo.cmpSeqSize); + if (constInfo.actualLenDimsCmpKV != 0) { + return actualSeqLengthsCmpKVGm.GetValue(bIdx); + } + return (constInfo.cmpSeqSize != 0) ? static_cast(constInfo.cmpSeqSize) : + actualOriS2Size / static_cast(constInfo.cmpRatio); + } +} + +template +__aicore__ inline int32_t SparseFlashMlaCsa::GetCmpMaskS2Size(uint32_t bIdx, int32_t actualOriS2Size, + int32_t actualCmpS2Size) +{ + (void)actualOriS2Size; + int32_t residual = 0; + if (constInfo.cmpResidualKVSize != 0) { + residual = cmpResidualKVGm.GetValue(bIdx); + } + return actualCmpS2Size * static_cast(constInfo.cmpRatio) + residual; +} + +template +__aicore__ inline void SparseFlashMlaCsa::GetSparseActualSeqLen() +{ + // 行无效通过ori部分判断, ori部分如果有行无效那么ori和cmp都有 + if (static_cast(tempLoopInfo.s1EndIdx) < -(tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size)) { + tempLoopInfo.actOriS2Size = 0; + tempLoopInfo.actCmpS2Size = 0; + return; + } + + // 对于cmp部分还有top k, tempLoopInfo.actS2Size只针对cmp。 + // 因果可见长度是 q_idx+1,连续 indices(0..q_idx,-1)时有效条数等于该长度; + // 出现 sparse gap(如 0..70,88,89,90,-1)时有效条数更少,gather 与计算都按真实有效条数走。 + int32_t thresHold = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio; + int32_t bound = Min(tempLoopInfo.actCmpS2Size, + Min(constInfo.sparseBlockCount * constInfo.sparseBlockSize, Max(thresHold, 0))); + tempLoopInfo.actCmpS2Size = Min(bound, CountValidCmpSparseLen(bound)); +} + +template +__aicore__ inline int32_t SparseFlashMlaCsa::CountValidCmpSparseLen(int32_t maxLen) +{ + if (maxLen <= 0) { + return 0; + } + int64_t base; + if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND) { + base = (static_cast(tempLoopInfo.bIdx) * constInfo.qSeqSize + tempLoopInfo.s1StartIdx) * + constInfo.kvHeadNum * constInfo.sparseBlockCount + + static_cast(tempLoopInfo.n2Idx) * constInfo.sparseBlockCount; + } else { + base = (static_cast(tempLoopInfo.actualSeqQPrefixSum) + tempLoopInfo.s1StartIdx) * + constInfo.kvHeadNum * constInfo.sparseBlockCount + + static_cast(tempLoopInfo.n2Idx) * constInfo.sparseBlockCount; + } + int32_t lo = 0; + int32_t hi = maxLen; + while (lo < hi) { + int32_t mid = (lo + hi) >> 1; + if (topKGm.GetValue(base + mid) < 0) { + hi = mid; + } else { + lo = mid + 1; + } + } + return lo; +} + +template +__aicore__ inline void SparseFlashMlaCsa::UpdateInnerLoopCond() +{ + if ((tempLoopInfo.actCmpS2Size == 0 && tempLoopInfo.actOriS2Size == 0) || (tempLoopInfo.actS1Size == 0)) { + tempLoopInfo.curActSeqLenIsZero = true; + return; + } + tempLoopInfo.curActSeqLenIsZero = false; + tempLoopInfo.mBasicSizeTail = (tempLoopInfo.actS1Size * constInfo.gSize) % constInfo.mBaseSize; + tempLoopInfo.mBasicSizeTail = + (tempLoopInfo.mBasicSizeTail == 0) ? constInfo.mBaseSize : tempLoopInfo.mBasicSizeTail; +} + +template +__aicore__ inline void SparseFlashMlaCsa::Init( + __gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensKV, __gm__ uint8_t *cuSeqlensCmpKV, __gm__ uint8_t *seqUsedQ, + __gm__ uint8_t *seqUsedKV, __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata, + __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, __gm__ uint8_t *workspace, + const SparseFlashMlaTilingData *__restrict tiling, __gm__ uint8_t *gmTiling, TPipe *tPipe) +{ + if ASCEND_IS_AIV { + tmpBlockIdx = GetBlockIdx(); // vec:0-47 + aiCoreIdx = tmpBlockIdx / 2; + } else { + tmpBlockIdx = GetBlockIdx(); // cube:0-23 + aiCoreIdx = tmpBlockIdx; + } + + // init tiling data + tilingData = tiling; + InitTilingData(); + if (KV_LAYOUT_T == SMLA_LAYOUT::TND && LAYOUT_T == SMLA_LAYOUT::TND) { + InitActualSeqLen(cuSeqlensQ, cuSeqlensKV, cuSeqlensCmpKV); + } else if (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + InitActualSeqLen(seqUsedQ, cuSeqlensKV, cuSeqlensCmpKV); + } else if ((KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND || KV_LAYOUT_T == SMLA_LAYOUT::BSND) && + LAYOUT_T == SMLA_LAYOUT::TND) { + InitActualSeqLen(cuSeqlensQ, seqUsedKV, seqUsedCmpKV); + } else if ((KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND || KV_LAYOUT_T == SMLA_LAYOUT::BSND)) { + InitActualSeqLen(seqUsedQ, seqUsedKV, seqUsedCmpKV); + } + + metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata); + InitCalcParamsEach(); + + pipe = tPipe; + // init global buffer + queryGm.SetGlobalBuffer((__gm__ Q_T *)query); + oriKvGm.SetGlobalBuffer((__gm__ KV_T *)oriKV); + cmpKvGm.SetGlobalBuffer((__gm__ KV_T *)cmpKV); + if (constInfo.cmpResidualKVSize != 0) { + cmpResidualKVGm.SetGlobalBuffer((__gm__ int32_t *)cmpResidualKV, constInfo.cmpResidualKVSize); + } + + if (sinks != nullptr) { + sinksGm.SetGlobalBuffer((__gm__ SINKS_T *)sinks); + } + + attentionOutGm.SetGlobalBuffer((__gm__ OUT_T *)attentionOut); + softmaxLseGm.SetGlobalBuffer((__gm__ T *)softmaxLse); + + if ASCEND_IS_AIV { + if (LAYOUT_T != SMLA_LAYOUT::TND) { + if (constInfo.needInit) { + InitOutputSingleCore(); + } + } + } + + if constexpr (PAGE_ATTENTION) { + oriBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)oriBlockTable); + cmpBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)cmpBlockTable); + } + topKGm.SetGlobalBuffer((__gm__ int32_t *)cmpSparseIndices); + + // workspace 内存排布 + // |Q--|mm1ResGm|vec1ResGm|mm2ResGm|vec2ResGm + // |Core0_Q1-Core0_Q2-Core1_Q1-Core1_Q2....Core32_Q1-Core32_Q2|Core0_mmRes + uint64_t offset = 0; + mm1ResGm.SetGlobalBuffer( + (__gm__ MM1_OUT_T *)(workspace + offset + + aiCoreIdx * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(MM1_OUT_T))); + offset += GetBlockNum() * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(MM1_OUT_T); + + vec1ResGm.SetGlobalBuffer( + (__gm__ Q_T *)(workspace + offset + aiCoreIdx * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(KV_T))); + offset += GetBlockNum() * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(KV_T); + + mm2ResGm.SetGlobalBuffer( + (__gm__ MM2_OUT_T *)(workspace + offset + + aiCoreIdx * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(MM2_OUT_T))); + offset += GetBlockNum() * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(MM2_OUT_T); + + vec2ResGm.SetGlobalBuffer( + (__gm__ T *)(workspace + offset + aiCoreIdx * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(T))); + offset += GetBlockNum() * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(T); + + kvMergeGm_.SetGlobalBuffer( + (__gm__ KV_T *)(workspace + offset + aiCoreIdx * 512 * 512 * MERGE_CACHE_GM_BUF_NUM * sizeof(KV_T))); + offset += GetBlockNum() * 512 * 512 * 4 * sizeof(KV_T); + + if ASCEND_IS_AIV { + vectorBlock.InitParams(constInfo, tilingData); + vectorBlock.InitVec0GlobalTensor(kvMergeGm_, oriKvGm, cmpKvGm, oriBlockTableGm, cmpBlockTableGm); + vectorBlock.InitVec1GlobalTensor(mm1ResGm, vec1ResGm, actualSeqLengthsQGm, actualSeqLengthsKVGm, topKGm, + sinksGm, softmaxLseGm); + vectorBlock.InitVec2GlobalTensor(accumOutGm, vec2ResGm, mm2ResGm, attentionOutGm); + } + + if ASCEND_IS_AIC { + cubeBlock.InitParams(constInfo); + cubeBlock.InitMm1GlobalTensor(queryGm, oriKvGm, cmpKvGm, mm1ResGm); + cubeBlock.InitMm2GlobalTensor(vec1ResGm, mm2ResGm, attentionOutGm); + cubeBlock.InitPageAttentionInfo(oriKvGm, kvMergeGm_, oriBlockTableGm, cmpBlockTableGm); + } + // 要在InitParams之后执行 + if (pipe != nullptr) { + InitBuffers(); + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::InitCalcParamsEach() +{ + if (aiCoreIdx != 0) { + constInfo.bN2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_BN2_START_INDEX, false)); + constInfo.gS1Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_M_START_INDEX, false)); + constInfo.s2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_S2_START_INDEX, false)); + } + constInfo.bN2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_BN2_END_INDEX, false)); + constInfo.gS1End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_M_END_INDEX, false)); + constInfo.s2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_S2_END_INDEX, false)); +} + +template +__aicore__ inline void SparseFlashMlaCsa::CalcParams(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, + uint32_t s2LoopIdx, RunInfo &info) +{ + info.isValid = s2LoopIdx < tempLoopInfo.s2LoopTimes; + info.loop = loop; + info.cmpLoop = cmpLoop; + info.bIdx = tempLoopInfo.bIdx; + info.n2IdxReal = tempLoopInfo.n2Idx; + + info.gS1Idx = tempLoopInfo.gS1Idx; + info.s1Idx = tempLoopInfo.gS1Idx / constInfo.gSize; + info.s2Idx = s2LoopIdx; + info.curSInnerLoopTimes = tempLoopInfo.s2LoopTimes; + info.tndIsS2SplitCore = tempLoopInfo.tndIsS2SplitCore; + info.tndCoreStartKVSplitPos = tempLoopInfo.tndCoreStartKVSplitPos; + info.isBmm2Output = false; + info.actS1Size = tempLoopInfo.actS1Size; + + // M方向的尾块 + info.actMBaseSize = tempLoopInfo.mBasicSizeTail; + + if ASCEND_IS_AIV { + info.mSize = info.actMBaseSize; + info.mSizeV = (info.mSize <= 16) ? info.mSize : ((CeilDiv(info.mSize, 16) + 1) / 2 * 16); + info.mSizeVStart = 0; + if (tmpBlockIdx % 2 == 1) { + info.mSizeVStart = info.mSizeV; + info.mSizeV = info.mSize - info.mSizeV; + } + } + + info.isFirstSInnerLoop = s2LoopIdx == s2Start; + if (info.isFirstSInnerLoop) { + tempLoopInfo.bn2IdxInCurCore++; + } + info.isLastS2Loop = (s2LoopIdx == (tempLoopInfo.s2LoopTimes - 1)); + info.bn2IdxInCurCore = tempLoopInfo.bn2IdxInCurCore - 1; + + uint64_t tndBIdxOffsetForQ = tempLoopInfo.actualSeqQPrefixSum * constInfo.qHeadNum * constInfo.headDim; + uint64_t tndBIdxOffsetForKV = tempLoopInfo.actualSeqKVPrefixSum * constInfo.kvHeadNum * constInfo.headDim; + uint64_t tndBIdxOffsetForCmpKV = tempLoopInfo.actualSeqCmpKVPrefixSum * constInfo.kvHeadNum * constInfo.headDim; + + if (info.isFirstSInnerLoop) { + uint64_t s1HeadOffset = (info.gS1Idx / constInfo.gSize) * constInfo.qHeadNum; + uint64_t qHeadOffset = info.n2Idx * constInfo.gSize + info.gS1Idx % constInfo.gSize; + tensorACoreOffset = tndBIdxOffsetForQ + (s1HeadOffset + qHeadOffset) * constInfo.headDim; + tensorBCoreOffset = tndBIdxOffsetForKV + info.n2Idx * constInfo.headDim; // 当前为PA场景,该变量失效 + tensorCmpBCoreOffset = tndBIdxOffsetForCmpKV + info.n2Idx * constInfo.headDim; + if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND) { // B,S1,N2 K + topKBaseOffset = (info.bIdx * constInfo.qSeqSize + tempLoopInfo.s1StartIdx) * constInfo.kvHeadNum * + constInfo.sparseBlockCount + + info.n2Idx * constInfo.sparseBlockCount; + } else if (LAYOUT_T == SMLA_LAYOUT::TND) { // T N2 K + topKBaseOffset = (tempLoopInfo.actualSeqQPrefixSum + tempLoopInfo.s1StartIdx) * constInfo.kvHeadNum * + constInfo.sparseBlockCount + + info.n2Idx * constInfo.sparseBlockCount; + } + } + info.tensorAOffset = tensorACoreOffset; + info.tensorBOffset = tensorBCoreOffset; + info.tensorCmpBOffset = tensorCmpBCoreOffset; + info.attenOutOffset = tensorACoreOffset; + info.topKBaseOffset = topKBaseOffset; + + if (s2LoopIdx < tempLoopInfo.oriLoopTimes) { + // S2首次循环只能在ori_kv + info.isOriOnly = true; + info.relativeS2Idx = 0; + uint64_t s2Offset = info.s2Idx * constInfo.s2BaseSize; + if (s2LoopIdx + 1 == tempLoopInfo.oriLoopTimes) { + info.actualSingleProcessSInnerSize = (tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1) - s2Offset; + } else { + info.actualSingleProcessSInnerSize = constInfo.s2BaseSize; + } + info.s2StartPoint = tempLoopInfo.oriMaskLeft; + info.cmpS2IdLimit = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio; + } else { + info.isOriOnly = false; + info.relativeS2Idx = info.s2Idx - tempLoopInfo.oriLoopTimes; + uint64_t s2Offset = (info.s2Idx - tempLoopInfo.oriLoopTimes) * constInfo.s2BaseSize; + if (s2LoopIdx + 1 == tempLoopInfo.s2LoopTimes) { + info.actualSingleProcessSInnerSize = tempLoopInfo.actCmpS2Size - s2Offset; + } else { + info.actualSingleProcessSInnerSize = constInfo.s2BaseSize; + } + info.s2StartPoint = 0; + info.cmpS2IdLimit = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio; + if constexpr (HEAD_RATIO_ONE) { + info.v0S2Start = static_cast(s2Offset); + info.v0S2DealSize = static_cast(info.actualSingleProcessSInnerSize); + } else { + info.v0S2Start = 0; + if (s2LoopIdx + 1 == tempLoopInfo.s2LoopTimes && s2LoopIdx == 2) { + info.v0S2Start = 512; + } + info.v0S2DealSize = 512; + } + } + + info.actualSingleProcessSInnerSizeAlign = + SMLAAlign(info.actualSingleProcessSInnerSize, SMLAVectorBlock::BYTE_BLOCK); + if (info.isOriOnly) { + info.v0S2Start = 0; + info.v0S2DealSize = 0; + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::ComputeMm1(const RunInfo &info) +{ + uint32_t nBufferLoopTimes = CeilDiv(info.actMBaseSize, constInfo.nBufferMBaseSize); + uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize; + for (uint32_t i = 0; i < nBufferLoopTimes; i++) { + MSplitInfo mSplitInfo; + mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize; + mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail; + cubeBlock.ComputeMm1(info, mSplitInfo); + if constexpr (HEAD_RATIO_ONE) { + event_t eventIdFixWait = static_cast(GetTPipePtr()->FetchEventID(HardEvent::FIX_M)); + SetFlag(eventIdFixWait); + WaitFlag(eventIdFixWait); + } + CrossCoreSetFlag(constInfo.syncC1V1); + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::ComputeMm2(const RunInfo &info) +{ + uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize; + uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize; + for (uint32_t i = 0; i < nBufferLoopTimes; i++) { + MSplitInfo mSplitInfo; + mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize; + mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail; + if constexpr (HEAD_RATIO_ONE) { + CrossCoreWaitFlag(constInfo.syncV1C2); + } else { + CrossCoreWaitFlag(constInfo.syncV1C2); + } + cubeBlock.ComputeMm2(info, mSplitInfo); + if constexpr (HEAD_RATIO_ONE) { + event_t eventIdFixWait = static_cast(GetTPipePtr()->FetchEventID(HardEvent::FIX_M)); + SetFlag(eventIdFixWait); + WaitFlag(eventIdFixWait); + } + CrossCoreSetFlag(constInfo.syncC2V2); + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::Process() +{ + uint32_t hasLoad = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_CORE_ENABLE_INDEX, false)); + if (hasLoad == 0) { + return; + } + if ASCEND_IS_AIV { + vectorBlock.AllocEventID(); + vectorBlock.InitSoftmaxDefaultBuffer(); + } else { + cubeBlock.AllocEventID(); + } + ProcessBalance(); + if ASCEND_IS_AIV { + vectorBlock.FreeEventID(); + } else { + cubeBlock.FreeEventID(); + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::GetBN2Idx(uint32_t bN2Idx, uint32_t &bIdx, uint32_t &n2Idx) +{ + bIdx = bN2Idx / constInfo.kvHeadNum; + n2Idx = bN2Idx % constInfo.kvHeadNum; +} + +template +__aicore__ inline void SparseFlashMlaCsa::ProcessBalance() +{ + RunInfo extraInfo[SMLA_PRELOAD_TASK_CACHE_SIZE]; + uint32_t gloop = 0; + uint32_t cmpLoop = 0; + uint32_t gS1LoopEnd = 0; + bool globalLoopStart = true; + + if ASCEND_IS_AIC { + CrossCoreSetFlag(3); + CrossCoreSetFlag(3); + CrossCoreSetFlag(3); + CrossCoreSetFlag(3); + } + + // 适配左闭右开 + if (constInfo.bN2Start == constInfo.bN2End) { + if (constInfo.gS1Start != constInfo.gS1End || constInfo.s2Start != constInfo.s2End) { + constInfo.bN2End += 1; + } + } else if ((constInfo.gS1End != 0) || (constInfo.s2End != 0)) { + constInfo.bN2End += 1; + } + + for (uint32_t bN2LoopIdx = constInfo.bN2Start; bN2LoopIdx < constInfo.bN2End; bN2LoopIdx++) { + GetBN2Idx(bN2LoopIdx, tempLoopInfo.bIdx, tempLoopInfo.n2Idx); + tempLoopInfo.actS1Size = GetActualSeqLenQ(tempLoopInfo.bIdx); // 获取actualSeqLength + bool isS1ZeroAndLastBatch = (tempLoopInfo.actS1Size == 0) && ((constInfo.outputLayout == SMLA_LAYOUT::BSND) || + (bN2LoopIdx + 1 == constInfo.bN2End)); + uint32_t gS1SplitNum = CeilDiv(tempLoopInfo.actS1Size * constInfo.gSize, constInfo.mBaseSize); + + // 当处于最后一个BN2时, 且gS1End为0时, 说明当前BN2里的所有数据都在当前核处理 + gS1LoopEnd = (bN2LoopIdx + 1 == constInfo.bN2End && constInfo.gS1End != 0) ? constInfo.gS1End : gS1SplitNum; + // 当处于最后一个BN2且当前S1为0时,需要进入循环计算preload导致的未完成的部分 + gS1LoopEnd = isS1ZeroAndLastBatch ? gS1LoopEnd + 1 : gS1LoopEnd; + for (uint32_t gS1LoopIdx = constInfo.gS1Start; gS1LoopIdx < gS1LoopEnd; gS1LoopIdx++) { + tempLoopInfo.actOriS2Size = GetActualSeqLenKV(tempLoopInfo.bIdx); + tempLoopInfo.actCmpS2Size = GetActualSeqLenCmpKV(tempLoopInfo.bIdx, tempLoopInfo.actOriS2Size); + // 计算需要的数据, 避免重复计算 + tempLoopInfo.gS1Idx = gS1LoopIdx * constInfo.mBaseSize; + tempLoopInfo.s1StartIdx = tempLoopInfo.gS1Idx / constInfo.gSize; + tempLoopInfo.s1EndIdx = + Min((tempLoopInfo.s1StartIdx + constInfo.mBaseSize / constInfo.gSize - 1), tempLoopInfo.actS1Size - 1); + + // 此处均为闭区间 + tempLoopInfo.oriMaskRight = tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size + + static_cast(tempLoopInfo.s1EndIdx) + constInfo.oriWinRight; + tempLoopInfo.oriMaskLeft = Max(tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size + + static_cast(tempLoopInfo.s1EndIdx) - constInfo.oriWinLeft, + 0); + int32_t cmpMaskS2Size = + GetCmpMaskS2Size(tempLoopInfo.bIdx, tempLoopInfo.actOriS2Size, tempLoopInfo.actCmpS2Size); + tempLoopInfo.cmpMaskRight = cmpMaskS2Size - tempLoopInfo.actS1Size; + GetSparseActualSeqLen(); + UpdateInnerLoopCond(); + + uint32_t oriS2Size = tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1; + uint32_t oriSplitNum = 0; + uint32_t cmpSplitNum = 0; + uint32_t cmpS2Size = 0; + bool isEnd = (bN2LoopIdx + 1 == constInfo.bN2End) && (gS1LoopIdx + 1 == gS1LoopEnd); + if (tempLoopInfo.curActSeqLenIsZero) { + if ASCEND_IS_AIV { + InitAllZeroOutput(tempLoopInfo.bIdx, tempLoopInfo.s1StartIdx, tempLoopInfo.n2Idx); + } + if (!isEnd) { + continue; + } + } else { + oriSplitNum = CeilDiv(oriS2Size, constInfo.s2BaseSize); + cmpS2Size = tempLoopInfo.actCmpS2Size; + cmpSplitNum = CeilDiv(cmpS2Size, constInfo.s2BaseSize); + } + + uint32_t s2SplitNum = oriSplitNum + cmpSplitNum; + constexpr uint32_t V0_SPLIT = 32; // align to 32 + uint32_t v0OriSize = CeilDiv(oriS2Size * cmpS2Size, oriS2Size + cmpS2Size); + if (cmpS2Size > V0_SPLIT * oriSplitNum) { + v0OriSize = SMLAAlign(v0OriSize, V0_SPLIT * oriSplitNum); + } + uint32_t v0CmpSize = cmpS2Size - v0OriSize; + + tempLoopInfo.oriLoopTimes = oriSplitNum; + tempLoopInfo.cmpLoopTimes = cmpSplitNum; + tempLoopInfo.s2LoopTimes = s2SplitNum; + tempLoopInfo.v0OriSize = v0OriSize; + tempLoopInfo.v0CmpSize = v0CmpSize; + + uint32_t s2LoopEnd = (isEnd && constInfo.s2End != 0) ? constInfo.s2End : tempLoopInfo.s2LoopTimes; + tempLoopInfo.s2LoopTimes = s2LoopEnd; + // 分核修改后需要打开 + // 当前s2是否被切,决定了输出是否要写到attenOut上 + tempLoopInfo.tndIsS2SplitCore = ((constInfo.s2Start == 0) && (s2LoopEnd == s2SplitNum)) ? false : true; + tempLoopInfo.tndCoreStartKVSplitPos = globalLoopStart ? constInfo.coreStartKVSplitPos : 0; + uint32_t extraLoop = isEnd ? 2 : 0; + uint32_t curTopKIdx = 0; + for (uint32_t s2LoopIdx = constInfo.s2Start; s2LoopIdx < (s2LoopEnd + extraLoop); s2LoopIdx++) { + PreloadPipeline(gloop, cmpLoop, constInfo.s2Start, s2LoopIdx, extraInfo); + ++gloop; + if (s2LoopIdx >= tempLoopInfo.oriLoopTimes && s2LoopIdx < s2LoopEnd) { // 用于判断v0使用的循环GM的id + ++cmpLoop; + } + } + globalLoopStart = false; + constInfo.s2Start = 0; + } + constInfo.gS1Start = 0; + } + if ASCEND_IS_AIV { + CrossCoreWaitFlag(3); + CrossCoreWaitFlag(3); + CrossCoreWaitFlag(3); + CrossCoreWaitFlag(3); + } +} + +template +__aicore__ inline void SparseFlashMlaCsa::PreloadPipeline(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, + uint64_t s2LoopIdx, + RunInfo extraInfo[SMLA_PRELOAD_TASK_CACHE_SIZE]) +{ + RunInfo &extraInfo0 = extraInfo[loop % SMLA_PRELOAD_TASK_CACHE_SIZE]; // 本轮任务 + RunInfo &extraInfo2 = extraInfo[(loop + 2) % SMLA_PRELOAD_TASK_CACHE_SIZE]; // 上一轮任务 + RunInfo &extraInfo1 = extraInfo[(loop + 1) % SMLA_PRELOAD_TASK_CACHE_SIZE]; // 上两轮任务 + + CalcParams(loop, cmpLoop, s2Start, s2LoopIdx, extraInfo0); + if constexpr (!HEAD_RATIO_ONE) { + if (extraInfo0.isValid) { + if ASCEND_IS_AIC { + if (!extraInfo0.isOriOnly) { + CrossCoreWaitFlag(constInfo.syncV0C1); + } + ComputeMm1(extraInfo0); + } else { + if (extraInfo0.isFirstSInnerLoop) { + CrossCoreWaitFlag(3); + } + vectorBlock.ProcessVec0L(extraInfo0); + if (!extraInfo0.isOriOnly) { + CrossCoreSetFlag(constInfo.syncV0C1); + } + } + } + if (extraInfo2.isValid) { + if ASCEND_IS_AIV { + vectorBlock.ProcessVec1L(extraInfo2); + } + if ASCEND_IS_AIC { + ComputeMm2(extraInfo2); + if (extraInfo2.isLastS2Loop) { + CrossCoreSetFlag(3); + } + } + } + if (extraInfo1.isValid) { + if ASCEND_IS_AIV { + vectorBlock.ProcessVec2L(extraInfo1); + } + extraInfo1.isValid = false; + } + } else { + if (extraInfo0.isValid) { + if ASCEND_IS_AIV { + if (extraInfo0.isFirstSInnerLoop) { + CrossCoreWaitFlag(3); + } + vectorBlock.ProcessVec0L(extraInfo0); + if (!extraInfo0.isOriOnly) { + CrossCoreSetFlag(constInfo.syncV0C1); + } + } + } + if (extraInfo2.isValid) { + if ASCEND_IS_AIV { + vectorBlock.ProcessVec1L(extraInfo2); + } + if ASCEND_IS_AIC { + ComputeMm2(extraInfo2); + if (extraInfo2.isLastS2Loop) { + CrossCoreSetFlag(3); + } + } + } + if (extraInfo1.isValid) { + if ASCEND_IS_AIV { + vectorBlock.ProcessVec2L(extraInfo1); + } + extraInfo1.isValid = false; + } + if (extraInfo0.isValid) { + if ASCEND_IS_AIC { + if (!extraInfo0.isOriOnly) { + CrossCoreWaitFlag(constInfo.syncV0C1); + } + ComputeMm1(extraInfo0); + } + } + } +} + +} // namespace SMLAKernel +#endif // SPARSE_FLASH_MLA_CSA_KERNEL_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_cube.h b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_cube.h new file mode 100644 index 000000000000..c948393c7661 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_cube.h @@ -0,0 +1,1292 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_swa_block_cube.h + * \brief use 7 buffer for matmul l1, better pipeline + */ +#ifndef SPARSE_FLASH_MLA_SWA_BLOCK_CUBE_H +#define SPARSE_FLASH_MLA_SWA_BLOCK_CUBE_H + +#include "kernel_operator.h" +#include "kernel_operator_list_tensor_intf.h" +#include "kernel_tiling/kernel_tiling.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" +#include "sparse_flash_mla_common_arch22.h" + +namespace SMLAKernel { +template +class SWACubeBlock { +public: + // 中间计算数据类型为float, 高精度模式 + using T = float; + using Q_T = typename SMLAT::queryType; + using KV_T = typename SMLAT::kvType; + using OUT_T = typename SMLAT::outputType; + using MM_OUT_T = T; + + __aicore__ inline SWACubeBlock(){}; + __aicore__ inline void InitParams(const ConstInfo &constInfo); + __aicore__ inline void InitMm1GlobalTensor(GlobalTensor queryGm, GlobalTensor oriKvGm, + GlobalTensor cmpKV, GlobalTensor mm1ResGm); + __aicore__ inline void InitMm2GlobalTensor(GlobalTensor vec1ResGm, GlobalTensor mm2ResGm, + GlobalTensor attentionOutGm); + __aicore__ inline void InitPageAttentionInfo(GlobalTensor oriKvGm, GlobalTensor kvMergeGm, + GlobalTensor oriBlockTableGm, + GlobalTensor cmpBlockTableGm); + __aicore__ inline void InitBuffers(TPipe *pipe); + + __aicore__ inline void AllocEventID(); + __aicore__ inline void FreeEventID(); + __aicore__ inline void ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo); + __aicore__ inline void ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo); + +private: + static constexpr bool PAGE_ATTENTION = SMLAT::pageAttention; + static constexpr int TEMPLATE_MODE = SMLAT::templateMode; + static constexpr bool FLASH_DECODE = SMLAT::flashDecode; + static constexpr SMLA_LAYOUT LAYOUT_T = SMLAT::layout; + static constexpr SMLA_LAYOUT KV_LAYOUT_T = SMLAT::kvLayout; + + static constexpr uint32_t M_SPLIT_SIZE = 128; // m方向切分 + static constexpr uint32_t N_SPLIT_SIZE = 128; // n方向切分 + static constexpr uint32_t K_L0_SPLIT_SIZE = 128; // k方向L0切分 + static constexpr uint32_t K_L1_SPLIT_SIZE = 256; // k方向L1切分 + static constexpr uint32_t N_WORKSPACE_SIZE = 512; // n方向切分 + static constexpr uint32_t MERGE_CACHE_GM_BUF_NUM = 3; + static constexpr uint32_t D_SPLIT_SIZE = 256; // d轴切分 + + static constexpr uint32_t L1_BLOCK_SIZE = (64 * 512 * sizeof(Q_T)); + static constexpr uint32_t L1_BLOCK_OFFSET = 64 * 512; + + static constexpr uint32_t L0A_PP_SIZE = (32 * 1024); + static constexpr uint32_t L0B_PP_SIZE = (32 * 1024); + static constexpr uint32_t L0C_PP_SIZE = (64 * 1024); + + // mte2 <> mte1 EventID + // L1 3buf, 使用3个eventId + static constexpr uint32_t L1_EVENT0 = EVENT_ID2; + static constexpr uint32_t L1_EVENT1 = EVENT_ID3; + static constexpr uint32_t L1_EVENT2 = EVENT_ID4; + static constexpr uint32_t L1_EVENT3 = EVENT_ID5; + static constexpr uint32_t L1_EVENT4 = EVENT_ID6; + static constexpr uint32_t L1_EVENT5 = EVENT_ID7; + static constexpr uint32_t L1_EVENT6 = EVENT_ID1; + + // m <> mte1 EventID + static constexpr uint32_t L0AB_EVENT0 = EVENT_ID3; + static constexpr uint32_t L0AB_EVENT1 = EVENT_ID4; + + static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding; + static constexpr uint32_t mte21QPIds[4] = {L1_EVENT0, L1_EVENT1, L1_EVENT2, L1_EVENT3}; // mte12复用 + static constexpr uint32_t mte21KVIds[3] = {L1_EVENT4, L1_EVENT5, L1_EVENT6}; + + ConstInfo constInfo{}; + + // L1分成3块buf, 用于记录 + uint32_t qpL1BufIter = 0; + uint32_t kvL1BufIter = -1; + uint32_t abL0BufIter = 0; + uint32_t cL0BufIter = 0; + + // mm1 + GlobalTensor queryGm; + GlobalTensor keyGm; + GlobalTensor mm1ResGm; + GlobalTensor kvMergeGm_; + GlobalTensor oriKvGm; + GlobalTensor cmpKvGm; + + // mm2 + GlobalTensor vec1ResGm; + GlobalTensor valueGm; + GlobalTensor mm2ResGm; + GlobalTensor attentionOutGm; + + // block_table + GlobalTensor oriBlockTableGm; + GlobalTensor cmpBlockTableGm; + + TBuf bufQPL1; + TBuf bufKVL1; + TBuf tmpBufL0A; + TBuf tmpBufL0B; + TBuf tmpBufL0C; + + LocalTensor l1QPTensor; + LocalTensor l1KVTensor; + LocalTensor aL0TensorPingPong; + LocalTensor bL0TensorPingPong; + LocalTensor cL0TensorPingPong; + + // L0AB m <> mte1 EventID + __aicore__ inline uint32_t Mte1MmABEventId(uint32_t idx) + { + return (L0AB_EVENT0 + idx); + } + + __aicore__ inline uint32_t GetQPL1RealIdx(uint32_t mIdx, uint32_t k1Idx) + { + uint32_t idxMap[] = {0, 2}; // 确保0块和1块连在一起, 2和3块连在一起, 来保证同一m块的地址相连 + return idxMap[mIdx % 2] + k1Idx; + } + + __aicore__ inline void CopyGmToL1(LocalTensor &l1Tensor, GlobalTensor &gmSrcTensor, uint32_t srcN, + uint32_t srcD, uint32_t srcDstride); + __aicore__ inline void CopyInMm1AToL1(LocalTensor &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx, + uint32_t mSizeAct, uint32_t headSize, uint32_t headOffset); + __aicore__ inline void CopyInMm2AToL1(LocalTensor &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx, + uint32_t subMSizeAct, uint32_t nSize, uint32_t nOffset); + __aicore__ inline void LoadDataMm1A(LocalTensor &aL0Tensor, LocalTensor &aL1Tensor, uint32_t idx, + uint32_t kSplitSize, uint32_t mSize, uint32_t kSize); + __aicore__ inline void LoadDataMm1B(LocalTensor &bL0Tensor, LocalTensor &bL1Tensor, uint32_t idx, + uint32_t kSplitSize, uint32_t kSize, uint32_t nSize); +}; + +template +__aicore__ inline void SWACubeBlock::InitParams(const ConstInfo &constInfo) +{ + this->constInfo = constInfo; +} + +template +__aicore__ inline void SWACubeBlock::InitMm1GlobalTensor(GlobalTensor queryGm, GlobalTensor oriKvGm, + GlobalTensor cmpKvGm, + GlobalTensor mm1ResGm) +{ + // mm1 + this->queryGm = queryGm; + this->oriKvGm = oriKvGm; + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + this->cmpKvGm = cmpKvGm; + } + this->mm1ResGm = mm1ResGm; +} + +template +__aicore__ inline void SWACubeBlock::InitMm2GlobalTensor(GlobalTensor vec1ResGm, + GlobalTensor mm2ResGm, + GlobalTensor attentionOutGm) +{ + // mm2 + this->vec1ResGm = vec1ResGm; + this->mm2ResGm = mm2ResGm; + this->attentionOutGm = attentionOutGm; +} + +template +__aicore__ inline void SWACubeBlock::InitPageAttentionInfo(GlobalTensor oriKvGm, + GlobalTensor kvMergeGm, + GlobalTensor oriBlockTableGm, + GlobalTensor cmpBlockTableGm) +{ + this->oriKvGm = oriKvGm; + this->kvMergeGm_ = kvMergeGm; + this->oriBlockTableGm = oriBlockTableGm; + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + this->cmpBlockTableGm = cmpBlockTableGm; + } +} + +template +__aicore__ inline void SWACubeBlock::InitBuffers(TPipe *pipe) +{ + pipe->InitBuffer(bufQPL1, L1_BLOCK_SIZE * 4); + l1QPTensor = bufQPL1.Get(); + pipe->InitBuffer(bufKVL1, L1_BLOCK_SIZE * 3); + l1KVTensor = bufKVL1.Get(); + // L0A + pipe->InitBuffer(tmpBufL0A, L0A_PP_SIZE * 2); // 64K + aL0TensorPingPong = tmpBufL0A.Get(); + // L0B + pipe->InitBuffer(tmpBufL0B, L0B_PP_SIZE * 2); // 64K + bL0TensorPingPong = tmpBufL0B.Get(); + // L0C + pipe->InitBuffer(tmpBufL0C, L0C_PP_SIZE * 2); // 128K + cL0TensorPingPong = tmpBufL0C.Get(); +} + +template +__aicore__ inline void SWACubeBlock::AllocEventID() +{ + SetFlag(L1_EVENT0); + SetFlag(L1_EVENT1); + SetFlag(L1_EVENT2); + SetFlag(L1_EVENT3); + SetFlag(L1_EVENT4); + SetFlag(L1_EVENT5); + SetFlag(L1_EVENT6); + SetFlag(L0AB_EVENT0); + SetFlag(L0AB_EVENT1); +} + +template +__aicore__ inline void SWACubeBlock::FreeEventID() +{ + WaitFlag(L1_EVENT0); + WaitFlag(L1_EVENT1); + WaitFlag(L1_EVENT2); + WaitFlag(L1_EVENT3); + WaitFlag(L1_EVENT4); + WaitFlag(L1_EVENT5); + WaitFlag(L1_EVENT6); + WaitFlag(L0AB_EVENT0); + WaitFlag(L0AB_EVENT1); +} + +template +__aicore__ inline void SWACubeBlock::CopyGmToL1(LocalTensor &l1Tensor, GlobalTensor &gmSrcTensor, + uint32_t srcN, uint32_t srcD, uint32_t srcDstride) +{ + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = srcN; // 行数 + nd2nzPara.dValue = srcD; + nd2nzPara.srcDValue = srcDstride; + nd2nzPara.dstNzC0Stride = (srcN + 15) / 16 * 16; // 对齐到16 单位block + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(l1Tensor, gmSrcTensor, nd2nzPara); +} + +template +__aicore__ inline void SWACubeBlock::CopyInMm1AToL1(LocalTensor &l1Tensor, const RunInfo &info, + uint32_t mSeqIdx, uint32_t mSizeAct, uint32_t headSize, + uint32_t headOffset) +{ + auto srcGm = queryGm[info.tensorAOffset + mSeqIdx * constInfo.headDim + headOffset]; + CopyGmToL1(l1Tensor, srcGm, mSizeAct, headSize, constInfo.headDim); +} + +template +__aicore__ inline void SWACubeBlock::LoadDataMm1A(LocalTensor &aL0Tensor, LocalTensor &aL1Tensor, + uint32_t idx, uint32_t kSplitSize, uint32_t mSize, + uint32_t kSize) +{ + LocalTensor srcTensor = aL1Tensor[mSize * kSplitSize * idx]; + LoadData3DParamsV2 loadData3DParams; + // SetFmatrixParams + loadData3DParams.l1H = mSize / 16; // Hin=M1=8 + loadData3DParams.l1W = 16; // Win=M0 + loadData3DParams.padList[0] = 0; + loadData3DParams.padList[1] = 0; + loadData3DParams.padList[2] = 0; + loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果 + + // SetLoadToA0Params + loadData3DParams.mExtension = mSize; // M + loadData3DParams.kExtension = kSize; // K + loadData3DParams.mStartPt = 0; + loadData3DParams.kStartPt = 0; + loadData3DParams.strideW = 1; + loadData3DParams.strideH = 1; + loadData3DParams.filterW = 1; + loadData3DParams.filterSizeW = (1 >> 8) & 255; + loadData3DParams.filterH = 1; + loadData3DParams.filterSizeH = (1 >> 8) & 255; + loadData3DParams.dilationFilterW = 1; + loadData3DParams.dilationFilterH = 1; + loadData3DParams.enTranspose = 0; + loadData3DParams.fMatrixCtrl = 0; + loadData3DParams.channelSize = kSize; // Cin=K + LoadData(aL0Tensor, srcTensor, loadData3DParams); +} + +template +__aicore__ inline void SWACubeBlock::LoadDataMm1B(LocalTensor &l0Tensor, LocalTensor &l1Tensor, + uint32_t idx, uint32_t kSplitSize, uint32_t kSize, + uint32_t nSize) +{ + // N 方向全载 + LocalTensor srcTensor = l1Tensor[nSize * kSplitSize * idx]; + LoadData2DParams loadData2DParams; + loadData2DParams.startIndex = 0; + loadData2DParams.repeatTimes = (nSize + 15) / 16 * kSize / (32 / sizeof(KV_T)); + loadData2DParams.srcStride = 1; + loadData2DParams.dstGap = 0; + loadData2DParams.ifTranspose = false; + LoadData(l0Tensor, srcTensor, loadData2DParams); +} + +template +__aicore__ inline void SWACubeBlock::CopyInMm2AToL1(LocalTensor &aL1Tensor, const RunInfo &info, + uint32_t mSeqIdx, uint32_t subMSizeAct, uint32_t nSize, + uint32_t nOffset) +{ + auto srcGm = vec1ResGm[(info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize + + mSeqIdx * info.actualSingleProcessSInnerSizeAlign + nOffset]; + CopyGmToL1(aL1Tensor, srcGm, subMSizeAct, nSize, info.actualSingleProcessSInnerSizeAlign); +} + +template +__aicore__ inline void SWACubeBlock::ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo) +{ + uint32_t mSize = mSplitInfo.nBufferDealM; + uint32_t mL1Size = M_SPLIT_SIZE; + uint32_t mL1SizeAlign = SMLAAlign(M_SPLIT_SIZE, 16); + uint32_t mL1Loops = CeilDiv(mSize, M_SPLIT_SIZE); + + uint32_t nSize = info.actualSingleProcessSInnerSize; + uint32_t nL1Size = N_SPLIT_SIZE; + uint32_t nL1SizeAlign = SMLAAlign(N_SPLIT_SIZE, 16); + uint32_t nL1Loops = CeilDiv(nSize, N_SPLIT_SIZE); + + uint32_t kSize = 512; + uint32_t kL1Size = 256; + uint32_t kL1Loops = 2; + uint32_t kL0Size = 128; + uint32_t kL0Loops = CeilDiv(kL1Size, kL0Size); + + LocalTensor bL1Tensor; + LocalTensor kTensor; + uint32_t ka = 0, kb = 0; + uint32_t copyRowCnt = 0; + uint32_t copyRowCntTmp = 0; + // L1 切n切k + for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // L1切n, 512/128=4 + if (nL1 == (nL1Loops - 1)) { + // 尾块重新计算size + nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE; + nL1SizeAlign = SMLAAlign(nL1Size, 16); + } + + for (uint32_t kL1 = 0; kL1 < kL1Loops; kL1++) { + kvL1BufIter++; + uint32_t kb = kvL1BufIter % 3; + WaitFlag(mte21KVIds[kb]); + // 从k当中取当前的块 + bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET]; + uint32_t copyFinishRowCnt = 0; + + if (info.isOriOnly) { + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + uint64_t curS2Offset = static_cast(info.s2Idx) * constInfo.s2BaseSize + + info.s2StartPoint + nL1 * N_SPLIT_SIZE; + uint32_t copyFinishRowCnt = 0; + LocalTensor kTensor; + uint32_t copyRowCnt = 0; + + if (constInfo.hasOriSparseIndices) { + uint64_t mergeBase = static_cast(info.cmpLoop % MERGE_CACHE_GM_BUF_NUM) * + N_WORKSPACE_SIZE * constInfo.headDim; + uint32_t blockElementCnt = 32U / sizeof(KV_T); + for (uint32_t row = 0; row < nL1Size; ++row) { + uint64_t mergeOffset = mergeBase + + static_cast(nL1 * N_SPLIT_SIZE + row) * constInfo.headDim + + kL1 * D_SPLIT_SIZE; + GlobalTensor mergeSrcGm = kvMergeGm_[mergeOffset]; + kTensor = bL1Tensor[row * blockElementCnt]; + DataCopyGmNDToL1(kTensor, mergeSrcGm, 1, nL1SizeAlign, D_SPLIT_SIZE, + constInfo.headDim); + } + } else + while (copyFinishRowCnt < nL1Size) { + // 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位 + copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize; + if (copyFinishRowCnt + copyRowCnt > nL1Size) { + copyRowCnt = nL1Size - copyFinishRowCnt; + } + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + // 256、32等待7buf命名更改 + startPos.dIdx = kL1 * 256; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + PAShape shape; + shape.blockSize = constInfo.paOriBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.oriKvStride0; + shape.actHeadDim = 256; + shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = nL1SizeAlign; + kTensor = bL1Tensor[copyFinishRowCnt * 16]; + DataCopyPA(kTensor, oriKvGm, oriBlockTableGm, shape, startPos); + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; // 行数 + nd2nzPara.dValue = D_SPLIT_SIZE; // 256 + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + uint32_t headStride = constInfo.headDim; + uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim; + uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? + static_cast(constInfo.kvSeqSize) * seqStride : + constInfo.oriKvStride0; + + uint64_t curS2 = info.s2StartPoint + nL1 * N_SPLIT_SIZE; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE; + DataCopy(bL1Tensor, oriKvGm[offset], nd2nzPara); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + uint64_t curS2Offset = + static_cast(info.s2Idx) * static_cast(constInfo.s2BaseSize) + + info.s2StartPoint + nL1 * N_SPLIT_SIZE; + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; + nd2nzPara.dValue = constInfo.headDim >> 1; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor, + oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + kL1 * D_SPLIT_SIZE], + nd2nzPara); + } + } else if (info.isOriCmpMix) { + uint32_t cumNl1LoopSize = nL1 * N_SPLIT_SIZE; + uint32_t oriSizeCur = (info.actualSingleProcessSInnerOriSize < cumNl1LoopSize) ? + 0 : + (info.actualSingleProcessSInnerOriSize - cumNl1LoopSize); + oriSizeCur = (oriSizeCur > nL1Size) ? nL1Size : oriSizeCur; + uint32_t cmpSizeCur = nL1Size - oriSizeCur; + cmpSizeCur = (cmpSizeCur < info.actualSingleProcessSInnerCmpSize) ? + cmpSizeCur : + info.actualSingleProcessSInnerCmpSize; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + LocalTensor kTensor; + uint32_t copyRowCnt = 0; + + if (oriSizeCur > 0) { + uint32_t copyFinishRowCnt = 0; + uint64_t curS2Offset = info.s2StartPoint + nL1 * N_SPLIT_SIZE; + while (copyFinishRowCnt < oriSizeCur) { + // 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位 + copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize; + if (copyFinishRowCnt + copyRowCnt > oriSizeCur) { + copyRowCnt = oriSizeCur - copyFinishRowCnt; + } + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + // 256、32等待7buf命名更改 + startPos.dIdx = kL1 * 256; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + PAShape shape; + shape.blockSize = constInfo.paOriBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.oriKvStride0; + shape.actHeadDim = 256; + shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = nL1SizeAlign; + kTensor = bL1Tensor[copyFinishRowCnt * 16]; + DataCopyPA(kTensor, oriKvGm, oriBlockTableGm, shape, startPos); + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } + if (cmpSizeCur > 0) { + uint32_t cmpMixsizeCur = N_SPLIT_SIZE - info.actualSingleProcessSInnerOriSize % N_SPLIT_SIZE; + uint32_t oriOnlyLoopTimes = info.actualSingleProcessSInnerOriSize / N_SPLIT_SIZE; + uint32_t cmpLoopTimes = nL1 - oriOnlyLoopTimes; + uint64_t curS2Offset = + (cmpLoopTimes == 0) ? + 0 : + (cmpMixsizeCur + static_cast(cmpLoopTimes - 1) * N_SPLIT_SIZE); + uint32_t copyFinishRowCnt = 0; + while (copyFinishRowCnt < cmpSizeCur) { + // 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位 + copyRowCnt = constInfo.paCmpBlockSize - curS2Offset % constInfo.paCmpBlockSize; + if (copyFinishRowCnt + copyRowCnt > cmpSizeCur) { + copyRowCnt = cmpSizeCur - copyFinishRowCnt; + } + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + // 256、32等待7buf命名更改 + startPos.dIdx = kL1 * 256; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + PAShape shape; + shape.blockSize = constInfo.paCmpBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.cmpKvStride0; + shape.actHeadDim = 256; + shape.maxblockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = nL1SizeAlign; + kTensor = bL1Tensor[copyFinishRowCnt * 16 + oriSizeCur * 16]; + DataCopyPA(kTensor, cmpKvGm, cmpBlockTableGm, shape, startPos); + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; // 行数 + nd2nzPara.dValue = D_SPLIT_SIZE; // 256 + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + uint32_t headStride = constInfo.headDim; + uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim; + if (oriSizeCur > 0) { + nd2nzPara.nValue = oriSizeCur; + uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? + static_cast(constInfo.kvSeqSize) * seqStride : + constInfo.oriKvStride0; + uint64_t curS2 = info.s2StartPoint + nL1 * N_SPLIT_SIZE; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE; + DataCopy(bL1Tensor, oriKvGm[offset], nd2nzPara); + } + if (cmpSizeCur > 0) { + nd2nzPara.nValue = cmpSizeCur; + uint64_t batchStride = (constInfo.cmpKvStride0 == 0) ? + static_cast(constInfo.cmpSeqSize) * seqStride : + constInfo.cmpKvStride0; + uint32_t cmpMixsizeCur = N_SPLIT_SIZE - info.actualSingleProcessSInnerOriSize % N_SPLIT_SIZE; + uint32_t oriOnlyLoopTimes = info.actualSingleProcessSInnerOriSize / N_SPLIT_SIZE; + uint32_t cmpLoopTimes = nL1 - oriOnlyLoopTimes; + uint64_t curS2 = (cmpLoopTimes == 0) ? 0 : (cmpMixsizeCur + (cmpLoopTimes - 1) * N_SPLIT_SIZE); + uint64_t offset = (uint64_t)info.bIdx * batchStride + curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE; + DataCopy(bL1Tensor[oriSizeCur * 16], cmpKvGm[offset], nd2nzPara); + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; + nd2nzPara.dValue = constInfo.headDim >> 1; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + if (oriSizeCur > 0) { + nd2nzPara.nValue = oriSizeCur; + uint64_t curS2Offset = info.s2StartPoint + nL1 * N_SPLIT_SIZE; + DataCopy(bL1Tensor, + oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + kL1 * D_SPLIT_SIZE], + nd2nzPara); + } + if (cmpSizeCur > 0) { + nd2nzPara.nValue = cmpSizeCur; + uint32_t cmpMixsizeCur = N_SPLIT_SIZE - info.actualSingleProcessSInnerOriSize % N_SPLIT_SIZE; + uint32_t oriOnlyLoopTimes = info.actualSingleProcessSInnerOriSize / N_SPLIT_SIZE; + uint32_t cmpLoopTimes = nL1 - oriOnlyLoopTimes; + uint64_t curS2Offset = + (cmpLoopTimes == 0) ? + 0 : + (cmpMixsizeCur + static_cast(cmpLoopTimes - 1) * N_SPLIT_SIZE); + DataCopy(bL1Tensor[oriSizeCur * 16], + cmpKvGm[info.tensorCmpBOffset + curS2Offset * constInfo.headDim + kL1 * D_SPLIT_SIZE], + nd2nzPara); + } + } + } else { + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + uint64_t curS2Offset = static_cast(info.relativeS2Idx) * constInfo.s2BaseSize + + info.s2StartPoint + nL1 * N_SPLIT_SIZE; + while (copyFinishRowCnt < nL1Size) { + // 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位 + copyRowCnt = constInfo.paCmpBlockSize - curS2Offset % constInfo.paCmpBlockSize; + if (copyFinishRowCnt + copyRowCnt > nL1Size) { + copyRowCnt = nL1Size - copyFinishRowCnt; + } + + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + // 256、32等待7buf命名更改 + startPos.dIdx = kL1 * 256; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + + PAShape shape; + shape.blockSize = constInfo.paCmpBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.cmpKvStride0; + shape.actHeadDim = 256; + shape.maxblockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = nL1SizeAlign; + kTensor = bL1Tensor[copyFinishRowCnt * 16]; + DataCopyPA(kTensor, cmpKvGm, cmpBlockTableGm, shape, startPos); + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; // 行数 + nd2nzPara.dValue = D_SPLIT_SIZE; // 256 + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + uint32_t headStride = constInfo.headDim; + uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim; + uint64_t batchStride = (constInfo.cmpKvStride0 == 0) ? + static_cast(constInfo.cmpSeqSize) * seqStride : + constInfo.cmpKvStride0; + + uint64_t curS2 = static_cast(info.relativeS2Idx) * constInfo.s2BaseSize + + info.s2StartPoint + nL1 * N_SPLIT_SIZE; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE; + DataCopy(bL1Tensor, cmpKvGm[offset], nd2nzPara); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + uint64_t curS2Offset = static_cast(info.relativeS2Idx) * constInfo.s2BaseSize + + info.s2StartPoint + nL1 * N_SPLIT_SIZE; + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = nL1Size; + nd2nzPara.dValue = constInfo.headDim >> 1; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = nL1SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor, + cmpKvGm[info.tensorCmpBOffset + curS2Offset * constInfo.headDim + +kL1 * D_SPLIT_SIZE], + nd2nzPara); + } + } + + SetFlag(mte21KVIds[kb]); + WaitFlag(mte21KVIds[kb]); + mL1Size = M_SPLIT_SIZE; + mL1SizeAlign = SMLAAlign(M_SPLIT_SIZE, 16U); + for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) { + uint32_t aL1PaddingSize = 0; // 用于使左矩阵对齐到尾部, 以保证两块32K内存连续 + if (mL1 == (mL1Loops - 1)) { + mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE; + mL1SizeAlign = SMLAAlign(mL1Size, 16U); + aL1PaddingSize = (M_SPLIT_SIZE - mL1SizeAlign) * 256; + } + uint32_t mIdx = qpL1BufIter + mL1; + ka = GetQPL1RealIdx(mIdx, kL1); + LocalTensor aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET + (1 - kL1) * aL1PaddingSize]; + if (nL1 == 0) { + if (kL1 == 0) { + WaitFlag(mte21QPIds[ka]); + WaitFlag(mte21QPIds[ka + 1]); + CopyInMm1AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 256, 0); + } else { + LocalTensor qTmpTensor = aL1Tensor; + CopyInMm1AToL1(qTmpTensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 256, + 256); + } + SetFlag(mte21QPIds[ka]); + WaitFlag(mte21QPIds[ka]); + } + // 使用unitflag同步 + LocalTensor cL0Tensor = + cL0TensorPingPong[(cL0BufIter % 2) * + (L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致 + for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) { + WaitFlag(Mte1MmABEventId(abL0BufIter % 2)); + LocalTensor aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE / sizeof(KV_T))]; + LoadDataMm1A(aL0Tensor, aL1Tensor, kL0, kL0Size, mL1SizeAlign, kL0Size); + LocalTensor bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE / sizeof(KV_T))]; + if (mL1 == 0) { + LoadDataMm1B(bL0Tensor, bL1Tensor, kL0, kL0Size, kL0Size, nL1SizeAlign); + } + SetFlag(Mte1MmABEventId(abL0BufIter % 2)); + WaitFlag(Mte1MmABEventId(abL0BufIter % 2)); + + MmadParams mmadParams; + mmadParams.m = mL1SizeAlign; + mmadParams.n = nL1SizeAlign; + mmadParams.k = kL0Size; + mmadParams.cmatrixInitVal = (kL1 == 0 && kL0 == 0); + mmadParams.cmatrixSource = false; + mmadParams.unitFlag = + (kL1 == 1 && kL0 == (kL0Loops - 1)) ? 0b11 : 0b10; // 累加最后一次翻转flag, 表示可以搬出 + Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams); + if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) { + PipeBarrier(); + } + SetFlag(Mte1MmABEventId(abL0BufIter % 2)); + abL0BufIter++; + } + + if (nL1 == (nL1Loops - 1)) { + SetFlag(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完 + } + + if (kL1 == 1) { // 最后一轮kL1循环 + FixpipeParamsV220 fixParams; + fixParams.nSize = nL1SizeAlign; + fixParams.mSize = mL1SizeAlign; + fixParams.srcStride = mL1SizeAlign; + // 改成nSizeAlign + fixParams.dstStride = info.actualSingleProcessSInnerSizeAlign; // mm1ResGm两行之间的间隔 + fixParams.unitFlag = 0b11; + fixParams.ndNum = 1; // 输出ND + + Fixpipe(mm1ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.mmResUbSize + nL1 * N_SPLIT_SIZE + + (mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) * + info.actualSingleProcessSInnerSizeAlign], + cL0Tensor, fixParams); + } + if (mL1Loops == 2) { + cL0BufIter++; + } + } + + SetFlag(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完 + } + if (mL1Loops == 1) { + cL0BufIter++; + } + } + qpL1BufIter += mL1Loops; +} + +template +__aicore__ inline void SWACubeBlock::ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo) +{ + uint32_t mSize = mSplitInfo.nBufferDealM; + uint32_t mSizeAlign = (mSize + 16 - 1) / 16; + uint32_t mL1Loops = (mSize + M_SPLIT_SIZE - 1) / M_SPLIT_SIZE; + uint32_t mL1SizeAlign = M_SPLIT_SIZE; // 16对齐 + uint32_t mL1Size = M_SPLIT_SIZE; // m的实际大小 + + uint32_t nSize = BlockAlign(constInfo.headDim); + uint32_t nL1Loops = (nSize + N_SPLIT_SIZE - 1) / N_SPLIT_SIZE; + uint32_t nL1SizeAlign = N_SPLIT_SIZE; // 16对齐 + uint32_t nL1Size = N_SPLIT_SIZE; // n的实际大小 + + uint32_t kSize = info.actualSingleProcessSInnerSize; + uint32_t kL1Size = K_L1_SPLIT_SIZE; + uint32_t kL1SizeAlign = SMLAAlign(kL1Size, 16U); + uint32_t kL1Loops = (kSize + kL1Size - 1) / kL1Size; + uint32_t kL0Size = K_L0_SPLIT_SIZE; + uint32_t kL0Loops = (kL1Size + kL0Size - 1) / kL0Size; + uint32_t kL0SizeAlign = kL0Size; + LocalTensor bL1Tensor; + LocalTensor subvTensor; + // ka表示左矩阵4buf选择哪一块buf, kb表示右矩阵3buf选择哪一块buf + uint32_t ka = 0, kb = 0; + uint32_t mBaseIdx = qpL1BufIter; + for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // n切L1 + if (nL1 == (nL1Loops - 1)) { + // 尾块 + nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE; + nL1SizeAlign = SMLAAlign(nL1Size, 16U); + } + // k l1写成一个循环, 和mm1保持一致 + kL1Size = K_L1_SPLIT_SIZE; + kL1SizeAlign = SMLAAlign(kL1Size, 16U); + uint32_t copyRowCnt = 0; + for (uint32_t k1 = 0; k1 < kL1Loops; k1++) { // k切L1, 这里套了一层l0来操作 + if (k1 == (kL1Loops - 1)) { + // 尾块 + kL1Size = kSize - (kL1Loops - 1) * K_L1_SPLIT_SIZE; + kL1SizeAlign = SMLAAlign(kL1Size, 16U); + } + kvL1BufIter++; + uint32_t kb = kvL1BufIter % 3; + WaitFlag(mte21KVIds[kb]); + bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET]; + uint32_t kOffset = k1 * kL0Loops; + kL0Size = K_L0_SPLIT_SIZE; + // 此处必须先初始化kL0Size, 再求kL0Loops, 否则由于循环会改变kL0Size大小, 导致kL0Loops错误 + kL0Loops = (kL1Size + kL0Size - 1) / kL0Size; + kL0SizeAlign = kL0Size; + for (uint32_t kL1 = kOffset; kL1 < kL0Loops + kOffset; kL1++) { // 128 循环搬pa + if (kL1 == kOffset + kL0Loops - 1) { + // 尾块 + kL0Size = kL1Size - (kL0Loops - 1) * kL0Size; + kL0SizeAlign = SMLAAlign(kL0Size, 16U); + } + + uint32_t copyFinishRowCnt = 0; + if (info.isOriOnly) { + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + uint64_t curS2Offset = static_cast(info.s2Idx) * constInfo.s2BaseSize + + info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + if (constInfo.hasOriSparseIndices) { + uint64_t mergeBase = static_cast(info.cmpLoop % MERGE_CACHE_GM_BUF_NUM) * + N_WORKSPACE_SIZE * constInfo.headDim; + uint32_t blockElementCnt = 32U / sizeof(KV_T); + subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE]; + for (uint32_t row = 0; row < kL0Size; ++row) { + uint64_t mergeOffset = + mergeBase + static_cast(kL1 * K_L0_SPLIT_SIZE + row) * constInfo.headDim + + nL1 * N_SPLIT_SIZE; + GlobalTensor mergeSrcGm = kvMergeGm_[mergeOffset]; + LocalTensor dstTensor = subvTensor[row * blockElementCnt]; + DataCopyGmNDToL1(dstTensor, mergeSrcGm, 1, kL0SizeAlign, nL1Size, + constInfo.headDim); + } + } else + while (copyFinishRowCnt < kL0Size) { + copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize; + if (copyFinishRowCnt + copyRowCnt > kL0Size) { + copyRowCnt = kL0Size - copyFinishRowCnt; + } + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + startPos.dIdx = + nL1 * N_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + PAShape shape; + shape.blockSize = constInfo.paOriBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.oriKvStride0; + shape.actHeadDim = nL1Size; + shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = kL0SizeAlign; + subvTensor = + bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE + copyFinishRowCnt * 16]; + DataCopyPA(subvTensor, oriKvGm, oriBlockTableGm, shape, startPos); + + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE]; + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = nL1Size; // 256 + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + uint32_t headStride = constInfo.headDim; + uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim; + uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? + static_cast(constInfo.kvSeqSize) * seqStride : + constInfo.oriKvStride0; + + uint64_t curS2 = info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE; + DataCopy(subvTensor, oriKvGm[offset], nd2nzPara); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + uint64_t curS2Offset = static_cast(info.s2Idx) * constInfo.s2BaseSize + + info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE], + oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + nL1 * N_SPLIT_SIZE], + nd2nzPara); + } + } else if (info.isOriCmpMix) { + uint32_t cumKl0LoopSize = kL1 * K_L0_SPLIT_SIZE; + uint32_t oriSizeCur = (info.actualSingleProcessSInnerOriSize < cumKl0LoopSize) ? + 0 : + (info.actualSingleProcessSInnerOriSize - cumKl0LoopSize); + oriSizeCur = (oriSizeCur > kL0Size) ? kL0Size : oriSizeCur; + uint32_t cmpSizeCur = kL0Size - oriSizeCur; + cmpSizeCur = (cmpSizeCur < info.actualSingleProcessSInnerCmpSize) ? + cmpSizeCur : + info.actualSingleProcessSInnerCmpSize; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + if (oriSizeCur > 0) { + copyFinishRowCnt = 0; + uint64_t curS2Offset = info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + while (copyFinishRowCnt < oriSizeCur) { + copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize; + if (copyFinishRowCnt + copyRowCnt > oriSizeCur) { + copyRowCnt = oriSizeCur - copyFinishRowCnt; + } + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + startPos.dIdx = + nL1 * N_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + PAShape shape; + shape.blockSize = constInfo.paOriBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.oriKvStride0; + shape.actHeadDim = nL1Size; + shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = kL0SizeAlign; + subvTensor = + bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE + copyFinishRowCnt * 16]; + DataCopyPA(subvTensor, oriKvGm, oriBlockTableGm, shape, startPos); + + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } + if (cmpSizeCur > 0) { + copyFinishRowCnt = 0; + uint32_t cmpMixsizeCur = + K_L0_SPLIT_SIZE - info.actualSingleProcessSInnerOriSize % K_L0_SPLIT_SIZE; + uint32_t oriOnlyLoopTimes = info.actualSingleProcessSInnerOriSize / K_L0_SPLIT_SIZE; + uint32_t cmpLoopTimes = kL1 - oriOnlyLoopTimes; + uint64_t curS2Offset = + (cmpLoopTimes == 0) ? + 0 : + (cmpMixsizeCur + static_cast(cmpLoopTimes - 1) * K_L0_SPLIT_SIZE); + while (copyFinishRowCnt < cmpSizeCur) { + copyRowCnt = constInfo.paCmpBlockSize - curS2Offset % constInfo.paCmpBlockSize; + if (copyFinishRowCnt + copyRowCnt > cmpSizeCur) { + copyRowCnt = cmpSizeCur - copyFinishRowCnt; + } + + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + // 256、32等待7buf命名更改 + startPos.dIdx = + nL1 * N_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + + PAShape shape; + shape.blockSize = constInfo.paCmpBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.cmpKvStride0; + shape.actHeadDim = nL1Size; + shape.maxblockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = kL0SizeAlign; + subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE + + oriSizeCur * 16 + copyFinishRowCnt * 16]; + DataCopyPA(subvTensor, cmpKvGm, cmpBlockTableGm, shape, startPos); + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = nL1Size; // 256 + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + uint32_t headStride = constInfo.headDim; + uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim; + if (oriSizeCur > 0) { + subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE]; + nd2nzPara.nValue = oriSizeCur; + uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? + static_cast(constInfo.kvSeqSize) * seqStride : + constInfo.oriKvStride0; + uint64_t curS2 = info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE; + DataCopy(subvTensor, oriKvGm[offset], nd2nzPara); + } + if (cmpSizeCur > 0) { + nd2nzPara.nValue = cmpSizeCur; + subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE + oriSizeCur * 16]; + uint64_t batchStride = (constInfo.cmpKvStride0 == 0) ? + static_cast(constInfo.cmpSeqSize) * seqStride : + constInfo.cmpKvStride0; + uint32_t cmpMixsizeCur = + K_L0_SPLIT_SIZE - info.actualSingleProcessSInnerOriSize % K_L0_SPLIT_SIZE; + uint32_t oriOnlyLoopTimes = info.actualSingleProcessSInnerOriSize / K_L0_SPLIT_SIZE; + uint32_t cmpLoopTimes = kL1 - oriOnlyLoopTimes; + uint64_t curS2 = + (cmpLoopTimes == 0) ? + 0 : + (cmpMixsizeCur + static_cast(cmpLoopTimes - 1) * K_L0_SPLIT_SIZE); + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE; + DataCopy(subvTensor, cmpKvGm[offset], nd2nzPara); + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + if (oriSizeCur > 0) { + nd2nzPara.nValue = oriSizeCur; + uint64_t curS2Offset = info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + DataCopy(bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE], + oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + nL1 * N_SPLIT_SIZE], + nd2nzPara); + } + if (cmpSizeCur > 0) { + nd2nzPara.nValue = cmpSizeCur; + uint32_t cmpMixsizeCur = + K_L0_SPLIT_SIZE - info.actualSingleProcessSInnerOriSize % K_L0_SPLIT_SIZE; + uint32_t oriOnlyLoopTimes = info.actualSingleProcessSInnerOriSize / K_L0_SPLIT_SIZE; + uint32_t cmpLoopTimes = kL1 - oriOnlyLoopTimes; + uint64_t curS2Offset = + (cmpLoopTimes == 0) ? + 0 : + (cmpMixsizeCur + static_cast(cmpLoopTimes - 1) * K_L0_SPLIT_SIZE); + DataCopy( + bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE + oriSizeCur * 16], + cmpKvGm[info.tensorCmpBOffset + curS2Offset * constInfo.headDim + nL1 * N_SPLIT_SIZE], + nd2nzPara); + } + } + } else { + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + uint64_t curS2Offset = static_cast(info.relativeS2Idx) * constInfo.s2BaseSize + + info.s2StartPoint + K_L0_SPLIT_SIZE * kL1; + while (copyFinishRowCnt < kL0Size) { + copyRowCnt = constInfo.paCmpBlockSize - curS2Offset % constInfo.paCmpBlockSize; + if (copyFinishRowCnt + copyRowCnt > kL0Size) { + copyRowCnt = kL0Size - copyFinishRowCnt; + } + + Position startPos; + startPos.bIdx = info.bIdx; + startPos.n2Idx = info.n2Idx; + startPos.s2Idx = curS2Offset; + // 256、32等待7buf命名更改 + startPos.dIdx = + nL1 * N_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分 + + PAShape shape; + shape.blockSize = constInfo.paCmpBlockSize; + shape.headNum = constInfo.kvHeadNum; + shape.headDim = constInfo.headDim; + shape.kvStride = constInfo.cmpKvStride0; + shape.actHeadDim = nL1Size; + shape.maxblockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch; + shape.copyRowNum = copyRowCnt; + shape.copyRowNumAlign = kL0SizeAlign; + subvTensor = + bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE + copyFinishRowCnt * 16]; + DataCopyPA(subvTensor, cmpKvGm, cmpBlockTableGm, shape, startPos); + // 更新循环变量 + copyFinishRowCnt += copyRowCnt; + curS2Offset += copyRowCnt; + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE]; + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = nL1Size; // 256 + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + + uint32_t headStride = constInfo.headDim; + uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim; + uint64_t batchStride = (constInfo.cmpKvStride0 == 0) ? + static_cast(constInfo.cmpSeqSize) * seqStride : + constInfo.cmpKvStride0; + + uint64_t curS2 = (uint64_t)info.relativeS2Idx * constInfo.s2BaseSize + info.s2StartPoint + + K_L0_SPLIT_SIZE * kL1; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE; + DataCopy(subvTensor, cmpKvGm[offset], nd2nzPara); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + uint64_t curS2Offset = static_cast(info.relativeS2Idx) * constInfo.s2BaseSize + + info.s2StartPoint + K_L0_SPLIT_SIZE * kL1; + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = kL0Size; // 行数 + nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim; + nd2nzPara.srcDValue = constInfo.headDim; + nd2nzPara.dstNzC0Stride = kL0SizeAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE], + cmpKvGm[info.tensorCmpBOffset + curS2Offset * constInfo.headDim + nL1 * N_SPLIT_SIZE], + nd2nzPara); + } + } + } + SetFlag(mte21KVIds[kb]); + WaitFlag(mte21KVIds[kb]); + mL1SizeAlign = M_SPLIT_SIZE; + mL1Size = M_SPLIT_SIZE; // m的实际大小 + for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) { + if (mL1 == (mL1Loops - 1)) { + // 尾块 + mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE; + mL1SizeAlign = SMLAAlign(mL1Size, 16U); + } + + uint32_t mIdx = mBaseIdx + mL1; + ka = GetQPL1RealIdx(mIdx, k1); + LocalTensor aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET]; + if (nL1 == 0) { + WaitFlag(mte21QPIds[ka]); + CopyInMm2AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, kL1Size, + 256 * k1); + SetFlag(mte21QPIds[ka]); + WaitFlag(mte21QPIds[ka]); + } + + LocalTensor cL0Tensor = + cL0TensorPingPong[(cL0BufIter % 2) * + (L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致 + uint32_t baseK = 128; + uint32_t baseN = 128; + kL0Size = 128; + kL0SizeAlign = kL0Size; + for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) { + if (kL0 + 1 == kL0Loops) { + kL0Size = kL1Size - (kL0Loops - 1) * kL0Size; + kL0SizeAlign = SMLAAlign(kL0Size, 16U); + } + WaitFlag(Mte1MmABEventId(abL0BufIter % 2)); + LocalTensor bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE / sizeof(KV_T))]; + LoadData3DParamsV2 loadData3DParamsForB; + loadData3DParamsForB.l1H = kL0SizeAlign / 16; // 源操作数height + loadData3DParamsForB.l1W = 16; // 源操作数weight=16,目的height=l1H*L1W + loadData3DParamsForB.padList[0] = 0; + loadData3DParamsForB.padList[1] = 0; + loadData3DParamsForB.padList[2] = 0; + loadData3DParamsForB.padList[3] = 255; // 尾部数据不影响滑窗的结果 + + loadData3DParamsForB.mExtension = kL0SizeAlign; // 在目的操作数height维度的传输长度 + loadData3DParamsForB.kExtension = nL1SizeAlign; // 在目的操作数width维度的传输长度 + loadData3DParamsForB.mStartPt = 0; // 卷积核在目的操作数width维度的起点 + loadData3DParamsForB.kStartPt = 0; // 卷积核在目的操作数height维度的起点 + loadData3DParamsForB.strideW = 1; + loadData3DParamsForB.strideH = 1; + loadData3DParamsForB.filterW = 1; + loadData3DParamsForB.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素 + loadData3DParamsForB.filterH = 1; + loadData3DParamsForB.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素 + loadData3DParamsForB.dilationFilterW = 1; // 卷积核width膨胀系数 + loadData3DParamsForB.dilationFilterH = 1; // 卷积核height膨胀系数 + loadData3DParamsForB.enTranspose = 1; // 是否启用转置功能 + loadData3DParamsForB.fMatrixCtrl = + 0; // 使用FMATRIX_LEFT还是使用FMATRIX_RIGHT,=0使用FMATRIX_LEFT,=1使用FMATRIX_RIGHT 1 + loadData3DParamsForB.channelSize = + nL1SizeAlign; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize + LoadData(bL0Tensor, bL1Tensor[kL0 * baseK * baseN], loadData3DParamsForB); + + LocalTensor aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE / sizeof(KV_T))]; + LoadData3DParamsV2 loadData3DParamsForA; + loadData3DParamsForA.l1H = mL1SizeAlign / 16; // 源操作数height + loadData3DParamsForA.l1W = 16; // 源操作数weight + loadData3DParamsForA.padList[0] = 0; + loadData3DParamsForA.padList[1] = 0; + loadData3DParamsForA.padList[2] = 0; + loadData3DParamsForA.padList[3] = 255; // 尾部数据不影响滑窗的结果 + + loadData3DParamsForA.mExtension = mL1SizeAlign; // 在目的操作数height维度的传输长度 + loadData3DParamsForA.kExtension = kL0SizeAlign; // 在目的操作数width维度的传输长度 + loadData3DParamsForA.mStartPt = 0; // 卷积核在目的操作数width维度的起点 + loadData3DParamsForA.kStartPt = 0; // 卷积核在目的操作数height维度的起点 + loadData3DParamsForA.strideW = 1; // 卷积核在源操作数width维度滑动的步长 + loadData3DParamsForA.strideH = 1; // 卷积核在源操作数height维度滑动的步长 + loadData3DParamsForA.filterW = 1; // 卷积核width + loadData3DParamsForA.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素 + loadData3DParamsForA.filterH = 1; // 卷积核height + loadData3DParamsForA.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素 + loadData3DParamsForA.dilationFilterW = 1; // 卷积核width膨胀系数 + loadData3DParamsForA.dilationFilterH = 1; // 卷积核height膨胀系数 + loadData3DParamsForA.enTranspose = 0; // 是否启用转置功能,对整个目标矩阵进行转置 + loadData3DParamsForA.fMatrixCtrl = 0; + loadData3DParamsForA.channelSize = + kL0SizeAlign; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize + LoadData(aL0Tensor, aL1Tensor[kL0 * baseK * mL1SizeAlign], + loadData3DParamsForA); + SetFlag(Mte1MmABEventId(abL0BufIter % 2)); + WaitFlag(Mte1MmABEventId(abL0BufIter % 2)); + + MmadParams mmadParams; + mmadParams.m = mL1SizeAlign; + mmadParams.n = nL1SizeAlign; + mmadParams.k = kL0Size; + mmadParams.cmatrixInitVal = (kL0 == 0 && k1 == 0); + mmadParams.cmatrixSource = false; + mmadParams.unitFlag = ((k1 == (kL1Loops - 1)) && (kL0 == (kL0Loops - 1))) ? 0b11 : 0b10; + + Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams); + if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) { + PipeBarrier(); + } + SetFlag(Mte1MmABEventId(abL0BufIter % 2)); + abL0BufIter++; + } + + if (nL1 == (nL1Loops - 1)) { // nL1最后一轮, 需要将B驻留在L1中, 用于下一轮的计算? + SetFlag(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完 + } + + if (k1 == (kL1Loops - 1)) { + // ND + FixpipeParamsV220 fixParams; + fixParams.nSize = nL1SizeAlign; + fixParams.mSize = mL1SizeAlign; + fixParams.srcStride = mL1SizeAlign; + fixParams.dstStride = nSize; // mm2ResGm两行之间的间隔 + fixParams.ndNum = 1; // 输出ND + fixParams.unitFlag = 0b11; + + uint64_t mm2Offset = (mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) * nSize + nL1 * N_SPLIT_SIZE; + Fixpipe(mm2ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.bmm2ResUbSize + mm2Offset], + cL0Tensor, fixParams); + } + + if (mL1Loops == 2) { + cL0BufIter++; + } + } + SetFlag(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完 + } + // cL0BufIter已经不在使用 + if (mL1Loops == 1) { + cL0BufIter++; + } + } + qpL1BufIter += mL1Loops; +} +} // namespace SMLAKernel +#endif diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_vector.h b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_vector.h new file mode 100644 index 000000000000..301cc40fe801 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_vector.h @@ -0,0 +1,1269 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_swa_block_vector.h + * \brief + */ +#ifndef SPARSE_FLASH_MLA_SWA_BLOCK_VECTOR_H +#define SPARSE_FLASH_MLA_SWA_BLOCK_VECTOR_H + +#include "kernel_operator.h" +#include "kernel_operator_list_tensor_intf.h" +#include "kernel_tiling/kernel_tiling.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" +#include "sparse_flash_mla_common_arch22.h" + +namespace SMLAKernel { +using AscendC::CrossCoreSetFlag; +using AscendC::CrossCoreWaitFlag; + +template +class SWAVectorBlock { +public: + // 中间计算数据类型为float,高精度模式 + using T = float; + using KV_T = typename SMLAT::kvType; + using OUT_T = typename SMLAT::outputType; + using UPDATE_T = T; + using SINKS_T = T; + using MM1_OUT_T = float; + using MM2_OUT_T = float; + + __aicore__ inline SWAVectorBlock(){}; + __aicore__ inline void ProcessVec0L(const RunInfo &runInfo); + __aicore__ inline void ProcessVec1L(const RunInfo &info); + __aicore__ inline void ProcessVec2L(const RunInfo &info); + __aicore__ inline void InitBuffers(TPipe *pipe); + __aicore__ inline void InitParams(const struct ConstInfo &constInfo, + const SparseFlashMlaTilingData *__restrict tilingData); + __aicore__ inline void InitVec0GlobalTensor(GlobalTensor kvMergeGm, GlobalTensor oriKvGm, + GlobalTensor oriBlockTableGm, + GlobalTensor oriSparseIndicesGm); + __aicore__ inline void InitVec1GlobalTensor(GlobalTensor mm1ResGm, GlobalTensor vec1ResGm, + GlobalTensor actualSeqLengthsQGm, + GlobalTensor actualSeqLengthsKVGm, GlobalTensor sinksGm, + GlobalTensor softmaxLseGm, GlobalTensor oriSparseIndicesGm, + GlobalTensor oriTopkLengthGm); + __aicore__ inline void InitVec2GlobalTensor(GlobalTensor accumOutGm, GlobalTensor vec2ResGm, + GlobalTensor mm2ResGm, GlobalTensor attentionOutGm); + __aicore__ inline void AllocEventID(); + __aicore__ inline void FreeEventID(); + __aicore__ inline void CopySinksIn(); + __aicore__ inline void SliceAndContactSinksValue(uint32_t nIdx, uint32_t dealRowCount); + __aicore__ inline void InitSoftmaxDefaultBuffer(); + // ================================Base Vector========================================== + __aicore__ inline void RowDivs(LocalTensor dstUb, LocalTensor src0Ub, LocalTensor src1Ub, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + __aicore__ inline void RowMuls(LocalTensor dstUb, LocalTensor src0Ub, LocalTensor src1Ub, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + // ================================Vector1========================================== + __aicore__ inline void ProcessVec1SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo); + __aicore__ inline void DealBmm1ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t startRow, + uint32_t dealRowCount, uint32_t columnCount, uint32_t loopId); + __aicore__ inline void SoftmaxFlashV2Compute(const RunInfo &info, const MSplitInfo &mSplitInfo, + LocalTensor &mmResUb, LocalTensor &softmaxTmpUb, + uint32_t startRow, uint32_t dealRowCount, uint32_t columnCount, + uint32_t actualColumnCount); + __aicore__ inline void ElewiseCompute(const RunInfo &info, const MSplitInfo &mSplitInfo, + const LocalTensor &mmResUb, uint32_t dealRowCount, uint32_t startRow, + uint32_t columnCount); + __aicore__ inline void ProcessLse(const RunInfo &info, const MSplitInfo &mSplitInfo); + // ================================Vecotr2========================================== + __aicore__ inline void ProcessVec2SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo); + __aicore__ inline void DealBmm2ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t startRow, + uint32_t dealRowCount, uint32_t columnCount, + uint32_t actualColumnCount); + __aicore__ inline void ProcessVec2Inner(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t mStartRow, + uint32_t mDealSize); + __aicore__ inline void Bmm2DataCopyOutTrans(const RunInfo &info, LocalTensor &attenOutUb, uint32_t wsMStart, + uint32_t dealRowCount, uint32_t columnCount, + uint32_t actualColumnCount); + __aicore__ inline void Bmm2ResCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, uint32_t wsMStart, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + __aicore__ inline void Bmm2CastAndCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, uint32_t wsMStart, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + __aicore__ inline void Bmm2FDDataCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, uint32_t wsMStart, + uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount); + __aicore__ inline uint64_t CalcAccumOffset(uint32_t bN2Idx, uint32_t gS1Idx); + __aicore__ inline void SetInfInBlk(const LocalTensor &mmResUb, uint32_t dealRowCount, uint32_t columnCount, + int64_t startId, int64_t endId); + + // BLOCK和REPEAT的字节数 + static constexpr uint64_t BYTE_BLOCK = 32UL; + static constexpr uint32_t REPEAT_BLOCK_BYTE = 256U; + // BLOCK和REPEAT的FP32元素数 + static constexpr uint32_t FP32_BLOCK_ELEMENT_NUM = BYTE_BLOCK / sizeof(float); + static constexpr uint32_t FP32_REPEAT_ELEMENT_NUM = REPEAT_BLOCK_BYTE / sizeof(float); + // repeat stride不能超过256 + static constexpr uint32_t REPEATE_STRIDE_UP_BOUND = 256; + static constexpr uint64_t MERGE_CACHE_GM_BUF_NUM = 3; + static constexpr uint64_t MERGE_WORKSPACE_ROW_NUM = 512; + +private: + __aicore__ inline int64_t GetOriSparseKeyGmOffset(int32_t logicalIdx, const RunInfo &runInfo); + __aicore__ inline void CopyInSingleOriSparseRow(int64_t &mte2Size, int64_t mte3Size, int64_t mergeMte3Idx, + int32_t logicalIdx, const RunInfo &runInfo); + __aicore__ inline void CopyInOriSparseKv(int64_t &mte2Size, int64_t mte3Size, int64_t mergeMte3Idx, + int32_t logicalIdx0, int32_t logicalIdx1, const RunInfo &runInfo); + __aicore__ inline void CopyOutOriSparseMerge(int64_t mte2Size, int64_t mte3Size, int64_t s2GmStartOffset, + int64_t mergeMte3Idx, const RunInfo &runInfo); + __aicore__ inline int32_t AlignOriSparseIndexLoadCount(int32_t loadCount); + __aicore__ inline void LoadOriSparseIndicesGmToUb(uint64_t qTokenOffset, uint32_t n2Idx, uint32_t gmColStart, + int32_t loadCount, const LocalTensor &dstUb); + static constexpr bool PAGE_ATTENTION = SMLAT::pageAttention; + static constexpr bool FLASH_DECODE = SMLAT::flashDecode; + static constexpr SMLA_LAYOUT LAYOUT_T = SMLAT::layout; + static constexpr SMLA_LAYOUT KV_LAYOUT_T = SMLAT::kvLayout; + + static constexpr uint64_t SYNC_INPUT_BUF1_FLAG = 2; + static constexpr uint64_t SYNC_INPUT_BUF1_PONG_FLAG = 3; + static constexpr uint64_t SYNC_INPUT_BUF2_FLAG = 4; + static constexpr uint64_t SYNC_INPUT_BUF2_PONG_FLAG = 5; + static constexpr uint64_t SYNC_OUTPUT_BUF1_FLAG = 4; + static constexpr uint64_t SYNC_OUTPUT_BUF2_FLAG = 5; + static constexpr uint64_t SYNC_SINKS_BUF_FLAG = 6; + static constexpr uint32_t INPUT1_BUFFER_OFFSET = ConstInfo::BUFFER_SIZE_BYTE_32K; + static constexpr uint32_t SOFTMAX_TMP_BUFFER_OFFSET = ConstInfo::BUFFER_SIZE_BYTE_1K; + static constexpr uint32_t BASE_BLOCK_MAX_ELEMENT_NUM = ConstInfo::BUFFER_SIZE_BYTE_32K / sizeof(T); // 32768/4=8096 + static constexpr uint32_t BLOCK_ELEMENT_NUM = BYTE_BLOCK / sizeof(T); // 32/4=8 + static constexpr uint32_t MAX_N1_SIZE = 128U; + static constexpr T SOFTMAX_MIN_NUM = -2e38; + static constexpr SINKS_T R0 = 1.0f; + + const SparseFlashMlaTilingData *__restrict tilingData; + + uint32_t pingpongFlag = 0U; + ConstInfo constInfo = {}; + + GlobalTensor mm1ResGm; + GlobalTensor vec1ResGm; + GlobalTensor softmaxMaxGm; + GlobalTensor softmaxSumGm; + GlobalTensor sinksGm; + + GlobalTensor actualSeqLengthsQGm; + GlobalTensor actualSeqLengthsKVGm; + GlobalTensor oriSparseIndicesGm; + GlobalTensor oriTopkLengthGm; + GlobalTensor vec2ResGm; + GlobalTensor mm2ResGm; + GlobalTensor accumOutGm; + GlobalTensor attentionOutGm; + GlobalTensor blkTableGm_; + GlobalTensor keyGm_; + GlobalTensor kvValidSizeGm_; + GlobalTensor oriKvGm_; + GlobalTensor cmpKvGm_; + GlobalTensor oriBlockTableGm_; + GlobalTensor cmpBlockTableGm_; + GlobalTensor kvMergeGm_; + GlobalTensor softmaxLseGm; + + // ================================Local Buffer区==================================== + TBuf<> inputBuff1; // 32K + TBuf<> inputBuff2; // 8K + TBuf<> outputBuff1; // 32K + TBuf<> outputBuff2; // 4K + + TBuf<> tmpBuff1; // 32K + TBuf<> v0ValidSizeBuff; // 8K + TBuf<> sinksBuff; // 1K + TBuf<> sinksBrcbBuff; // 12K + + TBuf<> softmaxMaxBuff; // PRE_LOAD_NUM * 2K + TBuf<> softmaxExpBuff; // PRE_LOAD_NUM * 2K + TBuf<> softmaxSumBuff; // PRE_LOAD_NUM * 2K + TBuf<> softmaxMaxDefaultBuff; // 2K + TBuf<> softmaxSumDefaultBuff; // 2K + + LocalTensor softmaxMaxDefaultUb; + LocalTensor softmaxSumDefaultUb; + + LocalTensor softmaxMaxUb; + LocalTensor softmaxSumUb; + LocalTensor softmaxExpUb; + LocalTensor sinksUb; + LocalTensor sinksBrcbUb; + LocalTensor kvMergUb_; + uint32_t mergeMte3Idx = 0; +}; + +// ============================== init ============================================== +template +__aicore__ inline void SWAVectorBlock::InitBuffers(TPipe *pipe) +{ + pipe->InitBuffer(inputBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K * 2); // 2:pingpong + pipe->InitBuffer(inputBuff2, ConstInfo::BUFFER_SIZE_BYTE_8K * 2); // 2:pingpong + pipe->InitBuffer(outputBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K); + pipe->InitBuffer(outputBuff2, ConstInfo::BUFFER_SIZE_BYTE_4K); + + pipe->InitBuffer(tmpBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K); + pipe->InitBuffer(v0ValidSizeBuff, ConstInfo::BUFFER_SIZE_BYTE_8K); + // M_MAX = 512/2vector = 256, 256 * sizeof(T) * N_Buffer + pipe->InitBuffer(softmaxMaxBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum); + pipe->InitBuffer(softmaxExpBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum); + pipe->InitBuffer(softmaxSumBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum); + + pipe->InitBuffer(softmaxMaxDefaultBuff, ConstInfo::BUFFER_SIZE_BYTE_1K); + pipe->InitBuffer(softmaxSumDefaultBuff, ConstInfo::BUFFER_SIZE_BYTE_1K); + + pipe->InitBuffer(sinksBuff, MAX_N1_SIZE * sizeof(SINKS_T)); + // 分配256+N1大小内存,其中256是m轴VEC最大切块 + pipe->InitBuffer(sinksBrcbBuff, MAX_N1_SIZE * sizeof(SINKS_T) * BLOCK_ELEMENT_NUM * 3U); + + softmaxMaxUb = softmaxMaxBuff.Get(); + softmaxSumUb = softmaxSumBuff.Get(); + softmaxExpUb = softmaxExpBuff.Get(); + + softmaxMaxDefaultUb = softmaxMaxDefaultBuff.Get(); + softmaxSumDefaultUb = softmaxSumDefaultBuff.Get(); + + sinksUb = sinksBuff.Get(); + sinksBrcbUb = sinksBrcbBuff.Get(); + kvMergUb_ = inputBuff1.Get(); +} + +template +__aicore__ inline void SWAVectorBlock::InitParams(const struct ConstInfo &constInfo, + const SparseFlashMlaTilingData *__restrict tilingData) +{ + this->constInfo = constInfo; + this->tilingData = tilingData; +} + +template +__aicore__ inline void SWAVectorBlock::InitVec0GlobalTensor(GlobalTensor kvMergeGm, + GlobalTensor oriKvGm, + GlobalTensor oriBlockTableGm, + GlobalTensor oriSparseIndicesGm) +{ + this->kvMergeGm_ = kvMergeGm; + this->oriKvGm_ = oriKvGm; + this->oriBlockTableGm_ = oriBlockTableGm; + this->oriSparseIndicesGm = oriSparseIndicesGm; +} + +template +__aicore__ inline int64_t SWAVectorBlock::GetOriSparseKeyGmOffset(int32_t logicalIdx, const RunInfo &runInfo) +{ + if (logicalIdx < 0) { + return -1; + } + int32_t oriLenLimit = actualSeqLengthsKVGm.GetValue(runInfo.bIdx); + if (logicalIdx >= oriLenLimit) { + return -1; + } + uint64_t blockTableIdx = static_cast(logicalIdx) / constInfo.paOriBlockSize; + uint64_t inBlockIdx = static_cast(logicalIdx) % constInfo.paOriBlockSize; + uint64_t idInBlockTable = + oriBlockTableGm_.GetValue(runInfo.bIdx * constInfo.oriMaxBlockNumPerBatch + blockTableIdx); + return static_cast(idInBlockTable * constInfo.oriKvStride0 + + static_cast(runInfo.n2IdxReal) * constInfo.headDim * + constInfo.paOriBlockSize + + inBlockIdx * constInfo.headDim); +} + +template +__aicore__ inline void SWAVectorBlock::CopyInSingleOriSparseRow(int64_t &mte2Size, int64_t mte3Size, + int64_t mergeMte3Idx, int32_t logicalIdx, + const RunInfo &runInfo) +{ + int64_t keyOffset = GetOriSparseKeyGmOffset(logicalIdx, runInfo); + if (keyOffset >= 0) { + int64_t ubRowOffset = + mergeMte3Idx % 2 * INPUT1_BUFFER_OFFSET / sizeof(KV_T) + (mte2Size - mte3Size) * constInfo.headDim; + DataCopyExtParams copyInParams; + copyInParams.blockCount = 1; + copyInParams.blockLen = constInfo.headDim * sizeof(KV_T); + copyInParams.srcStride = 0; + copyInParams.dstStride = 0; + DataCopyPadExtParams padInParams{false, 0, 0, 0}; + DataCopyPad(kvMergUb_[ubRowOffset], oriKvGm_[keyOffset], copyInParams, padInParams); + } + mte2Size += constInfo.sparseBlockSize; +} + +template +__aicore__ inline void SWAVectorBlock::CopyInOriSparseKv(int64_t &mte2Size, int64_t mte3Size, + int64_t mergeMte3Idx, int32_t logicalIdx0, + int32_t logicalIdx1, const RunInfo &runInfo) +{ + int64_t keyOffset1 = GetOriSparseKeyGmOffset(logicalIdx0, runInfo); + int64_t keyOffset2 = GetOriSparseKeyGmOffset(logicalIdx1, runInfo); + if (logicalIdx1 < 0) { + CopyInSingleOriSparseRow(mte2Size, mte3Size, mergeMte3Idx, logicalIdx0, runInfo); + return; + } + if (unlikely(keyOffset1 < 0 && keyOffset2 < 0)) { + // invalid 行仅占 merge 行位,不搬 KV;V1 按索引刷 -inf,值无关 + mte2Size += 2 * constInfo.sparseBlockSize; + return; + } + + int64_t keySrcStride = ((keyOffset1 > keyOffset2 ? (keyOffset1 - keyOffset2) : (keyOffset2 - keyOffset1)) - + constInfo.sparseBlockSize * constInfo.headDim) * + static_cast(sizeof(KV_T)); + if (unlikely(keySrcStride >= INT32_MAX || keySrcStride < 0 || keyOffset1 < 0 || keyOffset2 < 0)) { + CopyInSingleOriSparseRow(mte2Size, mte3Size, mergeMte3Idx, logicalIdx0, runInfo); + CopyInSingleOriSparseRow(mte2Size, mte3Size, mergeMte3Idx, logicalIdx1, runInfo); + } else { + DataCopyExtParams intriParams; + intriParams.blockLen = constInfo.sparseBlockSize * constInfo.headDim * sizeof(KV_T); + intriParams.blockCount = (keyOffset1 >= 0) + (keyOffset2 >= 0); + intriParams.dstStride = 0; + intriParams.srcStride = keySrcStride; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + int64_t startGmOffset = keyOffset1 > -1 ? keyOffset1 : keyOffset2; + if (keyOffset2 > -1 && keyOffset2 < keyOffset1) { + startGmOffset = keyOffset2; + } + DataCopyPad(kvMergUb_[mergeMte3Idx % 2 * INPUT1_BUFFER_OFFSET / sizeof(KV_T) + + (mte2Size - mte3Size) * constInfo.headDim], + oriKvGm_[startGmOffset], intriParams, padParams); + mte2Size += 2 * constInfo.sparseBlockSize; + } +} + +template +__aicore__ inline void SWAVectorBlock::CopyOutOriSparseMerge(int64_t mte2Size, int64_t mte3Size, + int64_t s2GmStartOffset, int64_t mergeMte3Idx, + const RunInfo &runInfo) +{ + if (mte2Size <= mte3Size) { + return; + } + SetFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + WaitFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = static_cast(mte2Size - mte3Size); + dataCopyParams.blockLen = constInfo.headDim * sizeof(KV_T); + dataCopyParams.srcStride = 0; + dataCopyParams.dstStride = 0; + DataCopyPad(kvMergeGm_[runInfo.cmpLoop % MERGE_CACHE_GM_BUF_NUM * MERGE_WORKSPACE_ROW_NUM * constInfo.headDim + + (s2GmStartOffset + mte3Size) * constInfo.headDim], + kvMergUb_[mergeMte3Idx % 2 * INPUT1_BUFFER_OFFSET / sizeof(KV_T)], dataCopyParams); +} + +template +__aicore__ inline int32_t SWAVectorBlock::AlignOriSparseIndexLoadCount(int32_t loadCount) +{ + if (loadCount <= 0) { + return 0; + } + // MTE blockLen must be 32B-aligned; int32 slot count aligns to 8 elements. + return ((loadCount + 7) / 8) * 8; +} + +template +__aicore__ inline void SWAVectorBlock::LoadOriSparseIndicesGmToUb(uint64_t qTokenOffset, uint32_t n2Idx, + uint32_t gmColStart, int32_t loadCount, + const LocalTensor &dstUb) +{ + if (loadCount <= 0 || constInfo.oriSparseIndexWidth == 0) { + return; + } + int32_t remain = static_cast(constInfo.oriSparseIndexWidth) - static_cast(gmColStart); + int32_t validCount = (remain < loadCount) ? remain : loadCount; + if (validCount <= 0) { + return; + } + int32_t alignedCount = AlignOriSparseIndexLoadCount(loadCount); + uint64_t gmOffset = (qTokenOffset * constInfo.kvHeadNum + n2Idx) * constInfo.oriSparseIndexWidth + gmColStart; + DataCopyExtParams copyParams; + copyParams.blockCount = 1; + copyParams.blockLen = static_cast(validCount) * sizeof(int32_t); + copyParams.srcStride = 0; + copyParams.dstStride = 0; + DataCopyPadExtParams padParams; + padParams.isPad = true; + padParams.leftPadding = 0; + padParams.rightPadding = static_cast(alignedCount - validCount); + padParams.paddingValue = -1; + DataCopyPad(dstUb, oriSparseIndicesGm[gmOffset], copyParams, padParams); + event_t mte2ToS = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_S)); + SetFlag(mte2ToS); + WaitFlag(mte2ToS); +} + +template +__aicore__ inline void SWAVectorBlock::ProcessVec0L(const RunInfo &runInfo) +{ + // V0 shares inputBuff1 with V1. Complete all prior pipe activity before reuse. + PipeBarrier(); + int64_t s2ProcessSize = runInfo.v0S2DealSize; + if (s2ProcessSize <= 0) { + return; + } + uint32_t sparseColStart = runInfo.s2Idx * constInfo.s2BaseSize + static_cast(runInfo.v0S2Start); + int64_t s2Pair = CeilDiv(s2ProcessSize, 2 * constInfo.sparseBlockSize); + int64_t s2SplitPoint = SMLAAlign(s2Pair, 2) * constInfo.sparseBlockSize; + int64_t s2GmStartOffset = GetSubBlockIdx() == 0 ? 0 : s2SplitPoint; + int64_t s2GmLimit = GetSubBlockIdx() == 0 ? s2SplitPoint : s2ProcessSize; + if (s2GmLimit > s2ProcessSize) { + s2GmLimit = s2ProcessSize; + } + if (s2GmStartOffset >= s2GmLimit) { + return; + } + + LocalTensor sliceUb = v0ValidSizeBuff.Get(); + LoadOriSparseIndicesGmToUb(runInfo.qTokenOffset, runInfo.n2IdxReal, + sparseColStart + static_cast(s2GmStartOffset), + static_cast(s2GmLimit - s2GmStartOffset), sliceUb); + + int64_t s2LocalSize = s2GmLimit - s2GmStartOffset; + int64_t mte2Size = 0; + int64_t mte3Size = 0; + bool needWaitMte3ToMte2 = true; + for (int64_t s2GmOffsetArray = 0; s2GmOffsetArray < s2LocalSize; s2GmOffsetArray += 2 * constInfo.sparseBlockSize) { + if (needWaitMte3ToMte2) { + WaitFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + needWaitMte3ToMte2 = false; + } + int32_t logicalIdx0 = sliceUb.GetValue(static_cast(s2GmOffsetArray)); + int32_t logicalIdx1 = -1; + if (s2GmOffsetArray + constInfo.sparseBlockSize < s2LocalSize) { + logicalIdx1 = sliceUb.GetValue(static_cast(s2GmOffsetArray + constInfo.sparseBlockSize)); + } + CopyInOriSparseKv(mte2Size, mte3Size, mergeMte3Idx, logicalIdx0, logicalIdx1, runInfo); + if ((mte2Size - mte3Size + 2 * constInfo.sparseBlockSize > 16) || + s2GmOffsetArray + 2 * constInfo.sparseBlockSize >= s2LocalSize) { + CopyOutOriSparseMerge(mte2Size, mte3Size, s2GmStartOffset, mergeMte3Idx, runInfo); + mte3Size = mte2Size; + SetFlag(mergeMte3Idx % 2 + SYNC_INPUT_BUF2_FLAG); + mergeMte3Idx++; + needWaitMte3ToMte2 = true; + } + } + // V1 may overwrite inputBuff1 only after V0's final MTE3 copy is complete. + PipeBarrier(); +} + +template +__aicore__ inline void SWAVectorBlock::InitVec1GlobalTensor( + GlobalTensor mm1ResGm, GlobalTensor vec1ResGm, GlobalTensor actualSeqLengthsQGm, + GlobalTensor actualSeqLengthsKVGm, GlobalTensor sinksGm, GlobalTensor softmaxLseGm, + GlobalTensor oriSparseIndicesGm, GlobalTensor oriTopkLengthGm) +{ + this->mm1ResGm = mm1ResGm; + this->vec1ResGm = vec1ResGm; + this->actualSeqLengthsQGm = actualSeqLengthsQGm; + this->actualSeqLengthsKVGm = actualSeqLengthsKVGm; + this->sinksGm = sinksGm; + this->softmaxLseGm = softmaxLseGm; + this->oriSparseIndicesGm = oriSparseIndicesGm; + this->oriTopkLengthGm = oriTopkLengthGm; +} + +template +__aicore__ inline void SWAVectorBlock::InitVec2GlobalTensor(GlobalTensor accumOutGm, + GlobalTensor vec2ResGm, + GlobalTensor mm2ResGm, + GlobalTensor attentionOutGm) +{ + this->accumOutGm = accumOutGm; + this->vec2ResGm = vec2ResGm; + this->mm2ResGm = mm2ResGm; + this->attentionOutGm = attentionOutGm; +} + +template +__aicore__ inline void SWAVectorBlock::AllocEventID() +{ + SetFlag(SYNC_INPUT_BUF1_FLAG); + SetFlag(SYNC_INPUT_BUF1_PONG_FLAG); + SetFlag(SYNC_INPUT_BUF2_FLAG); + SetFlag(SYNC_INPUT_BUF2_PONG_FLAG); + if (constInfo.hasOriSparseIndices) { + SetFlag(SYNC_INPUT_BUF2_FLAG); + SetFlag(SYNC_INPUT_BUF2_PONG_FLAG); + } + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + SetFlag(SYNC_OUTPUT_BUF2_FLAG); +} + +template +__aicore__ inline void SWAVectorBlock::FreeEventID() +{ + WaitFlag(SYNC_INPUT_BUF1_FLAG); + WaitFlag(SYNC_INPUT_BUF1_PONG_FLAG); + WaitFlag(SYNC_INPUT_BUF2_FLAG); + WaitFlag(SYNC_INPUT_BUF2_PONG_FLAG); + if (constInfo.hasOriSparseIndices) { + WaitFlag(SYNC_INPUT_BUF2_FLAG); + WaitFlag(SYNC_INPUT_BUF2_PONG_FLAG); + } + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF2_FLAG); +} + +template +__aicore__ inline void SWAVectorBlock::CopySinksIn() +{ + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1U; + dataCopyParams.blockLen = constInfo.qHeadNum * sizeof(T); + dataCopyParams.srcStride = 0U; + dataCopyParams.dstStride = 0U; + DataCopyPadExtParams padParams; + DataCopyPad(sinksUb, sinksGm, dataCopyParams, padParams); + SetFlag(SYNC_SINKS_BUF_FLAG); + WaitFlag(SYNC_SINKS_BUF_FLAG); + uint32_t repeatTimes = (constInfo.qHeadNum + BLOCK_ELEMENT_NUM - 1U) / BLOCK_ELEMENT_NUM; // 每次处理 8 datablocks + Brcb(sinksBrcbUb, sinksUb, repeatTimes, {1, BLOCK_ELEMENT_NUM}); + PipeBarrier(); + + DataCopyParams repeatParams; + repeatParams.blockCount = 1; // 搬到有一个块超过单个vec核减分核M轴大小即可,核间切分每个vec256 + repeatParams.blockLen = constInfo.qHeadNum; + repeatParams.srcStride = 0U; + repeatParams.dstStride = 0U; + for (uint32_t i = 1U; i <= 256U / constInfo.qHeadNum; i++) { + DataCopy(sinksBrcbUb[constInfo.qHeadNum * BLOCK_ELEMENT_NUM * i], sinksBrcbUb, repeatParams); + } + PipeBarrier(); +} + +template +__aicore__ inline void SWAVectorBlock::SliceAndContactSinksValue(uint32_t nIdx, uint32_t dealRowCount) +{ + // 由于WholeReduceMax接口中repeatTimes支持范围(0,255),因此需要分多次调用WholeReduceMax,这里就使用每次repeatTime=128 + uint32_t repeatTimesOnce = 128; + uint32_t loopTimes = (dealRowCount + repeatTimesOnce - 1) / repeatTimesOnce; + uint32_t repeatTimes = repeatTimesOnce; + + for (uint32_t loop = 0; loop < loopTimes; ++loop) { + if (loop == loopTimes - 1) { + repeatTimes = dealRowCount - loop * repeatTimesOnce; + } + WholeReduceMax(softmaxMaxDefaultUb[loop * repeatTimesOnce], + sinksBrcbUb[(nIdx + loop * repeatTimesOnce) * BLOCK_ELEMENT_NUM], + BLOCK_ELEMENT_NUM * BLOCK_ELEMENT_NUM, repeatTimes, 1, 0, 1, ReduceOrder::ORDER_ONLY_VALUE); + PipeBarrier(); + } +} + +template +__aicore__ inline void SWAVectorBlock::InitSoftmaxDefaultBuffer() +{ + CopySinksIn(); + Duplicate(softmaxMaxDefaultUb, SOFTMAX_MIN_NUM, SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T)); + Duplicate(softmaxSumDefaultUb, R0, SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T)); +} + +template +__aicore__ inline void SWAVectorBlock::ElewiseCompute(const RunInfo &info, const MSplitInfo &mSplitInfo, + const LocalTensor &mmResUb, uint32_t startRow, + uint32_t dealRowCount, uint32_t columnCount) +{ + Muls(mmResUb, mmResUb, static_cast(tilingData->baseParams.softmaxScale), dealRowCount * columnCount); + uint32_t gs1StartIdx = info.gS1Idx + mSplitInfo.nBufferStartM + mSplitInfo.vecStartM + startRow; + uint32_t gs1EndIdx = gs1StartIdx + dealRowCount; + uint32_t s1StartIdx = gs1StartIdx / constInfo.gSize; + uint32_t s1EndIdx = gs1EndIdx / constInfo.gSize; + uint32_t gStartIdx = gs1StartIdx % constInfo.gSize; + uint32_t gEndIdx = gs1EndIdx % constInfo.gSize; + uint32_t dealTempSize = 0; + uint32_t ubOffset = 0; + if (info.isOriOnly) { + if (constInfo.hasOriSparseIndices) { + int32_t oriLenLimit = actualSeqLengthsKVGm.GetValue(info.bIdx); + uint32_t validCols = Min(info.actualSingleProcessSInnerSize, columnCount); + uint32_t sparseColStart = info.s2Idx * constInfo.s2BaseSize; + for (uint32_t i = s1StartIdx; i <= s1EndIdx; i++) { + dealTempSize = constInfo.gSize - gStartIdx; + if (i == s1EndIdx) { + dealTempSize = gEndIdx - gStartIdx; + } + if (dealTempSize == 0) { + continue; + } + uint64_t qTokenOffsetForS1 = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + qTokenOffsetForS1 = + static_cast(actualSeqLengthsQGm.GetValue(info.bIdx)) + static_cast(i); + } else { + qTokenOffsetForS1 = + static_cast(info.bIdx) * constInfo.qSeqSize + static_cast(i); + } + uint64_t topkLenOffset = qTokenOffsetForS1 * constInfo.kvHeadNum + info.n2IdxReal; + int32_t rowLen = oriTopkLengthGm.GetValue(topkLenOffset); + if (rowLen <= 0) { + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, 0, static_cast(columnCount - 1)); + ubOffset += dealTempSize * columnCount; + gStartIdx = 0; + continue; + } + uint32_t colEnd = 0; + uint32_t tailStart = 0; + if (static_cast(rowLen) > sparseColStart) { + colEnd = Min(validCols, static_cast(rowLen) - sparseColStart); + tailStart = static_cast(rowLen) - sparseColStart; + } + if (colEnd > 0) { + LocalTensor rowIndexUb = tmpBuff1.Get(); + LoadOriSparseIndicesGmToUb(qTokenOffsetForS1, info.n2IdxReal, sparseColStart, + static_cast(colEnd), rowIndexUb); + for (uint32_t k = 0; k < colEnd; ++k) { + int32_t logicalIdx = rowIndexUb.GetValue(k); + if (logicalIdx < 0 || logicalIdx >= oriLenLimit) { + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, static_cast(k), + static_cast(k)); + } + } + } + if (tailStart < validCols) { + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, static_cast(tailStart), + static_cast(validCols - 1)); + } + if (info.actualSingleProcessSInnerSize < columnCount) { + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, + static_cast(info.actualSingleProcessSInnerSize), + static_cast(columnCount - 1)); + } + ubOffset += dealTempSize * columnCount; + gStartIdx = 0; + } + } else { + int32_t right = info.oriDealSize + s1StartIdx - info.gS1Idx / constInfo.gSize; + int32_t left = Max(0, static_cast(right) - constInfo.oriWinLeft + constInfo.oriWinRight); + for (uint32_t i = s1StartIdx; i <= s1EndIdx; i++) { + dealTempSize = constInfo.gSize - gStartIdx; + if (i == s1EndIdx) { + dealTempSize = gEndIdx - gStartIdx; + } + if (dealTempSize == 0) { + continue; + } + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, 0, left - 1); + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, right + 1, columnCount - 1); + right = Min(right + 1, columnCount - 1); + left = Max(0, static_cast(right) - constInfo.oriWinLeft + constInfo.oriWinRight); + ubOffset += dealTempSize * columnCount; + gStartIdx = 0; + } + } + } else if (info.isOriCmpMix) { + int32_t oriRight = info.oriDealSize + s1StartIdx - info.gS1Idx / constInfo.gSize; + int32_t oriLeft = Max(0, static_cast(oriRight) - constInfo.oriWinLeft + constInfo.oriWinRight); + int32_t noMaskCmpSize = info.cmpMaskRight + s1StartIdx + 1; + int32_t actNoMaskCmpSize = 0; + int32_t cmpRight = 0; + for (uint32_t i = s1StartIdx; i <= s1EndIdx; i++) { + dealTempSize = constInfo.gSize - gStartIdx; + if (i == s1EndIdx) { + dealTempSize = gEndIdx - gStartIdx; + } + if (dealTempSize == 0) { + continue; + } + actNoMaskCmpSize = noMaskCmpSize / constInfo.cmpRatio; + if (actNoMaskCmpSize <= 0) { + cmpRight = 0; + } else { + cmpRight = Min(actNoMaskCmpSize, constInfo.s2BaseSize - info.actualSingleProcessSInnerOriSize); + } + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, 0, oriLeft - 1); + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, oriRight + 1, + info.actualSingleProcessSInnerOriSize - 1); + oriRight = Min(oriRight + 1, info.actualSingleProcessSInnerOriSize); + oriLeft = Max(0, static_cast(oriRight) - constInfo.oriWinLeft + constInfo.oriWinRight); + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, cmpRight + info.actualSingleProcessSInnerOriSize, + columnCount - 1); + noMaskCmpSize += 1; + ubOffset += dealTempSize * columnCount; + gStartIdx = 0; + } + } else { + int32_t noMaskCmpSize = info.cmpMaskRight + s1StartIdx + 1; + int32_t actNoMaskCmpSize = 0; + int32_t right = 0; + // relativeS2Idx is the cmp tile index relative to the current batch. + // s2Idx may include the global task index. + int64_t cmpTileStart = + static_cast(info.relativeS2Idx) * constInfo.s2BaseSize + static_cast(info.s2StartPoint); + for (uint32_t i = s1StartIdx; i <= s1EndIdx; i++) { + actNoMaskCmpSize = noMaskCmpSize / constInfo.cmpRatio; + if (actNoMaskCmpSize <= 0) { + right = 0; + } else { + right = actNoMaskCmpSize - cmpTileStart; + } + dealTempSize = constInfo.gSize - gStartIdx; + if (i == s1EndIdx) { + dealTempSize = gEndIdx - gStartIdx; + } + if (dealTempSize == 0) { + continue; + } + SetInfInBlk(mmResUb[ubOffset], dealTempSize, columnCount, right, columnCount - 1); + noMaskCmpSize += 1; + ubOffset += dealTempSize * columnCount; + gStartIdx = 0; + } + } +} + +template +__aicore__ inline void SWAVectorBlock::SetInfInBlk(const LocalTensor &mmResUb, uint32_t dealRowCount, + uint32_t columnCount, int64_t startId, int64_t endId) +{ + // startId endId + // x x x 0 0 0 x x x + // 从startId到endId部分置-inf, endId、startId为endId一个blk内部的下标 + // 左闭右闭 + if (startId > endId) { + return; + } + int64_t start = startId < 0 ? 0 : startId; + int64_t end = endId >= static_cast(columnCount) ? static_cast(columnCount) - 1 : endId; + if (start > end) { + return; + } + + uint64_t curStart = static_cast(start); + uint64_t curEnd = static_cast(end); + while (curStart <= curEnd) { + uint64_t blockStart = curStart / BLOCK_ELEMENT_NUM * BLOCK_ELEMENT_NUM; + uint64_t blockEnd = blockStart + BLOCK_ELEMENT_NUM - 1; + blockEnd = blockEnd > curEnd ? curEnd : blockEnd; + + uint64_t preMask = (1llu << (curStart - blockStart)) - 1; + uint64_t postMask = ~((1llu << (blockEnd - blockStart + 1)) - 1); + uint64_t mask[1] = {~(preMask | postMask)}; + Duplicate(mmResUb[blockStart], SOFTMAX_MIN_NUM, mask, dealRowCount, 1, columnCount / BLOCK_ELEMENT_NUM); + curStart = blockEnd + 1; + } +} + +template +__aicore__ inline void SWAVectorBlock::SoftmaxFlashV2Compute(const RunInfo &info, const MSplitInfo &mSplitInfo, + LocalTensor &mmResUb, + LocalTensor &softmaxTmpUb, + uint32_t startRow, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + LocalTensor inSumTensor; + LocalTensor inMaxTensor; + uint32_t baseOffset = mSplitInfo.nBufferStartM / 2 + startRow; + uint32_t outIdx = info.loop % (constInfo.preLoadNum); + uint32_t softmaxOutOffset = outIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset; + if (info.isFirstSInnerLoop) { + inMaxTensor = softmaxMaxDefaultUb[startRow]; + inSumTensor = softmaxSumDefaultUb; + } else { + uint32_t inIdx = (info.loop - 1) % (constInfo.preLoadNum); + inMaxTensor = softmaxMaxUb[inIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset]; + inSumTensor = softmaxSumUb[inIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset]; + } + if (actualColumnCount != 0) { + SoftMaxShapeInfo srcShape{dealRowCount, columnCount, dealRowCount, actualColumnCount}; + SoftMaxTiling newTiling = + SoftMaxFlashV2TilingFunc(srcShape, sizeof(T), sizeof(T), softmaxTmpUb.GetSize(), true, false); + SoftmaxFlashV2( + mmResUb, softmaxSumUb[softmaxOutOffset], softmaxMaxUb[softmaxOutOffset], mmResUb, + softmaxExpUb[softmaxOutOffset], inSumTensor, inMaxTensor, softmaxTmpUb, newTiling, srcShape); + } else { + uint32_t dealRowCountAlign = SMLAAlign(dealRowCount, FP32_BLOCK_ELEMENT_NUM); + DataCopy(softmaxSumUb[softmaxOutOffset], inSumTensor, dealRowCountAlign); + PipeBarrier(); + DataCopy(softmaxMaxUb[softmaxOutOffset], inMaxTensor, dealRowCountAlign); + } +} + +template +__aicore__ inline void SWAVectorBlock::ProcessLse(const RunInfo &info, const MSplitInfo &mSplitInfo) +{ + if (mSplitInfo.vecDealM == 0) { + return; + } + uint64_t lseOffset; + if (constInfo.outputLayout == SMLA_LAYOUT::TND) { + uint32_t tBase = actualSeqLengthsQGm.GetValue(info.bIdx); + lseOffset = (tBase + info.s1Idx) * constInfo.gSize + info.n2IdxReal * constInfo.qSeqSize * constInfo.gSize; + } else if (constInfo.outputLayout == SMLA_LAYOUT::BSND) { + lseOffset = info.bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize + + info.n2IdxReal * constInfo.qSeqSize * constInfo.gSize + info.s1Idx * constInfo.gSize; + } + lseOffset = lseOffset + mSplitInfo.nBufferStartM + mSplitInfo.vecStartM; + uint32_t baseOffset = mSplitInfo.nBufferStartM / 2; + uint32_t outIdx = info.loop % (constInfo.preLoadNum); + uint32_t softmaxOffset = outIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset; + auto sumTensor = softmaxSumUb[softmaxOffset]; + auto maxTensor = softmaxMaxUb[softmaxOffset]; + auto outLSETensor = outputBuff2.Get(); + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1; + dataCopyParams.blockLen = mSplitInfo.vecDealM * sizeof(T); + dataCopyParams.srcStride = 0; + dataCopyParams.dstStride = 0; + + WaitFlag(SYNC_OUTPUT_BUF2_FLAG); + PipeBarrier(); + Log(outLSETensor, sumTensor, mSplitInfo.vecDealM); + PipeBarrier(); + Add(outLSETensor, outLSETensor, maxTensor, mSplitInfo.vecDealM); + SetFlag(SYNC_OUTPUT_BUF2_FLAG); + WaitFlag(SYNC_OUTPUT_BUF2_FLAG); + + DataCopyPad(softmaxLseGm[lseOffset], outLSETensor, dataCopyParams); + SetFlag(SYNC_OUTPUT_BUF2_FLAG); +} + +template +__aicore__ inline void SWAVectorBlock::DealBmm1ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, + uint32_t startRow, uint32_t dealRowCount, + uint32_t columnCount, uint32_t loopId) +{ + uint32_t computeSize = dealRowCount * columnCount; + uint64_t inOutGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize + + (mSplitInfo.nBufferStartM + mSplitInfo.vecStartM + startRow) * columnCount; + LocalTensor mmResUb = inputBuff1.Get(); + mmResUb = mmResUb[pingpongFlag * INPUT1_BUFFER_OFFSET / sizeof(MM1_OUT_T)]; + WaitFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + + DataCopy(mmResUb, mm1ResGm[inOutGmOffset], computeSize); + SetFlag(SYNC_INPUT_BUF1_FLAG); + WaitFlag(SYNC_INPUT_BUF1_FLAG); + + ElewiseCompute(info, mSplitInfo, mmResUb, startRow, dealRowCount, columnCount); + + PipeBarrier(); + LocalTensor tmpAFloorUb = tmpBuff1.Get(); + LocalTensor softmaxTmpUb = tmpAFloorUb.template ReinterpretCast(); + + SoftmaxFlashV2Compute(info, mSplitInfo, mmResUb, softmaxTmpUb, startRow, dealRowCount, columnCount, + info.actualSingleProcessSInnerSize); + PipeBarrier(); + LocalTensor tmpMMResCastTensor = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + + Cast(tmpMMResCastTensor, mmResUb, AscendC::RoundMode::CAST_ROUND, computeSize); + SetFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + DataCopy(vec1ResGm[inOutGmOffset], tmpMMResCastTensor, computeSize); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); +} + +template +__aicore__ inline void SWAVectorBlock::ProcessVec1SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo) +{ + if (mSplitInfo.vecDealM == 0) { + return; + } + uint32_t mSplitSize = info.actualSingleProcessSInnerSize == 0 ? + 16 : + BASE_BLOCK_MAX_ELEMENT_NUM / info.actualSingleProcessSInnerSizeAlign; + // 1. 向下8对齐是因为UB操作至少32B + // 2. info.actualSingleProcessSInnerSizeAlign最大512, mSplitSize可以确保最小为16 + mSplitSize = mSplitSize / 8 * 8; + + if (mSplitSize > mSplitInfo.vecDealM) { + mSplitSize = mSplitInfo.vecDealM; + } + uint32_t loopCount = (mSplitInfo.vecDealM + mSplitSize - 1) / mSplitSize; + uint32_t tailSplitSize = mSplitInfo.vecDealM - (loopCount - 1) * mSplitSize; + + uint32_t sinkHeadIdx = + (info.n2IdxReal * constInfo.gSize + mSplitInfo.nBufferStartM + mSplitInfo.vecStartM) % constInfo.qHeadNum; + SliceAndContactSinksValue(sinkHeadIdx, mSplitInfo.vecDealM); + + for (uint32_t i = 0, dealSize = mSplitSize; i < loopCount; i++) { + if (i == (loopCount - 1)) { + dealSize = tailSplitSize; + } + DealBmm1ResBaseBlock(info, mSplitInfo, i * mSplitSize, dealSize, info.actualSingleProcessSInnerSizeAlign, i); + pingpongFlag ^= 1; // pingpong 0 1切换 + } +} + +// =======================vec1============================= + +template +__aicore__ inline void SWAVectorBlock::ProcessVec1L(const RunInfo &info) +{ + uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize; + uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize; + for (uint32_t i = 0; i < nBufferLoopTimes; i++) { + MSplitInfo mSplitInfo; + mSplitInfo.nBufferIdx = i; + mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize; + mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail; + + mSplitInfo.vecDealM = (mSplitInfo.nBufferDealM <= 16) ? mSplitInfo.nBufferDealM : + (((mSplitInfo.nBufferDealM + 15) / 16 + 1) / 2 * 16); + mSplitInfo.vecStartM = 0; + if (GetBlockIdx() % 2 == 1) { + mSplitInfo.vecStartM = mSplitInfo.vecDealM; + mSplitInfo.vecDealM = mSplitInfo.nBufferDealM - mSplitInfo.vecDealM; + } + + CrossCoreWaitFlag(constInfo.syncC1V1); + // vec1 compute + ProcessVec1SingleBuf(info, mSplitInfo); + CrossCoreSetFlag(constInfo.syncV1C2); + + // move lse for flash decode or FA + if (constInfo.returnSoftmaxLse && info.s2Idx == info.curSInnerLoopTimes - 1) { + ProcessLse(info, mSplitInfo); + } + } +} + +// =======================vec2============================= + +template +__aicore__ inline uint64_t SWAVectorBlock::CalcAccumOffset(uint32_t bN2Idx, uint32_t gS1Idx) +{ + return 0; +} + +template +__aicore__ inline void SWAVectorBlock::ProcessVec2SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo) +{ + if (mSplitInfo.vecDealM == 0) { + return; + } + + ProcessVec2Inner(info, mSplitInfo, 0, mSplitInfo.vecDealM); +} + +template +__aicore__ inline void SWAVectorBlock::ProcessVec2L(const RunInfo &info) +{ + uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize; + uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize; + for (uint32_t i = 0; i < nBufferLoopTimes; i++) { + MSplitInfo mSplitInfo; + mSplitInfo.nBufferIdx = i; + mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize; + mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail; + + mSplitInfo.vecDealM = (mSplitInfo.nBufferDealM <= 16) ? mSplitInfo.nBufferDealM : + (((mSplitInfo.nBufferDealM + 15) / 16 + 1) / 2 * 16); + mSplitInfo.vecStartM = 0; + if (GetBlockIdx() % 2 == 1) { + mSplitInfo.vecStartM = mSplitInfo.vecDealM; + mSplitInfo.vecDealM = mSplitInfo.nBufferDealM - mSplitInfo.vecDealM; + } + CrossCoreWaitFlag(constInfo.syncC2V2); + ProcessVec2SingleBuf(info, mSplitInfo); + } +} + +template +__aicore__ inline void SWAVectorBlock::ProcessVec2Inner(const RunInfo &info, const MSplitInfo &mSplitInfo, + uint32_t mStartRow, uint32_t mDealSize) +{ + uint32_t mSplitSize = BASE_BLOCK_MAX_ELEMENT_NUM / constInfo.headDim; + if (mSplitSize > mDealSize) { + mSplitSize = mDealSize; + } + + uint32_t loopCount = (mDealSize + mSplitSize - 1) / mSplitSize; + uint32_t tailSplitSize = mDealSize - (loopCount - 1) * mSplitSize; + for (uint32_t i = 0, dealSize = mSplitSize; i < loopCount; i++) { + if (i == (loopCount - 1)) { + dealSize = tailSplitSize; + } + DealBmm2ResBaseBlock(info, mSplitInfo, i * mSplitSize + mStartRow, dealSize, constInfo.headDim, + constInfo.headDim); + pingpongFlag ^= 1; // pingpong 0 1切换 + } +} + +template +__aicore__ inline void SWAVectorBlock::Bmm2FDDataCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, + uint32_t wsMStart, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + LocalTensor tmp = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + DataCopy(tmp, bmm2ResUb, columnCount * dealRowCount); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + uint64_t accumTmpOutNum = CalcAccumOffset(info.bIdx, info.gS1Idx); + uint64_t offset = + accumTmpOutNum * constInfo.kvHeadNum * constInfo.mBaseSize * constInfo.headDim + // taskoffset + info.tndCoreStartKVSplitPos * constInfo.kvHeadNum * constInfo.mBaseSize * constInfo.headDim + // 份数offset + wsMStart * actualColumnCount; // m轴offset + GlobalTensor dst = accumOutGm[offset]; + if (info.actualSingleProcessSInnerSize == 0) { + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = dealRowCount; + dataCopyParams.blockLen = actualColumnCount * sizeof(T); + dataCopyParams.srcStride = (columnCount - actualColumnCount) / (BYTE_BLOCK / sizeof(T)); + dataCopyParams.dstStride = 0; + DataCopyPad(dst, tmp, dataCopyParams); + } else { + matmul::InitOutput(dst, dealRowCount * actualColumnCount, ConstInfo::FLOAT_ZERO); + } + SetFlag(SYNC_OUTPUT_BUF1_FLAG); +} + +template +__aicore__ inline void SWAVectorBlock::Bmm2DataCopyOutTrans(const RunInfo &info, LocalTensor &attenOutUb, + uint32_t wsMStart, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = dealRowCount; + dataCopyParams.blockLen = actualColumnCount * sizeof(OUT_T); + dataCopyParams.srcStride = (columnCount - actualColumnCount) / (BYTE_BLOCK / sizeof(OUT_T)); + dataCopyParams.dstStride = 0; + DataCopyPad(attentionOutGm[info.attenOutOffset + wsMStart * actualColumnCount], attenOutUb, dataCopyParams); + return; +} + +template +__aicore__ inline void SWAVectorBlock::Bmm2CastAndCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, + uint32_t wsMStart, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + LocalTensor tmpBmm2ResCastTensor = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + if constexpr (IsSameType::value) { // bf16 采取四舍六入五成双模式 + Cast(tmpBmm2ResCastTensor, bmm2ResUb, AscendC::RoundMode::CAST_RINT, dealRowCount * columnCount); + } else { + Cast(tmpBmm2ResCastTensor, bmm2ResUb, AscendC::RoundMode::CAST_ROUND, dealRowCount * columnCount); + } + + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + Bmm2DataCopyOutTrans(info, tmpBmm2ResCastTensor, wsMStart, dealRowCount, columnCount, actualColumnCount); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); +} + +template +__aicore__ inline void SWAVectorBlock::Bmm2ResCopyOut(const RunInfo &info, LocalTensor &bmm2ResUb, + uint32_t wsMStart, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + if constexpr (FLASH_DECODE) { + if (info.tndIsS2SplitCore) { + Bmm2FDDataCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount); + } else { + Bmm2CastAndCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount); + } + } else { + Bmm2CastAndCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount); + } +} + +template +__aicore__ inline void SWAVectorBlock::DealBmm2ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, + uint32_t startRow, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + uint32_t vec2ComputeSize = dealRowCount * columnCount; + uint32_t mStart = mSplitInfo.nBufferStartM + mSplitInfo.vecStartM + startRow; + uint64_t srcGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + mStart * columnCount; + LocalTensor tmpBmm2ResUb = inputBuff1.Get(); + tmpBmm2ResUb = tmpBmm2ResUb[pingpongFlag * INPUT1_BUFFER_OFFSET / sizeof(MM2_OUT_T)]; + WaitFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + DataCopy(tmpBmm2ResUb, mm2ResGm[srcGmOffset], vec2ComputeSize); + + SetFlag(SYNC_INPUT_BUF1_FLAG); + WaitFlag(SYNC_INPUT_BUF1_FLAG); + + LocalTensor bmm2ResUb = tmpBuff1.Get(); + bmm2ResUb.SetSize(vec2ComputeSize); + DataCopy(bmm2ResUb, tmpBmm2ResUb, vec2ComputeSize); + SetFlag(SYNC_INPUT_BUF1_FLAG + pingpongFlag); + + uint32_t inOutBaseOffset = mStart * columnCount; + uint32_t baseOffset = mSplitInfo.nBufferStartM / 2 + startRow; + + // 除第一个循环外,均需要更新中间计算结果 + if (!info.isFirstSInnerLoop) { + event_t eventIdMte2WaitMte3 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2)); + SetFlag(eventIdMte2WaitMte3); + WaitFlag(eventIdMte2WaitMte3); + + LocalTensor bmm2ResPreUb = inputBuff2.Get(); + WaitFlag(SYNC_INPUT_BUF2_FLAG); + + uint64_t vec2ResGmOffset = ((info.loop - 1) % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + inOutBaseOffset; + DataCopy(bmm2ResPreUb, vec2ResGm[vec2ResGmOffset], vec2ComputeSize); + + SetFlag(SYNC_INPUT_BUF2_FLAG); + WaitFlag(SYNC_INPUT_BUF2_FLAG); + + uint32_t idx = info.loop % (constInfo.preLoadNum); + LocalTensor expUb = v0ValidSizeBuff.Get()[384]; // sumUb用临时内存 16 * 32B = 512B + Brcb(expUb, softmaxExpUb[idx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset], (dealRowCount + 7) / 8, + {1, 8}); + PipeBarrier(); + + RowMuls(bmm2ResPreUb, bmm2ResPreUb, expUb, dealRowCount, columnCount, actualColumnCount); + AscendC::PipeBarrier(); + Add(bmm2ResUb, bmm2ResUb, bmm2ResPreUb, vec2ComputeSize); + AscendC::PipeBarrier(); + + SetFlag(SYNC_INPUT_BUF2_FLAG); + } + + // 最后一次输出计算结果,否则将中间结果暂存至workspace + if (info.isLastS2Loop) { + uint32_t idx = info.loop % (constInfo.preLoadNum); + LocalTensor tmpSumUb = v0ValidSizeBuff.Get()[384]; // sumUb用临时内存 16 * 32B = 512B + Brcb(tmpSumUb, softmaxSumUb[idx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset], (dealRowCount + 7) / 8, + {1, 8}); + PipeBarrier(); + RowDivs(bmm2ResUb, bmm2ResUb, tmpSumUb, dealRowCount, columnCount, actualColumnCount); + PipeBarrier(); + Bmm2ResCopyOut(info, bmm2ResUb, mStart, dealRowCount, columnCount, actualColumnCount); + } else { + LocalTensor outUb = outputBuff1.Get(); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + DataCopy(outUb, bmm2ResUb, dealRowCount * columnCount); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + WaitFlag(SYNC_OUTPUT_BUF1_FLAG); + uint64_t vec2ResGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + inOutBaseOffset; + DataCopy(vec2ResGm[vec2ResGmOffset], outUb, vec2ComputeSize); + SetFlag(SYNC_OUTPUT_BUF1_FLAG); + } +} + +template +__aicore__ inline void SWAVectorBlock::RowDivs(LocalTensor dstUb, LocalTensor src0Ub, + LocalTensor src1Ub, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + // divs by row, 每行的元素除以相同的元素 + // dstUb[i, (j * 8) : (j * 8 + 7)] = src0Ub[i, (j * 8) : (j * 8 + 7)] / src1Ub[i, 0 : 7] + // src0Ub:[dealRowCount, columnCount], src1Ub:[dealRowCount, FP32_BLOCK_ELEMENT_NUM] dstUb:[dealRowCount, + // columnCount] + uint32_t dtypeMask = FP32_REPEAT_ELEMENT_NUM; + uint32_t dLoop = actualColumnCount / dtypeMask; + uint32_t dRemain = actualColumnCount % dtypeMask; + + BinaryRepeatParams repeatParamsDiv; + repeatParamsDiv.src0BlkStride = 1; + repeatParamsDiv.src1BlkStride = 0; + repeatParamsDiv.dstBlkStride = 1; + repeatParamsDiv.src0RepStride = columnCount / FP32_BLOCK_ELEMENT_NUM; + repeatParamsDiv.src1RepStride = 1; + repeatParamsDiv.dstRepStride = columnCount / FP32_BLOCK_ELEMENT_NUM; + uint32_t columnRepeatCount = dLoop; + if (columnRepeatCount <= dealRowCount) { + uint32_t offset = 0; + for (uint32_t i = 0; i < dLoop; i++) { + Div(dstUb[offset], src0Ub[offset], src1Ub, dtypeMask, dealRowCount, repeatParamsDiv); + offset += dtypeMask; + } + } else { + BinaryRepeatParams columnRepeatParams; + columnRepeatParams.src0BlkStride = 1; + columnRepeatParams.src1BlkStride = 0; + columnRepeatParams.dstBlkStride = 1; + columnRepeatParams.src0RepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + columnRepeatParams.src1RepStride = 0; + columnRepeatParams.dstRepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + uint32_t offset = 0; + for (uint32_t i = 0; i < dealRowCount; i++) { + Div(dstUb[offset], src0Ub[offset], src1Ub[i * FP32_BLOCK_ELEMENT_NUM], dtypeMask, columnRepeatCount, + columnRepeatParams); + offset += columnCount; + } + } + if (dRemain > 0) { + Div(dstUb[dLoop * dtypeMask], src0Ub[dLoop * dtypeMask], src1Ub, dRemain, dealRowCount, repeatParamsDiv); + } +} + +template +__aicore__ inline void SWAVectorBlock::RowMuls(LocalTensor dstUb, LocalTensor src0Ub, + LocalTensor src1Ub, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + // muls by row, 每行的元素乘以相同的元素 + // dstUb[i, (j * 8) : (j * 8 + 7)] = src0Ub[i, (j * 8) : (j * 8 + 7)] * src1Ub[i, 0 : 7] + // src0Ub:[dealRowCount, columnCount] src1Ub:[dealRowCount, FP32_BLOCK_ELEMENT_NUM] dstUb:[dealRowCount, + // columnCount] + // dealRowCount is repeat times, must be less 256 + uint32_t repeatElementNum = FP32_REPEAT_ELEMENT_NUM; + uint32_t blockElementNum = FP32_BLOCK_ELEMENT_NUM; + + if constexpr (std::is_same::value) { + // 此限制由于每个repeat至多连续读取256B数据 + repeatElementNum = FP32_REPEAT_ELEMENT_NUM * 2; // 256/4 * 2=128 + blockElementNum = FP32_BLOCK_ELEMENT_NUM * 2; // 32/4 * 2 = 16 + } + + // 每次只能连续读取256B的数据进行计算,故每次只能处理256B/sizeof(dType)= + // 列方向分dLoop次,每次处理8列数据 + uint32_t dLoop = actualColumnCount / repeatElementNum; + uint32_t dRemain = actualColumnCount % repeatElementNum; + // REPEATE_STRIDE_UP_BOUND=256, 此限制由于src0RepStride数据类型为uint8之多256个datablock间距 + if (columnCount < REPEATE_STRIDE_UP_BOUND * blockElementNum) { + BinaryRepeatParams repeatParams; + repeatParams.src0BlkStride = 1; + repeatParams.src1BlkStride = 0; + repeatParams.dstBlkStride = 1; + repeatParams.src0RepStride = columnCount / blockElementNum; + repeatParams.src1RepStride = 1; + repeatParams.dstRepStride = columnCount / blockElementNum; + + // 如果以列为repeat所处理的次数小于行处理次数,则以列方式处理。反之则以行进行repeat处理 + if (dLoop <= dealRowCount) { + uint32_t offset = 0; + for (uint32_t i = 0; i < dLoop; i++) { + Mul(dstUb[offset], src0Ub[offset], src1Ub, repeatElementNum, dealRowCount, repeatParams); + offset += repeatElementNum; + } + } else { + BinaryRepeatParams columnRepeatParams; + columnRepeatParams.src0BlkStride = 1; + columnRepeatParams.src1BlkStride = 0; + columnRepeatParams.dstBlkStride = 1; + columnRepeatParams.src0RepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + columnRepeatParams.src1RepStride = 0; + columnRepeatParams.dstRepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + for (uint32_t i = 0; i < dealRowCount; i++) { + Mul(dstUb[i * columnCount], src0Ub[i * columnCount], src1Ub[i * blockElementNum], repeatElementNum, + dLoop, columnRepeatParams); + } + } + + // 最后一次完成[dealRowCount, dRemain] * [dealRowCount, blockElementNum] 只计算有效部分 + if (dRemain > 0) { + Mul(dstUb[dLoop * repeatElementNum], src0Ub[dLoop * repeatElementNum], src1Ub, dRemain, dealRowCount, + repeatParams); + } + } else { + BinaryRepeatParams repeatParams; + repeatParams.src0RepStride = 8; // 每个repeat为256B数据,正好8个datablock + repeatParams.src0BlkStride = 1; + repeatParams.src1RepStride = 0; + repeatParams.src1BlkStride = 0; + repeatParams.dstRepStride = 8; + repeatParams.dstBlkStride = 1; + // 每次计算一行,共计算dealRowCount行 + for (uint32_t i = 0; i < dealRowCount; i++) { + // 计算一行中的dLoop个repeat, 每个repeat计算256/block_size 个data_block + Mul(dstUb[i * columnCount], src0Ub[i * columnCount], src1Ub[i * blockElementNum], repeatElementNum, dLoop, + repeatParams); + // 计算一行中的尾块 + if (dRemain > 0) { + Mul(dstUb[i * columnCount + dLoop * repeatElementNum], + src0Ub[i * columnCount + dLoop * repeatElementNum], src1Ub[i * blockElementNum], dRemain, 1, + repeatParams); + } + } + } +} +} // namespace SMLAKernel +#endif diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_kernel.h b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_kernel.h new file mode 100644 index 000000000000..52560699fbb2 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_kernel.h @@ -0,0 +1,1008 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_swa_kernel.h + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_SWA_KERNEL_H +#define SPARSE_FLASH_MLA_SWA_KERNEL_H + +#include "kernel_operator.h" +#include "kernel_operator_list_tensor_intf.h" +#include "kernel_tiling/kernel_tiling.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" +#include "sparse_flash_mla_common_arch22.h" +#include "sparse_flash_mla_swa_block_cube.h" +#include "sparse_flash_mla_swa_block_vector.h" +#include "sparse_flash_mla_arch22_metadata.h" + +namespace SMLAKernel { +using namespace matmul; +using namespace optiling; +using AscendC::CrossCoreSetFlag; +using AscendC::CrossCoreWaitFlag; + +// 由于S2循环前,RunInfo还没有赋值,使用Bngs1Param临时存放B、N、S1轴相关的信息;同时减少重复计算 +struct SwaTempLoopInfo { + uint32_t bn2IdxInCurCore = 0; + uint32_t bIdx = 0U; + uint32_t n2Idx = 0U; + uint32_t s2LoopTimes = 0U; // S2方向循环的总次数,无论TND还是BXXD都是等于实际次数,不用减1 + uint64_t s2BasicSizeTail = 0U; // S2方向循环的尾基本块大小 + + int32_t actS1Size = 0; // TND场景下当前Batch循环处理的S1轴的大小 + int32_t actOriS2Size = 0; + int32_t actCmpS2Size = 0; + + bool curActSeqLenIsZero = false; + bool hasInvalidRow = false; + bool tndIsS2SplitCore = false; + bool resv; + uint32_t tndCoreStartKVSplitPos = 0; + uint32_t gS1Idx = 0U; + uint32_t s1StartIdx = 0; + uint32_t s1EndIdx = 0; + uint64_t mBasicSizeTail = 0U; // gS1方向循环的尾基本块大小 + uint32_t cmpLoopTimes = 0; + uint32_t oriLoopTimes = 0; + uint32_t oriCmpMixLoopTimes = 0; + uint32_t cmpMixSize = 0; + + int32_t oriMaskRight = 0; + int32_t oriMaskLeft = 0; + int32_t cmpMaskRight = 0; + + uint64_t actualSeqQPrefixSum = 0; + uint64_t actualSeqKVPrefixSum = 0; + uint64_t actualSeqCmpKVPrefixSum = 0; +}; + +template +class SparseFlashMlaSwa { +public: + // 中间计算数据类型为float,高精度模式 + using T = float; + using Q_T = typename SMLAT::queryType; + using KV_T = typename SMLAT::kvType; + using OUT_T = typename SMLAT::outputType; + using SINKS_T = float; + using UPDATE_T = T; + using MM1_OUT_T = T; + using MM2_OUT_T = T; + + __aicore__ inline SparseFlashMlaSwa(){}; + __aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, + __gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensKV, __gm__ uint8_t *cuSeqlensCmpKV, + __gm__ uint8_t *seqUsedQ, __gm__ uint8_t *seqUsedKV, __gm__ uint8_t *seqUsedCmpKV, + __gm__ uint8_t *cmpResidualKV, __gm__ uint8_t *oriTopkLength, + __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata, + __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, __gm__ uint8_t *workspace, + const SparseFlashMlaTilingData *__restrict tiling, __gm__ uint8_t *gmTiling, + TPipe *tPipe); + + __aicore__ inline void Process(); + +private: + static constexpr bool PAGE_ATTENTION = SMLAT::pageAttention; + static constexpr int TEMPLATE_MODE = SMLAT::templateMode; + static constexpr bool FLASH_DECODE = SMLAT::flashDecode; + static constexpr SMLA_LAYOUT LAYOUT_T = SMLAT::layout; + static constexpr SMLA_LAYOUT KV_LAYOUT_T = SMLAT::kvLayout; + + static constexpr uint32_t PRELOAD_NUM = 2; + static constexpr uint32_t N_BUFFER_M_BASIC_SIZE = 256; + static constexpr uint32_t SMLA_PRELOAD_TASK_CACHE_SIZE = 3; + + static constexpr uint32_t SYNC_V0_C1_FLAG = 6; + static constexpr uint32_t SYNC_C1_V1_FLAG = 7; + static constexpr uint32_t SYNC_V1_C2_FLAG = 8; + static constexpr uint32_t SYNC_C2_V2_FLAG = 9; + static constexpr uint32_t MERGE_CACHE_GM_BUF_NUM = 3; + + static constexpr uint64_t SYNC_MM2RES_BUF1_FLAG = 10; + static constexpr uint64_t SYNC_MM2RES_BUF2_FLAG = 11; + static constexpr uint64_t SYNC_FDOUTPUT_BUF_FLAG = 12; + + static constexpr uint64_t headDim = 512ULL; + static constexpr uint64_t headDimAlign = 512ULL; + static constexpr uint32_t msdIterNum = 2U; + + static constexpr uint32_t dbWorkspaceRatio = PRELOAD_NUM; + + const SparseFlashMlaTilingData *__restrict tilingData = nullptr; + + TPipe *pipe = nullptr; + GlobalTensor metadataGm; + + uint64_t mSizeVStart = 0ULL; + int64_t threshold = 0; + uint64_t s2BatchBaseOffset = 0; + uint64_t tensorACoreOffset = 0ULL; + uint64_t tensorBCoreOffset = 0ULL; + uint64_t tensorCmpBCoreOffset = 0ULL; + uint64_t attenOutOffset = 0ULL; + + uint32_t tmpBlockIdx = 0U; + uint32_t aiCoreIdx = 0U; + + ConstInfo constInfo{}; + SwaTempLoopInfo tempLoopInfo{}; + + SWACubeBlock cubeBlock; + SWAVectorBlock vectorBlock; + + GlobalTensor queryGm; + GlobalTensor oriKvGm; + GlobalTensor cmpKvGm; + GlobalTensor sinksGm; + + GlobalTensor attentionOutGm; + GlobalTensor softmaxLseGm; + + GlobalTensor oriBlockTableGm; + GlobalTensor cmpBlockTableGm; + GlobalTensor oriSparseIndicesGm; + GlobalTensor oriTopkLengthGm; + + GlobalTensor actualSeqLengthsQGm; + GlobalTensor actualSeqLengthsKVGm; + GlobalTensor actualSeqLengthsCmpKVGm; + GlobalTensor cmpResidualKVGm; + + // workspace + GlobalTensor mm1ResGm; + GlobalTensor vec1ResGm; + GlobalTensor mm2ResGm; + + GlobalTensor vec2ResGm; + + GlobalTensor kvMergeGm_; + GlobalTensor accumOutGm; + // ================================Init functions================================== + __aicore__ inline void InitTilingData(); + __aicore__ inline void InitCalcParamsEach(); + __aicore__ inline void InitBuffers(); + __aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKv); + __aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKV, + __gm__ uint8_t *actualSeqLengthsCmpKV); + __aicore__ inline void InitOutputSingleCore(); + // ================================Process functions================================ + __aicore__ inline void ProcessBalance(); + __aicore__ inline void PreloadPipeline(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, uint64_t s2LoopIdx, + RunInfo extraInfo[SMLA_PRELOAD_TASK_CACHE_SIZE]); + // ================================Offset Calc===================================== + __aicore__ inline void GetSparseActualSeqLen(); + __aicore__ inline void UpdateInnerLoopCond(); + __aicore__ inline void CalcParams(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, uint32_t s2LoopIdx, + RunInfo &info); + __aicore__ inline int32_t GetActualSeqLenQ(uint32_t bIdx); + __aicore__ inline int32_t GetActualSeqLenKV(uint32_t bIdx); + __aicore__ inline uint32_t GetOriSparseActualSeqLen(); + __aicore__ inline int32_t GetActualSeqLenCmpKV(uint32_t bIdx, int32_t actualOriS2Size); + __aicore__ inline int32_t GetCmpMaskS2Size(uint32_t bIdx, int32_t actualOriS2Size, int32_t actualCmpS2Size); + __aicore__ inline void GetBN2Idx(uint32_t bN2Idx, uint32_t &bIdx, uint32_t &n2Idx); + // ================================Mm1============================================== + __aicore__ inline void ComputeMm1(const RunInfo &info); + // ================================Mm2============================================== + __aicore__ inline void ComputeMm2(const RunInfo &info); + __aicore__ inline void InitAllZeroOutput(uint32_t bIdx, uint32_t inValidRowS1StartIdx, int32_t inValidRowCount, + uint32_t n2Idx); +}; + +template +__aicore__ inline void SparseFlashMlaSwa::InitTilingData() +{ + // singleCoreParams + // singleCoreTensorSize + constInfo.mmResUbSize = tilingData->baseParams.mmResUbSize; + constInfo.bmm2ResUbSize = tilingData->baseParams.bmm2ResUbSize; + constInfo.usedCoreNum = tilingData->baseParams.usedCoreNum; + // baseParams + constInfo.batchSize = tilingData->baseParams.batchSize; + constInfo.gSize = tilingData->baseParams.nNumOfQInOneGroup; + constInfo.kvHeadNum = (tilingData->baseParams.kvHeadNum == 0) ? 1 : tilingData->baseParams.kvHeadNum; + constInfo.qHeadNum = constInfo.gSize * constInfo.kvHeadNum; + constInfo.kvSeqSize = tilingData->baseParams.kvSeqSize; + constInfo.qSeqSize = tilingData->baseParams.qSeqSize; + constInfo.oriMaxBlockNumPerBatch = tilingData->baseParams.oriMaxBlockNumPerBatch; + constInfo.kvCacheBlockSize = tilingData->baseParams.paBlockSize; + + constInfo.paOriBlockSize = tilingData->baseParams.oriBlockSize; + constInfo.paCmpBlockSize = tilingData->baseParams.cmpBlockSize; + constInfo.outputLayout = static_cast(tilingData->baseParams.outputLayout); + constInfo.headDim = headDim; + constInfo.oriMaskMode = tilingData->baseParams.oriMaskMode; + constInfo.oriKvStride0 = tilingData->baseParams.oriKvStride0; + constInfo.oriWinLeft = tilingData->baseParams.oriWinLeft; + constInfo.oriWinRight = tilingData->baseParams.oriWinRight; + constInfo.hasOriSparseIndices = tilingData->baseParams.hasOriSparseIndices != 0; + constInfo.oriSparseIndexWidth = tilingData->baseParams.oriSparseIndexWidth; + constInfo.returnSoftmaxLse = tilingData->baseParams.returnSoftmaxLse; + + constInfo.actualLenDimsQ = tilingData->baseParams.actualLenDimsQ; + constInfo.actualLenDimsKV = tilingData->baseParams.actualLenDimsKV; + constInfo.actualLenDimsCmpKV = tilingData->baseParams.actualLenDimsCmpKV; + constInfo.cmpResidualKVSize = tilingData->baseParams.cmpResidualKVSize; + + // innerSplitParams + constInfo.mBaseSize = tilingData->baseParams.mBaseSize; + constInfo.s2BaseSize = tilingData->baseParams.s2BaseSize; + constInfo.sparseBlockSize = tilingData->baseParams.sparseBlockSize; + + constInfo.preLoadNum = PRELOAD_NUM; + constInfo.nBufferMBaseSize = N_BUFFER_M_BASIC_SIZE; + constInfo.syncV0C1 = SYNC_V0_C1_FLAG; + constInfo.syncC1V1 = SYNC_C1_V1_FLAG; + constInfo.syncV1C2 = SYNC_V1_C2_FLAG; + constInfo.syncC2V2 = SYNC_C2_V2_FLAG; + constInfo.templateMode = TEMPLATE_MODE; + + // cmp + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + constInfo.cmpRatio = tilingData->cmpParams.cmpRatio; + constInfo.cmpMaskMode = tilingData->cmpParams.cmpMaskMode; + constInfo.cmpKvStride0 = tilingData->cmpParams.cmpKvStride0; + constInfo.cmpMaxBlockNumPerBatch = tilingData->cmpParams.cmpMaxBlockNumPerBatch; + constInfo.cmpSeqSize = tilingData->cmpParams.cmpKvSeqSize; + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::InitBuffers() +{ + if ASCEND_IS_AIV { + vectorBlock.InitBuffers(pipe); + } else { + cubeBlock.InitBuffers(pipe); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, + __gm__ uint8_t *actualSeqLengthsKv) +{ + if (constInfo.actualLenDimsKV != 0) { + actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsKv, constInfo.actualLenDimsKV); + } + if (constInfo.actualLenDimsQ != 0) { + actualSeqLengthsQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ, constInfo.actualLenDimsQ); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, + __gm__ uint8_t *actualSeqLengthsKV, + __gm__ uint8_t *actualSeqLengthsCmpKV) +{ + if (constInfo.actualLenDimsKV != 0) { + actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsKV, constInfo.actualLenDimsKV); + } + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + if (constInfo.actualLenDimsCmpKV != 0) { + actualSeqLengthsCmpKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsCmpKV, + constInfo.actualLenDimsCmpKV); + } + } + if (constInfo.actualLenDimsQ != 0) { + actualSeqLengthsQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ, constInfo.actualLenDimsQ); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::InitAllZeroOutput(uint32_t bIdx, uint32_t inValidRowS1StartIdx, + int32_t inValidRowCount, uint32_t n2Idx) +{ + if (constInfo.outputLayout == SMLA_LAYOUT::TND) { + if (tempLoopInfo.actS1Size == 0) { + return; + } + uint32_t tBase = actualSeqLengthsQGm.GetValue(bIdx); + uint64_t attenOutOffset = + (tBase + inValidRowS1StartIdx) * constInfo.kvHeadNum * constInfo.gSize * constInfo.headDim + + n2Idx * constInfo.gSize * constInfo.headDim; // N2轴偏移 + uint64_t lseOffset = (tBase + inValidRowS1StartIdx) * constInfo.gSize + // T轴、s1轴偏移 + n2Idx * constInfo.qSeqSize * constInfo.gSize; // N2轴偏移 + if (constInfo.kvHeadNum == 1 || inValidRowCount <= 1) { + matmul::InitOutput(attentionOutGm[attenOutOffset], + inValidRowCount * constInfo.gSize * constInfo.headDim, 0); + } else { + uint64_t attenOutRowStride = constInfo.qHeadNum * constInfo.headDim; + for (int32_t rowIdx = 0; rowIdx < inValidRowCount; ++rowIdx) { + matmul::InitOutput( + attentionOutGm[attenOutOffset + static_cast(rowIdx) * attenOutRowStride], + constInfo.gSize * constInfo.headDim, 0); + } + } + if (constInfo.returnSoftmaxLse) { + matmul::InitOutput(softmaxLseGm[lseOffset], inValidRowCount * constInfo.gSize, 0); + } + } else if (constInfo.outputLayout == SMLA_LAYOUT::BSND) { + uint64_t attenOutOffset = + bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize * constInfo.headDim + + inValidRowS1StartIdx * constInfo.kvHeadNum * constInfo.gSize * constInfo.headDim + + n2Idx * constInfo.gSize * constInfo.headDim; // N2轴偏移 + uint64_t lseOffset = bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize + // B轴偏移 + n2Idx * constInfo.qSeqSize * constInfo.gSize + // N2轴偏移 + inValidRowS1StartIdx * constInfo.gSize; // S1轴偏移 + if (constInfo.kvHeadNum == 1 || inValidRowCount <= 1) { + matmul::InitOutput(attentionOutGm[attenOutOffset], + inValidRowCount * constInfo.gSize * constInfo.headDim, 0); + } else { + uint64_t attenOutRowStride = constInfo.qHeadNum * constInfo.headDim; + for (int32_t rowIdx = 0; rowIdx < inValidRowCount; ++rowIdx) { + matmul::InitOutput( + attentionOutGm[attenOutOffset + static_cast(rowIdx) * attenOutRowStride], + constInfo.gSize * constInfo.headDim, 0); + } + } + if (constInfo.returnSoftmaxLse) { + matmul::InitOutput(softmaxLseGm[lseOffset], inValidRowCount * constInfo.gSize, 0); + } + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::InitOutputSingleCore() +{ + uint32_t coreNum = GetBlockNum(); + if (coreNum != 0) { + uint64_t totalOutputSize = constInfo.batchSize * constInfo.qHeadNum * constInfo.qSeqSize * constInfo.headDim; + uint64_t singleCoreSize = (totalOutputSize + (2 * coreNum) - 1) / (2 * coreNum); // 2 means c:v = 1:2 + uint64_t tailSize = totalOutputSize - tmpBlockIdx * singleCoreSize; + uint64_t singleInitOutputSize = tailSize < singleCoreSize ? tailSize : singleCoreSize; + if (singleInitOutputSize > 0) { + matmul::InitOutput(attentionOutGm[tmpBlockIdx * singleCoreSize], singleInitOutputSize, 0); + } + SyncAll(); + } +} + +template +__aicore__ inline int32_t SparseFlashMlaSwa::GetActualSeqLenQ(uint32_t bIdx) +{ + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + int32_t actualSeqQPrefixSum = actualSeqLengthsQGm.GetValue(bIdx); + int32_t actualSeqQNextSum = actualSeqLengthsQGm.GetValue(bIdx + 1); + tempLoopInfo.actualSeqQPrefixSum = static_cast(actualSeqQPrefixSum); + return actualSeqQNextSum - actualSeqQPrefixSum; + } else { + tempLoopInfo.actualSeqQPrefixSum = static_cast(bIdx * constInfo.qSeqSize); + if (constInfo.actualLenDimsQ == 0) { + return static_cast(constInfo.qSeqSize); + } else { + return actualSeqLengthsQGm.GetValue(bIdx); + } + } +} + +template +__aicore__ inline int32_t SparseFlashMlaSwa::GetActualSeqLenKV(uint32_t bIdx) +{ + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + tempLoopInfo.actualSeqKVPrefixSum = static_cast(bIdx * constInfo.kvSeqSize); + if (constInfo.actualLenDimsKV == 0) { + return static_cast(constInfo.kvSeqSize); + } + return actualSeqLengthsKVGm.GetValue(bIdx); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::BSND) { + tempLoopInfo.actualSeqKVPrefixSum = static_cast(bIdx * constInfo.kvSeqSize); + if (constInfo.actualLenDimsKV != 0) { + return actualSeqLengthsKVGm.GetValue(bIdx); + } + return static_cast(constInfo.kvSeqSize); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + int32_t actualSeqKVPrefixSum = actualSeqLengthsKVGm.GetValue(bIdx); + int32_t actualSeqKVNextSum = actualSeqLengthsKVGm.GetValue(bIdx + 1); + tempLoopInfo.actualSeqKVPrefixSum = actualSeqKVPrefixSum; + return actualSeqKVNextSum - actualSeqKVPrefixSum; + } +} + +template +__aicore__ inline int32_t SparseFlashMlaSwa::GetActualSeqLenCmpKV(uint32_t bIdx, int32_t actualOriS2Size) +{ + (void)actualOriS2Size; + if constexpr (TEMPLATE_MODE != HCA_TEMPLATE) { + return 0; + } + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + int32_t actualSeqCmpKVPrefixSum = actualSeqLengthsCmpKVGm.GetValue(bIdx); + int32_t actualSeqCmpKVNextSum = actualSeqLengthsCmpKVGm.GetValue(bIdx + 1); + tempLoopInfo.actualSeqCmpKVPrefixSum = actualSeqCmpKVPrefixSum; + return actualSeqCmpKVNextSum - actualSeqCmpKVPrefixSum; + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + tempLoopInfo.actualSeqCmpKVPrefixSum = static_cast(bIdx * constInfo.cmpSeqSize); + if (constInfo.actualLenDimsCmpKV != 0) { + return actualSeqLengthsCmpKVGm.GetValue(bIdx); + } + return static_cast(constInfo.cmpSeqSize); + } else { + tempLoopInfo.actualSeqCmpKVPrefixSum = static_cast(bIdx * constInfo.cmpSeqSize); + if (constInfo.actualLenDimsCmpKV != 0) { + return actualSeqLengthsCmpKVGm.GetValue(bIdx); + } + return (constInfo.cmpSeqSize != 0) ? static_cast(constInfo.cmpSeqSize) : + actualOriS2Size / static_cast(constInfo.cmpRatio); + } +} + +template +__aicore__ inline int32_t SparseFlashMlaSwa::GetCmpMaskS2Size(uint32_t bIdx, int32_t actualOriS2Size, + int32_t actualCmpS2Size) +{ + (void)actualOriS2Size; + if constexpr (TEMPLATE_MODE != HCA_TEMPLATE) { + return actualOriS2Size; + } + int32_t residual = 0; + if (constInfo.cmpResidualKVSize != 0) { + residual = cmpResidualKVGm.GetValue(bIdx); + } + return actualCmpS2Size * static_cast(constInfo.cmpRatio) + residual; +} + +template +__aicore__ inline uint32_t SparseFlashMlaSwa::GetOriSparseActualSeqLen() +{ + if (!constInfo.hasOriSparseIndices || constInfo.oriSparseIndexWidth == 0 || tempLoopInfo.actOriS2Size <= 0) { + return 0; + } + uint64_t qTokenOffset = tempLoopInfo.actualSeqQPrefixSum + tempLoopInfo.s1StartIdx; + uint64_t topkLenOffset = qTokenOffset * constInfo.kvHeadNum + tempLoopInfo.n2Idx; + if (!constInfo.hasOriTopkLength) { + return 0; + } + int32_t topkLen = oriTopkLengthGm.GetValue(topkLenOffset); + if (topkLen <= 0) { + return 0; + } + return Min(static_cast(topkLen), constInfo.oriSparseIndexWidth); +} + +template +__aicore__ inline void SparseFlashMlaSwa::GetSparseActualSeqLen() +{ + // 行无效通过ori部分判断, ori部分如果有行无效那么ori和cmp都有 + if (static_cast(tempLoopInfo.s1EndIdx) < -(tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size)) { + tempLoopInfo.actOriS2Size = 0; + tempLoopInfo.actCmpS2Size = 0; + return; + } + + // 对于cmp部分还有top k, tempLoopInfo.actS2Size只针对cmp + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + int32_t thresHold = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio; + tempLoopInfo.actCmpS2Size = Min(tempLoopInfo.actCmpS2Size, Max(thresHold, 0)); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::UpdateInnerLoopCond() +{ + if ((tempLoopInfo.actCmpS2Size == 0 && tempLoopInfo.actOriS2Size == 0) || (tempLoopInfo.actS1Size == 0)) { + tempLoopInfo.curActSeqLenIsZero = true; + return; + } + tempLoopInfo.curActSeqLenIsZero = false; + tempLoopInfo.mBasicSizeTail = + ((tempLoopInfo.s1EndIdx - tempLoopInfo.s1StartIdx + 1) * constInfo.gSize) % constInfo.mBaseSize; + tempLoopInfo.mBasicSizeTail = + (tempLoopInfo.mBasicSizeTail == 0) ? constInfo.mBaseSize : tempLoopInfo.mBasicSizeTail; +} + +template +__aicore__ inline void SparseFlashMlaSwa::Init( + __gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensKV, __gm__ uint8_t *cuSeqlensCmpKV, __gm__ uint8_t *seqUsedQ, + __gm__ uint8_t *seqUsedKV, __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata, + __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, __gm__ uint8_t *workspace, + const SparseFlashMlaTilingData *__restrict tiling, __gm__ uint8_t *gmTiling, TPipe *tPipe) +{ + if ASCEND_IS_AIV { + tmpBlockIdx = GetBlockIdx(); // vec:0-47 + aiCoreIdx = tmpBlockIdx / 2; + } else { + tmpBlockIdx = GetBlockIdx(); // cube:0-23 + aiCoreIdx = tmpBlockIdx; + } + + // init tiling data + tilingData = tiling; + + InitTilingData(); + constInfo.hasOriTopkLength = (oriTopkLength != nullptr); + if (constInfo.hasOriTopkLength) { + oriTopkLengthGm.SetGlobalBuffer((__gm__ int32_t *)oriTopkLength); + } + (void)cmpTopkLength; + if (KV_LAYOUT_T == SMLA_LAYOUT::TND && LAYOUT_T == SMLA_LAYOUT::TND) { + InitActualSeqLen(cuSeqlensQ, cuSeqlensKV, cuSeqlensCmpKV); + } else if (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + InitActualSeqLen(seqUsedQ, cuSeqlensKV, cuSeqlensCmpKV); + } else if ((KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND || KV_LAYOUT_T == SMLA_LAYOUT::BSND) && + LAYOUT_T == SMLA_LAYOUT::TND) { + InitActualSeqLen(cuSeqlensQ, seqUsedKV, seqUsedCmpKV); + } else if ((KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND || KV_LAYOUT_T == SMLA_LAYOUT::BSND)) { + InitActualSeqLen(seqUsedQ, seqUsedKV, seqUsedCmpKV); + } + metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata); + InitCalcParamsEach(); + + pipe = tPipe; + // init global buffer + queryGm.SetGlobalBuffer((__gm__ Q_T *)query); + oriKvGm.SetGlobalBuffer((__gm__ KV_T *)oriKV); + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + cmpKvGm.SetGlobalBuffer((__gm__ KV_T *)cmpKV); + if (constInfo.cmpResidualKVSize != 0) { + cmpResidualKVGm.SetGlobalBuffer((__gm__ int32_t *)cmpResidualKV, constInfo.cmpResidualKVSize); + } + } + + if (sinks != nullptr) { + sinksGm.SetGlobalBuffer((__gm__ SINKS_T *)sinks); + } + + attentionOutGm.SetGlobalBuffer((__gm__ OUT_T *)attentionOut); + softmaxLseGm.SetGlobalBuffer((__gm__ T *)softmaxLse); + + if ASCEND_IS_AIV { + if (LAYOUT_T != SMLA_LAYOUT::TND) { + if (constInfo.needInit) { + InitOutputSingleCore(); + } + } + } + + if constexpr (PAGE_ATTENTION) { + oriBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)oriBlockTable); + if (constInfo.hasOriSparseIndices) { + oriSparseIndicesGm.SetGlobalBuffer((__gm__ int32_t *)oriSparseIndices); + } + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + cmpBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)cmpBlockTable); + } + } + + // workspace 内存排布 + // |Q--|mm1ResGm|vec1ResGm|mm2ResGm|vec2ResGm + // |Core0_Q1-Core0_Q2-Core1_Q1-Core1_Q2....Core32_Q1-Core32_Q2|Core0_mmRes + uint64_t offset = 0; + mm1ResGm.SetGlobalBuffer( + (__gm__ MM1_OUT_T *)(workspace + offset + + aiCoreIdx * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(MM1_OUT_T))); + offset += GetBlockNum() * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(MM1_OUT_T); + + vec1ResGm.SetGlobalBuffer( + (__gm__ Q_T *)(workspace + offset + aiCoreIdx * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(KV_T))); + offset += GetBlockNum() * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(KV_T); + + mm2ResGm.SetGlobalBuffer( + (__gm__ MM2_OUT_T *)(workspace + offset + + aiCoreIdx * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(MM2_OUT_T))); + offset += GetBlockNum() * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(MM2_OUT_T); + + vec2ResGm.SetGlobalBuffer( + (__gm__ T *)(workspace + offset + aiCoreIdx * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(T))); + offset += GetBlockNum() * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(T); + + if (constInfo.hasOriSparseIndices) { + kvMergeGm_.SetGlobalBuffer( + (__gm__ KV_T *)(workspace + offset + aiCoreIdx * 512 * 512 * MERGE_CACHE_GM_BUF_NUM * sizeof(KV_T))); + offset += static_cast(GetBlockNum()) * 512 * 512 * MERGE_CACHE_GM_BUF_NUM * sizeof(KV_T); + } + + if ASCEND_IS_AIV { + vectorBlock.InitParams(constInfo, tilingData); + if (constInfo.hasOriSparseIndices) { + vectorBlock.InitVec0GlobalTensor(kvMergeGm_, oriKvGm, oriBlockTableGm, oriSparseIndicesGm); + } + vectorBlock.InitVec1GlobalTensor(mm1ResGm, vec1ResGm, actualSeqLengthsQGm, actualSeqLengthsKVGm, sinksGm, + softmaxLseGm, oriSparseIndicesGm, oriTopkLengthGm); + vectorBlock.InitVec2GlobalTensor(accumOutGm, vec2ResGm, mm2ResGm, attentionOutGm); + } + + if ASCEND_IS_AIC { + cubeBlock.InitParams(constInfo); + cubeBlock.InitMm1GlobalTensor(queryGm, oriKvGm, cmpKvGm, mm1ResGm); + cubeBlock.InitMm2GlobalTensor(vec1ResGm, mm2ResGm, attentionOutGm); + cubeBlock.InitPageAttentionInfo(oriKvGm, kvMergeGm_, oriBlockTableGm, cmpBlockTableGm); + } + // 要在InitParams之后执行 + if (pipe != nullptr) { + InitBuffers(); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::InitCalcParamsEach() +{ + if (aiCoreIdx != 0) { + constInfo.bN2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_BN2_START_INDEX, false)); + constInfo.gS1Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_M_START_INDEX, false)); + constInfo.s2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_S2_START_INDEX, false)); + } + constInfo.bN2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_BN2_END_INDEX, false)); + constInfo.gS1End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_M_END_INDEX, false)); + constInfo.s2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_S2_END_INDEX, false)); +} + +template +__aicore__ inline void SparseFlashMlaSwa::CalcParams(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, + uint32_t s2LoopIdx, RunInfo &info) +{ + info.isValid = s2LoopIdx < tempLoopInfo.s2LoopTimes; + info.loop = loop; + info.cmpLoop = cmpLoop; + info.bIdx = tempLoopInfo.bIdx; + info.gS1Idx = tempLoopInfo.gS1Idx; + info.s1Idx = tempLoopInfo.gS1Idx / constInfo.gSize; + info.s2Idx = s2LoopIdx; + info.n2IdxReal = tempLoopInfo.n2Idx; + + info.curSInnerLoopTimes = tempLoopInfo.s2LoopTimes; + info.tndIsS2SplitCore = tempLoopInfo.tndIsS2SplitCore; + info.tndCoreStartKVSplitPos = tempLoopInfo.tndCoreStartKVSplitPos; + info.isBmm2Output = false; + info.actS1Size = tempLoopInfo.actS1Size; + info.oriDealSize = + tempLoopInfo.oriMaskRight + tempLoopInfo.s1StartIdx - tempLoopInfo.oriMaskLeft - tempLoopInfo.s1EndIdx; + info.cmpMaskRight = tempLoopInfo.cmpMaskRight; + + // M方向的尾块 + info.actMBaseSize = tempLoopInfo.mBasicSizeTail; + + if ASCEND_IS_AIV { + info.mSize = info.actMBaseSize; + info.mSizeV = (info.mSize <= 16) ? info.mSize : ((CeilDiv(info.mSize, 16) + 1) / 2 * 16); + info.mSizeVStart = 0; + if (tmpBlockIdx % 2 == 1) { + info.mSizeVStart = info.mSizeV; + info.mSizeV = info.mSize - info.mSizeV; + } + } + + info.isFirstSInnerLoop = s2LoopIdx == s2Start; + if (info.isFirstSInnerLoop) { + tempLoopInfo.bn2IdxInCurCore++; + } + info.isLastS2Loop = (s2LoopIdx == (tempLoopInfo.s2LoopTimes - 1)); + info.bn2IdxInCurCore = tempLoopInfo.bn2IdxInCurCore - 1; + + uint64_t tndBIdxOffsetForQ = tempLoopInfo.actualSeqQPrefixSum * constInfo.qHeadNum * constInfo.headDim; + uint64_t tndBIdxOffsetForKV = tempLoopInfo.actualSeqKVPrefixSum * constInfo.kvHeadNum * constInfo.headDim; + uint64_t tndBIdxOffsetForCmpKV = tempLoopInfo.actualSeqCmpKVPrefixSum * constInfo.kvHeadNum * constInfo.headDim; + + if (info.isFirstSInnerLoop) { + uint64_t s1HeadOffset = (info.gS1Idx / constInfo.gSize) * constInfo.qHeadNum; + uint64_t qHeadOffset = info.n2Idx * constInfo.gSize + info.gS1Idx % constInfo.gSize; + tensorACoreOffset = tndBIdxOffsetForQ + (s1HeadOffset + qHeadOffset) * constInfo.headDim; + tensorBCoreOffset = tndBIdxOffsetForKV + info.n2Idx * constInfo.headDim; + tensorCmpBCoreOffset = tndBIdxOffsetForCmpKV + info.n2Idx * constInfo.headDim; + } + info.tensorAOffset = tensorACoreOffset; + info.tensorBOffset = tensorBCoreOffset; + info.tensorCmpBOffset = tensorCmpBCoreOffset; + info.attenOutOffset = tensorACoreOffset; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + info.qTokenOffset = tempLoopInfo.actualSeqQPrefixSum + info.s1Idx; + } else { + info.qTokenOffset = info.bIdx * constInfo.qSeqSize + info.s1Idx; + } + + if constexpr (TEMPLATE_MODE == SWA_TEMPLATE) { + // SWA只有ori_kv + info.isOriOnly = true; + info.isOriCmpMix = false; + info.relativeS2Idx = 0; + uint64_t s2Offset = info.s2Idx * constInfo.s2BaseSize; + if (s2LoopIdx + 1 == tempLoopInfo.oriLoopTimes) { + info.actualSingleProcessSInnerOriSize = + (tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1) - s2Offset; + } else { + info.actualSingleProcessSInnerOriSize = constInfo.s2BaseSize; + } + info.actualSingleProcessSInnerSize = info.actualSingleProcessSInnerOriSize; + info.s2StartPoint = constInfo.hasOriSparseIndices ? 0 : tempLoopInfo.oriMaskLeft; + info.cmpS2IdLimit = 0; + } else { // HCA_TEMPLATE场景 + if (s2LoopIdx < tempLoopInfo.oriLoopTimes) { + // S2首次循环只能在ori_kv + info.isOriOnly = true; + info.isOriCmpMix = false; + info.relativeS2Idx = 0; + uint64_t s2Offset = info.s2Idx * constInfo.s2BaseSize; + if (s2LoopIdx + 1 == tempLoopInfo.oriLoopTimes) { // HCA场景可能处理oriLen / cmpRatio等于0的场景 + info.actualSingleProcessSInnerOriSize = + (tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1) - s2Offset; + } else { + info.actualSingleProcessSInnerOriSize = constInfo.s2BaseSize; + } + info.actualSingleProcessSInnerSize = info.actualSingleProcessSInnerOriSize; + info.s2StartPoint = tempLoopInfo.oriMaskLeft; + info.cmpS2IdLimit = 0; + } else if (s2LoopIdx - tempLoopInfo.oriLoopTimes < tempLoopInfo.oriCmpMixLoopTimes) { + uint64_t s2Offset = info.s2Idx * constInfo.s2BaseSize; + info.isOriOnly = false; + info.isOriCmpMix = true; + info.relativeS2Idx = info.s2Idx - tempLoopInfo.oriLoopTimes; + info.actualSingleProcessSInnerOriSize = + (tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1) - s2Offset; + info.s2StartPoint = tempLoopInfo.oriMaskLeft + tempLoopInfo.oriLoopTimes * constInfo.s2BaseSize; + info.actualSingleProcessSInnerCmpSize = tempLoopInfo.cmpMixSize; + info.cmpS2IdLimit = tempLoopInfo.cmpMixSize; + info.actualSingleProcessSInnerSize = + info.actualSingleProcessSInnerOriSize + info.actualSingleProcessSInnerCmpSize; + } else { + info.isOriOnly = false; + info.isOriCmpMix = false; + info.relativeS2Idx = info.s2Idx - tempLoopInfo.oriLoopTimes - tempLoopInfo.oriCmpMixLoopTimes; + uint64_t s2Offset = info.relativeS2Idx * constInfo.s2BaseSize + tempLoopInfo.cmpMixSize; + if (s2LoopIdx + 1 == tempLoopInfo.s2LoopTimes) { + info.actualSingleProcessSInnerSize = tempLoopInfo.actCmpS2Size - s2Offset; + } else { + info.actualSingleProcessSInnerSize = constInfo.s2BaseSize; + } + info.s2StartPoint = tempLoopInfo.cmpMixSize; + info.cmpS2IdLimit = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio; + } + } + + info.inValidRowCount = tempLoopInfo.actS1Size - tempLoopInfo.actOriS2Size - tempLoopInfo.s1StartIdx; + info.actualSingleProcessSInnerSizeAlign = + SMLAAlign(info.actualSingleProcessSInnerSize, SMLAVectorBlock::BYTE_BLOCK); + info.actualSingleProcessSInnerOriAlignSize = + SMLAAlign(info.actualSingleProcessSInnerOriSize, SMLAVectorBlock::BYTE_BLOCK); + info.actualSingleProcessSInnerCmpAlignSize = + SMLAAlign(info.actualSingleProcessSInnerCmpSize, SMLAVectorBlock::BYTE_BLOCK); + if (constInfo.hasOriSparseIndices && info.isValid) { + info.v0S2Start = 0; + info.v0S2DealSize = static_cast(info.actualSingleProcessSInnerOriSize); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::ComputeMm1(const RunInfo &info) +{ + uint32_t nBufferLoopTimes = CeilDiv(info.actMBaseSize, constInfo.nBufferMBaseSize); + uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize; + for (uint32_t i = 0; i < nBufferLoopTimes; i++) { + MSplitInfo mSplitInfo; + mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize; + mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail; + cubeBlock.ComputeMm1(info, mSplitInfo); + CrossCoreSetFlag(constInfo.syncC1V1); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::ComputeMm2(const RunInfo &info) +{ + uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize; + uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize; + for (uint32_t i = 0; i < nBufferLoopTimes; i++) { + MSplitInfo mSplitInfo; + mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize; + mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail; + CrossCoreWaitFlag(constInfo.syncV1C2); + cubeBlock.ComputeMm2(info, mSplitInfo); + CrossCoreSetFlag(constInfo.syncC2V2); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::Process() +{ + uint32_t hasLoad = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_CORE_ENABLE_INDEX, false)); + if (hasLoad == 0) { + return; + } + if ASCEND_IS_AIV { + vectorBlock.AllocEventID(); + vectorBlock.InitSoftmaxDefaultBuffer(); + } else { + cubeBlock.AllocEventID(); + } + ProcessBalance(); + if ASCEND_IS_AIV { + vectorBlock.FreeEventID(); + } else { + cubeBlock.FreeEventID(); + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::GetBN2Idx(uint32_t bN2Idx, uint32_t &bIdx, uint32_t &n2Idx) +{ + bIdx = bN2Idx / constInfo.kvHeadNum; + n2Idx = bN2Idx % constInfo.kvHeadNum; +} + +template +__aicore__ inline void SparseFlashMlaSwa::ProcessBalance() +{ + RunInfo extraInfo[SMLA_PRELOAD_TASK_CACHE_SIZE]; + uint32_t gloop = 0; + uint32_t cmpLoop = 0; + uint32_t gS1LoopEnd = 0; + bool globalLoopStart = true; + // 适配左闭右开 + if (constInfo.bN2Start == constInfo.bN2End) { + if (constInfo.gS1Start != constInfo.gS1End || constInfo.s2Start != constInfo.s2End) { + constInfo.bN2End += 1; + } + } else if ((constInfo.gS1End != 0) || (constInfo.s2End != 0)) { + constInfo.bN2End += 1; + } + for (uint32_t bN2LoopIdx = constInfo.bN2Start; bN2LoopIdx < constInfo.bN2End; bN2LoopIdx++) { + GetBN2Idx(bN2LoopIdx, tempLoopInfo.bIdx, tempLoopInfo.n2Idx); + tempLoopInfo.actS1Size = GetActualSeqLenQ(tempLoopInfo.bIdx); // 获取actualSeqLength + tempLoopInfo.actOriS2Size = GetActualSeqLenKV(tempLoopInfo.bIdx); // 获取actualSeqLengthKV + // 判断是否存在行无效场景,直接全部刷零(DSpark 走 sparse topk,不做 dense 行无效清零) + int32_t inValidRowCount = 0; + int32_t inValidRowS1StartIdx = 0; + if (!constInfo.hasOriSparseIndices && + inValidRowS1StartIdx < (tempLoopInfo.actS1Size - tempLoopInfo.actOriS2Size)) { + inValidRowCount = tempLoopInfo.actS1Size - tempLoopInfo.actOriS2Size; + tempLoopInfo.hasInvalidRow = true; + if ASCEND_IS_AIV { + InitAllZeroOutput(tempLoopInfo.bIdx, inValidRowS1StartIdx, inValidRowCount, tempLoopInfo.n2Idx); + } + } + bool isS1S2ZeroAndLastBatch = + (tempLoopInfo.actS1Size == 0 || tempLoopInfo.actOriS2Size == 0) && + ((constInfo.outputLayout == SMLA_LAYOUT::BSND) || (bN2LoopIdx + 1 == constInfo.bN2End)); + uint32_t gS1SplitNum = CeilDiv((tempLoopInfo.actS1Size - inValidRowCount) * constInfo.gSize, + constInfo.mBaseSize); // gS1轴上有效基本块数量 + + // 当处于最后一个BN2时, 且gS1End为0时, 说明当前BN2里的所有数据都在当前核处理 + gS1LoopEnd = (bN2LoopIdx == constInfo.bN2End - 1 && constInfo.gS1End != 0) ? constInfo.gS1End : gS1SplitNum; + // 当处于最后一个BN2且当前S1为0时,需要进入循环计算preload导致的未完成的部分 + gS1LoopEnd = isS1S2ZeroAndLastBatch ? gS1LoopEnd + 1 : gS1LoopEnd; + for (uint32_t gS1LoopIdx = constInfo.gS1Start; gS1LoopIdx < gS1LoopEnd; gS1LoopIdx++) { + tempLoopInfo.actOriS2Size = GetActualSeqLenKV(tempLoopInfo.bIdx); + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + tempLoopInfo.actCmpS2Size = GetActualSeqLenCmpKV(tempLoopInfo.bIdx, tempLoopInfo.actOriS2Size); + } + // 对于各轴上的真实的idx, 采用左闭右闭的方案 + // 跳过行无效部分,从有效行开始后续计算 + tempLoopInfo.gS1Idx = inValidRowCount * constInfo.gSize + gS1LoopIdx * constInfo.mBaseSize; + tempLoopInfo.s1StartIdx = tempLoopInfo.gS1Idx / constInfo.gSize; + tempLoopInfo.s1EndIdx = + Min((tempLoopInfo.s1StartIdx + constInfo.mBaseSize / constInfo.gSize - 1), tempLoopInfo.actS1Size - 1); + // 此处均为闭区间 + // oriMaskMode 0 (DSpark): effective ori S2 from sparse topk_length, not band window. + if (constInfo.hasOriSparseIndices) { + uint32_t sparseOriS2Size = GetOriSparseActualSeqLen(); + tempLoopInfo.oriMaskLeft = 0; + tempLoopInfo.oriMaskRight = sparseOriS2Size == 0 ? -1 : static_cast(sparseOriS2Size - 1U); + } else { + // oriMaskMode 4 (Band): sliding window on dense ori_kv. + tempLoopInfo.oriMaskRight = tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size + + static_cast(tempLoopInfo.s1EndIdx) + constInfo.oriWinRight; + tempLoopInfo.oriMaskRight = Min(tempLoopInfo.oriMaskRight, tempLoopInfo.actOriS2Size - 1); + tempLoopInfo.oriMaskLeft = Max(tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size + + static_cast(tempLoopInfo.s1StartIdx) - constInfo.oriWinLeft, + 0); + } + if constexpr (TEMPLATE_MODE == HCA_TEMPLATE) { + int32_t cmpMaskS2Size = + GetCmpMaskS2Size(tempLoopInfo.bIdx, tempLoopInfo.actOriS2Size, tempLoopInfo.actCmpS2Size); + tempLoopInfo.cmpMaskRight = cmpMaskS2Size - tempLoopInfo.actS1Size; + } + GetSparseActualSeqLen(); + UpdateInnerLoopCond(); + bool isEnd = (bN2LoopIdx + 1 == constInfo.bN2End) && (gS1LoopIdx + 1 == gS1LoopEnd); + if (tempLoopInfo.curActSeqLenIsZero && !isEnd) { + continue; + } + if constexpr (TEMPLATE_MODE == SWA_TEMPLATE) { + uint32_t oriS2Size = + (tempLoopInfo.oriMaskRight >= tempLoopInfo.oriMaskLeft) ? + static_cast(tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1) : + 0U; + tempLoopInfo.oriLoopTimes = CeilDiv(oriS2Size, constInfo.s2BaseSize); + tempLoopInfo.oriCmpMixLoopTimes = 0; + tempLoopInfo.cmpLoopTimes = 0; + tempLoopInfo.s2LoopTimes = tempLoopInfo.oriLoopTimes; + } else { // HCA_TEMPLATE + uint32_t oriLen = tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1; + if (tempLoopInfo.actCmpS2Size == 0) { // ori/cmp_ratio长度等于0的场景 + tempLoopInfo.oriLoopTimes = (oriLen + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize; + tempLoopInfo.oriCmpMixLoopTimes = 0; + tempLoopInfo.cmpLoopTimes = 0; + } else { + tempLoopInfo.oriLoopTimes = oriLen / constInfo.s2BaseSize; + tempLoopInfo.oriCmpMixLoopTimes = (oriLen % constInfo.s2BaseSize == 0) ? 0 : 1; + if (tempLoopInfo.oriCmpMixLoopTimes == 0) { + tempLoopInfo.cmpMixSize = 0; + } else { + uint32_t cmpLeftSize = + constInfo.s2BaseSize - (oriLen - tempLoopInfo.oriLoopTimes * constInfo.s2BaseSize); + tempLoopInfo.cmpMixSize = + (cmpLeftSize < tempLoopInfo.actCmpS2Size) ? cmpLeftSize : tempLoopInfo.actCmpS2Size; + } + tempLoopInfo.cmpLoopTimes = + CeilDiv(tempLoopInfo.actCmpS2Size - tempLoopInfo.cmpMixSize, constInfo.s2BaseSize); + } + tempLoopInfo.s2LoopTimes = + tempLoopInfo.oriLoopTimes + tempLoopInfo.oriCmpMixLoopTimes + tempLoopInfo.cmpLoopTimes; + } + + tempLoopInfo.tndIsS2SplitCore = false; // 当前不支持核间切S2 + tempLoopInfo.tndCoreStartKVSplitPos = 0; + uint32_t extraLoop = isEnd ? PRELOAD_NUM : 0; + + for (uint32_t s2LoopIdx = constInfo.s2Start; s2LoopIdx < (tempLoopInfo.s2LoopTimes + extraLoop); + s2LoopIdx++) { + PreloadPipeline(gloop, cmpLoop, constInfo.s2Start, s2LoopIdx, extraInfo); + ++gloop; + if (constInfo.hasOriSparseIndices && s2LoopIdx < tempLoopInfo.s2LoopTimes) { + ++cmpLoop; + } + } + globalLoopStart = false; + constInfo.s2Start = 0; + } + constInfo.gS1Start = 0; + tempLoopInfo.hasInvalidRow = false; + } +} + +template +__aicore__ inline void SparseFlashMlaSwa::PreloadPipeline(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, + uint64_t s2LoopIdx, + RunInfo extraInfo[SMLA_PRELOAD_TASK_CACHE_SIZE]) +{ + RunInfo &extraInfo0 = extraInfo[loop % SMLA_PRELOAD_TASK_CACHE_SIZE]; // 本轮任务 + RunInfo &extraInfo2 = extraInfo[(loop + 2) % SMLA_PRELOAD_TASK_CACHE_SIZE]; // 上一轮任务 + RunInfo &extraInfo1 = extraInfo[(loop + 1) % SMLA_PRELOAD_TASK_CACHE_SIZE]; // 上两轮任务 + + CalcParams(loop, cmpLoop, s2Start, s2LoopIdx, extraInfo0); + if (extraInfo0.isValid) { + if (constInfo.hasOriSparseIndices) { + if ASCEND_IS_AIV { + vectorBlock.ProcessVec0L(extraInfo0); + CrossCoreSetFlag(constInfo.syncV0C1); + } + } else if ASCEND_IS_AIC { + ComputeMm1(extraInfo0); + } + } + if (extraInfo2.isValid) { + if ASCEND_IS_AIV { + vectorBlock.ProcessVec1L(extraInfo2); + } + if ASCEND_IS_AIC { + ComputeMm2(extraInfo2); + } + } + if (extraInfo1.isValid) { + if ASCEND_IS_AIV { + vectorBlock.ProcessVec2L(extraInfo1); + } + extraInfo1.isValid = false; + } + if (extraInfo0.isValid && constInfo.hasOriSparseIndices) { + if ASCEND_IS_AIC { + CrossCoreWaitFlag(constInfo.syncV0C1); + ComputeMm1(extraInfo0); + } + } +} +} // namespace SMLAKernel +#endif // SPARSE_FLASH_MLA_SWA_KERNEL_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/buffers_policy_3buff_sfa.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/buffers_policy_3buff_sfa.h new file mode 100644 index 000000000000..75b9730b2492 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/buffers_policy_3buff_sfa.h @@ -0,0 +1,151 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. +*/ + +/*! + * \file buffer_policy_sfa.h + * \brief + */ +#ifndef BUFFER_POLICY_SFA_H +#define BUFFER_POLICY_SFA_H + +#if __has_include("../../../../common/op_kernel/buffers_policy.h") +#include "../../../../common/op_kernel/buffers_policy.h" +#elif __has_include("../../../common/buffers_policy.h") +#include "../../../common/buffers_policy.h" +#endif + +namespace fa_base_matmul { +// 申请3个buffer, 轮转 +template +class BuffersPolicy3buffSFA { +public: + __aicore__ inline void Init(BufferManager &bufferManager, uint32_t size, uint32_t aId = 0U, + uint32_t bId = 0U, uint32_t cId = 0U) + { + a_ = bufferManager.template AllocBuffer(size); + b_ = bufferManager.template AllocBuffer(size); + c_ = bufferManager.template AllocBuffer(size); + + if constexpr (idSource == IdSource::INTERNAL) { + a_.template Init(); + b_.template Init(); + c_.template Init(); + } else if constexpr (idSource == IdSource::EXTERNAL) { + a_.template Init(aId); + b_.template Init(bId); + c_.template Init(cId); + } + } + + __aicore__ inline void Uninit(BufferManager &bufferManager) + { + a_.template UnInit(); + b_.template UnInit(); + c_.template UnInit(); + + bufferManager.template FreeBuffer(a_); + bufferManager.template FreeBuffer(b_); + bufferManager.template FreeBuffer(c_); + } + + __aicore__ inline Buffer &Get() + { + if (flag1_ == 0) { + flag1_ = 1; + return a_; + } else if (flag1_ == 1) { + flag1_ = NUM_2; + return b_; + } else { + flag1_ = 0; + return c_; + } + } + + __aicore__ inline Buffer &Get(uint32_t id) + { + uint32_t flag = id % 3; + if (flag == 0) { + return a_; + } else if (flag == 1) { + return b_; + } else { + return c_; + } + } + + __aicore__ inline Buffer &GetVec() + { // mixcore architecture + if (flag1_vec1_ == 0) { + flag1_vec1_ = 1; + return a_; + } else if (flag1_vec1_ == 1) { + flag1_vec1_ = NUM_2; + return b_; + } else { + flag1_vec1_ = 0; + return c_; + } + } + + __aicore__ inline Buffer &GetCube() + { // mixcore architecture + if (flag1_bmm2_ == 0) { + flag1_bmm2_ = 1; + return a_; + } else if (flag1_bmm2_ == 1) { + flag1_bmm2_ = NUM_2; + return b_; + } else { + flag1_bmm2_ = 0; + return c_; + } + } + + // Q复用 + __aicore__ inline Buffer &GetPre() + { + if (flag1_ == 0) { + return c_; + } else if (flag1_ == 1) { + return a_; + } else { + return b_; + } + } + + // KV复用 + __aicore__ inline Buffer &GetReused() + { + if (flag2_ == 0) { + flag2_ = 1; + return a_; + } else if (flag2_ == 1) { + flag2_ = NUM_2; + return b_; + } else { + flag2_ = 0; + return c_; + } + } + +private: + Buffer a_; + Buffer b_; + Buffer c_; + uint32_t flag1_ = 0; + uint32_t flag1_vec1_ = 0; + uint32_t flag1_bmm2_ = 0; + uint32_t flag2_ = 0; +}; + +} // namespace fa_base_matmul +#endif diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/flash_decode.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/flash_decode.h new file mode 100644 index 000000000000..4f8d727a94d8 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/flash_decode.h @@ -0,0 +1,514 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file flash_decode.h + * \brief S2-split Flash Decode data-plane: staging layout, Vec1 LSE staging, + * Vec2 partial-O staging, and FD chunk reduction. + * + * GM staging layout (three contiguous regions): + * [partial O slots][max slots][sum slots] + * + * Per-slot sizes: + * partial O : stagingM * dAlign * sizeof(float) + * max : stagingM * broadcastElems * sizeof(float) (each row is broadcastElems identical floats) + * sum : stagingM * broadcastElems * sizeof(float) + * + * Numerics: + * Vec2 must write normalized p_i = sum(exp(score-m_i)*V) / s_i (after final division). + * FD Reduce computes M = max(m_i), G = sum(exp(m_i-M)*s_i), w_i = exp(m_i-M)*s_i/G, O = sum(w_i*p_i). + * + * Synchronization: + * - All participating cores must enter SyncAll() after FA completes before FD reads staging. + * - SyncAll() is owned by the operator kernel, not this file. + * - Event IDs (V_MTE3, MTE3_V, V_MTE2, MTE2_V) are passed in by the caller. + * + * FD_MAX_S2_SPLIT_NUM is the ordinary per-task staging limit. Batch-consistency paths may use a different + * number of staging slots; their host allocation and metadata must agree with workspaceNum. + */ +#ifndef SPARSE_FLASH_MLA_FLASH_DECODE_H +#define SPARSE_FLASH_MLA_FLASH_DECODE_H + +#include + +#if __has_include("../../../../common/op_kernel/arch35/vf/vf_flash_decode_arch35.h") +#include "../../../../common/op_kernel/arch35/vf/vf_flash_decode_arch35.h" +#elif __has_include("../../../common/arch35/vf/vf_flash_decode_arch35.h") +#include "../../../common/arch35/vf/vf_flash_decode_arch35.h" +#endif +#include "static_buffer.h" + +namespace AttentionCommon { + +constexpr int64_t FD_MAX_S2_SPLIT_NUM = 2U; +constexpr int64_t FD_INCREMENTAL_MERGE_INPUT_NUM = 2U; +constexpr int64_t FD_BROADCAST_ELEMS_PER_ROW = 8U; +constexpr int64_t FD_REDUCE_CHUNK_ROWS = 16U; +static constexpr int64_t FD_BUFFER_SIZE_BYTE_32B = 32; + +struct FdRunInfo { + bool coreEnable = false; + int64_t bn2Idx = 0; + int64_t mIdx = 0; + int64_t workspaceIdx = 0; + int64_t workspaceNum = 0; + int64_t mStartIdx = 0; + int64_t mNum = 0; +}; + +template +struct FdBuffers { + BufferType accumOut; + BufferType blockMax; + BufferType blockSum; + BufferType lseExp; + BufferType partialO; +}; + +template +__aicore__ inline void InitFDBuffers(const FdRunInfo &fdRunInfo, PipeType *tPipe, FdBuffers &buffers) +{ + tPipe->Reset(); + int64_t maxSumTotal = static_cast(fdRunInfo.workspaceNum) * FD_REDUCE_CHUNK_ROWS * + FD_BROADCAST_ELEMS_PER_ROW * sizeof(float); + int64_t lseExpSize = FD_REDUCE_CHUNK_ROWS * FD_BROADCAST_ELEMS_PER_ROW * sizeof(float); + int64_t accumOutSize = static_cast(fdRunInfo.mNum) * D_ALIGN * sizeof(T); + int64_t partialOSize = FD_REDUCE_CHUNK_ROWS * D_ALIGN * sizeof(T); + tPipe->InitBuffer(buffers.accumOut, accumOutSize); + tPipe->InitBuffer(buffers.blockMax, maxSumTotal); + tPipe->InitBuffer(buffers.blockSum, maxSumTotal); + tPipe->InitBuffer(buffers.lseExp, lseExpSize); + tPipe->InitBuffer(buffers.partialO, partialOSize); +} + +// 静态 tensor 版本的 FD buffer 初始化:不调用 tPipe->Reset(),从 ubBaseAddr 顺序排布。 +// FD 在主流程 SyncAll() 之后独立执行,可复用主流程 UB 地址空间。 +template +__aicore__ inline void InitFDBuffersStatic(const FdRunInfo &fdRunInfo, uint32_t ubBaseAddr, + FdBuffers> &buffers) +{ + uint32_t ubAddr = ubBaseAddr; + int64_t maxSumTotal = static_cast(fdRunInfo.workspaceNum) * FD_REDUCE_CHUNK_ROWS * + FD_BROADCAST_ELEMS_PER_ROW * sizeof(float); + int64_t lseExpSize = FD_REDUCE_CHUNK_ROWS * FD_BROADCAST_ELEMS_PER_ROW * sizeof(float); + int64_t accumOutSize = static_cast(fdRunInfo.mNum) * D_ALIGN * sizeof(T); + int64_t partialOSize = FD_REDUCE_CHUNK_ROWS * D_ALIGN * sizeof(T); + + buffers.accumOut = {LocalTensor(TPosition::VECIN, ubAddr, accumOutSize), 0}; + ubAddr += accumOutSize; + buffers.blockMax = {LocalTensor(TPosition::VECIN, ubAddr, maxSumTotal), 0}; + ubAddr += maxSumTotal; + buffers.blockSum = {LocalTensor(TPosition::VECIN, ubAddr, maxSumTotal), 0}; + ubAddr += maxSumTotal; + buffers.lseExp = {LocalTensor(TPosition::VECIN, ubAddr, lseExpSize), 0}; + ubAddr += lseExpSize; + buffers.partialO = {LocalTensor(TPosition::VECIN, ubAddr, partialOSize), 0}; +} + +// The three regions are contiguous: partial O, max, then sum. +// slotCount = maxSplits * physicalCoreSlots (already includes maxSplits). +// broadcastElems: each max/sum row is stored as broadcastElems identical floats (currently 8). +// chunkRows: FD reduction chunk width (currently 16). +struct S2SplitFdStagingLayout { + int64_t stagingM; + int64_t dAlign; + int64_t slotCount; + int64_t broadcastElems; + int64_t chunkRows; + + __aicore__ inline int64_t StagingAttenOutElems() const + { + return stagingM * dAlign; + } + + __aicore__ inline int64_t StagingMaxSumBytes() const + { + return stagingM * broadcastElems * sizeof(float); + } + + __aicore__ inline __gm__ uint8_t *AttenOutRegion(__gm__ uint8_t *base) const + { + return base; + } + + __aicore__ inline __gm__ uint8_t *MaxRegion(__gm__ uint8_t *base) const + { + return AttenOutRegion(base) + slotCount * StagingAttenOutElems() * sizeof(float); + } + + __aicore__ inline __gm__ uint8_t *SumRegion(__gm__ uint8_t *base) const + { + return MaxRegion(base) + slotCount * StagingMaxSumBytes(); + } +}; + +// Stage max/sum that are already stored as one broadcast block per row. +__aicore__ inline void StageBroadcastMaxSum(const S2SplitFdStagingLayout &layout, __gm__ uint8_t *stagingBase, + int64_t workspaceIdx, int64_t stagingMOffset, int64_t validRows, + LocalTensor &maxBroadcastUb, LocalTensor &sumBroadcastUb, + uint8_t vToMte3Id, uint8_t mte3ToVId) +{ + __gm__ uint8_t *maxRegion = layout.MaxRegion(stagingBase); + __gm__ uint8_t *sumRegion = layout.SumRegion(stagingBase); + GlobalTensor maxGm; + maxGm.SetGlobalBuffer((__gm__ float *)maxRegion); + GlobalTensor sumGm; + sumGm.SetGlobalBuffer((__gm__ float *)sumRegion); + int64_t maxSumBytes = layout.StagingMaxSumBytes(); + int64_t floatOffset = workspaceIdx * (maxSumBytes / sizeof(float)) + stagingMOffset * layout.broadcastElems; + DataCopyExtParams copyParams{static_cast(validRows), + static_cast(layout.broadcastElems * sizeof(float)), 0, 0, 0}; + SetFlag(vToMte3Id); + WaitFlag(vToMte3Id); + DataCopyPad(sumGm[floatOffset], sumBroadcastUb, copyParams); + DataCopyPad(maxGm[floatOffset], maxBroadcastUb, copyParams); + SetFlag(mte3ToVId); + WaitFlag(mte3ToVId); +} + +// Stage Vec1 max/sum to GM staging. +// tmpUb must hold at least 2 * stagingM * broadcastElems floats (max + sum broadcast blocks). +__aicore__ inline void StageVec1Lse(const S2SplitFdStagingLayout &layout, __gm__ uint8_t *stagingBase, + int64_t workspaceIdx, int64_t stagingMOffset, int64_t validRows, + LocalTensor &maxUb, LocalTensor &sumUb, LocalTensor &tmpUb, + uint8_t vToMte3Id, uint8_t mte3ToVId) +{ + LocalTensor tmpMaxBlockUb = tmpUb; + LocalTensor tmpSumBlockUb = tmpUb[1024 / sizeof(float)]; + int64_t mSizeAlign8 = (validRows + layout.broadcastElems - 1) / layout.broadcastElems * layout.broadcastElems; + int64_t brcbRepeat = mSizeAlign8 / layout.broadcastElems; + Brcb(tmpMaxBlockUb, maxUb, brcbRepeat, {1, static_cast(layout.broadcastElems)}); + Brcb(tmpSumBlockUb, sumUb, brcbRepeat, {1, static_cast(layout.broadcastElems)}); + StageBroadcastMaxSum(layout, stagingBase, workspaceIdx, stagingMOffset, validRows, tmpMaxBlockUb, tmpSumBlockUb, + vToMte3Id, mte3ToVId); +} + +// Stage Vec2 normalized partial O to GM staging. +// vec2ResUb must contain FP32 partial O after final division. +// stagingOut points to the base of the atten-out staging region; workspaceIdx offset is folded into the element offset. +template +__aicore__ inline void StageVec2PartialO(const S2SplitFdStagingLayout &layout, GlobalTensor &stagingOut, + int64_t workspaceIdx, int64_t stagingMOffset, int64_t validRows, + int64_t dValid, LocalTensor &vec2ResUb, uint8_t vToMte3Id, + uint8_t mte3ToVId) +{ + int64_t offset = workspaceIdx * layout.StagingAttenOutElems() + stagingMOffset * dValid; + SetFlag(vToMte3Id); + WaitFlag(vToMte3Id); + DataCopyExtParams outParams; + outParams.blockLen = dValid * sizeof(float); + outParams.srcStride = static_cast((layout.dAlign - dValid) >> 3); + outParams.dstStride = 0; + outParams.blockCount = validRows; + DataCopyPad(stagingOut[offset], vec2ResUb, outParams); +} + +// Stage normalized partial O and wait for MTE3 completion before the source UB can be reused. +template +__aicore__ inline void StageVec2PartialOAndWait(const S2SplitFdStagingLayout &layout, GlobalTensor &stagingOut, + int64_t workspaceIdx, int64_t stagingMOffset, int64_t validRows, + int64_t dValid, LocalTensor &vec2ResUb, uint8_t vToMte3Id, + uint8_t mte3ToVId) +{ + StageVec2PartialO(layout, stagingOut, workspaceIdx, stagingMOffset, validRows, dValid, vec2ResUb, vToMte3Id, + mte3ToVId); + SetFlag(mte3ToVId); + WaitFlag(mte3ToVId); +} + +// FD chunk reduction: read all splits from staging, compute weights, reduce partial O. +// D_ALIGN is a compile-time constant (e.g. 512) used by ReduceFinalRes_const_VF. +// workspaceNum must match the number of contiguous slots allocated by the host and emitted by metadata. +template +__aicore__ inline void ReduceWithLse(const S2SplitFdStagingLayout &layout, __gm__ uint8_t *stagingBase, + int64_t workspaceIdx, int64_t workspaceNum, int64_t fdMOffset, int64_t mNum, + int64_t dValid, LocalTensor &accumulatedO, LocalTensor &lseExpUb, + LocalTensor &blockMaxUb, LocalTensor &blockSumUb, + LocalTensor &partialOFp32, bool softmaxLseFlag, GlobalTensor &softmaxLseGm, + int64_t softmaxLseOffset, uint8_t vToMte2Id0, uint8_t vToMte2Id1, + uint8_t mte2ToVId, uint8_t vToMte3LseOutId, uint8_t mte3ToVLseOutId) +{ + constexpr int64_t outputElemsPerBlock = FD_BUFFER_SIZE_BYTE_32B / sizeof(float); + int64_t attenOutElems = layout.StagingAttenOutElems(); + int64_t maxSumBytes = layout.StagingMaxSumBytes(); + __gm__ uint8_t *maxRegion = layout.MaxRegion(stagingBase); + __gm__ uint8_t *sumRegion = layout.SumRegion(stagingBase); + + GlobalTensor maxGm; + maxGm.SetGlobalBuffer((__gm__ float *)(maxRegion + workspaceIdx * maxSumBytes)); + GlobalTensor sumGm; + sumGm.SetGlobalBuffer((__gm__ float *)(sumRegion + workspaceIdx * maxSumBytes)); + int64_t splitStride = maxSumBytes / sizeof(float); + GlobalTensor stagingOutGm; + stagingOutGm.SetGlobalBuffer( + (__gm__ float *)(layout.AttenOutRegion(stagingBase) + workspaceIdx * attenOutElems * sizeof(float))); + int64_t outSplitStride = attenOutElems; + LocalTensor sinkUb; + int64_t mChunks = (mNum + layout.chunkRows - 1) / layout.chunkRows; + int64_t startRow = 0; + for (int64_t chunkIdx = 0; chunkIdx < mChunks; chunkIdx++) { + int64_t dealRowCount = layout.chunkRows; + if (startRow + dealRowCount > mNum) { + dealRowCount = mNum - startRow; + } + WaitFlag(vToMte2Id0); + int64_t dealRowsElems = dealRowCount * layout.broadcastElems; + int64_t srcOffset = (fdMOffset + startRow) * layout.broadcastElems; + int64_t dstOffset = 0; + for (int64_t splitIdx = 0; splitIdx < workspaceNum; splitIdx++) { + DataCopy(blockMaxUb[dstOffset], maxGm[srcOffset], dealRowsElems); + DataCopy(blockSumUb[dstOffset], sumGm[srcOffset], dealRowsElems); + srcOffset += splitStride; + dstOffset += dealRowsElems; + } + SetFlag(mte2ToVId); + WaitFlag(mte2ToVId); + + FaVectorApi::ComputeScaleValue_VF(sinkUb, blockMaxUb, blockSumUb, lseExpUb, dealRowCount, workspaceNum, + softmaxLseFlag, false); + PipeBarrier(); + if (softmaxLseFlag) { + DataCopyExtParams lseParams; + lseParams.blockCount = static_cast(dealRowCount); + lseParams.blockLen = sizeof(float); + lseParams.srcStride = 0; + lseParams.dstStride = 0; + WaitFlag(mte3ToVLseOutId); + SetFlag(vToMte3LseOutId); + WaitFlag(vToMte3LseOutId); + DataCopyPad(softmaxLseGm[softmaxLseOffset + startRow], lseExpUb, lseParams); + SetFlag(mte3ToVLseOutId); + } + + LocalTensor chunkAccumO = accumulatedO[startRow * D_ALIGN]; + int64_t outSrcOffset = (fdMOffset + startRow) * dValid; + for (int64_t splitIdx = 0; splitIdx < workspaceNum; splitIdx++) { + DataCopyExtParams inParams; + inParams.blockLen = dValid * sizeof(float); + inParams.srcStride = 0; + inParams.dstStride = static_cast((D_ALIGN - dValid) / outputElemsPerBlock); + inParams.blockCount = dealRowCount; + DataCopyPadExtParams padParams{true, 0, + static_cast((D_ALIGN - dValid) % outputElemsPerBlock), 0}; + WaitFlag(vToMte2Id1); + DataCopyPad(partialOFp32, stagingOutGm[outSrcOffset], inParams, padParams); + SetFlag(mte2ToVId); + WaitFlag(mte2ToVId); + FaVectorApi::ReduceFinalRes_const_VF(chunkAccumO, blockSumUb, partialOFp32, dealRowCount, + splitIdx); + SetFlag(vToMte2Id1); + outSrcOffset += outSplitStride; + } + + SetFlag(vToMte2Id0); + startRow += layout.chunkRows; + } +} + +template +__aicore__ inline void Reduce(const S2SplitFdStagingLayout &layout, __gm__ uint8_t *stagingBase, int64_t workspaceIdx, + int64_t workspaceNum, int64_t fdMOffset, int64_t mNum, int64_t dValid, + LocalTensor &accumulatedO, LocalTensor &lseExpUb, + LocalTensor &blockMaxUb, LocalTensor &blockSumUb, + LocalTensor &partialOFp32, uint8_t vToMte2Id0, uint8_t vToMte2Id1, uint8_t mte2ToVId) +{ + GlobalTensor unusedSoftmaxLseGm; + ReduceWithLse(layout, stagingBase, workspaceIdx, workspaceNum, fdMOffset, mNum, dValid, accumulatedO, + lseExpUb, blockMaxUb, blockSumUb, partialOFp32, false, unusedSoftmaxLseGm, 0, vToMte2Id0, + vToMte2Id1, mte2ToVId, 0, 0); +} + +// Merge two normalized partial results. mergedO may alias leftO; rightO must remain +// readable until the second reduction finishes. The result is (LSE, 1, normalized O). +template +__aicore__ inline void MergeTwoInputsWithLse(const S2SplitFdStagingLayout &layout, int64_t mNum, + LocalTensor &mergedO, LocalTensor &leftO, LocalTensor &rightO, + LocalTensor &mergedLseUb, LocalTensor &mergedSumUb, + LocalTensor &maxUb, LocalTensor &sumUb, + LocalTensor &sinkUb) +{ + FaVectorApi::ComputeScaleValue_VF(sinkUb, maxUb, sumUb, mergedLseUb, mNum, FD_INCREMENTAL_MERGE_INPUT_NUM, + true, false); + FaVectorApi::ReduceFinalRes_const_VF(mergedO, sumUb, leftO, mNum, 0); + PipeBarrier(); + FaVectorApi::ReduceFinalRes_const_VF(mergedO, sumUb, rightO, mNum, 1); + PipeBarrier(); + Duplicate(mergedSumUb, static_cast(1.0), mNum * layout.broadcastElems); + PipeBarrier(); +} + +// Deterministic FD reduction: merge contiguous staging slots in a fixed left fold. +// Each step uses the same two-input UB reduction as the intra-core batch-consistency path: +// state = ((slot0 merge slot1) merge slot2) ... +// accumulatedO holds the left-fold state; partialOUb streams one new input at a time. +template +__aicore__ inline void ReducePairwiseWithLse(const S2SplitFdStagingLayout &layout, __gm__ uint8_t *stagingBase, + int64_t workspaceIdx, int64_t workspaceNum, int64_t fdMOffset, + int64_t mNum, int64_t dValid, LocalTensor &accumulatedO, + LocalTensor &lseExpUb, LocalTensor &blockMaxUb, + LocalTensor &blockSumUb, LocalTensor &partialOUb, + bool softmaxLseFlag, GlobalTensor &softmaxLseGm, + int64_t softmaxLseOffset, uint8_t vToMte2Id0, uint8_t vToMte2Id1, + uint8_t mte2ToVId, uint8_t vToMte3LseOutId, uint8_t mte3ToVLseOutId) +{ + if (workspaceNum <= FD_INCREMENTAL_MERGE_INPUT_NUM) { + ReduceWithLse(layout, stagingBase, workspaceIdx, workspaceNum, fdMOffset, mNum, dValid, + accumulatedO, lseExpUb, blockMaxUb, blockSumUb, partialOUb, softmaxLseFlag, + softmaxLseGm, softmaxLseOffset, vToMte2Id0, vToMte2Id1, mte2ToVId, vToMte3LseOutId, + mte3ToVLseOutId); + return; + } + + constexpr int64_t outputElemsPerBlock = FD_BUFFER_SIZE_BYTE_32B / sizeof(float); + int64_t attenOutElems = layout.StagingAttenOutElems(); + int64_t maxSumBytes = layout.StagingMaxSumBytes(); + GlobalTensor maxGm; + maxGm.SetGlobalBuffer((__gm__ float *)(layout.MaxRegion(stagingBase) + workspaceIdx * maxSumBytes)); + GlobalTensor sumGm; + sumGm.SetGlobalBuffer((__gm__ float *)(layout.SumRegion(stagingBase) + workspaceIdx * maxSumBytes)); + GlobalTensor stagingOutGm; + stagingOutGm.SetGlobalBuffer( + (__gm__ float *)(layout.AttenOutRegion(stagingBase) + workspaceIdx * attenOutElems * sizeof(float))); + int64_t splitStride = maxSumBytes / sizeof(float); + int64_t outSplitStride = attenOutElems; + + LocalTensor sinkUb; + int64_t mChunks = (mNum + layout.chunkRows - 1) / layout.chunkRows; + int64_t startRow = 0; + for (int64_t chunkIdx = 0; chunkIdx < mChunks; chunkIdx++) { + int64_t dealRowCount = layout.chunkRows; + if (startRow + dealRowCount > mNum) { + dealRowCount = mNum - startRow; + } + int64_t dealRowsElems = dealRowCount * layout.broadcastElems; + int64_t maxSumRowOffset = (fdMOffset + startRow) * layout.broadcastElems; + int64_t outRowOffset = (fdMOffset + startRow) * dValid; + DataCopyExtParams inParams; + inParams.blockLen = dValid * sizeof(float); + inParams.srcStride = 0; + inParams.dstStride = static_cast((D_ALIGN - dValid) / outputElemsPerBlock); + inParams.blockCount = dealRowCount; + DataCopyPadExtParams padParams{true, 0, static_cast((D_ALIGN - dValid) % outputElemsPerBlock), + 0}; + + // Seed the fold with slot0 and slot1. + if (softmaxLseFlag) { + WaitFlag(mte3ToVLseOutId); + } + WaitFlag(vToMte2Id0); + WaitFlag(vToMte2Id1); + DataCopy(blockMaxUb, maxGm[maxSumRowOffset], dealRowsElems); + DataCopy(blockSumUb, sumGm[maxSumRowOffset], dealRowsElems); + DataCopy(blockMaxUb[dealRowsElems], maxGm[maxSumRowOffset + splitStride], dealRowsElems); + DataCopy(blockSumUb[dealRowsElems], sumGm[maxSumRowOffset + splitStride], dealRowsElems); + LocalTensor chunkAccumO = accumulatedO[startRow * D_ALIGN]; + DataCopyPad(chunkAccumO, stagingOutGm[outRowOffset], inParams, padParams); + DataCopyPad(partialOUb, stagingOutGm[outRowOffset + outSplitStride], inParams, padParams); + SetFlag(mte2ToVId); + WaitFlag(mte2ToVId); + + MergeTwoInputsWithLse(layout, dealRowCount, chunkAccumO, chunkAccumO, partialOUb, blockMaxUb, + blockSumUb, blockMaxUb, blockSumUb, sinkUb); + + // Convert the accumulated state to (LSE, 1, normalized O), then merge one slot at a time. + for (int64_t splitIdx = FD_INCREMENTAL_MERGE_INPUT_NUM; splitIdx < workspaceNum; splitIdx++) { + SetFlag(vToMte2Id0); + SetFlag(vToMte2Id1); + WaitFlag(vToMte2Id0); + WaitFlag(vToMte2Id1); + + int64_t maxSumOffset = maxSumRowOffset + splitIdx * splitStride; + DataCopy(blockMaxUb[dealRowsElems], maxGm[maxSumOffset], dealRowsElems); + DataCopy(blockSumUb[dealRowsElems], sumGm[maxSumOffset], dealRowsElems); + DataCopyPad(partialOUb, stagingOutGm[outRowOffset + splitIdx * outSplitStride], inParams, padParams); + SetFlag(mte2ToVId); + WaitFlag(mte2ToVId); + MergeTwoInputsWithLse(layout, dealRowCount, chunkAccumO, chunkAccumO, partialOUb, blockMaxUb, + blockSumUb, blockMaxUb, blockSumUb, sinkUb); + } + + if (softmaxLseFlag) { + DataCopyExtParams lseParams; + lseParams.blockCount = static_cast(dealRowCount); + lseParams.blockLen = sizeof(float); + lseParams.srcStride = 0; + lseParams.dstStride = 0; + SetFlag(vToMte3LseOutId); + WaitFlag(vToMte3LseOutId); + DataCopyPad(softmaxLseGm[softmaxLseOffset + startRow], blockMaxUb, lseParams); + SetFlag(mte3ToVLseOutId); + } + + SetFlag(vToMte2Id0); + SetFlag(vToMte2Id1); + startRow += layout.chunkRows; + } +} + +// Merge one previously staged result with one current normalized result. +// dealRowCount must not exceed layout.chunkRows. +// currentMaxUb/currentSumUb contain one scalar per row; currentPartialO uses D_ALIGN elements per row. +// blockMaxUb/blockSumUb layout: [previous, current][dealRowCount][broadcastElems]. +// partialOTmpUb holds the previous normalized O and then the merged O. +// The merged state is returned as normalized partial O plus (LSE, 1), so it can be staged and merged again. +// Both V_MTE2 events must be set before the first call. This function restores them before returning. +template +__aicore__ inline void MergeStagedAndCurrentChunk( + const S2SplitFdStagingLayout &layout, __gm__ uint8_t *stagingBase, int64_t workspaceIdx, int64_t stagingMOffset, + int64_t dealRowCount, int64_t dValid, LocalTensor ¤tMaxUb, LocalTensor ¤tSumUb, + LocalTensor ¤tPartialO, LocalTensor &blockMaxUb, LocalTensor &blockSumUb, + LocalTensor &partialOTmpUb, LocalTensor &mergedLseBroadcastUb, LocalTensor &mergedSumBroadcastUb, + LocalTensor &sinkUb, uint8_t maxSumVToMte2Id, uint8_t partialOVToMte2Id, uint8_t mte2ToVId) +{ + constexpr int64_t outputElemsPerBlock = FD_BUFFER_SIZE_BYTE_32B / sizeof(float); + int64_t dealRowsElems = dealRowCount * layout.broadcastElems; + int64_t brcbRepeat = (dealRowCount + layout.broadcastElems - 1) / layout.broadcastElems; + Brcb(blockMaxUb[dealRowsElems], currentMaxUb, brcbRepeat, {1, static_cast(layout.broadcastElems)}); + Brcb(blockSumUb[dealRowsElems], currentSumUb, brcbRepeat, {1, static_cast(layout.broadcastElems)}); + + int64_t maxSumBytes = layout.StagingMaxSumBytes(); + GlobalTensor previousMaxGm; + previousMaxGm.SetGlobalBuffer((__gm__ float *)(layout.MaxRegion(stagingBase) + workspaceIdx * maxSumBytes)); + GlobalTensor previousSumGm; + previousSumGm.SetGlobalBuffer((__gm__ float *)(layout.SumRegion(stagingBase) + workspaceIdx * maxSumBytes)); + int64_t maxSumOffset = stagingMOffset * layout.broadcastElems; + WaitFlag(maxSumVToMte2Id); + DataCopy(blockMaxUb, previousMaxGm[maxSumOffset], dealRowsElems); + DataCopy(blockSumUb, previousSumGm[maxSumOffset], dealRowsElems); + SetFlag(mte2ToVId); + WaitFlag(mte2ToVId); + + GlobalTensor previousPartialOGm; + previousPartialOGm.SetGlobalBuffer((__gm__ float *)(layout.AttenOutRegion(stagingBase) + + workspaceIdx * layout.StagingAttenOutElems() * sizeof(float))); + int64_t partialOOffset = stagingMOffset * dValid; + DataCopyExtParams inParams; + inParams.blockLen = static_cast(dValid * sizeof(float)); + inParams.srcStride = 0; + inParams.dstStride = static_cast((D_ALIGN - dValid) / outputElemsPerBlock); + inParams.blockCount = static_cast(dealRowCount); + DataCopyPadExtParams padParams{true, 0, static_cast((D_ALIGN - dValid) % outputElemsPerBlock), 0}; + WaitFlag(partialOVToMte2Id); + DataCopyPad(partialOTmpUb, previousPartialOGm[partialOOffset], inParams, padParams); + SetFlag(mte2ToVId); + WaitFlag(mte2ToVId); + + DataCopy(partialOTmpUb[dealRowCount * D_ALIGN], currentPartialO, dealRowCount * D_ALIGN); + PipeBarrier(); + LocalTensor currentPartialOTmp = partialOTmpUb[dealRowCount * D_ALIGN]; + MergeTwoInputsWithLse(layout, dealRowCount, currentPartialO, partialOTmpUb, currentPartialOTmp, + mergedLseBroadcastUb, mergedSumBroadcastUb, blockMaxUb, blockSumUb, sinkUb); + SetFlag(partialOVToMte2Id); + SetFlag(maxSumVToMte2Id); +} + +} // namespace AttentionCommon + +#endif // SPARSE_FLASH_MLA_FLASH_DECODE_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/static_buffer.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/static_buffer.h new file mode 100644 index 000000000000..80de3a4b698c --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/static_buffer.h @@ -0,0 +1,61 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file static_buffer.h + * \brief 静态 tensor buffer 管理:StaticBuffer 携带显式 bufferId,RingBuffer 提供轮转。 + * 与 TPipe 管理的 TBuf/Buffer 不同,本文件的 buffer 直接从首地址偏移排布, + * 由使用者手动指定地址与 bufferId。 + */ +#ifndef SPARSE_FLASH_MLA_STATIC_BUFFER_H +#define SPARSE_FLASH_MLA_STATIC_BUFFER_H + +#include +#if ASC_DEVKIT_MAJOR >= 9 +#include "kernel_basic_intf.h" +#else +#include "kernel_operator.h" +#endif +using namespace AscendC; + +namespace fa_base_matmul { + +template +struct StaticBuffer { + LocalTensor tensor; + uint32_t idx; +}; + +template +struct RingBuffer { + StaticBuffer *bufs; + uint32_t bufNum; + uint32_t curId; + + __aicore__ inline RingBuffer() + : bufs(nullptr), + bufNum(0), + curId(0) + {} + __aicore__ inline RingBuffer(StaticBuffer *b, uint32_t n) + : bufs(b), + bufNum(n), + curId(n - 1) + {} + + __aicore__ inline StaticBuffer &GetNext() + { + curId = (curId + 1) % bufNum; + return bufs[curId]; + } +}; + +} // namespace fa_base_matmul +#endif // SPARSE_FLASH_MLA_STATIC_BUFFER_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/static_matmul.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/static_matmul.h new file mode 100644 index 000000000000..a005d0d40f00 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/common/static_matmul.h @@ -0,0 +1,158 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file matmul_static.h + * \brief 静态 tensor 版本的 MatmulK/MatmulN。 + * 与 common/op_kernel/matmul.h 中的 MatmulK/MatmulN 等价,但 L0A/L0B 通过 RingBuffer + * 管理,且核内同步使用显式 bufferId 的 SetFlag/WaitFlag(M_MTE1 / MTE1_M), + * 不再依赖 BuffersPolicyDB + AllocEventID。 + * L0A/L0B 合并为单个 bufferId(INNERCORE_L0AB 语义),因为 A/B 总是成对加载、 + * 被同一个 M 管道消费。 + */ +#ifndef SPARSE_FLASH_MLA_MATMUL_STATIC_H +#define SPARSE_FLASH_MLA_MATMUL_STATIC_H + +#if __has_include("../../../../common/op_kernel/matmul.h") +#include "../../../../common/op_kernel/matmul.h" +#elif __has_include("../../../common/op_kernel/matmul.h") +#include "../../../common/op_kernel/matmul.h" +#elif __has_include("../../common/op_kernel/matmul.h") +#include "../../common/op_kernel/matmul.h" +#else +#include "../common/matmul.h" +#endif +#include "static_buffer.h" +using namespace AscendC; + +// L0A/L0B 合并槽位对应的 flag id 映射,默认直接用槽位号 (0/1)。 +// 若使用方需要将 M_MTE1/MTE1_M 的 flag id 偏移,可在 include 前覆盖此宏。 +#ifndef MATMUL_STATIC_L0AB_ID +#define MATMUL_STATIC_L0AB_ID(s) (s) +#endif + +namespace fa_base_matmul { + +// 切K +template +__aicore__ inline void MatmulKStatic(const LocalTensor &aL1Tensor, const LocalTensor &bL1Tensor, + RingBuffer &l0A, RingBuffer &l0B, + const LocalTensor &cL0Tensor, const MMParam ¶m, + const LocalTensor &aScaleL1Tensor = LocalTensor(), + const LocalTensor &bScaleL1Tensor = LocalTensor()) +{ + uint32_t kLoops = (param.singleK + baseK - 1) / baseK; + uint32_t tailSize = param.singleK % baseK; + uint32_t tailK = tailSize ? tailSize : baseK; + uint64_t L1Aoffset = param.isLeftTranspose ? baseK << 4 : ((param.singleM + 15) >> 4 << 4) * baseK; + uint64_t L1Boffset = param.isRightTranspose ? ((param.singleN + 15) >> 4 << 4) * baseK : baseK << 4; + + for (uint32_t k = 0; k < kLoops; k++) { + uint32_t tileK = (k == (kLoops - 1)) ? tailK : baseK; + + StaticBuffer &aBuf = l0A.GetNext(); + StaticBuffer &bBuf = l0B.GetNext(); + + WaitFlag(MATMUL_STATIC_L0AB_ID(aBuf.idx)); // 等 M 用完该 AB 槽 + LocalTensor L0ATensor = aBuf.tensor.template ReinterpretCast(); + LoadDataToL0A(L0ATensor, aL1Tensor, param, k * L1Aoffset, tileK, param.singleM); + + LocalTensor L0BTensor = bBuf.tensor.template ReinterpretCast(); + uint64_t loopNum = param.isRightTranspose ? 1 : kLoops; + LoadDataToL0B(L0BTensor, bL1Tensor, param, k * L1Boffset, tileK, param.singleN, loopNum); + + SetFlag(MATMUL_STATIC_L0AB_ID(aBuf.idx)); // MTE1 搬运完,通知 M + WaitFlag(MATMUL_STATIC_L0AB_ID(aBuf.idx)); // M 等数据就绪 + + MmadParams mmadParams; + mmadParams.m = param.singleM; + if (param.realM != 0) { + mmadParams.m = param.realM; + } + mmadParams.n = param.singleN; + mmadParams.k = tileK; + if (mmadParams.m == 1) { + mmadParams.m = 16; + } + mmadParams.cmatrixInitVal = param.isOutKFisrt && (k == 0); + mmadParams.cmatrixSource = false; + if (param.unitFlag != 0) { + mmadParams.unitFlag = (param.unitFlag == UNITFLAG_EN_OUTER_LAST) && (k == kLoops - 1) ? + UNITFLAG_EN_OUTER_LAST : + UNITFLAG_ENABLE; + } + Mmad(cL0Tensor, L0ATensor, L0BTensor, mmadParams); + + SetFlag(MATMUL_STATIC_L0AB_ID(aBuf.idx)); // M 用完,释放该 AB 槽 + } +} + +// 切N +template +__aicore__ inline void MatmulNStatic(const LocalTensor &aL1Tensor, const LocalTensor &bL1Tensor, + RingBuffer &l0A, RingBuffer &l0B, + const LocalTensor &cL0Tensor, const MMParam ¶m, + const LocalTensor &aScaleL1Tensor = LocalTensor(), + const LocalTensor &bScaleL1Tensor = LocalTensor()) +{ + uint32_t nLoops = (param.singleN + baseN - 1) / baseN; + uint32_t tailSize = param.singleN % baseN; + uint32_t tailN = tailSize ? tailSize : baseN; + uint64_t L1Boffset = param.isRightTranspose ? (baseN << 4) : ((param.singleK + 15) >> 4 << 4) * baseN; + uint64_t L0Coffset = ((param.singleM + 15) >> 4 << 4) * baseN; + if (param.realM != 0) { + L0Coffset = ((param.realM + 15) >> 4 << 4) * baseN; + } + + MmadParams mmadParams; + mmadParams.m = param.singleM; + if (param.realM != 0) { + mmadParams.m = param.realM; + } + mmadParams.k = param.singleK; + if (mmadParams.m == 1) { + mmadParams.m = FP16_ONE_FRACTAL_ELEMENT; + } + mmadParams.cmatrixInitVal = param.isOutKFisrt; + mmadParams.cmatrixSource = false; + mmadParams.unitFlag = param.unitFlag; + + for (uint32_t n = 0; n < nLoops; n++) { + uint32_t tileN = (n == (nLoops - 1)) ? tailN : baseN; + mmadParams.n = tileN; + + StaticBuffer &aBuf = l0A.GetNext(); + StaticBuffer &bBuf = l0B.GetNext(); + + WaitFlag(MATMUL_STATIC_L0AB_ID(aBuf.idx)); // 等 M 用完该 AB 槽 + if (n == 0 || n == 1) { // 每个 ping-pong slot 首次访问各装一份 A, 之后复用不再覆盖 + LocalTensor L0ATensor = aBuf.tensor.template ReinterpretCast(); + LoadDataToL0A(L0ATensor, aL1Tensor, param, 0, param.singleK, param.singleM); + } + + LocalTensor L0BTensor = bBuf.tensor.template ReinterpretCast(); + uint64_t loopNum = param.isRightTranspose ? nLoops : 1; + LoadDataToL0B(L0BTensor, bL1Tensor, param, n * L1Boffset, param.singleK, tileN, loopNum); + + SetFlag(MATMUL_STATIC_L0AB_ID(aBuf.idx)); // MTE1 搬运完,通知 M + WaitFlag(MATMUL_STATIC_L0AB_ID(aBuf.idx)); // M 等数据就绪 + + Mmad(cL0Tensor[n * L0Coffset], aBuf.tensor, bBuf.tensor, mmadParams); + + SetFlag(MATMUL_STATIC_L0AB_ID(aBuf.idx)); // M 用完,释放该 AB 槽 + } +} + +} // namespace fa_base_matmul +#endif // SPARSE_FLASH_MLA_MATMUL_STATIC_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_common_arch35.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_common_arch35.h new file mode 100644 index 000000000000..47a6ea78afe4 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_common_arch35.h @@ -0,0 +1,159 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_common_arch35.h + * \brief + */ +#ifndef SPARSE_FLASH_MLA_COMMON_ARCH35_H +#define SPARSE_FLASH_MLA_COMMON_ARCH35_H +#include +#include "kernel_tiling/kernel_tiling.h" +#include "../sparse_flash_mla_common.h" +#if __has_include("common/static_buffer.h") +#include "common/static_buffer.h" +#endif + +constexpr uint64_t BLOCK_BYTE = 32; +constexpr uint32_t NEGATIVE_MIN_VAULE_FP32 = 0xFF7FFFFF; + +// ===== C 侧 buffer 元素个数 (tensor 偏移/地址递增用) ===== +constexpr uint32_t L1Q_ELEM_PER_BUF = 16384; // Q_T, 32KB +constexpr uint32_t L1_RIGHT_ELEM_PER_BLOCK = 65536; // Q_T, 128KB +constexpr uint32_t L0A_ELEM_PER_BUF = 8192; // Q_T, 16KB +constexpr uint32_t L0B_ELEM_PER_BUF = 16384; // Q_T, 32KB +constexpr uint32_t L0C_ELEM_PER_BUF = 32768; // T , 128KB + +// ===== C 侧核内 flag id (各 HardEvent 命名空间独立) ===== +#define INNERCORE_L0AB(s) (s) // 0,1 M_MTE1 / MTE1_M +#define INNERCORE_L0C(s) (s) // 0,1 FIX_M / M_FIX +#define INNERCORE_L1Q(s) (s) // 0,1,2 MTE1_MTE2 / MTE2_MTE1 +#define INNERCORE_L1KV(s) (3 + (s)) // 3,4,5 MTE1_MTE2 / MTE2_MTE1 + +// ===== V 侧主流程 ===== +#define INNERCORE_STAGE1(s) (4 + (s)) // 4,5 V_MTE3 / MTE3_V (stage1->L1) +#define INNERCORE_STAGE2 (6) // 6 V_MTE3 / MTE3_V (vec2结果+attentionOut拷出+staging, 串行复用) +#define INNERCORE_STAGE0OUT_MTE3_MTE2(s) (s) // 0,1 Vec0 stage0OutBuf (原 mte3ToMte2) +#define INNERCORE_STAGE0OUT_MTE2_MTE3(s) (s) // 0,1 Vec0 stage0OutBuf (原 mte2ToMte3) +#define INNERCORE_SINKS_SYNC (7) // 7 MTE2_V / V_MTE2 + +// ===== GetKVPhyAddr 独立相位 (保留原值; 与主流程相位分离, 可复用) ===== +#define INNERCORE_PHYADDR_BLKTABLE_FREE (3) // V_MTE2 +#define INNERCORE_PHYADDR_BLKTABLE_READY (8) // MTE2_V +#define INNERCORE_PHYADDR_SPARSEIDX_FREE (4) // V_MTE2 +#define INNERCORE_PHYADDR_SPARSEIDX_READY (6) // MTE2_V +#define INNERCORE_PHYADDR_KVADDR_READY (5) // V_MTE3 +#define INNERCORE_PHYADDR_KVADDR_FREE (7) // MTE3_V + +// ===== batch-consistency / LSE / FD / init ===== +#define INNERCORE_REDUCE_MAXSUM_V_MTE2 (2) // V_MTE2 +#define INNERCORE_INTRAPARTIALO_V_MTE2 (3) // V_MTE2 +#define INNERCORE_FD_V_MTE2(s) (4 + (s)) // 4,5 V_MTE2 +#define INNERCORE_REDUCE_MTE2_V (2) // MTE2_V +#define INNERCORE_FD_MTE2_V (3) // MTE2_V +#define INNERCORE_LSE_V_MTE3 (1) // V_MTE3 +#define INNERCORE_STAGE_FD_MTE3_V (7) // MTE3_V (Stage* staging) +#define INNERCORE_LSE_MTE3_V (0) // MTE3_V +#define INNERCORE_FD_MTE3_V (1) // MTE3_V +#define INNERCORE_INITOUT_MTE3_V (0) // MTE3_V (init, 与 LSE 相位分离) +#define INNERCORE_INTRALSE_MTE3_MTE2(s) (2 + (s)) // 2,3 +#define INNERCORE_INTRAATTN_MTE3_MTE2(s) (4 + (s)) // 4,5 +#define INNERCORE_FD_MTE3_MTE2 (6) + +// ===== 跨核 flag id (mode 4) ===== +#define CROSSCORE_L1P(s) (s) // 0,1 +#define CROSSCORE_BMM2 (2) +#define CROSSCORE_BMM1(s) (3 + (s)) // 3,4 +#define CROSSCORE_V0RES(s) (5 + (s)) // 5,6,7 (仅 CSA kernel) + +constexpr uint32_t L0AB_SHARED_SIZE_64K = 65536; // 65536表示64*1024 +constexpr uint32_t L0C_SHARED_SIZE_256K = 262144; // 262144表示256 * 1024 + +constexpr uint32_t BUFFER_SIZE_8K = 8192; // 8192表示8 * 1024 +constexpr uint32_t BUFFER_SIZE_16K = 16384; // 16384表示16 * 1024 +constexpr uint32_t BUFFER_SIZE_32K = 32768; // 32768表示32 * 1024 +constexpr uint32_t BUFFER_SIZE_64K = 65536; // 65536表示64 * 1024 +constexpr uint32_t BUFFER_SIZE_96K = 98304; // 98304表示96 * 1024 +constexpr uint32_t BUFFER_SIZE_128K = 131072; // 131072表示128 * 1024 +constexpr uint32_t BUFFER_SIZE_256K = 262144; // 262144表示256 * 1024 + +constexpr uint32_t CV_RATIO = 2; +constexpr uint64_t SYNC_MODE = 4; +constexpr uint32_t BATCH_CONSISTENCY_MAX_REDUCE_BLOCK_NUM = 33U; + +namespace SMLAKernel { +__aicore__ constexpr uint64_t Align2Func(uint64_t data) +{ + return (data + 1UL) >> 1UL << 1UL; // 向上2对齐, +1移位2 +} + +__aicore__ constexpr uint64_t Align8Func(uint64_t data) +{ + return (data + 7UL) >> 3UL << 3UL; // 向上8对齐, +7移位3 +} + +__aicore__ constexpr uint64_t Align16Func(uint64_t data) +{ + return (data + 15UL) >> 4UL << 4UL; // 向上16对齐, +15移位4 +} + +__aicore__ constexpr uint64_t Align64Func(uint64_t data) +{ + return (data + 63UL) >> 6UL << 6UL; // 向上64对齐, +63移位6 +} +} // namespace SMLAKernel + +#define TEMPLATE_INTF \ + template + +#define TEMPLATE_INTF_ARGS \ + Q_T, KV_T, T, OUTPUT_T, IS_FD, LAYOUT_T, KV_LAYOUT_T, TEMPLATE_MODE, IS_SPLIT_G, IS_BATCH_CONSISTENCY, \ + IS_VEC_S2PHYADDR + +#define CUBE_BLOCK_TRAITS_TYPE_FIELDS(X) \ + X(Q_T) \ + X(KV_T) \ + X(T) \ + X(OUTPUT_T) + +#define CUBE_BLOCK_TRAITS_CONST_FIELDS(X) \ + X(IS_FD, bool, false) \ + X(LAYOUT_T, SMLA_LAYOUT, SMLA_LAYOUT::BSND) \ + X(KV_LAYOUT_T, SMLA_LAYOUT, SMLA_LAYOUT::PA_BBND) \ + X(TEMPLATE_MODE, SMLATemplateMode, SMLATemplateMode::CSA_TEMPLATE_MODE) \ + X(IS_SPLIT_G, bool, false) \ + X(IS_BATCH_CONSISTENCY, bool, false) \ + X(IS_VEC_S2PHYADDR, bool, false) + +/* 1. 生成带默认值的模版Template */ +#define GEN_TYPE_PARAM(name) typename name, +#define GEN_CONST_PARAM(name, type, default_val) type name = default_val, + +#define TEMPLATES_DEF \ + template + +/* 2. 生成不带带默认值的模版Template */ +#define GEN_TEMPLATE_TYPE_NODEF(name) typename name, +#define GEN_TEMPLATE_CONST_NODEF(name, type, default_val) type name, +#define TEMPLATES_DEF_NO_DEFAULT \ + template + +/* 3. 生成有默认值的Args */ +#define GEN_ARG_NAME(name, ...) name, +#define TEMPLATE_ARGS \ + CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARG_NAME) \ + CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARG_NAME) \ + end + +#endif diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_block_cube_arch35.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_block_cube_arch35.h new file mode 100644 index 000000000000..bb8e3b608847 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_block_cube_arch35.h @@ -0,0 +1,608 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_csa_block_cube_arch35.h + * \brief + */ +#ifndef SPARSE_FLASH_MLA_CSA_BLOCK_CUBE_H__ARCH35 +#define SPARSE_FLASH_MLA_CSA_BLOCK_CUBE_H__ARCH35 +#include "kernel_operator_list_tensor_intf.h" +#include "util_regbase.h" +#include "sparse_flash_mla_common_arch35.h" +#include "common/static_matmul.h" + +#if __has_include("../../common/op_kernel/offset_calculator.h") +#include "../../common/op_kernel/offset_calculator.h" +#else +#include "../common/offset_calculator.h" +#endif +#if __has_include("../../common/op_kernel/matmul.h") +#include "../../common/op_kernel/matmul.h" +#else +#include "../common/matmul.h" +#endif +#if __has_include("../../common/op_kernel/FixpipeOut.h") +#include "../../common/op_kernel/FixpipeOut.h" +#else +#include "../common/FixpipeOut.h" +#endif +#if __has_include("../../common/op_kernel/CopyInL1.h") +#include "../../common/op_kernel/CopyInL1.h" +#else +#include "../common/CopyInL1.h" +#endif + +using namespace AscendC; +using namespace AscendC::Impl::Detail; +using namespace regbaseutil; +using namespace fa_base_matmul; +namespace SMLAKernel { +struct CubeCoordInfo { + uint32_t curBIdx; + uint32_t s1Coord; + uint32_t s2Coord; +}; + +template +__aicore__ inline constexpr GmFormat GetQueryGmFormat() +{ + if constexpr (LAYOUT == SMLA_LAYOUT::BSND) { + return GmFormat::BSNGD; + } else { + return GmFormat::TNGD; + } +} + +template +__aicore__ inline constexpr GmFormat GetKvGmFormat() +{ + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + return GmFormat::PA_BnBsND; + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + return GmFormat::TND; + } else { // BSND + return GmFormat::BSND; + } +} + +TEMPLATES_DEF +class CSABlockCube { +public: + /* =================编译期常量的基本块信息================= */ + static constexpr uint32_t s1BaseSize = 64; + static constexpr uint32_t s2BaseSize = 128; + static constexpr uint32_t dBaseSize = 512; + static constexpr uint32_t dBaseMatmulSize = 128; + static constexpr uint32_t rightBufNum = 3; + static constexpr uint32_t rightBufSingleSize = s2BaseSize * dBaseSize; + static constexpr uint32_t rightBufTotalSize = rightBufSingleSize * rightBufNum; + static constexpr uint32_t l1QBufNum = 3; // L1 Q 三缓冲 + static constexpr uint32_t l1KBufNum = 3; // L1 K 三缓冲 + static constexpr uint32_t qHalfNum = 2; // Q 沿 d 轴切半 + static constexpr uint32_t crossCoreMte2SyncFlagId = 15; // IS_SPLIT_G 核间 MTE2 同步 flag ID + + __aicore__ inline CSABlockCube(){}; + __aicore__ inline void InitLocalBuffer(uint32_t l1BaseAddr); + __aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *query, __gm__ uint8_t *oriKv, __gm__ uint8_t *cmpKv, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, + __gm__ uint8_t *cmpBlockTable, __gm__ uint8_t *sequsedQ, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *seqUsedOriKV, + __gm__ uint8_t *seqUsedCmpKV, const ConstInfo &constInfo); + + // SWA/HCA场景 + __aicore__ inline void IterateLoadQK(RunInfo &runInfo, ConstInfo &constInfo, bool isFirstLoop); + + // CSA场景 + __aicore__ inline void IterateLoadQK(Buffer &v0ResGm, + RunInfo &runInfo, ConstInfo &constInfo, bool isFirstLoop); + __aicore__ inline void IterateBmm1(StaticBuffer &output, bool notLastTwoLoop, RunInfo &runInfoNext, + RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void IterateBmm2(StaticBuffer &outputBuf, StaticBuffer &l1PBuffer, const RunInfo &runInfo, + const ConstInfo &constInfo); + __aicore__ inline void FreeEvent(); + +private: + __aicore__ inline void InitGmTensor(__gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *sequsedQ, + __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, + const ConstInfo &constInfo); + __aicore__ inline void CalcS2Coord(const RunInfo &runInfo, const ConstInfo &constInfo); + __aicore__ inline void CopyQGmToL1(RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void LoadKGmToL1(LocalTensor &inputRightTensor, const RunInfo &runInfo, + const ConstInfo &constInfo); + + __aicore__ inline void IterateBmm1Impl(StaticBuffer &outputBuf, bool notLastTwoLoop, RunInfo &runInfoNext, + RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void IterateBmm2Impl(StaticBuffer &outputBuf, StaticBuffer &l1PBuffer, + const RunInfo &runInfo, const ConstInfo &constInfo); + + /* =====================GM变量==================== */ + static constexpr GmFormat Q_FORMAT = GetQueryGmFormat(); + static constexpr GmFormat KV_FORMAT = GetKvGmFormat(); + static constexpr bool Q_WITH_ZERO_HEAD = (LAYOUT_T == SMLA_LAYOUT::TND); + FaGmTensor queryGm; + static constexpr bool KV_WITH_ZERO_HEAD = (KV_LAYOUT_T == SMLA_LAYOUT::TND); + FaGmTensor oriKvGm; + FaGmTensor cmpKvGm; + GlobalTensor cmpSparseIndicesGm; + GlobalTensor oriBlockTableGm; + GlobalTensor cmpBlockTableGm; + GlobalTensor blockTableGm; + FaGmTensor curKvGm; + GlobalTensor cuSeqlensQGm; + + /* =====================运行时变量==================== */ + CubeCoordInfo coordInfo[3]; + uint32_t kvCacheBlockSize = 0; + uint32_t maxBlockNumPerBatch = 0; + uint32_t l1QBufId = 0; // 3 buffer, 0-2 (轮转游标) + uint32_t l1KLoadBufId = 0; + uint32_t l1KMatmul1BufId = 0; + uint32_t l1KMatmul2BufId = 0; + /* =====================LocalBuffer变量==================== */ + StaticBuffer l1QBufs[3]; + StaticBuffer l1RightBufs[3]; + StaticBuffer l0ABufs[2]; + RingBuffer l0A; + StaticBuffer l0BBufs[2]; + RingBuffer l0B; + StaticBuffer l0CBufs[2]; + RingBuffer l0C; +}; + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::InitLocalBuffer(uint32_t l1BaseAddr) +{ + if ASCEND_IS_AIC { + uint32_t l1Addr = l1BaseAddr; + + l1QBufs[0] = {LocalTensor(TPosition::A1, l1Addr, L1Q_ELEM_PER_BUF), 0}; + l1Addr += L1Q_ELEM_PER_BUF * sizeof(Q_T); + l1QBufs[1] = {LocalTensor(TPosition::A1, l1Addr, L1Q_ELEM_PER_BUF), 1}; + l1Addr += L1Q_ELEM_PER_BUF * sizeof(Q_T); + l1QBufs[2] = {LocalTensor(TPosition::A1, l1Addr, L1Q_ELEM_PER_BUF), 2}; + l1Addr += L1Q_ELEM_PER_BUF * sizeof(Q_T); + + l1RightBufs[0] = {LocalTensor(TPosition::B1, l1Addr, L1_RIGHT_ELEM_PER_BLOCK), 0}; + l1Addr += L1_RIGHT_ELEM_PER_BLOCK * sizeof(Q_T); + l1RightBufs[1] = {LocalTensor(TPosition::B1, l1Addr, L1_RIGHT_ELEM_PER_BLOCK), 1}; + l1Addr += L1_RIGHT_ELEM_PER_BLOCK * sizeof(Q_T); + l1RightBufs[2] = {LocalTensor(TPosition::B1, l1Addr, L1_RIGHT_ELEM_PER_BLOCK), 2}; + l1Addr += L1_RIGHT_ELEM_PER_BLOCK * sizeof(Q_T); + + uint32_t l0aAddr = 0; + l0ABufs[0] = {LocalTensor(TPosition::A2, l0aAddr, L0A_ELEM_PER_BUF), 0}; + l0aAddr += L0A_ELEM_PER_BUF * sizeof(Q_T); + l0ABufs[1] = {LocalTensor(TPosition::A2, l0aAddr, L0A_ELEM_PER_BUF), 1}; + + uint32_t l0bAddr = 0; + l0BBufs[0] = {LocalTensor(TPosition::B2, l0bAddr, L0B_ELEM_PER_BUF), 0}; + l0bAddr += L0B_ELEM_PER_BUF * sizeof(Q_T); + l0BBufs[1] = {LocalTensor(TPosition::B2, l0bAddr, L0B_ELEM_PER_BUF), 1}; + + uint32_t l0cAddr = 0; + l0CBufs[0] = {LocalTensor(TPosition::CO1, l0cAddr, L0C_ELEM_PER_BUF), 0}; + l0cAddr += L0C_ELEM_PER_BUF * sizeof(T); + l0CBufs[1] = {LocalTensor(TPosition::CO1, l0cAddr, L0C_ELEM_PER_BUF), 1}; + + l0A = RingBuffer(l0ABufs, 2); + l0B = RingBuffer(l0BBufs, 2); + l0C = RingBuffer(l0CBufs, 2); + + SetFlag(INNERCORE_L0C(0)); + SetFlag(INNERCORE_L0C(1)); + SetFlag(INNERCORE_L0AB(0)); + SetFlag(INNERCORE_L0AB(1)); + SetFlag(INNERCORE_L1Q(0)); + SetFlag(INNERCORE_L1Q(1)); + SetFlag(INNERCORE_L1Q(2)); + SetFlag(INNERCORE_L1KV(0)); + SetFlag(INNERCORE_L1KV(1)); + SetFlag(INNERCORE_L1KV(2)); + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::InitGlobalBuffer( + __gm__ uint8_t *query, __gm__ uint8_t *oriKv, __gm__ uint8_t *cmpKv, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *cuSeqlensQ, + __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *seqUsedOriKV, + __gm__ uint8_t *seqUsedCmpKV, const ConstInfo &constInfo) +{ + if ASCEND_IS_AIC { + this->queryGm.gmTensor.SetGlobalBuffer((__gm__ Q_T *)query); + this->oriKvGm.gmTensor.SetGlobalBuffer((__gm__ KV_T *)oriKv); + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + this->oriBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)oriBlockTable); + } + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + this->cmpKvGm.gmTensor.SetGlobalBuffer((__gm__ KV_T *)cmpKv); + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + this->cmpBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)cmpBlockTable); + } + } + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE) { + this->cmpSparseIndicesGm.SetGlobalBuffer((__gm__ int32_t *)cmpSparseIndices); + } + InitGmTensor(cuSeqlensQ, sequsedQ, cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedOriKV, seqUsedCmpKV, constInfo); + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::FreeEvent() +{ + WaitFlag(INNERCORE_L0AB(0)); + WaitFlag(INNERCORE_L0AB(1)); + WaitFlag(INNERCORE_L0C(0)); + WaitFlag(INNERCORE_L0C(1)); + WaitFlag(INNERCORE_L1Q(0)); + WaitFlag(INNERCORE_L1Q(1)); + WaitFlag(INNERCORE_L1Q(2)); // 2 for l1q buffer id + WaitFlag(INNERCORE_L1KV(0)); + WaitFlag(INNERCORE_L1KV(1)); + WaitFlag(INNERCORE_L1KV(2)); // 2 for l1kv buffer id +} + +/* 初始化GmTensor,设置shape信息并计算strides */ +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::InitGmTensor(__gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *sequsedQ, + __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *seqUsedOriKV, + __gm__ uint8_t *seqUsedCmpKV, + const ConstInfo &constInfo) +{ + if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND) { + this->queryGm.offsetCalculator.Init(constInfo.bSize, constInfo.n2Size, constInfo.gSize, constInfo.s1Size, + constInfo.dSize); + } else { // SMLA_LAYOUT::TND + cuSeqlensQGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensQ); + uint32_t sequsedQSize = (sequsedQ == nullptr) ? 0 : constInfo.bSize; + ActualSeqLensParser parser; + parser.Init(cuSeqlensQ, constInfo.bSize + 1, sequsedQ, sequsedQSize); + this->queryGm.offsetCalculator.Init(constInfo.n2Size, constInfo.gSize, constInfo.dSize, parser); + } + + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + this->oriKvGm.offsetCalculator.Init(constInfo.n2Size, constInfo.oriBlockSize, constInfo.dSize, + this->oriBlockTableGm, constInfo.oriMaxBlockNumPerBatch); + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE) { + this->cmpKvGm.offsetCalculator.Init(constInfo.n2Size, constInfo.cmpBlockSize, constInfo.dSize, + this->cmpBlockTableGm, constInfo.cmpMaxBlockNumPerBatch); + } + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + uint32_t seqUsedOriKvSize = (seqUsedOriKV == nullptr) ? 0 : constInfo.bSize; + ActualSeqLensParser parser; + parser.Init(cuSeqlensOriKv, constInfo.actualSeqLenSize + 1, seqUsedOriKV, seqUsedOriKvSize); + this->oriKvGm.offsetCalculator.Init(constInfo.n2Size, constInfo.dSize, parser); + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + uint32_t seqUseCmpKvSize = (seqUsedCmpKV == nullptr) ? 0 : constInfo.bSize; + ActualSeqLensParser parser; + parser.Init(cuSeqlensCmpKv, constInfo.actualSeqLenSize + 1, seqUsedCmpKV, seqUseCmpKvSize); + this->cmpKvGm.offsetCalculator.Init(constInfo.n2Size, constInfo.dSize, parser); + } + } else { + // BSND不需要初始化 + this->oriKvGm.offsetCalculator.Init(constInfo.bSize, constInfo.n2Size, constInfo.s2Size, constInfo.dSize); + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE) { + this->cmpKvGm.offsetCalculator.Init(constInfo.bSize, constInfo.n2Size, constInfo.cmpS2Size, + constInfo.dSize); // 替换constInfo.s2Size / constInfo.cmpRatio + } + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::CalcS2Coord(const RunInfo &runInfo, const ConstInfo &constInfo) +{ + // 计算s2方向偏移 + coordInfo[runInfo.taskIdMod3].curBIdx = runInfo.boIdx; + if (runInfo.s2LoopCount >= runInfo.oriKvLoopEndIdx) { + kvCacheBlockSize = constInfo.cmpBlockSize; + maxBlockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch; + coordInfo[runInfo.taskIdMod3].s2Coord = + runInfo.s2StartIdx + (runInfo.s2LoopCount - runInfo.oriKvLoopEndIdx) * s2BaseSize; + blockTableGm = cmpBlockTableGm; + curKvGm = cmpKvGm; + } else { + kvCacheBlockSize = constInfo.oriBlockSize; + maxBlockNumPerBatch = constInfo.oriMaxBlockNumPerBatch; + coordInfo[runInfo.taskIdMod3].s2Coord = runInfo.s2StartIdx + runInfo.s2LoopCount * s2BaseSize; + blockTableGm = oriBlockTableGm; + curKvGm = oriKvGm; + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::CopyQGmToL1(RunInfo &runInfo, ConstInfo &constInfo) +{ + uint64_t gmOffset = this->queryGm.offsetCalculator.GetOffset(runInfo.boIdx, runInfo.n2oIdx, runInfo.goIdx, + runInfo.s1oIdx * runInfo.qSNumInOneBlock, 0); + for (uint32_t i = 0; i < qHalfNum; i++) { + uint32_t curL1QBufId = (l1QBufId + i) % l1QBufNum; + WaitFlag(INNERCORE_L1Q(curL1QBufId)); + uint64_t curGmOffset = gmOffset + i * (constInfo.dSize >> 1); + CopyToL1Nd2Nz(l1QBufs[curL1QBufId].tensor, this->queryGm.gmTensor[curGmOffset], runInfo.mRealSize, + constInfo.dSize >> 1, constInfo.mm1Ka); + SetFlag(INNERCORE_L1Q(curL1QBufId)); + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::LoadKGmToL1(LocalTensor &inputRightTensor, + const RunInfo &runInfo, const ConstInfo &constInfo) +{ + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + Position startPos; + startPos.bIdx = runInfo.boIdx; + startPos.n2Idx = runInfo.n2oIdx; + startPos.s2Offset = coordInfo[runInfo.taskIdMod3].s2Coord; + startPos.dIdx = 0; + PAShape shape; + shape.blockSize = kvCacheBlockSize; + shape.headNum = constInfo.n2Size; + shape.headDim = constInfo.dSize; + shape.actHeadDim = constInfo.dSize; + shape.maxblockNumPerBatch = maxBlockNumPerBatch; + shape.copyRowNum = runInfo.s2RealSize; + shape.copyRowNumAlign = (runInfo.s2RealSize + 15) >> 4 << 4; // 15,4:进行16字节对齐处理 + shape.pageStride = runInfo.isCmp ? constInfo.cmpKeyStride0 : constInfo.oriKeyStride0; + GmCopyInToL1PA(inputRightTensor, curKvGm.gmTensor, blockTableGm, KVLAYOUT::BBH, shape, startPos); + } else { + int64_t keyOffset = this->curKvGm.offsetCalculator.GetOffset( + coordInfo[runInfo.taskIdMod3].curBIdx, runInfo.n2oIdx, coordInfo[runInfo.taskIdMod3].s2Coord, 0); + CopyToL1Nd2Nz(inputRightTensor, curKvGm.gmTensor[keyOffset], runInfo.s2RealSize, constInfo.dSize, + constInfo.mm1Kb); + } +} + +// SWA/HCA场景: K从GM直接搬运 +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::IterateLoadQK(RunInfo &runInfo, ConstInfo &constInfo, + bool isFirstLoop) +{ + if (unlikely(isFirstLoop)) { + CopyQGmToL1(runInfo, constInfo); + } + + // 加载当前轮的右矩阵到L1 + CalcS2Coord(runInfo, constInfo); + WaitFlag(INNERCORE_L1KV(l1KLoadBufId)); + LocalTensor dst = l1RightBufs[runInfo.taskIdMod3].tensor; + LoadKGmToL1(dst, runInfo, constInfo); + SetFlag(INNERCORE_L1KV(l1KLoadBufId)); + l1KLoadBufId = (l1KLoadBufId + 1) % l1KBufNum; +} + +// CSA场景: K来源取决于isSparse +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::IterateLoadQK( + Buffer &v0ResGm, RunInfo &runInfo, ConstInfo &constInfo, + bool isFirstLoop) +{ + if (unlikely(isFirstLoop)) { + CopyQGmToL1(runInfo, constInfo); + } + + // ORI_SPARSE、ORI_CMP_SPARSE及CSA的cmpKv为v0稀疏搬运,CSA的oriKv为cube搬运 + bool isSparse = false; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + isSparse = true; + } else if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE) { + isSparse = runInfo.isCmp ? true : false; + } + + WaitFlag(INNERCORE_L1KV(l1KLoadBufId)); + LocalTensor dst = l1RightBufs[runInfo.taskIdMod3].tensor; + if (!isSparse) { + if constexpr (IS_SPLIT_G) { + CrossCoreSetFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + CrossCoreWaitFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + } + // cube直接从kv cache搬运K + CalcS2Coord(runInfo, constInfo); + LoadKGmToL1(dst, runInfo, constInfo); + } else { + v0ResGm.WaitCrossCore(); + if constexpr (IS_SPLIT_G) { + CrossCoreSetFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + CrossCoreWaitFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + } + GlobalTensor v0ResGmTensor = v0ResGm.template GetTensor(); + CopyToL1Nd2Nz(dst, v0ResGmTensor, runInfo.s2RealSize, constInfo.dSize, constInfo.mm1Kb); + } + SetFlag(INNERCORE_L1KV(l1KLoadBufId)); + l1KLoadBufId = (l1KLoadBufId + 1) % l1KBufNum; +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::IterateBmm1(StaticBuffer &outputBuf, bool notLastTwoLoop, + RunInfo &runInfoNext, RunInfo &runInfo, + ConstInfo &constInfo) +{ + IterateBmm1Impl(outputBuf, notLastTwoLoop, runInfoNext, runInfo, constInfo); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::IterateBmm2(StaticBuffer &outputBuf, + StaticBuffer &l1PBuffer, const RunInfo &runInfo, + const ConstInfo &constInfo) +{ + IterateBmm2Impl(outputBuf, l1PBuffer, runInfo, constInfo); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::IterateBmm1Impl(StaticBuffer &outputBuf, bool notLastTwoLoop, + RunInfo &runInfoNext, RunInfo &runInfo, + ConstInfo &constInfo) +{ + LocalTensor curL1RightTensor = l1RightBufs[runInfo.taskIdMod3].tensor; + WaitFlag(INNERCORE_L1KV(l1KMatmul1BufId)); + l1KMatmul1BufId = (l1KMatmul1BufId + 1) % l1KBufNum; + + StaticBuffer &cBuf = l0C.GetNext(); + WaitFlag(INNERCORE_L0C(cBuf.idx)); + MMParam param = { + static_cast(runInfo.mRealSize), // singleM + static_cast(runInfo.s2RealSize), // singleN + static_cast(constInfo.dSize >> 1), // singleK + 0, // isLeftTranspose + 1 // isRightTranspose + }; + uint32_t curL1QBufId = l1QBufId; + if (unlikely(runInfo.s2LoopCount == 0)) { + WaitFlag(INNERCORE_L1Q(curL1QBufId)); + } + + // m,n不切,k切128,mm1B直接用tensor的数据 + MatmulKStatic( + l1QBufs[curL1QBufId].tensor, curL1RightTensor, l0A, l0B, cBuf.tensor, param); + + curL1QBufId = (curL1QBufId + 1) % l1QBufNum; + if (unlikely(runInfo.s2LoopCount == 0)) { + WaitFlag(INNERCORE_L1Q(curL1QBufId)); + } + param.singleK = constInfo.dSize - param.singleK; + param.isOutKFisrt = false; + + // m,n不切,k切128, mm1B直接用tensor的数据 + MatmulKStatic( + l1QBufs[curL1QBufId].tensor, curL1RightTensor[(constInfo.dSize >> 1) * Align16Func(runInfo.s2RealSize)], l0A, + l0B, cBuf.tensor, param); + + if (unlikely(runInfo.s2LoopCount == runInfo.s2LoopLimit)) { + SetFlag(INNERCORE_L1Q(l1QBufId)); + SetFlag(INNERCORE_L1Q(curL1QBufId)); + l1QBufId = (l1QBufId + qHalfNum) % l1QBufNum; + if (notLastTwoLoop) { + CopyQGmToL1(runInfoNext, constInfo); + } + } + + SetFlag(INNERCORE_L0C(cBuf.idx)); + WaitFlag(INNERCORE_L0C(cBuf.idx)); + + CrossCoreWaitFlag(CROSSCORE_BMM1(outputBuf.idx)); + CrossCoreWaitFlag(CROSSCORE_BMM1(outputBuf.idx) + AIV0_AIV1_OFFSET); + FixpipeParamsC310 fixpipeParams; // L0C→UB + // L0C上的bmm1结果矩阵N方向的size大小; 同mmadParams.n; 为什么要8个元素对齐(32B对齐) // 128 + fixpipeParams.nSize = Align8Func(runInfo.s2RealSize); + // 有效数据不足16行,只需要输出部分行即可; L0C上的bmm1结果矩阵M方向的size大小(必须为偶数) // 128 + fixpipeParams.mSize = Align2Func(runInfo.mRealSize); + // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔), 单位为16*sizeof(T) // + // 源Nz矩阵中相邻大Z排布的起始地址偏移 + fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize); + // mmResUb上两行之间的间隔,单位:element。 // 128:根据比对dump文件得到, ND方案(S1*S2)时脏数据用mask剔除 + fixpipeParams.dstStride = s2BaseSize; + // 双目标模式,按M维度拆分,M / 2 * N写入每个UB, M必须为2的倍数 + fixpipeParams.dualDstCtl = 1; + fixpipeParams.params.ndNum = 1; + fixpipeParams.params.srcNdStride = 0; + fixpipeParams.params.dstNdStride = 0; + + // 将matmul结果从L0C搬运到UB + Fixpipe(outputBuf.tensor, cBuf.tensor, fixpipeParams); + SetFlag(INNERCORE_L0C(cBuf.idx)); + CrossCoreSetFlag(CROSSCORE_BMM1(outputBuf.idx)); + CrossCoreSetFlag(CROSSCORE_BMM1(outputBuf.idx) + AIV0_AIV1_OFFSET); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockCube::IterateBmm2Impl(StaticBuffer &outputBuf, + StaticBuffer &l1PBuffer, + const RunInfo &runInfo, const ConstInfo &constInfo) +{ + LocalTensor curL1RightTensor = l1RightBufs[runInfo.taskIdMod3].tensor; + CrossCoreWaitFlag(CROSSCORE_L1P(l1PBuffer.idx)); + CrossCoreWaitFlag(CROSSCORE_L1P(l1PBuffer.idx) + AIV0_AIV1_OFFSET); + + StaticBuffer &cBuf = l0C.GetNext(); + WaitFlag(INNERCORE_L0C(cBuf.idx)); + MMParam param = { + static_cast(runInfo.mRealSize), // singleM + static_cast(constInfo.dSizeV), // singleN 512 + static_cast(runInfo.s2RealSize), // singleK 128 + 0, // isLeftTranspose + 0 // isRightTranspose + }; + MatmulNStatic( + l1PBuffer.tensor, curL1RightTensor, l0A, l0B, cBuf.tensor, param); + + SetFlag(INNERCORE_L0C(cBuf.idx)); + WaitFlag(INNERCORE_L0C(cBuf.idx)); + SetFlag(INNERCORE_L1KV(l1KMatmul2BufId)); + l1KMatmul2BufId = (l1KMatmul2BufId + 1) % l1KBufNum; + + CrossCoreWaitFlag(CROSSCORE_BMM2); + CrossCoreWaitFlag(CROSSCORE_BMM2 + AIV0_AIV1_OFFSET); + // L0C→UB;FixpipeParamsM300:L0C→UB + FixpipeParamsC310 fixpipeParams; + // L0C上的bmm1结果矩阵N方向的size大小, 分档计算且vector2中通过mask筛选出实际有效值 + fixpipeParams.nSize = Align8Func(constInfo.dSizeV); + // 有效数据不足16行,只需要输出部分行即可; L0C上的bmm1结果矩阵M方向的size大小; 同mmadParams.m + fixpipeParams.mSize = Align2Func(runInfo.mRealSize); + // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔) + fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize); + fixpipeParams.dstStride = Align16Func(constInfo.dSizeV); + fixpipeParams.dualDstCtl = 1; + fixpipeParams.params.ndNum = 1; + fixpipeParams.params.srcNdStride = 0; + fixpipeParams.params.dstNdStride = 0; + Fixpipe(outputBuf.tensor, cBuf.tensor, fixpipeParams); // 将matmul结果从L0C搬运到UB + SetFlag(INNERCORE_L0C(cBuf.idx)); + + CrossCoreSetFlag(CROSSCORE_BMM2); + CrossCoreSetFlag(CROSSCORE_BMM2 + AIV0_AIV1_OFFSET); +} + +TEMPLATES_DEF +class CSABlockCubeDummy { +public: + __aicore__ inline CSABlockCubeDummy(){}; + __aicore__ inline void InitLocalBuffer(uint32_t l1BaseAddr) {} + __aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *query, __gm__ uint8_t *oriKv, __gm__ uint8_t *cmpKv, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, + __gm__ uint8_t *cmpBlockTable, __gm__ uint8_t *sequsedQ, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *seqUsedOriKV, + __gm__ uint8_t *seqUsedCmpKV, const ConstInfo &constInfo) + {} + __aicore__ inline void FreeEvent() {} +}; + +template +struct CubeBlockTraits; // 声明 + +/* 生成CubeBlockTraits */ +#define GEN_TRAIT_TYPE(name, ...) using name##_TRAITS = name; +#define GEN_TRAIT_CONST(name, type, ...) static constexpr type name##Traits = name; + +#define DEFINE_CUBE_BLOCK_TRAITS(CUBE_BLOCK_CLASS) \ + TEMPLATES_DEF_NO_DEFAULT \ + struct CubeBlockTraits> { \ + CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TRAIT_TYPE) \ + CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TRAIT_CONST) \ + } + +DEFINE_CUBE_BLOCK_TRAITS(CSABlockCube); +DEFINE_CUBE_BLOCK_TRAITS(CSABlockCubeDummy); + +// /* 生成Arg Traits, kernel中只需要调用ARGS_TRAITS就可以获取所有CubeBlock中的模板参数 */ +#define GEN_ARGS_TYPE(name, ...) using name = typename CubeBlockTraits::name##_TRAITS; +#define GEN_ARGS_CONST(name, type, ...) static constexpr type name = CubeBlockTraits::name##Traits; +#define ARGS_TRAITS \ + CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARGS_TYPE) \ + CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARGS_CONST) +} // namespace SMLAKernel +#endif // FLASH_ATTENTION_SCORE_BLOCK_CUBE_H_ diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_block_vector_arch35.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_block_vector_arch35.h new file mode 100644 index 000000000000..2246d9a97d6a --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_block_vector_arch35.h @@ -0,0 +1,2339 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_csa_block_vector_arch35.h + * \brief + */ +#ifndef SPARSE_FLASH_MLA_CSA_BLOCK_VECTOR_ARCH35_H +#define SPARSE_FLASH_MLA_CSA_BLOCK_VECTOR_ARCH35_H + +#include "util_regbase.h" +#include "sparse_flash_mla_common_arch35.h" +#include "kernel_operator_list_tensor_intf.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" + +using AscendC::Reg::StoreDist; + +#include "common/flash_decode.h" + +#if __has_include("../../common/op_kernel/arch35/vf/vf_flash_decode_arch35.h") +#include "../../common/op_kernel/arch35/vf/vf_flash_decode_arch35.h" +#else +#include "../common/arch35/vf/vf_flash_decode_arch35.h" +#endif + +#if __has_include("../../common/op_kernel/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h") +#include "../../common/op_kernel/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h" +#else +#include "../../common/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h" +#endif + +#if __has_include("../../common/op_kernel/arch35/vf/vf_flashupdate_new.h") +#include "../../common/op_kernel/arch35/vf/vf_flashupdate_new.h" +#else +#include "../../common/arch35/vf/vf_flashupdate_new.h" +#endif + +#if __has_include("../../common/op_kernel/buffers_policy.h") +#include "../../common/op_kernel/buffers_policy.h" +#else +#include "../common/buffers_policy.h" +#endif +#if __has_include("../../common/op_kernel/attn_buffer_manager.h") +#include "../../common/op_kernel/attn_buffer_manager.h" +#else +#include "../common/attn_buffer_manager.h" +#endif +#if __has_include("../../common/op_kernel/attn_buffer.h") +#include "../../common/op_kernel/attn_buffer.h" +#else +#include "../common/attn_buffer.h" +#endif +#if __has_include("../../common/op_kernel/init_output.h") +#include "../../common/op_kernel/init_output.h" +#else +#include "../common/init_output.h" +#endif + +using namespace AscendC; +using namespace FaVectorApi; +using namespace AscendC::Impl::Detail; +using namespace regbaseutil; +using namespace matmul; +using namespace fa_base_matmul; +using AttentionCommon::FdRunInfo; + +namespace SMLAKernel { + +// 统一窗口公式 +struct PhyAddrValidInfo { + static constexpr int64_t BIAS_UNBOUND = 0x7FFFFFFF; // INT32_MAX + int64_t oriLeftBias = BIAS_UNBOUND; + int64_t oriRightBias = BIAS_UNBOUND; + int32_t oriS2Act = 0; + bool oriTopkMode = false; // oriMaskMode==0: 走topkLength语义(保持原行为) + bool cmpTopkMode = true; // cmpMaskMode==0: 走topkLength语义(保持原行为) + int64_t cmpBase = 0; // restoredSize - actualS1Size + 1 +}; + +TEMPLATES_DEF +class CSABlockVec { +public: + // BUFFER的字节数 + static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32; + /* =================编译期常量的基本块信息================= */ + static constexpr uint32_t s1BaseSize = 64; + static constexpr uint32_t s2BaseSize = 128; + static constexpr uint32_t vec1Srcstride = (s1BaseSize >> 1) + 1; + static constexpr uint32_t dVTemplateType = 512; + static constexpr uint32_t dTemplateAlign64 = Align64Func(dVTemplateType); + static constexpr float R0 = 1.0f; + static constexpr uint32_t initOutputEventId = + INNERCORE_INITOUT_MTE3_V; // attenOut和lse,刷无效行会用到剩余ub,需要加同步 + // Sparse KV搬入/拷出块大小:每次搬入8行,每16行做一次拷出 + static constexpr int64_t KV_COPYIN_UNIT = 8; // 每次搬入8行 + static constexpr int64_t KV_PROCESS_UNIT = 16; // 每次拷出16行 + + // ==================== Functions ====================== + __aicore__ inline CSABlockVec(){}; + __aicore__ inline void InitVecBlock(__gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *seqUsedOriKV, + __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV) + { + if ASCEND_IS_AIV { + if (cuSeqlensQ != nullptr) { + cuSeqlensQGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensQ); + } + if (cuSeqlensOriKv != nullptr) { + cuSeqlensOriKvGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensOriKv); + } + if (cuSeqlensCmpKv != nullptr) { + cuSeqlensCmpKvGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensCmpKv); + } + if (seqUsedOriKV != nullptr) { + actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)seqUsedOriKV); + } + if (seqUsedCmpKV != nullptr) { + actualSeqLengthsCmpKVGm.SetGlobalBuffer((__gm__ int32_t *)seqUsedCmpKV); + } + if (cmpResidualKV != nullptr) { + cmpResidualKVGm.SetGlobalBuffer((__gm__ int32_t *)cmpResidualKV); + } + this->GetExtremeValue(this->negativeFloatScalar); + } + } + + // 初始化LocalTensor + __aicore__ inline void InitLocalBuffer(ConstInfo &constInfo, uint32_t ubBaseAddr); + // 初始化attentionOutGM + __aicore__ inline void CleanOutput(__gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, ConstInfo &constInfo); + __aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, + __gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sinks, + __gm__ uint8_t *sequsedOriKv, __gm__ uint8_t *sequsedCmpKv, + __gm__ uint8_t *cmpResidualKv); + __aicore__ inline void InitOutputSingleCore(ConstInfo &constInfo); + __aicore__ inline void ProcessVec0(Buffer &v0ResGm, + const RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void ProcessVec1(StaticBuffer &outputBuf, StaticBuffer &bmm1ResBuf, RunInfo &runInfo, + ConstInfo &constInfo); + __aicore__ inline void InitS2SplitStaging(Buffer &fdStaging) + { + fdStagingBase = fdStaging.template GetTensor().GetPhyAddr(0); + stagingOutGm = fdStaging.template GetTensor(); + } + __aicore__ inline void InitS2SplitStaging(Buffer &intraCoreCombine, + Buffer &crossCoreCombine) + { + intraCoreCombineBase = intraCoreCombine.template GetTensor().GetPhyAddr(0); + intraCoreCombineGm = intraCoreCombine.template GetTensor(); + crossCoreCombineBase = crossCoreCombine.template GetTensor().GetPhyAddr(0); + crossCoreCombineGm = crossCoreCombine.template GetTensor(); + fdStagingBase = crossCoreCombineBase; + stagingOutGm = crossCoreCombineGm; + } + __aicore__ inline void InitFDBuffers(FdRunInfo &fdRunInfo); + __aicore__ inline void ProcessFlashDecode(FdRunInfo &fdRunInfo, ConstInfo &constInfo); + __aicore__ inline void ProcessVec2(StaticBuffer &bmm2ResBuf, RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void GetKVPhyAddr(uint32_t hasLoad, uint32_t bN2StartIdx, uint32_t bN2EndIdx, + uint32_t gS1StartIdx, uint32_t nextGs1Idx, bool hasActualSeqQlen, + bool hasCuSeqlensQ, bool hasActualSeqOriKvlen, bool hasCuSeqlensOriKv, + GlobalTensor actualSeqOriKvlenGm, + GlobalTensor cuSeqlensOriKvGm, GlobalTensor oriTopkLengthGm, + bool hasActualSeqCmpKvlen, bool hasCuSeqlensCmpKv, + GlobalTensor actualSeqCmpKvlenGm, + GlobalTensor cuSeqlensCmpKvGm, GlobalTensor cmpTopkLengthGm, + GlobalTensor cmpResidualKvGm, GlobalTensor actualSeqQlenGm, + GlobalTensor cuSeqlensQGm, __gm__ uint8_t *workspace, + ConstInfo &constInfo); + __aicore__ inline void FreeEvent(ConstInfo &constInfo); + +private: + template + __aicore__ inline void ComputeVec1Softmax(LocalTensor &stage1CastTensor, LocalTensor &mmRes, + LocalTensor &sumUb, LocalTensor &maxUb, + LocalTensor &apiTmpBuffer, RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void InitVec1SoftmaxFromSinks(LocalTensor &sumUb, LocalTensor &maxUb, + RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void CopyVec1ResultToL1(StaticBuffer &outputBuf, LocalTensor &stage1CastTensor, + RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void StageCrossCoreVec1Lse(LocalTensor &maxUb, LocalTensor &sumUb, RunInfo &runInfo, + ConstInfo &constInfo); + __aicore__ inline void StageBatchConsistencyVec1Lse(LocalTensor &maxUb, LocalTensor &sumUb, + RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void StageLegacyVec1Lse(LocalTensor &maxUb, LocalTensor &sumUb, RunInfo &runInfo, + ConstInfo &constInfo); + __aicore__ inline void CopyOutVec1Lse(LocalTensor &maxUb, LocalTensor &sumUb, RunInfo &runInfo, + ConstInfo &constInfo); + + __aicore__ inline uint32_t GetStagingSlotNum(bool isInner = false) const + { + if constexpr (IS_BATCH_CONSISTENCY) { + if (isInner) { + if constexpr (IS_SPLIT_G) { + return GetBlockNum(); + } + return GetBlockNum() << 1U; + } + if constexpr (IS_SPLIT_G) { + return BATCH_CONSISTENCY_MAX_REDUCE_BLOCK_NUM * (GetBlockNum() >> 1U); + } + return BATCH_CONSISTENCY_MAX_REDUCE_BLOCK_NUM * GetBlockNum(); + } + if constexpr (IS_SPLIT_G) { + return AttentionCommon::FD_MAX_S2_SPLIT_NUM * (GetBlockNum() >> 1U); + } else { + return AttentionCommon::FD_MAX_S2_SPLIT_NUM * GetBlockNum(); + } + } + + __aicore__ inline uint32_t GetIntraCoreWorkspaceIdx(const RunInfo &runInfo, const ConstInfo &constInfo) const + { + uint32_t coreIdx; + if constexpr (IS_SPLIT_G) { + coreIdx = static_cast(constInfo.aivIdx >> 2U); + } else { + coreIdx = static_cast(constInfo.aivIdx >> 1U); + } + return (coreIdx << 1U) + runInfo.multiCoreIdxMod2; + } + + __aicore__ inline uint32_t GetCrossCoreWorkspaceIdx(const RunInfo &runInfo) const + { + return static_cast(runInfo.firstFdDataWorkspaceIdx + runInfo.s2SplitIdx); + } + + __aicore__ inline int64_t GetFaStagingMOffset(const RunInfo &runInfo, const ConstInfo &constInfo) const + { + int64_t stagingMOffset = (constInfo.subBlockIdx == 1) ? static_cast(runInfo.firstHalfMRealSize) : 0L; + if constexpr (IS_SPLIT_G) { + stagingMOffset += static_cast(runInfo.goIdx); + } + return stagingMOffset; + } + + __aicore__ inline void ProcessSparseKv(Buffer &v0ResGm, + const RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void CalSparseCalSize(const RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline int64_t GetkeyOffset(int64_t s2Idx, const RunInfo &runInfo, ConstInfo &constInfo); + template + __aicore__ inline void GetRealCmpS2Idx(int64_t *tokenData, int64_t s2IdxInBase, const RunInfo &runInfo, + ConstInfo &constInfo); + template + __aicore__ inline void CopyInKvSparse(LocalTensor kvInUb, int64_t startRow, int64_t *tokenData, + const RunInfo &runInfo, ConstInfo &constInfo); + template + __aicore__ inline void CopyIn8Block(LocalTensor kvInUb, int64_t startRow, int64_t &s2, const RunInfo &runInfo, + ConstInfo &constInfo); + __aicore__ inline void CopyToOutUb(LocalTensor kvNzUb, LocalTensor srcTensor, int64_t dealRow, + ConstInfo &constInfo); + __aicore__ inline void CopyOutKvUb2Gm(Buffer &v0ResGm, + LocalTensor kvOutUb, int64_t dealRow, int64_t s2StartIdx, + const RunInfo &runInfo, ConstInfo &constInfo); + __aicore__ inline void CopyInSingleKv(LocalTensor kvInUb, int64_t startRow, int64_t keyOffset, + ConstInfo &constInfo); + template + __aicore__ inline void GetRealS2Addr(int64_t *tokenData, int64_t s2IdxInBase, const RunInfo &runInfo, + ConstInfo &constInfo); + __aicore__ inline void GetKVPhyAddrForKvType( + uint32_t bN2StartIdx, uint32_t bN2EndIdx, uint32_t gS1StartIdx, uint32_t nextGs1Idx, bool hasActualSeqQlen, + bool hasCuSeqlensQ, bool hasActualSeqKvlen, bool hasCuSeqlensKv, GlobalTensor actualSeqQlenGm, + GlobalTensor cuSeqlensQGm, GlobalTensor actualSeqKvlenGm, GlobalTensor cuSeqlensKvGm, + GlobalTensor topkLengthGm, GlobalTensor cmpResidualKvGm, ConstInfo &constInfo, + GlobalTensor &blockTableGm, GlobalTensor &sparseIndicesGm, GlobalTensor &phyAddrGm, + uint32_t kvStride, uint32_t blockSize, uint32_t maxBlockNumPerBatch, uint32_t sparseBlockCount, + uint32_t alignedSparseBlockCount, bool isOriKv); + __aicore__ inline int32_t GetSeqLen(int32_t bIdx, bool hasActualSeq, bool hasCuSeqlens, + GlobalTensor &actualSeqGm, GlobalTensor &cuSeqlensGm, + int64_t defaultSize); + __aicore__ inline PhyAddrValidInfo CalcPhyAddrValidInfo(bool isOriKv, int32_t actualS1Size, int32_t actualOriS2Size, + int32_t restoredSize, ConstInfo &constInfo); + __aicore__ inline int32_t CalcCurValidS2(uint32_t bIdx, int32_t s1Idx, int32_t actualS1Size, bool isOriKv, + GlobalTensor &cuSeqlensQGm, GlobalTensor &topkLengthGm, + ConstInfo &constInfo, int32_t sparseBlockCount, + const PhyAddrValidInfo &validInfo); + __aicore__ inline void CopyPhyAddrToGm(LocalTensor kvPhyAddrUb, int64_t bS1Idx, int64_t s1Idx, + int64_t validS2, int64_t alignNum, GlobalTensor &phyAddrGm, + uint32_t alignedSparseBlockCount); + __aicore__ inline void CopyPaTableToUb(LocalTensor blkTableUb, int64_t bIdx, + GlobalTensor &blockTableGm, uint32_t maxBlockNumPerBatch); + __aicore__ inline void CopySparseIdxToUb(LocalTensor sparseIdxUb, int64_t bS1Idx, int64_t s1Idx, + int64_t validS2, GlobalTensor &sparseIndicesGm, + uint32_t sparseBlockCount); + /* VEC2_RES_T 表示bmm2ResUb当前的类型,VEC2_RES_T = Q_T那么不需要做Cast。另外,无效行场景当前默认需要做Cast */ + template + __aicore__ inline void Bmm2DataCopyOut(RunInfo &runInfo, ConstInfo &constInfo, LocalTensor &vec2ResUb, + int64_t vec2S1Idx, int64_t vec2CalcSize = 0); + template + __aicore__ inline void CopyOutAttentionOut(RunInfo &runInfo, ConstInfo &constInfo, + LocalTensor &vec2ResUb, int64_t vec2S1Idx, + int64_t vec2CalcSize); + __aicore__ inline void SoftmaxInitBuffer(uint32_t &ubAddr); + __aicore__ inline void GetExtremeValue(T &negativeScalar); + __aicore__ inline void InitSinksBuffer(ConstInfo &constInfo); + __aicore__ inline void ReduceIntraBlockAndStage(RunInfo &runInfo, ConstInfo &constInfo, LocalTensor &vec2ResUb, + LocalTensor &partialTmpUb); + + GlobalTensor attentionOutGm; + GlobalTensor softmaxLseGm; + GlobalTensor oriKVGm; + GlobalTensor cmpKVGm; + GlobalTensor keyGm; + GlobalTensor cuSeqlensKvGm; + GlobalTensor oriSparseIndicesGm; + GlobalTensor cmpSparseIndicesGm; + GlobalTensor sparseIndicesGm; + GlobalTensor oriBlockTableGm; + GlobalTensor cmpBlockTableGm; + GlobalTensor blockTableGm; + GlobalTensor sinksGm; + GlobalTensor cuSeqlensQGm; + GlobalTensor cuSeqlensOriKvGm; + GlobalTensor cuSeqlensCmpKvGm; + GlobalTensor actualSeqLengthsKVGm; + GlobalTensor actualSeqLengthsCmpKVGm; + GlobalTensor cmpResidualKVGm; + GlobalTensor oriKvPhyAddrGm; + GlobalTensor cmpKvPhyAddrGm; + + StaticBuffer commonUb; + StaticBuffer sinksUb; + StaticBuffer stage1OutBufs[2]; + StaticBuffer stage2OutBufs; + StaticBuffer stage0OutBufs[2]; + StaticBuffer softmaxMaxBufs[2]; + StaticBuffer softmaxSumBufs[2]; + StaticBuffer softmaxFinalMaxBufs[2]; + StaticBuffer softmaxFinalSumBufs[2]; + StaticBuffer softmaxExpBufs[2]; + StaticBuffer batchReduceTmpUb; + StaticBuffer outLseUbs[2]; + TBuf<> vselrIndexesBuf[2]; + AttentionCommon::FdBuffers> fdBuffers; + uint32_t pingPongV0 = 0; + __gm__ uint8_t *fdStagingBase = nullptr; + GlobalTensor stagingOutGm; + __gm__ uint8_t *intraCoreCombineBase = nullptr; + GlobalTensor intraCoreCombineGm; + __gm__ uint8_t *crossCoreCombineBase = nullptr; + GlobalTensor crossCoreCombineGm; + + T negativeFloatScalar; + bool isSinks = false; + uint32_t maxBlockNumPerBatch; + uint32_t blockSize; + int64_t sparseCalSize; + int64_t sparseS2Start; + int64_t sparseS2End; +}; + +TEMPLATES_DEF_NO_DEFAULT +template +__aicore__ inline void CSABlockVec::GetRealCmpS2Idx(int64_t *tokenData, int64_t s2IdxInBase, + const RunInfo &runInfo, ConstInfo &constInfo) +{ + int64_t curSparseS2End = this->sparseS2End; + int64_t sparseBlockCount = 0; + int64_t curS2LoopCnt = runInfo.s2LoopCount; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE) { + sparseBlockCount = constInfo.cmpSparseBlockCount; + curS2LoopCnt -= runInfo.oriKvLoopEndIdx; + } else if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + sparseBlockCount = constInfo.oriSparseBlockCount; + } else if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (runInfo.isCmp) { + sparseBlockCount = constInfo.cmpSparseBlockCount; + curS2LoopCnt -= runInfo.oriKvLoopEndIdx; + } else { + sparseBlockCount = constInfo.oriSparseBlockCount; + } + } + uint64_t topkBS1Idx = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + uint64_t actualSeqQPrefixSum = cuSeqlensQGm.GetValue(runInfo.boIdx); + topkBS1Idx += (actualSeqQPrefixSum + runInfo.s1oIdx) * sparseBlockCount; // T, N2(1), K + } else { + topkBS1Idx += + runInfo.boIdx * constInfo.s1Size * sparseBlockCount + runInfo.s1oIdx * sparseBlockCount; // B, S1, N2(1), K + } + + uint64_t topkKIdx = s2IdxInBase + curS2LoopCnt * constInfo.s2BaseSize; + for (uint64_t i = 0; i < KV_COPYIN_UNIT; ++i) { // 每次处理8个数据块 + uint64_t idx = topkBS1Idx + runInfo.s2StartIdx + topkKIdx + i; + if constexpr (!IS_FULL) { + // 尾块:保留边界判断,防止越界读取 + if (likely(s2IdxInBase + i < curSparseS2End)) { + tokenData[i] = sparseIndicesGm.GetValue(idx); + } else { + break; + } + } else { + // 非尾块:8行均有效,直接读取 + tokenData[i] = sparseIndicesGm.GetValue(idx); + } + } +} + +TEMPLATES_DEF_NO_DEFAULT +template +__aicore__ inline void CSABlockVec::GetRealS2Addr(int64_t *tokenData, int64_t s2IdxInBase, + const RunInfo &runInfo, ConstInfo &constInfo) +{ + int64_t curSparseS2End = this->sparseS2End; + uint32_t alignedSparseBlockCount = + runInfo.isCmp ? constInfo.alignedCmpSparseBlockCount : constInfo.alignedOriSparseBlockCount; + int64_t curS2LoopCnt = runInfo.s2LoopCount; + GlobalTensor phyAddrGm64; + if (runInfo.isCmp) { + curS2LoopCnt -= runInfo.oriKvLoopEndIdx; + phyAddrGm64 = cmpKvPhyAddrGm.template ReinterpretCast(); + } else { + phyAddrGm64 = oriKvPhyAddrGm.template ReinterpretCast(); + } + + uint64_t topkBS1Idx = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + uint64_t actualSeqQPrefixSum = cuSeqlensQGm.GetValue(runInfo.boIdx); + topkBS1Idx += (actualSeqQPrefixSum + runInfo.s1oIdx) * alignedSparseBlockCount; + } else { + topkBS1Idx += + runInfo.boIdx * constInfo.s1Size * alignedSparseBlockCount + runInfo.s1oIdx * alignedSparseBlockCount; + } + uint64_t topkKIdx = s2IdxInBase + curS2LoopCnt * constInfo.s2BaseSize; + for (uint64_t i = 0; i < KV_COPYIN_UNIT; ++i) { // 每次处理8个数据块 + uint64_t idx = topkBS1Idx + runInfo.s2StartIdx + topkKIdx + i; + if constexpr (!IS_FULL) { + // 尾块:保留边界判断,防止越界读取 + if (likely(s2IdxInBase + i < curSparseS2End)) { + tokenData[i] = phyAddrGm64.GetValue(idx); + } else { + break; + } + } else { + // 非尾块:8行均有效,直接读取 + tokenData[i] = phyAddrGm64.GetValue(idx); + } + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline int64_t CSABlockVec::GetkeyOffset(int64_t s2Idx, const RunInfo &runInfo, + ConstInfo &constInfo) +{ + if (s2Idx < 0) { + return -1; + } + int64_t realkeyOffset = 0; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + int64_t blkTableIdx = s2Idx / blockSize; + int64_t blkTableOffset = s2Idx % blockSize; + int64_t paBlockStride = runInfo.isCmp ? constInfo.cmpKeyStride0 : constInfo.oriKeyStride0; + realkeyOffset = blockTableGm.GetValue(runInfo.boIdx * maxBlockNumPerBatch + blkTableIdx) * paBlockStride + + blkTableOffset * constInfo.dSizeVInput; // BlockNum, BlockSize, N(1), D + } else if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND) { + if (runInfo.isCmp) { + realkeyOffset = runInfo.boIdx * constInfo.n2Size * constInfo.cmpS2Size * constInfo.dSize + + runInfo.n2oIdx * constInfo.cmpS2Size * constInfo.dSize + s2Idx * constInfo.dSize; // BSN(1)D + } else { + realkeyOffset = runInfo.boIdx * constInfo.n2Size * constInfo.s2Size * constInfo.dSize + + runInfo.n2oIdx * constInfo.s2Size * constInfo.dSize + s2Idx * constInfo.dSize; // BSN(1)D + } + } else if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + realkeyOffset = (cuSeqlensKvGm.GetValue(runInfo.boIdx) + s2Idx) * constInfo.n2Size * constInfo.dSize + + runInfo.n2oIdx * constInfo.dSize; // TN(1)D + } + return realkeyOffset; +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CopyInSingleKv(LocalTensor kvInUb, int64_t startRow, + int64_t keyOffset, ConstInfo &constInfo) +{ + if (keyOffset < 0) { + return; + } + DataCopyExtParams intriParams; + intriParams.blockCount = 1; + intriParams.dstStride = 0; + intriParams.srcStride = 0; + intriParams.blockLen = constInfo.dSize * sizeof(KV_T); + + DataCopyPadExtParams padParams; + padParams.isPad = true; + padParams.leftPadding = 0; + padParams.rightPadding = + (CeilAlign(constInfo.dSize * sizeof(KV_T), BUFFER_SIZE_BYTE_32B) - constInfo.dSize * sizeof(KV_T)) / + sizeof(KV_T); + padParams.paddingValue = 0; + DataCopyPad(kvInUb[startRow * constInfo.dSize], keyGm[keyOffset], intriParams, padParams); +} + +TEMPLATES_DEF_NO_DEFAULT +template +__aicore__ inline void CSABlockVec::CopyInKvSparse(LocalTensor kvInUb, int64_t startRow, + int64_t *tokenData, const RunInfo &runInfo, + ConstInfo &constInfo) +{ + for (uint32_t i = 0; i < 8; i += 2) { // 遍历8个元素的数组/缓冲区,每次处理2个元素 + int64_t keyOffset0; + int64_t keyOffset1; + if constexpr (IS_VEC_S2PHYADDR) { + keyOffset0 = tokenData[i]; + keyOffset1 = tokenData[i + 1]; + } else { + keyOffset0 = GetkeyOffset(tokenData[i], runInfo, constInfo); + keyOffset1 = GetkeyOffset(tokenData[i + 1], runInfo, constInfo); + } + if constexpr (!IS_FULL) { + // 尾块:提前返回判断 + if (unlikely(keyOffset0 < 0 && keyOffset1 < 0)) { + return; + } + } + int64_t combineBytes = constInfo.dSizeVInput * sizeof(KV_T); + int64_t keySrcStride = + (keyOffset0 > keyOffset1 ? (keyOffset0 - keyOffset1) : (keyOffset1 - keyOffset0)) * sizeof(KV_T) - + combineBytes; + if (unlikely(keySrcStride >= INT32_MAX || keySrcStride < 0) || constInfo.sparseBlockSize > 1) { + // stride溢出、stride为负数、s2超长等异常场景,还原成2条搬运指令 + CopyInSingleKv(kvInUb, startRow, keyOffset0, constInfo); + CopyInSingleKv(kvInUb, startRow + 1, keyOffset1, constInfo); + } else { + DataCopyExtParams intriParams; + if constexpr (!IS_FULL) { + // 尾块:根据实际有效条目数设置blockCount,且此处仅有可能存在keyOffset1为-1的情况 + intriParams.blockCount = 1 + (keyOffset1 >= 0); + } else { + // 非尾块:两条均有效,blockCount恒为2 + intriParams.blockCount = 2; + } + intriParams.blockLen = combineBytes; + intriParams.dstStride = 0; + intriParams.srcStride = keySrcStride; + DataCopyPadExtParams padParams; + padParams.isPad = true; + padParams.leftPadding = 0; + padParams.rightPadding = (CeilAlign(combineBytes, BUFFER_SIZE_BYTE_32B) - combineBytes) / sizeof(KV_T); + padParams.paddingValue = 0; + + int64_t keyOffset; + if constexpr (!IS_FULL) { + keyOffset = (keyOffset1 > -1 && keyOffset1 < keyOffset0) ? keyOffset1 : keyOffset0; + } else { + // 非尾块:两条均有效,取较小地址作为起始 + keyOffset = keyOffset0 < keyOffset1 ? keyOffset0 : keyOffset1; + } + DataCopyPad(kvInUb[startRow * constInfo.dSize], keyGm[keyOffset], intriParams, padParams); + } + startRow += 2; // 每次迭代处理2个输入元素 + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CopyToOutUb(LocalTensor kvOutUb, LocalTensor srcTensor, + int64_t dealRow, ConstInfo &constInfo) +{ + LocalTensor kvNdUb = srcTensor.template ReinterpretCast(); + DataCopy(kvOutUb, kvNdUb, dealRow * constInfo.dSize); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CopyOutKvUb2Gm( + Buffer &v0ResGm, LocalTensor kvOutUb, int64_t dealRow, + int64_t s2StartIdx, const RunInfo &runInfo, ConstInfo &constInfo) +{ + GlobalTensor v0ResGmTensor = v0ResGm.template GetTensor(); + DataCopy(v0ResGmTensor[s2StartIdx * constInfo.dSize], kvOutUb, dealRow * constInfo.dSize); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CalSparseCalSize(const RunInfo &runInfo, ConstInfo &constInfo) +{ + if constexpr (IS_SPLIT_G) { + uint32_t aicIdx = constInfo.aivIdx >> 1U; + uint32_t v0S2SizeFirstCore = CeilDiv(runInfo.s2RealSize, 2); + uint32_t v0S2SizeSecondCore = runInfo.s2RealSize - v0S2SizeFirstCore; + int32_t vecCnt = (aicIdx % 2U == 0) ? + (GetSubBlockIdx() == 0 ? 0 : 1) : + (GetSubBlockIdx() == 0 ? 2 : 3); // 2,3:根据核心索引和子块索引设置处理参数 + if (aicIdx % 2 == 0) { // 2:根据aicIdx的奇偶性来区分不同的处理逻辑 + if (GetSubBlockIdx() == 0) { + sparseCalSize = CeilDiv(v0S2SizeFirstCore, 2); // 2:处理大小为v0S2SizeFirstCore的一半 + sparseS2Start = 0; + } else { + sparseCalSize = v0S2SizeFirstCore - CeilDiv(v0S2SizeFirstCore, 2); // 2:处理剩余部分 + sparseS2Start = CeilDiv(v0S2SizeFirstCore, 2); // 2:起始位置为v0S2SizeFirstCore的一半 + } + } else { + if (GetSubBlockIdx() == 0) { + sparseCalSize = CeilDiv(v0S2SizeSecondCore, 2); // 2:处理大小为v0S2SizeSecondCore的一半 + sparseS2Start = v0S2SizeFirstCore; + } else { + sparseCalSize = v0S2SizeSecondCore - CeilDiv(v0S2SizeSecondCore, 2); // 2:处理剩余部分 + sparseS2Start = + v0S2SizeFirstCore + + CeilDiv(v0S2SizeSecondCore, 2); // 2:起始位置为v0S2SizeFirstCore加上v0S2SizeSecondCore的一半 + } + } + sparseS2End = sparseS2Start + sparseCalSize; + } else { + uint32_t v0S2SizeFirstCore = CeilDiv(runInfo.s2RealSize, 2); // 2:平均分配给两个子块 + sparseCalSize = GetSubBlockIdx() == 0 ? v0S2SizeFirstCore : runInfo.s2RealSize - v0S2SizeFirstCore; + sparseS2Start = GetSubBlockIdx() == 0 ? 0 : v0S2SizeFirstCore; + sparseS2End = sparseS2Start + sparseCalSize; + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::ProcessVec0( + Buffer &v0ResGm, const RunInfo &runInfo, ConstInfo &constInfo) +{ + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE) { + if (runInfo.s2LoopCount < runInfo.oriKvLoopEndIdx) { + return; + } + keyGm = cmpKVGm; + cuSeqlensKvGm = cuSeqlensCmpKvGm; + sparseIndicesGm = cmpSparseIndicesGm; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + blockTableGm = cmpBlockTableGm; + blockSize = constInfo.cmpBlockSize; + maxBlockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch; + } + CalSparseCalSize(runInfo, constInfo); + ProcessSparseKv(v0ResGm, runInfo, constInfo); + v0ResGm.SetCrossCore(); + } else if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (!runInfo.isCmp) { + keyGm = oriKVGm; + cuSeqlensKvGm = cuSeqlensOriKvGm; + sparseIndicesGm = oriSparseIndicesGm; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + blockTableGm = oriBlockTableGm; + blockSize = constInfo.oriBlockSize; + maxBlockNumPerBatch = constInfo.oriMaxBlockNumPerBatch; + } + } else { + keyGm = cmpKVGm; + cuSeqlensKvGm = cuSeqlensCmpKvGm; + sparseIndicesGm = cmpSparseIndicesGm; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + blockTableGm = cmpBlockTableGm; + blockSize = constInfo.cmpBlockSize; + maxBlockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch; + } + } + CalSparseCalSize(runInfo, constInfo); + ProcessSparseKv(v0ResGm, runInfo, constInfo); + v0ResGm.SetCrossCore(); + } +} + +TEMPLATES_DEF_NO_DEFAULT +template +__aicore__ inline void CSABlockVec::CopyIn8Block(LocalTensor kvInUb, int64_t startRow, int64_t &s2, + const RunInfo &runInfo, ConstInfo &constInfo) +{ + // tokenData元素为-1表示无效token,尾块场景下GetReal*仅填充有效区间,其余保持-1 + int64_t tokenData[KV_COPYIN_UNIT] = {-1, -1, -1, -1, -1, -1, -1, -1}; + if constexpr (IS_VEC_S2PHYADDR) { + GetRealS2Addr(tokenData, s2, runInfo, constInfo); + } else { + GetRealCmpS2Idx(tokenData, s2, runInfo, constInfo); + } + s2 += KV_COPYIN_UNIT; + CopyInKvSparse(kvInUb, startRow, tokenData, runInfo, constInfo); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::ProcessSparseKv( + Buffer &v0ResGm, const RunInfo &runInfo, ConstInfo &constInfo) +{ + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + int64_t curSparseCalSize = this->sparseCalSize; + if (curSparseCalSize == 0) { + return; + } + + // 前置计算:16行拷出循环数、尾块行数、尾块内8行搬入次数及剩余行数 + int64_t process16LoopCnt = curSparseCalSize / KV_PROCESS_UNIT; + int64_t tail16Rows = curSparseCalSize % KV_PROCESS_UNIT; + int64_t tail8FullCnt = tail16Rows / KV_COPYIN_UNIT; + int64_t tail8Remain = tail16Rows % KV_COPYIN_UNIT; + int64_t s2 = sparseS2Start; + + // 阶段1:完整16行块(非尾块,8行均有效,无需判断) + for (int64_t i = 0; i < process16LoopCnt; i++) { + int64_t s2StartIdx = sparseS2Start + i * KV_PROCESS_UNIT; + // 1、copy kv in, gm -> ub + LocalTensor stage0OutUb = this->stage0OutBufs[pingPongV0].tensor; + WaitFlag(INNERCORE_STAGE0OUT_MTE3_MTE2(pingPongV0)); + CopyIn8Block(stage0OutUb, 0, s2, runInfo, constInfo); + CopyIn8Block(stage0OutUb, KV_COPYIN_UNIT, s2, runInfo, constInfo); + // 2、copy kv out, ub -> l1 + SetFlag(INNERCORE_STAGE0OUT_MTE2_MTE3(pingPongV0)); + WaitFlag(INNERCORE_STAGE0OUT_MTE2_MTE3(pingPongV0)); + CopyOutKvUb2Gm(v0ResGm, stage0OutUb, KV_PROCESS_UNIT, s2StartIdx, runInfo, constInfo); + SetFlag(INNERCORE_STAGE0OUT_MTE3_MTE2(pingPongV0)); + pingPongV0 ^= 1; + } + + // 阶段2:尾块(不足16行,保留判断逻辑) + if (tail16Rows > 0) { + int64_t s2StartIdx = sparseS2Start + process16LoopCnt * KV_PROCESS_UNIT; + int64_t dealRow = 0; + // 1、copy kv in, gm -> ub + LocalTensor stage0OutUb = this->stage0OutBufs[pingPongV0].tensor; + WaitFlag(INNERCORE_STAGE0OUT_MTE3_MTE2(pingPongV0)); + // 尾块内满8行搬入(8行均有效,无需判断) + for (int64_t j = 0; j < tail8FullCnt; j++) { + CopyIn8Block(stage0OutUb, dealRow, s2, runInfo, constInfo); + dealRow += KV_COPYIN_UNIT; + } + // 尾块内不足8行(需判断边界,与当前逻辑一致) + if (tail8Remain > 0) { + CopyIn8Block(stage0OutUb, dealRow, s2, runInfo, constInfo); + dealRow += tail8Remain; + } + // 2、copy kv out, ub -> l1 + SetFlag(INNERCORE_STAGE0OUT_MTE2_MTE3(pingPongV0)); + WaitFlag(INNERCORE_STAGE0OUT_MTE2_MTE3(pingPongV0)); + CopyOutKvUb2Gm(v0ResGm, stage0OutUb, tail16Rows, s2StartIdx, runInfo, constInfo); + SetFlag(INNERCORE_STAGE0OUT_MTE3_MTE2(pingPongV0)); + pingPongV0 ^= 1; + } + } +} + +TEMPLATES_DEF_NO_DEFAULT +template +__aicore__ inline void CSABlockVec::ComputeVec1Softmax(LocalTensor &stage1CastTensor, + LocalTensor &mmRes, LocalTensor &sumUb, + LocalTensor &maxUb, + LocalTensor &apiTmpBuffer, RunInfo &runInfo, + ConstInfo &constInfo) +{ + if (likely(runInfo.s2RealSize == 128 && runInfo.s2RealSizeUpdate == 128)) { + ProcessVec1Vf( + stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, vselrIndexesBuf, runInfo.halfMRealSize, + runInfo.s2RealSizeUpdate, static_cast(constInfo.softmaxScale), negativeFloatScalar); + } else if (runInfo.s2RealSize <= 64) { + ProcessVec1Vf( + stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, vselrIndexesBuf, runInfo.halfMRealSize, + runInfo.s2RealSizeUpdate, static_cast(constInfo.softmaxScale), negativeFloatScalar); + } else if (runInfo.s2RealSize < 128 || runInfo.s2RealSizeUpdate < 128) { + ProcessVec1Vf( + stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, vselrIndexesBuf, runInfo.halfMRealSize, + runInfo.s2RealSizeUpdate, static_cast(constInfo.softmaxScale), negativeFloatScalar); + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::InitVec1SoftmaxFromSinks(LocalTensor &sumUb, + LocalTensor &maxUb, RunInfo &runInfo, + ConstInfo &constInfo) +{ + bool includeSink = (!runInfo.isCrossCoreSplit) || runInfo.isFirstS2SplitCore; + if constexpr (IS_BATCH_CONSISTENCY) { + includeSink = includeSink && (runInfo.reduceBlockId == 0); + } + if (!includeSink) { + Duplicate(maxUb, this->negativeFloatScalar, runInfo.halfMRealSize); + Duplicate(sumUb, static_cast(0), runInfo.halfMRealSize); + return; + } + int64_t sinksOffset = 0; + if constexpr (!IS_SPLIT_G) { + sinksOffset = GetBlockIdx() % 2 == 0 ? 0 : runInfo.firstHalfMRealSize; // 2:判断块索引的奇偶性 + } else { + sinksOffset = runInfo.goIdx; + if (constInfo.subBlockIdx == 1) { + sinksOffset += runInfo.firstHalfMRealSize; + } + } + LocalTensor sinksUb = this->sinksUb.tensor; + InitSoftmaxFromSinks(sumUb, maxUb, sinksUb, sinksOffset, R0, runInfo.halfMRealSize); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CopyVec1ResultToL1(StaticBuffer &outputBuf, + LocalTensor &stage1CastTensor, + RunInfo &runInfo, ConstInfo &constInfo) +{ + int64_t stage1Offset = runInfo.taskIdMod2; + SetFlag(INNERCORE_STAGE1(stage1Offset)); + WaitFlag(INNERCORE_STAGE1(stage1Offset)); + LocalTensor mm2AL1Tensor = outputBuf.tensor; + if (likely(runInfo.halfMRealSize != 0)) { + DataCopy(mm2AL1Tensor[constInfo.subBlockIdx * (BLOCK_BYTE / sizeof(Q_T)) * + (runInfo.mRealSize - runInfo.halfMRealSize)], + stage1CastTensor, + {s2BaseSize / 16, static_cast(runInfo.halfMRealSize), + static_cast(vec1Srcstride - runInfo.halfMRealSize), + static_cast(Align16Func(runInfo.mRealSize) - runInfo.halfMRealSize)}); + } + SetFlag(INNERCORE_STAGE1(stage1Offset)); + CrossCoreSetFlag(CROSSCORE_L1P(outputBuf.idx)); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::StageCrossCoreVec1Lse(LocalTensor &maxUb, + LocalTensor &sumUb, RunInfo &runInfo, + ConstInfo &constInfo) +{ + AttentionCommon::S2SplitFdStagingLayout stagingLayout = { + constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(false), AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, + AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + LocalTensor tmpUb = this->batchReduceTmpUb.tensor; + AttentionCommon::StageVec1Lse(stagingLayout, crossCoreCombineBase, GetCrossCoreWorkspaceIdx(runInfo), + GetFaStagingMOffset(runInfo, constInfo), runInfo.halfMRealSize, maxUb, sumUb, tmpUb, + INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::StageBatchConsistencyVec1Lse(LocalTensor &maxUb, + LocalTensor &sumUb, + RunInfo &runInfo, ConstInfo &constInfo) +{ + if (!runInfo.isLastBase) { + return; + } + if (runInfo.halfMRealSize > 0) { + LocalTensor finalMaxUb = this->softmaxFinalMaxBufs[runInfo.taskIdMod2].tensor; + LocalTensor finalSumUb = this->softmaxFinalSumBufs[runInfo.taskIdMod2].tensor; + uint64_t snapshotElems = Align8Func(runInfo.halfMRealSize); + DataCopy(finalMaxUb, maxUb, snapshotElems); + DataCopy(finalSumUb, sumUb, snapshotElems); + } + if (runInfo.isCrossCoreSplit && !runInfo.isFirstS2SplitCore) { + StageCrossCoreVec1Lse(maxUb, sumUb, runInfo, constInfo); + } else if (runInfo.isFirstS2SplitCore && runInfo.reduceBlockId == 0 && runInfo.s2LoopCount < runInfo.s2LoopLimit) { + AttentionCommon::S2SplitFdStagingLayout stagingLayout = { + constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(true), AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, + AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + LocalTensor tmpUb = this->batchReduceTmpUb.tensor; + AttentionCommon::StageVec1Lse(stagingLayout, intraCoreCombineBase, GetIntraCoreWorkspaceIdx(runInfo, constInfo), + GetFaStagingMOffset(runInfo, constInfo), runInfo.halfMRealSize, maxUb, sumUb, + tmpUb, INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); + SetFlag(INNERCORE_INTRALSE_MTE3_MTE2(runInfo.multiCoreIdxMod2)); + } else if (runInfo.isCrossCoreSplit && runInfo.isFirstS2SplitCore && runInfo.reduceBlockId == 0) { + StageCrossCoreVec1Lse(maxUb, sumUb, runInfo, constInfo); + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::StageLegacyVec1Lse(LocalTensor &maxUb, + LocalTensor &sumUb, RunInfo &runInfo, + ConstInfo &constInfo) +{ + if (!runInfo.isCrossCoreSplit || runInfo.halfMRealSize <= 0 || runInfo.s2LoopCount != runInfo.s2LoopLimit) { + return; + } + AttentionCommon::S2SplitFdStagingLayout stagingLayout = {constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(), + AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, + AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + LocalTensor tmpUb = this->stage2OutBufs.tensor.template ReinterpretCast(); + AttentionCommon::StageVec1Lse(stagingLayout, fdStagingBase, GetCrossCoreWorkspaceIdx(runInfo), + GetFaStagingMOffset(runInfo, constInfo), static_cast(runInfo.halfMRealSize), + maxUb, sumUb, tmpUb, INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CopyOutVec1Lse(LocalTensor &maxUb, LocalTensor &sumUb, + RunInfo &runInfo, ConstInfo &constInfo) +{ + bool copyOutLse = + constInfo.returnSoftmaxLse && runInfo.halfMRealSize > 0 && runInfo.s2LoopCount == runInfo.s2LoopLimit; + if constexpr (IS_BATCH_CONSISTENCY) { + copyOutLse = copyOutLse && !runInfo.isCrossCoreSplit && !runInfo.needReduce; + } + if (!copyOutLse) { + return; + } + LocalTensor outLse = this->outLseUbs[runInfo.multiCoreIdxMod2].tensor; + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1; + dataCopyParams.blockLen = sizeof(float) * runInfo.halfMRealSize; + dataCopyParams.srcStride = 0; + dataCopyParams.dstStride = 0; + WaitFlag(INNERCORE_LSE_MTE3_V); + ComputeLse(outLse, sumUb, maxUb, runInfo.halfMRealSize); + SetFlag(INNERCORE_LSE_V_MTE3); + WaitFlag(INNERCORE_LSE_V_MTE3); + DataCopyPad(this->softmaxLseGm[runInfo.softmaxLseOffset], outLse, dataCopyParams); + SetFlag(INNERCORE_LSE_MTE3_V); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::ProcessVec1(StaticBuffer &outputBuf, + StaticBuffer &bmm1ResBuf, RunInfo &runInfo, + ConstInfo &constInfo) +{ + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1ResBuf.idx)); + + LocalTensor sumUb = this->softmaxSumBufs[runInfo.multiCoreIdxMod2].tensor; + LocalTensor maxUb = this->softmaxMaxBufs[runInfo.multiCoreIdxMod2].tensor; + LocalTensor expUb = this->softmaxExpBufs[runInfo.taskIdMod2].tensor; + int64_t stage1Offset = runInfo.taskIdMod2; + WaitFlag(INNERCORE_STAGE1(stage1Offset)); + LocalTensor stage1CastTensor = this->stage1OutBufs[stage1Offset].tensor; + + LocalTensor apiTmpBuffer = this->commonUb.tensor; + LocalTensor mmRes = bmm1ResBuf.tensor; + + runInfo.s2RealSizeUpdate = runInfo.s2RealSize; + + bool isFirstSoftmaxBase = runInfo.s2LoopCount == 0; + if constexpr (IS_BATCH_CONSISTENCY) { + isFirstSoftmaxBase = runInfo.isFirstBase; + } + // loopCount = 0 但传入sinks时走update分支,maxUb通过sinks初始化,sumUb初始化为1.0 + if (isFirstSoftmaxBase && !isSinks) { + ComputeVec1Softmax(stage1CastTensor, mmRes, sumUb, maxUb, apiTmpBuffer, runInfo, constInfo); + } else { + if (isFirstSoftmaxBase && isSinks) { + InitVec1SoftmaxFromSinks(sumUb, maxUb, runInfo, constInfo); + } + ComputeVec1Softmax(stage1CastTensor, mmRes, sumUb, maxUb, apiTmpBuffer, runInfo, constInfo); + } + CrossCoreSetFlag(CROSSCORE_BMM1(bmm1ResBuf.idx)); + CopyVec1ResultToL1(outputBuf, stage1CastTensor, runInfo, constInfo); + if (!isFirstSoftmaxBase || isSinks) { + SFAUpdateExpSumAndExpMax(sumUb, maxUb, expUb, sumUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize); + } + if constexpr (IS_BATCH_CONSISTENCY) { + StageBatchConsistencyVec1Lse(maxUb, sumUb, runInfo, constInfo); + } else { + StageLegacyVec1Lse(maxUb, sumUb, runInfo, constInfo); + } + CopyOutVec1Lse(maxUb, sumUb, runInfo, constInfo); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::ReduceIntraBlockAndStage(RunInfo &runInfo, ConstInfo &constInfo, + LocalTensor &vec2ResUb, + LocalTensor &partialTmpUb) +{ + AttentionCommon::S2SplitFdStagingLayout intraLayout = {constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(true), + AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, + AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + AttentionCommon::S2SplitFdStagingLayout crossLayout = {constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(false), + AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, + AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + uint32_t intraWorkspaceIdx = GetIntraCoreWorkspaceIdx(runInfo, constInfo); + uint32_t crossWorkspaceIdx = + static_cast(runInfo.firstFdDataWorkspaceIdx + runInfo.s2SplitIdx - runInfo.reduceBlockId); + int64_t stagingMOffset = GetFaStagingMOffset(runInfo, constInfo); + LocalTensor tmpUb = this->batchReduceTmpUb.tensor; + LocalTensor blockMaxUb = tmpUb; + LocalTensor blockSumUb = tmpUb[256]; + LocalTensor lseBroadcastUb = tmpUb[512]; + LocalTensor sumBroadcastUb = tmpUb[640]; + LocalTensor maxUb = this->softmaxFinalMaxBufs[runInfo.taskIdMod2].tensor; + LocalTensor sumUb = this->softmaxFinalSumBufs[runInfo.taskIdMod2].tensor; + bool copyOutMergedLse = + constInfo.returnSoftmaxLse && !runInfo.isCrossCoreSplit && runInfo.s2LoopCount == runInfo.s2LoopLimit; + + WaitFlag(INNERCORE_INTRALSE_MTE3_MTE2(runInfo.multiCoreIdxMod2)); + WaitFlag(INNERCORE_INTRAATTN_MTE3_MTE2(runInfo.multiCoreIdxMod2)); + LocalTensor sinkUb; + int64_t startRow = 0; + while (startRow < runInfo.vec2MRealSize) { + int64_t dealRowCount = intraLayout.chunkRows; + if (startRow + dealRowCount > runInfo.vec2MRealSize) { + dealRowCount = runInfo.vec2MRealSize - startRow; + } + LocalTensor chunkCurrent = vec2ResUb[startRow * dTemplateAlign64]; + LocalTensor chunkMaxUb = maxUb[startRow]; + LocalTensor chunkSumUb = sumUb[startRow]; + if (copyOutMergedLse) { + WaitFlag(INNERCORE_LSE_MTE3_V); + } + AttentionCommon::MergeStagedAndCurrentChunk( + intraLayout, intraCoreCombineBase, intraWorkspaceIdx, stagingMOffset + startRow, dealRowCount, + static_cast(constInfo.dSizeV), chunkMaxUb, chunkSumUb, chunkCurrent, blockMaxUb, blockSumUb, + partialTmpUb, lseBroadcastUb, sumBroadcastUb, sinkUb, INNERCORE_REDUCE_MAXSUM_V_MTE2, + INNERCORE_INTRAPARTIALO_V_MTE2, INNERCORE_REDUCE_MTE2_V); + + AttentionCommon::StageBroadcastMaxSum(intraLayout, intraCoreCombineBase, intraWorkspaceIdx, + stagingMOffset + startRow, dealRowCount, lseBroadcastUb, sumBroadcastUb, + INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); + if (copyOutMergedLse) { + DataCopyExtParams lseParams; + lseParams.blockCount = static_cast(dealRowCount); + lseParams.blockLen = sizeof(float); + lseParams.srcStride = 0; + lseParams.dstStride = 0; + SetFlag(INNERCORE_LSE_V_MTE3); + WaitFlag(INNERCORE_LSE_V_MTE3); + DataCopyPad(this->softmaxLseGm[runInfo.softmaxLseOffset + startRow], lseBroadcastUb, lseParams); + SetFlag(INNERCORE_LSE_MTE3_V); + } + if (runInfo.isCrossCoreSplit && runInfo.s2LoopCount == runInfo.s2LoopLimit) { + AttentionCommon::StageBroadcastMaxSum(crossLayout, crossCoreCombineBase, crossWorkspaceIdx, + stagingMOffset + startRow, dealRowCount, lseBroadcastUb, + sumBroadcastUb, INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); + } + startRow += intraLayout.chunkRows; + } + + AttentionCommon::StageVec2PartialOAndWait(intraLayout, intraCoreCombineGm, intraWorkspaceIdx, stagingMOffset, + runInfo.vec2MRealSize, static_cast(constInfo.dSizeV), + vec2ResUb, INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); + if (runInfo.isCrossCoreSplit && runInfo.s2LoopCount == runInfo.s2LoopLimit) { + AttentionCommon::StageVec2PartialOAndWait(crossLayout, crossCoreCombineGm, crossWorkspaceIdx, stagingMOffset, + runInfo.vec2MRealSize, static_cast(constInfo.dSizeV), + vec2ResUb, INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); + } + if (runInfo.s2LoopCount < runInfo.s2LoopLimit) { + SetFlag(INNERCORE_INTRALSE_MTE3_MTE2(runInfo.multiCoreIdxMod2)); + SetFlag(INNERCORE_INTRAATTN_MTE3_MTE2(runInfo.multiCoreIdxMod2)); + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::ProcessVec2(StaticBuffer &bmm2ResBuf, RunInfo &runInfo, + ConstInfo &constInfo) +{ + CrossCoreWaitFlag(CROSSCORE_BMM2); + if (unlikely(runInfo.vec2MBaseSize == 0)) { + CrossCoreSetFlag(CROSSCORE_BMM2); + return; + } + + runInfo.vec2MRealSize = runInfo.vec2MBaseSize; + int64_t vec2CalcSize = runInfo.vec2MRealSize * dTemplateAlign64; + LocalTensor vec2ResUb = this->stage2OutBufs.tensor; + LocalTensor mmRes = bmm2ResBuf.tensor; + WaitFlag(INNERCORE_STAGE2); + bool needIntraBlockReduce = false; + if constexpr (IS_BATCH_CONSISTENCY) { + needIntraBlockReduce = runInfo.isLastBase && runInfo.isFirstS2SplitCore && runInfo.reduceBlockId > 0; + if (needIntraBlockReduce) { + WaitFlag(INNERCORE_INTRAPARTIALO_V_MTE2); + WaitFlag(INNERCORE_REDUCE_MAXSUM_V_MTE2); + } + } + bool isFirstVec2Base = runInfo.s2LoopCount == 0; + if constexpr (IS_BATCH_CONSISTENCY) { + isFirstVec2Base = runInfo.isFirstBase; + } + if (unlikely(isFirstVec2Base)) { + DataCopy(vec2ResUb, mmRes, vec2CalcSize); + } else { + if (runInfo.s2RealSizeUpdate > 0) { + LocalTensor expUb = softmaxExpBufs[runInfo.taskIdMod2].tensor; + bool isLastVec2Base = (runInfo.s2LoopCount == runInfo.s2LoopLimit); + if constexpr (IS_BATCH_CONSISTENCY) { + isLastVec2Base = runInfo.isLastBase; + } + if (isLastVec2Base) { + LocalTensor sumUb; + if constexpr (IS_BATCH_CONSISTENCY) { + sumUb = this->softmaxFinalSumBufs[runInfo.taskIdMod2].tensor; + } else { + sumUb = this->softmaxSumBufs[runInfo.multiCoreIdxMod2].tensor; + } + FlashUpdateLastNew( + vec2ResUb, mmRes, vec2ResUb, expUb, expUb, sumUb, runInfo.vec2MRealSize, dTemplateAlign64, 1.0, + 1.0); + } else { + FlashUpdateNew( + vec2ResUb, mmRes, vec2ResUb, expUb, expUb, runInfo.vec2MRealSize, dTemplateAlign64, 1.0, 1.0); + } + } else { + bool isLastVec2Base = runInfo.s2LoopCount >= runInfo.s2LoopLimit; + if constexpr (IS_BATCH_CONSISTENCY) { + isLastVec2Base = runInfo.isLastBase; + } + if (isLastVec2Base) { + LocalTensor sumUb; + if constexpr (IS_BATCH_CONSISTENCY) { + sumUb = this->softmaxFinalSumBufs[runInfo.taskIdMod2].tensor; + } else { + sumUb = this->softmaxSumBufs[runInfo.multiCoreIdxMod2].tensor; + } + LastDivNew(vec2ResUb, vec2ResUb, sumUb, + runInfo.vec2MRealSize, dTemplateAlign64, 1.0); + } + } + } + + if constexpr (IS_BATCH_CONSISTENCY) { + if (runInfo.isLastBase) { + if (unlikely(isFirstVec2Base)) { + LocalTensor sumUb = this->softmaxFinalSumBufs[runInfo.taskIdMod2].tensor; + LastDivNew(vec2ResUb, vec2ResUb, sumUb, + runInfo.vec2MRealSize, dTemplateAlign64, 1.0); + } + if (needIntraBlockReduce) { + SetFlag(INNERCORE_INTRAPARTIALO_V_MTE2); + SetFlag(INNERCORE_REDUCE_MAXSUM_V_MTE2); + ReduceIntraBlockAndStage(runInfo, constInfo, vec2ResUb, mmRes); + if (!runInfo.isCrossCoreSplit && runInfo.s2LoopCount == runInfo.s2LoopLimit) { + this->CopyOutAttentionOut(runInfo, constInfo, vec2ResUb, 0, vec2CalcSize); + } + } else { + if (runInfo.isFirstS2SplitCore && runInfo.reduceBlockId == 0 && + runInfo.s2LoopCount < runInfo.s2LoopLimit) { + AttentionCommon::S2SplitFdStagingLayout stagingLayout = { + constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(true), + AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + int64_t stagingMOffset = GetFaStagingMOffset(runInfo, constInfo); + AttentionCommon::StageVec2PartialOAndWait( + stagingLayout, intraCoreCombineGm, GetIntraCoreWorkspaceIdx(runInfo, constInfo), stagingMOffset, + runInfo.vec2MRealSize, static_cast(constInfo.dSizeV), vec2ResUb, INNERCORE_STAGE2, + INNERCORE_STAGE_FD_MTE3_V); + SetFlag(INNERCORE_INTRAATTN_MTE3_MTE2(runInfo.multiCoreIdxMod2)); + } + if (runInfo.isCrossCoreSplit && + (!runInfo.isFirstS2SplitCore || (runInfo.isFirstS2SplitCore && runInfo.reduceBlockId == 0 && + runInfo.s2LoopCount == runInfo.s2LoopLimit))) { + AttentionCommon::S2SplitFdStagingLayout stagingLayout = { + constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(false), + AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + uint32_t workspaceIdx = GetCrossCoreWorkspaceIdx(runInfo); + int64_t stagingMOffset = GetFaStagingMOffset(runInfo, constInfo); + AttentionCommon::StageVec2PartialOAndWait(stagingLayout, crossCoreCombineGm, workspaceIdx, + stagingMOffset, runInfo.vec2MRealSize, + static_cast(constInfo.dSizeV), vec2ResUb, + INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); + } else if (!runInfo.isCrossCoreSplit && runInfo.s2LoopCount == runInfo.s2LoopLimit) { + this->CopyOutAttentionOut(runInfo, constInfo, vec2ResUb, 0, vec2CalcSize); + } + } + } + } else if (runInfo.s2LoopCount == runInfo.s2LoopLimit) { + if (unlikely(runInfo.s2LoopCount == 0)) { + LocalTensor sumUb = this->softmaxSumBufs[runInfo.multiCoreIdxMod2].tensor; + LastDivNew(vec2ResUb, vec2ResUb, sumUb, runInfo.vec2MRealSize, + dTemplateAlign64, 1.0); + } + if (runInfo.isCrossCoreSplit) { + AttentionCommon::S2SplitFdStagingLayout stagingLayout = { + constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(), AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, + AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + uint32_t workspaceIdx = GetCrossCoreWorkspaceIdx(runInfo); + int64_t stagingMOffset = GetFaStagingMOffset(runInfo, constInfo); + AttentionCommon::StageVec2PartialO( + stagingLayout, stagingOutGm, workspaceIdx, stagingMOffset, static_cast(runInfo.vec2MRealSize), + static_cast(constInfo.dSizeV), vec2ResUb, INNERCORE_STAGE2, INNERCORE_STAGE_FD_MTE3_V); + } else { + this->CopyOutAttentionOut(runInfo, constInfo, vec2ResUb, 0, vec2CalcSize); + } + } + CrossCoreSetFlag(CROSSCORE_BMM2); + SetFlag(INNERCORE_STAGE2); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::InitFDBuffers(FdRunInfo &fdRunInfo) +{ + FdRunInfo fdBufferInfo = fdRunInfo; + if (fdBufferInfo.mNum > AttentionCommon::FD_REDUCE_CHUNK_ROWS) { + fdBufferInfo.mNum = AttentionCommon::FD_REDUCE_CHUNK_ROWS; + } + AttentionCommon::InitFDBuffersStatic(fdBufferInfo, 0, fdBuffers); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::ProcessFlashDecode(FdRunInfo &fdRunInfo, ConstInfo &constInfo) +{ + InitFDBuffers(fdRunInfo); + int64_t seqOffset = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + seqOffset = this->cuSeqlensQGm.GetValue(fdRunInfo.bn2Idx); + } else { + seqOffset = fdRunInfo.bn2Idx * constInfo.s1Size; + } + int64_t attentionOutOffset = + seqOffset * constInfo.n2GDv + fdRunInfo.mIdx * constInfo.n2GDv + fdRunInfo.mStartIdx * constInfo.dSizeV; + int64_t softmaxLseOffset = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + softmaxLseOffset = (seqOffset + fdRunInfo.mIdx) * constInfo.gSize + fdRunInfo.mStartIdx; + } else { + softmaxLseOffset = + (fdRunInfo.bn2Idx * constInfo.s1Size + fdRunInfo.mIdx) * constInfo.gSize + fdRunInfo.mStartIdx; + } + LocalTensor accumulatedO = this->fdBuffers.accumOut.tensor.template ReinterpretCast(); + LocalTensor lseExpUb = this->fdBuffers.lseExp.tensor.template ReinterpretCast(); + LocalTensor blockMaxUb = this->fdBuffers.blockMax.tensor.template ReinterpretCast(); + LocalTensor blockSumUb = this->fdBuffers.blockSum.tensor.template ReinterpretCast(); + LocalTensor partialOFp32 = this->fdBuffers.partialO.tensor.template ReinterpretCast(); + AttentionCommon::S2SplitFdStagingLayout stagingLayout = {constInfo.gSize, dTemplateAlign64, GetStagingSlotNum(), + AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW, + AttentionCommon::FD_REDUCE_CHUNK_ROWS}; + int64_t attentionOutRowStride = + static_cast(constInfo.dSizeV) + static_cast(constInfo.attentionOutStride) / sizeof(OUTPUT_T); + int64_t startRow = 0; + while (startRow < fdRunInfo.mNum) { + int64_t dealRowCount = AttentionCommon::FD_REDUCE_CHUNK_ROWS; + if (startRow + dealRowCount > fdRunInfo.mNum) { + dealRowCount = fdRunInfo.mNum - startRow; + } + WaitFlag(INNERCORE_FD_MTE3_V); + if constexpr (IS_BATCH_CONSISTENCY) { + WaitFlag(INNERCORE_FD_MTE3_MTE2); + AttentionCommon::ReducePairwiseWithLse( + stagingLayout, fdStagingBase, fdRunInfo.workspaceIdx, fdRunInfo.workspaceNum, + static_cast(fdRunInfo.mStartIdx + startRow), dealRowCount, + static_cast(constInfo.dSizeV), accumulatedO, lseExpUb, blockMaxUb, blockSumUb, partialOFp32, + constInfo.returnSoftmaxLse, softmaxLseGm, softmaxLseOffset + startRow, INNERCORE_FD_V_MTE2(0), + INNERCORE_FD_V_MTE2(1), INNERCORE_FD_MTE2_V, INNERCORE_LSE_V_MTE3, INNERCORE_LSE_MTE3_V); + } else { + AttentionCommon::ReduceWithLse( + stagingLayout, fdStagingBase, fdRunInfo.workspaceIdx, fdRunInfo.workspaceNum, + static_cast(fdRunInfo.mStartIdx + startRow), dealRowCount, + static_cast(constInfo.dSizeV), accumulatedO, lseExpUb, blockMaxUb, blockSumUb, partialOFp32, + constInfo.returnSoftmaxLse, softmaxLseGm, softmaxLseOffset + startRow, INNERCORE_FD_V_MTE2(0), + INNERCORE_FD_V_MTE2(1), INNERCORE_FD_MTE2_V, INNERCORE_LSE_V_MTE3, INNERCORE_LSE_MTE3_V); + } + RunInfo runInfo; + runInfo.vec2MRealSize = dealRowCount; + runInfo.attentionOutOffset = attentionOutOffset + startRow * attentionOutRowStride; + int64_t vec2CalcSize = dealRowCount * dTemplateAlign64; + this->CopyOutAttentionOut(runInfo, constInfo, accumulatedO, 0, vec2CalcSize); + if constexpr (IS_BATCH_CONSISTENCY) { + SetFlag(INNERCORE_FD_MTE3_MTE2); + } + SetFlag(INNERCORE_FD_MTE3_V); + startRow += dealRowCount; + } +} + +TEMPLATES_DEF_NO_DEFAULT +template +__aicore__ inline void CSABlockVec::Bmm2DataCopyOut(RunInfo &runInfo, ConstInfo &constInfo, + LocalTensor &vec2ResUb, + int64_t vec2S1Idx, int64_t vec2CalcSize) +{ + LocalTensor attenOut; + int64_t dSizeAligned64 = (int64_t)dTemplateAlign64; + + attenOut.SetAddr(vec2ResUb.address_); + Cast(attenOut, vec2ResUb, RoundMode::CAST_ROUND, vec2CalcSize); + SetFlag(INNERCORE_STAGE2); + WaitFlag(INNERCORE_STAGE2); + + DataCopyExtParams dataCopyParams; + dataCopyParams.blockLen = constInfo.dSizeV * sizeof(OUTPUT_T); + dataCopyParams.srcStride = (dSizeAligned64 - constInfo.dSizeV) >> 4; // 以32B为单位偏移,bf16类型即偏移16个数,右移4 + dataCopyParams.dstStride = constInfo.attentionOutStride; + dataCopyParams.blockCount = runInfo.vec2MRealSize; + + DataCopyPad(this->attentionOutGm[runInfo.attentionOutOffset], attenOut, dataCopyParams); +} + +TEMPLATES_DEF_NO_DEFAULT +template +__aicore__ inline void CSABlockVec::CopyOutAttentionOut(RunInfo &runInfo, ConstInfo &constInfo, + LocalTensor &vec2ResUb, + int64_t vec2S1Idx, int64_t vec2CalcSize) +{ + this->Bmm2DataCopyOut(runInfo, constInfo, vec2ResUb, vec2S1Idx, vec2CalcSize); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::InitOutputSingleCore(ConstInfo &constInfo) +{ + uint32_t coreNum = GetBlockNum(); + uint32_t vecCoreNum = CV_RATIO * coreNum; + uint64_t totalOutputSize = 0; + + // n2 = 1, n1 = gn2 = gSize + if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND) { + totalOutputSize = constInfo.bSize * constInfo.gSize * constInfo.s1Size * constInfo.dSizeV; + } else if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + totalOutputSize = constInfo.s1Size * constInfo.gSize * constInfo.dSizeV; + } + + static constexpr uint32_t ATTEN_OUT_POP_BUF_START_ADDR = 184U * 1024U; + static constexpr uint32_t ATTEN_OUT_POP_BUF_ELE_SIZE = (32U * 1024U) / sizeof(OUTPUT_T); + if (coreNum != 0 && totalOutputSize > 0) { + AttentionCommon::InitOutput(this->attentionOutGm, totalOutputSize, + vecCoreNum, static_cast(0)); + } + if (constInfo.returnSoftmaxLse) { + uint64_t totalReturnSoftmaxSize = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND) { + totalReturnSoftmaxSize = constInfo.bSize * constInfo.n2Size * constInfo.s1Size * constInfo.gSize; + } else if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + totalReturnSoftmaxSize = constInfo.n2Size * constInfo.s1Size * constInfo.gSize; // (N2,T1,G) + } + static constexpr uint32_t LSE_POP_BUF_START_ADDR = 216U * 1024U; + static constexpr uint32_t LSE_POP_BUF_ELE_SIZE = (32U * 1024U) / sizeof(float); + if (coreNum != 0 && totalReturnSoftmaxSize > 0) { + AttentionCommon::InitOutput( + this->softmaxLseGm, totalReturnSoftmaxSize, vecCoreNum, static_cast(0)); + } + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CleanOutput(__gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, + ConstInfo &constInfo) +{ + if ASCEND_IS_AIV { + this->attentionOutGm.SetGlobalBuffer((__gm__ OUTPUT_T *)attentionOut); + this->softmaxLseGm.SetGlobalBuffer((__gm__ T *)softmaxLse); + if (constInfo.needInit == 1) { + InitOutputSingleCore(constInfo); + } + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::InitGlobalBuffer( + __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sinks, + __gm__ uint8_t *sequsedOriKv, __gm__ uint8_t *sequsedCmpKv, __gm__ uint8_t *cmpResidualKv) +{ + oriKVGm.SetGlobalBuffer((__gm__ KV_T *)(oriKV)); + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + oriBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)oriBlockTable); + } + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + cmpKVGm.SetGlobalBuffer((__gm__ KV_T *)cmpKV); + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + cmpBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)cmpBlockTable); + } + cmpSparseIndicesGm.SetGlobalBuffer((__gm__ int32_t *)cmpSparseIndices); + } + + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + oriSparseIndicesGm.SetGlobalBuffer((__gm__ int32_t *)oriSparseIndices); + } + + if (sinks != nullptr) { + sinksGm.SetGlobalBuffer((__gm__ T *)sinks); + this->isSinks = true; + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::SoftmaxInitBuffer(uint32_t &ubAddr) +{ + constexpr uint32_t softmaxBufSize = 256; // VF单次操作256Byte + constexpr uint32_t softmaxElems = softmaxBufSize / sizeof(float); + softmaxSumBufs[0] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxElems), 0}; + ubAddr += softmaxBufSize; + softmaxSumBufs[1] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxElems), 1}; + ubAddr += softmaxBufSize; + softmaxMaxBufs[0] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxElems), 0}; + ubAddr += softmaxBufSize; + softmaxMaxBufs[1] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxElems), 1}; + ubAddr += softmaxBufSize; + if constexpr (IS_BATCH_CONSISTENCY) { + softmaxFinalSumBufs[0] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxElems), 0}; + ubAddr += softmaxBufSize; + softmaxFinalSumBufs[1] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxElems), 1}; + ubAddr += softmaxBufSize; + softmaxFinalMaxBufs[0] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxElems), 0}; + ubAddr += softmaxBufSize; + softmaxFinalMaxBufs[1] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxElems), 1}; + ubAddr += softmaxBufSize; + batchReduceTmpUb = {LocalTensor(TPosition::VECIN, ubAddr, 768), + 0}; // 768:batchReduceTmpUb申请内存大小为768个float + ubAddr += 768U * sizeof(float); + } + softmaxExpBufs[0] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxBufSize / sizeof(T)), 0}; + ubAddr += softmaxBufSize; + softmaxExpBufs[1] = {LocalTensor(TPosition::VECIN, ubAddr, softmaxBufSize / sizeof(T)), 1}; + ubAddr += softmaxBufSize; +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::InitSinksBuffer(ConstInfo &constInfo) +{ + LocalTensor sinksUb = this->sinksUb.tensor; + const uint32_t maxN = constInfo.gSize; // N最大支持128, sink shape是[N] + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1U; + dataCopyParams.blockLen = maxN * sizeof(T); + dataCopyParams.srcStride = 0U; + dataCopyParams.dstStride = 0U; + DataCopyPadExtParams padParams; + DataCopyPad(sinksUb, this->sinksGm, dataCopyParams, padParams); + SetFlag(INNERCORE_SINKS_SYNC); + WaitFlag(INNERCORE_SINKS_SYNC); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::InitLocalBuffer(ConstInfo &constInfo, uint32_t ubBaseAddr) +{ + uint32_t ubAddr = ubBaseAddr; + + SoftmaxInitBuffer(ubAddr); + + commonUb = {LocalTensor(TPosition::VECIN, ubAddr, 512 / sizeof(T)), 0}; // 512 for common ub size + ubAddr += 512; // 512 for common ub offset + sinksUb = {LocalTensor(TPosition::VECIN, ubAddr, 512 / sizeof(T)), 0}; // 512 for sinks ub size + ubAddr += 512; // 512 for sinks ub offset + if (this->isSinks) { + InitSinksBuffer(constInfo); + } + + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + stage0OutBufs[0] = {LocalTensor(TPosition::VECIN, ubAddr, dVTemplateType * 16U), + 0}; // 输出缓冲区处理16个seq + ubAddr += dVTemplateType * 16U * sizeof(KV_T); + stage0OutBufs[1] = {LocalTensor(TPosition::VECIN, ubAddr, dVTemplateType * 16U), + 1}; // 输出缓冲区处理16个seq + ubAddr += dVTemplateType * 16U * sizeof(KV_T); + } + if (constInfo.returnSoftmaxLse) { + outLseUbs[0] = {LocalTensor(TPosition::VECIN, ubAddr, 256 / sizeof(float)), + 0}; // outLseBuf[0]内存申请256B + ubAddr += 256U; + outLseUbs[1] = {LocalTensor(TPosition::VECIN, ubAddr, 256 / sizeof(float)), + 1}; // outLseBuf[1]内存申请256B + ubAddr += 256U; + } + + stage1OutBufs[0] = {LocalTensor(TPosition::VECIN, ubAddr, vec1Srcstride * s2BaseSize), 0}; + ubAddr += vec1Srcstride * s2BaseSize * sizeof(Q_T); + stage1OutBufs[1] = {LocalTensor(TPosition::VECIN, ubAddr, vec1Srcstride * s2BaseSize), 1}; + ubAddr += vec1Srcstride * s2BaseSize * sizeof(Q_T); + + stage2OutBufs = {LocalTensor(TPosition::VECIN, ubAddr, (s1BaseSize / CV_RATIO) * dTemplateAlign64), 0}; + + // 显式 flag 初始化 (替代 AllocEventID + 初始 SetFlag) + SetFlag(INNERCORE_STAGE2); + if constexpr (IS_BATCH_CONSISTENCY) { + SetFlag(INNERCORE_INTRAPARTIALO_V_MTE2); + SetFlag(INNERCORE_REDUCE_MAXSUM_V_MTE2); + } + if (constInfo.returnSoftmaxLse) { + SetFlag(INNERCORE_LSE_MTE3_V); + } + SetFlag(INNERCORE_STAGE0OUT_MTE3_MTE2(0)); + SetFlag(INNERCORE_STAGE0OUT_MTE3_MTE2(1)); + SetFlag(INNERCORE_STAGE1(0)); + SetFlag(INNERCORE_STAGE1(1)); + SetFlag(INNERCORE_FD_V_MTE2(0)); + SetFlag(INNERCORE_FD_V_MTE2(1)); + SetFlag(INNERCORE_FD_MTE3_V); + if constexpr (IS_BATCH_CONSISTENCY) { + SetFlag(INNERCORE_FD_MTE3_MTE2); + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::FreeEvent(ConstInfo &constInfo) +{ + if constexpr (IS_BATCH_CONSISTENCY) { + WaitFlag(INNERCORE_INTRAPARTIALO_V_MTE2); + WaitFlag(INNERCORE_REDUCE_MAXSUM_V_MTE2); + } + WaitFlag(INNERCORE_STAGE2); + if (constInfo.returnSoftmaxLse) { + WaitFlag(INNERCORE_LSE_MTE3_V); + } + WaitFlag(INNERCORE_STAGE0OUT_MTE3_MTE2(0)); + WaitFlag(INNERCORE_STAGE0OUT_MTE3_MTE2(1)); + WaitFlag(INNERCORE_STAGE1(0)); + WaitFlag(INNERCORE_STAGE1(1)); + WaitFlag(INNERCORE_FD_V_MTE2(0)); + WaitFlag(INNERCORE_FD_V_MTE2(1)); + WaitFlag(INNERCORE_FD_MTE3_V); + if constexpr (IS_BATCH_CONSISTENCY) { + WaitFlag(INNERCORE_FD_MTE3_MTE2); + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::GetExtremeValue(T &negativeScalar) +{ + uint32_t tmp1 = NEGATIVE_MIN_VAULE_FP32; + negativeScalar = *((float *)&tmp1); +} + +template +__simd_vf__ void GetKVPhyAddrVFPaImpl(__ubuf__ uint32_t *kvPhyAddrUb, __ubuf__ int32_t *sparseIdxUb, + __ubuf__ int32_t *blkTableUb, const uint16_t s2Loop, uint32_t s2Tail, + const uint32_t blockSize, const int16_t shiftRightNum, + const uint32_t sparseBlockSize, const uint32_t kvDim, const uint32_t kvStride) +{ + static const uint16_t s2_num_per_loop = 128; + static const uint16_t s2_num_per_reg = 64; + static const uint16_t out_offset_per_loop = 256; + static const uint16_t out_offset_per_reg = 128; + static const uint32_t invalid_value = 0xFFFFFFFF; + Reg::MaskReg preg_all_b32 = Reg::CreateMask(); + Reg::MaskReg add_carry_l_1; + Reg::MaskReg add_carry_h_1; + Reg::MaskReg add_carry_l_2; + Reg::MaskReg add_carry_h_2; + Reg::MaskReg preg_tail_neg_1_b32; + Reg::MaskReg preg_tail_neg_2_b32; + + Reg::RegTensor vreg_kv_stride; + Reg::RegTensor vreg_sparse_idx_1; + Reg::RegTensor vreg_sparse_idx_2; + Reg::RegTensor vreg_block_size; + Reg::RegTensor vreg_shift_rights_num; + Reg::RegTensor vreg_pa_blk_idx_1; + Reg::RegTensor vreg_pa_blk_idx_2; + Reg::RegTensor vreg_pa_tmp_1; + Reg::RegTensor vreg_pa_tmp_2; + Reg::RegTensor vreg_pa_offset_1; + Reg::RegTensor vreg_pa_offset_2; + Reg::RegTensor vreg_phy_offset_1; + Reg::RegTensor vreg_phy_offset_2; + Reg::RegTensor vreg_phy_blk_idx_1; + Reg::RegTensor vreg_phy_blk_idx_2; + + Reg::RegTensor vreg_blk_id_mul_stride_h_1; + Reg::RegTensor vreg_blk_id_mul_stride_tmp_h_1; + Reg::RegTensor vreg_blk_id_mul_stride_l_1; + Reg::RegTensor vreg_mul_overflow_l_1; + Reg::RegTensor vreg_total_offset_l_1; + Reg::RegTensor vreg_total_offset_h_1; + + Reg::RegTensor vreg_blk_id_mul_stride_h_2; + Reg::RegTensor vreg_blk_id_mul_stride_tmp_h_2; + Reg::RegTensor vreg_blk_id_mul_stride_l_2; + Reg::RegTensor vreg_mul_overflow_l_2; + Reg::RegTensor vreg_total_offset_l_2; + Reg::RegTensor vreg_total_offset_h_2; + + Reg::RegTensor vreg_zero; + Reg::Duplicate(vreg_zero, 0); + Reg::Duplicate(vreg_kv_stride, kvStride); + + for (; s2Loop > 1;) { + for (uint16_t i = 0; i < s2Loop - 1; i++) { + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_1, + sparseIdxUb + i * s2_num_per_loop); + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_2, + sparseIdxUb + s2_num_per_reg + i * s2_num_per_loop); + // * sparseBlockSize + Reg::Muls(vreg_sparse_idx_1, vreg_sparse_idx_1, sparseBlockSize, preg_all_b32); + Reg::Muls(vreg_sparse_idx_2, vreg_sparse_idx_2, sparseBlockSize, preg_all_b32); + // 计算右移位数 + // 右移 -> 除blockSize 得到paBlockIdx,vreg_sparse_idx - pa_idx * blocksize -> pa offset + Reg::ShiftRights(vreg_pa_blk_idx_1, vreg_sparse_idx_1, shiftRightNum, preg_all_b32); + Reg::ShiftRights(vreg_pa_blk_idx_2, vreg_sparse_idx_2, shiftRightNum, preg_all_b32); + + Reg::Muls(vreg_pa_tmp_1, vreg_pa_blk_idx_1, blockSize, preg_all_b32); + Reg::Muls(vreg_pa_tmp_2, vreg_pa_blk_idx_2, blockSize, preg_all_b32); + // offset + Reg::Sub(vreg_pa_offset_1, vreg_sparse_idx_1, vreg_pa_tmp_1, preg_all_b32); + Reg::Sub(vreg_pa_offset_2, vreg_sparse_idx_2, vreg_pa_tmp_2, preg_all_b32); + // 物理页内offset + Reg::Muls(vreg_phy_offset_1, vreg_pa_offset_1, kvDim, preg_all_b32); + Reg::Muls(vreg_phy_offset_2, vreg_pa_offset_2, kvDim, preg_all_b32); + + // int32 paBlockId -> 物理id + DataCopyGather(vreg_phy_blk_idx_1, blkTableUb, vreg_pa_blk_idx_1, preg_all_b32); + DataCopyGather(vreg_phy_blk_idx_2, blkTableUb, vreg_pa_blk_idx_2, preg_all_b32); + + // 分高低32位计算int64物理地址 -- 乘 stride + // 低位乘 带进位 + Reg::Mull(vreg_blk_id_mul_stride_l_1, vreg_mul_overflow_l_1, vreg_phy_blk_idx_1, vreg_kv_stride, + preg_all_b32); + Reg::Mull(vreg_blk_id_mul_stride_l_2, vreg_mul_overflow_l_2, vreg_phy_blk_idx_2, vreg_kv_stride, + preg_all_b32); + + // 分高低32位计算int64物理地址 -- 加 offset + Reg::Add(add_carry_l_1, vreg_total_offset_l_1, vreg_blk_id_mul_stride_l_1, vreg_phy_offset_1, preg_all_b32); + Reg::Add(add_carry_l_2, vreg_total_offset_l_2, vreg_blk_id_mul_stride_l_2, vreg_phy_offset_2, preg_all_b32); + + Reg::AddC(add_carry_h_1, vreg_total_offset_h_1, vreg_mul_overflow_l_1, vreg_zero, add_carry_l_1, + preg_all_b32); + Reg::AddC(add_carry_h_2, vreg_total_offset_h_2, vreg_mul_overflow_l_2, vreg_zero, add_carry_l_2, + preg_all_b32); + + // 搬出 由于拆分为了int32类型,元素个数翻倍 + Reg::StoreAlign( + kvPhyAddrUb + i * out_offset_per_loop, vreg_total_offset_l_1, vreg_total_offset_h_1, preg_all_b32); + Reg::StoreAlign( + kvPhyAddrUb + out_offset_per_reg + i * out_offset_per_loop, vreg_total_offset_l_2, + vreg_total_offset_h_2, preg_all_b32); + } + break; + } + + for (uint16_t i = s2Loop - 1; i < s2Loop; i++) { + Reg::MaskReg preg_tail_1_b32 = Reg::UpdateMask(s2Tail); + Reg::MaskReg preg_tail_2_b32 = Reg::UpdateMask(s2Tail); + Reg::Not(preg_tail_neg_1_b32, preg_tail_1_b32, preg_all_b32); + Reg::Not(preg_tail_neg_2_b32, preg_tail_2_b32, preg_all_b32); + + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_1, + sparseIdxUb + i * s2_num_per_loop); + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_2, + sparseIdxUb + s2_num_per_reg + i * s2_num_per_loop); + // * sparseBlockSize + Reg::Muls(vreg_sparse_idx_1, vreg_sparse_idx_1, sparseBlockSize, preg_tail_1_b32); + Reg::Muls(vreg_sparse_idx_2, vreg_sparse_idx_2, sparseBlockSize, preg_tail_2_b32); + // 计算右移位数 + // 右移 -> 除blockSize 得到paBlockIdx,vreg_sparse_idx - pa_idx * blocksize -> pa offset + Reg::ShiftRights(vreg_pa_blk_idx_1, vreg_sparse_idx_1, shiftRightNum, preg_tail_1_b32); + Reg::ShiftRights(vreg_pa_blk_idx_2, vreg_sparse_idx_2, shiftRightNum, preg_tail_2_b32); + + Reg::Muls(vreg_pa_tmp_1, vreg_pa_blk_idx_1, blockSize, preg_tail_1_b32); + Reg::Muls(vreg_pa_tmp_2, vreg_pa_blk_idx_2, blockSize, preg_tail_2_b32); + // offset + Reg::Sub(vreg_pa_offset_1, vreg_sparse_idx_1, vreg_pa_tmp_1, preg_tail_1_b32); + Reg::Sub(vreg_pa_offset_2, vreg_sparse_idx_2, vreg_pa_tmp_2, preg_tail_2_b32); + // 物理页内offset + Reg::Muls(vreg_phy_offset_1, vreg_pa_offset_1, kvDim, preg_tail_1_b32); + Reg::Muls(vreg_phy_offset_2, vreg_pa_offset_2, kvDim, preg_tail_2_b32); + + // int32 paBlockId -> 物理id + DataCopyGather(vreg_phy_blk_idx_1, blkTableUb, vreg_pa_blk_idx_1, preg_tail_1_b32); + DataCopyGather(vreg_phy_blk_idx_2, blkTableUb, vreg_pa_blk_idx_2, preg_tail_2_b32); + + // 分高低32位计算int64物理地址 -- 乘 stride + // 低位乘 带进位 + Reg::Mull(vreg_blk_id_mul_stride_l_1, vreg_mul_overflow_l_1, vreg_phy_blk_idx_1, vreg_kv_stride, + preg_tail_1_b32); + Reg::Mull(vreg_blk_id_mul_stride_l_2, vreg_mul_overflow_l_2, vreg_phy_blk_idx_2, vreg_kv_stride, + preg_tail_2_b32); + + // 分高低32位计算int64物理地址 -- 加 offset + Reg::Add(add_carry_l_1, vreg_total_offset_l_1, vreg_blk_id_mul_stride_l_1, vreg_phy_offset_1, preg_tail_1_b32); + Reg::Add(add_carry_l_2, vreg_total_offset_l_2, vreg_blk_id_mul_stride_l_2, vreg_phy_offset_2, preg_tail_2_b32); + + Reg::AddC(add_carry_h_1, vreg_total_offset_h_1, vreg_mul_overflow_l_1, vreg_zero, add_carry_l_1, + preg_tail_1_b32); + Reg::AddC(add_carry_h_2, vreg_total_offset_h_2, vreg_mul_overflow_l_2, vreg_zero, add_carry_l_2, + preg_tail_2_b32); + + // 无效值填充-1(0xFFFFFFFF) + Reg::Duplicate(vreg_total_offset_l_1, invalid_value, + preg_tail_neg_1_b32); + Reg::Duplicate(vreg_total_offset_h_1, invalid_value, + preg_tail_neg_1_b32); + Reg::Duplicate(vreg_total_offset_l_2, invalid_value, + preg_tail_neg_2_b32); + Reg::Duplicate(vreg_total_offset_h_2, invalid_value, + preg_tail_neg_2_b32); + Reg::StoreAlign( + kvPhyAddrUb + i * out_offset_per_loop, vreg_total_offset_l_1, vreg_total_offset_h_1, preg_all_b32); + Reg::StoreAlign( + kvPhyAddrUb + out_offset_per_reg + i * out_offset_per_loop, vreg_total_offset_l_2, vreg_total_offset_h_2, + preg_all_b32); + } +} + +template +__aicore__ inline void GetKVPhyAddrVFPa(LocalTensor kvPhyAddrTensor, LocalTensor sparseIdxTensor, + LocalTensor blkTableTensor, const uint16_t s2Loop, + const uint32_t s2Tail, const uint32_t blockSize, const int16_t shiftRightNum, + const uint32_t sparseBlockSize, const uint32_t kvDim, const uint32_t kvStride) +{ + __ubuf__ uint32_t *kv_phy_addr_ub = (__ubuf__ uint32_t *)(kvPhyAddrTensor.GetPhyAddr()); + __ubuf__ int32_t *sparse_idx_ub = (__ubuf__ int32_t *)(sparseIdxTensor.GetPhyAddr()); + __ubuf__ int32_t *blk_table_ub = (__ubuf__ int32_t *)(blkTableTensor.GetPhyAddr()); + GetKVPhyAddrVFPaImpl(kv_phy_addr_ub, sparse_idx_ub, blk_table_ub, s2Loop, s2Tail, blockSize, + shiftRightNum, sparseBlockSize, kvDim, kvStride); +} + +template +__simd_vf__ void GetKVPhyAddrVFTndImpl(__ubuf__ uint32_t *kvPhyAddrUb, __ubuf__ int32_t *sparseIdxUb, + const uint16_t s2Loop, uint32_t s2Tail, const uint32_t sparseBlockSize, + const uint32_t kvDim, const uint32_t kvPrefix) +{ + static const uint16_t s2_num_per_loop = 128; + static const uint16_t s2_num_per_reg = 64; + static const uint16_t out_offset_per_loop = 256; + static const uint16_t out_offset_per_reg = 128; + static const uint32_t invalid_value = 0xFFFFFFFF; + Reg::MaskReg preg_all_b32 = Reg::CreateMask(); + Reg::MaskReg preg_tail_neg_1_b32; + Reg::MaskReg preg_tail_neg_2_b32; + + Reg::RegTensor vreg_sparse_idx_1; + Reg::RegTensor vreg_sparse_idx_2; + Reg::RegTensor vreg_kv_prefix; + Reg::RegTensor vreg_kv_dim; + Reg::RegTensor vreg_sum_1; + Reg::RegTensor vreg_sum_2; + Reg::RegTensor vreg_mul_overflow_l_1; + Reg::RegTensor vreg_mul_overflow_l_2; + Reg::RegTensor vreg_total_offset_l_1; + Reg::RegTensor vreg_total_offset_h_1; + Reg::RegTensor vreg_total_offset_l_2; + Reg::RegTensor vreg_total_offset_h_2; + + Reg::Duplicate(vreg_kv_prefix, kvPrefix); + Reg::Duplicate(vreg_kv_dim, kvDim); + + for (; s2Loop > 1;) { + for (uint16_t i = 0; i < s2Loop - 1; i++) { + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_1, + sparseIdxUb + i * s2_num_per_loop); + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_2, + sparseIdxUb + s2_num_per_reg + i * s2_num_per_loop); + // * sparseBlockSize + Reg::Muls(vreg_sparse_idx_1, vreg_sparse_idx_1, sparseBlockSize, preg_all_b32); + Reg::Muls(vreg_sparse_idx_2, vreg_sparse_idx_2, sparseBlockSize, preg_all_b32); + // (kvPrefix + sparseIdx) * kvDim -> int64 物理地址 + Reg::Add(vreg_sum_1, vreg_sparse_idx_1, vreg_kv_prefix, preg_all_b32); + Reg::Add(vreg_sum_2, vreg_sparse_idx_2, vreg_kv_prefix, preg_all_b32); + // 带进位乘法 + Reg::Mull(vreg_total_offset_l_1, vreg_total_offset_h_1, vreg_sum_1, vreg_kv_dim, preg_all_b32); + Reg::Mull(vreg_total_offset_l_2, vreg_total_offset_h_2, vreg_sum_2, vreg_kv_dim, preg_all_b32); + // 搬出 + Reg::StoreAlign( + kvPhyAddrUb + i * out_offset_per_loop, vreg_total_offset_l_1, vreg_total_offset_h_1, preg_all_b32); + Reg::StoreAlign( + kvPhyAddrUb + out_offset_per_reg + i * out_offset_per_loop, vreg_total_offset_l_2, + vreg_total_offset_h_2, preg_all_b32); + } + break; + } + + for (uint16_t i = s2Loop - 1; i < s2Loop; i++) { + Reg::MaskReg preg_tail_1_b32 = Reg::UpdateMask(s2Tail); + Reg::MaskReg preg_tail_2_b32 = Reg::UpdateMask(s2Tail); + Reg::Not(preg_tail_neg_1_b32, preg_tail_1_b32, preg_all_b32); + Reg::Not(preg_tail_neg_2_b32, preg_tail_2_b32, preg_all_b32); + + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_1, + sparseIdxUb + i * s2_num_per_loop); + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_2, + sparseIdxUb + s2_num_per_reg + i * s2_num_per_loop); + // * sparseBlockSize + Reg::Muls(vreg_sparse_idx_1, vreg_sparse_idx_1, sparseBlockSize, preg_tail_1_b32); + Reg::Muls(vreg_sparse_idx_2, vreg_sparse_idx_2, sparseBlockSize, preg_tail_2_b32); + // (kvPrefix + sparseIdx) * kvDim -> int64 物理地址 + Reg::Add(vreg_sum_1, vreg_sparse_idx_1, vreg_kv_prefix, preg_tail_1_b32); + Reg::Add(vreg_sum_2, vreg_sparse_idx_2, vreg_kv_prefix, preg_tail_2_b32); + // 带进位乘法 + Reg::Mull(vreg_total_offset_l_1, vreg_total_offset_h_1, vreg_sum_1, vreg_kv_dim, preg_tail_1_b32); + Reg::Mull(vreg_total_offset_l_2, vreg_total_offset_h_2, vreg_sum_2, vreg_kv_dim, preg_tail_2_b32); + // 无效值填充-1(0xFFFFFFFF) + Reg::Duplicate(vreg_total_offset_l_1, invalid_value, + preg_tail_neg_1_b32); + Reg::Duplicate(vreg_total_offset_h_1, invalid_value, + preg_tail_neg_1_b32); + Reg::Duplicate(vreg_total_offset_l_2, invalid_value, + preg_tail_neg_2_b32); + Reg::Duplicate(vreg_total_offset_h_2, invalid_value, + preg_tail_neg_2_b32); + Reg::StoreAlign( + kvPhyAddrUb + i * out_offset_per_loop, vreg_total_offset_l_1, vreg_total_offset_h_1, preg_all_b32); + Reg::StoreAlign( + kvPhyAddrUb + out_offset_per_reg + i * out_offset_per_loop, vreg_total_offset_l_2, vreg_total_offset_h_2, + preg_all_b32); + } +} + +template +__aicore__ inline void GetKVPhyAddrVFTnd(LocalTensor kvPhyAddrTensor, LocalTensor sparseIdxTensor, + const uint16_t s2Loop, const uint32_t s2Tail, const uint32_t sparseBlockSize, + const uint32_t kvDim, const uint32_t kvPrefix) +{ + __ubuf__ uint32_t *kv_phy_addr_ub = (__ubuf__ uint32_t *)(kvPhyAddrTensor.GetPhyAddr()); + __ubuf__ int32_t *sparse_idx_ub = (__ubuf__ int32_t *)(sparseIdxTensor.GetPhyAddr()); + GetKVPhyAddrVFTndImpl(kv_phy_addr_ub, sparse_idx_ub, s2Loop, s2Tail, sparseBlockSize, kvDim, kvPrefix); +} + +template +__simd_vf__ void GetKVPhyAddrVFBsndImpl(__ubuf__ uint32_t *kvPhyAddrUb, __ubuf__ int32_t *sparseIdxUb, + const uint16_t s2Loop, uint32_t s2Tail, const uint32_t sparseBlockSize, + const uint32_t kvDim, const uint32_t bS2BaseLow, const uint32_t bS2BaseHigh) +{ + static const uint16_t s2_num_per_loop = 128; + static const uint16_t s2_num_per_reg = 64; + static const uint16_t out_offset_per_loop = 256; + static const uint16_t out_offset_per_reg = 128; + static const uint32_t invalid_value = 0xFFFFFFFF; + Reg::MaskReg preg_all_b32 = Reg::CreateMask(); + Reg::MaskReg add_carry_l_1; + Reg::MaskReg add_carry_h_1; + Reg::MaskReg add_carry_l_2; + Reg::MaskReg add_carry_h_2; + Reg::MaskReg preg_tail_neg_1_b32; + Reg::MaskReg preg_tail_neg_2_b32; + + Reg::RegTensor vreg_sparse_idx_1; + Reg::RegTensor vreg_sparse_idx_2; + Reg::RegTensor vreg_kv_dim; + Reg::RegTensor vreg_b_s2_base_low; + Reg::RegTensor vreg_b_s2_base_high; + Reg::RegTensor vreg_s2_offset_l_1; + Reg::RegTensor vreg_s2_offset_l_2; + Reg::RegTensor vreg_mul_overflow_l_1; + Reg::RegTensor vreg_mul_overflow_l_2; + Reg::RegTensor vreg_total_offset_l_1; + Reg::RegTensor vreg_total_offset_h_1; + Reg::RegTensor vreg_total_offset_l_2; + Reg::RegTensor vreg_total_offset_h_2; + Reg::RegTensor vreg_zero; + + Reg::Duplicate(vreg_zero, 0); + Reg::Duplicate(vreg_kv_dim, kvDim); + Reg::Duplicate(vreg_b_s2_base_low, bS2BaseLow); + Reg::Duplicate(vreg_b_s2_base_high, bS2BaseHigh); + + for (; s2Loop > 1;) { + for (uint16_t i = 0; i < s2Loop - 1; i++) { + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_1, + sparseIdxUb + i * s2_num_per_loop); + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_2, + sparseIdxUb + s2_num_per_reg + i * s2_num_per_loop); + // * sparseBlockSize + Reg::Muls(vreg_sparse_idx_1, vreg_sparse_idx_1, sparseBlockSize, preg_all_b32); + Reg::Muls(vreg_sparse_idx_2, vreg_sparse_idx_2, sparseBlockSize, preg_all_b32); + // sparseIdx * kvDim (带进位乘法) + Reg::Mull(vreg_s2_offset_l_1, vreg_mul_overflow_l_1, vreg_sparse_idx_1, vreg_kv_dim, preg_all_b32); + Reg::Mull(vreg_s2_offset_l_2, vreg_mul_overflow_l_2, vreg_sparse_idx_2, vreg_kv_dim, preg_all_b32); + // s2_offset + bS2Base (int64 + int64) + Reg::Add(add_carry_l_1, vreg_total_offset_l_1, vreg_s2_offset_l_1, vreg_b_s2_base_low, preg_all_b32); + Reg::Add(add_carry_l_2, vreg_total_offset_l_2, vreg_s2_offset_l_2, vreg_b_s2_base_low, preg_all_b32); + Reg::AddC(add_carry_h_1, vreg_total_offset_h_1, vreg_mul_overflow_l_1, vreg_b_s2_base_high, add_carry_l_1, + preg_all_b32); + Reg::AddC(add_carry_h_2, vreg_total_offset_h_2, vreg_mul_overflow_l_2, vreg_b_s2_base_high, add_carry_l_2, + preg_all_b32); + // 搬出 + Reg::StoreAlign( + kvPhyAddrUb + i * out_offset_per_loop, vreg_total_offset_l_1, vreg_total_offset_h_1, preg_all_b32); + Reg::StoreAlign( + kvPhyAddrUb + out_offset_per_reg + i * out_offset_per_loop, vreg_total_offset_l_2, + vreg_total_offset_h_2, preg_all_b32); + } + break; + } + + for (uint16_t i = s2Loop - 1; i < s2Loop; i++) { + Reg::MaskReg preg_tail_1_b32 = Reg::UpdateMask(s2Tail); + Reg::MaskReg preg_tail_2_b32 = Reg::UpdateMask(s2Tail); + Reg::Not(preg_tail_neg_1_b32, preg_tail_1_b32, preg_all_b32); + Reg::Not(preg_tail_neg_2_b32, preg_tail_2_b32, preg_all_b32); + + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_1, + sparseIdxUb + i * s2_num_per_loop); + Reg::LoadAlign((Reg::RegTensor &)vreg_sparse_idx_2, + sparseIdxUb + s2_num_per_reg + i * s2_num_per_loop); + // * sparseBlockSize + Reg::Muls(vreg_sparse_idx_1, vreg_sparse_idx_1, sparseBlockSize, preg_tail_1_b32); + Reg::Muls(vreg_sparse_idx_2, vreg_sparse_idx_2, sparseBlockSize, preg_tail_2_b32); + // sparseIdx * kvDim (带进位乘法) + Reg::Mull(vreg_s2_offset_l_1, vreg_mul_overflow_l_1, vreg_sparse_idx_1, vreg_kv_dim, preg_tail_1_b32); + Reg::Mull(vreg_s2_offset_l_2, vreg_mul_overflow_l_2, vreg_sparse_idx_2, vreg_kv_dim, preg_tail_2_b32); + // s2_offset + bS2Base (int64 + int64) + Reg::Add(add_carry_l_1, vreg_total_offset_l_1, vreg_s2_offset_l_1, vreg_b_s2_base_low, preg_tail_1_b32); + Reg::Add(add_carry_l_2, vreg_total_offset_l_2, vreg_s2_offset_l_2, vreg_b_s2_base_low, preg_tail_2_b32); + Reg::AddC(add_carry_h_1, vreg_total_offset_h_1, vreg_mul_overflow_l_1, vreg_b_s2_base_high, add_carry_l_1, + preg_tail_1_b32); + Reg::AddC(add_carry_h_2, vreg_total_offset_h_2, vreg_mul_overflow_l_2, vreg_b_s2_base_high, add_carry_l_2, + preg_tail_2_b32); + // 无效值填充-1(0xFFFFFFFF) + Reg::Duplicate(vreg_total_offset_l_1, invalid_value, + preg_tail_neg_1_b32); + Reg::Duplicate(vreg_total_offset_h_1, invalid_value, + preg_tail_neg_1_b32); + Reg::Duplicate(vreg_total_offset_l_2, invalid_value, + preg_tail_neg_2_b32); + Reg::Duplicate(vreg_total_offset_h_2, invalid_value, + preg_tail_neg_2_b32); + Reg::StoreAlign( + kvPhyAddrUb + i * out_offset_per_loop, vreg_total_offset_l_1, vreg_total_offset_h_1, preg_all_b32); + Reg::StoreAlign( + kvPhyAddrUb + out_offset_per_reg + i * out_offset_per_loop, vreg_total_offset_l_2, vreg_total_offset_h_2, + preg_all_b32); + } +} + +template +__aicore__ inline void GetKVPhyAddrVFBsnd(LocalTensor kvPhyAddrTensor, LocalTensor sparseIdxTensor, + const uint16_t s2Loop, const uint32_t s2Tail, const uint32_t sparseBlockSize, + const uint32_t kvDim, const uint32_t bS2BaseLow, const uint32_t bS2BaseHigh) +{ + __ubuf__ uint32_t *kv_phy_addr_ub = (__ubuf__ uint32_t *)(kvPhyAddrTensor.GetPhyAddr()); + __ubuf__ int32_t *sparse_idx_ub = (__ubuf__ int32_t *)(sparseIdxTensor.GetPhyAddr()); + GetKVPhyAddrVFBsndImpl(kv_phy_addr_ub, sparse_idx_ub, s2Loop, s2Tail, sparseBlockSize, kvDim, bS2BaseLow, + bS2BaseHigh); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline int32_t CSABlockVec::GetSeqLen(int32_t bIdx, bool hasActualSeq, bool hasCuSeqlens, + GlobalTensor &actualSeqGm, + GlobalTensor &cuSeqlensGm, int64_t defaultSize) +{ + if (hasActualSeq) { + return actualSeqGm.GetValue(bIdx); + } else if (hasCuSeqlens) { + return cuSeqlensGm.GetValue(bIdx + 1) - cuSeqlensGm.GetValue(bIdx); + } else { + return defaultSize; + } +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline PhyAddrValidInfo CSABlockVec::CalcPhyAddrValidInfo(bool isOriKv, int32_t actualS1Size, + int32_t actualOriS2Size, + int32_t restoredSize, + ConstInfo &constInfo) +{ + // per-batch执行一次, per-s1循环内不再判断maskmode + PhyAddrValidInfo validInfo; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + validInfo.oriS2Act = actualOriS2Size; + if (isOriKv) { + if (constInfo.oriMaskMode == 0U) { + validInfo.oriTopkMode = true; + } else if (constInfo.oriMaskMode == 3U) { + validInfo.oriRightBias = 0; + } else { + validInfo.oriLeftBias = + (constInfo.oriWinLeft == -1) ? PhyAddrValidInfo::BIAS_UNBOUND : constInfo.oriWinLeft + 1; + validInfo.oriRightBias = + (constInfo.oriWinRight == -1) ? PhyAddrValidInfo::BIAS_UNBOUND : constInfo.oriWinRight; + } + } + } + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (!isOriKv) { + validInfo.cmpTopkMode = (constInfo.cmpMaskMode == 0U); + validInfo.cmpBase = restoredSize - actualS1Size + 1; + } + } + return validInfo; +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline int32_t CSABlockVec::CalcCurValidS2(uint32_t bIdx, int32_t s1Idx, int32_t actualS1Size, + bool isOriKv, GlobalTensor &cuSeqlensQGm, + GlobalTensor &topkLengthGm, + ConstInfo &constInfo, int32_t sparseBlockCount, + const PhyAddrValidInfo &validInfo) +{ + bool topkMode = false; + bool hasTopk = false; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (isOriKv) { + topkMode = validInfo.oriTopkMode; + hasTopk = constInfo.hasOriTopkLength; + } + } + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (!isOriKv) { + topkMode = validInfo.cmpTopkMode; + hasTopk = constInfo.hasCmpTopkLength; + } + } + if (topkMode) { + uint64_t topkIdx = + (LAYOUT_T == SMLA_LAYOUT::TND) ? (cuSeqlensQGm.GetValue(bIdx) + s1Idx) : (bIdx * constInfo.s1Size + s1Idx); + int32_t topkLen = hasTopk ? topkLengthGm.GetValue(topkIdx) : sparseBlockCount; + return Min(topkLen, sparseBlockCount); + } + + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (isOriKv) { + int64_t thr = validInfo.oriS2Act - actualS1Size + 1 + s1Idx; + int64_t leftBound = Max(thr - validInfo.oriLeftBias, 0); + int64_t rightBound = Min(thr + validInfo.oriRightBias, static_cast(validInfo.oriS2Act)); + return Min(static_cast(Max(0, rightBound - leftBound)), sparseBlockCount); + } + } + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + int64_t numerator = Max(validInfo.cmpBase + s1Idx, 0); + return Min(sparseBlockCount, static_cast(numerator / static_cast(constInfo.cmpRatio))); + } + return 0; +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CopyPhyAddrToGm(LocalTensor kvPhyAddrUb, int64_t bS1Idx, + int64_t s1Idx, int64_t validS2, int64_t alignNum, + GlobalTensor &phyAddrGm, + uint32_t alignedSparseBlockCount) +{ + constexpr int64_t numPerBlock = 32; + DataCopyParams dataCopyParams; + dataCopyParams.blockCount = 1U; + dataCopyParams.blockLen = ((validS2 + alignNum - 1) / alignNum * alignNum) * sizeof(int64_t) / numPerBlock; + dataCopyParams.srcGap = 0U; + dataCopyParams.dstGap = 0U; + DataCopy(phyAddrGm[(bS1Idx + s1Idx) * alignedSparseBlockCount * 2], kvPhyAddrUb, dataCopyParams); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CopyPaTableToUb(LocalTensor blkTableUb, int64_t bIdx, + GlobalTensor &blockTableGm, + uint32_t maxBlockNumPerBatch) +{ + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1U; + dataCopyParams.blockLen = maxBlockNumPerBatch * sizeof(int32_t); + dataCopyParams.srcStride = 0U; + dataCopyParams.dstStride = 0U; + DataCopyPadExtParams padParams; + DataCopyPad(blkTableUb, blockTableGm[bIdx * maxBlockNumPerBatch], dataCopyParams, padParams); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::CopySparseIdxToUb(LocalTensor sparseIdxUb, int64_t bS1Idx, + int64_t s1Idx, int64_t validS2, + GlobalTensor &sparseIndicesGm, + uint32_t sparseBlockCount) +{ + DataCopyExtParams dataCopyParams; + dataCopyParams.blockCount = 1U; + dataCopyParams.blockLen = validS2 * sizeof(int32_t); + dataCopyParams.srcStride = 0U; + dataCopyParams.dstStride = 0U; + DataCopyPadExtParams padParams; + DataCopyPad(sparseIdxUb, sparseIndicesGm[(bS1Idx + s1Idx) * sparseBlockCount], dataCopyParams, padParams); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::GetKVPhyAddrForKvType( + uint32_t bN2StartIdx, uint32_t bN2EndIdx, uint32_t gS1StartIdx, uint32_t nextGs1Idx, bool hasActualSeqQlen, + bool hasCuSeqlensQ, bool hasActualSeqKvlen, bool hasCuSeqlensKv, GlobalTensor actualSeqQlenGm, + GlobalTensor cuSeqlensQGm, GlobalTensor actualSeqKvlenGm, GlobalTensor cuSeqlensKvGm, + GlobalTensor topkLengthGm, GlobalTensor cmpResidualKvGm, ConstInfo &constInfo, + GlobalTensor &blockTableGm, GlobalTensor &sparseIndicesGm, GlobalTensor &phyAddrGm, + uint32_t kvStride, uint32_t blockSize, uint32_t maxBlockNumPerBatch, uint32_t sparseBlockCount, + uint32_t alignedSparseBlockCount, bool isOriKv) +{ + static constexpr uint16_t s2NumPerLoop = 128; + static constexpr uint32_t vecCoreNum = IS_SPLIT_G ? 4 : 2; + uint32_t vecCoreIdx = IS_SPLIT_G ? constInfo.aivIdx % 4 : constInfo.aivIdx % 2; + uint32_t phyAddrUb = 0; + int16_t shiftRightNum = 0; + LocalTensor blkTableUb; + + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + int32_t blkSize = static_cast(blockSize); + while (blkSize > 1) { + blkSize >>= 1; + shiftRightNum++; + } + blkTableUb = LocalTensor(TPosition::VECIN, phyAddrUb, maxBlockNumPerBatch); + phyAddrUb = CeilAlign(phyAddrUb + maxBlockNumPerBatch * sizeof(int32_t), BUFFER_SIZE_BYTE_32B); + } + LocalTensor sparseIdxUb(TPosition::VECIN, phyAddrUb, alignedSparseBlockCount); + phyAddrUb += alignedSparseBlockCount * sizeof(int32_t); + LocalTensor kvPhyAddrUb(TPosition::VECIN, phyAddrUb, alignedSparseBlockCount * 2); // 2 for ori/cmp kv + + // 第一遍: 统计totalValidS1 + int64_t totalValidS1 = 0; + uint32_t tmpGS1Start = gS1StartIdx; + for (uint32_t bIdx = bN2StartIdx; bIdx < bN2EndIdx; ++bIdx) { + bool lastBN = (bIdx == bN2EndIdx - 1); + int32_t actualS1Size = + GetSeqLen(bIdx, hasActualSeqQlen, hasCuSeqlensQ, actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); + int32_t s1End = actualS1Size; + if (lastBN && nextGs1Idx != 0) { + s1End = nextGs1Idx; + } + + int64_t bS1IdxBase = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + bS1IdxBase = hasCuSeqlensQ ? cuSeqlensQGm.GetValue(bIdx) : constInfo.s1Size * bIdx; + } else { + bS1IdxBase = constInfo.s1Size * bIdx; + } + + int32_t restoredSize = 0; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (!isOriKv && constInfo.cmpMaskMode != 0) { + int32_t actualKvSize = GetSeqLen(bIdx, hasActualSeqKvlen, hasCuSeqlensKv, actualSeqKvlenGm, + cuSeqlensKvGm, constInfo.cmpS2Size); + restoredSize = actualKvSize * static_cast(constInfo.cmpRatio) + cmpResidualKvGm.GetValue(bIdx); + } + } + int32_t actualOriS2Size = 0; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (isOriKv && constInfo.oriMaskMode != 0) { + actualOriS2Size = GetSeqLen(bIdx, hasActualSeqKvlen, hasCuSeqlensKv, actualSeqKvlenGm, cuSeqlensKvGm, + constInfo.s2Size); + } + } + PhyAddrValidInfo validInfo = + CalcPhyAddrValidInfo(isOriKv, actualS1Size, actualOriS2Size, restoredSize, constInfo); + + for (int32_t s1Idx = tmpGS1Start; s1Idx < s1End; ++s1Idx) { + int32_t curValidS2 = CalcCurValidS2(bIdx, s1Idx, actualS1Size, isOriKv, cuSeqlensQGm, topkLengthGm, + constInfo, static_cast(sparseBlockCount), validInfo); + if (curValidS2 > 0) { + totalValidS1++; + } + } + tmpGS1Start = 0; + } + + int64_t s1PerVecCore = totalValidS1 / vecCoreNum; + int64_t s1Tail = totalValidS1 % vecCoreNum; + int64_t curStart = s1PerVecCore * vecCoreIdx + Min((int64_t)vecCoreIdx, s1Tail); + int64_t curCount = s1PerVecCore + (vecCoreIdx < (uint32_t)s1Tail ? 1 : 0); + + if (curCount == 0) { + return; + } + + // 第二遍: 实际计算 + int64_t validCounter = 0; + int64_t processedCount = 0; + tmpGS1Start = gS1StartIdx; + bool done = false; + + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + SetFlag(INNERCORE_PHYADDR_BLKTABLE_FREE); + } + SetFlag(INNERCORE_PHYADDR_SPARSEIDX_FREE); + SetFlag(INNERCORE_PHYADDR_KVADDR_FREE); + for (uint32_t bIdx = bN2StartIdx; bIdx < bN2EndIdx && !done; ++bIdx) { + bool lastBN = (bIdx == bN2EndIdx - 1); + int32_t actualS1Size = + GetSeqLen(bIdx, hasActualSeqQlen, hasCuSeqlensQ, actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); + int64_t bS1Idx = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + bS1Idx = hasCuSeqlensQ ? cuSeqlensQGm.GetValue(bIdx) : constInfo.s1Size * bIdx; + } else { + bS1Idx = constInfo.s1Size * bIdx; + } + + int32_t s1End = actualS1Size; + if (lastBN && nextGs1Idx != 0) { + s1End = nextGs1Idx; + } + + // per-batch 参数预计算 + uint32_t kvPrefix = 0; + uint32_t bS2BaseLow = 0; + uint32_t bS2BaseHigh = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + kvPrefix = static_cast(cuSeqlensKvGm.GetValue(bIdx)); + } else { + uint32_t s2Size = + isOriKv ? static_cast(constInfo.s2Size) : static_cast(constInfo.cmpS2Size); + uint64_t bS2Base = static_cast(bIdx) * s2Size * static_cast(constInfo.dSize); + bS2BaseLow = static_cast(bS2Base); + bS2BaseHigh = static_cast(bS2Base >> 32U); + } + + int32_t restoredSize = 0; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (!isOriKv && constInfo.cmpMaskMode != 0) { + int32_t actualKvSize = GetSeqLen(bIdx, hasActualSeqKvlen, hasCuSeqlensKv, actualSeqKvlenGm, + cuSeqlensKvGm, constInfo.cmpS2Size); + restoredSize = actualKvSize * static_cast(constInfo.cmpRatio) + cmpResidualKvGm.GetValue(bIdx); + } + } + int32_t actualOriS2Size = 0; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (isOriKv && constInfo.oriMaskMode != 0) { + actualOriS2Size = GetSeqLen(bIdx, hasActualSeqKvlen, hasCuSeqlensKv, actualSeqKvlenGm, cuSeqlensKvGm, + constInfo.s2Size); + } + } + PhyAddrValidInfo validInfo = + CalcPhyAddrValidInfo(isOriKv, actualS1Size, actualOriS2Size, restoredSize, constInfo); + + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + WaitFlag(INNERCORE_PHYADDR_BLKTABLE_FREE); + CopyPaTableToUb(blkTableUb, bIdx, blockTableGm, maxBlockNumPerBatch); + SetFlag(INNERCORE_PHYADDR_BLKTABLE_READY); + WaitFlag(INNERCORE_PHYADDR_BLKTABLE_READY); + } + + for (int32_t s1Idx = tmpGS1Start; s1Idx < s1End; ++s1Idx) { + int32_t curValidS2 = CalcCurValidS2(bIdx, s1Idx, actualS1Size, isOriKv, cuSeqlensQGm, topkLengthGm, + constInfo, static_cast(sparseBlockCount), validInfo); + if (curValidS2 <= 0) { + continue; + } + + if (validCounter < curStart || validCounter >= curStart + curCount) { + validCounter++; + continue; + } + validCounter++; + + uint16_t s2Loop = (curValidS2 + s2NumPerLoop - 1) / s2NumPerLoop; + int32_t s2Tail = curValidS2 - (s2Loop - 1) * s2NumPerLoop; + WaitFlag(INNERCORE_PHYADDR_SPARSEIDX_FREE); + CopySparseIdxToUb(sparseIdxUb, bS1Idx, s1Idx, curValidS2, sparseIndicesGm, sparseBlockCount); + SetFlag(INNERCORE_PHYADDR_SPARSEIDX_READY); + + WaitFlag(INNERCORE_PHYADDR_SPARSEIDX_READY); + WaitFlag(INNERCORE_PHYADDR_KVADDR_FREE); + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + GetKVPhyAddrVFPa(kvPhyAddrUb, sparseIdxUb, blkTableUb, s2Loop, s2Tail, blockSize, + shiftRightNum, constInfo.sparseBlockSize, + static_cast(constInfo.dSize), kvStride); + } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + GetKVPhyAddrVFTnd(kvPhyAddrUb, sparseIdxUb, s2Loop, s2Tail, constInfo.sparseBlockSize, + static_cast(constInfo.dSize), kvPrefix); + } else { + GetKVPhyAddrVFBsnd(kvPhyAddrUb, sparseIdxUb, s2Loop, s2Tail, constInfo.sparseBlockSize, + static_cast(constInfo.dSize), bS2BaseLow, bS2BaseHigh); + } + SetFlag(INNERCORE_PHYADDR_SPARSEIDX_FREE); + SetFlag(INNERCORE_PHYADDR_KVADDR_READY); + WaitFlag(INNERCORE_PHYADDR_KVADDR_READY); + CopyPhyAddrToGm(kvPhyAddrUb, bS1Idx, s1Idx, curValidS2, s2NumPerLoop, phyAddrGm, alignedSparseBlockCount); + SetFlag(INNERCORE_PHYADDR_KVADDR_FREE); + + processedCount++; + if (processedCount >= curCount) { + done = true; + break; + } + } + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + SetFlag(INNERCORE_PHYADDR_BLKTABLE_FREE); + } + tmpGS1Start = 0; + } + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + WaitFlag(INNERCORE_PHYADDR_BLKTABLE_FREE); + } + WaitFlag(INNERCORE_PHYADDR_SPARSEIDX_FREE); + WaitFlag(INNERCORE_PHYADDR_KVADDR_FREE); +} + +TEMPLATES_DEF_NO_DEFAULT +__aicore__ inline void CSABlockVec::GetKVPhyAddr( + uint32_t hasLoad, uint32_t bN2StartIdx, uint32_t bN2EndIdx, uint32_t gS1StartIdx, uint32_t nextGs1Idx, + bool hasActualSeqQlen, bool hasCuSeqlensQ, bool hasActualSeqOriKvlen, bool hasCuSeqlensOriKv, + GlobalTensor actualSeqOriKvlenGm, GlobalTensor cuSeqlensOriKvGm, + GlobalTensor oriTopkLengthGm, bool hasActualSeqCmpKvlen, bool hasCuSeqlensCmpKv, + GlobalTensor actualSeqCmpKvlenGm, GlobalTensor cuSeqlensCmpKvGm, + GlobalTensor cmpTopkLengthGm, GlobalTensor cmpResidualKvGm, GlobalTensor actualSeqQlenGm, + GlobalTensor cuSeqlensQGm, __gm__ uint8_t *workspace, ConstInfo &constInfo) +{ + if (hasLoad == 0) { + return; + } + + // GM分配: ori在前, cmp在后 + int64_t v0TotalOffset = 0; + uint32_t v0ResSize = constInfo.s2BaseSize * constInfo.dSize * sizeof(Q_T); + if constexpr (IS_SPLIT_G) { + v0TotalOffset = v0ResSize * 3 * (GetBlockNum() >> 1U); + } else { + v0TotalOffset = v0ResSize * 3 * GetBlockNum(); + } + + // SMLA特有: 加上s2RealBuf大小 + constexpr uint32_t TRIPLE_BUFFER_NUM = 3; + constexpr uint32_t S2_REAL_BUF_LEN = 128; + v0TotalOffset += TRIPLE_BUFFER_NUM * S2_REAL_BUF_LEN * sizeof(int32_t) * GetBlockNum(); + + uint32_t totalBS1 = (LAYOUT_T == SMLA_LAYOUT::TND) ? constInfo.s1Size : (constInfo.bSize * constInfo.s1Size); + + uint64_t oriPhyAddrSize = 0; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + oriPhyAddrSize = static_cast(totalBS1) * constInfo.alignedOriSparseBlockCount * sizeof(int64_t); + this->oriKvPhyAddrGm.SetGlobalBuffer((__gm__ uint32_t *)(workspace + v0TotalOffset)); + } + + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + uint64_t cmpPhyAddrSize = + static_cast(totalBS1) * constInfo.alignedCmpSparseBlockCount * sizeof(int64_t); + this->cmpKvPhyAddrGm.SetGlobalBuffer((__gm__ uint32_t *)(workspace + v0TotalOffset + oriPhyAddrSize)); + } + + // ori部分 (先计算) + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + GetKVPhyAddrForKvType(bN2StartIdx, bN2EndIdx, gS1StartIdx, nextGs1Idx, hasActualSeqQlen, hasCuSeqlensQ, + hasActualSeqOriKvlen, hasCuSeqlensOriKv, actualSeqQlenGm, cuSeqlensQGm, + actualSeqOriKvlenGm, cuSeqlensOriKvGm, oriTopkLengthGm, cmpResidualKvGm, constInfo, + oriBlockTableGm, oriSparseIndicesGm, oriKvPhyAddrGm, constInfo.oriKeyStride0, + constInfo.oriBlockSize, constInfo.oriMaxBlockNumPerBatch, constInfo.oriSparseBlockCount, + constInfo.alignedOriSparseBlockCount, true); + } + + // cmp部分 (后计算) + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + GetKVPhyAddrForKvType(bN2StartIdx, bN2EndIdx, gS1StartIdx, nextGs1Idx, hasActualSeqQlen, hasCuSeqlensQ, + hasActualSeqCmpKvlen, hasCuSeqlensCmpKv, actualSeqQlenGm, cuSeqlensQGm, + actualSeqCmpKvlenGm, cuSeqlensCmpKvGm, cmpTopkLengthGm, cmpResidualKvGm, constInfo, + cmpBlockTableGm, cmpSparseIndicesGm, cmpKvPhyAddrGm, constInfo.cmpKeyStride0, + constInfo.cmpBlockSize, constInfo.cmpMaxBlockNumPerBatch, constInfo.cmpSparseBlockCount, + constInfo.alignedCmpSparseBlockCount, false); + } +} + +TEMPLATES_DEF +class CSABlockVecDummy { +public: + __aicore__ inline CSABlockVecDummy(){}; + __aicore__ inline void CleanOutput(__gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, ConstInfo &constInfo) + {} + __aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, + __gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sinks, + __gm__ uint8_t *sequsedOriKv, __gm__ uint8_t *sequsedCmpKv, + __gm__ uint8_t *cmpResidualKv) + {} + __aicore__ inline void InitVecBlock(__gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *seqUsedOriKV, + __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV) {}; + __aicore__ inline void InitS2SplitStaging(Buffer &fdStaging) {} + __aicore__ inline void InitS2SplitStaging(Buffer &intraCoreCombine, + Buffer &crossCoreCombine) + {} + __aicore__ inline void InitLocalBuffer(ConstInfo &constInfo, uint32_t ubBaseAddr) {} + __aicore__ inline void InitFDBuffers(FdRunInfo &fdRunInfo) {} + __aicore__ inline void ProcessFlashDecode(FdRunInfo &fdRunInfo, ConstInfo &constInfo) {} + __aicore__ inline void ProcessVec1(StaticBuffer &outputBuf, StaticBuffer &bmm1ResBuf, RunInfo &runInfo, + ConstInfo &constInfo) + {} + __aicore__ inline void ProcessVec2(StaticBuffer &bmm2ResBuf, RunInfo &runInfo, ConstInfo &constInfo) {} + __aicore__ inline void FreeEvent(ConstInfo &constInfo) {} +}; +} // namespace SMLAKernel +#endif // SPARSE_FLASH_MLA_CSA_BLOCK_VECTOR_ARCH35_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_kernel_arch35.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_kernel_arch35.h new file mode 100644 index 000000000000..432604f0bdf4 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_csa_kernel_arch35.h @@ -0,0 +1,961 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_csa_kernel_arch35.h + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_CSA_KERNEL_ARCH35_H +#define SPARSE_FLASH_MLA_CSA_KERNEL_ARCH35_H +#include "sparse_flash_mla_common_arch35.h" +#include "sparse_flash_mla_kvcache.h" +#include "sparse_flash_mla_csa_block_cube_arch35.h" +#include "sparse_flash_mla_csa_block_vector_arch35.h" +#include "kernel_operator.h" +#include "../sparse_flash_mla_kernel_metadata.h" + +#if __has_include("../../common/op_kernel/matmul.h") +#include "../../common/op_kernel/matmul.h" +#else +#include "../common/matmul.h" +#endif +#if __has_include("../../common/op_kernel/FixpipeOut.h") +#include "../../common/op_kernel/FixpipeOut.h" +#else +#include "../common/FixpipeOut.h" +#endif +#if __has_include("../../common/op_kernel/CopyInL1.h") +#include "../../common/op_kernel/CopyInL1.h" +#else +#include "../common/CopyInL1.h" +#endif +#if __has_include("common/buffers_policy_3buff_sfa.h") +#include "common/buffers_policy_3buff_sfa.h" +#endif + +#include "kernel_operator_list_tensor_intf.h" + +using matmul::MatmulType; +using namespace AscendC; +using namespace optiling; +using namespace optiling::detail; +using namespace AscendC::Impl::Detail; +using namespace regbaseutil; +using AttentionCommon::FdRunInfo; + +namespace SMLAKernel { +template +class SparseFlashMlaCsaKernel { +public: + ARGS_TRAITS; + __aicore__ inline SparseFlashMlaCsaKernel(){}; + + __aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, + __gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *seqUsedOriKV, + __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, + __gm__ uint8_t *metadata, __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, + __gm__ uint8_t *workspace, const SparseFlashMlaTilingData *__restrict tiling); + __aicore__ inline void Process(); + +private: + __aicore__ inline void ProcessMainLoop(); + __aicore__ inline int64_t GetSeqLen(int32_t bIdx, bool hasActualSeq, bool hasCuSeqlens, + GlobalTensor &actualSeqGm, GlobalTensor &cuSeqlensGm, + int64_t defaultSize); + __aicore__ inline void ParseTilingData(__gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *sequsedQ, + __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, + __gm__ uint8_t *cmpResidualKV); + __aicore__ inline void InitGlobalBuffer( + __gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *sequsedQ, __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, + __gm__ uint8_t *cmpResidualKV, __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, + __gm__ uint8_t *sinks, __gm__ uint8_t *workspace, const SparseFlashMlaTilingData *__restrict tiling); + __aicore__ inline void InitLocalBuffer(); + __aicore__ inline void FreeEvent(); + __aicore__ inline void InitMMResBuf(__gm__ uint8_t *workspace); + __aicore__ inline void ComputeConstexpr(); + __aicore__ inline void SetRunInfo(RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount, + int64_t s2LoopLimit, int64_t multiCoreInnerIdx); + __aicore__ inline void ComputeBmm1Tail(RunInfo &runInfo, RunParamStr &runParam); + __aicore__ inline void ComputeAxisIdxByBnAndGs1(int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam); + __aicore__ inline void InitUniqueRunInfo(const RunParamStr &runParam, RunInfo &runInfo); + __aicore__ inline void ParseFdRunInfo(FdRunInfo &fdRunInfo); + __aicore__ inline int64_t ConvertS2MetadataBlockToToken(const RunParamStr &runParam, const ConstInfo &constInfo, + uint32_t s2BlockIdx); + __aicore__ inline bool ApplyS2MetadataRange(RunParamStr &runParam, ConstInfo &constInfo, int64_t s2StartPoint, + int64_t s2EndPoint, bool isFirstS2RangeTask, bool isLastS2RangeTask); + const SparseFlashMlaTilingData *__restrict tilingData; + /* 编译期常量的基本块信息 */ + static constexpr uint32_t PRELOAD_NUM = 3; + static constexpr uint32_t crossCoreMte2SyncFlagId = 15; // IS_SPLIT_G 核间 MTE2 同步 flag ID + static constexpr uint32_t SPARSE_BLOCK_ALIGN_NUM = 128; + + /* 核间通道 */ + BufferManager v0ResGmBufferManager; + + StaticBuffer bmm1Buffers[2]; + StaticBuffer bmm2Buffers; + uint32_t bmm1GetFlag = 0; + uint32_t vUbBase = 0; + + // mm2左矩阵P + StaticBuffer l1PBuffers[2]; + uint32_t l1PGetFlag = 0; + uint32_t l1CubeBase = 0; + GlobalTensor metadataGm; + GlobalTensor cuSeqlensQGm; + GlobalTensor cuSeqlensOriKvGm; + GlobalTensor cuSeqlensCmpKvGm; + GlobalTensor actualSeqOriKvlenGm; + GlobalTensor actualSeqCmpKvlenGm; + GlobalTensor cmpResidualKvGm; + GlobalTensor actualSeqQlenGm; + GlobalTensor oriTopkLengthGm; + GlobalTensor cmpTopkLengthGm; + + bool hasCuSeqlensQ = false; + bool hasCuSeqlensOriKv = false; + bool hasCuSeqlensCmpKv = false; + bool hasActualSeqQlen = false; + bool hasActualSeqOriKvlen = false; + bool hasActualSeqCmpKvlen = false; + /* workspace 空间 */ + BuffersPolicy3buffSFA v0ResGmBuffers; + BufferManager fdStagingBufferManager; + BuffersPolicySingleBuffer fdStagingBuffer; + BuffersPolicySingleBuffer intraCoreCombineBuffer; + BuffersPolicySingleBuffer crossCoreCombineBuffer; + /* 核Index信息 */ + int32_t aicIdx; + + /* Init阶段metadata解析结果 */ + uint32_t bN2StartIdx; + uint32_t gS1StartIdx; + uint32_t bN2EndIdx; + uint32_t nextGs1Idx; + uint32_t hasLoad; + + /* 初始化后不变的信息 */ + ConstInfo constInfo; + + /* 模板库Block */ + CubeBlockType cubeBlock; + VecBlockType vecBlock; +}; + +template +__aicore__ inline void SparseFlashMlaCsaKernel::Init( + __gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *sequsedQ, __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata, + __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, __gm__ uint8_t *workspace, + const SparseFlashMlaTilingData *__restrict tiling) +{ + fa_base_matmul::ResetIdCounter(); + constInfo.subBlockIdx = GetSubBlockIdx(); + if ASCEND_IS_AIC { + this->aicIdx = GetBlockIdx(); + constInfo.aivIdx = 0; + this->tilingData = tiling; + } else { + constInfo.aivIdx = GetBlockIdx(); + this->aicIdx = constInfo.aivIdx >> 1; + this->tilingData = tiling; + } + + if (metadata == nullptr) { + return; + } + this->metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata); + + constInfo.s1BaseSize = 64; + constInfo.s2BaseSize = 128; + + this->ParseTilingData(cuSeqlensQ, sequsedQ, cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedOriKV, seqUsedCmpKV, + cmpResidualKV); + vecBlock.InitVecBlock(cuSeqlensQ, cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedOriKV, seqUsedCmpKV, cmpResidualKV); + vecBlock.CleanOutput(attentionOut, softmaxLse, constInfo); + + // 从meta data解析分核信息 + bN2StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_BN2_START_INDEX, false)); + gS1StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_M_START_INDEX, false)); + bN2EndIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_BN2_END_INDEX, false)); + nextGs1Idx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_M_END_INDEX, false)); + hasLoad = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_CORE_ENABLE_INDEX, false)); + if (nextGs1Idx != 0) { + bN2EndIdx++; + } + + this->InitGlobalBuffer(query, oriKV, cmpKV, oriSparseIndices, cmpSparseIndices, oriBlockTable, cmpBlockTable, + cuSeqlensQ, cuSeqlensOriKv, cuSeqlensCmpKv, sequsedQ, seqUsedOriKV, seqUsedCmpKV, + cmpResidualKV, oriTopkLength, cmpTopkLength, sinks, workspace, tiling); // gm设置 + + if ASCEND_IS_AIV { + if constexpr ((TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) && + IS_VEC_S2PHYADDR) { + this->vecBlock.GetKVPhyAddr(hasLoad, bN2StartIdx, bN2EndIdx, gS1StartIdx, nextGs1Idx, hasActualSeqQlen, + hasCuSeqlensQ, hasActualSeqOriKvlen, hasCuSeqlensOriKv, actualSeqOriKvlenGm, + cuSeqlensOriKvGm, oriTopkLengthGm, hasActualSeqCmpKvlen, hasCuSeqlensCmpKv, + actualSeqCmpKvlenGm, cuSeqlensCmpKvGm, cmpTopkLengthGm, cmpResidualKvGm, + actualSeqQlenGm, cuSeqlensQGm, workspace, constInfo); + } + } + + InitMMResBuf(workspace); + if constexpr (IS_BATCH_CONSISTENCY) { + vecBlock.InitS2SplitStaging(intraCoreCombineBuffer.Get(), crossCoreCombineBuffer.Get()); + } else { + vecBlock.InitS2SplitStaging(fdStagingBuffer.Get()); + } + this->ComputeConstexpr(); + this->InitLocalBuffer(); +} + +template +__aicore__ inline int64_t SparseFlashMlaCsaKernel::GetSeqLen( + int32_t bIdx, bool hasActualSeq, bool hasCuSeqlens, GlobalTensor &actualSeqGm, + GlobalTensor &cuSeqlensGm, int64_t defaultSize) +{ + if (hasActualSeq) { + return actualSeqGm.GetValue(bIdx); + } else if (hasCuSeqlens) { + return cuSeqlensGm.GetValue(bIdx + 1) - cuSeqlensGm.GetValue(bIdx); + } else { + return defaultSize; + } +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::ParseTilingData( + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, + __gm__ uint8_t *cmpResidualKV) +{ + auto &sparseFlashMLABaseParams = this->tilingData->baseParams; + auto &sparseFlashMLACmpParams = this->tilingData->cmpParams; + constInfo.bSize = sparseFlashMLABaseParams.batchSize; + constInfo.n2Size = 1; + constInfo.gSize = sparseFlashMLABaseParams.nNumOfQInOneGroup; + constInfo.s1Size = sparseFlashMLABaseParams.qSeqSize; + constInfo.s2Size = sparseFlashMLABaseParams.kvSeqSize; + constInfo.cmpS2Size = sparseFlashMLACmpParams.cmpKvSeqSize; + constInfo.oriSparseBlockCount = sparseFlashMLABaseParams.oriSparseBlockCount; + constInfo.cmpSparseBlockCount = sparseFlashMLACmpParams.cmpSparseBlockCount; + constInfo.alignedOriSparseBlockCount = + (constInfo.oriSparseBlockCount + SPARSE_BLOCK_ALIGN_NUM - 1) / SPARSE_BLOCK_ALIGN_NUM * SPARSE_BLOCK_ALIGN_NUM; + constInfo.alignedCmpSparseBlockCount = + (constInfo.cmpSparseBlockCount + SPARSE_BLOCK_ALIGN_NUM - 1) / SPARSE_BLOCK_ALIGN_NUM * SPARSE_BLOCK_ALIGN_NUM; + constInfo.cmpRatio = sparseFlashMLACmpParams.cmpRatio; + constInfo.oriMaskMode = sparseFlashMLABaseParams.oriMaskMode; + constInfo.cmpMaskMode = sparseFlashMLACmpParams.cmpMaskMode; + constInfo.oriWinLeft = sparseFlashMLABaseParams.oriWinLeft; + constInfo.oriWinRight = sparseFlashMLABaseParams.oriWinRight; + constInfo.layoutType = sparseFlashMLABaseParams.outputLayout; + constInfo.returnSoftmaxLse = sparseFlashMLABaseParams.returnSoftmaxLse; + constInfo.tileSize = 0; + constInfo.dSizeRope = 64; + constInfo.oriKeyStride0 = sparseFlashMLABaseParams.oriKeyStride0; + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE) { + constInfo.cmpKeyStride0 = sparseFlashMLACmpParams.cmpKeyStride0; + } + if ASCEND_IS_AIV { + constInfo.softmaxScale = sparseFlashMLABaseParams.softmaxScale; + } + constInfo.dSize = 512; + constInfo.dSizeV = constInfo.dSize; + constInfo.dSizeVInput = constInfo.dSize; + constInfo.dSizeNope = constInfo.dSize - constInfo.dSizeRope; + constInfo.sparseBlockSize = 1; + constInfo.actualSeqLenSize = constInfo.bSize + 1; + constInfo.actualLenDimsOriKV = sparseFlashMLABaseParams.actualLenDimsOriKV; + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE) { + constInfo.actualLenDimsCmpKV = sparseFlashMLABaseParams.actualLenDimsCmpKV; + constInfo.cmpResidualKVSize = sparseFlashMLABaseParams.cmpResidualKVSize; + } + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + this->constInfo.isActualLenDimsOriKVNull = 0U; + } else { + this->constInfo.isActualLenDimsOriKVNull = (seqUsedOriKV == nullptr); + } + + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + constInfo.oriBlockSize = sparseFlashMLABaseParams.oriBlockSize; + constInfo.cmpBlockSize = sparseFlashMLABaseParams.cmpBlockSize; + constInfo.oriMaxBlockNumPerBatch = sparseFlashMLABaseParams.oriMaxBlockNumPerBatch; + constInfo.cmpMaxBlockNumPerBatch = sparseFlashMLACmpParams.cmpMaxBlockNumPerBatch; + } + + if (cuSeqlensQ != nullptr) { + cuSeqlensQGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensQ); + hasCuSeqlensQ = true; + } + if (cuSeqlensOriKv != nullptr) { + cuSeqlensOriKvGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensOriKv); + hasCuSeqlensOriKv = true; + } + + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + if (cuSeqlensCmpKv != nullptr) { + cuSeqlensCmpKvGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensCmpKv); + hasCuSeqlensCmpKv = true; + } + } + + if (sequsedQ != nullptr) { + actualSeqQlenGm.SetGlobalBuffer((__gm__ int32_t *)sequsedQ); + hasActualSeqQlen = true; + } + if (seqUsedOriKV != nullptr) { + actualSeqOriKvlenGm.SetGlobalBuffer((__gm__ int32_t *)seqUsedOriKV); + hasActualSeqOriKvlen = true; + } + + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + if (seqUsedCmpKV != nullptr) { + actualSeqCmpKvlenGm.SetGlobalBuffer((__gm__ int32_t *)seqUsedCmpKV); + hasActualSeqCmpKvlen = true; + } + if (cmpResidualKV != nullptr) { + cmpResidualKvGm.SetGlobalBuffer((__gm__ int32_t *)cmpResidualKV); + } + } + + constInfo.needInit = 0; + if (TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE && constInfo.oriMaskMode != 0) { + for (uint32_t bIdx = 0; bIdx < constInfo.bSize; bIdx++) { + int64_t s2Size; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + s2Size = actualSeqOriKvlenGm.GetValue(bIdx); + } else { + s2Size = GetSeqLen(bIdx, hasActualSeqOriKvlen, hasCuSeqlensOriKv, actualSeqOriKvlenGm, cuSeqlensOriKvGm, + constInfo.s2Size); + } + int64_t s1Size = + GetSeqLen(bIdx, hasActualSeqQlen, hasCuSeqlensQ, actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); + int64_t expectQs; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + expectQs = GetSeqLen(bIdx, false, hasCuSeqlensQ, actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); + } else { + expectQs = constInfo.s1Size; + } + if (s1Size > s2Size || s1Size < expectQs) { + constInfo.needInit = 1; + break; + } + } + } else { + constInfo.needInit = 1; + } +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::InitGlobalBuffer( + __gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *sequsedQ, __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, __gm__ uint8_t *workspace, + const SparseFlashMlaTilingData *__restrict tiling) +{ + vecBlock.InitGlobalBuffer(oriKV, cmpKV, oriSparseIndices, cmpSparseIndices, oriBlockTable, cmpBlockTable, sequsedQ, + sinks, seqUsedOriKV, seqUsedCmpKV, cmpResidualKV); + cubeBlock.InitGlobalBuffer(query, oriKV, cmpKV, cmpSparseIndices, oriBlockTable, cmpBlockTable, sequsedQ, + cuSeqlensQ, cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedOriKV, seqUsedCmpKV, constInfo); + + if (oriTopkLength != nullptr) { + constInfo.hasOriTopkLength = true; + oriTopkLengthGm.SetGlobalBuffer((__gm__ int32_t *)oriTopkLength); + } else { + constInfo.hasOriTopkLength = false; + } + if (cmpTopkLength != nullptr) { + constInfo.hasCmpTopkLength = true; + cmpTopkLengthGm.SetGlobalBuffer((__gm__ int32_t *)cmpTopkLength); + } else { + constInfo.hasCmpTopkLength = false; + } +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::InitMMResBuf(__gm__ uint8_t *workspace) +{ + // L1: [l1P x2][cube L1], l1P 必须放在最前面保证与 vec 申请地址相同 + uint32_t mm2LeftSize = constInfo.s1BaseSize * constInfo.s2BaseSize; + uint32_t l1PAddr = 0; + l1PBuffers[0] = {LocalTensor(TPosition::A1, l1PAddr, mm2LeftSize), 0}; + l1PAddr += (mm2LeftSize * sizeof(Q_T)); + l1PBuffers[1] = {LocalTensor(TPosition::A1, l1PAddr, mm2LeftSize), 1}; + l1PAddr += (mm2LeftSize * sizeof(Q_T)); + l1CubeBase = l1PAddr; + + // UB: [bmm2][bmm1 x2][vec UB] + uint32_t mm1ResultSize = constInfo.s1BaseSize / CV_RATIO * constInfo.s2BaseSize; + uint32_t mm2ResultSize = constInfo.s1BaseSize / CV_RATIO * 512; + uint32_t ubAddr = 0; + bmm2Buffers = {LocalTensor(TPosition::VECIN, ubAddr, mm2ResultSize), 0}; + ubAddr += (mm2ResultSize * sizeof(T)); + bmm1Buffers[0] = {LocalTensor(TPosition::VECIN, ubAddr, mm1ResultSize), 0}; + ubAddr += (mm1ResultSize * sizeof(T)); + bmm1Buffers[1] = {LocalTensor(TPosition::VECIN, ubAddr, mm1ResultSize), 1}; + ubAddr += (mm1ResultSize * sizeof(T)); + vUbBase = ubAddr; + + if ASCEND_IS_AIV { + CrossCoreSetFlag(CROSSCORE_BMM1(bmm1Buffers[0].idx)); + CrossCoreSetFlag(CROSSCORE_BMM1(bmm1Buffers[1].idx)); + CrossCoreSetFlag(CROSSCORE_BMM2); + } + + if constexpr (IS_SPLIT_G || TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + uint32_t v0ResSize = constInfo.s2BaseSize * 512U * sizeof(Q_T); + int64_t v0ResTotalOffset; + if constexpr (IS_SPLIT_G) { + v0ResTotalOffset = v0ResSize * 3 * (aicIdx >> 1U); + } else { + v0ResTotalOffset = v0ResSize * 3 * aicIdx; + } + v0ResGmBufferManager.Init(workspace + v0ResTotalOffset); + v0ResGmBuffers.Init(v0ResGmBufferManager, v0ResSize); + v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, CROSSCORE_V0RES(0)); + v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, CROSSCORE_V0RES(1)); + v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, CROSSCORE_V0RES(2)); + } + int64_t fdStagingOffset = 0LL; + if constexpr (IS_SPLIT_G || TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + constexpr int64_t TRIPLE_BUFFER_NUM = 3LL; + int64_t v0ResSize = static_cast(constInfo.s2BaseSize) * constInfo.dSize * sizeof(Q_T); + uint32_t v0LogicalSlotCount = IS_SPLIT_G ? (GetBlockNum() >> 1U) : GetBlockNum(); + fdStagingOffset = v0ResSize * TRIPLE_BUFFER_NUM * v0LogicalSlotCount; + fdStagingOffset += TRIPLE_BUFFER_NUM * constInfo.s2BaseSize * sizeof(int32_t) * GetBlockNum(); + if constexpr (IS_VEC_S2PHYADDR) { + int64_t totalBS1 = (LAYOUT_T == SMLA_LAYOUT::TND) ? + static_cast(constInfo.s1Size) : + static_cast(constInfo.bSize) * constInfo.s1Size; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + fdStagingOffset += totalBS1 * constInfo.alignedOriSparseBlockCount * sizeof(int64_t); + } + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + fdStagingOffset += totalBS1 * constInfo.alignedCmpSparseBlockCount * sizeof(int64_t); + } + } + } + fdStagingBufferManager.Init(workspace + fdStagingOffset); + constexpr uint32_t FD_MAX_SUM_REGION_NUM = 2U; + uint32_t gSize = static_cast(constInfo.gSize); + uint32_t combineElemSize = + gSize * constInfo.dSize + + FD_MAX_SUM_REGION_NUM * gSize * static_cast(AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW); + if constexpr (IS_BATCH_CONSISTENCY) { + uint32_t intraCoreSlotNum = IS_SPLIT_G ? GetBlockNum() : (GetBlockNum() << 1U); + uint32_t intraCoreCombineSize = intraCoreSlotNum * combineElemSize * sizeof(float); + uint32_t crossCoreCombineSize = + GetBlockNum() * BATCH_CONSISTENCY_MAX_REDUCE_BLOCK_NUM * combineElemSize * sizeof(float); + intraCoreCombineBuffer.Init(fdStagingBufferManager, intraCoreCombineSize); + crossCoreCombineBuffer.Init(fdStagingBufferManager, crossCoreCombineSize); + } else { + uint32_t fdSlotCount = static_cast(AttentionCommon::FD_MAX_S2_SPLIT_NUM) * + (IS_SPLIT_G ? (GetBlockNum() >> 1U) : GetBlockNum()); + fdStagingBuffer.Init(fdStagingBufferManager, fdSlotCount * combineElemSize * sizeof(float)); + } +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::InitLocalBuffer() +{ + vecBlock.InitLocalBuffer(constInfo, vUbBase); + cubeBlock.InitLocalBuffer(l1CubeBase); +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::ComputeConstexpr() +{ + // 计算轴的乘积 + + constInfo.s1S2 = constInfo.s1Size * constInfo.s2Size; + constInfo.gS1 = constInfo.gSize * constInfo.s1Size; + constInfo.n2G = constInfo.n2Size * constInfo.gSize; + + constInfo.s1Dv = constInfo.s1Size * constInfo.dSizeV; + constInfo.s2Dv = constInfo.s2Size * constInfo.dSizeV; + constInfo.n2Dv = constInfo.n2Size * constInfo.dSizeV; + constInfo.gDv = constInfo.gSize * constInfo.dSizeV; + constInfo.gS1Dv = constInfo.gSize * constInfo.s1Dv; + constInfo.n2S2Dv = constInfo.n2Size * constInfo.s2Dv; + constInfo.n2GDv = constInfo.n2Size * constInfo.gDv; + constInfo.s2BaseN2Dv = constInfo.s2BaseSize * constInfo.n2Dv; + constInfo.n2GS1Dv = constInfo.n2Size * constInfo.gS1Dv; + + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + // (BS)ND + constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv; + + constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize; + constInfo.mm1Kb = constInfo.n2Size * constInfo.dSize; + if ASCEND_IS_AIV { + constInfo.attentionOutStride = (constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T); + } + } else if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND) { + // BSH/BSNGD + constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv; + constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize; + constInfo.mm1Kb = constInfo.n2Size * constInfo.dSize; + if ASCEND_IS_AIV { + constInfo.attentionOutStride = (constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T); + } + } +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::Process() +{ + // SyncAll Cube和Vector都需要调用 + if constexpr (IS_VEC_S2PHYADDR) { + SyncAll(); + } else if (this->constInfo.needInit) { + SyncAll(); + } + FdRunInfo fdRunInfo; + if ASCEND_IS_AIV { + ParseFdRunInfo(fdRunInfo); + } + ProcessMainLoop(); + if ASCEND_IS_AIV { + SyncAll(); + if (fdRunInfo.coreEnable) { + this->vecBlock.ProcessFlashDecode(fdRunInfo, this->constInfo); + } + } + FreeEvent(); +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::ProcessMainLoop() +{ + uint32_t hasLoad = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_CORE_ENABLE_INDEX, false)); + int64_t maxS2LoopCnt = 0; + if constexpr (IS_SPLIT_G) { + maxS2LoopCnt = static_cast(metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_S2_MAX_NUM, false))); + } + if (hasLoad == 0) { + if ASCEND_IS_AIC { + if constexpr (IS_SPLIT_G) { + for (int64_t loopCnt = 0; loopCnt < maxS2LoopCnt; loopCnt++) { + CrossCoreSetFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + CrossCoreWaitFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + } + } + } + return; + } + + // 从meta data解析分核信息 + uint32_t bN2StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_BN2_START_INDEX, false)); + uint32_t gS1StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_M_START_INDEX, false)); + uint32_t s2StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_S2_START_INDEX, false)); + uint32_t bN2EndIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_BN2_END_INDEX, false)); + uint32_t nextGs1Idx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_M_END_INDEX, false)); + uint32_t s2EndIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_S2_END_INDEX, false)); + uint32_t firstFdDataWorkspaceIdx = + metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX, false)); + + uint32_t s2LoopLimit = 0; + + if (nextGs1Idx != 0 || s2EndIdx != 0) { + bN2EndIdx++; + } + + int64_t taskId = 0; + bool notLast = true; + bool isFirstLoop = true; + RunInfo runInfo[4]; + RunParamStr runParam; + runParam.firstFdDataWorkspaceIdx = firstFdDataWorkspaceIdx; + int64_t multiCoreInnerIdx = 1; + int64_t s2SplitIdxCounter = 0; + for (int64_t bnIdx = bN2StartIdx; bnIdx < bN2EndIdx; bnIdx++) { + bool lastBN = (bnIdx == bN2EndIdx - 1); + runParam.boIdx = bnIdx; + runParam.n2oIdx = 0; + ComputeParamBatch( + runParam, this->constInfo, this->cuSeqlensQGm, this->cuSeqlensOriKvGm, this->cuSeqlensCmpKvGm, + this->actualSeqQlenGm, this->actualSeqOriKvlenGm, this->actualSeqCmpKvlenGm, this->cmpResidualKvGm, + this->hasActualSeqQlen, this->hasActualSeqOriKvlen, this->hasActualSeqCmpKvlen, this->hasCuSeqlensCmpKv); + ComputeS1LoopInfo(runParam, this->constInfo, lastBN, nextGs1Idx, gS1StartIdx, s2EndIdx); + + int64_t gS1LoopEnd = lastBN ? (runParam.gs1LoopEndIdx + PRELOAD_NUM) : runParam.gs1LoopEndIdx; + for (int64_t gS1Index = runParam.gs1LoopStartIdx; gS1Index < gS1LoopEnd; gS1Index++) { + bool notLastThreeLoop = true; + bool notLastTwoLoop = true; + if (lastBN) { + int32_t extraGS1 = gS1Index - runParam.gs1LoopEndIdx; + switch (extraGS1) { + case 0: + notLastThreeLoop = false; + break; + case 1: + notLastTwoLoop = false; + notLastThreeLoop = false; + break; + case 2: + notLast = false; + notLastTwoLoop = false; + notLastThreeLoop = false; + break; + default: + break; + } + } + if (notLastThreeLoop) { + this->ComputeAxisIdxByBnAndGs1(bnIdx, gS1Index, runParam); + bool s1NoNeedCalc = + ComputeParamS1(runParam, this->constInfo, gS1Index, this->cuSeqlensQGm); + bool s2NoNeedCalc = ComputeS2LoopInfo( + bnIdx, gS1Index, this->cuSeqlensQGm, oriTopkLengthGm, cmpTopkLengthGm, runParam, this->constInfo); + if constexpr (IS_BATCH_CONSISTENCY) { + int64_t oriLoad = runParam.s2OriLineEndIdx - runParam.s2OriLineStartIdx; + int64_t cmpLoad = runParam.s2CmpLineEndIdx - runParam.s2CmpLineStartIdx; + int64_t totalLoad = oriLoad + cmpLoad; + int64_t s2BaseSize = static_cast(constInfo.s2BaseSize); + int64_t rawReductionBlockSize = totalLoad / 32LL; + int64_t reductionBlockSize = (rawReductionBlockSize + s2BaseSize - 1LL) / s2BaseSize * s2BaseSize; + runParam.baseBlockNumPerReductionBlock = + reductionBlockSize > 0 ? reductionBlockSize / s2BaseSize : 1LL; + } + if (!s2NoNeedCalc) { + bool isFirstS2RangeTask = (bnIdx == bN2StartIdx && gS1Index == runParam.gs1LoopStartIdx); + bool isLastS2RangeTask = (lastBN && gS1Index == runParam.gs1LoopEndIdx - 1); + int64_t s2StartPoint = ConvertS2MetadataBlockToToken(runParam, this->constInfo, s2StartIdx); + int64_t s2EndPoint = (isLastS2RangeTask && s2EndIdx == 0) ? + 0 : + ConvertS2MetadataBlockToToken(runParam, this->constInfo, s2EndIdx); + s2NoNeedCalc = ApplyS2MetadataRange(runParam, this->constInfo, s2StartPoint, s2EndPoint, + isFirstS2RangeTask, isLastS2RangeTask); + } else { + runParam.isCrossCoreSplit = false; + } + // s1和s2有任意一个不需要算, 则continue, 如果是当前核最后一次循环,则补充计算taskIdx+2的部分 + if (s1NoNeedCalc || s2NoNeedCalc) { + continue; + } + if constexpr (!IS_BATCH_CONSISTENCY) { + if (runParam.isCrossCoreSplit) { + runParam.s2SplitIdx = s2SplitIdxCounter++; + } + } + s2LoopLimit = runParam.s2LoopEndIdx - 1; + if constexpr (IS_SPLIT_G) { + maxS2LoopCnt -= (s2LoopLimit + 1); + } + } else { + s2LoopLimit = 0; + } + for (int64_t s2LoopCount = 0; s2LoopCount <= s2LoopLimit; ++s2LoopCount) { + if constexpr (IS_BATCH_CONSISTENCY) { + int64_t safeBaseBlockNum = + runParam.baseBlockNumPerReductionBlock > 0 ? runParam.baseBlockNumPerReductionBlock : 1LL; + if (runParam.isCrossCoreSplit && s2LoopCount % safeBaseBlockNum == 0) { + runParam.s2SplitIdx = s2SplitIdxCounter++; + } + } + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (notLastThreeLoop) { + RunInfo &runInfo1 = runInfo[taskId % 4]; + this->SetRunInfo(runInfo1, runParam, taskId, s2LoopCount, s2LoopLimit, multiCoreInnerIdx); + } + if ASCEND_IS_AIV { + if (notLastThreeLoop) { + RunInfo &runInfo1 = runInfo[taskId % 4]; + this->vecBlock.ProcessVec0(this->v0ResGmBuffers.Get(runInfo1.taskIdMod3), runInfo1, + this->constInfo); + } + if (taskId > 1 && notLast) { + uint32_t bmm1Slot = bmm1GetFlag; + bmm1GetFlag ^= 1; + uint32_t l1PSlot = l1PGetFlag; + l1PGetFlag ^= 1; + auto &runInfo2 = runInfo[(taskId + 2) % 4]; + this->vecBlock.ProcessVec1(this->l1PBuffers[l1PSlot], this->bmm1Buffers[bmm1Slot], runInfo2, + this->constInfo); + } + if (taskId > 2) { + RunInfo &runInfo3 = runInfo[(taskId + 1) % 4]; + this->vecBlock.ProcessVec2(this->bmm2Buffers, runInfo3, this->constInfo); + } + } else { + if (taskId > 0 && notLastTwoLoop) { + RunInfo &runInfo1 = runInfo[(taskId + 3) % 4]; + this->cubeBlock.IterateLoadQK(this->v0ResGmBuffers.Get(runInfo1.taskIdMod3), runInfo1, + this->constInfo, isFirstLoop); + isFirstLoop = false; + } else { + if constexpr (IS_SPLIT_G) { + if (taskId > 0 && maxS2LoopCnt > 0) { + maxS2LoopCnt--; + CrossCoreSetFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + CrossCoreWaitFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + } + } + } + if (taskId > 1 && notLast) { + uint32_t bmm1Slot = bmm1GetFlag; + bmm1GetFlag ^= 1; + auto &runInfo2 = runInfo[(taskId + 2) % 4]; + RunInfo &runInfoNext = runInfo[(taskId + 3) % 4]; + this->cubeBlock.IterateBmm1(this->bmm1Buffers[bmm1Slot], notLastTwoLoop, runInfoNext, + runInfo2, this->constInfo); + } + if (taskId > 2) { + uint32_t l1PSlot = l1PGetFlag; + l1PGetFlag ^= 1; + RunInfo &runInfo3 = runInfo[(taskId + 1) % 4]; + this->cubeBlock.IterateBmm2(this->bmm2Buffers, this->l1PBuffers[l1PSlot], runInfo3, + this->constInfo); + } + } + } + ++taskId; + } + ++multiCoreInnerIdx; + } + gS1StartIdx = 0; + } + if ASCEND_IS_AIC { + if constexpr (IS_SPLIT_G) { + for (int64_t loopCnt = 0; loopCnt < maxS2LoopCnt; loopCnt++) { + CrossCoreSetFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + CrossCoreWaitFlag<0, PIPE_MTE2>(crossCoreMte2SyncFlagId); + } + } + } +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::ComputeAxisIdxByBnAndGs1( + int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam) +{ + // GS1合轴, 不切G, 只切S1 + runParam.s1oIdx = gS1Index * runParam.qSNumInOneBlock; + if constexpr (IS_SPLIT_G) { + int64_t halfG = (constInfo.gSize + 1) / 2; // ceil(gSize/2), 第一个AIC多处理一行 + runParam.goIdx = (aicIdx % 2 == 0) ? 0 : halfG; + runParam.gSplitSize = (aicIdx % 2 == 0) ? halfG : (constInfo.gSize - halfG); + } else { + runParam.goIdx = 0; + runParam.gSplitSize = constInfo.gSize; + } +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::SetRunInfo( + RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount, int64_t s2LoopLimit, + int64_t multiCoreInnerIdx) +{ + if (s2LoopCount < runParam.oriKvLoopEndIdx) { + runInfo.s2StartIdx = runParam.s2OriLineStartIdx; + runInfo.s2EndIdx = runParam.s2OriLineEndIdx; + } else { + runInfo.s2StartIdx = runParam.s2CmpLineStartIdx; + runInfo.s2EndIdx = runParam.s2CmpLineEndIdx; + } + runInfo.s2LoopCount = s2LoopCount; + if (runInfo.multiCoreInnerIdx != multiCoreInnerIdx) { + runInfo.s1oIdx = runParam.s1oIdx; + runInfo.boIdx = runParam.boIdx; + runInfo.n2oIdx = runParam.n2oIdx; + runInfo.goIdx = runParam.goIdx; + runInfo.multiCoreInnerIdx = multiCoreInnerIdx; + runInfo.multiCoreIdxMod2 = multiCoreInnerIdx & 1; + runInfo.multiCoreIdxMod3 = multiCoreInnerIdx % 3; + } + + runInfo.taskId = taskId; + runInfo.taskIdMod2 = taskId & 1; + runInfo.taskIdMod3 = taskId % 3; + runInfo.s2LoopLimit = s2LoopLimit; + + runInfo.actualS1Size = runParam.actualS1Size; + runInfo.attentionOutOffset = runParam.attentionOutOffset; + runInfo.sOuterOffset = runParam.sOuterOffset; + runInfo.firstFdDataWorkspaceIdx = runParam.firstFdDataWorkspaceIdx; + runInfo.isCrossCoreSplit = runParam.isCrossCoreSplit; + runInfo.s2SplitIdx = runParam.s2SplitIdx; + runInfo.isFirstS2SplitCore = runParam.isFirstS2SplitCore; + int64_t safeBaseBlockNum = + runParam.baseBlockNumPerReductionBlock > 0 ? runParam.baseBlockNumPerReductionBlock : 1LL; + int64_t baseBlockIdInReduceBlock = s2LoopCount % safeBaseBlockNum; + runInfo.reduceBlockId = s2LoopCount / safeBaseBlockNum; + runInfo.isFirstBase = baseBlockIdInReduceBlock == 0; + runInfo.isLastBase = baseBlockIdInReduceBlock == safeBaseBlockNum - 1LL || s2LoopCount == s2LoopLimit; + runInfo.needReduce = runInfo.reduceBlockId > 0; + this->ComputeBmm1Tail(runInfo, runParam); + InitUniqueRunInfo(runParam, runInfo); +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::InitUniqueRunInfo( + const RunParamStr &runParam, RunInfo &runInfo) +{ + InitTaskParamByRun(runParam, runInfo, constInfo); +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::ComputeBmm1Tail(RunInfo &runInfo, + RunParamStr &runParam) +{ + // ------------------------S1 Base Related--------------------------- + runInfo.s1RealSize = runParam.s1RealSize; + runInfo.halfS1RealSize = runParam.halfS1RealSize; + runInfo.firstHalfS1RealSize = runParam.firstHalfS1RealSize; + runInfo.mRealSize = runParam.mRealSize; + runInfo.halfMRealSize = runParam.halfMRealSize; + runInfo.firstHalfMRealSize = runParam.firstHalfMRealSize; + + runInfo.vec2MBaseSize = runInfo.halfMRealSize; + + // ------------------------S2 Base Related---------------------------- + runInfo.s2RealSize = constInfo.s2BaseSize; + runInfo.s2AlignedSize = runInfo.s2RealSize; + int64_t curS2LoopCnt = (runInfo.s2LoopCount >= runParam.oriKvLoopEndIdx) ? + (runInfo.s2LoopCount - runParam.oriKvLoopEndIdx) : + runInfo.s2LoopCount; + if (runInfo.s2StartIdx + (curS2LoopCnt + 1) * runInfo.s2RealSize > runInfo.s2EndIdx) { + runInfo.s2RealSize = runInfo.s2EndIdx - curS2LoopCnt * runInfo.s2RealSize - runInfo.s2StartIdx; + runInfo.s2AlignedSize = Align(runInfo.s2RealSize); + } +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::ParseFdRunInfo(FdRunInfo &fdRunInfo) +{ + uint32_t aivIdx = static_cast(this->constInfo.aivIdx); + fdRunInfo.coreEnable = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_CORE_ENABLE_INDEX, true)) != 0; + if (!fdRunInfo.coreEnable) { + return; + } + fdRunInfo.bn2Idx = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_BN2_IDX_INDEX, true)); + fdRunInfo.mIdx = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_M_IDX_INDEX, true)); + fdRunInfo.workspaceIdx = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_WORKSPACE_IDX_INDEX, true)); + fdRunInfo.workspaceNum = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_WORKSPACE_NUM_INDEX, true)); + fdRunInfo.mStartIdx = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_M_START_INDEX, true)); + fdRunInfo.mNum = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_M_NUM_INDEX, true)); +} + +template +__aicore__ inline int64_t SparseFlashMlaCsaKernel::ConvertS2MetadataBlockToToken( + const RunParamStr &runParam, const ConstInfo &constInfo, uint32_t s2BlockIdx) +{ + int64_t s2BaseSize = static_cast(constInfo.s2BaseSize); + int64_t oriLen = runParam.s2OriLineEndIdx - runParam.s2OriLineStartIdx; + int64_t cmpLen = runParam.s2CmpLineEndIdx - runParam.s2CmpLineStartIdx; + int64_t reductionBlockSize = runParam.baseBlockNumPerReductionBlock * s2BaseSize; + int64_t oriReductionBlockNum = (oriLen + reductionBlockSize - 1) / reductionBlockSize; + int64_t reductionBlockIdx = static_cast(s2BlockIdx); + if (reductionBlockIdx < oriReductionBlockNum) { + int64_t oriToken = reductionBlockIdx * reductionBlockSize; + return oriToken < oriLen ? oriToken : oriLen; + } + int64_t cmpToken = (reductionBlockIdx - oriReductionBlockNum) * reductionBlockSize; + return oriLen + (cmpToken < cmpLen ? cmpToken : cmpLen); +} + +template +__aicore__ inline void SparseFlashMlaCsaKernel::FreeEvent() +{ + if ASCEND_IS_AIC { + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1Buffers[0].idx)); + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1Buffers[0].idx) + AIV0_AIV1_OFFSET); + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1Buffers[1].idx)); + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1Buffers[1].idx) + AIV0_AIV1_OFFSET); + CrossCoreWaitFlag(CROSSCORE_BMM2); + CrossCoreWaitFlag(CROSSCORE_BMM2 + AIV0_AIV1_OFFSET); + this->cubeBlock.FreeEvent(); + } else { + this->vecBlock.FreeEvent(constInfo); + } +} + +template +__aicore__ inline bool SparseFlashMlaCsaKernel::ApplyS2MetadataRange( + RunParamStr &runParam, ConstInfo &constInfo, int64_t s2StartPoint, int64_t s2EndPoint, bool isFirstS2RangeTask, + bool isLastS2RangeTask) +{ + int64_t oriStart = runParam.s2OriLineStartIdx; + int64_t oriEnd = runParam.s2OriLineEndIdx; + int64_t oriLen = oriEnd - oriStart; + int64_t cmpStart = runParam.s2CmpLineStartIdx; + int64_t cmpEnd = runParam.s2CmpLineEndIdx; + int64_t cmpLen = cmpEnd - cmpStart; + int64_t totalLen = oriLen + cmpLen; + + int64_t effectiveS2EndPoint = (isLastS2RangeTask && s2EndPoint == 0) ? totalLen : s2EndPoint; + int64_t rangeStart = isFirstS2RangeTask ? s2StartPoint : 0; + rangeStart = rangeStart < 0 ? 0 : rangeStart; + rangeStart = rangeStart < totalLen ? rangeStart : totalLen; + int64_t rangeEnd = isLastS2RangeTask ? effectiveS2EndPoint : totalLen; + rangeEnd = rangeEnd < 0 ? 0 : rangeEnd; + rangeEnd = rangeEnd < totalLen ? rangeEnd : totalLen; + if (rangeEnd <= rangeStart) { + runParam.oriKvLoopEndIdx = 0; + runParam.cmpKvLoopEndIdx = 0; + runParam.s2LoopEndIdx = 0; + runParam.isCrossCoreSplit = false; + return true; + } + + bool hasPrevCore = rangeStart > 0; + bool hasNextCore = rangeEnd < totalLen; + runParam.isCrossCoreSplit = hasPrevCore || hasNextCore; + runParam.isFirstS2SplitCore = !hasPrevCore; + + int64_t oriRangeStart = rangeStart < oriLen ? rangeStart : oriLen; + int64_t oriRangeEnd = rangeEnd < oriLen ? rangeEnd : oriLen; + runParam.s2OriLineStartIdx = oriStart + oriRangeStart; + runParam.s2OriLineEndIdx = oriStart + oriRangeEnd; + + int64_t cmpRangeStart = rangeStart > oriLen ? rangeStart - oriLen : 0; + cmpRangeStart = cmpRangeStart < cmpLen ? cmpRangeStart : cmpLen; + int64_t cmpRangeEnd = rangeEnd > oriLen ? rangeEnd - oriLen : 0; + cmpRangeEnd = cmpRangeEnd < cmpLen ? cmpRangeEnd : cmpLen; + runParam.s2CmpLineStartIdx = cmpStart + cmpRangeStart; + runParam.s2CmpLineEndIdx = cmpStart + cmpRangeEnd; + + int64_t s2BaseSize = static_cast(constInfo.s2BaseSize); + int64_t oriRangeLen = runParam.s2OriLineEndIdx - runParam.s2OriLineStartIdx; + int64_t cmpRangeLen = runParam.s2CmpLineEndIdx - runParam.s2CmpLineStartIdx; + runParam.oriKvLoopEndIdx = (oriRangeLen + s2BaseSize - 1) / s2BaseSize; + runParam.cmpKvLoopEndIdx = (cmpRangeLen + s2BaseSize - 1) / s2BaseSize; + runParam.s2LoopEndIdx = runParam.oriKvLoopEndIdx + runParam.cmpKvLoopEndIdx; + return runParam.s2LoopEndIdx == 0; +} +} // namespace SMLAKernel +#endif // SPARSE_FLASH_MLA_CSA_KERNEL_ARCH35_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_kvcache.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_kvcache.h new file mode 100644 index 000000000000..a95e12e6967b --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_kvcache.h @@ -0,0 +1,415 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_kvcache.h + * \brief + */ +#ifndef SPARSE_FLASH_MLA_KVCACHE_H +#define SPARSE_FLASH_MLA_KVCACHE_H + +#include "kernel_operator.h" +#include "kernel_operator_list_tensor_intf.h" +#include "sparse_flash_mla_common_arch35.h" +#include "util_regbase.h" + +using namespace matmul; +using namespace regbaseutil; +using namespace AscendC; +using namespace AscendC::Impl::Detail; +using namespace SMLAKernel; + +TEMPLATE_INTF +__aicore__ inline void GetSingleCoreParam(RunParamStr &runParam, const ConstInfo &constInfo, + GlobalTensor &cuSeqlensQGm, GlobalTensor &cuSeqlensOriKvGm, + GlobalTensor &cuSeqlensCmpKvGm, + GlobalTensor &actualSeqQlenGm, + GlobalTensor &actualSeqOriKvlenGm, + GlobalTensor &actualSeqCmpKvlenGm, + GlobalTensor &cmpResidualKvGm, bool hasActualSeqQlen, + bool hasActualSeqOriKvlen, bool hasActualSeqCmpKvlen, bool hasCuSeqlensCmpKv) +{ + int32_t actualS1Size = 0; + int32_t actualS2OriSize = 0; + int32_t actualS2CmpSize = 0; + int32_t bIdx = runParam.boIdx; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + actualS1Size = (!hasActualSeqQlen) ? (cuSeqlensQGm.GetValue(bIdx + 1) - cuSeqlensQGm.GetValue(bIdx)) : + actualSeqQlenGm.GetValue(bIdx); + } else { + actualS1Size = (!hasActualSeqQlen) ? constInfo.s1Size : actualSeqQlenGm.GetValue(bIdx); + } + + if (constInfo.isActualLenDimsOriKVNull) { + actualS2OriSize = constInfo.s2Size; + } else { + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + if (hasActualSeqOriKvlen) { + actualS2OriSize = actualSeqOriKvlenGm.GetValue(bIdx); + } else { + actualS2OriSize = cuSeqlensOriKvGm.GetValue(bIdx + 1) - cuSeqlensOriKvGm.GetValue(bIdx); + } + } else { + actualS2OriSize = actualSeqOriKvlenGm.GetValue(bIdx); + } + } + + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + if (hasActualSeqCmpKvlen) { + actualS2CmpSize = actualSeqCmpKvlenGm.GetValue(bIdx); + } else if (hasCuSeqlensCmpKv) { + actualS2CmpSize = cuSeqlensCmpKvGm.GetValue(bIdx + 1) - cuSeqlensCmpKvGm.GetValue(bIdx); + } + } else { + if (!hasActualSeqCmpKvlen) { + actualS2CmpSize = constInfo.cmpS2Size; + } else { + actualS2CmpSize = actualSeqCmpKvlenGm.GetValue(bIdx); + } + } + } + + runParam.actualS1Size = actualS1Size; + runParam.actualS2OriSize = actualS2OriSize; + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + runParam.actualS2CmpSize = actualS2CmpSize; + if (constInfo.cmpMaskMode == 0) { + runParam.nextTokensPerBatchCmp = runParam.actualS2CmpSize * constInfo.cmpRatio; + } else { + runParam.cmpResidual = (cmpResidualKvGm.GetPhyAddr() != nullptr) ? cmpResidualKvGm.GetValue(bIdx) : 0; + runParam.nextTokensPerBatchCmp = + (int64_t)runParam.actualS2CmpSize * constInfo.cmpRatio + runParam.cmpResidual - runParam.actualS1Size; + } + } + + if (constInfo.oriMaskMode == 3) { // 3: RightDownCausal模式 + runParam.nextTokensPerBatchOri = runParam.actualS2OriSize - runParam.actualS1Size; + runParam.preTokensPerBatchOri = runParam.actualS1Size; + } else if (constInfo.oriMaskMode == 4) { // 4: Band模式 + const int64_t casualOffset = runParam.actualS2OriSize - runParam.actualS1Size; + + runParam.preTokensPerBatchOri = + (constInfo.oriWinLeft == -1) ? runParam.actualS1Size : constInfo.oriWinLeft - casualOffset; + + runParam.nextTokensPerBatchOri = + (constInfo.oriWinRight == -1) ? runParam.actualS2OriSize : casualOffset + constInfo.oriWinRight; + + runParam.preTokensPerBatchOri = Min(runParam.preTokensPerBatchOri, static_cast(runParam.actualS1Size)); + } else if (constInfo.oriMaskMode == 0) { + runParam.nextTokensPerBatchOri = runParam.actualS2OriSize; + runParam.preTokensPerBatchOri = runParam.actualS1Size; + } +} + +TEMPLATE_INTF +__aicore__ inline void ComputeParamBatch(RunParamStr &runParam, const ConstInfo &constInfo, + GlobalTensor &cuSeqlensQGm, GlobalTensor &cuSeqlensOriKvGm, + GlobalTensor &cuSeqlensCmpKvGm, + GlobalTensor &actualSeqQlenGm, + GlobalTensor &actualSeqOriKvlenGm, + GlobalTensor &actualSeqCmpKvlenGm, + GlobalTensor &cmpResidualKvGm, bool hasActualSeqQlen, + bool hasActualSeqOriKvlen, bool hasActualSeqCmpKvlen, bool hasCuSeqlensCmpKv) +{ + GetSingleCoreParam(runParam, constInfo, cuSeqlensQGm, cuSeqlensOriKvGm, cuSeqlensCmpKvGm, + actualSeqQlenGm, actualSeqOriKvlenGm, actualSeqCmpKvlenGm, cmpResidualKvGm, + hasActualSeqQlen, hasActualSeqOriKvlen, hasActualSeqCmpKvlen, + hasCuSeqlensCmpKv); +} + +TEMPLATE_INTF +__aicore__ inline void ComputeS1LoopInfo(RunParamStr &runParam, const ConstInfo &constInfo, bool lastBN, + int64_t nextGs1Idx, int64_t gS1StartIdx, int64_t s2EndIdx = 0) +{ + runParam.qSNumInOneBlock = 1; + runParam.gs1LoopStartIdx = gS1StartIdx; + if (TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if constexpr (TEMPLATE_MODE == SMLATemplateMode::HCA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE) { + int64_t skipThreshold = 0; + if (runParam.nextTokensPerBatchOri < 0 && runParam.nextTokensPerBatchCmp < 0) { + skipThreshold = Min(-runParam.nextTokensPerBatchOri, -runParam.nextTokensPerBatchCmp); + } + if (skipThreshold > 0) { + int64_t gs1LoopStartIdx = skipThreshold / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock; + if (gs1LoopStartIdx > gS1StartIdx) { + runParam.gs1LoopStartIdx = gs1LoopStartIdx; + } + } + } else { + if (runParam.nextTokensPerBatchOri < 0) { + int64_t gs1LoopStartIdx = + runParam.nextTokensPerBatchOri * (-1) / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock; + if (gs1LoopStartIdx > gS1StartIdx) { + runParam.gs1LoopStartIdx = gs1LoopStartIdx; + } + } + } + } + + int32_t gs1LoopEndIdx = 0; + if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + gs1LoopEndIdx = runParam.actualS1Size; + } else { // SWA/HCA + // 不需要取topk, 每次计算gSize行, 循环qs次 + gs1LoopEndIdx = (runParam.actualS1Size + runParam.qSNumInOneBlock - 1) / runParam.qSNumInOneBlock; + } + // 不是最后一个bn, 赋值souterBlockNum + if (!lastBN) { + runParam.gs1LoopEndIdx = gs1LoopEndIdx; + } else { // 最后一个bn, 从数组下一个元素取值 + uint32_t actualNextGs1Idx = s2EndIdx == 0 ? nextGs1Idx : nextGs1Idx + 1; + runParam.gs1LoopEndIdx = (nextGs1Idx == 0 && s2EndIdx == 0) ? gs1LoopEndIdx : actualNextGs1Idx; + } + + if (runParam.gs1LoopStartIdx > runParam.gs1LoopEndIdx) { + runParam.gs1LoopStartIdx = runParam.gs1LoopEndIdx; + } +} + +TEMPLATE_INTF +__aicore__ inline void ComputeSouterParam(RunParamStr &runParam, const ConstInfo &constInfo, uint32_t sOuterLoopIdx) +{ + int64_t cubeSOuterOffset = sOuterLoopIdx * runParam.qSNumInOneBlock; + if (runParam.actualS1Size == 0) { + runParam.s1RealSize = 0; + runParam.mRealSize = 0; + } else { + runParam.s1RealSize = Min(runParam.qSNumInOneBlock, runParam.actualS1Size - cubeSOuterOffset); + runParam.mRealSize = runParam.s1RealSize * constInfo.gSize; + if constexpr (IS_SPLIT_G) { + runParam.mRealSize = runParam.s1RealSize * runParam.gSplitSize; + } + } + + runParam.cubeMOuterOffset = cubeSOuterOffset * constInfo.gSize; + runParam.halfMRealSize = (runParam.mRealSize + 1) >> 1; + runParam.firstHalfMRealSize = runParam.halfMRealSize; + if (constInfo.subBlockIdx == 1) { + runParam.halfMRealSize = runParam.mRealSize - runParam.halfMRealSize; + runParam.mOuterOffset = runParam.cubeMOuterOffset + runParam.firstHalfMRealSize; + } else { + runParam.mOuterOffset = runParam.cubeMOuterOffset; + } + + runParam.halfS1RealSize = (runParam.s1RealSize + 1) >> 1; + runParam.firstHalfS1RealSize = runParam.halfS1RealSize; + if (constInfo.subBlockIdx == 1) { + runParam.halfS1RealSize = runParam.s1RealSize - runParam.halfS1RealSize; + runParam.sOuterOffset = cubeSOuterOffset + runParam.firstHalfMRealSize / constInfo.gSize; + } else { + runParam.sOuterOffset = cubeSOuterOffset; + } + runParam.cubeSOuterOffset = cubeSOuterOffset; +} + +TEMPLATE_INTF +__aicore__ inline void LoopSOuterOffsetInit(RunParamStr &runParam, const ConstInfo &constInfo, int32_t sIdx, + GlobalTensor &cuSeqlensQGm) +{ + if ASCEND_IS_AIV { + int64_t seqOffset = 0; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + seqOffset = cuSeqlensQGm.GetValue(sIdx); + } else { + seqOffset = sIdx * constInfo.s1Size; + } + + int64_t attentionOutSeqOffset = seqOffset * constInfo.n2GDv; + if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND || LAYOUT_T == SMLA_LAYOUT::TND) { + runParam.attentionOutOffset = attentionOutSeqOffset + runParam.sOuterOffset * constInfo.n2GDv + + runParam.n2oIdx * constInfo.gDv + runParam.goIdx * constInfo.dSizeV; + } + if (constInfo.subBlockIdx == 1) { + runParam.attentionOutOffset += runParam.firstHalfMRealSize * constInfo.dSizeV; + } + if (constInfo.returnSoftmaxLse) { + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + // [N2, T, G] (TND) + runParam.softmaxLseOffset = runParam.n2oIdx * constInfo.s1Size * constInfo.gSize + + (seqOffset + runParam.sOuterOffset) * constInfo.gSize; + } else { + // [B, N2, S1, G] (BSND) + runParam.softmaxLseOffset = sIdx * constInfo.n2Size * constInfo.s1Size * constInfo.gSize + + runParam.n2oIdx * constInfo.s1Size * constInfo.gSize + + runParam.sOuterOffset * constInfo.gSize; + } + if (IS_SPLIT_G) { + runParam.softmaxLseOffset += runParam.goIdx; + } + if (constInfo.subBlockIdx == 1) { + runParam.softmaxLseOffset += runParam.firstHalfMRealSize; + } + } + } +} + +TEMPLATE_INTF +__aicore__ inline bool ComputeParamS1(RunParamStr &runParam, const ConstInfo &constInfo, uint32_t sOuterLoopIdx, + GlobalTensor &cuSeqlensQGm) +{ + if (TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if constexpr (TEMPLATE_MODE == SMLATemplateMode::HCA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE) { + int64_t skipThreshold = 0; + if (runParam.nextTokensPerBatchOri < 0 && runParam.nextTokensPerBatchCmp < 0) { + skipThreshold = Min(-runParam.nextTokensPerBatchOri, -runParam.nextTokensPerBatchCmp); + } + if (skipThreshold > 0) { + if (runParam.s1oIdx < skipThreshold / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock) { + return true; + } + } + } else { + if (runParam.nextTokensPerBatchOri < 0) { + if (runParam.s1oIdx < + (runParam.nextTokensPerBatchOri * (-1)) / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock) { + return true; + } + } + } + } + + ComputeSouterParam(runParam, constInfo, sOuterLoopIdx); + + LoopSOuterOffsetInit(runParam, constInfo, runParam.boIdx, cuSeqlensQGm); + return false; +} + +TEMPLATE_INTF +__aicore__ inline bool ComputeLastBN(RunParamStr &runParam, GlobalTensor &cuSeqlensQGm) +{ + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + // TND格式下 相邻Batch中当actualSeqQlen相等时则返回true + if (runParam.boIdx > 0 && + cuSeqlensQGm.GetValue(runParam.boIdx + 1) - cuSeqlensQGm.GetValue(runParam.boIdx) == 0) { + return true; + } + } + return false; +} + +TEMPLATE_INTF +__aicore__ inline int64_t ClipSInnerTokenCube(int64_t sInnerToken, int64_t minValue, int64_t maxValue) +{ + sInnerToken = sInnerToken > minValue ? sInnerToken : minValue; + sInnerToken = sInnerToken < maxValue ? sInnerToken : maxValue; + return sInnerToken; +} + +TEMPLATE_INTF +__aicore__ inline bool ComputeS2LoopInfo(int64_t bnIndex, int64_t gS1Index, GlobalTensor &cuSeqlensQGm, + GlobalTensor &oriTopkLengthGm, GlobalTensor &cmpTopkLengthGm, + RunParamStr &runParam, const ConstInfo &constInfo) +{ + if (TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + if (runParam.actualS2OriSize == 0) { + runParam.oriKvLoopEndIdx = 0; + runParam.cmpKvLoopEndIdx = 0; + runParam.s2LoopEndIdx = 0; + runParam.s2CmpLineStartIdx = 0; + return true; + } + } + uint32_t s2BaseSize = constInfo.s2BaseSize; + + // 计算topk length + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + uint64_t actualSeqQPrefixSum = cuSeqlensQGm.GetValue(runParam.boIdx); + runParam.oriSparseBlockCount = + constInfo.hasOriTopkLength ? + Min(oriTopkLengthGm.GetValue(actualSeqQPrefixSum + runParam.s1oIdx), constInfo.oriSparseBlockCount) : + constInfo.oriSparseBlockCount; + runParam.cmpSparseBlockCount = + constInfo.hasCmpTopkLength ? + Min(cmpTopkLengthGm.GetValue(actualSeqQPrefixSum + runParam.s1oIdx), constInfo.cmpSparseBlockCount) : + constInfo.cmpSparseBlockCount; + } else { + uint64_t bsndTopkIdx = runParam.boIdx * constInfo.s1Size + runParam.s1oIdx; + runParam.oriSparseBlockCount = constInfo.hasOriTopkLength ? + Min(oriTopkLengthGm.GetValue(bsndTopkIdx), constInfo.oriSparseBlockCount) : + constInfo.oriSparseBlockCount; + runParam.cmpSparseBlockCount = constInfo.hasCmpTopkLength ? + Min(cmpTopkLengthGm.GetValue(bsndTopkIdx), constInfo.cmpSparseBlockCount) : + constInfo.cmpSparseBlockCount; + } + + // orikv + runParam.s2OriLineStartIdx = ClipSInnerTokenCube( + runParam.cubeSOuterOffset - runParam.preTokensPerBatchOri, 0, runParam.actualS2OriSize); + runParam.s2OriLineEndIdx = ClipSInnerTokenCube( + runParam.cubeSOuterOffset + runParam.nextTokensPerBatchOri + runParam.s1RealSize, 0, runParam.actualS2OriSize); + if constexpr (TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + int64_t oriSparseRangeLen = runParam.s2OriLineEndIdx - runParam.s2OriLineStartIdx; + runParam.s2OriLineStartIdx = 0; + runParam.s2OriLineEndIdx = Min(oriSparseRangeLen, runParam.oriSparseBlockCount); + runParam.s2OriLineEndIdx = Min(runParam.s2OriLineEndIdx, runParam.actualS2OriSize); + } + runParam.oriKvLoopEndIdx = (runParam.s2OriLineEndIdx - runParam.s2OriLineStartIdx + s2BaseSize - 1) / s2BaseSize; + + // cmpkv + if constexpr (TEMPLATE_MODE == SMLATemplateMode::SWA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + runParam.s2CmpLineStartIdx = 0; + runParam.s2CmpLineEndIdx = 0; + runParam.cmpKvLoopEndIdx = 0; + } else if constexpr (TEMPLATE_MODE == SMLATemplateMode::HCA_TEMPLATE_MODE) { + runParam.s2CmpLineStartIdx = 0; + runParam.s2CmpLineEndIdx = ClipSInnerTokenCube( + (runParam.cubeSOuterOffset + runParam.s1RealSize + runParam.nextTokensPerBatchCmp) / constInfo.cmpRatio, 0, + runParam.actualS2CmpSize); + runParam.s2CmpLineEndIdx = Min(runParam.s2CmpLineEndIdx, runParam.actualS2CmpSize); + runParam.cmpKvLoopEndIdx = (runParam.s2CmpLineEndIdx + s2BaseSize - 1) / s2BaseSize; + } else if constexpr (TEMPLATE_MODE == SMLATemplateMode::CSA_TEMPLATE_MODE || + TEMPLATE_MODE == SMLATemplateMode::ORI_CMP_SPARSE_TEMPLATE_MODE) { + runParam.s2CmpLineStartIdx = 0; + runParam.s2CmpLineEndIdx = ClipSInnerTokenCube( + (runParam.cubeSOuterOffset + runParam.s1RealSize + runParam.nextTokensPerBatchCmp) / constInfo.cmpRatio, 0, + runParam.actualS2CmpSize); + runParam.s2CmpLineEndIdx = Min(runParam.s2CmpLineEndIdx, runParam.cmpSparseBlockCount); + runParam.s2CmpLineEndIdx = Min(runParam.s2CmpLineEndIdx, runParam.actualS2CmpSize); + runParam.cmpKvLoopEndIdx = (runParam.s2CmpLineEndIdx + s2BaseSize - 1) / s2BaseSize; + } + + runParam.s2LoopEndIdx = runParam.oriKvLoopEndIdx + runParam.cmpKvLoopEndIdx; + return (runParam.s2LoopEndIdx == 0); +} + +TEMPLATE_INTF +__aicore__ inline void InitTaskParamByRun(const RunParamStr &runParam, RunInfo &runInfo, const ConstInfo &constInfo) +{ + runInfo.boIdx = runParam.boIdx; + runInfo.preTokensPerBatchOri = runParam.preTokensPerBatchOri; + runInfo.nextTokensPerBatchOri = runParam.nextTokensPerBatchOri; + runInfo.actualS1Size = runParam.actualS1Size; + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE && + TEMPLATE_MODE != SMLATemplateMode::ORI_SPARSE_TEMPLATE_MODE) { + runInfo.actualS2CmpSize = runParam.actualS2CmpSize; + runInfo.cmpResidual = runParam.cmpResidual; + } + runInfo.softmaxLseOffset = runParam.softmaxLseOffset; + runInfo.qSNumInOneBlock = runParam.qSNumInOneBlock; + runInfo.oriKvLoopEndIdx = runParam.oriKvLoopEndIdx; + runInfo.cmpKvLoopEndIdx = runParam.cmpKvLoopEndIdx; + runInfo.isCmp = runInfo.s2LoopCount >= runInfo.oriKvLoopEndIdx; + runInfo.oriSparseBlockCount = runParam.oriSparseBlockCount; + runInfo.cmpSparseBlockCount = runParam.cmpSparseBlockCount; +} + +#endif // SPARSE_FLASH_MLA_KVCACHE_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_swa_kernel_arch35.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_swa_kernel_arch35.h new file mode 100644 index 000000000000..1c412057ea36 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_swa_kernel_arch35.h @@ -0,0 +1,809 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_swa_kernel_arch35.h + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_SWA_KERNEL_ARCH35_H +#define SPARSE_FLASH_MLA_SWA_KERNEL_ARCH35_H +#include "sparse_flash_mla_common_arch35.h" +#include "sparse_flash_mla_kvcache.h" +#include "sparse_flash_mla_csa_block_cube_arch35.h" +#include "sparse_flash_mla_csa_block_vector_arch35.h" +#include "kernel_operator.h" +#include "../sparse_flash_mla_kernel_metadata.h" + +#if __has_include("../../common/op_kernel/matmul.h") +#include "../../common/op_kernel/matmul.h" +#else +#include "../common/matmul.h" +#endif +#if __has_include("common/buffers_policy_3buff_sfa.h") +#include "common/buffers_policy_3buff_sfa.h" +#endif +#if __has_include("../../common/op_kernel/FixpipeOut.h") +#include "../../common/op_kernel/FixpipeOut.h" +#else +#include "../common/FixpipeOut.h" +#endif +#if __has_include("../../common/op_kernel/CopyInL1.h") +#include "../../common/op_kernel/CopyInL1.h" +#else +#include "../common/CopyInL1.h" +#endif + +#include "kernel_operator_list_tensor_intf.h" + +using matmul::MatmulType; +using namespace AscendC; +using namespace optiling; +using namespace optiling::detail; +using namespace AscendC::Impl::Detail; +using namespace regbaseutil; +using AttentionCommon::FdRunInfo; + +namespace SMLAKernel { +template +class SparseFlashMlaSwaKernel { +public: + ARGS_TRAITS; + __aicore__ inline SparseFlashMlaSwaKernel(){}; + + __aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, + __gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *seqUsedOriKV, + __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, + __gm__ uint8_t *metadata, __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, + __gm__ uint8_t *workspace, const SparseFlashMlaTilingData *__restrict tiling); + __aicore__ inline void Process(); + +private: + __aicore__ inline void ProcessMainLoop(); + __aicore__ inline int64_t GetSeqLen(int32_t bIdx, bool hasActualSeq, bool hasCuSeqlens, + GlobalTensor &actualSeqGm, GlobalTensor &cuSeqlensGm, + int64_t defaultSize); + __aicore__ inline void ParseTilingData(__gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *sequsedQ, + __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, + __gm__ uint8_t *cmpResidualKV); + __aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, + __gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, + __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *sequsedQ, + __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, + __gm__ uint8_t *cmpResidualKV, __gm__ uint8_t *sinks, + __gm__ uint8_t *workspace, + const SparseFlashMlaTilingData *__restrict tiling); + __aicore__ inline void InitLocalBuffer(); + __aicore__ inline void FreeEvent(); + __aicore__ inline void InitMMResBuf(__gm__ uint8_t *workspace); + __aicore__ inline void ComputeConstexpr(); + __aicore__ inline void SetRunInfo(RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount, + int64_t s2LoopLimit, int64_t multiCoreInnerIdx); + __aicore__ inline void ComputeBmm1Tail(RunInfo &runInfo, RunParamStr &runParam); + __aicore__ inline void ComputeAxisIdxByBnAndGs1(int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam); + __aicore__ inline void InitUniqueRunInfo(const RunParamStr &runParam, RunInfo &runInfo); + __aicore__ inline void ParseFdRunInfo(FdRunInfo &fdRunInfo); + __aicore__ inline int64_t ConvertS2MetadataBlockToToken(const RunParamStr &runParam, const ConstInfo &constInfo, + uint32_t s2BlockIdx); + __aicore__ inline bool ApplyS2MetadataRange(RunParamStr &runParam, ConstInfo &constInfo, int64_t s2StartPoint, + int64_t s2EndPoint, bool isFirstS2RangeTask, bool isLastS2RangeTask); + const SparseFlashMlaTilingData *__restrict tilingData; + /* 编译期常量的基本块信息 */ + static constexpr uint32_t PRELOAD_NUM = 2; + + StaticBuffer bmm1Buffers[2]; + StaticBuffer bmm2Buffers; + uint32_t bmm1GetFlag = 0; + uint32_t vUbBase = 0; + + // mm2左矩阵P + StaticBuffer l1PBuffers[2]; + uint32_t l1PGetFlag = 0; + uint32_t l1CubeBase = 0; + /* GM信息 */ + GlobalTensor metadataGm; + GlobalTensor cuSeqlensQGm; + GlobalTensor cuSeqlensOriKvGm; + GlobalTensor cuSeqlensCmpKvGm; + GlobalTensor actualSeqOriKvlenGm; + GlobalTensor actualSeqCmpKvlenGm; + GlobalTensor cmpResidualKvGm; + GlobalTensor actualSeqQlenGm; + + bool hasCuSeqlensQ = false; + bool hasCuSeqlensOriKv = false; + bool hasCuSeqlensCmpKv = false; + bool hasActualSeqQlen = false; + bool hasActualSeqOriKvlen = false; + bool hasActualSeqCmpKvlen = false; + BufferManager fdStagingBufferManager; + BuffersPolicySingleBuffer fdStagingBuffer; + BuffersPolicySingleBuffer intraCoreCombineBuffer; + BuffersPolicySingleBuffer crossCoreCombineBuffer; + /* 核Index信息 */ + int32_t aicIdx; + + /* 初始化后不变的信息 */ + ConstInfo constInfo; + + /* 模板库Block */ + CubeBlockType cubeBlock; + VecBlockType vecBlock; +}; + +template +__aicore__ inline void SparseFlashMlaSwaKernel::Init( + __gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *sequsedQ, __gm__ uint8_t *seqUsedOriKv, __gm__ uint8_t *seqUsedCmpKv, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata, + __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, __gm__ uint8_t *workspace, + const SparseFlashMlaTilingData *__restrict tiling) +{ + fa_base_matmul::ResetIdCounter(); + constInfo.subBlockIdx = GetSubBlockIdx(); + if ASCEND_IS_AIC { + this->aicIdx = GetBlockIdx(); + constInfo.aivIdx = 0; + this->tilingData = tiling; + } else { + constInfo.aivIdx = GetBlockIdx(); + this->aicIdx = constInfo.aivIdx >> 1; + this->tilingData = tiling; + } + + if (metadata == nullptr) { + return; + } + this->metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata); + + constInfo.s1BaseSize = 64; + constInfo.s2BaseSize = 128; + constInfo.hasOriTopkLength = (oriTopkLength != nullptr); + constInfo.hasCmpTopkLength = (cmpTopkLength != nullptr); + + this->ParseTilingData(cuSeqlensQ, sequsedQ, cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedOriKv, seqUsedCmpKv, + cmpResidualKV); + vecBlock.InitVecBlock(cuSeqlensQ, cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedOriKv, seqUsedCmpKv, cmpResidualKV); + vecBlock.CleanOutput(attentionOut, softmaxLse, constInfo); + InitMMResBuf(workspace); + if constexpr (IS_BATCH_CONSISTENCY) { + vecBlock.InitS2SplitStaging(intraCoreCombineBuffer.Get(), crossCoreCombineBuffer.Get()); + } else { + vecBlock.InitS2SplitStaging(fdStagingBuffer.Get()); + } + this->ComputeConstexpr(); + this->InitGlobalBuffer(query, oriKV, cmpKV, oriSparseIndices, cmpSparseIndices, oriBlockTable, cmpBlockTable, + cuSeqlensQ, cuSeqlensOriKv, cuSeqlensCmpKv, sequsedQ, seqUsedOriKv, seqUsedCmpKv, + cmpResidualKV, sinks, workspace, tiling); // gm设置 + this->InitLocalBuffer(); +} + +template +__aicore__ inline int64_t SparseFlashMlaSwaKernel::GetSeqLen( + int32_t bIdx, bool hasActualSeq, bool hasCuSeqlens, GlobalTensor &actualSeqGm, + GlobalTensor &cuSeqlensGm, int64_t defaultSize) +{ + if (hasActualSeq) { + return actualSeqGm.GetValue(bIdx); + } else if (hasCuSeqlens) { + return cuSeqlensGm.GetValue(bIdx + 1) - cuSeqlensGm.GetValue(bIdx); + } else { + return defaultSize; + } +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::ParseTilingData( + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *cuSeqlensOriKv, + __gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, + __gm__ uint8_t *cmpResidualKV) +{ + auto &sparseFlashMLABaseParams = this->tilingData->baseParams; + auto &sparseFlashMLACmpParams = this->tilingData->cmpParams; + constInfo.bSize = sparseFlashMLABaseParams.batchSize; + constInfo.n2Size = 1; + constInfo.gSize = sparseFlashMLABaseParams.nNumOfQInOneGroup; + constInfo.s1Size = sparseFlashMLABaseParams.qSeqSize; + constInfo.s2Size = sparseFlashMLABaseParams.kvSeqSize; + constInfo.cmpS2Size = sparseFlashMLACmpParams.cmpKvSeqSize; + constInfo.oriSparseBlockCount = sparseFlashMLABaseParams.oriSparseBlockCount; + constInfo.cmpSparseBlockCount = sparseFlashMLACmpParams.cmpSparseBlockCount; + constInfo.cmpRatio = sparseFlashMLACmpParams.cmpRatio; + constInfo.oriMaskMode = sparseFlashMLABaseParams.oriMaskMode; + constInfo.cmpMaskMode = sparseFlashMLACmpParams.cmpMaskMode; + constInfo.oriWinLeft = sparseFlashMLABaseParams.oriWinLeft; + constInfo.oriWinRight = sparseFlashMLABaseParams.oriWinRight; + constInfo.layoutType = sparseFlashMLABaseParams.outputLayout; + constInfo.returnSoftmaxLse = sparseFlashMLABaseParams.returnSoftmaxLse; + constInfo.tileSize = 0; + constInfo.dSizeRope = 64; + constInfo.softmaxScale = sparseFlashMLABaseParams.softmaxScale; + constInfo.dSize = 512; + constInfo.dSizeV = constInfo.dSize; + constInfo.dSizeVInput = constInfo.dSize; + constInfo.dSizeNope = constInfo.dSize - constInfo.dSizeRope; + constInfo.sparseBlockSize = 1; + constInfo.actualSeqLenSize = constInfo.bSize + 1; + constInfo.actualSeqLenKVSize = constInfo.bSize; + constInfo.oriKeyStride0 = sparseFlashMLABaseParams.oriKeyStride0; + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE) { + constInfo.cmpKeyStride0 = sparseFlashMLACmpParams.cmpKeyStride0; + } + constInfo.actualLenDimsOriKV = sparseFlashMLABaseParams.actualLenDimsOriKV; + if constexpr (TEMPLATE_MODE != SMLATemplateMode::SWA_TEMPLATE_MODE) { + constInfo.actualLenDimsCmpKV = sparseFlashMLABaseParams.actualLenDimsCmpKV; + constInfo.cmpResidualKVSize = sparseFlashMLABaseParams.cmpResidualKVSize; + } + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { + this->constInfo.isActualLenDimsOriKVNull = 0U; + } else { + this->constInfo.isActualLenDimsOriKVNull = (seqUsedOriKV == nullptr); + } + + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + constInfo.oriBlockSize = sparseFlashMLABaseParams.oriBlockSize; + constInfo.cmpBlockSize = sparseFlashMLABaseParams.cmpBlockSize; + constInfo.oriMaxBlockNumPerBatch = sparseFlashMLABaseParams.oriMaxBlockNumPerBatch; + constInfo.cmpMaxBlockNumPerBatch = sparseFlashMLACmpParams.cmpMaxBlockNumPerBatch; + } + + if (cuSeqlensQ != nullptr) { + cuSeqlensQGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensQ); + hasCuSeqlensQ = true; + } + if (cuSeqlensOriKv != nullptr) { + cuSeqlensOriKvGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensOriKv); + hasCuSeqlensOriKv = true; + } + if (cuSeqlensCmpKv != nullptr) { + cuSeqlensCmpKvGm.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensCmpKv); + hasCuSeqlensCmpKv = true; + } + if (seqUsedOriKV != nullptr) { + actualSeqOriKvlenGm.SetGlobalBuffer((__gm__ int32_t *)seqUsedOriKV); + hasActualSeqOriKvlen = true; + } + if (seqUsedCmpKV != nullptr) { + actualSeqCmpKvlenGm.SetGlobalBuffer((__gm__ int32_t *)seqUsedCmpKV); + hasActualSeqCmpKvlen = true; + } + if (cmpResidualKV != nullptr) { + cmpResidualKvGm.SetGlobalBuffer((__gm__ int32_t *)cmpResidualKV); + } + if (sequsedQ != nullptr) { + actualSeqQlenGm.SetGlobalBuffer((__gm__ int32_t *)sequsedQ); + hasActualSeqQlen = true; + } + + constInfo.needInit = 0; + if (constInfo.oriMaskMode != 0) { + for (uint32_t bIdx = 0; bIdx < constInfo.bSize; bIdx++) { + int64_t s2Size; + if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { + s2Size = actualSeqOriKvlenGm.GetValue(bIdx); + } else { + s2Size = GetSeqLen(bIdx, hasActualSeqOriKvlen, hasCuSeqlensOriKv, actualSeqOriKvlenGm, cuSeqlensOriKvGm, + constInfo.s2Size); + } + int64_t s1Size = + GetSeqLen(bIdx, hasActualSeqQlen, hasCuSeqlensQ, actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); + int64_t expectQs; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + expectQs = GetSeqLen(bIdx, false, hasCuSeqlensQ, actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); + } else { + expectQs = constInfo.s1Size; + } + if (s1Size > s2Size || s1Size < expectQs) { + constInfo.needInit = 1; + break; + } + } + } else { + constInfo.needInit = 1; + } +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::InitGlobalBuffer( + __gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *sequsedQ, __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *sinks, __gm__ uint8_t *workspace, const SparseFlashMlaTilingData *__restrict tiling) +{ + vecBlock.InitGlobalBuffer(oriKV, cmpKV, oriSparseIndices, cmpSparseIndices, oriBlockTable, cmpBlockTable, sequsedQ, + sinks, seqUsedOriKV, seqUsedCmpKV, cmpResidualKV); + cubeBlock.InitGlobalBuffer(query, oriKV, cmpKV, cmpSparseIndices, oriBlockTable, cmpBlockTable, sequsedQ, + cuSeqlensQ, cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedOriKV, seqUsedCmpKV, constInfo); +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::InitMMResBuf(__gm__ uint8_t *workspace) +{ + // L1: [l1P x2][cube L1], l1P 必须放在最前面保证与 vec 申请地址相同 + uint32_t mm2LeftSize = constInfo.s1BaseSize * constInfo.s2BaseSize; + uint32_t l1PAddr = 0; + l1PBuffers[0] = {LocalTensor(TPosition::A1, l1PAddr, mm2LeftSize), 0}; + l1PAddr += (mm2LeftSize * sizeof(Q_T)); + l1PBuffers[1] = {LocalTensor(TPosition::A1, l1PAddr, mm2LeftSize), 1}; + l1PAddr += (mm2LeftSize * sizeof(Q_T)); + l1CubeBase = l1PAddr; + + // UB: [bmm2][bmm1 x2][vec UB] + uint32_t mm1ResultSize = constInfo.s1BaseSize / CV_RATIO * constInfo.s2BaseSize; + uint32_t mm2ResultSize = constInfo.s1BaseSize / CV_RATIO * 512; + uint32_t ubAddr = 0; + bmm2Buffers = {LocalTensor(TPosition::VECIN, ubAddr, mm2ResultSize), 0}; + ubAddr += (mm2ResultSize * sizeof(T)); + bmm1Buffers[0] = {LocalTensor(TPosition::VECIN, ubAddr, mm1ResultSize), 0}; + ubAddr += (mm1ResultSize * sizeof(T)); + bmm1Buffers[1] = {LocalTensor(TPosition::VECIN, ubAddr, mm1ResultSize), 1}; + ubAddr += (mm1ResultSize * sizeof(T)); + vUbBase = ubAddr; + + if ASCEND_IS_AIV { + CrossCoreSetFlag(CROSSCORE_BMM1(bmm1Buffers[0].idx)); + CrossCoreSetFlag(CROSSCORE_BMM1(bmm1Buffers[1].idx)); + CrossCoreSetFlag(CROSSCORE_BMM2); + } + + int64_t fdStagingOffset = 0U; + if constexpr (IS_SPLIT_G) { + constexpr uint32_t TRIPLE_BUFFER_NUM = 3U; + int64_t v0ResSize = constInfo.s2BaseSize * constInfo.dSize * sizeof(Q_T); + int64_t v0LogicalSlotCount = GetBlockNum() >> 1U; + fdStagingOffset = v0ResSize * TRIPLE_BUFFER_NUM * v0LogicalSlotCount; + fdStagingOffset += TRIPLE_BUFFER_NUM * constInfo.s2BaseSize * sizeof(int32_t) * GetBlockNum(); + } + fdStagingBufferManager.Init(workspace + fdStagingOffset); + constexpr uint32_t FD_MAX_SUM_REGION_NUM = 2U; + uint32_t gSize = static_cast(constInfo.gSize); + uint32_t combineElemSize = + gSize * constInfo.dSize + + FD_MAX_SUM_REGION_NUM * gSize * static_cast(AttentionCommon::FD_BROADCAST_ELEMS_PER_ROW); + if constexpr (IS_BATCH_CONSISTENCY) { + uint32_t intraCoreSlotNum = IS_SPLIT_G ? GetBlockNum() : (GetBlockNum() << 1U); + uint32_t intraCoreCombineSize = intraCoreSlotNum * combineElemSize * sizeof(float); + uint32_t crossCoreCombineSize = + GetBlockNum() * BATCH_CONSISTENCY_MAX_REDUCE_BLOCK_NUM * combineElemSize * sizeof(float); + intraCoreCombineBuffer.Init(fdStagingBufferManager, intraCoreCombineSize); + crossCoreCombineBuffer.Init(fdStagingBufferManager, crossCoreCombineSize); + } else { + uint32_t fdSlotCount = static_cast(AttentionCommon::FD_MAX_S2_SPLIT_NUM) * + (IS_SPLIT_G ? (GetBlockNum() >> 1U) : GetBlockNum()); + fdStagingBuffer.Init(fdStagingBufferManager, fdSlotCount * combineElemSize * sizeof(float)); + } +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::InitLocalBuffer() +{ + vecBlock.InitLocalBuffer(constInfo, vUbBase); + cubeBlock.InitLocalBuffer(l1CubeBase); +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::ComputeConstexpr() +{ + // 计算轴的乘积 + constInfo.s1S2 = constInfo.s1Size * constInfo.s2Size; + constInfo.gS1 = constInfo.gSize * constInfo.s1Size; + constInfo.n2G = constInfo.n2Size * constInfo.gSize; + + constInfo.s1Dv = constInfo.s1Size * constInfo.dSizeV; + constInfo.s2Dv = constInfo.s2Size * constInfo.dSizeV; + constInfo.n2Dv = constInfo.n2Size * constInfo.dSizeV; + constInfo.gDv = constInfo.gSize * constInfo.dSizeV; + constInfo.gS1Dv = constInfo.gSize * constInfo.s1Dv; + constInfo.n2S2Dv = constInfo.n2Size * constInfo.s2Dv; + constInfo.n2GDv = constInfo.n2Size * constInfo.gDv; + constInfo.s2BaseN2Dv = constInfo.s2BaseSize * constInfo.n2Dv; + constInfo.n2GS1Dv = constInfo.n2Size * constInfo.gS1Dv; + + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + // (BS)ND + constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv; + constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize; + constInfo.mm1Kb = constInfo.n2Size * constInfo.dSize; + if ASCEND_IS_AIV { + constInfo.attentionOutStride = (constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T); + } + } else if constexpr (LAYOUT_T == SMLA_LAYOUT::BSND) { + // BSH/BSNGD + constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv; + constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize; + constInfo.mm1Kb = constInfo.n2Size * constInfo.dSize; + if ASCEND_IS_AIV { + constInfo.attentionOutStride = (constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T); + } + } +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::Process() +{ + // SyncAll Cube和Vector都需要调用 + if (this->constInfo.needInit) { + SyncAll(); + } + FdRunInfo fdRunInfo; + if ASCEND_IS_AIV { + ParseFdRunInfo(fdRunInfo); + } + ProcessMainLoop(); + if ASCEND_IS_AIV { + SyncAll(); + if (fdRunInfo.coreEnable) { + this->vecBlock.ProcessFlashDecode(fdRunInfo, this->constInfo); + } + } + FreeEvent(); +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::ProcessMainLoop() +{ + uint32_t hasLoad = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_CORE_ENABLE_INDEX, false)); + if (hasLoad == 0) { + return; + } + + // 从meta data解析分核信息 + uint32_t bN2StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_BN2_START_INDEX, false)); + uint32_t gS1StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_M_START_INDEX, false)); + uint32_t s2StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_S2_START_INDEX, false)); + uint32_t bN2EndIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_BN2_END_INDEX, false)); + uint32_t nextGs1Idx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_M_END_INDEX, false)); + uint32_t s2EndIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_S2_END_INDEX, false)); + uint32_t firstFdDataWorkspaceIdx = + metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX, false)); + uint32_t s2LoopLimit = 0; + if (nextGs1Idx != 0 || s2EndIdx != 0) { + bN2EndIdx++; + } + + int64_t taskId = 0; + bool notLast = true; + bool isFirstLoop = true; + RunInfo runInfo[3]; + RunParamStr runParam; + runParam.firstFdDataWorkspaceIdx = firstFdDataWorkspaceIdx; + int64_t multiCoreInnerIdx = 1; + int64_t s2SplitIdxCounter = 0; + for (int64_t bnIdx = bN2StartIdx; bnIdx < bN2EndIdx; bnIdx++) { + bool lastBN = (bnIdx == bN2EndIdx - 1); + runParam.boIdx = bnIdx; + runParam.n2oIdx = 0; + ComputeParamBatch( + runParam, this->constInfo, this->cuSeqlensQGm, this->cuSeqlensOriKvGm, this->cuSeqlensCmpKvGm, + this->actualSeqQlenGm, this->actualSeqOriKvlenGm, this->actualSeqCmpKvlenGm, this->cmpResidualKvGm, + this->hasActualSeqQlen, this->hasActualSeqOriKvlen, this->hasActualSeqCmpKvlen, this->hasCuSeqlensCmpKv); + ComputeS1LoopInfo(runParam, this->constInfo, lastBN, nextGs1Idx, gS1StartIdx, s2EndIdx); + + int64_t gS1LoopEnd = lastBN ? (runParam.gs1LoopEndIdx + PRELOAD_NUM) : runParam.gs1LoopEndIdx; + for (int64_t gS1Index = runParam.gs1LoopStartIdx; gS1Index < gS1LoopEnd; gS1Index++) { + bool notLastTwoLoop = true; + if (lastBN) { + int32_t extraGS1 = gS1Index - runParam.gs1LoopEndIdx; + switch (extraGS1) { + case 0: + notLastTwoLoop = false; + break; + case 1: + notLast = false; + notLastTwoLoop = false; + break; + default: + break; + } + } + if (notLastTwoLoop) { + this->ComputeAxisIdxByBnAndGs1(bnIdx, gS1Index, runParam); + bool s1NoNeedCalc = + ComputeParamS1(runParam, this->constInfo, gS1Index, this->cuSeqlensQGm); + GlobalTensor tmpTensor; + bool s2NoNeedCalc = ComputeS2LoopInfo( + bnIdx, gS1Index, this->cuSeqlensQGm, tmpTensor, tmpTensor, runParam, this->constInfo); + if constexpr (IS_BATCH_CONSISTENCY) { + int64_t oriLoad = runParam.s2OriLineEndIdx - runParam.s2OriLineStartIdx; + int64_t cmpLoad = runParam.s2CmpLineEndIdx - runParam.s2CmpLineStartIdx; + int64_t totalLoad = oriLoad + cmpLoad; + int64_t s2BaseSize = static_cast(constInfo.s2BaseSize); + int64_t rawReductionBlockSize = totalLoad / 32LL; + int64_t reductionBlockSize = (rawReductionBlockSize + s2BaseSize - 1LL) / s2BaseSize * s2BaseSize; + runParam.baseBlockNumPerReductionBlock = + reductionBlockSize > 0 ? reductionBlockSize / s2BaseSize : 1LL; + } + if (!s2NoNeedCalc) { + bool isFirstS2RangeTask = (bnIdx == bN2StartIdx && gS1Index == runParam.gs1LoopStartIdx); + bool isLastS2RangeTask = (lastBN && gS1Index == runParam.gs1LoopEndIdx - 1); + int64_t s2StartPoint = ConvertS2MetadataBlockToToken(runParam, this->constInfo, s2StartIdx); + int64_t s2EndPoint = (isLastS2RangeTask && s2EndIdx == 0) ? + 0 : + ConvertS2MetadataBlockToToken(runParam, this->constInfo, s2EndIdx); + s2NoNeedCalc = ApplyS2MetadataRange(runParam, this->constInfo, s2StartPoint, s2EndPoint, + isFirstS2RangeTask, isLastS2RangeTask); + } else { + runParam.isCrossCoreSplit = false; + } + // s1和s2有任意一个不需要算, 则continue, 如果是当前核最后一次循环,则补充计算taskIdx+2的部分 + if (s1NoNeedCalc || s2NoNeedCalc) { + continue; + } + if constexpr (!IS_BATCH_CONSISTENCY) { + if (runParam.isCrossCoreSplit) { + runParam.s2SplitIdx = s2SplitIdxCounter++; + } + } + s2LoopLimit = runParam.s2LoopEndIdx - 1; + } else { + s2LoopLimit = 0; + } + for (int64_t s2LoopCount = 0; s2LoopCount <= s2LoopLimit; ++s2LoopCount) { + if constexpr (IS_BATCH_CONSISTENCY) { + int64_t safeBaseBlockNum = + runParam.baseBlockNumPerReductionBlock > 0 ? runParam.baseBlockNumPerReductionBlock : 1LL; + if (runParam.isCrossCoreSplit && s2LoopCount % safeBaseBlockNum == 0) { + runParam.s2SplitIdx = s2SplitIdxCounter++; + } + } + if (notLastTwoLoop) { + RunInfo &runInfo1 = runInfo[taskId % 3]; + this->SetRunInfo(runInfo1, runParam, taskId, s2LoopCount, s2LoopLimit, multiCoreInnerIdx); + if ASCEND_IS_AIC { + this->cubeBlock.IterateLoadQK(runInfo1, this->constInfo, isFirstLoop); + isFirstLoop = false; + } + } + if (taskId > 0 && notLast) { + auto &runInfo2 = runInfo[(taskId + 2) % 3]; + if ASCEND_IS_AIV { + uint32_t bmm1Slot = bmm1GetFlag; + bmm1GetFlag ^= 1; + uint32_t l1PSlot = l1PGetFlag; + l1PGetFlag ^= 1; + this->vecBlock.ProcessVec1(this->l1PBuffers[l1PSlot], this->bmm1Buffers[bmm1Slot], runInfo2, + this->constInfo); + } else { + uint32_t bmm1Slot = bmm1GetFlag; + bmm1GetFlag ^= 1; + RunInfo &runInfoNext = runInfo[taskId % 3]; + this->cubeBlock.IterateBmm1(this->bmm1Buffers[bmm1Slot], notLastTwoLoop, runInfoNext, runInfo2, + this->constInfo); + } + } + if (taskId > 1) { + RunInfo &runInfo3 = runInfo[(taskId + 1) % 3]; + if ASCEND_IS_AIV { + this->vecBlock.ProcessVec2(this->bmm2Buffers, runInfo3, this->constInfo); + } else { + uint32_t l1PSlot = l1PGetFlag; + l1PGetFlag ^= 1; + this->cubeBlock.IterateBmm2(this->bmm2Buffers, this->l1PBuffers[l1PSlot], runInfo3, + this->constInfo); + } + } + ++taskId; + } + ++multiCoreInnerIdx; + } + gS1StartIdx = 0; + } +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::ComputeAxisIdxByBnAndGs1( + int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam) +{ + // GS1合轴, 不切G, 只切S1 + runParam.s1oIdx = gS1Index * runParam.qSNumInOneBlock; + if constexpr (IS_SPLIT_G) { + int64_t halfG = (constInfo.gSize + 1) / 2; // ceil(gSize/2), 第一个AIC多处理一行 + runParam.goIdx = (aicIdx % 2 == 0) ? 0 : halfG; + runParam.gSplitSize = (aicIdx % 2 == 0) ? halfG : (constInfo.gSize - halfG); + } else { + runParam.goIdx = 0; + runParam.gSplitSize = constInfo.gSize; + } +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::SetRunInfo( + RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount, int64_t s2LoopLimit, + int64_t multiCoreInnerIdx) +{ + if (s2LoopCount < runParam.oriKvLoopEndIdx) { + runInfo.s2StartIdx = runParam.s2OriLineStartIdx; + runInfo.s2EndIdx = runParam.s2OriLineEndIdx; + } else { + runInfo.s2StartIdx = runParam.s2CmpLineStartIdx; + runInfo.s2EndIdx = runParam.s2CmpLineEndIdx; + } + runInfo.s2LoopCount = s2LoopCount; + if (runInfo.multiCoreInnerIdx != multiCoreInnerIdx) { + runInfo.s1oIdx = runParam.s1oIdx; + runInfo.boIdx = runParam.boIdx; + runInfo.n2oIdx = runParam.n2oIdx; + runInfo.goIdx = runParam.goIdx; + runInfo.multiCoreInnerIdx = multiCoreInnerIdx; + runInfo.multiCoreIdxMod2 = multiCoreInnerIdx & 1; + runInfo.multiCoreIdxMod3 = multiCoreInnerIdx % 3; + } + + runInfo.taskId = taskId; + runInfo.taskIdMod2 = taskId & 1; + runInfo.taskIdMod3 = taskId % 3; + runInfo.s2LoopLimit = s2LoopLimit; + + runInfo.actualS1Size = runParam.actualS1Size; + runInfo.attentionOutOffset = runParam.attentionOutOffset; + runInfo.sOuterOffset = runParam.sOuterOffset; + runInfo.firstFdDataWorkspaceIdx = runParam.firstFdDataWorkspaceIdx; + runInfo.isCrossCoreSplit = runParam.isCrossCoreSplit; + runInfo.s2SplitIdx = runParam.s2SplitIdx; + runInfo.isFirstS2SplitCore = runParam.isFirstS2SplitCore; + int64_t safeBaseBlockNum = + runParam.baseBlockNumPerReductionBlock > 0 ? runParam.baseBlockNumPerReductionBlock : 1LL; + int64_t baseBlockIdInReduceBlock = s2LoopCount % safeBaseBlockNum; + runInfo.reduceBlockId = s2LoopCount / safeBaseBlockNum; + runInfo.isFirstBase = baseBlockIdInReduceBlock == 0; + runInfo.isLastBase = baseBlockIdInReduceBlock == safeBaseBlockNum - 1LL || s2LoopCount == s2LoopLimit; + runInfo.needReduce = runInfo.reduceBlockId > 0; + this->ComputeBmm1Tail(runInfo, runParam); + InitUniqueRunInfo(runParam, runInfo); +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::InitUniqueRunInfo( + const RunParamStr &runParam, RunInfo &runInfo) +{ + InitTaskParamByRun(runParam, runInfo, constInfo); +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::ComputeBmm1Tail(RunInfo &runInfo, + RunParamStr &runParam) +{ + // ------------------------S1 Base Related--------------------------- + runInfo.s1RealSize = runParam.s1RealSize; + runInfo.halfS1RealSize = runParam.halfS1RealSize; + runInfo.firstHalfS1RealSize = runParam.firstHalfS1RealSize; + runInfo.mRealSize = runParam.mRealSize; + runInfo.halfMRealSize = runParam.halfMRealSize; + runInfo.firstHalfMRealSize = runParam.firstHalfMRealSize; + + runInfo.vec2MBaseSize = runInfo.halfMRealSize; + + // ------------------------S2 Base Related---------------------------- + runInfo.s2RealSize = constInfo.s2BaseSize; + runInfo.s2AlignedSize = runInfo.s2RealSize; + int64_t curS2LoopCnt = (runInfo.s2LoopCount >= runParam.oriKvLoopEndIdx) ? + (runInfo.s2LoopCount - runParam.oriKvLoopEndIdx) : + runInfo.s2LoopCount; + if (runInfo.s2StartIdx + (curS2LoopCnt + 1) * runInfo.s2RealSize > runInfo.s2EndIdx) { + runInfo.s2RealSize = runInfo.s2EndIdx - curS2LoopCnt * runInfo.s2RealSize - runInfo.s2StartIdx; + runInfo.s2AlignedSize = Align(runInfo.s2RealSize); + } +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::ParseFdRunInfo(FdRunInfo &fdRunInfo) +{ + uint32_t aivIdx = static_cast(this->constInfo.aivIdx); + fdRunInfo.coreEnable = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_CORE_ENABLE_INDEX, true)) != 0; + if (!fdRunInfo.coreEnable) { + return; + } + fdRunInfo.bn2Idx = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_BN2_IDX_INDEX, true)); + fdRunInfo.mIdx = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_M_IDX_INDEX, true)); + fdRunInfo.workspaceIdx = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_WORKSPACE_IDX_INDEX, true)); + fdRunInfo.workspaceNum = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_WORKSPACE_NUM_INDEX, true)); + fdRunInfo.mStartIdx = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_M_START_INDEX, true)); + fdRunInfo.mNum = metadataGm.GetValue(GetAttrAbsIndex(aivIdx, FD_M_NUM_INDEX, true)); +} + +template +__aicore__ inline int64_t SparseFlashMlaSwaKernel::ConvertS2MetadataBlockToToken( + const RunParamStr &runParam, const ConstInfo &constInfo, uint32_t s2BlockIdx) +{ + int64_t s2BaseSize = static_cast(constInfo.s2BaseSize); + int64_t oriLen = runParam.s2OriLineEndIdx - runParam.s2OriLineStartIdx; + int64_t cmpLen = runParam.s2CmpLineEndIdx - runParam.s2CmpLineStartIdx; + int64_t reductionBlockSize = runParam.baseBlockNumPerReductionBlock * s2BaseSize; + int64_t oriReductionBlockNum = (oriLen + reductionBlockSize - 1) / reductionBlockSize; + int64_t reductionBlockIdx = static_cast(s2BlockIdx); + if (reductionBlockIdx < oriReductionBlockNum) { + int64_t oriToken = reductionBlockIdx * reductionBlockSize; + return oriToken < oriLen ? oriToken : oriLen; + } + int64_t cmpToken = (reductionBlockIdx - oriReductionBlockNum) * reductionBlockSize; + return oriLen + (cmpToken < cmpLen ? cmpToken : cmpLen); +} + +template +__aicore__ inline void SparseFlashMlaSwaKernel::FreeEvent() +{ + if ASCEND_IS_AIC { + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1Buffers[0].idx)); + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1Buffers[0].idx) + AIV0_AIV1_OFFSET); + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1Buffers[1].idx)); + CrossCoreWaitFlag(CROSSCORE_BMM1(bmm1Buffers[1].idx) + AIV0_AIV1_OFFSET); + CrossCoreWaitFlag(CROSSCORE_BMM2); + CrossCoreWaitFlag(CROSSCORE_BMM2 + AIV0_AIV1_OFFSET); + this->cubeBlock.FreeEvent(); + } else { + this->vecBlock.FreeEvent(constInfo); + } +} + +template +__aicore__ inline bool SparseFlashMlaSwaKernel::ApplyS2MetadataRange( + RunParamStr &runParam, ConstInfo &constInfo, int64_t s2StartPoint, int64_t s2EndPoint, bool isFirstS2RangeTask, + bool isLastS2RangeTask) +{ + int64_t oriStart = runParam.s2OriLineStartIdx; + int64_t oriEnd = runParam.s2OriLineEndIdx; + int64_t oriLen = oriEnd - oriStart; + int64_t cmpStart = runParam.s2CmpLineStartIdx; + int64_t cmpEnd = runParam.s2CmpLineEndIdx; + int64_t cmpLen = cmpEnd - cmpStart; + int64_t totalLen = oriLen + cmpLen; + + int64_t effectiveS2EndPoint = (isLastS2RangeTask && s2EndPoint == 0) ? totalLen : s2EndPoint; + int64_t rangeStart = isFirstS2RangeTask ? s2StartPoint : 0; + rangeStart = rangeStart < 0 ? 0 : rangeStart; + rangeStart = rangeStart < totalLen ? rangeStart : totalLen; + int64_t rangeEnd = isLastS2RangeTask ? effectiveS2EndPoint : totalLen; + rangeEnd = rangeEnd < 0 ? 0 : rangeEnd; + rangeEnd = rangeEnd < totalLen ? rangeEnd : totalLen; + if (rangeEnd <= rangeStart) { + runParam.oriKvLoopEndIdx = 0; + runParam.cmpKvLoopEndIdx = 0; + runParam.s2LoopEndIdx = 0; + runParam.isCrossCoreSplit = false; + return true; + } + + bool hasPrevCore = rangeStart > 0; + bool hasNextCore = rangeEnd < totalLen; + runParam.isCrossCoreSplit = hasPrevCore || hasNextCore; + runParam.isFirstS2SplitCore = !hasPrevCore; + + int64_t oriRangeStart = rangeStart < oriLen ? rangeStart : oriLen; + int64_t oriRangeEnd = rangeEnd < oriLen ? rangeEnd : oriLen; + runParam.s2OriLineStartIdx = oriStart + oriRangeStart; + runParam.s2OriLineEndIdx = oriStart + oriRangeEnd; + + int64_t cmpRangeStart = rangeStart > oriLen ? rangeStart - oriLen : 0; + cmpRangeStart = cmpRangeStart < cmpLen ? cmpRangeStart : cmpLen; + int64_t cmpRangeEnd = rangeEnd > oriLen ? rangeEnd - oriLen : 0; + cmpRangeEnd = cmpRangeEnd < cmpLen ? cmpRangeEnd : cmpLen; + runParam.s2CmpLineStartIdx = cmpStart + cmpRangeStart; + runParam.s2CmpLineEndIdx = cmpStart + cmpRangeEnd; + + int64_t s2BaseSize = static_cast(constInfo.s2BaseSize); + int64_t oriRangeLen = runParam.s2OriLineEndIdx - runParam.s2OriLineStartIdx; + int64_t cmpRangeLen = runParam.s2CmpLineEndIdx - runParam.s2CmpLineStartIdx; + runParam.oriKvLoopEndIdx = (oriRangeLen + s2BaseSize - 1) / s2BaseSize; + runParam.cmpKvLoopEndIdx = (cmpRangeLen + s2BaseSize - 1) / s2BaseSize; + runParam.s2LoopEndIdx = runParam.oriKvLoopEndIdx + runParam.cmpKvLoopEndIdx; + return runParam.s2LoopEndIdx == 0; +} +} // namespace SMLAKernel +#endif // SPARSE_FLASH_MLA_SWA_KERNEL_ARCH35_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/arch35/util_regbase.h b/csrc/attention/sparse_flash_mla/op_kernel/arch35/util_regbase.h new file mode 100644 index 000000000000..5d20ea51d44c --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/arch35/util_regbase.h @@ -0,0 +1,275 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file util_regbase.h + * \brief + */ + +#ifndef SMLA_UTIL_REGBASE_H +#define SMLA_UTIL_REGBASE_H + +#include "util.h" + +using AscendC::QuePosition; +using AscendC::TQue; + +namespace regbaseutil { +constexpr int64_t MAX_PRE_NEXT_TOKENS = 0x7FFFFFFF; +enum class VselrIndexEnum { + GT_64_AND_LTE_128_INDEX = 0, + GT_0_AND_LTE_64_INDEX = 1 +}; + +#define COMMON_RUN_PARAM \ + int64_t boIdx; \ + int64_t s1oIdx; \ + int64_t n2oIdx; \ + int64_t goIdx; \ + int64_t gSplitSize; /* split-G模式下当前AIC处理的G轴行数 */ \ + int64_t s2LoopEndIdx; /* S2方向的循环控制信息 souter层确定 */ \ + int64_t s2LineStartIdx = 0; /* S2方向按行的起始位置 */ \ + int64_t s2LineEndIdx; /* S2方向按行的结束位置 */ \ + int64_t s2OriLoopEndIdx; /* S2方向的循环控制信息 souter层确定 */ \ + int64_t s2OriLineStartIdx = 0; /* S2方向按行的起始位置 */ \ + int64_t s2OriLineEndIdx; /* S2方向按行的结束位置 */ \ + int64_t s2CmpLoopEndIdx; \ + int64_t s2CmpLineStartIdx = 0; \ + int64_t s2CmpLineEndIdx; \ + /* cube视角的sOuter,在SAMEAB场景中cubeSOuterSize为两倍的 halfS1RealSize souter层确定 */ \ + uint32_t s1RealSize; \ + uint32_t halfS1RealSize; \ + uint32_t firstHalfS1RealSize; \ + uint32_t mRealSize; \ + uint32_t halfMRealSize; \ + uint32_t firstHalfMRealSize; \ + int64_t attentionOutOffset; /* attentionOut的offset souter层确定 */ \ + int32_t actualS1Size; /* Q的actualSeqLength */ \ + int32_t actualS2OriSize; /* ori_kv的真实使用长度 */ \ + int32_t actualS2CmpSize; /* cmp_kv的真实使用长度 */ \ + int32_t cmpResidual /* cmp的余数,用于mask计算 */ + +struct RunParamStr { // 分核与切块需要使用到参数 + COMMON_RUN_PARAM; + /* 推理新增 */ + int64_t gs1LoopStartIdx; + int64_t gs1LoopEndIdx; + // BN循环生产的数据 + + int64_t preTokensPerBatchOri = MAX_PRE_NEXT_TOKENS; // 左上顶点的pretoken + int64_t nextTokensPerBatchOri = MAX_PRE_NEXT_TOKENS; // 左上顶点的nexttoken + + int64_t preTokensPerBatchCmp = MAX_PRE_NEXT_TOKENS; // 左上顶点的pretoken + int64_t nextTokensPerBatchCmp = MAX_PRE_NEXT_TOKENS; // 左上顶点的nexttoken + + // NBS1循环生产的数据 + int64_t sOuterOffset; // 单个S内 souter的 souterIdx * halfS1RealSize souter层确定 + int64_t cubeSOuterOffset; // 单个S内 souter的 souterIdx * halfS1RealSize souter层确定 + int64_t mOuterOffset; + int64_t cubeMOuterOffset; + uint32_t oriSparseBlockCount; + uint32_t cmpSparseBlockCount; + + // lse 输出offset + int64_t softmaxLseOffset; // souter层确定 + + int64_t qSNumInOneBlock; + int64_t oriKvLoopEndIdx; + int64_t cmpKvLoopEndIdx; + // FD S2-split + int64_t firstFdDataWorkspaceIdx = 0; + bool isCrossCoreSplit = false; + int64_t s2SplitIdx = 0; + bool isFirstS2SplitCore = true; + int64_t baseBlockNumPerReductionBlock = 1; +}; + +#define COMMON_RUN_INFO \ + int64_t s2StartIdx; /* s2的起始位置,sparse场景下可能不是0 */ \ + int64_t s2EndIdx; \ + int64_t s2LoopCount; /* s2循环当前的循环index */ \ + int64_t s2LoopLimit; \ + int64_t s1oIdx = 0; /* s1轴的index */ \ + int64_t loop = 0; /* for v0 perload loop */ \ + int64_t boIdx = 0; /* b轴的index */ \ + int64_t n2oIdx = 0; /* n2轴的index */ \ + int64_t goIdx = 0; /* g轴的index */ \ + int32_t s1RealSize; \ + int32_t halfS1RealSize; /* vector侧实际的s1基本块大小,如果Cube基本块=128,那么halfS1RealSize=64 */ \ + int32_t \ + firstHalfS1RealSize; /* 当s1RealSize不是2的整数倍时,v0比v1少计算一行,计算subblock偏移的时候需要使用v0的s1 \ + size */ \ + int32_t mRealSize; \ + int32_t halfMRealSize; \ + int32_t firstHalfMRealSize; \ + int32_t s2RealSize; /* s2方向基本块的真实长度 */ \ + int32_t s2RealSizeUpdate; \ + int64_t s2AlignedSize; /* s2方向基本块对齐到16之后的长度 */ \ + int32_t vec2MBaseSize; \ + int32_t vec2MRealSize; \ + int64_t taskId; \ + int64_t multiCoreInnerIdx = 0; \ + int64_t attentionOutOffset; \ + int32_t actualS1Size; /* 非TND场景=总s1Size, Tnd场景下当前batch对应的s1 */ \ + int32_t actualS2CmpSize; /* cmp_kv的真实使用长度 */ \ + int32_t cmpResidual; /* cmp的余数,用于mask计算 */ \ + int64_t preTokensPerBatchOri; /* vector2 左上顶点的pretoken */ \ + int64_t nextTokensPerBatchOri; /* vector2 ori 左上顶点的nexttoken */ \ + int64_t nextTokensPerBatchCmp; /* vector2 cmp 左上顶点的nexttoken */ \ + uint8_t taskIdMod2; \ + uint8_t taskIdMod3; \ + uint8_t multiCoreIdxMod2 = 0; \ + uint8_t multiCoreIdxMod3 = 0; \ + bool isCmp; \ + uint8_t resv[3]; \ + int64_t sOuterOffset; \ + int64_t mOuterOffset; \ + bool isCrossCoreSplit = false; \ + int64_t s2SplitIdx = 0; \ + bool isFirstS2SplitCore = true; \ + int64_t baseBlockNumPerReductionBlock = 1 + +struct RunInfo { + COMMON_RUN_INFO; + // 推理新增 + // lse 输出offset + int64_t softmaxLseOffset; + + int64_t qSNumInOneBlock; + int64_t oriKvLoopEndIdx; + int64_t cmpKvLoopEndIdx; + uint32_t oriSparseBlockCount; + uint32_t cmpSparseBlockCount; + int64_t firstFdDataWorkspaceIdx = 0; + bool isFirstBase = true; + bool isLastBase = true; + bool needReduce = false; + int64_t reduceBlockId = 0; +}; + +#define COMMON_CONST_INFO \ + /* 全局的基本块信息 */ \ + uint32_t bSize; \ + uint32_t needInit; \ + uint32_t s1BaseSize; \ + uint32_t s2BaseSize; \ + int64_t dSize; /* query d 512 */ \ + int64_t dSizeV; /* key d 512 */ \ + int64_t dSizeVInput; /* key inpue d 640 = rope + nope + scale + pad */ \ + int64_t dSizeNope; /* key nope d 448 */ \ + int64_t dSizeRope; /* key rope d 64 */ \ + int64_t tileSize; /* 64 */ \ + int64_t sparseMode = 3; \ + int64_t gSize; /* g轴的大小 */ \ + int64_t n2Size; \ + int64_t s1Size; /* s1总大小 */ \ + int64_t s2Size; /* s2总大小 */ \ + int64_t cmpS2Size; /* s2总大小 */ \ + /* 轴的乘积 */ \ + int64_t s1D; \ + int64_t gS1D; \ + int64_t n2GS1D; \ + int64_t s2D; \ + int64_t n2S2D; \ + int64_t s1Dv; \ + int64_t gS1Dv; \ + int64_t n2GS1Dv; \ + int64_t s2Dv; \ + int64_t n2S2Dv; \ + int64_t s1S2; \ + int64_t gS1; \ + int64_t gD; \ + int64_t n2D; \ + int64_t bN2D; \ + int64_t gDv; \ + int64_t n2Dv; \ + int64_t bN2Dv; \ + int64_t n2G; \ + int64_t n2GD; \ + int64_t bN2GD; \ + int64_t n2GDv; \ + int64_t bN2GDv; \ + int64_t gS2; \ + int64_t s1Dr; \ + int64_t gS1Dr; \ + int64_t n2GS1Dr; \ + int64_t s2Dr; \ + int64_t n2S2Dr; \ + int64_t gDr; \ + int64_t n2Dr; \ + int64_t bN2Dr; \ + int64_t n2GDr; \ + int64_t bN2GDr; \ + int32_t s2BaseN2D; \ + int32_t s1BaseN2GD; \ + int64_t s2BaseBN2D; \ + int64_t s1BaseBN2GD; \ + int32_t s1BaseD; \ + int32_t s2BaseD; \ + int64_t s2BaseN2Dv; \ + int64_t s2BaseBN2Dv; \ + int64_t s1BaseN2GDv; \ + int64_t s1BaseBN2GDv; \ + int32_t s1BaseDv; \ + int32_t s2BaseDv; \ + /* matmul跳读参数 */ \ + int64_t mm1Ka; \ + int64_t mm1Kb; \ + /* dq 或者attentionOut的Stride */ \ + int64_t attentionOutStride; \ + uint32_t aivIdx; \ + uint8_t layoutType; \ + uint8_t subBlockIdx; \ + bool hasOriTopkLength; \ + bool hasCmpTopkLength; \ + /* nonContiguous */ \ + int64_t oriKeyStride0; \ + int64_t cmpKeyStride0 + +#define INFER_CONST_INFO \ + /* 推理 */ \ + bool isActualLenDimsNull; /* 判断是否有actualseq */ \ + bool isActualLenDimsKVNull; /* 判断是否有actualseq_kv */ \ + bool isActualLenDimsOriKVNull; /* 判断是否有actualseq_kv */ \ + bool isActualLenDimsCmpKVNull; /* 判断是否有actualseq_kv */ \ + bool cmpResidualKVNull; /* 判断是否有actualseq_kv */ \ + bool isSoftmaxLseEnable; \ + bool rsvd1; \ + bool returnSoftmaxLse; \ + uint32_t oriSparseBlockCount; \ + uint32_t cmpSparseBlockCount; \ + uint32_t alignedOriSparseBlockCount; \ + uint32_t alignedCmpSparseBlockCount; \ + uint32_t actualSeqLenSize; /* 用户输入的actualseq的长度 */ \ + uint32_t actualLenDimsOriKV; /* seqused_ori_kv的维度 */ \ + uint32_t actualLenDimsCmpKV; /* seqused_cmp_kv的维度 */ \ + uint32_t cmpResidualKVSize; /* cmp_residual_kv的长度 */ \ + uint32_t actualSeqLenKVSize; /* 用户输入的actualseq_kv的长度 */ \ + /* service mm1 mm2 pageAttention */ \ + uint32_t oriBlockSize; \ + uint32_t cmpBlockSize; \ + uint32_t paLayoutType; \ + uint32_t oriMaxBlockNumPerBatch; \ + uint32_t cmpMaxBlockNumPerBatch; \ + int32_t oriWinLeft; \ + int32_t oriWinRight; \ + uint32_t sparseBlockSize; \ + uint32_t cmpRatio; \ + float softmaxScale; \ + uint32_t oriMaskMode; \ + uint32_t cmpMaskMode + +struct ConstInfo { + COMMON_CONST_INFO; + INFER_CONST_INFO; +}; +} // namespace regbaseutil + +#endif // SMLA_UTIL_REGBASE_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla.cpp b/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla.cpp new file mode 100644 index 000000000000..b57d80d763af --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla.cpp @@ -0,0 +1,131 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla.cpp + * \brief + */ + +#if (__CCE_AICORE__ == 310) +#include "kernel_operator.h" +#include "lib/matmul_intf.h" +#include "sparse_flash_mla_template_tiling_key.h" +#include "arch35/sparse_flash_mla_csa_kernel_arch35.h" +#include "arch35/sparse_flash_mla_swa_kernel_arch35.h" +#include "sparse_flash_mla_kernel_metadata.h" +#else +#include "kernel_operator.h" +#include "lib/matmul_intf.h" +#include "sparse_flash_mla_template_tiling_key.h" +#include "arch22/sparse_flash_mla_csa_kernel.h" +#include "arch22/sparse_flash_mla_swa_kernel.h" +#include "arch22/sparse_flash_mla_arch22_metadata.h" +#endif + +using namespace AscendC; +using namespace optiling::detail; +using namespace SMLAKernel; + +#if (__CCE_AICORE__ == 310) +#define SMLA_OP_IMPL(templateClass, tilingdataClass, ...) \ + do { \ + using CubeBlockType = \ + typename std::conditional, \ + SMLAKernel::CSABlockCubeDummy<__VA_ARGS__>>::type; \ + using VecBlockType = \ + typename std::conditional, \ + SMLAKernel::CSABlockVec<__VA_ARGS__>>::type; \ + templateClass op; \ + GET_TILING_DATA_WITH_STRUCT(tilingdataClass, tilingDataIn, tiling); \ + const tilingdataClass *__restrict tilingData = &tilingDataIn; \ + op.Init(query, oriKV, cmpKV, oriSparseIndices, cmpSparseIndices, oriBlockTable, cmpBlockTable, cuSeqlensQ, \ + cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedQ, seqUsedOriKV, seqUsedCmpKV, cmpResidualKV, oriTopkLength, \ + cmpTopkLength, sinks, metadata, attentionOut, softmaxLse, user, tilingData); \ + op.Process(); \ + } while (0) +#else +#define SMLA_OP_IMPL(templateClass, tilingdataClass, ...) \ + do { \ + templateClass> op; \ + GET_TILING_DATA_WITH_STRUCT(tilingdataClass, tiling_data_in, tiling); \ + const tilingdataClass *__restrict tiling_data = &tiling_data_in; \ + op.Init(query, oriKV, cmpKV, oriSparseIndices, cmpSparseIndices, oriBlockTable, cmpBlockTable, cuSeqlensQ, \ + cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedQ, seqUsedOriKV, seqUsedCmpKV, cmpResidualKV, oriTopkLength, \ + cmpTopkLength, sinks, metadata, attentionOut, softmaxLse, user, tiling_data, tiling, &tPipe); \ + op.Process(); \ + } while (0) +#endif + +template +__global__ __aicore__ void sparse_flash_mla( + __gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *oriSparseIndices, + __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, + __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv, __gm__ uint8_t *cuSeqlensCmpKv, + __gm__ uint8_t *seqUsedQ, __gm__ uint8_t *seqUsedOriKV, __gm__ uint8_t *seqUsedCmpKV, __gm__ uint8_t *cmpResidualKV, + __gm__ uint8_t *oriTopkLength, __gm__ uint8_t *cmpTopkLength, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata, + __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxLse, __gm__ uint8_t *workspace, __gm__ uint8_t *tiling) +{ + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2); + __gm__ uint8_t *user = GetUserWorkspace(workspace); + +#if (__CCE_AICORE__ == 310) + if constexpr (ORIG_DTYPE_Q == DT_FLOAT16 && ORIG_DTYPE_ORI_KV == DT_FLOAT16 && ORIG_DTYPE_ATTN_OUT == DT_FLOAT16) { + if constexpr (TEMPLATE_MODE == CSA_TEMPLATE || TEMPLATE_MODE == ORI_SPARSE_TEMPLATE || + TEMPLATE_MODE == ORI_CMP_SPARSE_TEMPLATE) { + SMLA_OP_IMPL(SMLAKernel::SparseFlashMlaCsaKernel, SparseFlashMlaTilingData, half, half, float, half, + FLASH_DECODE, static_cast(LAYOUT_T), static_cast(KV_LAYOUT_T), + static_cast(TEMPLATE_MODE), SPLIT_G, BATCH_CONSISTENCY, IS_VEC_S2PHYADDR); + } else { + SMLA_OP_IMPL(SMLAKernel::SparseFlashMlaSwaKernel, SparseFlashMlaTilingData, half, half, float, half, + FLASH_DECODE, static_cast(LAYOUT_T), static_cast(KV_LAYOUT_T), + static_cast(TEMPLATE_MODE), SPLIT_G, BATCH_CONSISTENCY, IS_VEC_S2PHYADDR); + } + } + if constexpr (ORIG_DTYPE_Q == DT_BF16 && ORIG_DTYPE_ORI_KV == DT_BF16 && ORIG_DTYPE_ATTN_OUT == DT_BF16) { + if constexpr (TEMPLATE_MODE == CSA_TEMPLATE || TEMPLATE_MODE == ORI_SPARSE_TEMPLATE || + TEMPLATE_MODE == ORI_CMP_SPARSE_TEMPLATE) { + SMLA_OP_IMPL(SMLAKernel::SparseFlashMlaCsaKernel, SparseFlashMlaTilingData, bfloat16_t, bfloat16_t, float, + bfloat16_t, FLASH_DECODE, static_cast(LAYOUT_T), + static_cast(KV_LAYOUT_T), static_cast(TEMPLATE_MODE), SPLIT_G, + BATCH_CONSISTENCY, IS_VEC_S2PHYADDR); + } else { + SMLA_OP_IMPL(SMLAKernel::SparseFlashMlaSwaKernel, SparseFlashMlaTilingData, bfloat16_t, bfloat16_t, float, + bfloat16_t, FLASH_DECODE, static_cast(LAYOUT_T), + static_cast(KV_LAYOUT_T), static_cast(TEMPLATE_MODE), SPLIT_G, + BATCH_CONSISTENCY, IS_VEC_S2PHYADDR); + } + } +#else + TPipe tPipe; + if constexpr (ORIG_DTYPE_Q == DT_FLOAT16 && ORIG_DTYPE_ORI_KV == DT_FLOAT16 && ORIG_DTYPE_ATTN_OUT == DT_FLOAT16) { + if constexpr (TEMPLATE_MODE == CSA_TEMPLATE) { + SMLA_OP_IMPL(SparseFlashMlaCsa, SparseFlashMlaTilingData, half, half, half, FLASH_DECODE, + static_cast(LAYOUT_T), static_cast(KV_LAYOUT_T), TEMPLATE_MODE, + static_cast(HEAD_RATIO_ONE)); + } else { + SMLA_OP_IMPL(SparseFlashMlaSwa, SparseFlashMlaTilingData, half, half, half, FLASH_DECODE, + static_cast(LAYOUT_T), static_cast(KV_LAYOUT_T), TEMPLATE_MODE, + static_cast(HEAD_RATIO_ONE)); + } + } + if constexpr (ORIG_DTYPE_Q == DT_BF16 && ORIG_DTYPE_ORI_KV == DT_BF16 && ORIG_DTYPE_ATTN_OUT == DT_BF16) { + if constexpr (TEMPLATE_MODE == CSA_TEMPLATE) { + SMLA_OP_IMPL(SparseFlashMlaCsa, SparseFlashMlaTilingData, bfloat16_t, bfloat16_t, bfloat16_t, FLASH_DECODE, + static_cast(LAYOUT_T), static_cast(KV_LAYOUT_T), TEMPLATE_MODE, + static_cast(HEAD_RATIO_ONE)); + } else { + SMLA_OP_IMPL(SparseFlashMlaSwa, SparseFlashMlaTilingData, bfloat16_t, bfloat16_t, bfloat16_t, FLASH_DECODE, + static_cast(LAYOUT_T), static_cast(KV_LAYOUT_T), TEMPLATE_MODE, + static_cast(HEAD_RATIO_ONE)); + } + } +#endif +} diff --git a/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_common.h b/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_common.h new file mode 100644 index 000000000000..d246617d8a87 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_common.h @@ -0,0 +1,336 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_common.h + * \brief + */ + +#ifndef SPARSE_ATTN_SHAREDKV_COMMON_H +#define SPARSE_ATTN_SHAREDKV_COMMON_H + +#include "kernel_operator.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" + +namespace SMLAKernel { +using namespace AscendC; + +enum class SMLATemplateMode { + SWA_TEMPLATE_MODE = 0, + HCA_TEMPLATE_MODE = 1, + CSA_TEMPLATE_MODE = 2, + ORI_SPARSE_TEMPLATE_MODE = 3, + ORI_CMP_SPARSE_TEMPLATE_MODE = 4 +}; + +enum class SMLA_LAYOUT { + BSND = 0, + TND = 1, + PA_BBND = 2 +}; + +#if (__CCE_AICORE__ != 310) +// 将isCheckTiling设置为false, 输入输出的max&sum&exp的shape为(m, 1) +constexpr SoftmaxConfig SMLA_SOFTMAX_FLASHV2_CFG_WITHOUT_BRC = {false, 0, 0, SoftmaxMode::SOFTMAX_OUTPUT_WITHOUT_BRC}; + +template +struct SMLAType { + using queryType = Q_T; + using kvType = KV_T; + using outputType = OUT_T; + static constexpr bool flashDecode = FLASH_DECODE; + static constexpr SMLA_LAYOUT layout = LAYOUT_T; + static constexpr SMLA_LAYOUT kvLayout = KV_LAYOUT_T; + static constexpr bool pageAttention = (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND); + static constexpr int templateMode = TEMPLATE_MODE; + static constexpr bool headRatioOne = HEAD_RATIO_ONE; +}; + +// ================================Util functions================================== +template +__aicore__ inline T1 SMLAAlign(T1 num, T2 rnd) +{ + return (rnd == 0) ? 0 : ((num + rnd - 1) / rnd * rnd); +} + +template +__aicore__ inline T1 CeilDiv(T1 num, T2 rnd) +{ + return (rnd == 0) ? 0 : ((num + rnd - 1) / rnd); +} + +template +__aicore__ inline T1 Min(T1 a, T2 b) +{ + return (a > b) ? b : a; +} + +template +__aicore__ inline T1 Max(T1 a, T2 b) +{ + return (a > b) ? a : b; +} + +template +__aicore__ inline size_t BlockAlign(size_t s) +{ + if constexpr (IsSameType::value) { + return (s + 63) / 64 * 64; + } + size_t n = (32 / sizeof(T)); + return (s + n - 1) / n * n; +} + +struct PAShape { + uint32_t blockSize; + uint32_t headNum; // 一般为kv的head num,对应n2 + uint32_t headDim; // 512 对应d + uint32_t maxblockNumPerBatch; // block table 每一行的最大个数 + uint32_t actHeadDim; // 实际拷贝col大小,考虑到N切块 s*d, 对应d + uint32_t copyRowNum; // 总共要拷贝的行数 + uint32_t copyRowNumAlign; +}; + +struct Position { + uint32_t bIdx; + uint32_t n2Idx; + uint32_t s2Idx; + uint32_t dIdx; +}; + +// 场景:query、key、value GM to L1 +// GM按ND格式存储 +// L1按NZ格式存储 +// GM的行、列、列的stride +template +__aicore__ inline void DataCopyGmNDToL1(LocalTensor &l1Tensor, GlobalTensor &gmTensor, uint32_t rowAct, + uint32_t rowAlign, + uint32_t col, // D + uint32_t colStride) // D or N*D +{ + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = rowAct; // nd矩阵的行数 + // T为int4场景下,dValue = col / 2,srcDValue = colStride / 2 + nd2nzPara.dValue = col; // nd矩阵的列数 + nd2nzPara.srcDValue = colStride; // 同一nd矩阵相邻行起始地址间的偏移 + nd2nzPara.dstNzC0Stride = rowAlign; + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(l1Tensor, gmTensor, nd2nzPara); +} + +/* + 适用PA数据从GM拷贝到L1,支持ND、NZ数据; + PA的layout分 BBND(blockNum,N,blockSize,D) BBH(blockNum,blockSize,N*D + BSH\BSND\TND 为BBH + shape.copyRowNumAlign 需要16字节对齐,如拷贝k矩阵,一次拷贝128*512,遇到尾块 10*512 需对齐到16*512 +*/ +template +__aicore__ inline void DataCopyPA(LocalTensor &dstTensor, // l1 + GlobalTensor &srcTensor, // gm + GlobalTensor &blockTableGm, + const PAShape &shape, // blockSize, headNum, headDim + const Position &startPos) // bacthIdx nIdx curSeqIdx +{ + uint32_t copyFinishRowCnt = 0; + uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; + uint32_t curS2Idx = startPos.s2Idx; + uint32_t blockElementCnt = 32 / sizeof(T); + while (copyFinishRowCnt < shape.copyRowNum) { + uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引 + uint64_t reaminRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数 + uint64_t idInBlockTable = + blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上的获取编号 + uint32_t copyRowCnt = shape.blockSize - reaminRowCnt; // 一次只能处理一个Block + if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) { + copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满 + } + uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim; // PA的偏移 + + uint64_t dStride = shape.headDim; + if constexpr (SRC_LAYOUT == SMLA_LAYOUT::BSND || SRC_LAYOUT == SMLA_LAYOUT::TND) { + offset += (uint64_t)(startPos.n2Idx * shape.headDim) + reaminRowCnt * shape.headDim * shape.headNum + + startPos.dIdx; + dStride = shape.headDim * shape.headNum; + } else { + offset += (uint64_t)(startPos.n2Idx * shape.headDim * shape.blockSize) + reaminRowCnt * shape.headDim + + startPos.dIdx; + } + + uint32_t dValue = shape.actHeadDim; + uint32_t srcDValue = dStride; + LocalTensor tmpDstTensor = dstTensor[copyFinishRowCnt * blockElementCnt]; + GlobalTensor tmpSrcTensor = srcTensor[offset]; + + DataCopyGmNDToL1(tmpDstTensor, tmpSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue); + copyFinishRowCnt += copyRowCnt; + curS2Idx += copyRowCnt; + } +} + +struct RunInfo { + uint32_t loop = 0; + uint32_t cmpLoop = 0; // 用于判断取 用于merge的4块GM 中的哪一块 + uint32_t bIdx = 0; + uint32_t gIdx = 0; + uint32_t s1Idx = 0; + uint32_t s2Idx = 0; + uint32_t relativeS2Idx = 0; + uint32_t bn2IdxInCurCore = 0; + uint32_t curSInnerLoopTimes = 0; + uint64_t tndBIdxOffsetForQ = 0; + uint64_t tndBIdxOffsetForKV = 0; + uint64_t tensorAOffset = 0; + uint64_t tensorBOffset = 0; + uint64_t attenOutOffset = 0; + uint64_t attenMaskOffset = 0; + uint64_t topKBaseOffset = 0; + uint32_t actualSingleProcessSInnerSize = 0; + uint32_t actualSingleProcessSInnerSizeAlign = 0; + bool isFirstSInnerLoop = false; + uint32_t s2BatchOffset = 0; + uint32_t gSize = 0; + uint32_t s1Size = 0; + uint32_t s2Size = 0; + uint32_t mSize = 0; + uint32_t mSizeV = 0; + uint32_t mSizeVStart = 0; + uint32_t tndIsS2SplitCore = 0; + uint32_t tndCoreStartKVSplitPos = 0; + bool isBmm2Output = false; + bool isValid = false; + + static constexpr uint32_t n2Idx = 0; + uint64_t actS1Size = 1; + uint64_t actS2SizeOri = 0ULL; + uint32_t gS1Idx = 0; + uint64_t actS2Size = 1; + uint64_t actOriS2Size = 1; + uint32_t actMBaseSize = 0; + bool isLastS2Loop = 0; + int32_t nextTokensPerBatch = 0; + int64_t threshold = 0; + uint32_t curTopKIdx = 0; + uint64_t curOffsetInSparseBlock = 0; + bool isOri = true; // 判断当前块是在Ori部分还是Cmp部分 + uint64_t s2StartPoint = 0; + int64_t cmpS2IdLimit = 0; + int32_t v0S2DealSize = 0; + int32_t v0S2Start = 0; +}; + +struct ConstInfo { + // CUBE与VEC核间同步的模式 + static constexpr uint32_t SMLA_SYNC_MODE2 = 2; + // BUFFER的字节数 + static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32; + static constexpr uint32_t BUFFER_SIZE_BYTE_64B = 64; + static constexpr uint32_t BUFFER_SIZE_BYTE_256B = 256; + static constexpr uint32_t BUFFER_SIZE_BYTE_512B = 512; + static constexpr uint32_t BUFFER_SIZE_BYTE_1K = 1024; + static constexpr uint32_t BUFFER_SIZE_BYTE_2K = 2048; + static constexpr uint32_t BUFFER_SIZE_BYTE_4K = 4096; + static constexpr uint32_t BUFFER_SIZE_BYTE_8K = 8192; + static constexpr uint32_t BUFFER_SIZE_BYTE_16K = 16384; + static constexpr uint32_t BUFFER_SIZE_BYTE_32K = 32768; + // FP32的0值和极大值 + static constexpr float FLOAT_ZERO = 0; + static constexpr float FLOAT_MAX = 3.402823466e+38F; + + // preLoad的总次数 + uint32_t preLoadNum = 0U; + uint32_t nBufferMBaseSize = 0U; + // CUBE和VEC的核间同步EventID + uint32_t syncV0C1 = 0U; + uint32_t syncC1V1 = 0U; + uint32_t syncV1C2 = 0U; + uint32_t syncC2V2 = 0U; + + uint32_t mmResUbSize = 0U; // Matmul1输出结果GM上的大小 + uint32_t vec1ResUbSize = 0U; // Vector1输出结果GM上的大小 + uint32_t bmm2ResUbSize = 0U; // Matmul2输出结果GM上的大小 + uint64_t batchSize = 0ULL; + uint64_t gSize = 0ULL; + uint64_t qHeadNum = 0ULL; + uint64_t kvHeadNum = 0; + uint64_t headDim = 0; + uint64_t kvSeqSize = 0ULL; // kv最大S长度 + uint64_t qSeqSize = 1ULL; // q最大S长度 + int64_t kvCacheBlockSize = 0; // PA场景的block size + uint64_t paCmpBlockSize = 0; + uint64_t paOriBlockSize = 0; + int64_t orikvCacheBlockSize = 0; + int64_t cmpkvCacheBlockSize = 0; + uint32_t oriMaxBlockNumPerBatch = 0; // PA场景的最大单batch block number + uint32_t cmpMaxBlockNumPerBatch = 0; + uint32_t splitKVNum = 0U; // S2核间切分的切分份数 + SMLA_LAYOUT outputLayout; // 输出的Transpose格式 + uint32_t oriMaskMode = 0; + uint32_t cmpMaskMode = 0; + bool needInit = false; + uint32_t templateMode = 0; + + // FlashDecoding + uint32_t actualCombineLoopSize = 0U; // FlashDecoding场景, S2在核间切分的最大份数 + uint64_t combineLseOffset = 0ULL; + uint64_t combineAccumOutOffset = 0ULL; + + uint32_t actualLenDimsQ = 0U; // query的actualSeqLength 的维度 + uint32_t actualLenDimsKV = 0U; // KV 的actualSeqLength 的维度 + + // TND + uint32_t s2Start = 0U; // TND场景下,S2的起始位置 + uint32_t s2End = 0U; // 单核TND场景下S2循环index上限 + + uint32_t bN2Start = 0U; + uint32_t bN2End = 0U; + uint32_t gS1Start = 0U; + uint32_t gS1End = 0U; + + uint32_t tndFDCoreArrLen = 0U; // TNDFlashDecoding相关分核信息array的长度 + uint32_t coreStartKVSplitPos = 0U; // TNDFlashDecoding kv起始位置 + + uint32_t mBaseSize = 1ULL; + uint32_t s2BaseSize = 1ULL; + + uint32_t subBlockIdx = 0; + uint32_t aivIdx = 0; + + // sparse attr + int64_t sparseBlockSize = 0; + uint32_t sparseBlockCount = 0; + uint32_t oriSparseBlockCount = 0; + uint32_t cmpSparseBlockCount = 0; + bool hasOriTopkLength = false; + bool hasCmpTopkLength = false; + + // cmp attr + int64_t cmpRatio = 0; + + // win + int32_t oriWinRight = 0; + int32_t oriWinLeft = 128; +}; + +struct MSplitInfo { + uint32_t nBufferIdx = 0U; + uint32_t nBufferStartM = 0U; + uint32_t nBufferDealM = 0U; + uint32_t vecStartM = 0U; + uint32_t vecDealM = 0U; +}; +#endif +} // namespace SMLAKernel +#endif // SPARSE_ATTN_SHAREDKV_COMMON_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_kernel_metadata.h b/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_kernel_metadata.h new file mode 100644 index 000000000000..2459b2eab37f --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_kernel_metadata.h @@ -0,0 +1,80 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_kernel_metadata.h + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_KERNEL_METADATA_H +#define SPARSE_FLASH_MLA_KERNEL_METADATA_H + +#include + +namespace optiling { + +// Constants +constexpr uint32_t AIC_CORE_MAX_NUM = 36; +constexpr uint32_t AIV_CORE_MAX_NUM = 72; +constexpr uint32_t SMLA_METADATA_TOTAL_SIZE = 1024; +using SMLA_METADATA_T = int32_t; + +constexpr uint32_t FA_METADATA_SIZE = 9; +constexpr uint32_t FD_METADATA_SIZE = 8; + +// FA Metadata Index Definitions +constexpr uint32_t FA_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FA_BN2_START_INDEX = 1; +constexpr uint32_t FA_M_START_INDEX = 2; +constexpr uint32_t FA_S2_START_INDEX = 3; +constexpr uint32_t FA_BN2_END_INDEX = 4; +constexpr uint32_t FA_M_END_INDEX = 5; +constexpr uint32_t FA_S2_END_INDEX = 6; +constexpr uint32_t FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX = 7; +constexpr uint32_t FA_S2_MAX_NUM = 8; + +// FD Metadata Index Definitions +constexpr uint32_t FD_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FD_BN2_IDX_INDEX = 1; +constexpr uint32_t FD_M_IDX_INDEX = 2; +constexpr uint32_t FD_WORKSPACE_IDX_INDEX = 3; +constexpr uint32_t FD_WORKSPACE_NUM_INDEX = 4; +constexpr uint32_t FD_M_START_INDEX = 5; +constexpr uint32_t FD_M_NUM_INDEX = 6; + +/** + * @brief 获取属性的绝对索引 + * @param coreIdx 核索引 + * @param metaIdx 元数据索引 + * @param isAIV 是否为AIV数据,默认为false + * @return 返回属性的绝对索引 + */ +#ifdef __CCE_AICORE__ +__aicore__ inline uint32_t GetAttrAbsIndex(uint32_t coreIdx, uint32_t metaIdx, bool isAIV = false) +{ + if (isAIV) { + return FA_METADATA_SIZE * AIC_CORE_MAX_NUM + FD_METADATA_SIZE * coreIdx + metaIdx; + } else { + return FA_METADATA_SIZE * coreIdx + metaIdx; + } +} +#endif + +namespace detail { +struct SmlaMetadata { + uint32_t faMetadata[AIC_CORE_MAX_NUM][FA_METADATA_SIZE]; + uint32_t fdMetadata[AIV_CORE_MAX_NUM][FD_METADATA_SIZE]; +}; +}; // namespace detail + +static_assert(SMLA_METADATA_TOTAL_SIZE * sizeof(SMLA_METADATA_T) >= sizeof(detail::SmlaMetadata)); +}; // namespace optiling + +#endif // SPARSE_FLASH_MLA_KERNEL_METADATA_H diff --git a/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_template_tiling_key.h b/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_template_tiling_key.h new file mode 100644 index 000000000000..cf76aa9454d7 --- /dev/null +++ b/csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_template_tiling_key.h @@ -0,0 +1,107 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_template_tiling_key.h + * \brief + */ + +#ifndef SPARSE_ATTN_SHARED_TEMPLATE_TILING_KEY_H +#define SPARSE_ATTN_SHARED_TEMPLATE_TILING_KEY_H + +#include "ascendc/host_api/tiling/template_argument.h" + +#define SMLA_LAYOUT_BSND 0 +#define SMLA_LAYOUT_TND 1 +#define SMLA_LAYOUT_PA_BBND 2 + +#define ASCENDC_TPL_4_BW 4 + +#define SWA_TEMPLATE 0 +#define HCA_TEMPLATE 1 +#define CSA_TEMPLATE 2 +#define ORI_SPARSE_TEMPLATE 3 +#define ORI_CMP_SPARSE_TEMPLATE 4 +// 模板参数支持的范围定义 +ASCENDC_TPL_ARGS_DECL(SparseFlashMla, // 算子OpType + ASCENDC_TPL_BOOL_DECL(FLASH_DECODE, 0, 1), + ASCENDC_TPL_UINT_DECL(LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_BSND, + SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_DECL(KV_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_BSND, + SMLA_LAYOUT_TND, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE, + HCA_TEMPLATE, CSA_TEMPLATE, ORI_SPARSE_TEMPLATE, ORI_CMP_SPARSE_TEMPLATE), + ASCENDC_TPL_BOOL_DECL(SPLIT_G, 0, 1), ASCENDC_TPL_BOOL_DECL(HEAD_RATIO_ONE, 0, 1), + ASCENDC_TPL_BOOL_DECL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_DECL(IS_VEC_S2PHYADDR, 0, 1)); + +// Aurora always uses TND Q and PA_BBND KV: C0 -> SWA, C1/C2 -> CSA. +// Keep the declaration above unchanged to preserve the host/kernel key ABI. +// Only CSA on A2/A3 uses HEAD_RATIO_ONE. Split-G and vectorized sparse +// addressing belong to A5. Both deterministic levels remain reachable. +#if !defined(__CCE_AICORE__) +// Host validation accepts the union of the two device selections (14 keys). +ASCENDC_TPL_SEL( + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0), + ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE, CSA_TEMPLATE), + ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0), ASCENDC_TPL_BOOL_SEL(HEAD_RATIO_ONE, 0), + ASCENDC_TPL_BOOL_SEL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_SEL(IS_VEC_S2PHYADDR, 0)), + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0), + ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, CSA_TEMPLATE), + ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0), ASCENDC_TPL_BOOL_SEL(HEAD_RATIO_ONE, 1), + ASCENDC_TPL_BOOL_SEL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_SEL(IS_VEC_S2PHYADDR, 0)), + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0), + ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE, CSA_TEMPLATE), + ASCENDC_TPL_BOOL_SEL(SPLIT_G, 1), ASCENDC_TPL_BOOL_SEL(HEAD_RATIO_ONE, 0), + ASCENDC_TPL_BOOL_SEL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_SEL(IS_VEC_S2PHYADDR, 0)), + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0), + ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, CSA_TEMPLATE), + ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0, 1), ASCENDC_TPL_BOOL_SEL(HEAD_RATIO_ONE, 0), + ASCENDC_TPL_BOOL_SEL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_SEL(IS_VEC_S2PHYADDR, 1))); +#elif (__CCE_AICORE__ == 310) +// A5: keep split-G and CSA vectorization; never compile HEAD_RATIO_ONE=1. +ASCENDC_TPL_SEL( + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0), + ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE), + ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0, 1), ASCENDC_TPL_BOOL_SEL(HEAD_RATIO_ONE, 0), + ASCENDC_TPL_BOOL_SEL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_SEL(IS_VEC_S2PHYADDR, 0)), + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0), + ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, CSA_TEMPLATE), + ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0, 1), ASCENDC_TPL_BOOL_SEL(HEAD_RATIO_ONE, 0), + ASCENDC_TPL_BOOL_SEL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_SEL(IS_VEC_S2PHYADDR, 0, 1))); +#else +// A2/A3: keep the CSA single-head kernel; never compile A5-only flags. +ASCENDC_TPL_SEL( + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0), + ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE), + ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0), ASCENDC_TPL_BOOL_SEL(HEAD_RATIO_ONE, 0), + ASCENDC_TPL_BOOL_SEL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_SEL(IS_VEC_S2PHYADDR, 0)), + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0), + ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_TND), + ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SMLA_LAYOUT_PA_BBND), + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, CSA_TEMPLATE), + ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0), ASCENDC_TPL_BOOL_SEL(HEAD_RATIO_ONE, 0, 1), + ASCENDC_TPL_BOOL_SEL(BATCH_CONSISTENCY, 0, 1), ASCENDC_TPL_BOOL_SEL(IS_VEC_S2PHYADDR, 0))); +#endif + +#endif // TEMPLATE_TILING_KEY diff --git a/csrc/attention/sparse_flash_mla/sparse_flash_mla_torch_adpt.h b/csrc/attention/sparse_flash_mla/sparse_flash_mla_torch_adpt.h new file mode 100644 index 000000000000..3679c7be2cbc --- /dev/null +++ b/csrc/attention/sparse_flash_mla/sparse_flash_mla_torch_adpt.h @@ -0,0 +1,182 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ +/*! + * \file sparse_flash_mla.cpp + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_TORCH_ADPT_H +#define SPARSE_FLASH_MLA_TORCH_ADPT_H + +namespace vllm_ascend { +constexpr int64_t DIM_0 = 0; +constexpr int64_t DIM_1 = 1; +constexpr int64_t DIM_2 = 2; +constexpr int64_t DIM_3 = 3; +constexpr int64_t DIM_4 = 4; + +constexpr int64_t SMLA_METADATA_SIZE = 1024; + +inline at::Tensor GetValidSparseFlashMlaTensor( + const c10::optional &tensor, const at::Device &device) +{ + return tensor.has_value() + ? tensor.value() + : at::empty({0}, at::TensorOptions().dtype(at::kInt).device(device)); +} + +at::Tensor npu_sparse_flash_mla_metadata( + int64_t numHeadsQ, int64_t numHeadsKv, int64_t headDim, const c10::optional &cuSeqlensQ, + const c10::optional &cuSeqlensOriKv, const c10::optional &cuSeqlensCmpKv, + const c10::optional &sequsedQ, const c10::optional &sequsedOriKv, + const c10::optional &sequsedCmpKv, const c10::optional &cmpResidualKv, + const c10::optional &oriTopkLength, const c10::optional &cmpTopkLength, int64_t batchSize, + int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, int64_t maxSeqlenCmpKv, int64_t oriTopk, int64_t cmpTopk, + int64_t cmpRatio, int64_t oriMaskMode, int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight, + c10::string_view layoutQ, c10::string_view layoutKv, bool hasOriKv, bool hasCmpKv) +{ + at::Device outputDevice = at::Device(std::string("npu")); + if (cuSeqlensQ.has_value()) { + outputDevice = cuSeqlensQ.value().device(); + } else if (cuSeqlensOriKv.has_value()) { + outputDevice = cuSeqlensOriKv.value().device(); + } else if (cuSeqlensCmpKv.has_value()) { + outputDevice = cuSeqlensCmpKv.value().device(); + } else if (sequsedQ.has_value()) { + outputDevice = sequsedQ.value().device(); + } else if (sequsedOriKv.has_value()) { + outputDevice = sequsedOriKv.value().device(); + } else if (sequsedCmpKv.has_value()) { + outputDevice = sequsedCmpKv.value().device(); + } else if (cmpResidualKv.has_value()) { + outputDevice = cmpResidualKv.value().device(); + } else if (oriTopkLength.has_value()) { + outputDevice = oriTopkLength.value().device(); + } else if (cmpTopkLength.has_value()) { + outputDevice = cmpTopkLength.value().device(); + } + + at::Tensor output = torch::empty({SMLA_METADATA_SIZE}, torch::dtype(torch::kInt32).device(outputDevice)); + auto cuSeqlensQVal = GetValidSparseFlashMlaTensor(cuSeqlensQ, outputDevice); + auto cuSeqlensOriKvVal = GetValidSparseFlashMlaTensor(cuSeqlensOriKv, outputDevice); + auto cuSeqlensCmpKvVal = GetValidSparseFlashMlaTensor(cuSeqlensCmpKv, outputDevice); + auto sequsedQVal = GetValidSparseFlashMlaTensor(sequsedQ, outputDevice); + auto sequsedOriKvVal = GetValidSparseFlashMlaTensor(sequsedOriKv, outputDevice); + auto sequsedCmpKvVal = GetValidSparseFlashMlaTensor(sequsedCmpKv, outputDevice); + auto cmpResidualKvVal = GetValidSparseFlashMlaTensor(cmpResidualKv, outputDevice); + auto oriTopkLengthVal = GetValidSparseFlashMlaTensor(oriTopkLength, outputDevice); + auto cmpTopkLengthVal = GetValidSparseFlashMlaTensor(cmpTopkLength, outputDevice); + + // convert str + std::string layoutQStr = std::string(layoutQ); + std::string layoutKvStr = std::string(layoutKv); + char *layoutQPtr = const_cast(layoutQStr.c_str()); + char *layoutKvPtr = const_cast(layoutKvStr.c_str()); + + EXEC_NPU_CMD(aclnnSparseFlashMlaMetadata, cuSeqlensQVal, cuSeqlensOriKvVal, cuSeqlensCmpKvVal, sequsedQVal, + sequsedOriKvVal, sequsedCmpKvVal, cmpResidualKvVal, oriTopkLengthVal, cmpTopkLengthVal, numHeadsQ, + numHeadsKv, headDim, batchSize, maxSeqlenQ, maxSeqlenOriKv, maxSeqlenCmpKv, oriTopk, cmpTopk, cmpRatio, + oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, layoutQPtr, layoutKvPtr, hasOriKv, hasCmpKv, output); + return output; +} + +void CheckQueryShape(const at::Tensor &q, const std::string &layoutQStr) +{ + TORCH_CHECK(layoutQStr == "BSND" || layoutQStr == "TND", "The layout of query only support BSND and TND, but got ", + layoutQStr); + TORCH_CHECK(q.numel() > 0, "Tensor query is empty."); + for (int64_t i = 0; i < q.dim(); i++) { + TORCH_CHECK(q.size(i) > 0, + "All values within query's shape should be greater " + "than 0, but shape[", + i, "] is ", q.size(i)); + } + if (layoutQStr == "BSND") { + TORCH_CHECK(q.dim() == DIM_4, "When the layout of query is BSND, the query dimension must be 4, but got ", + q.dim()); + } else { + TORCH_CHECK(q.dim() == DIM_3, "When the layout of query is TND, the query dimension must be 3, but got ", + q.dim()); + } +} + +int64_t GetKvHeadNum(const c10::optional &oriKv, const c10::optional &cmpKv, + const std::string &layoutKvStr) +{ + TORCH_CHECK(oriKv.has_value() || cmpKv.has_value(), + "ori_kv or cmp_kv is required when return_softmax_lse is true."); + const at::Tensor &kv = oriKv.has_value() ? oriKv.value() : cmpKv.value(); + if (layoutKvStr == "TND") { + return kv.size(DIM_1); + } + return kv.size(DIM_2); +} + +std::tuple MakeSparseFlashMlaOutputs(const at::Tensor &q, + const c10::optional &oriKv, + const c10::optional &cmpKv, + const std::string &layoutQStr, + const std::string &layoutKvStr, bool returnSoftmaxLse) +{ + CheckQueryShape(q, layoutQStr); + at::Tensor attenOut = at::empty_like(q); + at::Tensor softmaxLse; + if (!returnSoftmaxLse) { + softmaxLse = at::empty({0}, q.options().dtype(torch::kFloat32)); + return {attenOut, softmaxLse}; + } + + int64_t kvHeadNum = GetKvHeadNum(oriKv, cmpKv, layoutKvStr); + TORCH_CHECK(kvHeadNum > 0, "head num of ori_kv or cmp_kv must be greater than 0, but got ", kvHeadNum); + if (layoutQStr == "BSND") { + softmaxLse = at::empty({q.size(DIM_0), kvHeadNum, q.size(DIM_1), q.size(DIM_2) / kvHeadNum}, + q.options().dtype(torch::kFloat32)); + } else { + softmaxLse = + at::empty({kvHeadNum, q.size(DIM_0), q.size(DIM_1) / kvHeadNum}, q.options().dtype(torch::kFloat32)); + } + return {attenOut, softmaxLse}; +} + +std::tuple npu_sparse_flash_mla( + const at::Tensor &q, const c10::optional &oriKv, const c10::optional &cmpKv, + const c10::optional &oriSparseIndices, const c10::optional &cmpSparseIndices, + const c10::optional &oriBlockTable, const c10::optional &cmpBlockTable, + const c10::optional &cuSeqlensQ, const c10::optional &cuSeqlensOriKv, + const c10::optional &cuSeqlensCmpKv, const c10::optional &sequsedQ, + const c10::optional &sequsedOriKv, const c10::optional &sequsedCmpKv, + const c10::optional &cmpResidualKv, const c10::optional &oriTopkLength, + const c10::optional &cmpTopkLength, const c10::optional &sinks, + const c10::optional &metadata, double softmaxScale, int64_t cmpRatio, int64_t oriMaskMode, + int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight, c10::string_view layoutQ, c10::string_view layoutKv, + int64_t topkValueMode, bool returnSoftmaxLse) +{ + std::string layoutQStr = std::string(layoutQ); + std::string layoutKvStr = std::string(layoutKv); + // convert str + char *layoutQPtr = const_cast(layoutQStr.c_str()); + char *layoutKvPtr = const_cast(layoutKvStr.c_str()); + + // construct the atten_out tensor + std::tuple sparseFlashMlaAttenOut = + MakeSparseFlashMlaOutputs(q, oriKv, cmpKv, layoutQStr, layoutKvStr, returnSoftmaxLse); + at::Tensor attenOut = std::get<0>(sparseFlashMlaAttenOut); + at::Tensor softmaxLse = std::get<1>(sparseFlashMlaAttenOut); + + EXEC_NPU_CMD(aclnnSparseFlashMla, q, oriKv, cmpKv, oriSparseIndices, cmpSparseIndices, oriBlockTable, cmpBlockTable, + cuSeqlensQ, cuSeqlensOriKv, cuSeqlensCmpKv, sequsedQ, sequsedOriKv, sequsedCmpKv, cmpResidualKv, + oriTopkLength, cmpTopkLength, sinks, metadata, softmaxScale, cmpRatio, oriMaskMode, cmpMaskMode, + oriWinLeft, oriWinRight, layoutQPtr, layoutKvPtr, topkValueMode, returnSoftmaxLse, attenOut, softmaxLse); + return std::tuple(attenOut, softmaxLse); +} + +} // namespace vllm_ascend + +#endif // SPARSE_FLASH_MLA_TORCH_ADPT_H diff --git a/csrc/attention/sparse_flash_mla_metadata/CMakeLists.txt b/csrc/attention/sparse_flash_mla_metadata/CMakeLists.txt new file mode 100644 index 000000000000..e99a153f311b --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/CMakeLists.txt @@ -0,0 +1,19 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +if(NOT ENABLE_TEST AND NOT BENCHMARK) + list(REMOVE_ITEM CURRENT_DIRS tests) +endif() +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/sparse_flash_mla_metadata/README.md b/csrc/attention/sparse_flash_mla_metadata/README.md new file mode 100644 index 000000000000..b6a9aa42a5dd --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/README.md @@ -0,0 +1,299 @@ +# SparseFlashMlaMetadata + +## 产品支持情况 + +| 产品 | 是否支持 | +| ------------------------------------------------------------ | :------: | +|Ascend 950PR/Ascend 950DT | √ | +|Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ | +|Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ | +|Atlas 200I/500 A2推理系列产品 | × | +|Atlas 推理系列产品 | × | +|Atlas 训练系列产品 | × | + +## 功能说明 + +- 算子功能:`SparseFlashMlaMetadata`是`SparseFlashMla`算子的前置算子,用于后续Attention计算生成负载均衡的任务划分方案。本算子不执行实际的Attention计算,而是根据输入参数在AI CPU计算出每个AI Core应处理的Attention计算起止范围,从而最大化计算资源的利用率,避免各Core间负载不均衡的问题。 + +- 场景简称:SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)。 + +## 参数说明 + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出/属性描述数据类型数据格式
num_heads_q属性表示`q`的头数,支持[1, 128]。INT-
num_heads_kv属性表示`ori_kv`和`cmp_kv`的头数,仅支持1。INT-
head_dim属性注意力头的维度,仅支持512。INT-
cu_seqlens_q可选输入表示TND布局下不同batch中`q`的累积序列长度,shape为(B+1, )。INT32ND
cu_seqlens_ori_kv可选输入表示TND布局下不同batch中`ori_kv`的累积序列长度,shape为(B+1, )。INT32ND
cu_seqlens_cmp_kv可选输入表示TND布局下不同batch中`cmp_kv`的累积序列长度,shape为(B+1, )。INT32ND
seqused_q可选输入表示不同batch中`q`实际参与计算的token数,shape为(B, )。INT32ND
seqused_ori_kv可选输入表示不同batch中`ori_kv`实际参与计算的token数,shape为(B, )。INT32ND
seqused_cmp_kv可选输入表示不同batch中`cmp_kv`实际参与计算的token数,shape为(B, )。INT32ND
cmp_residual_kv可选输入表示压缩KV余数,用于恢复cmp侧mask使用的压缩前KV长度,shape为(B, )。INT32ND
ori_topk_length可选输入SWA稀疏ori_kv场景表示不同q token对应的ori_kv部分关键稀疏token的个数,必须传入,shape为(B, S1, N2)或(T1, N2)。INT32ND
cmp_topk_length可选输入表示不同q token对应的cmp_kv部分关键稀疏token的个数,shape为(B, S1, N2)或(T1, N2)。INT32ND
batch_size可选属性表示输入样本批量大小;传入0时表示由接口推导,默认值为0。INT-
max_seqlen_q可选属性表示所有batch中`q`的最大有效token数;传入0时表示由接口推导,默认值为0。INT-
max_seqlen_ori_kv可选属性表示所有batch中`ori_kv`的最大有效token数;传入0时表示由接口推导,默认值为0。INT-
max_seqlen_cmp_kv可选属性表示所有batch中`cmp_kv`的最大有效token数;传入0时表示由接口推导,默认值为0。INT-
ori_topk可选属性表示从`ori_kv`中筛选出的关键稀疏token个数;SWA稀疏ori_kv场景为主算子`ori_sparse_indices`最后一维K且必须大于0,其他场景默认值为0。INT-
cmp_topk可选属性表示从`cmp_kv`中筛选出的关键稀疏token个数,默认值为0。INT-
cmp_ratio可选属性表示`cmp_kv`相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度;仅传入`ori_kv`时不参与压缩KV计算。传入`cmp_kv`时支持[1, 128],默认值为0。INT-
ori_mask_mode可选属性表示`q`和`ori_kv`计算的mask模式,默认值为0。
0: No Mask。
3: RightDownCausal模式。
4: Band模式。
INT-
cmp_mask_mode可选属性表示`q`和`cmp_kv`计算的mask模式,默认值为0。
0: No Mask。
3: RightDownCausal模式。
INT-
ori_win_left可选属性表示`q`和`ori_kv`计算中`q`对过去token计算的数量,支持-1或非负数,其中-1表示窗口不受限,默认值为-1。INT-
ori_win_right可选属性表示`q`和`ori_kv`计算中`q`对未来token计算的数量,支持-1或非负数,其中-1表示窗口不受限,默认值为-1。INT-
layout_q可选属性表示输入`q`的数据排布格式,支持"BSND"和"TND",默认值为"BSND"。STRING-
layout_kv可选属性表示输入`ori_kv`和`cmp_kv`的数据排布格式,支持"BSND"、"TND"和"PA_BBND",默认值为"BSND"。STRING-
has_ori_kv可选属性表示`SparseFlashMla`主算子是否传入`ori_kv`,默认值为true。BOOL-
has_cmp_kv可选属性表示`SparseFlashMla`主算子是否传入`cmp_kv`,默认值为true。BOOL-
metadata输出表示`SparseFlashMla`主算子使用的任务切分结果,shape固定为(1024, )。INT32ND
+
    +
  • Atlas A3 训练系列产品/Atlas A3 推理系列产品 :num_heads_q/num_heads_kv仅支持1、2、4、8、16、32、64、128,不支持seqused_q、cmp_topk_length;SWA稀疏ori_kv场景支持ori_topk_length、ori_topk大于0及ori_mask_mode为0,ori_win_left和ori_win_right支持非负数;其他SWA场景ori_topk为0、ori_mask_mode为4、ori_win_left为127、ori_win_right为0;cmp_topk仅支持0、512、1024,cmp_mask_mode仅支持3,cmp_ratio在SWA支持0、CSA支持1、2或4、HCA支持128。
  • +
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品 :num_heads_q/num_heads_kv仅支持1、2、4、8、16、32、64、128,不支持seqused_q、cmp_topk_length;SWA稀疏ori_kv场景支持ori_topk_length、ori_topk大于0及ori_mask_mode为0,ori_win_left和ori_win_right支持非负数;其他SWA场景ori_topk为0、ori_mask_mode为4、ori_win_left为127、ori_win_right为0;cmp_topk仅支持0、512、1024,cmp_mask_mode仅支持3,cmp_ratio在SWA支持0、CSA支持1、2或4、HCA支持128。
  • +
+ +## 约束说明 + +- 该接口支持训练、推理场景下使用。 +- 该接口支持aclgraph模式。 +- 通用规格约束如下: + - B(Batch)表示输入样本批量大小,q、ori_kv、cmp_kv为配套的SparseFlashMla算子的入参,S1表示layout_q=BSND时,q shape中的S轴的大小,T1表示layout_q=TND时,q shape中的T轴的大小,S2表示layout_kv=BSND时,ori_kv shape中的S轴的大小,S3表示layout_kv=BSND时,cmp_kv shape中的S轴的大小,N2表示ori_kv、cmp_kv shape中的N轴的大小。 + - 参数`cu_seqlens_q`、`cu_seqlens_ori_kv`及`cu_seqlens_cmp_kv`要求其值为当前Batch与前序Batch有效token数的累加值,第一个元素固定为0,后一个元素的值必须大于等于前一个元素的值。 + - 参数`seqused_q`、`seqused_ori_kv`、`seqused_cmp_kv`要求其值表示每个Batch中的有效token数。 + - `layout_q`和`layout_kv`组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下`layout_q`和`layout_kv`必须一致。 + - 参数`cmp_residual_kv`需满足`cmp_residual_kv`[i] < `cmp_ratio`。 +- Ascend 950PR/Ascend 950DT约束: + - has_ori_kv为true时,ori_topk大于0认为ori_kv部分是稀疏的,ori_topk为0则认为ori_kv部分是非稀疏的。 + - has_cmp_kv为true时,cmp_topk大于0认为cmp_kv部分是稀疏的,cmp_topk为0则认为cmp_kv部分是非稀疏的。 + - has_ori_kv为true,ori_topk不为0且ori_mask_mode为0时,ori_topk_length必须传入,此时取ori_mask_mode规则与ori_topk_length元素的最小值作为当前q token对应的ori_kv的有效seqlen,其他ori_kv稀疏场景取ori_mask_mode规则与ori_topk的最小值作为当前q token对应的ori_kv的有效seqlen。 + - has_cmp_kv为true,cmp_topk不为0且cmp_mask_mode为0时,cmp_topk_length必须传入,此时取cmp_mask_mode规则与cmp_topk_length元素的最小值作为当前q token对应的cmp_kv的有效seqlen,其他cmp_kv稀疏场景取cmp_mask_mode规则与cmp_topk的最小值作为当前q token对应的cmp_kv的有效seqlen。 + - layout_q=BSND场景 + - max_seqlen_q必须传入S1的值。 + - layout_kv=BSND场景 + - has_ori_kv为true时,max_seqlen_ori_kv必须传入S2的值。 + - has_cmp_kv为true时,max_seqlen_cmp_kv必须传入S3的值。 + - layout_q=TND场景 + - cu_seqlens_q必须传入。 + - layout_kv=TND场景 + - has_ori_kv为true时,cu_seqlens_ori_kv必须传入。 + - has_cmp_kv为true时,cu_seqlens_cmp_kv必须传入。 + - layout_kv=PA_BBND场景 + - has_ori_kv为true,ori_topk不为0且ori_mask_mode为0时(ori_topk_length必传场景),seqused_ori_kv可选传入,其他场景seqused_ori_kv必须传入。 + - has_cmp_kv为true,cmp_topk不为0且cmp_mask_mode为0时(cmp_topk_length必传场景),seqused_cmp_kv可选传入,其他场景seqused_cmp_kv必须传入。 + - Batch取值规则 + - layout_q为BSND时,优先通过seqused_q的shape推导batch,seqused_q未传入则通过batch_size获取batch数。 + - layout_q为TND时,优先通过seqused_q的shape推导batch,seqused_q未传入则通过cu_seqlens_q的shape推导batch。 + - q Seqlen取值规则 + - layout_q为BSND时,优先通过seqused_q中的元素获取seqlen,seqused_q未传入则通过max_seqlen_q获取seqlen。 + - layout_q为TND时,优先通过seqused_q中的元素获取seqlen,seqused_q未传入则通过cu_seqlens_q中的元素获取seqlen。 + - ori_kv Seqlen取值规则 + - layout_kv为BSND时,优先通过seqused_ori_kv中的元素获取seqlen,seqused_ori_kv未传入则通过max_seqlen_ori_kv获取seqlen。 + - layout_kv为TND时,优先通过seqused_ori_kv中的元素获取seqlen,seqused_ori_kv未传入则通过cu_seqlens_ori_kv中的元素获取seqlen。 + - layout_kv为PA_BBND时,优先通过seqused_ori_kv中的元素获取seqlen,seqused_ori_kv未传入则通过ori_topk_length获取seqlen。 + - cmp_kv Seqlen取值规则 + - layout_kv为BSND时,优先通过seqused_cmp_kv中的元素获取seqlen,seqused_cmp_kv未传入则通过max_seqlen_cmp_kv获取seqlen。 + - layout_kv为TND时,优先通过seqused_cmp_kv中的元素获取seqlen,seqused_cmp_kv未传入则通过cu_seqlens_cmp_kv中的元素获取seqlen。 + - layout_kv为PA_BBND时,优先通过seqused_cmp_kv中的元素获取seqlen,seqused_cmp_kv未传入则通过cmp_topk_length获取seqlen。 +- Atlas A3 训练系列产品/Atlas A3 推理系列产品约束: + - SWA稀疏ori_kv场景下,仅支持SWA模板,`has_ori_kv`为true、`has_cmp_kv`为false、`ori_topk`大于0、`ori_mask_mode`为0,`ori_win_left`和`ori_win_right`为非负数,且必须传入`ori_topk_length`。`ori_topk`应与配套主算子`ori_sparse_indices`最后一维K保持一致;`ori_topk_length`表示每个q token和KV head的左对齐有效索引条目数,取值应在[0, K]范围内;Metadata仅使用`ori_topk_length`生成任务切分。配套主算子在PA_BBND场景仍要求传入`seqused_ori_kv`。 + - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时传0且不参与压缩KV计算。CSA场景传1、2或4,HCA场景传128。 + - `cmp_topk`在CSA场景支持512或1024,SWA、HCA场景传0。 +- Atlas A2 训练系列产品/Atlas A2 推理系列产品约束: + - SWA稀疏ori_kv场景下,仅支持SWA模板,`has_ori_kv`为true、`has_cmp_kv`为false、`ori_topk`大于0、`ori_mask_mode`为0,`ori_win_left`和`ori_win_right`为非负数,且必须传入`ori_topk_length`。`ori_topk`应与配套主算子`ori_sparse_indices`最后一维K保持一致;`ori_topk_length`表示每个q token和KV head的左对齐有效索引条目数,取值应在[0, K]范围内;Metadata仅使用`ori_topk_length`生成任务切分。配套主算子在PA_BBND场景仍要求传入`seqused_ori_kv`。 + - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时传0且不参与压缩KV计算。CSA场景传1、2或4,HCA场景传128。 + - `cmp_topk`在CSA场景支持512或1024,SWA、HCA场景传0。 + +## 调用说明 + +| 调用方式 | 样例代码 | 说明 | +| --------- | ------------------------------------------------------------ | ------------------------------------------------------------ | +| aclnn API | [test_aclnn_sparse_flash_mla_metadata](./examples/test_aclnn_sparse_flash_mla_metadata.cpp) | 通过[aclnnSparseFlashMlaMetadata](./docs/aclnnSparseFlashMlaMetadata.md)调用SparseFlashMlaMetadata算子 | +| PyTorch API | [test_torch_sparse_flash_mla_metadata](./examples/test_torch_sparse_flash_mla_metadata.py) | 通过[sparse_flash_mla_metadata](../../torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md)接口生成SparseFlashMla主算子使用的metadata | diff --git a/csrc/attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md b/csrc/attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md new file mode 100644 index 000000000000..d7af8d9410ce --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md @@ -0,0 +1,1103 @@ +# aclnnSparseFlashMlaMetadata + +## 产品支持情况 + + +- Ascend 950PR/Ascend 950DT:支持 + + +- Atlas A3 训练系列产品/Atlas A3 推理系列产品:支持 + + +- Atlas A2 训练系列产品/Atlas A2 推理系列产品:支持 + + +- Atlas 200I/500 A2 推理产品:不支持 + + +- Atlas 推理系列产品:不支持 + + +- Atlas 训练系列产品:不支持 + + +## 功能说明 + +- 算子功能:`aclnnSparseFlashMlaMetadata`是`aclnnSparseFlashMla`算子的前置算子,用于后续Attention计算生成负载均衡的任务划分方案。本算子不执行实际的Attention计算,而是根据输入参数在AI CPU计算出每个AI Core应处理的Attention计算起止范围,从而最大化计算资源的利用率,避免各Core间负载不均衡的问题。 + + **该算子不建议单独使用,建议与aclnnSparseFlashMla算子配合使用,形成完整的工作流。** +- 场景简称:SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)。 +- 计算公式: + + 该算子为AICPU调度算子,不涉及数值计算。核心流程为:解析各Batch的Q/KV序列长度 → 根据mask模式计算每个S1G块的有效S2范围 → 基于开销模型进行负载均衡分核 → 输出分核元数据。 + + 输出metadata tensor的shape为(1024,),数据类型为INT32,内部结构如下: + + - FA Metadata区域(AIC_CORE_NUM × 9个INT32),每个AICore的FA阶段任务信息: + + | 索引 | 含义 | + | :--- | :--- | + | 0 | core_enable,该核是否启用 | + | 1 | bn2_start,BN2起始索引 | + | 2 | m_start,M(S1G)起始索引 | + | 3 | s2_start,S2起始索引 | + | 4 | bn2_end,BN2结束索引 | + | 5 | m_end,M结束索引 | + | 6 | s2_end,S2结束索引 | + | 7 | first_fd_data_workspace_idx,第一份FD归约数据的workspace偏移 | + | 8 | max_s2_block_num,单核上分配到的最多的s2 block数 | + + - FD Metadata区域(AIV_CORE_NUM × 8个INT32),每个AIVCore的FD归约任务信息: + + | 索引 | 含义 | + | :--- | :--- | + | 0 | core_enable,该核是否启用 | + | 1 | bn2_idx,归约任务的BN2索引 | + | 2 | m_idx,归约任务的M索引 | + | 3 | workspace_idx,归约数据在workspace中的存放位置 | + | 4 | workspace_num,S2核间切分份数 | + | 5 | m_start,M轴起点 | + | 6 | m_num,M轴行数 | + +- 符号说明 + + | 符号 | 含义 | + | ------------------- | --------------------------------------------------------- | + | B | Batch Size | + | N1/N2 | Query/KV头数 | + | D | 每个注意力头的维度 | + | G | GQA分组比,G=N1/N2 | + | S1/S2 | Query/KV序列长度 | + | S1G | S1×G方向的分块索引 | + | mBaseSize | M轴基本块大小,等于G | + | s2BaseSize | S2轴基本块大小,固定为512 | + +## 函数原型 + +每个算子分为[两段式接口](../../../docs/zh/context/two_phase_api.md),必须先调用`aclnnSparseFlashMlaMetadataGetWorkspaceSize`接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用`aclnnSparseFlashMlaMetadata`执行实际计算。 + +```c++ +aclnnStatus aclnnSparseFlashMlaMetadataGetWorkspaceSize( + const aclTensor *cuSeqlensQOptional, + const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, + const aclTensor *sequsedQOptional, + const aclTensor *sequsedOriKvOptional, + const aclTensor *sequsedCmpKvOptional, + const aclTensor *cmpResidualKvOptional, + const aclTensor *oriTopkLengthOptional, + const aclTensor *cmpTopkLengthOptional, + int64_t numHeadsQ, + int64_t numHeadsKv, + int64_t headDim, + int64_t batchSize, + int64_t maxSeqlenQ, + int64_t maxSeqlenOriKv, + int64_t maxSeqlenCmpKv, + int64_t oriTopk, + int64_t cmpTopk, + int64_t cmpRatio, + int64_t oriMaskMode, + int64_t cmpMaskMode, + int64_t oriWinLeft, + int64_t oriWinRight, + const char *layoutQOptional, + const char *layoutKvOptional, + bool hasOriKv, + bool hasCmpKv, + const aclTensor *metaData, + uint64_t *workspaceSize, + aclOpExecutor **executor) +``` + +```c++ +aclnnStatus aclnnSparseFlashMlaMetadata( + void *workspace, + uint64_t workspaceSize, + aclOpExecutor *executor, + aclrtStream stream) +``` + +## aclnnSparseFlashMlaMetadataGetWorkspaceSize + +- **参数说明** + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensor
cuSeqlensQOptional(aclTensor*)输入表示不同Batch中q的有效token数(前缀和形式)。
  • 支持空Tensor。
  • shape固定为(B+1, )。
INT32ND1维√
cuSeqlensOriKvOptional(aclTensor*)输入表示不同Batch中oriKv的有效token数(前缀和形式)。
  • 支持空Tensor。
  • shape固定为(B+1, )。
INT32ND1维√
cuSeqlensCmpKvOptional(aclTensor*)输入表示不同Batch中cmpKv的有效token数(前缀和形式)。
  • 支持空Tensor。
  • shape固定为(B+1, )。
INT32ND1维√
sequsedQOptional(aclTensor*)输入表示不同Batch中q实际参与运算的token数。
  • 支持空Tensor。
  • shape固定为(B, )。
INT32ND(B,)√
sequsedOriKvOptional(aclTensor*)输入表示不同Batch中oriKv实际参与运算的token数。
  • 支持空Tensor。
  • shape固定为(B, )。
INT32ND1维√
sequsedCmpKvOptional(aclTensor*)输入表示不同Batch中cmpKv实际参与运算的token数。
  • 支持空Tensor。
  • shape固定为(B, )。
INT32ND1维√
cmpResidualKvOptional(aclTensor*)输入压缩KV余数,用于按cmp_len * cmpRatio + residual恢复cmp侧mask使用的压缩前KV长度。在CSA、HCA、cmpRatio不等于1且cmpMaskMode为3场景必传,layoutKvOptional为BSND、TND、PA_BBND时均可使用。INT32ND1维√
oriTopkLengthOptional(aclTensor*)输入表示不同q token对应的oriKvOptional部分关键稀疏token的个数。
  • SWA稀疏ori_kv场景必须传入,其他场景支持空Tensor。
  • shape为(B, S1, N2)或(T1, N2)。
INT32ND2维、3维√
cmpTopkLengthOptional(aclTensor*)输入表示不同q token对应的cmpKvOptional部分关键稀疏token的个数。
  • 支持空Tensor。
  • shape为(B, S1, N2)或(T1, N2)。
INT32ND2维、3维√
numHeadsQ(int64_t)输入Query的多头数。支持[1, 128]。----
numHeadsKv(int64_t)输入KV的多头数。仅支持1。----
headDim(int64_t)输入每个注意力头的维度。仅支持512。----
batchSize(int64_t)输入输入样本批量大小。layoutQOptional为TND时无需手动指定,建议值为0。----
maxSeqlenQ(int64_t)输入所有Batch中q的最大有效token数。传入0时表示由接口推导,建议值为0。----
maxSeqlenOriKv(int64_t)输入所有Batch中oriKv的最大有效token数。传入0时表示由接口推导,建议值为0。----
maxSeqlenCmpKv(int64_t)输入所有Batch中cmpKv的最大有效token数。传入0时表示由接口推导,建议值为0。----
oriTopk(int64_t)输入从oriKv中筛选的稀疏token个数。SWA稀疏ori_kv场景为主算子oriSparseIndicesOptional最后一维K,且必须大于0;其他场景建议值为0。----
cmpTopk(int64_t)输入从cmpKv中筛选的稀疏token个数。CSA场景下仅支持512或1024,SWA、HCA场景下为0,建议值为0。----
cmpRatio(int64_t)输入cmpKv相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度。传入cmpKv时支持[1, 128];仅传入oriKv时传0;CSA场景传1、2或4,HCA场景传128,建议值为0。----
oriMaskMode(int64_t)输入q和oriKv计算的mask模式。0: No Mask。
3: RightDownCausal模式。
4: Band模式。
建议值为0。
----
cmpMaskMode(int64_t)输入q和cmpKv计算的mask模式。0: No Mask。
3: RightDownCausal模式。
建议值为0。
----
oriWinLeft(int64_t)输入滑动窗口向左扩展的token数。支持-1或非负数,其中-1表示窗口不受限,建议值为-1。----
oriWinRight(int64_t)输入滑动窗口向右扩展的token数。支持-1或非负数,其中-1表示窗口不受限,建议值为-1。----
layoutQOptional(char*)输入标识输入q的数据排布格式。支持"BSND"和"TND",建议值为"BSND"。----
layoutKvOptional(char*)输入标识输入KV的数据排布格式。支持"PA_BBND"、"BSND"和"TND",建议值为"BSND"。----
hasOriKv(bool)输入是否传入oriKv。根据是否传入oriKv设置,建议值为true。----
hasCmpKv(bool)输入是否传入cmpKv。SWA场景为false,CSA、HCA场景为true。根据是否传入cmpKv设置,建议值为true。----
metaData(aclTensor*)输出分核元数据输出,供SparseFlashMla算子使用。shape固定为(1024, )。INT32ND1维×
workspaceSize(uint64_t*)输出返回需要在Device侧申请的workspace大小。-----
executor(aclOpExecutor**)输出返回op执行器,包含了算子计算流程。-----
+ +
    + +
  • Atlas A3 训练系列产品/Atlas A3 推理系列产品 :不支持sequsedQOptional、cmpTopkLengthOptional,numHeadsQ/numHeadsKv仅支持1、2、4、8、16、32、64、128;SWA稀疏ori_kv场景支持oriTopkLengthOptional、oriTopk大于0及oriMaskMode为0,oriWinLeft和oriWinRight支持非负数;其他SWA场景oriTopk为0、oriMaskMode为4、oriWinLeft为127、oriWinRight为0;cmpTopk仅支持0、512、1024,cmpMaskMode仅支持3,cmpRatio在SWA支持0、CSA支持1、2或4、HCA支持128。
  • + + +
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品 :不支持sequsedQOptional、cmpTopkLengthOptional,numHeadsQ/numHeadsKv仅支持1、2、4、8、16、32、64、128;SWA稀疏ori_kv场景支持oriTopkLengthOptional、oriTopk大于0及oriMaskMode为0,oriWinLeft和oriWinRight支持非负数;其他SWA场景oriTopk为0、oriMaskMode为4、oriWinLeft为127、oriWinRight为0;cmpTopk仅支持0、512、1024,cmpMaskMode仅支持3,cmpRatio在SWA支持0、CSA支持1、2或4、HCA支持128。
  • + +
+ +- **返回值** + + aclnnStatus:返回状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。 + + 第一段接口完成入参校验,出现以下场景时报错: + + + - Ascend 950PR/Ascend 950DT: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
返回值错误码描述
ACLNN_ERR_INNER_CREATE_EXECUTOR561101创建aclOpExecutor失败。
ACLNN_ERR_INNER_NULLPTR561103workspaceSize或executor为空指针;可选输入做连续化处理后为空指针;或添加SparseFlashMlaMetadata AICPU任务失败。
ACLNN_ERR_PARAM_INVALID161002batchSize或maxSeqlenQ为负数。
numHeadsQ不在[1,128]范围内,numHeadsKv不为1,numHeadsQ不能被numHeadsKv整除,或numHeadsQ/numHeadsKv不在[1,128]范围内。
headDim不为512。
oriMaskMode不为0、3、4,或cmpMaskMode不为0、3。
oriWinLeft或oriWinRight小于-1。
hasCmpKv为true时,cmpRatio不在[1,128]范围内。
hasCmpKv为true时,cmpTopk为负数。
layoutQOptional、layoutKvOptional、cuSeqlens、seqused或metaData的shape、数据类型、必选关系不在支持范围内。
+ + + - Atlas A3 训练系列产品/Atlas A3 推理系列产品: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
返回值错误码描述
ACLNN_ERR_INNER_CREATE_EXECUTOR561101创建aclOpExecutor失败。
ACLNN_ERR_INNER_NULLPTR561103workspaceSize或executor为空指针;可选输入做连续化处理后为空指针;或添加SparseFlashMlaMetadata AICPU任务失败。
ACLNN_ERR_PARAM_INVALID161002batchSize或maxSeqlenQ为负数。
numHeadsQ不在[1,128]范围内,numHeadsKv不为1,numHeadsQ不能被numHeadsKv整除,或numHeadsQ/numHeadsKv不是[1,128]范围内的2的幂。
headDim不为512。
非SWA稀疏ori_kv场景oriMaskMode不为4,SWA稀疏ori_kv场景oriMaskMode不为0,或cmpMaskMode不为3。
非SWA稀疏ori_kv场景oriWinLeft不为127,或oriWinRight不为0;SWA稀疏ori_kv场景oriWinLeft或oriWinRight为负数。
SWA场景cmpRatio不为0,或cmpRatio与CSA、HCA场景不匹配。
cmpTopk不为0、512或1024。
SWA稀疏ori_kv场景未传入oriTopkLengthOptional,或oriTopkLengthOptional的shape、数据类型不符合规格;cmpTopkLengthOptional传入非空Tensor。
layoutQOptional、layoutKvOptional、cuSeqlens、seqused或metaData的shape、数据类型、必选关系不在支持范围内。
+ + + - Atlas A2 训练系列产品/Atlas A2 推理系列产品: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
返回值错误码描述
ACLNN_ERR_INNER_CREATE_EXECUTOR561101创建aclOpExecutor失败。
ACLNN_ERR_INNER_NULLPTR561103workspaceSize或executor为空指针;可选输入做连续化处理后为空指针;或添加SparseFlashMlaMetadata AICPU任务失败。
ACLNN_ERR_PARAM_INVALID161002batchSize或maxSeqlenQ为负数。
numHeadsQ不在[1,128]范围内,numHeadsKv不为1,numHeadsQ不能被numHeadsKv整除,或numHeadsQ/numHeadsKv不是[1,128]范围内的2的幂。
headDim不为512。
非SWA稀疏ori_kv场景oriMaskMode不为4,SWA稀疏ori_kv场景oriMaskMode不为0,或cmpMaskMode不为3。
非SWA稀疏ori_kv场景oriWinLeft不为127,或oriWinRight不为0;SWA稀疏ori_kv场景oriWinLeft或oriWinRight为负数。
SWA场景cmpRatio不为0,或cmpRatio与CSA、HCA场景不匹配。
cmpTopk不为0、512或1024。
SWA稀疏ori_kv场景未传入oriTopkLengthOptional,或oriTopkLengthOptional的shape、数据类型不符合规格;cmpTopkLengthOptional传入非空Tensor。
layoutQOptional、layoutKvOptional、cuSeqlens、seqused或metaData的shape、数据类型、必选关系不在支持范围内。
+ + +## aclnnSparseFlashMlaMetadata + +- **参数说明** + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
参数名输入/输出描述
workspace输入在Device侧申请的workspace内存地址。
workspaceSize输入在Device侧申请的workspace大小,由第一段接口aclnnSparseFlashMlaMetadataGetWorkspaceSize获取。
executor输入op执行器,包含了算子计算流程。
stream输入指定执行任务的Stream。
+ +- **返回值** + + 返回aclnnStatus状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。 + +## 约束说明 + +- 确定性计算 + + - aclnnSparseFlashMlaMetadata默认采用确定性实现,相同输入多次调用结果一致。 + +- 通用规格约束 + - B(Batch)表示输入样本批量大小,q、oriKvOptional、cmpKvOptional为配套的aclnnSparseFlashMla算子的入参,S1表示layoutQOptional=BSND时,q shape中的S轴的大小,T1表示layoutQOptional=TND时,q shape中的T轴的大小,S2表示layoutKvOptional=BSND时,oriKvOptional shape中的S轴的大小,S3表示layoutKvOptional=BSND时,cmpKvOptional shape中的S轴的大小,N2表示oriKvOptional、cmpKvOptional shape中的N轴的大小。 + - 参数cuSeqlensQOptional、cuSeqlensOriKvOptional、cuSeqlensCmpKvOptional要求其值为当前Batch与前序Batch有效token数的累加值,第一个元素固定为0,后一个元素的值必须大于等于前一个元素的值。 + - 参数sequsedQOptional、sequsedOriKvOptional、sequsedCmpKvOptional要求其值表示每个Batch中的有效token数。 + - layoutQOptional和layoutKvOptional组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下layoutQOptional和layoutKvOptional必须一致。 + - 参数cmpResidualKvOptional需满足cmpResidualKvOptional[i] < cmpRatio。 + +- Ascend 950PR/Ascend 950DT约束: + - hasOriKv为true时,oriTopk大于0认为oriKvOptional部分是稀疏的,oriTopk为0则认为oriKvOptional部分是非稀疏的。 + - hasCmpKv为true时,cmpTopk大于0认为cmpKvOptional部分是稀疏的,cmpTopk为0则认为cmpKvOptional部分是非稀疏的。 + - hasOriKv为true,oriTopk不为0且oriMaskMode为0时,oriTopkLengthOptional必须传入,此时取oriMaskMode规则与oriTopkLengthOptional元素的最小值作为当前q token对应的oriKvOptional的有效seqlen,其他oriKvOptional稀疏场景取oriMaskMode规则与oriTopk的最小值作为当前q token对应的oriKvOptional的有效seqlen。 + - hasCmpKv为true,cmpTopk不为0且cmpMaskMode为0时,cmpTopkLengthOptional必须传入,此时取cmpMaskMode规则与cmpTopkLengthOptional元素的最小值作为当前q token对应的cmpKvOptional的有效seqlen,其他cmpKvOptional稀疏场景取cmpMaskMode规则与cmpTopk的最小值作为当前q token对应的cmpKvOptional的有效seqlen。 + - layoutQOptional=BSND场景 + - maxSeqlenQ必须传入S1的值。 + - layoutKvOptional=BSND场景 + - hasOriKv为true时,maxSeqlenOriKv必须传入S2的值。 + - hasCmpKv为true时,maxSeqlenCmpKv必须传入S3的值。 + - layoutQOptional=TND场景 + - cuSeqlensQOptional必须传入。 + - layoutKvOptional=TND场景 + - hasOriKv为true时,cuSeqlensOriKvOptional必须传入。 + - hasCmpKv为true时,cuSeqlensCmpKvOptional必须传入。 + - layoutKvOptional=PA_BBND场景 + - hasOriKv为true,oriTopk不为0且oriMaskMode为0时(oriTopkLengthOptional必传场景),sequsedOriKvOptional可选传入,其他场景sequsedOriKvOptional必须传入。 + - hasCmpKv为true,cmpTopk不为0且cmpMaskMode为0时(cmpTopkLengthOptional必传场景),sequsedCmpKvOptional可选传入,其他场景sequsedCmpKvOptional必须传入。 + - Batch取值规则 + - layoutQOptional为BSND时,优先通过sequsedQOptional的shape推导batch,sequsedQOptional未传入则通过batch_size获取batch数。 + - layoutQOptional为TND时,优先通过sequsedQOptional的shape推导batch,sequsedQOptional未传入则通过cuSeqlensQOptional的shape推导batch。 + - q Seqlen取值规则 + - layoutQOptional为BSND时,优先通过sequsedQOptional中的元素获取seqlen,sequsedQOptional未传入则通过maxSeqlenQ获取seqlen。 + - layoutQOptional为TND时,优先通过sequsedQOptional中的元素获取seqlen,sequsedQOptional未传入则通过cuSeqlensQOptional中的元素获取seqlen。 + - oriKvOptional Seqlen取值规则 + - layoutKvOptional为BSND时,优先通过sequsedOriKvOptional中的元素获取seqlen,sequsedOriKvOptional未传入则通过maxSeqlenOriKv获取seqlen。 + - layoutKvOptional为TND时,优先通过sequsedOriKvOptional中的元素获取seqlen,sequsedOriKvOptional未传入则通过cuSeqlensOriKvOptional中的元素获取seqlen。 + - layoutKvOptional为PA_BBND时,优先通过sequsedOriKvOptional中的元素获取seqlen,sequsedOriKvOptional未传入则通过oriTopkLengthOptional获取seqlen。 + - cmpKvOptional Seqlen取值规则 + - layoutKvOptional为BSND时,优先通过sequsedCmpKvOptional中的元素获取seqlen,sequsedCmpKvOptional未传入则通过maxSeqlenCmpKv获取seqlen。 + - layoutKvOptional为TND时,优先通过sequsedCmpKvOptional中的元素获取seqlen,sequsedCmpKvOptional未传入则通过cuSeqlensCmpKvOptional中的元素获取seqlen。 + - layoutKvOptional为PA_BBND时,优先通过sequsedCmpKvOptional中的元素获取seqlen,sequsedCmpKvOptional未传入则通过cmpTopkLengthOptional获取seqlen。 + + +- Atlas A3 训练系列产品/Atlas A3 推理系列产品约束: + - SWA稀疏ori_kv场景下,仅支持SWA模板,`hasOriKv`为true、`hasCmpKv`为false、`oriTopk`大于0、`oriMaskMode`为0,`oriWinLeft`和`oriWinRight`为非负数,且必须传入`oriTopkLengthOptional`。`oriTopk`应与配套主算子oriSparseIndicesOptional最后一维K保持一致;`oriTopkLengthOptional`表示每个q token和KV head的左对齐有效索引条目数,取值应在[0, K]范围内;Metadata仅使用`oriTopkLengthOptional`生成任务切分。配套主算子在PA_BBND场景仍要求传入`sequsedOriKvOptional`。 + - layoutQOptional为TND时,`cuSeqlensQOptional`必须传入。 + - layoutKvOptional为PA_BBND时,`sequsedOriKvOptional`必须传入。BSND场景可选传入`sequsedOriKvOptional`覆盖每个batch的oriKv有效长度;TND场景使用`cuSeqlensOriKvOptional`表达oriKv序列边界。 + - layoutKvOptional为TND时,`cuSeqlensOriKvOptional`必须传入;若hasCmpKv为true,`cuSeqlensCmpKvOptional`也必须传入。 + - `sequsedCmpKvOptional`为所有layoutKvOptional下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度。 + - `cmpResidualKvOptional`为`aclnnSparseFlashMlaMetadata`和`aclnnSparseFlashMla`的可选输入,在CSA、HCA、cmpRatio不等于1且cmpMaskMode为3场景必传,用于恢复cmp侧mask使用的压缩前长度。 + + +- Atlas A2 训练系列产品/Atlas A2 推理系列产品约束: + - SWA稀疏ori_kv场景下,仅支持SWA模板,`hasOriKv`为true、`hasCmpKv`为false、`oriTopk`大于0、`oriMaskMode`为0,`oriWinLeft`和`oriWinRight`为非负数,且必须传入`oriTopkLengthOptional`。`oriTopk`应与配套主算子oriSparseIndicesOptional最后一维K保持一致;`oriTopkLengthOptional`表示每个q token和KV head的左对齐有效索引条目数,取值应在[0, K]范围内;Metadata仅使用`oriTopkLengthOptional`生成任务切分。配套主算子在PA_BBND场景仍要求传入`sequsedOriKvOptional`。 + - layoutQOptional为TND时,`cuSeqlensQOptional`必须传入。 + - layoutKvOptional为PA_BBND时,`sequsedOriKvOptional`必须传入。BSND场景可选传入`sequsedOriKvOptional`覆盖每个batch的oriKv有效长度;TND场景使用`cuSeqlensOriKvOptional`表达oriKv序列边界。 + - layoutKvOptional为TND时,`cuSeqlensOriKvOptional`必须传入;若hasCmpKv为true,`cuSeqlensCmpKvOptional`也必须传入。 + - `sequsedCmpKvOptional`为所有layoutKvOptional下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度。 + - `cmpResidualKvOptional`为`aclnnSparseFlashMlaMetadata`和`aclnnSparseFlashMla`的可选输入,在CSA、HCA、cmpRatio不等于1且cmpMaskMode为3场景必传,用于恢复cmp侧mask使用的压缩前长度。 + + +## 调用示例 + +调用示例代码如下,仅供参考,具体编译和执行过程请参考[编译与运行样例](../../../docs/zh/context/compile_and_run_sample.md)。 + +```c++ +/** + * @file test_aclnn_sparse_flash_mla_metadata.cpp + */ +#include +#include +#include +#include +#include +#include +#include +#include "acl/acl.h" +#include "aclnnop/aclnn_sparse_flash_mla_metadata.h" + +#define CHECK_LOG_RET(cond, ret_val, fmt, ...) \ + do { \ + if (!(cond)) { \ + printf(fmt "\n", ##__VA_ARGS__); \ + return (ret_val); \ + } \ + } while (0) + +// 参考 sparse_flash_mla_metadata.h +constexpr uint32_t AIC_CORE_MAX_NUM = 36; +constexpr uint32_t AIV_CORE_MAX_NUM = 72; +constexpr uint32_t SMLA_METADATA_TOTAL_SIZE = 1024; +constexpr uint32_t FA_METADATA_SIZE = 9; +constexpr uint32_t FD_METADATA_SIZE = 8; + +// FA Metadata Index Definitions +constexpr uint32_t FA_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FA_BN2_START_INDEX = 1; +constexpr uint32_t FA_M_START_INDEX = 2; +constexpr uint32_t FA_S2_START_INDEX = 3; +constexpr uint32_t FA_BN2_END_INDEX = 4; +constexpr uint32_t FA_M_END_INDEX = 5; +constexpr uint32_t FA_S2_END_INDEX = 6; +constexpr uint32_t FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX = 7; +constexpr uint32_t FA_S2_MAX_NUM = 8; + +// FD Metadata Index Definitions +constexpr uint32_t FD_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FD_BN2_IDX_INDEX = 1; +constexpr uint32_t FD_M_IDX_INDEX = 2; +constexpr uint32_t FD_WORKSPACE_IDX_INDEX = 3; +constexpr uint32_t FD_WORKSPACE_NUM_INDEX = 4; +constexpr uint32_t FD_M_START_INDEX = 5; +constexpr uint32_t FD_M_NUM_INDEX = 6; + +struct SmlaMetadata { + uint32_t faMetadata[AIC_CORE_MAX_NUM][FA_METADATA_SIZE]; + uint32_t fdMetadata[AIV_CORE_MAX_NUM][FD_METADATA_SIZE]; +}; + +struct ScopeGuard +{ + explicit ScopeGuard(std::function onExitScope) : m_exitFunc(std::move(onExitScope)), + m_isDismissed(false) {} + // 禁止拷贝 + ScopeGuard(const ScopeGuard&) = delete; + ScopeGuard& operator=(const ScopeGuard&) = delete; + + ~ScopeGuard() + { + if (!m_isDismissed) { + m_exitFunc(); + } + } + + void Dismiss() + { + m_isDismissed = true; + } + + std::function m_exitFunc; + bool m_isDismissed; +}; + +struct Tensor { + void *hostAddr { nullptr }; + void *deviceAddr { nullptr }; + aclTensor *data { nullptr }; +}; + +struct ArgScenario { + bool hasCuSeq { false }; + bool hasSeqused { false }; +}; + +struct ArgContext { + // required input + int64_t numHeadsQ { 0 }; + int64_t numHeadsKv { 0 }; + int64_t headDim { 0 }; + // optional input + Tensor cuSeqlensQOptional {}; + Tensor cuSeqlensOriKvOptional {}; + Tensor cuSeqlensCmpKvOptional {}; + Tensor sequsedQOptional {}; + Tensor sequsedOriKvOptional {}; + Tensor sequsedCmpKvOptional {}; + Tensor cmpResidualKvOptional {}; + Tensor oriTopkLengthOptional {}; + Tensor cmpTopkLengthOptional {}; + int64_t batchSize { 0 }; + int64_t maxSeqlenQ { 0 }; + int64_t maxSeqlenOriKv { 0 }; + int64_t maxSeqlenCmpKv { 0 }; + int64_t oriTopk { 0 }; + int64_t cmpTopk { 0 }; + int64_t cmpRatio { 0 }; + int64_t oriMaskMode { 0 }; + int64_t cmpMaskMode { 0 }; + int64_t oriWinLeft { -1 }; + int64_t oriWinRight { -1 }; + char *layoutQOptional { nullptr }; + char *layoutKvOptional { nullptr }; + bool hasOriKv { true }; + bool hasCmpKv { true }; + // output + Tensor metadata {}; +}; + +int64_t GetShapeSize(const std::vector& shape) +{ + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} + +aclnnStatus Init(int32_t deviceId, aclrtStream* stream) +{ + // 固定写法,初始化 + auto ret = aclInit(nullptr); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclInit failed. ERROR: %d", ret); + ret = aclrtSetDevice(deviceId); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSetDevice failed. ERROR: %d", ret); + ret = aclrtCreateStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtCreateStream failed. ERROR: %d", ret); + return ACL_SUCCESS; +} + +void Finalize(int32_t deviceId, aclrtStream stream) +{ + aclrtDestroyStream(stream); + aclrtResetDevice(deviceId); + aclFinalize(); +} + +aclnnStatus CreateTensor(aclDataType dataType, const std::vector &shape, Tensor &tensor) +{ + auto size = GetShapeSize(shape) * aclDataTypeSize(dataType); + // 调用aclrtMallocHost申请host侧内存 + auto ret = aclrtMallocHost(&(tensor.hostAddr), size); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMallocHost failed. ERROR: %d", ret); + memset(tensor.hostAddr, 0, size); + // 调用aclrtMalloc申请device侧内存 + ret = aclrtMalloc(&(tensor.deviceAddr), size, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMalloc failed. ERROR: %d", ret); + // 调用aclCreateTensor接口创建aclTensor + tensor.data = aclCreateTensor(shape.data(), shape.size(), dataType, nullptr, 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), tensor.deviceAddr); + CHECK_LOG_RET(tensor.data != nullptr, ACL_ERROR_FAILURE, "aclCreateTensor failed"); + // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 + ret = aclrtMemcpy(tensor.deviceAddr, size, tensor.hostAddr, size, ACL_MEMCPY_HOST_TO_DEVICE); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d", ret); + return ACL_SUCCESS; +} + +void DestroyTensor(Tensor &tensor) +{ + if (tensor.data != nullptr) { + aclDestroyTensor(tensor.data); + tensor.data = nullptr; + } + if (tensor.deviceAddr != nullptr) { + aclrtFree(tensor.deviceAddr); + tensor.deviceAddr = nullptr; + } + if (tensor.hostAddr != nullptr) { + aclrtFreeHost(tensor.hostAddr); + tensor.hostAddr = nullptr; + } +} + +void DestroyArgs(ArgContext &context) +{ + DestroyTensor(context.metadata); + DestroyTensor(context.cuSeqlensQOptional); + DestroyTensor(context.cuSeqlensOriKvOptional); + DestroyTensor(context.cuSeqlensCmpKvOptional); + DestroyTensor(context.sequsedQOptional); + DestroyTensor(context.sequsedOriKvOptional); + DestroyTensor(context.sequsedCmpKvOptional); + DestroyTensor(context.cmpResidualKvOptional); + DestroyTensor(context.oriTopkLengthOptional); + DestroyTensor(context.cmpTopkLengthOptional); + + if (context.layoutQOptional != nullptr) { + free(context.layoutQOptional); + context.layoutQOptional = nullptr; + } + if (context.layoutKvOptional != nullptr) { + free(context.layoutKvOptional); + context.layoutKvOptional = nullptr; + } +} + +aclnnStatus CreateArgs(const ArgScenario &scenario, ArgContext &context) +{ + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + aclnnStatus ret; + + context.numHeadsQ = 64; + context.numHeadsKv = 1; + context.headDim = 512; + ret = CreateTensor(aclDataType::ACL_INT32, { SMLA_METADATA_TOTAL_SIZE }, context.metadata); // 1024: Fix size + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create metadata failed. Error: %d", ret); + context.oriTopk = 0; + context.cmpTopk = 0; + context.cmpRatio = 128; + context.oriMaskMode = 4; + context.cmpMaskMode = 3; + context.oriWinLeft = 127; + context.oriWinRight = 0; + context.layoutQOptional = (char *)malloc(sizeof(char) * 16); + context.layoutKvOptional = (char *)malloc(sizeof(char) * 16); + CHECK_LOG_RET(context.layoutQOptional != nullptr, ACL_ERROR_FAILURE, "Create layoutQOptional failed"); + CHECK_LOG_RET(context.layoutKvOptional != nullptr, ACL_ERROR_FAILURE, "Create layoutKvOptional failed"); + strcpy(context.layoutQOptional, "BSND"); // BSND,TND + strcpy(context.layoutKvOptional, "BSND"); // BSND,TND,PA_BBND + context.hasOriKv = true; + context.hasCmpKv = true; + + context.batchSize = 4; + context.maxSeqlenOriKv = 1024; + context.maxSeqlenCmpKv = 1024; + context.maxSeqlenQ = 1024; + + if (scenario.hasCuSeq) { + // (B+1,), first element is always 0 + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensOriKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensOriKvOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensCmpKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensCmpKvOptional failed. Error: %d", ret); + } + + if (scenario.hasSeqused) { + // (B,) + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedOriKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedOriKvOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedCmpKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedCmpKvOptional failed. Error: %d", ret); + } + + if (context.hasCmpKv && context.cmpRatio != 1 && context.cmpMaskMode == 3) { + // (B,) + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.cmpResidualKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cmpResidualKvOptional failed. Error: %d", ret); + } + + argsGuard.Dismiss(); + return ACL_SUCCESS; +} + +int main() { + // 1. (固定写法)device/stream初始化,参考对外接口列表 + // 根据自己的实际device填写deviceId + int32_t deviceId = 0; + aclrtStream stream; + auto ret = Init(deviceId, &stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Init acl failed. ERROR: %d", ret); + ScopeGuard sysGuard([&] { Finalize(deviceId, stream); }); + + // 2. 构造输入与输出,需要根据API的接口定义构造 + ArgScenario scenario {}; + scenario.hasCuSeq = false; + scenario.hasSeqused = false; + ArgContext context {}; + ret = CreateArgs(scenario, context); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create input arguments failed. ERROR: %d", ret); + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + + // 3. 调用CANN算子库API,需要修改为具体的API + // 调用aclnnSparseFlashMlaMetadata第一段接口 + uint64_t workspaceSize = 0; + aclOpExecutor *executor = nullptr; + void *workspaceAddr = nullptr; + ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( + context.cuSeqlensQOptional.data, context.cuSeqlensOriKvOptional.data, context.cuSeqlensCmpKvOptional.data, + context.sequsedQOptional.data, context.sequsedOriKvOptional.data, context.sequsedCmpKvOptional.data, + context.cmpResidualKvOptional.data, context.oriTopkLengthOptional.data, context.cmpTopkLengthOptional.data, + context.numHeadsQ, context.numHeadsKv, context.headDim, context.batchSize, context.maxSeqlenQ, + context.maxSeqlenOriKv, context.maxSeqlenCmpKv, context.oriTopk, context.cmpTopk, context.cmpRatio, + context.oriMaskMode, context.cmpMaskMode, context.oriWinLeft, context.oriWinRight, context.layoutQOptional, + context.layoutKvOptional, context.hasOriKv, context.hasCmpKv, context.metadata.data, &workspaceSize, &executor); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, + "aclnnSparseFlashMlaMetadataGetWorkspaceSize failed. ERROR: %d\n", ret); + + if (workspaceSize > static_cast(0)) { + ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "allocate workspace failed. ERROR: %d\n", ret); + } + ScopeGuard workspaceGuard([&] { + if (workspaceAddr != nullptr) { + aclrtFree(workspaceAddr); + workspaceAddr = nullptr; + } + }); + + // 调用aclnnSparseFlashMlaMetadata第二段接口 + ret = aclnnSparseFlashMlaMetadata(workspaceAddr, workspaceSize, executor, stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclnnSparseFlashMlaMetadata failed. ERROR: %d\n", ret); + + // 4. (固定写法)同步等待任务执行结束 + ret = aclrtSynchronizeStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSynchronizeStream failed. ERROR: %d\n", ret); + + // 5. 打印输出 + SmlaMetadata result {}; + ret = aclrtMemcpy(&result, sizeof(result), context.metadata.deviceAddr, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d\n", ret); + + for (uint32_t i = 0; i < AIC_CORE_MAX_NUM; ++i) { + printf("AIC Core%u\n", i); + printf(" Core Enable : %u\n", result.faMetadata[i][FA_CORE_ENABLE_INDEX]); + printf(" Start BN2 : %u\n", result.faMetadata[i][FA_BN2_START_INDEX]); + printf(" Start M : %u\n", result.faMetadata[i][FA_M_START_INDEX]); + printf(" Start S2 : %u\n", result.faMetadata[i][FA_S2_START_INDEX]); + printf(" End BN2 : %u\n", result.faMetadata[i][FA_BN2_END_INDEX]); + printf(" End M : %u\n", result.faMetadata[i][FA_M_END_INDEX]); + printf(" End S2 : %u\n", result.faMetadata[i][FA_S2_END_INDEX]); + printf(" First Workspace Index : %u\n", result.faMetadata[i][FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX]); + printf(" Max S2 Block Num : %u\n", result.faMetadata[i][FA_S2_MAX_NUM]); + } + for (uint32_t i = 0; i < AIV_CORE_MAX_NUM; ++i) { + printf("AIV Core%u\n", i); + printf(" Core Enable : %u\n", result.fdMetadata[i][FD_CORE_ENABLE_INDEX]); + printf(" FD Task BN2 Idx : %u\n", result.fdMetadata[i][FD_BN2_IDX_INDEX]); + printf(" FD Task M Idx : %u\n", result.fdMetadata[i][FD_M_IDX_INDEX]); + printf(" FD Task Workspace Idx : %u\n", result.fdMetadata[i][FD_WORKSPACE_IDX_INDEX]); + printf(" FD Task Workspace Num : %u\n", result.fdMetadata[i][FD_WORKSPACE_NUM_INDEX]); + printf(" FD Subtask M Start : %u\n", result.fdMetadata[i][FD_M_START_INDEX]); + printf(" FD Subtask M Num : %u\n", result.fdMetadata[i][FD_M_NUM_INDEX]); + } + + return 0; +} +``` diff --git a/csrc/attention/sparse_flash_mla_metadata/examples/test_aclnn_sparse_flash_mla_metadata.cpp b/csrc/attention/sparse_flash_mla_metadata/examples/test_aclnn_sparse_flash_mla_metadata.cpp new file mode 100644 index 000000000000..9a95b9a6504c --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/examples/test_aclnn_sparse_flash_mla_metadata.cpp @@ -0,0 +1,364 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/** + * @file test_aclnn_sparse_flash_mla_metadata.cpp + */ +#include +#include +#include +#include +#include +#include +#include +#include "acl/acl.h" +#include "aclnnop/aclnn_sparse_flash_mla_metadata.h" + +#define CHECK_LOG_RET(cond, ret_val, fmt, ...) \ + do { \ + if (!(cond)) { \ + printf(fmt "\n", ##__VA_ARGS__); \ + return (ret_val); \ + } \ + } while (0) + +// 参考 sparse_flash_mla_metadata.h +constexpr uint32_t AIC_CORE_MAX_NUM = 36; +constexpr uint32_t AIV_CORE_MAX_NUM = 72; +constexpr uint32_t SMLA_METADATA_TOTAL_SIZE = 1024; +constexpr uint32_t FA_METADATA_SIZE = 9; +constexpr uint32_t FD_METADATA_SIZE = 8; + +// FA Metadata Index Definitions +constexpr uint32_t FA_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FA_BN2_START_INDEX = 1; +constexpr uint32_t FA_M_START_INDEX = 2; +constexpr uint32_t FA_S2_START_INDEX = 3; +constexpr uint32_t FA_BN2_END_INDEX = 4; +constexpr uint32_t FA_M_END_INDEX = 5; +constexpr uint32_t FA_S2_END_INDEX = 6; +constexpr uint32_t FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX = 7; +constexpr uint32_t FA_S2_MAX_NUM = 8; + +// FD Metadata Index Definitions +constexpr uint32_t FD_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FD_BN2_IDX_INDEX = 1; +constexpr uint32_t FD_M_IDX_INDEX = 2; +constexpr uint32_t FD_WORKSPACE_IDX_INDEX = 3; +constexpr uint32_t FD_WORKSPACE_NUM_INDEX = 4; +constexpr uint32_t FD_M_START_INDEX = 5; +constexpr uint32_t FD_M_NUM_INDEX = 6; + +struct SmlaMetadata { + uint32_t faMetadata[AIC_CORE_MAX_NUM][FA_METADATA_SIZE]; + uint32_t fdMetadata[AIV_CORE_MAX_NUM][FD_METADATA_SIZE]; +}; + +struct ScopeGuard { + explicit ScopeGuard(std::function onExitScope) + : m_exitFunc(std::move(onExitScope)), + m_isDismissed(false) + {} + // 禁止拷贝 + ScopeGuard(const ScopeGuard &) = delete; + ScopeGuard &operator=(const ScopeGuard &) = delete; + + ~ScopeGuard() + { + if (!m_isDismissed) { + m_exitFunc(); + } + } + + void Dismiss() + { + m_isDismissed = true; + } + + std::function m_exitFunc; + bool m_isDismissed; +}; + +struct Tensor { + void *hostAddr{nullptr}; + void *deviceAddr{nullptr}; + aclTensor *data{nullptr}; +}; + +struct ArgScenario { + bool hasCuSeq{false}; + bool hasSeqused{false}; +}; + +struct ArgContext { + // required input + int64_t numHeadsQ{0}; + int64_t numHeadsKv{0}; + int64_t headDim{0}; + // optional input + Tensor cuSeqlensQOptional{}; + Tensor cuSeqlensOriKvOptional{}; + Tensor cuSeqlensCmpKvOptional{}; + Tensor sequsedQOptional{}; + Tensor sequsedOriKvOptional{}; + Tensor sequsedCmpKvOptional{}; + Tensor cmpResidualKvOptional{}; + Tensor oriTopkLengthOptional{}; + Tensor cmpTopkLengthOptional{}; + int64_t batchSize{0}; + int64_t maxSeqlenQ{0}; + int64_t maxSeqlenOriKv{0}; + int64_t maxSeqlenCmpKv{0}; + int64_t oriTopk{0}; + int64_t cmpTopk{0}; + int64_t cmpRatio{0}; + int64_t oriMaskMode{0}; + int64_t cmpMaskMode{0}; + int64_t oriWinLeft{-1}; + int64_t oriWinRight{-1}; + char *layoutQOptional{nullptr}; + char *layoutKvOptional{nullptr}; + bool hasOriKv{true}; + bool hasCmpKv{true}; + // output + Tensor metadata{}; +}; + +int64_t GetShapeSize(const std::vector &shape) +{ + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} + +aclnnStatus Init(int32_t deviceId, aclrtStream *stream) +{ + // 固定写法,初始化 + auto ret = aclInit(nullptr); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclInit failed. ERROR: %d", ret); + ret = aclrtSetDevice(deviceId); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSetDevice failed. ERROR: %d", ret); + ret = aclrtCreateStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtCreateStream failed. ERROR: %d", ret); + return ACL_SUCCESS; +} + +void Finalize(int32_t deviceId, aclrtStream stream) +{ + aclrtDestroyStream(stream); + aclrtResetDevice(deviceId); + aclFinalize(); +} + +aclnnStatus CreateTensor(aclDataType dataType, const std::vector &shape, Tensor &tensor) +{ + auto size = GetShapeSize(shape) * aclDataTypeSize(dataType); + // 调用aclrtMallocHost申请host侧内存 + auto ret = aclrtMallocHost(&(tensor.hostAddr), size); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMallocHost failed. ERROR: %d", ret); + memset(tensor.hostAddr, 0, size); + // 调用aclrtMalloc申请device侧内存 + ret = aclrtMalloc(&(tensor.deviceAddr), size, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMalloc failed. ERROR: %d", ret); + // 调用aclCreateTensor接口创建aclTensor + tensor.data = aclCreateTensor(shape.data(), shape.size(), dataType, nullptr, 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), tensor.deviceAddr); + CHECK_LOG_RET(tensor.data != nullptr, ACL_ERROR_FAILURE, "aclCreateTensor failed"); + // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 + ret = aclrtMemcpy(tensor.deviceAddr, size, tensor.hostAddr, size, ACL_MEMCPY_HOST_TO_DEVICE); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d", ret); + return ACL_SUCCESS; +} + +void DestroyTensor(Tensor &tensor) +{ + if (tensor.data != nullptr) { + aclDestroyTensor(tensor.data); + tensor.data = nullptr; + } + if (tensor.deviceAddr != nullptr) { + aclrtFree(tensor.deviceAddr); + tensor.deviceAddr = nullptr; + } + if (tensor.hostAddr != nullptr) { + aclrtFreeHost(tensor.hostAddr); + tensor.hostAddr = nullptr; + } +} + +void DestroyArgs(ArgContext &context) +{ + DestroyTensor(context.metadata); + DestroyTensor(context.cuSeqlensQOptional); + DestroyTensor(context.cuSeqlensOriKvOptional); + DestroyTensor(context.cuSeqlensCmpKvOptional); + DestroyTensor(context.sequsedQOptional); + DestroyTensor(context.sequsedOriKvOptional); + DestroyTensor(context.sequsedCmpKvOptional); + DestroyTensor(context.cmpResidualKvOptional); + DestroyTensor(context.oriTopkLengthOptional); + DestroyTensor(context.cmpTopkLengthOptional); + + if (context.layoutQOptional != nullptr) { + free(context.layoutQOptional); + context.layoutQOptional = nullptr; + } + if (context.layoutKvOptional != nullptr) { + free(context.layoutKvOptional); + context.layoutKvOptional = nullptr; + } +} + +aclnnStatus CreateArgs(const ArgScenario &scenario, ArgContext &context) +{ + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + aclnnStatus ret; + + context.numHeadsQ = 64; + context.numHeadsKv = 1; + context.headDim = 512; + ret = CreateTensor(aclDataType::ACL_INT32, {SMLA_METADATA_TOTAL_SIZE}, context.metadata); // 1024: Fix size + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create metadata failed. Error: %d", ret); + context.oriTopk = 0; + context.cmpTopk = 0; + context.cmpRatio = 128; + context.oriMaskMode = 4; + context.cmpMaskMode = 3; + context.oriWinLeft = 127; + context.oriWinRight = 0; + context.layoutQOptional = (char *)malloc(sizeof(char) * 16); + context.layoutKvOptional = (char *)malloc(sizeof(char) * 16); + CHECK_LOG_RET(context.layoutQOptional != nullptr, ACL_ERROR_FAILURE, "Create layoutQOptional failed"); + CHECK_LOG_RET(context.layoutKvOptional != nullptr, ACL_ERROR_FAILURE, "Create layoutKvOptional failed"); + strcpy(context.layoutQOptional, "BSND"); // BSND,TND + strcpy(context.layoutKvOptional, "BSND"); // BSND,TND,PA_BBND + context.hasOriKv = true; + context.hasCmpKv = true; + + context.batchSize = 4; + context.maxSeqlenOriKv = 1024; + context.maxSeqlenCmpKv = 1024; + context.maxSeqlenQ = 1024; + + if (scenario.hasCuSeq) { + // (B+1,), first element is always 0 + ret = CreateTensor(aclDataType::ACL_INT32, {context.batchSize + 1}, context.cuSeqlensQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, {context.batchSize + 1}, context.cuSeqlensOriKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensOriKvOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, {context.batchSize + 1}, context.cuSeqlensCmpKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensCmpKvOptional failed. Error: %d", ret); + } + + if (scenario.hasSeqused) { + // (B,) + ret = CreateTensor(aclDataType::ACL_INT32, {context.batchSize}, context.sequsedQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, {context.batchSize}, context.sequsedOriKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedOriKvOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, {context.batchSize}, context.sequsedCmpKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedCmpKvOptional failed. Error: %d", ret); + } + + if (context.hasCmpKv && context.cmpRatio != 1 && context.cmpMaskMode == 3) { + // (B,) + ret = CreateTensor(aclDataType::ACL_INT32, {context.batchSize}, context.cmpResidualKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cmpResidualKvOptional failed. Error: %d", ret); + } + + argsGuard.Dismiss(); + return ACL_SUCCESS; +} + +int main() +{ + // 1. (固定写法)device/stream初始化,参考对外接口列表 + // 根据自己的实际device填写deviceId + int32_t deviceId = 0; + aclrtStream stream; + auto ret = Init(deviceId, &stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Init acl failed. ERROR: %d", ret); + ScopeGuard sysGuard([&] { Finalize(deviceId, stream); }); + + // 2. 构造输入与输出,需要根据API的接口定义构造 + ArgScenario scenario{}; + scenario.hasCuSeq = false; + scenario.hasSeqused = false; + ArgContext context{}; + ret = CreateArgs(scenario, context); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create input arguments failed. ERROR: %d", ret); + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + + // 3. 调用CANN算子库API,需要修改为具体的API + // 调用aclnnSparseFlashMlaMetadata第一段接口 + uint64_t workspaceSize = 0; + aclOpExecutor *executor = nullptr; + void *workspaceAddr = nullptr; + ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( + context.cuSeqlensQOptional.data, context.cuSeqlensOriKvOptional.data, context.cuSeqlensCmpKvOptional.data, + context.sequsedQOptional.data, context.sequsedOriKvOptional.data, context.sequsedCmpKvOptional.data, + context.cmpResidualKvOptional.data, context.oriTopkLengthOptional.data, context.cmpTopkLengthOptional.data, + context.numHeadsQ, context.numHeadsKv, context.headDim, context.batchSize, context.maxSeqlenQ, + context.maxSeqlenOriKv, context.maxSeqlenCmpKv, context.oriTopk, context.cmpTopk, context.cmpRatio, + context.oriMaskMode, context.cmpMaskMode, context.oriWinLeft, context.oriWinRight, context.layoutQOptional, + context.layoutKvOptional, context.hasOriKv, context.hasCmpKv, context.metadata.data, &workspaceSize, &executor); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclnnSparseFlashMlaMetadataGetWorkspaceSize failed. ERROR: %d\n", ret); + + if (workspaceSize > static_cast(0)) { + ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "allocate workspace failed. ERROR: %d\n", ret); + } + ScopeGuard workspaceGuard([&] { + if (workspaceAddr != nullptr) { + aclrtFree(workspaceAddr); + workspaceAddr = nullptr; + } + }); + + // 调用aclnnSparseFlashMlaMetadata第二段接口 + ret = aclnnSparseFlashMlaMetadata(workspaceAddr, workspaceSize, executor, stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclnnSparseFlashMlaMetadata failed. ERROR: %d\n", ret); + + // 4. (固定写法)同步等待任务执行结束 + ret = aclrtSynchronizeStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSynchronizeStream failed. ERROR: %d\n", ret); + + // 5. 打印输出 + SmlaMetadata result{}; + ret = aclrtMemcpy(&result, sizeof(result), context.metadata.deviceAddr, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d\n", ret); + + for (uint32_t i = 0; i < AIC_CORE_MAX_NUM; ++i) { + printf("AIC Core%u\n", i); + printf(" Core Enable : %u\n", result.faMetadata[i][FA_CORE_ENABLE_INDEX]); + printf(" Start BN2 : %u\n", result.faMetadata[i][FA_BN2_START_INDEX]); + printf(" Start M : %u\n", result.faMetadata[i][FA_M_START_INDEX]); + printf(" Start S2 : %u\n", result.faMetadata[i][FA_S2_START_INDEX]); + printf(" End BN2 : %u\n", result.faMetadata[i][FA_BN2_END_INDEX]); + printf(" End M : %u\n", result.faMetadata[i][FA_M_END_INDEX]); + printf(" End S2 : %u\n", result.faMetadata[i][FA_S2_END_INDEX]); + printf(" First Workspace Index : %u\n", result.faMetadata[i][FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX]); + printf(" Max S2 Block Num : %u\n", result.faMetadata[i][FA_S2_MAX_NUM]); + } + for (uint32_t i = 0; i < AIV_CORE_MAX_NUM; ++i) { + printf("AIV Core%u\n", i); + printf(" Core Enable : %u\n", result.fdMetadata[i][FD_CORE_ENABLE_INDEX]); + printf(" FD Task BN2 Idx : %u\n", result.fdMetadata[i][FD_BN2_IDX_INDEX]); + printf(" FD Task M Idx : %u\n", result.fdMetadata[i][FD_M_IDX_INDEX]); + printf(" FD Task Workspace Idx : %u\n", result.fdMetadata[i][FD_WORKSPACE_IDX_INDEX]); + printf(" FD Task Workspace Num : %u\n", result.fdMetadata[i][FD_WORKSPACE_NUM_INDEX]); + printf(" FD Subtask M Start : %u\n", result.fdMetadata[i][FD_M_START_INDEX]); + printf(" FD Subtask M Num : %u\n", result.fdMetadata[i][FD_M_NUM_INDEX]); + } + + return 0; +} diff --git a/csrc/attention/sparse_flash_mla_metadata/examples/test_torch_sparse_flash_mla_metadata.py b/csrc/attention/sparse_flash_mla_metadata/examples/test_torch_sparse_flash_mla_metadata.py new file mode 100644 index 000000000000..610c5b3dee39 --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/examples/test_torch_sparse_flash_mla_metadata.py @@ -0,0 +1,41 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2025 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +import torch + +metadata = torch.ops.cann_ops_transformer.sparse_flash_mla_metadata( + cu_seqlens_q=torch.tensor([0, 10], dtype=torch.int32).npu(), + cu_seqlens_ori_kv=None, + cu_seqlens_cmp_kv=None, + seqused_q=None, + seqused_ori_kv=torch.tensor([8192], dtype=torch.int32).npu(), + seqused_cmp_kv=torch.tensor([64], dtype=torch.int32).npu(), + cmp_residual_kv=torch.tensor([1], dtype=torch.int32).npu(), + ori_topk_length=None, + cmp_topk_length=None, + num_heads_q=128, + num_heads_kv=1, + head_dim=512, + batch_size=1, + max_seqlen_q=1, + max_seqlen_ori_kv=512, + max_seqlen_cmp_kv=32, + ori_topk=0, + cmp_topk=512, + cmp_ratio=4, + ori_mask_mode=4, + cmp_mask_mode=3, + ori_win_left=127, + ori_win_right=0, + layout_q="TND", + layout_kv="PA_BBND", + has_ori_kv=True, + has_cmp_kv=True, +) diff --git a/csrc/attention/sparse_flash_mla_metadata/op_host/CMakeLists.txt b/csrc/attention/sparse_flash_mla_metadata/op_host/CMakeLists.txt new file mode 100644 index 000000000000..3a8f103e8edb --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_host/CMakeLists.txt @@ -0,0 +1,12 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +add_op_to_compiled_list() + +add_modules_sources(OPTYPE sparse_flash_mla_metadata ACLNNTYPE aclnn) \ No newline at end of file diff --git a/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/aclnn_sparse_flash_mla_metadata.cpp b/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/aclnn_sparse_flash_mla_metadata.cpp new file mode 100644 index 000000000000..92f7a4310623 --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/aclnn_sparse_flash_mla_metadata.cpp @@ -0,0 +1,186 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file aclnn_sparse_flash_mla_metadata.cpp + * \brief + */ + +#include "aclnn_sparse_flash_mla_metadata.h" +#include "../sparse_flash_mla_metadata_check.h" +#include "sparse_flash_mla_metadata.h" +#include "aclnn_kernels/contiguous.h" +#include "aclnn_kernels/reshape.h" +#include "aclnn/aclnn_base.h" +#include "aclnn_kernels/common/op_error_check.h" +#include "opdev/common_types.h" +#include "opdev/data_type_utils.h" +#include "opdev/format_utils.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/tensor_view_utils.h" +#include "opdev/make_op_executor.h" +#include "acl/acl_rt.h" + +constexpr int64_t BATCH_CONSISTENCY_LEVEL = 3; + +#ifdef __cplusplus +extern "C" { +#endif + +aclnnStatus aclnnSparseFlashMlaMetadataGetWorkspaceSize( + const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, const aclTensor *sequsedQOptional, const aclTensor *sequsedOriKvOptional, + const aclTensor *sequsedCmpKvOptional, const aclTensor *cmpResidualKvOptional, + const aclTensor *oriTopkLengthOptional, const aclTensor *cmpTopkLengthOptional, int64_t numHeadsQ, + int64_t numHeadsKv, int64_t headDim, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, + int64_t maxSeqlenCmpKv, int64_t oriTopk, int64_t cmpTopk, int64_t cmpRatio, int64_t oriMaskMode, + int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight, const char *layoutQOptional, + const char *layoutKvOptional, bool hasOriKv, bool hasCmpKv, const aclTensor *metaData, uint64_t *workspaceSize, + aclOpExecutor **executor) +{ + if (workspaceSize == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "workspaceSize is nullptr"); + return ACLNN_ERR_INNER_NULLPTR; + } + if (executor == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "executor is nullptr"); + return ACLNN_ERR_INNER_NULLPTR; + } + L2_DFX_PHASE_1(aclnnSparseFlashMlaMetadata, + DFX_IN(cuSeqlensQOptional, cuSeqlensOriKvOptional, cuSeqlensCmpKvOptional, sequsedQOptional, + sequsedOriKvOptional, sequsedCmpKvOptional, cmpResidualKvOptional, oriTopkLengthOptional, + cmpTopkLengthOptional, numHeadsQ, numHeadsKv, headDim, batchSize, maxSeqlenQ, maxSeqlenOriKv, + maxSeqlenCmpKv, oriTopk, cmpTopk, cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, + layoutQOptional, layoutKvOptional, hasOriKv, hasCmpKv), + DFX_OUT(metaData)); + + auto uniqueExecutor = CREATE_EXECUTOR(); + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); + + const op::PlatformInfo &npuInfo = op::GetCurrentPlatformInfo(); + uint32_t aicCoreNum = npuInfo.GetCubeCoreNum(); + uint32_t aivCoreNum = npuInfo.GetVectorCoreNum(); + std::string socVersionStr = npuInfo.GetSocLongVersion(); + const char *socVersion = socVersionStr.c_str(); + + int64_t batchConsistencyLevel = 0; + aclError aclRet = aclrtGetSysParamOpt(ACL_OPT_DETERMINISTIC, &batchConsistencyLevel); + if (aclRet != ACL_SUCCESS) { + OP_LOGW("aclnnSparseFlashMlaMetadata unable to get system param batch consistency level."); + } + OP_LOGD("deterministic_level=%lld", batchConsistencyLevel); + bool isBatchConsistency = (batchConsistencyLevel == BATCH_CONSISTENCY_LEVEL); + auto ret = ParamsCheck(cuSeqlensQOptional, cuSeqlensOriKvOptional, cuSeqlensCmpKvOptional, sequsedQOptional, + sequsedOriKvOptional, sequsedCmpKvOptional, cmpResidualKvOptional, oriTopkLengthOptional, + cmpTopkLengthOptional, numHeadsQ, numHeadsKv, headDim, batchSize, maxSeqlenQ, maxSeqlenOriKv, + maxSeqlenCmpKv, oriTopk, cmpTopk, cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft, + oriWinRight, layoutQOptional, layoutKvOptional, hasOriKv, hasCmpKv, aicCoreNum, aivCoreNum, + socVersion, metaData); + CHECK_RET(ret == ACLNN_SUCCESS, ret); + + const aclTensor *cuSeqlensQOptionalContiguous = nullptr; + if (cuSeqlensQOptional != nullptr) { + cuSeqlensQOptionalContiguous = l0op::Contiguous(cuSeqlensQOptional, uniqueExecutor.get()); + if (cuSeqlensQOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "cu_seqlens_q contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *cuSeqlensOriKvOptionalContiguous = nullptr; + if (cuSeqlensOriKvOptional != nullptr) { + cuSeqlensOriKvOptionalContiguous = l0op::Contiguous(cuSeqlensOriKvOptional, uniqueExecutor.get()); + if (cuSeqlensOriKvOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "cu_seqlens_ori_kv contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *cuSeqlensCmpKvOptionalContiguous = nullptr; + if (cuSeqlensCmpKvOptional != nullptr) { + cuSeqlensCmpKvOptionalContiguous = l0op::Contiguous(cuSeqlensCmpKvOptional, uniqueExecutor.get()); + if (cuSeqlensCmpKvOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "cu_seqlens_cmp_kv contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *sequsedQOptionalContiguous = nullptr; + if (sequsedQOptional != nullptr) { + sequsedQOptionalContiguous = l0op::Contiguous(sequsedQOptional, uniqueExecutor.get()); + if (sequsedQOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "seqused_q contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *sequsedOriKvOptionalContiguous = nullptr; + if (sequsedOriKvOptional != nullptr) { + sequsedOriKvOptionalContiguous = l0op::Contiguous(sequsedOriKvOptional, uniqueExecutor.get()); + if (sequsedOriKvOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "seqused_ori_kv contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *sequsedCmpKvOptionalContiguous = nullptr; + if (sequsedCmpKvOptional != nullptr) { + sequsedCmpKvOptionalContiguous = l0op::Contiguous(sequsedCmpKvOptional, uniqueExecutor.get()); + if (sequsedCmpKvOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "seqused_cmp_kv contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *cmpResidualKvOptionalContiguous = nullptr; + if (cmpResidualKvOptional != nullptr) { + cmpResidualKvOptionalContiguous = l0op::Contiguous(cmpResidualKvOptional, uniqueExecutor.get()); + if (cmpResidualKvOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "cmp_residual_kv contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *oriTopkLengthOptionalContiguous = nullptr; + if (oriTopkLengthOptional != nullptr) { + oriTopkLengthOptionalContiguous = l0op::Contiguous(oriTopkLengthOptional, uniqueExecutor.get()); + if (oriTopkLengthOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "ori_topk_length contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + const aclTensor *cmpTopkLengthOptionalContiguous = nullptr; + if (cmpTopkLengthOptional != nullptr) { + cmpTopkLengthOptionalContiguous = l0op::Contiguous(cmpTopkLengthOptional, uniqueExecutor.get()); + if (cmpTopkLengthOptionalContiguous == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "cmp_topk_length contiguous is null"); + return ACLNN_ERR_INNER_NULLPTR; + } + } + + auto output = l0op::SparseFlashMlaMetadata( + cuSeqlensQOptionalContiguous, cuSeqlensOriKvOptionalContiguous, cuSeqlensCmpKvOptionalContiguous, + sequsedQOptionalContiguous, sequsedOriKvOptionalContiguous, sequsedCmpKvOptionalContiguous, + cmpResidualKvOptionalContiguous, oriTopkLengthOptionalContiguous, cmpTopkLengthOptionalContiguous, numHeadsQ, + numHeadsKv, headDim, batchSize, maxSeqlenQ, maxSeqlenOriKv, maxSeqlenCmpKv, oriTopk, cmpTopk, cmpRatio, + oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, layoutQOptional, layoutKvOptional, hasOriKv, hasCmpKv, + socVersion, aicCoreNum, aivCoreNum, isBatchConsistency, metaData, uniqueExecutor.get()); + CHECK_RET(output != nullptr, ACLNN_ERR_INNER_NULLPTR); + + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); + uniqueExecutor.ReleaseTo(executor); + return ACLNN_SUCCESS; +} + +__attribute__((visibility("default"))) aclnnStatus aclnnSparseFlashMlaMetadata( + void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream) +{ + L2_DFX_PHASE_2(aclnnSparseFlashMlaMetadata); + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); +} + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/aclnn_sparse_flash_mla_metadata.h b/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/aclnn_sparse_flash_mla_metadata.h new file mode 100644 index 000000000000..398856853c46 --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/aclnn_sparse_flash_mla_metadata.h @@ -0,0 +1,38 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef ACLNN_SPARSE_FLASH_MLA_METADATA_H +#define ACLNN_SPARSE_FLASH_MLA_METADATA_H + +#include "aclnn/aclnn_base.h" + +#ifdef __cplusplus +extern "C" { +#endif + +__attribute__((visibility("default"))) aclnnStatus aclnnSparseFlashMlaMetadataGetWorkspaceSize( + const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, const aclTensor *sequsedQOptional, const aclTensor *sequsedOriKvOptional, + const aclTensor *sequsedCmpKvOptional, const aclTensor *cmpResidualKvOptional, + const aclTensor *oriTopkLengthOptional, const aclTensor *cmpTopkLengthOptional, int64_t numHeadsQ, + int64_t numHeadsKv, int64_t headDim, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, + int64_t maxSeqlenCmpKv, int64_t oriTopk, int64_t cmpTopk, int64_t cmpRatio, int64_t oriMaskMode, + int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight, const char *layoutQOptional, + const char *layoutKvOptional, bool hasOriKv, bool hasCmpKv, const aclTensor *metaData, uint64_t *workspaceSize, + aclOpExecutor **executor); + +__attribute__((visibility("default"))) aclnnStatus aclnnSparseFlashMlaMetadata( + void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream); + +#ifdef __cplusplus +} +#endif + +#endif // ACLNN_SPARSE_FLASH_MLA_METADATA_H diff --git a/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/sparse_flash_mla_metadata.cpp b/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/sparse_flash_mla_metadata.cpp new file mode 100644 index 000000000000..f071f450af40 --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/sparse_flash_mla_metadata.cpp @@ -0,0 +1,67 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_metadata.cpp + * \brief + */ + +#include "sparse_flash_mla_metadata.h" +#include "opdev/aicpu/aicpu_task.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_def.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/shape_utils.h" + +using namespace op; +namespace l0op { +OP_TYPE_REGISTER(SparseFlashMlaMetadata); + +const aclTensor *SparseFlashMlaMetadata( + const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, const aclTensor *sequsedQOptional, const aclTensor *sequsedOriKvOptional, + const aclTensor *sequsedCmpKvOptional, const aclTensor *cmpResidualKvOptional, + const aclTensor *oriTopkLengthOptional, const aclTensor *cmpTopkLengthOptional, int64_t numHeadsQ, + int64_t numHeadsKv, int64_t headDim, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, + int64_t maxSeqlenCmpKv, int64_t oriTopk, int64_t cmpTopk, int64_t cmpRatio, int64_t oriMaskMode, + int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight, const char *layoutQOptional, + const char *layoutKvOptional, bool hasOriKv, bool hasCmpKv, const char *socVersion, int64_t aicCoreNum, + int64_t aivCoreNum, bool isBatchConsistency, const aclTensor *metaData, aclOpExecutor *executor) +{ + L0_DFX(SparseFlashMlaMetadata, cuSeqlensQOptional, cuSeqlensOriKvOptional, cuSeqlensCmpKvOptional, + sequsedQOptional, sequsedOriKvOptional, sequsedCmpKvOptional, cmpResidualKvOptional, oriTopkLengthOptional, + cmpTopkLengthOptional, numHeadsQ, numHeadsKv, headDim, batchSize, maxSeqlenQ, maxSeqlenOriKv, maxSeqlenCmpKv, + oriTopk, cmpTopk, cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, layoutQOptional, + layoutKvOptional, hasOriKv, hasCmpKv, socVersion, aicCoreNum, aivCoreNum, isBatchConsistency, metaData); + + static internal::AicpuTaskSpace space("SparseFlashMlaMetadata"); + + auto ret = ADD_TO_LAUNCHER_LIST_AICPU( + SparseFlashMlaMetadata, + OP_ATTR_NAMES({"num_heads_q", "num_heads_kv", "head_dim", "batch_size", "max_seqlen_q", "max_seqlen_ori_kv", + "max_seqlen_cmp_kv", "ori_topk", "cmp_topk", "cmp_ratio", "ori_mask_mode", "cmp_mask_mode", + "ori_win_left", "ori_win_right", "layout_q", "layout_kv", "has_ori_kv", "has_cmp_kv", + "soc_version", "aic_core_num", "aiv_core_num", "is_batch_consistency"}), + OP_INPUT(cuSeqlensQOptional, cuSeqlensOriKvOptional, cuSeqlensCmpKvOptional, sequsedQOptional, + sequsedOriKvOptional, sequsedCmpKvOptional, cmpResidualKvOptional, oriTopkLengthOptional, + cmpTopkLengthOptional), + OP_OUTPUT(metaData), + OP_ATTR(numHeadsQ, numHeadsKv, headDim, batchSize, maxSeqlenQ, maxSeqlenOriKv, maxSeqlenCmpKv, oriTopk, cmpTopk, + cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, layoutQOptional, layoutKvOptional, + hasOriKv, hasCmpKv, socVersion, aicCoreNum, aivCoreNum, isBatchConsistency)); + OP_CHECK(ret == ACL_SUCCESS, + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "SparseFlashMlaMetadata ADD_TO_LAUNCHER_LIST_AICPU failed."), + return nullptr); + return metaData; +} + +} // namespace l0op diff --git a/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/sparse_flash_mla_metadata.h b/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/sparse_flash_mla_metadata.h new file mode 100644 index 000000000000..f91938e0fd72 --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_host/op_api/sparse_flash_mla_metadata.h @@ -0,0 +1,31 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef L0_SPARSE_FLASH_MLA_METADATA_H +#define L0_SPARSE_FLASH_MLA_METADATA_H + +#include "opdev/op_executor.h" + +namespace l0op { +const aclTensor *SparseFlashMlaMetadata(const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, const aclTensor *sequsedQOptional, + const aclTensor *sequsedOriKvOptional, const aclTensor *sequsedCmpKvOptional, + const aclTensor *cmpResidualKvOptional, const aclTensor *oriTopkLengthOptional, + const aclTensor *cmpTopkLengthOptional, int64_t numHeadsQ, int64_t numHeadsKv, + int64_t headDim, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, + int64_t maxSeqlenCmpKv, int64_t oriTopk, int64_t cmpTopk, int64_t cmpRatio, + int64_t oriMaskMode, int64_t cmpMaskMode, int64_t oriWinLeft, + int64_t oriWinRight, const char *layoutQOptional, const char *layoutKvOptional, + bool hasOriKv, bool hasCmpKv, const char *socVersion, int64_t aicCoreNum, + int64_t aivCoreNum, bool isBatchConsistency, const aclTensor *metaData, + aclOpExecutor *executor); +} // namespace l0op + +#endif diff --git a/csrc/attention/sparse_flash_mla_metadata/op_host/sparse_flash_mla_metadata_check.h b/csrc/attention/sparse_flash_mla_metadata/op_host/sparse_flash_mla_metadata_check.h new file mode 100644 index 000000000000..4a8241cc5b20 --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_host/sparse_flash_mla_metadata_check.h @@ -0,0 +1,1073 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_metadata_check.h + * \brief + */ + +#include "log/log.h" +#include "opdev/format_utils.h" +#include "opdev/data_type_utils.h" +#include "opdev/tensor_view_utils.h" +#include "../../sparse_flash_mla/op_kernel/sparse_flash_mla_kernel_metadata.h" +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +namespace { + +static constexpr const char *SMLA_ACLNN_OP_NAME = "aclnnSparseFlashMlaMetadata"; + +enum class SparseModeSmla : uint8_t { + DEFAULT_MASK = 0, + ALL_MASK, + LEFT_UP_CAUSAL, + RIGHT_DOWN_CAUSAL, + BAND, + SPARSE_BUTT, +}; + +inline constexpr int64_t SMLA_CMP_RATIO_LOWER_BOUND = 1; +inline constexpr int64_t SMLA_CMP_RATIO_UPPER_BOUND = 128; +inline constexpr int64_t SMLA_NUM_HEADS_Q_LOWER_BOUND = 1; +inline constexpr int64_t SMLA_NUM_HEADS_Q_UPPER_BOUND = 128; + +inline bool IsPowerOfTwoInRangeSmla(int64_t value, int64_t minValue, int64_t maxValue) +{ + return value >= minValue && value <= maxValue && ((value & (value - 1)) == 0); +} + +inline bool IsCmpRatioSupportSmla(const char *socVersion, bool hasCmpKv, int64_t cmpTopk, int64_t cmpRatio) +{ + if (!hasCmpKv) { + return cmpRatio == 0; + } + if (socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) { + return cmpRatio >= SMLA_CMP_RATIO_LOWER_BOUND && cmpRatio <= SMLA_CMP_RATIO_UPPER_BOUND; + } + return (cmpTopk > 0) ? (cmpRatio == 1 || cmpRatio == 2 || cmpRatio == 4) : (cmpRatio == 128); +} + +inline bool IsTensorExistSmla(const aclTensor *tensor) +{ + return (tensor != nullptr) && (tensor->GetViewShape().GetDimNum() > 0) && (tensor->GetViewShape().GetDim(0) > 0); +} + +aclnnStatus CheckReservedOptionalTensorSmla(const aclTensor *tensor, const char *tensorName) +{ + if (!IsTensorExistSmla(tensor)) { + return ACLNN_SUCCESS; + } + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, tensorName, + std::string(tensorName) + " is reserved and does not support " + "non-empty tensor"); + return ACLNN_ERR_PARAM_INVALID; +} + +int64_t GetDimNumSmla(const aclTensor *tensor) +{ + if (tensor == nullptr) { + return -1; + } + return tensor->GetViewShape().GetDimNum(); +} + +aclDataType GetDataTypeSmla(const aclTensor *tensor) +{ + aclDataType dataType = aclDataType::ACL_DT_UNDEFINED; + if (tensor == nullptr) { + return dataType; + } + aclGetDataType(tensor, &dataType); + return dataType; +} + +inline bool IsTensorSourceSmla(const std::string &source) +{ + return source != "batch_size"; +} + +inline int64_t GetRawShapeSizeSmla(const std::string &source, int64_t batchValue) +{ + if (source.find("cu_seqlens") != std::string::npos) { + return batchValue + 1; + } + return batchValue; +} + +inline std::string GetSourceDescSmla(const std::string &source) +{ + if (source == "batch_size") { + return "batch_size"; + } + if (source.find("cu_seqlens") != std::string::npos) { + return "the shape size of " + source + " minus 1"; + } + return "the shape size of " + source; +} + +aclnnStatus CheckSingleParamSmla(int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, int64_t maxSeqlenCmpKv, + int64_t numHeadsQ, int64_t numHeadsKv, int64_t headDim, int64_t oriTopk, + int64_t cmpTopk, int64_t cmpRatio, int64_t oriMaskMode, int64_t cmpMaskMode, + int64_t oriWinLeft, int64_t oriWinRight, const char *layoutQOptional, + const char *layoutKvOptional, bool hasOriKv, bool hasCmpKv, uint32_t aicCoreNum, + uint32_t aivCoreNum, const char *socVersion) +{ + // batch_size >= 0 + if (batchSize < 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "batch_size", std::to_string(batchSize), + "The value of batch_size must be greater than or equal to 0"); + return ACLNN_ERR_PARAM_INVALID; + } + // max_seqlen_q >= 0 + if (maxSeqlenQ < 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "max_seqlen_q", std::to_string(maxSeqlenQ), + "The value of max_seqlen_q must be greater than or equal to 0"); + return ACLNN_ERR_PARAM_INVALID; + } + // max_seqlen_ori_kv >= 0 + if (maxSeqlenOriKv < 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "max_seqlen_ori_kv", std::to_string(maxSeqlenOriKv), + "The value of max_seqlen_ori_kv must be greater than or equal to 0"); + return ACLNN_ERR_PARAM_INVALID; + } + // max_seqlen_cmp_kv >= 0 + if (maxSeqlenCmpKv < 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "max_seqlen_cmp_kv", std::to_string(maxSeqlenCmpKv), + "The value of max_seqlen_cmp_kv must be greater than or equal to 0"); + return ACLNN_ERR_PARAM_INVALID; + } + // num_heads_q [1, 128] + if (numHeadsQ < SMLA_NUM_HEADS_Q_LOWER_BOUND || numHeadsQ > SMLA_NUM_HEADS_Q_UPPER_BOUND) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "num_heads_q", std::to_string(numHeadsQ), + "The current value is not within the valid range. " + "The valid range is [" + + std::to_string(SMLA_NUM_HEADS_Q_LOWER_BOUND) + ", " + + std::to_string(SMLA_NUM_HEADS_Q_UPPER_BOUND) + "]"); + return ACLNN_ERR_PARAM_INVALID; + } + // num_heads_kv: 1 + if (numHeadsKv != 1) { + OP_LOGE_FOR_INVALID_VALUE(SMLA_ACLNN_OP_NAME, "num_heads_kv", std::to_string(numHeadsKv), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + if (numHeadsQ % numHeadsKv != 0) { + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON(SMLA_ACLNN_OP_NAME, "num_heads_q, num_heads_kv", + std::to_string(numHeadsQ) + ", " + std::to_string(numHeadsKv), + "The value of num_heads_q must be divisible by " + "that of num_heads_kv"); + return ACLNN_ERR_PARAM_INVALID; + } + int64_t headRatio = numHeadsQ / numHeadsKv; + if (socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) { + if (headRatio < SMLA_NUM_HEADS_Q_LOWER_BOUND || headRatio > SMLA_NUM_HEADS_Q_UPPER_BOUND) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "num_heads_q / num_heads_kv", + std::to_string(headRatio), + "The current value is not within the valid range. " + "The valid range is [" + + std::to_string(SMLA_NUM_HEADS_Q_LOWER_BOUND) + ", " + + std::to_string(SMLA_NUM_HEADS_Q_UPPER_BOUND) + "]"); + return ACLNN_ERR_PARAM_INVALID; + } + } else if (!IsPowerOfTwoInRangeSmla(headRatio, SMLA_NUM_HEADS_Q_LOWER_BOUND, SMLA_NUM_HEADS_Q_UPPER_BOUND)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "num_heads_q / num_heads_kv", + std::to_string(headRatio), + "The current value is not within the valid range. " + "The valid range is power of two in [" + + std::to_string(SMLA_NUM_HEADS_Q_LOWER_BOUND) + ", " + + std::to_string(SMLA_NUM_HEADS_Q_UPPER_BOUND) + "]"); + return ACLNN_ERR_PARAM_INVALID; + } + // head_dim: 512 + if (headDim != 512) { + OP_LOGE_FOR_INVALID_VALUE(SMLA_ACLNN_OP_NAME, "head_dim", std::to_string(headDim), "512"); + return ACLNN_ERR_PARAM_INVALID; + } + if (hasOriKv) { + // ori_topk >= 0 + if (oriTopk < 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "ori_topk", std::to_string(oriTopk), + "When has_ori_kv is true, the value of ori_topk must be " + "greater than or equal to 0"); + return ACLNN_ERR_PARAM_INVALID; + } + if (!(socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) && oriTopk != 0 && hasCmpKv) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "ori_topk", std::to_string(oriTopk), + "ori_topk is reserved and " + "the value of ori_topk must be 0 when has_cmp_kv is true"); + return ACLNN_ERR_PARAM_INVALID; + } + if (socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) { + // ori_mask_mode: 0, 3, or 4 + if (oriMaskMode != static_cast(SparseModeSmla::DEFAULT_MASK) && + oriMaskMode != static_cast(SparseModeSmla::RIGHT_DOWN_CAUSAL) && + oriMaskMode != static_cast(SparseModeSmla::BAND)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "ori_mask_mode", std::to_string(oriMaskMode), + "When has_ori_kv is true, the value of ori_mask_mode " + "must be in [0, 3, 4]"); + return ACLNN_ERR_PARAM_INVALID; + } + // A5 treats -1 as unlimited window + if (oriWinLeft < -1 || oriWinRight < -1) { + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + SMLA_ACLNN_OP_NAME, "ori_win_left and ori_win_right", + std::to_string(oriWinLeft) + " and " + std::to_string(oriWinRight), + "When has_ori_kv is true, the value of ori_win_left, " + "ori_win_right must be greater than or equal to -1"); + return ACLNN_ERR_PARAM_INVALID; + } + } else { + if (oriTopk != 0) { + if (oriMaskMode != static_cast(SparseModeSmla::DEFAULT_MASK)) { + OP_LOGE_FOR_INVALID_VALUE(SMLA_ACLNN_OP_NAME, "ori_mask_mode", std::to_string(oriMaskMode), "0"); + return ACLNN_ERR_PARAM_INVALID; + } + if (oriWinLeft < 0 || oriWinRight < 0) { + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + SMLA_ACLNN_OP_NAME, "ori_win_left, ori_win_right", + std::to_string(oriWinLeft) + ", " + std::to_string(oriWinRight), + "When has_ori_kv is true and ori_topk is non-zero (DSpark), " + "ori_win_left and ori_win_right must be non-negative"); + return ACLNN_ERR_PARAM_INVALID; + } + } else { + if (oriMaskMode != static_cast(SparseModeSmla::BAND)) { + OP_LOGE_FOR_INVALID_VALUE(SMLA_ACLNN_OP_NAME, "ori_mask_mode", std::to_string(oriMaskMode), "4"); + return ACLNN_ERR_PARAM_INVALID; + } + if (oriWinLeft != 127 || oriWinRight != 0) { + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + SMLA_ACLNN_OP_NAME, "ori_win_left and ori_win_right", + std::to_string(oriWinLeft) + " and " + std::to_string(oriWinRight), + "When has_ori_kv is true, the value of ori_win_left " + "must be 127 and the value of ori_win_right must be 0"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } + } + if (hasCmpKv) { + // cmp_topk >= 0 + if (cmpTopk < 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_topk", std::to_string(cmpTopk), + "When has_cmp_kv is true, the value of cmp_topk must be " + "greater than or equal to 0"); + return ACLNN_ERR_PARAM_INVALID; + } + if (!(socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) && cmpTopk != 0 && cmpTopk != 512 && + cmpTopk != 1024) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_topk", std::to_string(cmpTopk), + "When has_cmp_kv is true, the value of cmp_topk must be " + "in [0, 512, 1024]"); + return ACLNN_ERR_PARAM_INVALID; + } + if (socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) { + // cmp_mask_mode: 0 or 3 + if (cmpMaskMode != static_cast(SparseModeSmla::DEFAULT_MASK) && + cmpMaskMode != static_cast(SparseModeSmla::RIGHT_DOWN_CAUSAL)) { + OP_LOGE_FOR_INVALID_VALUE(SMLA_ACLNN_OP_NAME, "cmp_mask_mode", std::to_string(cmpMaskMode), "0 or 3"); + return ACLNN_ERR_PARAM_INVALID; + } + } else { + if (cmpMaskMode != static_cast(SparseModeSmla::RIGHT_DOWN_CAUSAL)) { + OP_LOGE_FOR_INVALID_VALUE(SMLA_ACLNN_OP_NAME, "cmp_mask_mode", std::to_string(cmpMaskMode), "3"); + return ACLNN_ERR_PARAM_INVALID; + } + } + if (!IsCmpRatioSupportSmla(socVersion, hasCmpKv, cmpTopk, cmpRatio)) { + if (socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_ratio", std::to_string(cmpRatio), + "When has_cmp_kv is true, the current value is not " + "within the valid range. The valid range is [" + + std::to_string(SMLA_CMP_RATIO_LOWER_BOUND) + ", " + + std::to_string(SMLA_CMP_RATIO_UPPER_BOUND) + "]"); + } else { + int64_t expectedCmpRatio = (cmpTopk > 0) ? 4 : 128; + if (cmpTopk > 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_ratio", std::to_string(cmpRatio), + "When has_cmp_kv is true and cmp_topk is non-zero" + "(CSA with cmp_sparse_indices), " + "the value of cmp_ratio must be 1, 2 or " + + std::to_string(expectedCmpRatio)); + } else { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_ratio", std::to_string(cmpRatio), + "When has_cmp_kv is true and cmp_topk is 0" + "(HCA without cmp_sparse_indices), " + "the value of cmp_ratio must be " + + std::to_string(expectedCmpRatio)); + } + } + return ACLNN_ERR_PARAM_INVALID; + } + } else if (!(socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) && + !IsCmpRatioSupportSmla(socVersion, hasCmpKv, cmpTopk, cmpRatio)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_ratio", std::to_string(cmpRatio), + "When has_cmp_kv is false, the value of cmp_ratio must be 0"); + return ACLNN_ERR_PARAM_INVALID; + } + if (layoutQOptional == nullptr) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "layout_q", "layout_q cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + if (layoutKvOptional == nullptr) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "layout_kv", "layout_kv cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + // layout_q: BSND or TND + if (strcmp(layoutQOptional, "TND") != 0 && strcmp(layoutQOptional, "BSND") != 0) { + OP_LOGE_FOR_INVALID_VALUE(SMLA_ACLNN_OP_NAME, "layout_q", layoutQOptional, "TND or BSND"); + return ACLNN_ERR_PARAM_INVALID; + } + // layout_kv: BSND, TND, or PA_BBND + if (strcmp(layoutKvOptional, "BSND") != 0 && strcmp(layoutKvOptional, "TND") != 0 && + strcmp(layoutKvOptional, "PA_BBND") != 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "layout_kv", layoutKvOptional, + "The value of layout_kv must be in [TND, BSND, PA_BBND]"); + return ACLNN_ERR_PARAM_INVALID; + } + if (strcmp(layoutKvOptional, "PA_BBND") != 0 && strcmp(layoutQOptional, layoutKvOptional) != 0) { + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + SMLA_ACLNN_OP_NAME, "layout_q, layout_kv", + std::string(layoutQOptional) + ", " + std::string(layoutKvOptional), + "When layout_kv is not PA_BBND, the values of layout_q, layout_kv must be the same"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 layout_q 为 BSND 时,max_seqlen_q 必须大于 0 + if (strcmp(layoutQOptional, "BSND") == 0 && maxSeqlenQ <= 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "max_seqlen_q", std::to_string(maxSeqlenQ), + "When layout_q is BSND, the value of max_seqlen_q " + "must be equal to the size of the second axis of q"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 has_ori_kv 且 layout_kv 为 BSND 时,max_seqlen_ori_kv 必须大于 0 + if (hasOriKv && strcmp(layoutKvOptional, "BSND") == 0 && maxSeqlenOriKv <= 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "max_seqlen_ori_kv", std::to_string(maxSeqlenOriKv), + "When has_ori_kv is true and layout_kv is BSND, " + "the value of max_seqlen_ori_kv " + "must be equal to the size of the second axis of ori_kv"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 has_cmp_kv 且 layout_kv 为 BSND 时,max_seqlen_cmp_kv 必须大于 0 + if (hasCmpKv && strcmp(layoutKvOptional, "BSND") == 0 && maxSeqlenCmpKv <= 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "max_seqlen_cmp_kv", std::to_string(maxSeqlenCmpKv), + "When has_cmp_kv is true and layout_kv is BSND, " + "the value of max_seqlen_cmp_kv " + "must be equal to the size of the second axis of cmp_kv"); + return ACLNN_ERR_PARAM_INVALID; + } + // 核数校验 + if (aicCoreNum == 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "aic_core_num", std::to_string(aicCoreNum), + "The value of aic_core_num must be greater than 0"); + return ACLNN_ERR_PARAM_INVALID; + } + if (aicCoreNum > optiling::AIC_CORE_MAX_NUM) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "aic_core_num", std::to_string(aicCoreNum), + "The current value is not within the valid range. " + "The valid range is [1, " + + std::to_string(optiling::AIC_CORE_MAX_NUM) + "]"); + return ACLNN_ERR_PARAM_INVALID; + } + if (aivCoreNum == 0) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "aiv_core_num", std::to_string(aivCoreNum), + "The value of aiv_core_num must be greater than 0"); + return ACLNN_ERR_PARAM_INVALID; + } + if (aivCoreNum > optiling::AIV_CORE_MAX_NUM) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(SMLA_ACLNN_OP_NAME, "aiv_core_num", std::to_string(aivCoreNum), + "The current value is not within the valid range. " + "The valid range is [1, " + + std::to_string(optiling::AIV_CORE_MAX_NUM) + "]"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验切g模板核数 + if (numHeadsQ == 128) { + if (aicCoreNum == 1 || aivCoreNum == 1) { + OP_LOGE_FOR_INVALID_VALUES_WITH_REASON( + SMLA_ACLNN_OP_NAME, "num_heads_q and aic_core_num and aiv_core_num", + std::to_string(numHeadsQ) + " and " + std::to_string(aicCoreNum) + " and " + std::to_string(aivCoreNum), + "When num_heads_q is 128, the value of aic_core_num, " + "aiv_core_num cannot be 1"); + return ACLNN_ERR_PARAM_INVALID; + } + } + return ACLNN_SUCCESS; +} + +aclnnStatus CheckExistenceSmla(const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, const aclTensor *sequsedOriKvOptional, + const aclTensor *sequsedCmpKvOptional, const aclTensor *cmpResidualKvOptional, + const aclTensor *oriTopkLengthOptional, const aclTensor *cmpTopkLengthOptional, + int64_t oriTopk, int64_t cmpTopk, int64_t cmpRatio, int64_t oriMaskMode, + int64_t cmpMaskMode, bool hasOriKv, bool hasCmpKv, const char *layoutQOptional, + const char *layoutKvOptional, const char *socVersion, const aclTensor *metadata) +{ + // cu_seqlens_q 存在性校验 + if (strcmp(layoutQOptional, "TND") == 0) { + if (!IsTensorExistSmla(cuSeqlensQOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "cu_seqlens_q", + "When layout_q is TND, cu_seqlens_q cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + if (hasOriKv) { + // cu_seqlens_ori_kv 存在性校验 + if (strcmp(layoutKvOptional, "TND") == 0) { + if (!IsTensorExistSmla(cuSeqlensOriKvOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "cu_seqlens_ori_kv", + "When has_ori_kv is true and layout_kv is TND, " + "cu_seqlens_ori_kv cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // seqused_ori_kv 存在性校验 + if (socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) { + if ((oriMaskMode != 0 || oriTopk == 0) && strcmp(layoutKvOptional, "PA_BBND") == 0) { + if (!IsTensorExistSmla(sequsedOriKvOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "seqused_ori_kv", + "When has_ori_kv is true, ori_mask_mode != 0 or " + "ori_topk == 0, and layout_kv is PA_BBND, " + "seqused_ori_kv cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } else { + if (strcmp(layoutKvOptional, "PA_BBND") == 0) { + if (!IsTensorExistSmla(sequsedOriKvOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "seqused_ori_kv", + "When has_ori_kv is true and layout_kv is PA_BBND, " + "seqused_ori_kv cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } + // ori_topk_length 存在性校验 + if (oriTopk != 0 && oriMaskMode == static_cast(SparseModeSmla::DEFAULT_MASK)) { + if (!IsTensorExistSmla(oriTopkLengthOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "ori_topk_length", + "When has_ori_kv is true, ori_topk is not 0 and " + "ori_mask_mode is 0, ori_topk_length cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } + if (hasCmpKv) { + // cu_seqlens_cmp_kv 存在性校验 + if (strcmp(layoutKvOptional, "TND") == 0) { + if (!IsTensorExistSmla(cuSeqlensCmpKvOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "cu_seqlens_cmp_kv", + "When has_cmp_kv is true and layout_kv is TND, " + "cu_seqlens_cmp_kv cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // seqused_cmp_kv 存在性校验 + if (socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr) { + if ((cmpMaskMode != 0 || cmpTopk == 0) && strcmp(layoutKvOptional, "PA_BBND") == 0) { + if (!IsTensorExistSmla(sequsedCmpKvOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "seqused_cmp_kv", + "When has_cmp_kv is true, cmp_mask_mode != 0 or " + "cmp_topk == 0, and layout_kv is PA_BBND, " + "seqused_cmp_kv cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } else { + if (strcmp(layoutKvOptional, "PA_BBND") == 0) { + if (!IsTensorExistSmla(sequsedCmpKvOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "seqused_cmp_kv", + "When has_cmp_kv is true and layout_kv is PA_BBND, " + "seqused_cmp_kv cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } + // cmp_residual_kv 存在性校验 + if (cmpRatio != 1 && cmpMaskMode == static_cast(SparseModeSmla::RIGHT_DOWN_CAUSAL)) { + if (!IsTensorExistSmla(cmpResidualKvOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_residual_kv", + "When has_cmp_kv is true, cmp_ratio is not 1 and " + "cmp_mask_mode is 3, cmp_residual_kv cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // cmp_topk_length 存在性校验 + if (cmpTopk != 0 && cmpMaskMode == static_cast(SparseModeSmla::DEFAULT_MASK)) { + if (!IsTensorExistSmla(cmpTopkLengthOptional)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_topk_length", + "When has_cmp_kv is true, cmp_topk is not 0 and " + "cmp_mask_mode is 0, cmp_topk_length cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } + // metadata 存在性校验 + if (!IsTensorExistSmla(metadata)) { + OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON(SMLA_ACLNN_OP_NAME, "metadata", "metadata cannot be empty"); + return ACLNN_ERR_PARAM_INVALID; + } + return ACLNN_SUCCESS; +} + +int64_t GetQueryBatchSizeSmla(const aclTensor *sequsedQOptional, const aclTensor *cuSeqlensQOptional, + const char *layoutQOptional, int64_t batchSize, std::string *source) +{ + if (IsTensorExistSmla(sequsedQOptional)) { + *source = "seqused_q"; + return sequsedQOptional->GetViewShape().GetDim(0); + } + if (strcmp(layoutQOptional, "TND") == 0) { + if (IsTensorExistSmla(cuSeqlensQOptional)) { + *source = "cu_seqlens_q"; + return cuSeqlensQOptional->GetViewShape().GetDim(0) - 1; + } + } + *source = "batch_size"; + return batchSize; +} + +int64_t GetOriKvBatchSizeSmla(const aclTensor *sequsedOriKvOptional, const aclTensor *cuSeqlensOriKvOptional, + const char *layoutKvOptional, int64_t batchSize, std::string *source) +{ + if (IsTensorExistSmla(sequsedOriKvOptional)) { + *source = "seqused_ori_kv"; + return sequsedOriKvOptional->GetViewShape().GetDim(0); + } + if (strcmp(layoutKvOptional, "TND") == 0) { + if (IsTensorExistSmla(cuSeqlensOriKvOptional)) { + *source = "cu_seqlens_ori_kv"; + return cuSeqlensOriKvOptional->GetViewShape().GetDim(0) - 1; + } + } + *source = "batch_size"; + return batchSize; +} + +int64_t GetCmpKvBatchSizeSmla(const aclTensor *sequsedCmpKvOptional, const aclTensor *cuSeqlensCmpKvOptional, + const char *layoutKvOptional, int64_t batchSize, std::string *source) +{ + if (IsTensorExistSmla(sequsedCmpKvOptional)) { + *source = "seqused_cmp_kv"; + return sequsedCmpKvOptional->GetViewShape().GetDim(0); + } + if (strcmp(layoutKvOptional, "TND") == 0) { + if (IsTensorExistSmla(cuSeqlensCmpKvOptional)) { + *source = "cu_seqlens_cmp_kv"; + return cuSeqlensCmpKvOptional->GetViewShape().GetDim(0) - 1; + } + } + *source = "batch_size"; + return batchSize; +} + +std::string TopkLengthShapeToStringSmla(const aclTensor *topkLengthOptional) +{ + const auto &shape = topkLengthOptional->GetViewShape(); + std::string result; + for (size_t i = 0; i < shape.GetDimNum(); ++i) { + if (i != 0) { + result += ", "; + } + result += std::to_string(shape.GetDim(i)); + } + return result; +} + +aclnnStatus CheckTopkLengthFirstDimSmla(const aclTensor *topkLengthOptional, const std::string &topkLengthName, + int64_t queryBatchSize, const std::string &querySource) +{ + if (topkLengthOptional->GetViewShape().GetDim(0) == queryBatchSize) { + return ACLNN_SUCCESS; + } + std::string incorrectShape = TopkLengthShapeToStringSmla(topkLengthOptional); + if (IsTensorSourceSmla(querySource)) { + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(SMLA_ACLNN_OP_NAME, topkLengthName, incorrectShape, + "When layout_q is BSND, the size of the first axis of " + topkLengthName + + " must be equal to " + GetSourceDescSmla(querySource)); + } else { + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + SMLA_ACLNN_OP_NAME, topkLengthName, incorrectShape, + "When layout_q is BSND, the size of the first axis of " + topkLengthName + " must be equal to batch_size"); + } + return ACLNN_ERR_PARAM_INVALID; +} + +struct TopkLengthAxisSmla { + int64_t index; + const char *desc; +}; + +inline constexpr TopkLengthAxisSmla SMLA_TOPK_LENGTH_SECOND_AXIS{1, "second"}; +inline constexpr TopkLengthAxisSmla SMLA_TOPK_LENGTH_THIRD_AXIS{2, "third"}; + +aclnnStatus CheckTopkLengthSingleDimSmla(const aclTensor *topkLengthOptional, const std::string &topkLengthName, + TopkLengthAxisSmla axis, int64_t expectedValue, + const std::string &expectedDesc, const char *layoutQOptional) +{ + if (topkLengthOptional->GetViewShape().GetDim(axis.index) == expectedValue) { + return ACLNN_SUCCESS; + } + std::string incorrectShape = TopkLengthShapeToStringSmla(topkLengthOptional); + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(SMLA_ACLNN_OP_NAME, topkLengthName, incorrectShape, + "When layout_q is " + std::string(layoutQOptional) + ", the size of the " + + axis.desc + " axis of " + topkLengthName + " must be equal to " + + expectedDesc); + return ACLNN_ERR_PARAM_INVALID; +} + +aclnnStatus CheckConsistencySmla(const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, const aclTensor *sequsedQOptional, + const aclTensor *sequsedOriKvOptional, const aclTensor *sequsedCmpKvOptional, + const aclTensor *cmpResidualKvOptional, const aclTensor *oriTopkLengthOptional, + const aclTensor *cmpTopkLengthOptional, int64_t batchSize, const char *layoutQOptional, + const char *layoutKvOptional, bool hasOriKv, bool hasCmpKv, int64_t oriTopk, + int64_t cmpTopk, int64_t oriMaskMode, int64_t cmpMaskMode, int64_t maxSeqlenQ, + int64_t numHeadsKv, const char *socVersion, const aclTensor *metadata) +{ + aclDataType dataType = aclDataType::ACL_DT_UNDEFINED; + int64_t dimNum = -1; + if (!(socVersion != nullptr && strstr(socVersion, "Ascend950") != nullptr)) { + if (CheckReservedOptionalTensorSmla(cmpTopkLengthOptional, "cmp_topk_length") != ACLNN_SUCCESS) { + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 cu_seqlens_q + if (IsTensorExistSmla(cuSeqlensQOptional)) { + // 校验 cu_seqlens_q 维度 + dimNum = GetDimNumSmla(cuSeqlensQOptional); + if (dimNum != 1) { + OP_LOGE_FOR_INVALID_SHAPEDIM(SMLA_ACLNN_OP_NAME, "cu_seqlens_q", std::to_string(dimNum), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 cu_seqlens_q 数据类型 + dataType = GetDataTypeSmla(cuSeqlensQOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cu_seqlens_q", ToString(dataType).GetString(), + "The dtype of cu_seqlens_q must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 seqused_q + if (IsTensorExistSmla(sequsedQOptional)) { + // 校验 seqused_q 维度 + dimNum = GetDimNumSmla(sequsedQOptional); + if (dimNum != 1) { + OP_LOGE_FOR_INVALID_SHAPEDIM(SMLA_ACLNN_OP_NAME, "seqused_q", std::to_string(dimNum), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 seqused_q 数据类型 + dataType = GetDataTypeSmla(sequsedQOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "seqused_q", ToString(dataType).GetString(), + "The dtype of seqused_q must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // ori_kv部分 + if (hasOriKv) { + // 校验 cu_seqlens_ori_kv + if (IsTensorExistSmla(cuSeqlensOriKvOptional)) { + // 校验 cu_seqlens_ori_kv 维度 + dimNum = GetDimNumSmla(cuSeqlensOriKvOptional); + if (dimNum != 1) { + OP_LOGE_FOR_INVALID_SHAPEDIM(SMLA_ACLNN_OP_NAME, "cu_seqlens_ori_kv", std::to_string(dimNum), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 cu_seqlens_ori_kv 数据类型 + dataType = GetDataTypeSmla(cuSeqlensOriKvOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cu_seqlens_ori_kv", + ToString(dataType).GetString(), + "The dtype of cu_seqlens_ori_kv must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 seqused_ori_kv + if (IsTensorExistSmla(sequsedOriKvOptional)) { + // 校验 seqused_ori_kv 维度 + dimNum = GetDimNumSmla(sequsedOriKvOptional); + if (dimNum != 1) { + OP_LOGE_FOR_INVALID_SHAPEDIM(SMLA_ACLNN_OP_NAME, "seqused_ori_kv", std::to_string(dimNum), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 seqused_ori_kv 数据类型 + dataType = GetDataTypeSmla(sequsedOriKvOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "seqused_ori_kv", + ToString(dataType).GetString(), + "The dtype of seqused_ori_kv must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 ori_topk_length + if (oriTopk != 0 && oriMaskMode == static_cast(SparseModeSmla::DEFAULT_MASK) && + IsTensorExistSmla(oriTopkLengthOptional)) { + // 校验 ori_topk_length 维度 + dimNum = GetDimNumSmla(oriTopkLengthOptional); + if (strcmp(layoutQOptional, "TND") == 0) { + if (dimNum != 2) { + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(SMLA_ACLNN_OP_NAME, "ori_topk_length", + std::to_string(dimNum), + "The shape dim of ori_topk_length must be 2 " + "when layout_q is TND"); + return ACLNN_ERR_PARAM_INVALID; + } + } else if (strcmp(layoutQOptional, "BSND") == 0) { + if (dimNum != 3) { + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(SMLA_ACLNN_OP_NAME, "ori_topk_length", + std::to_string(dimNum), + "The shape dim of ori_topk_length must be 3 " + "when layout_q is BSND"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 ori_topk_length 数据类型 + dataType = GetDataTypeSmla(oriTopkLengthOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "ori_topk_length", + ToString(dataType).GetString(), + "The dtype of ori_topk_length must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } + // cmp_kv部分 + if (hasCmpKv) { + // 校验 cu_seqlens_cmp_kv + if (IsTensorExistSmla(cuSeqlensCmpKvOptional)) { + // 校验 cu_seqlens_cmp_kv 维度 + dimNum = GetDimNumSmla(cuSeqlensCmpKvOptional); + if (dimNum != 1) { + OP_LOGE_FOR_INVALID_SHAPEDIM(SMLA_ACLNN_OP_NAME, "cu_seqlens_cmp_kv", std::to_string(dimNum), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 cu_seqlens_cmp_kv 数据类型 + dataType = GetDataTypeSmla(cuSeqlensCmpKvOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cu_seqlens_cmp_kv", + ToString(dataType).GetString(), + "The dtype of cu_seqlens_cmp_kv must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 seqused_cmp_kv + if (IsTensorExistSmla(sequsedCmpKvOptional)) { + // 校验 seqused_cmp_kv 维度 + dimNum = GetDimNumSmla(sequsedCmpKvOptional); + if (dimNum != 1) { + OP_LOGE_FOR_INVALID_SHAPEDIM(SMLA_ACLNN_OP_NAME, "seqused_cmp_kv", std::to_string(dimNum), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 seqused_cmp_kv 数据类型 + dataType = GetDataTypeSmla(sequsedCmpKvOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "seqused_cmp_kv", + ToString(dataType).GetString(), + "The dtype of seqused_cmp_kv must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 cmp_residual_kv + if (IsTensorExistSmla(cmpResidualKvOptional)) { + // 校验 cmp_residual_kv 维度 + dimNum = GetDimNumSmla(cmpResidualKvOptional); + if (dimNum != 1) { + OP_LOGE_FOR_INVALID_SHAPEDIM(SMLA_ACLNN_OP_NAME, "cmp_residual_kv", std::to_string(dimNum), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 cmp_residual_kv 数据类型 + dataType = GetDataTypeSmla(cmpResidualKvOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_residual_kv", + ToString(dataType).GetString(), + "The dtype of cmp_residual_kv must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 cmp_topk_length + if (cmpTopk != 0 && cmpMaskMode == static_cast(SparseModeSmla::DEFAULT_MASK) && + IsTensorExistSmla(cmpTopkLengthOptional)) { + // 校验 cmp_topk_length 维度 + dimNum = GetDimNumSmla(cmpTopkLengthOptional); + if (strcmp(layoutQOptional, "TND") == 0) { + if (dimNum != 2) { + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_topk_length", + std::to_string(dimNum), + "The shape dim of cmp_topk_length must be 2 " + "when layout_q is TND"); + return ACLNN_ERR_PARAM_INVALID; + } + } else if (strcmp(layoutQOptional, "BSND") == 0) { + if (dimNum != 3) { + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_topk_length", + std::to_string(dimNum), + "The shape dim of cmp_topk_length must be 3 " + "when layout_q is BSND"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 cmp_topk_length 数据类型 + dataType = GetDataTypeSmla(cmpTopkLengthOptional); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "cmp_topk_length", + ToString(dataType).GetString(), + "The dtype of cmp_topk_length must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + } + // 校验 metadata + if (IsTensorExistSmla(metadata)) { + // 校验 metadata 维度 + dimNum = GetDimNumSmla(metadata); + if (dimNum != 1) { + OP_LOGE_FOR_INVALID_SHAPEDIM(SMLA_ACLNN_OP_NAME, "metadata", std::to_string(dimNum), "1"); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 metadata 元素数 + if (metadata->GetViewShape().GetDim(0) != optiling::SMLA_METADATA_TOTAL_SIZE) { + OP_LOGE_FOR_INVALID_SHAPESIZE(SMLA_ACLNN_OP_NAME, "metadata", + std::to_string(metadata->GetViewShape().GetDim(0)), + std::to_string(optiling::SMLA_METADATA_TOTAL_SIZE)); + return ACLNN_ERR_PARAM_INVALID; + } + // 校验 metadata 数据类型 + dataType = GetDataTypeSmla(metadata); + if (dataType != aclDataType::ACL_INT32) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(SMLA_ACLNN_OP_NAME, "metadata", ToString(dataType).GetString(), + "The dtype of metadata must be int32"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 q/kv 维度一致性 + std::string querySource; + int64_t queryBatchSize = + GetQueryBatchSizeSmla(sequsedQOptional, cuSeqlensQOptional, layoutQOptional, batchSize, &querySource); + // 校验TND场景q维度一致性 + if (strcmp(layoutQOptional, "TND") == 0 && IsTensorExistSmla(sequsedQOptional)) { + int64_t cuSeqlensQBatchSize = cuSeqlensQOptional->GetViewShape().GetDim(0) - 1; + if (cuSeqlensQBatchSize != queryBatchSize) { + OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(SMLA_ACLNN_OP_NAME, "cu_seqlens_q and seqused_q", + std::to_string(cuSeqlensQOptional->GetViewShape().GetDim(0)) + + " and " + + std::to_string(sequsedQOptional->GetViewShape().GetDim(0)), + "When layout_q is TND and seqused_q is passed, " + "the shape size of cu_seqlens_q minus 1 must be equal to " + "the shape size of seqused_q"); + return ACLNN_ERR_PARAM_INVALID; + } + } + if (hasOriKv) { + std::string oriKvSource; + int64_t oriKvBatchSize = GetOriKvBatchSizeSmla(sequsedOriKvOptional, cuSeqlensOriKvOptional, layoutKvOptional, + batchSize, &oriKvSource); + // 校验q与ori_kv维度一致性 + if (queryBatchSize != oriKvBatchSize) { + if (IsTensorSourceSmla(querySource) && IsTensorSourceSmla(oriKvSource)) { + OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON( + SMLA_ACLNN_OP_NAME, querySource + " and " + oriKvSource, + std::to_string(GetRawShapeSizeSmla(querySource, queryBatchSize)) + " and " + + std::to_string(GetRawShapeSizeSmla(oriKvSource, oriKvBatchSize)), + "When has_ori_kv is true, " + GetSourceDescSmla(querySource) + " must be equal to " + + GetSourceDescSmla(oriKvSource)); + } else if (IsTensorSourceSmla(querySource)) { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + SMLA_ACLNN_OP_NAME, querySource, std::to_string(GetRawShapeSizeSmla(querySource, queryBatchSize)), + "When has_ori_kv is true, " + GetSourceDescSmla(querySource) + " must be equal to batch_size"); + } else { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + SMLA_ACLNN_OP_NAME, oriKvSource, std::to_string(GetRawShapeSizeSmla(oriKvSource, oriKvBatchSize)), + "When has_ori_kv is true, " + GetSourceDescSmla(oriKvSource) + " must be equal to batch_size"); + } + return ACLNN_ERR_PARAM_INVALID; + } + // 校验TND场景ori_kv维度一致性 + if (strcmp(layoutKvOptional, "TND") == 0 && IsTensorExistSmla(sequsedOriKvOptional)) { + int64_t cuSeqlensOriKvBatchSize = cuSeqlensOriKvOptional->GetViewShape().GetDim(0) - 1; + if (cuSeqlensOriKvBatchSize != oriKvBatchSize) { + OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON( + SMLA_ACLNN_OP_NAME, "cu_seqlens_ori_kv and seqused_ori_kv", + std::to_string(cuSeqlensOriKvOptional->GetViewShape().GetDim(0)) + " and " + + std::to_string(sequsedOriKvOptional->GetViewShape().GetDim(0)), + "When layout_kv is TND and seqused_ori_kv is passed, " + "the shape size of cu_seqlens_ori_kv minus 1 must be " + "equal to the shape size of seqused_ori_kv"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 ori_topk_length 维度一致性 + if (oriTopk != 0 && oriMaskMode == static_cast(SparseModeSmla::DEFAULT_MASK) && + IsTensorExistSmla(oriTopkLengthOptional)) { + if (strcmp(layoutQOptional, "BSND") == 0) { + // 校验 ori_topk_length 第一个维度 + aclnnStatus ret = + CheckTopkLengthFirstDimSmla(oriTopkLengthOptional, "ori_topk_length", queryBatchSize, querySource); + if (ret != ACLNN_SUCCESS) { + return ret; + } + // 校验 ori_topk_length 第二个维度 + ret = + CheckTopkLengthSingleDimSmla(oriTopkLengthOptional, "ori_topk_length", SMLA_TOPK_LENGTH_SECOND_AXIS, + maxSeqlenQ, "max_seqlen_q", layoutQOptional); + if (ret != ACLNN_SUCCESS) { + return ret; + } + // 校验 ori_topk_length 第三个维度 + ret = + CheckTopkLengthSingleDimSmla(oriTopkLengthOptional, "ori_topk_length", SMLA_TOPK_LENGTH_THIRD_AXIS, + numHeadsKv, "num_heads_kv", layoutQOptional); + if (ret != ACLNN_SUCCESS) { + return ret; + } + } else if (strcmp(layoutQOptional, "TND") == 0) { + // 校验 ori_topk_length 第二个维度 + aclnnStatus ret = + CheckTopkLengthSingleDimSmla(oriTopkLengthOptional, "ori_topk_length", SMLA_TOPK_LENGTH_SECOND_AXIS, + numHeadsKv, "num_heads_kv", layoutQOptional); + if (ret != ACLNN_SUCCESS) { + return ret; + } + } + } + } + if (hasCmpKv) { + std::string cmpKvSource; + int64_t cmpKvBatchSize = GetCmpKvBatchSizeSmla(sequsedCmpKvOptional, cuSeqlensCmpKvOptional, layoutKvOptional, + batchSize, &cmpKvSource); + // 校验q与cmp_kv维度一致性 + if (queryBatchSize != cmpKvBatchSize) { + if (IsTensorSourceSmla(querySource) && IsTensorSourceSmla(cmpKvSource)) { + OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON( + SMLA_ACLNN_OP_NAME, querySource + " and " + cmpKvSource, + std::to_string(GetRawShapeSizeSmla(querySource, queryBatchSize)) + " and " + + std::to_string(GetRawShapeSizeSmla(cmpKvSource, cmpKvBatchSize)), + "When has_cmp_kv is true, " + GetSourceDescSmla(querySource) + " must be equal to " + + GetSourceDescSmla(cmpKvSource)); + } else if (IsTensorSourceSmla(querySource)) { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + SMLA_ACLNN_OP_NAME, querySource, std::to_string(GetRawShapeSizeSmla(querySource, queryBatchSize)), + "When has_cmp_kv is true, " + GetSourceDescSmla(querySource) + " must be equal to batch_size"); + } else { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + SMLA_ACLNN_OP_NAME, cmpKvSource, std::to_string(GetRawShapeSizeSmla(cmpKvSource, cmpKvBatchSize)), + "When has_cmp_kv is true, " + GetSourceDescSmla(cmpKvSource) + " must be equal to batch_size"); + } + return ACLNN_ERR_PARAM_INVALID; + } + // 校验TND场景cmp_kv维度一致性 + if (strcmp(layoutKvOptional, "TND") == 0 && IsTensorExistSmla(sequsedCmpKvOptional)) { + int64_t cuSeqlensCmpKvBatchSize = cuSeqlensCmpKvOptional->GetViewShape().GetDim(0) - 1; + if (cuSeqlensCmpKvBatchSize != cmpKvBatchSize) { + OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON( + SMLA_ACLNN_OP_NAME, "cu_seqlens_cmp_kv and seqused_cmp_kv", + std::to_string(cuSeqlensCmpKvOptional->GetViewShape().GetDim(0)) + " and " + + std::to_string(sequsedCmpKvOptional->GetViewShape().GetDim(0)), + "When layout_kv is TND and seqused_cmp_kv is passed, " + "the shape size of cu_seqlens_cmp_kv minus 1 must be " + "equal to the shape size of seqused_cmp_kv"); + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 cmp_residual_kv 元素数 + if (IsTensorExistSmla(cmpResidualKvOptional)) { + if (cmpResidualKvOptional->GetViewShape().GetDim(0) != queryBatchSize) { + if (IsTensorSourceSmla(querySource)) { + OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON( + SMLA_ACLNN_OP_NAME, "cmp_residual_kv and " + querySource, + std::to_string(cmpResidualKvOptional->GetViewShape().GetDim(0)) + " and " + + std::to_string(GetRawShapeSizeSmla(querySource, queryBatchSize)), + "The shape size of cmp_residual_kv must be equal to " + GetSourceDescSmla(querySource)); + } else { + OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON( + SMLA_ACLNN_OP_NAME, "cmp_residual_kv", + std::to_string(cmpResidualKvOptional->GetViewShape().GetDim(0)), + "The shape size of cmp_residual_kv must be equal " + "to batch_size"); + } + return ACLNN_ERR_PARAM_INVALID; + } + } + // 校验 cmp_topk_length 维度一致性 + if (cmpTopk != 0 && cmpMaskMode == static_cast(SparseModeSmla::DEFAULT_MASK) && + IsTensorExistSmla(cmpTopkLengthOptional)) { + if (strcmp(layoutQOptional, "BSND") == 0) { + // 校验 cmp_topk_length 第一个维度 + aclnnStatus ret = + CheckTopkLengthFirstDimSmla(cmpTopkLengthOptional, "cmp_topk_length", queryBatchSize, querySource); + if (ret != ACLNN_SUCCESS) { + return ret; + } + // 校验 cmp_topk_length 第二个维度 + ret = + CheckTopkLengthSingleDimSmla(cmpTopkLengthOptional, "cmp_topk_length", SMLA_TOPK_LENGTH_SECOND_AXIS, + maxSeqlenQ, "max_seqlen_q", layoutQOptional); + if (ret != ACLNN_SUCCESS) { + return ret; + } + // 校验 cmp_topk_length 第三个维度 + ret = + CheckTopkLengthSingleDimSmla(cmpTopkLengthOptional, "cmp_topk_length", SMLA_TOPK_LENGTH_THIRD_AXIS, + numHeadsKv, "num_heads_kv", layoutQOptional); + if (ret != ACLNN_SUCCESS) { + return ret; + } + } else if (strcmp(layoutQOptional, "TND") == 0) { + // 校验 cmp_topk_length 第二个维度 + aclnnStatus ret = + CheckTopkLengthSingleDimSmla(cmpTopkLengthOptional, "cmp_topk_length", SMLA_TOPK_LENGTH_SECOND_AXIS, + numHeadsKv, "num_heads_kv", layoutQOptional); + if (ret != ACLNN_SUCCESS) { + return ret; + } + } + } + } + return ACLNN_SUCCESS; +} + +static aclnnStatus ParamsCheck(const aclTensor *cuSeqlensQOptional, const aclTensor *cuSeqlensOriKvOptional, + const aclTensor *cuSeqlensCmpKvOptional, const aclTensor *sequsedQOptional, + const aclTensor *sequsedOriKvOptional, const aclTensor *sequsedCmpKvOptional, + const aclTensor *cmpResidualKvOptional, const aclTensor *oriTopkLengthOptional, + const aclTensor *cmpTopkLengthOptional, int64_t numHeadsQ, int64_t numHeadsKv, + int64_t headDim, int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenOriKv, + int64_t maxSeqlenCmpKv, int64_t oriTopk, int64_t cmpTopk, int64_t cmpRatio, + int64_t oriMaskMode, int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight, + const char *layoutQOptional, const char *layoutKvOptional, bool hasOriKv, bool hasCmpKv, + uint32_t aicCoreNum, uint32_t aivCoreNum, const char *socVersion, + const aclTensor *metaData) +{ + if (CheckSingleParamSmla(batchSize, maxSeqlenQ, maxSeqlenOriKv, maxSeqlenCmpKv, numHeadsQ, numHeadsKv, headDim, + oriTopk, cmpTopk, cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, + layoutQOptional, layoutKvOptional, hasOriKv, hasCmpKv, aicCoreNum, aivCoreNum, + socVersion) == ACLNN_SUCCESS && + CheckExistenceSmla(cuSeqlensQOptional, cuSeqlensOriKvOptional, cuSeqlensCmpKvOptional, sequsedOriKvOptional, + sequsedCmpKvOptional, cmpResidualKvOptional, oriTopkLengthOptional, cmpTopkLengthOptional, + oriTopk, cmpTopk, cmpRatio, oriMaskMode, cmpMaskMode, hasOriKv, hasCmpKv, layoutQOptional, + layoutKvOptional, socVersion, metaData) == ACLNN_SUCCESS && + CheckConsistencySmla(cuSeqlensQOptional, cuSeqlensOriKvOptional, cuSeqlensCmpKvOptional, sequsedQOptional, + sequsedOriKvOptional, sequsedCmpKvOptional, cmpResidualKvOptional, oriTopkLengthOptional, + cmpTopkLengthOptional, batchSize, layoutQOptional, layoutKvOptional, hasOriKv, hasCmpKv, + oriTopk, cmpTopk, oriMaskMode, cmpMaskMode, maxSeqlenQ, numHeadsKv, socVersion, + metaData) == ACLNN_SUCCESS) { + return ACLNN_SUCCESS; + } else { + return ACLNN_ERR_PARAM_INVALID; + } +} +} // namespace + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/CMakeLists.txt b/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/CMakeLists.txt new file mode 100644 index 000000000000..8577ab6088c5 --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/CMakeLists.txt @@ -0,0 +1,26 @@ +# --------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# --------------------------------------------------------------------------------------------------------- + +if (BUILD_WITH_INSTALLED_DEPENDENCY_CANN_PKG) + if (NOT (UT_TEST_ALL OR OP_KERNEL_AICPU_UT)) + add_definitions(-D_GLIBCXX_USE_CXX11_ABI=1) + set(CMAKE_CXX_COMPILER ${ASCEND_DIR}/toolkit/toolchain/hcc/bin/aarch64-target-linux-gnu-g++) + endif() + + file(GLOB_RECURSE JSON_FILE ${CMAKE_CURRENT_SOURCE_DIR}/*.json) + file(GLOB AICPU_SRC ${CMAKE_CURRENT_SOURCE_DIR}/*_aicpu*.cpp) + message(STATUS "[sparse_flash_mla_metadata] Found aicpu sources: ${AICPU_SRC}, ascend dir: ${ASCEND_DIR}, ophsot name: ${OPHOST_NAME}") + + add_aicpu_cust_kernel_modules(sparse_flash_mla_metadata ${AICPU_SRC} ${JSON_FILE}) +endif() + +if(UT_TEST_ALL OR OP_KERNEL_AICPU_UT) + AddAicpuOpTestCase(sparse_flash_mla_metadata) +endif() diff --git a/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.cpp b/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.cpp new file mode 100644 index 000000000000..99b456c5bc9e --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.cpp @@ -0,0 +1,1638 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_metadata_aicpu.cpp + * \brief + */ + +#include "log.h" +#include "status.h" +#include +#include +#include "sparse_flash_mla_metadata_aicpu.h" + +using namespace optiling; + +namespace aicpu { +uint32_t SparseFlashMlaMetadataCpuKernel::Compute(CpuKernelContext &ctx) +{ + bool success = Prepare(ctx); + if (!success) { + return KERNEL_STATUS_PARAM_INVALID; + } + SplitResult splitRes{aicCoreNum_, aivCoreNum_}; + success = BalanceSchedule(splitRes) && GenMetadata(splitRes); + return success ? KERNEL_STATUS_OK : KERNEL_STATUS_PARAM_INVALID; +} + +bool SparseFlashMlaMetadataCpuKernel::Prepare(CpuKernelContext &ctx) +{ + // input + cuSeqlensQ_ = ctx.Input(static_cast(ParamId::cuSeqlensQ)); + cuSeqlensOriKv_ = ctx.Input(static_cast(ParamId::cuSeqlensOriKv)); + cuSeqlensCmpKv_ = ctx.Input(static_cast(ParamId::cuSeqlensCmpKv)); + sequsedQ_ = ctx.Input(static_cast(ParamId::sequsedQ)); + sequsedOriKv_ = ctx.Input(static_cast(ParamId::sequsedOriKv)); + sequsedCmpKv_ = ctx.Input(static_cast(ParamId::sequsedCmpKv)); + cmpResidualKv_ = ctx.Input(static_cast(ParamId::cmpResidualKv)); + oriTopkLength_ = ctx.Input(static_cast(ParamId::oriTopkLength)); + cmpTopkLength_ = ctx.Input(static_cast(ParamId::cmpTopkLength)); + hasOriTopkLength_ = (oriTopkLength_ != nullptr && oriTopkLength_->GetData() != nullptr && + oriTopkLength_->NumElements() > 0); + // output + metadata_ = ctx.Output(static_cast(ParamId::metaData)); + + bool requiredAttrs = GetAttrValue(ctx, "num_heads_q", numHeadsQ_) && + GetAttrValue(ctx, "num_heads_kv", numHeadsKv_) && GetAttrValue(ctx, "head_dim", headDim_); + if (!requiredAttrs) { + return false; + } + + // attributes optional + GetAttrValueOpt(ctx, "soc_version", socVersion_); + GetAttrValueOpt(ctx, "aic_core_num", aicCoreNum_); + GetAttrValueOpt(ctx, "aiv_core_num", aivCoreNum_); + GetAttrValueOpt(ctx, "batch_size", batchSize_); + GetAttrValueOpt(ctx, "max_seqlen_q", maxSeqlenQ_); + GetAttrValueOpt(ctx, "max_seqlen_ori_kv", maxSeqlenOriKv_); + GetAttrValueOpt(ctx, "max_seqlen_cmp_kv", maxSeqlenCmpKv_); + GetAttrValueOpt(ctx, "ori_topk", oriTopK_); + GetAttrValueOpt(ctx, "cmp_topk", cmpTopK_); + GetAttrValueOpt(ctx, "cmp_ratio", cmpRatio_); + GetAttrValueOpt(ctx, "ori_mask_mode", oriMaskMode_); + GetAttrValueOpt(ctx, "cmp_mask_mode", cmpMaskMode_); + GetAttrValueOpt(ctx, "ori_win_left", oriWinLeft_); + GetAttrValueOpt(ctx, "ori_win_right", oriWinRight_); + GetAttrValueOpt(ctx, "layout_q", layoutQ_); + GetAttrValueOpt(ctx, "layout_kv", layoutKv_); + GetAttrValueOpt(ctx, "has_ori_kv", hasOriKv_); + GetAttrValueOpt(ctx, "has_cmp_kv", hasCmpKv_); + GetAttrValueOpt(ctx, "is_batch_consistency", isBatchConsistency_); + return (ParamsCheck() && ParamsInit()); +} + +bool SparseFlashMlaMetadataCpuKernel::ParamsCheck() +{ + // 校验输出 metadata 是否为空 + if (metadata_ == nullptr) { + KERNEL_LOG_ERROR("Output metadata is nullptr"); + return false; + } else if (metadata_->GetData() == nullptr) { + KERNEL_LOG_ERROR("Output metadata data is nullptr"); + return false; + } + int32_t batchSize = GetQueryBatchSize(); + // 校验 cu_seqlens_q 元素 + if (layoutQ_ == "TND") { + if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) { + const int32_t *cuSeqlensQPtr = static_cast(cuSeqlensQ_->GetData()); + // 校验 cu_seqlens_q 首元素为 0 + if (cuSeqlensQPtr[0] != 0) { + KERNEL_LOG_ERROR("The first element of cu_seqlens_q should be 0, but got %d", cuSeqlensQPtr[0]); + return false; + } + for (int i = 0; i < batchSize + 1; i++) { + // 校验 cu_seqlens_q 元素递增 + if (i > 0 && cuSeqlensQPtr[i - 1] > cuSeqlensQPtr[i]) { + KERNEL_LOG_ERROR("The elements in cu_seqlens_q must be in ascending order, " + "but got cu_seqlens_q[%d] = %d, cu_seqlens_q[%d] = %d", + i - 1, cuSeqlensQPtr[i - 1], i, cuSeqlensQPtr[i]); + return false; + } + } + } + } + // 校验 seqused_q 元素 + if (sequsedQ_ != nullptr && sequsedQ_->GetData() != nullptr) { + const int32_t *sequsedQPtr = static_cast(sequsedQ_->GetData()); + const int32_t *cuSeqlensQPtr = (layoutQ_ == "TND" && cuSeqlensQ_ != nullptr && + cuSeqlensQ_->GetData() != nullptr) ? + static_cast(cuSeqlensQ_->GetData()) : + nullptr; + for (int i = 0; i < batchSize; i++) { + // 校验 seqused_q 元素非负 + if (sequsedQPtr[i] < 0) { + KERNEL_LOG_ERROR("The elements in seqused_q should be >= 0, but got seqused_q[%d] = %d", i, + sequsedQPtr[i]); + return false; + } + // 校验 seqused_q 元素不大于 max_seqlen_q (BSND) 或 cu_seqlens_q 序列长度 (TND) + if (layoutQ_ == "BSND" && sequsedQPtr[i] > maxSeqlenQ_) { + KERNEL_LOG_ERROR("The elements in seqused_q should not be greater than max_seqlen_q %d, " + "but got seqused_q[%d] = %d", + maxSeqlenQ_, i, sequsedQPtr[i]); + return false; + } + if (cuSeqlensQPtr != nullptr) { + int32_t seqLen = cuSeqlensQPtr[i + 1] - cuSeqlensQPtr[i]; + if (sequsedQPtr[i] > seqLen) { + KERNEL_LOG_ERROR("The elements in seqused_q should not be greater than the sequence length " + "from cu_seqlens_q %d, but got seqused_q[%d] = %d", + seqLen, i, sequsedQPtr[i]); + return false; + } + } + } + } + if (hasOriKv_) { + // 校验 cu_seqlens_ori_kv 元素 + if (layoutKv_ == "TND") { + if (cuSeqlensOriKv_ != nullptr && cuSeqlensOriKv_->GetData() != nullptr) { + const int32_t *cuSeqlensOriKvPtr = static_cast(cuSeqlensOriKv_->GetData()); + // 校验 cu_seqlens_ori_kv 首元素为 0 + if (cuSeqlensOriKvPtr[0] != 0) { + KERNEL_LOG_ERROR("The first element of cu_seqlens_ori_kv should be 0, but got %d", + cuSeqlensOriKvPtr[0]); + return false; + } + for (int i = 0; i < batchSize + 1; i++) { + // 校验 cu_seqlens_ori_kv 元素递增 + if (i > 0 && cuSeqlensOriKvPtr[i - 1] > cuSeqlensOriKvPtr[i]) { + KERNEL_LOG_ERROR("The elements in cu_seqlens_ori_kv must be in ascending order, " + "but got cu_seqlens_ori_kv[%d] = %d, cu_seqlens_ori_kv[%d] = %d", + i - 1, cuSeqlensOriKvPtr[i - 1], i, cuSeqlensOriKvPtr[i]); + return false; + } + } + } + } + // 校验 seqused_ori_kv 元素 + if (sequsedOriKv_ != nullptr && sequsedOriKv_->GetData() != nullptr) { + const int32_t *sequsedOriKvPtr = static_cast(sequsedOriKv_->GetData()); + const int32_t *cuSeqlensOriKvPtr = (layoutKv_ == "TND" && cuSeqlensOriKv_ != nullptr && + cuSeqlensOriKv_->GetData() != nullptr) ? + static_cast(cuSeqlensOriKv_->GetData()) : + nullptr; + for (int i = 0; i < batchSize; i++) { + // 校验 seqused_ori_kv 元素非负 + if (sequsedOriKvPtr[i] < 0) { + KERNEL_LOG_ERROR("The elements in seqused_ori_kv should be >= 0, but got seqused_ori_kv[%d] = %d", + i, sequsedOriKvPtr[i]); + return false; + } + // 校验 seqused_ori_kv 元素不大于 max_seqlen_ori_kv (BSND) 或 cu_seqlens_ori_kv 序列长度 (TND) + if (layoutKv_ == "BSND" && sequsedOriKvPtr[i] > maxSeqlenOriKv_) { + KERNEL_LOG_ERROR("The elements in seqused_ori_kv should not be greater than " + "max_seqlen_ori_kv %d, but got seqused_ori_kv[%d] = %d", + maxSeqlenOriKv_, i, sequsedOriKvPtr[i]); + return false; + } + if (cuSeqlensOriKvPtr != nullptr) { + int32_t seqLen = cuSeqlensOriKvPtr[i + 1] - cuSeqlensOriKvPtr[i]; + if (sequsedOriKvPtr[i] > seqLen) { + KERNEL_LOG_ERROR("The elements in seqused_ori_kv should not be greater than the sequence " + "length from cu_seqlens_ori_kv %d, but got seqused_ori_kv[%d] = %d", + seqLen, i, sequsedOriKvPtr[i]); + return false; + } + } + } + } + // 校验 ori_topk_length 元素 + if (oriTopK_ != 0 && oriMaskMode_ == static_cast(SparseMode::DEFAULT_MASK) && + oriTopkLength_ != nullptr && oriTopkLength_->GetData() != nullptr) { + // 校验 ori_topk_length 元素数量 + int32_t sumOfQuerySeq = GetSumOfQuerySeq(); + const int32_t *oriTopkLengthPtr = static_cast(oriTopkLength_->GetData()); + auto oriTopkLengthShape = oriTopkLength_->GetTensorShape(); + int32_t oriTopkLengthSize = layoutQ_ == "TND" ? + oriTopkLengthShape->GetDimSize(0) * oriTopkLengthShape->GetDimSize(1) : + oriTopkLengthShape->GetDimSize(0) * oriTopkLengthShape->GetDimSize(1) * + oriTopkLengthShape->GetDimSize(2); + if (oriTopkLengthSize < sumOfQuerySeq) { + KERNEL_LOG_ERROR("The size of ori_topk_length %d should not be smaller than " + "the sum of query sequence %d!", + oriTopkLengthSize, sumOfQuerySeq); + return false; + } + // 校验 ori_topk_length 元素非负 + for (int i = 0; i < oriTopkLengthSize; i++) { + if (oriTopkLengthPtr[i] < 0) { + KERNEL_LOG_ERROR("The elements in ori_topk_length should be >= 0, but got ori_topk_length[%d] = %d", + i, oriTopkLengthPtr[i]); + return false; + } + } + } + } + if (hasCmpKv_) { + if (layoutKv_ == "TND") { + // 校验 cu_seqlens_cmp_kv 元素 + if (cuSeqlensCmpKv_ != nullptr && cuSeqlensCmpKv_->GetData() != nullptr) { + const int32_t *cuSeqlensCmpKvPtr = static_cast(cuSeqlensCmpKv_->GetData()); + // 校验 cu_seqlens_cmp_kv 首元素为 0 + if (cuSeqlensCmpKvPtr[0] != 0) { + KERNEL_LOG_ERROR("The first element of cu_seqlens_cmp_kv should be 0, but got %d", + cuSeqlensCmpKvPtr[0]); + return false; + } + for (int i = 0; i < batchSize + 1; i++) { + // 校验 cu_seqlens_cmp_kv 元素递增 + if (i > 0 && cuSeqlensCmpKvPtr[i - 1] > cuSeqlensCmpKvPtr[i]) { + KERNEL_LOG_ERROR("The elements in cu_seqlens_cmp_kv must be in ascending order, " + "but got cu_seqlens_cmp_kv[%d] = %d, cu_seqlens_cmp_kv[%d] = %d", + i - 1, cuSeqlensCmpKvPtr[i - 1], i, cuSeqlensCmpKvPtr[i]); + return false; + } + } + } + } + // 校验 seqused_cmp_kv 元素 + if (sequsedCmpKv_ != nullptr && sequsedCmpKv_->GetData() != nullptr) { + const int32_t *sequsedCmpKvPtr = static_cast(sequsedCmpKv_->GetData()); + const int32_t *cuSeqlensCmpKvPtr = (layoutKv_ == "TND" && cuSeqlensCmpKv_ != nullptr && + cuSeqlensCmpKv_->GetData() != nullptr) ? + static_cast(cuSeqlensCmpKv_->GetData()) : + nullptr; + for (int i = 0; i < batchSize; i++) { + // 校验 seqused_cmp_kv 元素非负 + if (sequsedCmpKvPtr[i] < 0) { + KERNEL_LOG_ERROR("The elements in seqused_cmp_kv should be >= 0, but got seqused_cmp_kv[%d] = %d", + i, sequsedCmpKvPtr[i]); + return false; + } + // 校验 seqused_cmp_kv 元素不大于 max_seqlen_cmp_kv (BSND) 或 cu_seqlens_cmp_kv 序列长度 (TND) + if (layoutKv_ == "BSND" && sequsedCmpKvPtr[i] > maxSeqlenCmpKv_) { + KERNEL_LOG_ERROR("The elements in seqused_cmp_kv should not be greater than " + "max_seqlen_cmp_kv %d, but got seqused_cmp_kv[%d] = %d", + maxSeqlenCmpKv_, i, sequsedCmpKvPtr[i]); + return false; + } + if (cuSeqlensCmpKvPtr != nullptr) { + int32_t seqLen = cuSeqlensCmpKvPtr[i + 1] - cuSeqlensCmpKvPtr[i]; + if (sequsedCmpKvPtr[i] > seqLen) { + KERNEL_LOG_ERROR("The elements in seqused_cmp_kv should not be greater than the sequence " + "length from cu_seqlens_cmp_kv %d, but got seqused_cmp_kv[%d] = %d", + seqLen, i, sequsedCmpKvPtr[i]); + return false; + } + } + } + } + // 校验 cmp_residual_kv 元素 + if (cmpResidualKv_ != nullptr && cmpResidualKv_->GetData() != nullptr) { + const int32_t *cmpResidualKvPtr = static_cast(cmpResidualKv_->GetData()); + for (int i = 0; i < batchSize; i++) { + if (cmpResidualKvPtr[i] < 0 || cmpResidualKvPtr[i] >= cmpRatio_) { + KERNEL_LOG_ERROR("The elements in cmp_residual_kv should be in [0, cmpRatio_(%d)), but got " + "cmp_residual_kv[%d] = %d", + cmpRatio_, + i, cmpResidualKvPtr[i]); + return false; + } + } + } + // 校验 cmp_topk_length 元素 + if (cmpTopK_ != 0 && cmpMaskMode_ == static_cast(SparseMode::DEFAULT_MASK) && + cmpTopkLength_ != nullptr && cmpTopkLength_->GetData() != nullptr) { + // 校验 cmp_topk_length 元素数量 + int32_t sumOfQuerySeq = GetSumOfQuerySeq(); + const int32_t *cmpTopkLengthPtr = static_cast(cmpTopkLength_->GetData()); + auto cmpTopkLengthShape = cmpTopkLength_->GetTensorShape(); + int32_t cmpTopkLengthSize = layoutQ_ == "TND" ? + cmpTopkLengthShape->GetDimSize(0) * cmpTopkLengthShape->GetDimSize(1) : + cmpTopkLengthShape->GetDimSize(0) * cmpTopkLengthShape->GetDimSize(1) * + cmpTopkLengthShape->GetDimSize(2); + if (cmpTopkLengthSize < sumOfQuerySeq) { + KERNEL_LOG_ERROR("The size of cmp_topk_length %d should not be smaller than " + "the sum of query sequence %d!", + cmpTopkLengthSize, sumOfQuerySeq); + return false; + } + // 校验 cmp_topk_length 元素非负 + for (int i = 0; i < cmpTopkLengthSize; i++) { + if (cmpTopkLengthPtr[i] < 0) { + KERNEL_LOG_ERROR("The elements in cmp_topk_length should be >= 0, but got cmp_topk_length[%d] = %d", + i, cmpTopkLengthPtr[i]); + return false; + } + } + } + } + return true; +} + +int32_t SparseFlashMlaMetadataCpuKernel::GetSumOfQuerySeq() +{ + int32_t batchSize = GetQueryBatchSize(); + // 如果sequsedQ_ 传了,使用sequsedQ_获取 BsSize + if (sequsedQ_ != nullptr && sequsedQ_->GetData() != nullptr) { + if (sequsedQ_->GetTensorShape() != nullptr) { + const int32_t *seqUsedPtr = static_cast(sequsedQ_->GetData()); + int32_t queryBsSize = 0; + for (int i = 0; i < batchSize; i++) { + queryBsSize += seqUsedPtr[i]; + } + return queryBsSize; + } + } + // sequsedQ_ 没传,判断 Layout + if (layoutQ_ == "TND") { + // 如果是 TND,尝试使用 cuSeqlensQ_获取 BsSize + if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) { + if (cuSeqlensQ_->GetTensorShape() != nullptr) { + const int32_t *s1Ptr = static_cast(cuSeqlensQ_->GetData()); + return s1Ptr[batchSize]; + } + } + } + // 如果不是 TND,或者 cuSeqlensQ_ 为空,使用shape信息计算 BsSize + return batchSize_ * maxSeqlenQ_; +} + +int32_t SparseFlashMlaMetadataCpuKernel::GetQueryBatchSize() +{ + // 1. 如果sequsedQ_ 传了,使用sequsedQ_获取BatchSize + if (sequsedQ_ != nullptr && sequsedQ_->GetData() != nullptr) { + if (sequsedQ_->GetTensorShape() != nullptr) { + return sequsedQ_->GetTensorShape()->GetDimSize(0); + } + } + // 2. sequsedQ_ 没传,判断 Layout + if (layoutQ_ == "TND") { + // 如果是 TND,尝试使用 cuSeqlensQ_获取BatchSize + if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) { + if (cuSeqlensQ_->GetTensorShape() != nullptr) { + return cuSeqlensQ_->GetTensorShape()->GetDimSize(0) - 1; + } + } + } + // 3. 如果不是 TND,或者 cuSeqlensQ_ 为空,使用batchSize_ + return batchSize_; +} + +void SparseFlashMlaMetadataCpuKernel::CalcOriMaskMode() +{ + if (oriMaskMode_ == static_cast(SparseMode::DEFAULT_MASK)) { + oriPreToken_ = INT64_MAX; + oriNextToken_ = INT64_MAX; + oriAttentionMode_ = NO_MASK; + } else if (oriMaskMode_ == static_cast(SparseMode::RIGHT_DOWN_CAUSAL)) { + oriPreToken_ = INT64_MAX; + oriNextToken_ = 0; + oriAttentionMode_ = HAS_MASK; + } else { // SparseMode = 4 + oriPreToken_ = (oriWinLeft_ > -1) ? oriWinLeft_ : INT64_MAX; + oriNextToken_ = (oriWinRight_ > -1) ? oriWinRight_ : INT64_MAX; + oriAttentionMode_ = HAS_MASK; + } +} + +void SparseFlashMlaMetadataCpuKernel::CalcCmpMaskMode() +{ + if (cmpMaskMode_ == static_cast(SparseMode::DEFAULT_MASK)) { + cmpPreToken_ = INT64_MAX; + cmpNextToken_ = INT64_MAX; + cmpAttentionMode_ = NO_MASK; + } else if (cmpMaskMode_ == static_cast(SparseMode::RIGHT_DOWN_CAUSAL)) { + cmpPreToken_ = INT64_MAX; + cmpNextToken_ = 0; + cmpAttentionMode_ = HAS_MASK; + } else { // SparseMode = 4 + cmpPreToken_ = (oriWinLeft_ > -1) ? oriWinLeft_ : INT64_MAX; + cmpNextToken_ = (oriWinRight_ > -1) ? oriWinRight_ : INT64_MAX; + cmpAttentionMode_ = HAS_MASK; + } +} + +ValidSocVersion SparseFlashMlaMetadataCpuKernel::ProcessSocVersion() +{ + const std::string ascend950 = "Ascend950"; + if (socVersion_.find(ascend950) != std::string::npos) { + return ValidSocVersion::ASCEND950; + } else { + return ValidSocVersion::ASCEND910; + } +} + +bool SparseFlashMlaMetadataCpuKernel::ParamsInit() +{ + batchSize_ = GetQueryBatchSize(); + CalcOriMaskMode(); + CalcCmpMaskMode(); + isS1G_ = (layoutQ_ == "BSND" || layoutQ_ == "BSH" || layoutQ_ == "TND"); + if (numHeadsKv_ == 0) { + KERNEL_LOG_ERROR("num_heads_kv should not be 0."); + return false; + } + ValidSocVersion validSocVersion = ProcessSocVersion(); + groupSize_ = numHeadsQ_ / numHeadsKv_; + if (hasOriKv_ && oriTopK_ != 0) { + isSparseOriKv_ = true; + } + if (hasCmpKv_ && cmpTopK_ != 0) { + isSparseCmpKv_ = true; + } + if (isBatchConsistency_) { + supportFd_ = true; + } + if (validSocVersion == ValidSocVersion::ASCEND910) { + if (isSparseOriKv_ && !isSparseCmpKv_) { + mBaseSize_ = groupSize_; + } else { + mBaseSize_ = isSparseCmpKv_ ? groupSize_ : (256U / groupSize_) * groupSize_; + } + s2BaseSize_ = 512U; + } else if (validSocVersion == ValidSocVersion::ASCEND950) { + if (groupSize_ > 64U) { + isSplitG_ = true; + aicCoreNum_ /= 2; + } + mBaseSize_ = groupSize_; + s2BaseSize_ = 128U; + } else { + mBaseSize_ = groupSize_; + s2BaseSize_ = 128U; + } + return true; +} + +uint32_t SparseFlashMlaMetadataCpuKernel::GetS1Idx(uint32_t s1Size, uint32_t s1GIdx) +{ + uint32_t s1GToken = s1GIdx * mBaseSize_; + uint32_t s1Idx = 0; + if (isS1G_) { + s1Idx = s1GToken / static_cast(groupSize_); + } else { + s1Idx = s1GToken % static_cast(s1Size); + } + return s1Idx; +} + +uint32_t SparseFlashMlaMetadataCpuKernel::GetBsStride(uint32_t bIdx, uint32_t s1Idx) +{ + uint32_t bsStride = 0; + if (layoutQ_ == "TND") { + if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) { + const int32_t *s1Ptr = static_cast(cuSeqlensQ_->GetData()); + bsStride = s1Ptr[bIdx] + s1Idx; + return bsStride; + } + } + bsStride = bIdx * static_cast(maxSeqlenQ_) + s1Idx; + return bsStride; +} + +uint32_t SparseFlashMlaMetadataCpuKernel::GetOriTopkLength(uint32_t bsStride) +{ + // 尝试使用 oriTopkLength_ + if (oriTopK_ != 0 && oriMaskMode_ == static_cast(SparseMode::DEFAULT_MASK) && + oriTopkLength_ != nullptr && oriTopkLength_->GetData() != nullptr) { + const int32_t *oriTopkPtr = static_cast(oriTopkLength_->GetData()); + return static_cast(oriTopkPtr[bsStride]); + } + // 如果不是 DEFAULT_MASK,使用 oriTopK_ + return static_cast(oriTopK_); +} + +uint32_t SparseFlashMlaMetadataCpuKernel::ReadOriTopkLengthAtRow(uint32_t bIdx, uint32_t s1Idx) const +{ + if (!hasOriTopkLength_) { + return static_cast(oriTopK_); + } + const int32_t *topkLenPtr = static_cast(oriTopkLength_->GetData()); + if (layoutQ_ == "TND") { + uint32_t tIdx = s1Idx; + if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) { + const int32_t *cuPtr = static_cast(cuSeqlensQ_->GetData()); + tIdx = static_cast(cuPtr[bIdx]) + s1Idx; + } else { + tIdx = bIdx * static_cast(maxSeqlenQ_) + s1Idx; + } + return static_cast(topkLenPtr[tIdx * static_cast(numHeadsKv_)]); + } + return static_cast( + topkLenPtr[(static_cast(bIdx) * static_cast(maxSeqlenQ_) + s1Idx) * + static_cast(numHeadsKv_)]); +} + +uint32_t SparseFlashMlaMetadataCpuKernel::GetOriTopkLength(uint32_t s1GIdx, const BatchCache &batchCache) const +{ + if (!hasOriTopkLength_) { + return static_cast(oriTopK_); + } + int64_t s1GFirstToken = static_cast(s1GIdx) * static_cast(mBaseSize_); + int64_t s1GLastToken = std::min(s1GFirstToken + static_cast(mBaseSize_), + static_cast(batchCache.s1Size) * static_cast(groupSize_)) - + 1; + int64_t s1First = s1GFirstToken / static_cast(groupSize_); + int64_t s1Last = s1GLastToken / static_cast(groupSize_); + uint32_t maxLen = 0U; + for (int64_t s1Idx = s1First; s1Idx <= s1Last; ++s1Idx) { + maxLen = std::max(maxLen, ReadOriTopkLengthAtRow(batchCache.bIdx, static_cast(s1Idx))); + } + return std::min(maxLen, static_cast(oriTopK_)); +} + +uint32_t SparseFlashMlaMetadataCpuKernel::GetCmpTopkLength(uint32_t bsStride) +{ + // 尝试使用 cmpTopkLength_ + if (cmpTopK_ != 0 && cmpMaskMode_ == static_cast(SparseMode::DEFAULT_MASK) && + cmpTopkLength_ != nullptr && cmpTopkLength_->GetData() != nullptr) { + const int32_t *cmpTopkPtr = static_cast(cmpTopkLength_->GetData()); + return static_cast(cmpTopkPtr[bsStride]); + } + // 如果不是 DEFAULT_MASK,使用 cmpTopK_ + return static_cast(cmpTopK_); +} + +uint32_t SparseFlashMlaMetadataCpuKernel::GetS1SeqSize(uint32_t bIdx) +{ + // 1. 如果 sequsedQ_ 传了,直接使用 + if (sequsedQ_ != nullptr && sequsedQ_->GetData() != nullptr) { + const int32_t *seqUsedPtr = static_cast(sequsedQ_->GetData()); + return static_cast(seqUsedPtr[bIdx]); + } + // 2. sequsedQ_ 没传,判断 Layout + if (layoutQ_ == "TND") { + // 如果是 TND,尝试使用 cuSeqlensQ_ + if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) { + const int32_t *s1Ptr = static_cast(cuSeqlensQ_->GetData()); + return static_cast(s1Ptr[bIdx + 1U] - s1Ptr[bIdx]); + } + } + // 3. 如果不是 TND,或者 cuSeqlensQ_ 为空,使用 maxSeqlenQ_ + return static_cast(maxSeqlenQ_); +} + +uint32_t SparseFlashMlaMetadataCpuKernel::GetOriS2SeqSize(uint32_t bIdx) +{ + // 如果 sequsedOriKv_ 传了,直接使用 + if (sequsedOriKv_ != nullptr && sequsedOriKv_->GetData() != nullptr) { + const int32_t *seqUsedPtr = static_cast(sequsedOriKv_->GetData()); + return static_cast(seqUsedPtr[bIdx]); + } + // sequsedOriKv_ 没传,判断 Layout + if (layoutKv_ == "TND") { + // 如果是 TND,尝试使用 cuSeqlensOriKv_ + if (cuSeqlensOriKv_ != nullptr && cuSeqlensOriKv_->GetData() != nullptr) { + const int32_t *s2Ptr = static_cast(cuSeqlensOriKv_->GetData()); + return static_cast(s2Ptr[bIdx + 1U] - s2Ptr[bIdx]); + } + } + // 如果是PA场景,或 max_seqlen_ori_kv 没传入,且 ori_kv 为稀疏的,则尝试从 topk 中获取 + if ((layoutKv_ == "PA_BBND" || maxSeqlenOriKv_ == 0) && isSparseOriKv_) { + return UINT32_MAX; + } + // 使用 max_seqlen_ori_kv + return static_cast(maxSeqlenOriKv_); +} + +uint32_t SparseFlashMlaMetadataCpuKernel::GetCmpS2SeqSize(uint32_t bIdx) +{ + // 如果 sequsedCmpKv_ 传了,直接使用 + if (sequsedCmpKv_ != nullptr && sequsedCmpKv_->GetData() != nullptr) { + const int32_t *seqUsedPtr = static_cast(sequsedCmpKv_->GetData()); + return static_cast(seqUsedPtr[bIdx]); + } + // sequsedCmpKv_ 没传,判断 Layout + if (layoutKv_ == "TND") { + // 如果是 TND,尝试使用 cuSeqlensCmpKv_ + if (cuSeqlensCmpKv_ != nullptr && cuSeqlensCmpKv_->GetData() != nullptr) { + const int32_t *s2Ptr = static_cast(cuSeqlensCmpKv_->GetData()); + return static_cast(s2Ptr[bIdx + 1U] - s2Ptr[bIdx]); + } + } + // 如果是PA场景,或 max_seqlen_cmp_kv 没传入,且 cmp_kv 为稀疏的,则尝试从topk中获取 + if ((layoutKv_ == "PA_BBND" || maxSeqlenCmpKv_ == 0) && isSparseCmpKv_) { + return UINT32_MAX; + } + // 使用 max_seqlen_cmp_kv + return static_cast(maxSeqlenCmpKv_); +} + +uint64_t SparseFlashMlaMetadataCpuKernel::GetRevertS2Size(uint32_t bIdx) +{ + uint32_t cmpS2Size = GetCmpS2SeqSize(bIdx); + if (cmpResidualKv_ != nullptr && cmpResidualKv_->GetData() != nullptr) { + const int32_t *residualPtr = static_cast(cmpResidualKv_->GetData()); + return static_cast(cmpS2Size) * static_cast(cmpRatio_) + residualPtr[bIdx]; + } else { + return static_cast(cmpS2Size) * static_cast(cmpRatio_); + } +} + +void SparseFlashMlaMetadataCpuKernel::CalcSplitInfo(SplitContext &splitContext) +{ + // 计算每个batch的切分,统计是否为空batch,记录最后有效batch(每个batch的每个N2切分是一样的) + SplitInfo &splitInfo = splitContext.splitInfo; + for (uint32_t bIdx = 0; bIdx < batchSize_; bIdx++) { + uint32_t s1Size = GetS1SeqSize(bIdx); + splitInfo.s1GBaseNum[bIdx] = (static_cast(s1Size) * groupSize_ + (mBaseSize_ - 1U)) / mBaseSize_; + splitInfo.s1GTailSize[bIdx] = (static_cast(s1Size) * groupSize_) % mBaseSize_; + if (hasOriKv_) { + uint32_t curOriS2Size = GetOriS2SeqSize(bIdx); + splitInfo.oriS2BaseNum[bIdx] = (static_cast(curOriS2Size) + s2BaseSize_ - 1U) / s2BaseSize_; + } + if (hasCmpKv_) { + uint32_t curCmpS2Size = GetCmpS2SeqSize(bIdx); + splitInfo.cmpS2BaseNum[bIdx] = (static_cast(curCmpS2Size) + s2BaseSize_ - 1U) / s2BaseSize_; + } + if (splitInfo.s1GBaseNum[bIdx] != 0U && + (splitInfo.oriS2BaseNum[bIdx] != 0U || splitInfo.cmpS2BaseNum[bIdx] != 0U)) { + splitInfo.isKvSeqAllZero = false; + } + } +} + +int64_t SparseFlashMlaMetadataCpuKernel::CalcOriPreTokenLeftUp(uint32_t s1Size, uint32_t s2Size) +{ + auto mode = static_cast(oriMaskMode_); + if (mode == SparseMode::BAND) { + return oriPreToken_ == INT64_MAX ? INT64_MAX : + static_cast(s1Size) - static_cast(s2Size) + oriPreToken_; + } + return oriPreToken_; +} + +int64_t SparseFlashMlaMetadataCpuKernel::CalcOriNextTokenLeftUp(uint32_t s1Size, uint32_t s2Size) +{ + auto mode = static_cast(oriMaskMode_); + switch (mode) { + case SparseMode::DEFAULT_MASK: + case SparseMode::ALL_MASK: + case SparseMode::LEFT_UP_CAUSAL: + return oriNextToken_; + case SparseMode::RIGHT_DOWN_CAUSAL: + return static_cast(s2Size) - static_cast(s1Size); + case SparseMode::BAND: + return oriNextToken_ == INT64_MAX ? + INT64_MAX : + static_cast(s2Size) - static_cast(s1Size) + oriNextToken_; + default: + return oriNextToken_; + } +} + +int64_t SparseFlashMlaMetadataCpuKernel::CalcCmpPreTokenLeftUp(uint32_t s1Size, uint64_t s2Size) +{ + auto mode = static_cast(cmpMaskMode_); + if (mode == SparseMode::BAND) { + return cmpPreToken_ == INT64_MAX ? INT64_MAX : + static_cast(s1Size) - static_cast(s2Size) + cmpPreToken_; + } + return cmpPreToken_; +} + +int64_t SparseFlashMlaMetadataCpuKernel::CalcCmpNextTokenLeftUp(uint32_t s1Size, uint64_t s2Size) +{ + auto mode = static_cast(cmpMaskMode_); + switch (mode) { + case SparseMode::DEFAULT_MASK: + case SparseMode::ALL_MASK: + case SparseMode::LEFT_UP_CAUSAL: + return cmpNextToken_; + case SparseMode::RIGHT_DOWN_CAUSAL: + return static_cast(s2Size) - static_cast(s1Size); + case SparseMode::BAND: + return cmpNextToken_ == INT64_MAX ? + INT64_MAX : + static_cast(s2Size) - static_cast(s1Size) + cmpNextToken_; + default: + return cmpNextToken_; + } +} + +int64_t SparseFlashMlaMetadataCpuKernel::OriCalcCost(uint32_t basicM, uint32_t basicS2) +{ + uint32_t oriAlignCoefM = 16U; + uint32_t oriAlignCoefS2 = 64U; + uint32_t oriAlignBasicM = (basicM + oriAlignCoefM - 1U) / oriAlignCoefM; + uint32_t oriAlignBasicS2 = (basicS2 + oriAlignCoefS2 - 1U) / oriAlignCoefS2; + return static_cast(COST_WEIGHT_M * oriAlignBasicM + COST_WEIGHT_S2 * oriAlignBasicS2); +} + +int64_t SparseFlashMlaMetadataCpuKernel::CmpCalcCost(uint32_t basicM, uint32_t basicS2) +{ + uint32_t cmpAlignCoefM = 16U; + uint32_t cmpAlignCoefS2 = 64U; + uint32_t cmpAlignBasicM = (basicM + cmpAlignCoefM - 1U) / cmpAlignCoefM; + uint32_t cmpAlignBasicS2 = (basicS2 + cmpAlignCoefS2 - 1U) / cmpAlignCoefS2; + return static_cast(COST_WEIGHT_M * cmpAlignBasicM + COST_WEIGHT_S2 * cmpAlignBasicS2); +} + +void SparseFlashMlaMetadataCpuKernel::CalcCostTable(uint32_t s1GTailSize, uint32_t reductionBlockSize, + uint32_t oriS2TailSize, uint32_t cmpS2TailSize) +{ + uint32_t normalS2Size = isBatchConsistency_ && reductionBlockSize > 0U ? reductionBlockSize : s2BaseSize_; + // ori 部分 cost + if (hasOriKv_) { + typeCost_[ORI_NORMAL_BLOCK][ORI_NORMAL_BLOCK] = OriCalcCost(mBaseSize_, normalS2Size); + typeCost_[ORI_TAIL_BLOCK][ORI_NORMAL_BLOCK] = + (s1GTailSize == 0U) ? 0U : OriCalcCost(s1GTailSize, normalS2Size); + typeCost_[ORI_NORMAL_BLOCK][ORI_TAIL_BLOCK] = (oriS2TailSize == 0U) ? 0U : + OriCalcCost(mBaseSize_, oriS2TailSize); + typeCost_[ORI_TAIL_BLOCK][ORI_TAIL_BLOCK] = (s1GTailSize == 0U || oriS2TailSize == 0U) ? 0U : + OriCalcCost(s1GTailSize, oriS2TailSize); + } + // cmp 部分 cost + if (hasCmpKv_) { + typeCost_[CMP_NORMAL_BLOCK][CMP_NORMAL_BLOCK] = CmpCalcCost(mBaseSize_, normalS2Size); + typeCost_[CMP_TAIL_BLOCK][CMP_NORMAL_BLOCK] = + (s1GTailSize == 0U) ? 0U : CmpCalcCost(s1GTailSize, normalS2Size); + typeCost_[CMP_NORMAL_BLOCK][CMP_TAIL_BLOCK] = (cmpS2TailSize == 0U) ? 0U : + CmpCalcCost(mBaseSize_, cmpS2TailSize); + typeCost_[CMP_TAIL_BLOCK][CMP_TAIL_BLOCK] = (s1GTailSize == 0U || cmpS2TailSize == 0U) ? 0U : + CmpCalcCost(s1GTailSize, cmpS2TailSize); + } +} + +Range SparseFlashMlaMetadataCpuKernel::CalcS2TokenRange(uint32_t s1GIdx, const BatchCache &batchCache, + bool isCmpKv) +{ + // actual seq == 0 + if (!isCmpKv) { + if (batchCache.s1Size == 0U || batchCache.oriS2Size == 0U) { + return std::make_pair(0, 0); + } + if (isSparseOriKv_ && !hasCmpKv_) { + uint32_t oriTopkSize = GetOriTopkLength(s1GIdx, batchCache); + if (oriTopkSize == 0U) { + return std::make_pair(0, 0); + } + return std::make_pair(0, static_cast(oriTopkSize) - 1); + } + } else { + if (batchCache.s1Size == 0U || batchCache.cmpRevertS2Size == 0U) { + return std::make_pair(0, 0); + } + } + + // no mask + uint32_t hasMask = 1; + int64_t s2Size = + isCmpKv ? static_cast(batchCache.cmpRevertS2Size) : static_cast(batchCache.oriS2Size); + hasMask = isCmpKv ? cmpAttentionMode_ : oriAttentionMode_; + if (!hasMask) { + return std::make_pair(0, s2Size - 1); + } + + // 1. calc index of s2FirstToken, s2LastToken by index of s1GFirstToken, s1GLastToken + int64_t s1GFirstToken = static_cast(s1GIdx) * static_cast(mBaseSize_); + int64_t s1GLastToken = std::min(s1GFirstToken + static_cast(mBaseSize_), + static_cast(batchCache.s1Size) * static_cast(groupSize_)) - + 1; + + int64_t s1FirstToken = 0; + int64_t s1LastToken = 0; + if (isS1G_) { + s1FirstToken = s1GFirstToken / static_cast(groupSize_); + s1LastToken = s1GLastToken / static_cast(groupSize_); + } else { + if (s1GFirstToken / batchCache.s1Size == s1GLastToken / batchCache.s1Size) { + // start and end locate in one G + s1FirstToken = s1GFirstToken % static_cast(batchCache.s1Size); + s1LastToken = s1GLastToken % static_cast(batchCache.s1Size); + } else { + // start and end locate in tow or more G, but working same as crossing a complete block + s1FirstToken = 0; + s1LastToken = batchCache.s1Size; + } + } + + int64_t s2FirstToken = 0; + int64_t s2LastToken = 0; + if (!isCmpKv) { + s2FirstToken = s1FirstToken - batchCache.oriPreTokenLeftUp; + s2LastToken = + batchCache.oriNextTokenLeftUp == INT64_MAX ? INT64_MAX : s1LastToken + batchCache.oriNextTokenLeftUp; + } else { + s2FirstToken = s1FirstToken - batchCache.cmpPreTokenLeftUp; + s2LastToken = + batchCache.cmpNextTokenLeftUp == INT64_MAX ? INT64_MAX : s1LastToken + batchCache.cmpNextTokenLeftUp; + } + return std::make_pair(s2FirstToken, s2LastToken); +} + +void SparseFlashMlaMetadataCpuKernel::CalcBatchCache(uint32_t bIdx, const SplitContext &splitContext, + BatchCache &batchCache) +{ + const SplitInfo &splitInfo = splitContext.splitInfo; + + batchCache.bIdx = bIdx; + batchCache.s1Size = GetS1SeqSize(bIdx); + if (hasOriKv_) { + batchCache.oriS2Size = GetOriS2SeqSize(bIdx); + batchCache.oriPreTokenLeftUp = CalcOriPreTokenLeftUp(batchCache.s1Size, batchCache.oriS2Size); + batchCache.oriNextTokenLeftUp = CalcOriNextTokenLeftUp(batchCache.s1Size, batchCache.oriS2Size); + } + if (hasCmpKv_) { + batchCache.cmpRevertS2Size = GetRevertS2Size(bIdx); + batchCache.cmpPreTokenLeftUp = CalcCmpPreTokenLeftUp(batchCache.s1Size, batchCache.cmpRevertS2Size); + batchCache.cmpNextTokenLeftUp = CalcCmpNextTokenLeftUp(batchCache.s1Size, batchCache.cmpRevertS2Size); + } +} + +void SparseFlashMlaMetadataCpuKernel::CalcOriS1GCache(S1GCache &s1GCache, const SplitInfo &splitInfo) +{ + // 处理 ori 部分 block 信息 + if (s1GCache.oriS2Start >= s1GCache.oriS2End) { + // ori 范围无效, 则整个 s1g 行等效为空行 + s1GCache.oriS1GBlock = 0; + s1GCache.oriS1GCost = 0; + s1GCache.oriS1GLastBlockCost = 0; + s1GCache.oriS1GNormalBlockCost = 0; + } else { + // 计算 ori 方向 Block 数量及 Cost + s1GCache.oriS1GBlock = s1GCache.oriS2End - s1GCache.oriS2Start; + // 判断 ori S2 方向是否包含尾块 + uint32_t curOriTailS2Num = (s1GCache.oriS2TailSize != 0U) ? 1U : 0U; + uint32_t curOriNormalS2Num = s1GCache.oriS1GBlock - curOriTailS2Num; + if (s1GCache.s1GIdx == (splitInfo.s1GBaseNum[s1GCache.bIdx] - 1U) && + splitInfo.s1GTailSize[s1GCache.bIdx] != 0U) { + s1GCache.oriS1GCost = typeCost_[ORI_TAIL_BLOCK][ORI_NORMAL_BLOCK] * curOriNormalS2Num + + typeCost_[ORI_TAIL_BLOCK][ORI_TAIL_BLOCK] * curOriTailS2Num; + s1GCache.oriS1GLastBlockCost = curOriTailS2Num > 0U ? typeCost_[ORI_TAIL_BLOCK][ORI_TAIL_BLOCK] : + typeCost_[ORI_TAIL_BLOCK][ORI_NORMAL_BLOCK]; + s1GCache.oriS1GNormalBlockCost = typeCost_[ORI_TAIL_BLOCK][ORI_NORMAL_BLOCK]; + } else { + s1GCache.oriS1GCost = typeCost_[ORI_NORMAL_BLOCK][ORI_NORMAL_BLOCK] * curOriNormalS2Num + + typeCost_[ORI_NORMAL_BLOCK][ORI_TAIL_BLOCK] * curOriTailS2Num; + s1GCache.oriS1GLastBlockCost = curOriTailS2Num > 0U ? typeCost_[ORI_NORMAL_BLOCK][ORI_TAIL_BLOCK] : + typeCost_[ORI_NORMAL_BLOCK][ORI_NORMAL_BLOCK]; + s1GCache.oriS1GNormalBlockCost = typeCost_[ORI_NORMAL_BLOCK][ORI_NORMAL_BLOCK]; + } + } +} + +void SparseFlashMlaMetadataCpuKernel::CalcCmpS1GCache(S1GCache &s1GCache, const SplitInfo &splitInfo) +{ + // 处理cmp部分block信息 + if (s1GCache.cmpS2Start >= s1GCache.cmpS2End) { + // Cmp范围无效, Cost保持为 0 + s1GCache.cmpS1GBlock = 0; + s1GCache.cmpS1GCost = 0; + s1GCache.cmpS1GLastBlockCost = 0; + s1GCache.cmpS1GNormalBlockCost = 0; + } else { + // 计算 cmp 方向 Block 数量及 Cost + s1GCache.cmpS1GBlock = s1GCache.cmpS2End - s1GCache.cmpS2Start; + // 判断 Cmp S2 方向是否包含尾块 + uint32_t curCmpTailS2Num = (s1GCache.cmpS2TailSize != 0U) ? 1U : 0U; // Updated check using local var + uint32_t curCmpNormalS2Num = s1GCache.cmpS1GBlock - curCmpTailS2Num; + if (s1GCache.s1GIdx == (splitInfo.s1GBaseNum[s1GCache.bIdx] - 1U) && + splitInfo.s1GTailSize[s1GCache.bIdx] != 0U) { + s1GCache.cmpS1GCost = typeCost_[CMP_TAIL_BLOCK][CMP_NORMAL_BLOCK] * curCmpNormalS2Num + + typeCost_[CMP_TAIL_BLOCK][CMP_TAIL_BLOCK] * curCmpTailS2Num; + s1GCache.cmpS1GLastBlockCost = curCmpTailS2Num > 0U ? typeCost_[CMP_TAIL_BLOCK][CMP_TAIL_BLOCK] : + typeCost_[CMP_TAIL_BLOCK][CMP_NORMAL_BLOCK]; + s1GCache.cmpS1GNormalBlockCost = typeCost_[CMP_TAIL_BLOCK][CMP_NORMAL_BLOCK]; + } else { + s1GCache.cmpS1GCost = typeCost_[CMP_NORMAL_BLOCK][CMP_NORMAL_BLOCK] * curCmpNormalS2Num + + typeCost_[CMP_NORMAL_BLOCK][CMP_TAIL_BLOCK] * curCmpTailS2Num; + s1GCache.cmpS1GLastBlockCost = curCmpTailS2Num > 0U ? typeCost_[CMP_NORMAL_BLOCK][CMP_TAIL_BLOCK] : + typeCost_[CMP_NORMAL_BLOCK][CMP_NORMAL_BLOCK]; + s1GCache.cmpS1GNormalBlockCost = typeCost_[CMP_NORMAL_BLOCK][CMP_NORMAL_BLOCK]; + } + } +} + +void SparseFlashMlaMetadataCpuKernel::CalcOriBlockRange(const Range &oriS2TokenRange, + const BatchCache &batchCache, S1GCache &s1GCache) +{ + int64_t oriS2FirstToken = oriS2TokenRange.first; + int64_t oriS2LastToken = oriS2TokenRange.second; + s1GCache.oriS2Start = 0; + // ori 部分 s2 起止和 tailSize + if (oriS2FirstToken >= static_cast(batchCache.oriS2Size) || oriS2LastToken < 0 || + oriS2LastToken < oriS2FirstToken) { + s1GCache.oriS2End = 0; + s1GCache.oriS2TailSize = 0; + } else { + oriS2FirstToken = + Clip(oriS2FirstToken, static_cast(0), static_cast(batchCache.oriS2Size - 1U)); + oriS2LastToken = Clip(oriS2LastToken, static_cast(0), static_cast(batchCache.oriS2Size - 1U)); + // oriS2LastToken 与 topk 取最小 + uint32_t s1Idx = GetS1Idx(batchCache.s1Size, s1GCache.s1GIdx); + uint32_t bsStride = GetBsStride(s1GCache.bIdx, s1Idx); + uint32_t oriTopkSize = isSparseOriKv_ ? + GetOriTopkLength(s1GCache.s1GIdx, batchCache) : + GetOriTopkLength(GetBsStride(s1GCache.bIdx, GetS1Idx(batchCache.s1Size, s1GCache.s1GIdx))); + s1GCache.actOriS2Size = isSparseOriKv_ ? + std::min(static_cast(oriS2LastToken - oriS2FirstToken + 1), oriTopkSize) : + static_cast(oriS2LastToken - oriS2FirstToken + 1); + s1GCache.oriS2End = s1GCache.actOriS2Size == 0 ? + 0 : + (s1GCache.actOriS2Size - 1U) / s2BaseSize_ + 1U; + s1GCache.oriS2TailSize = s1GCache.actOriS2Size % s2BaseSize_; + } +} + +void SparseFlashMlaMetadataCpuKernel::CalcCmpBlockRange(const Range &cmpRevertS2TokenRange, + const BatchCache &batchCache, S1GCache &s1GCache) +{ + int64_t cmpRevertS2FirstToken = cmpRevertS2TokenRange.first; + int64_t cmpRevertS2LastToken = cmpRevertS2TokenRange.second; + s1GCache.cmpS2Start = s1GCache.oriS2End; + // cmp 部分 s2 起止和 tailSize + if (cmpRevertS2FirstToken >= static_cast(batchCache.cmpRevertS2Size) || cmpRevertS2LastToken < 0 || + cmpRevertS2LastToken < cmpRevertS2FirstToken) { + s1GCache.cmpS2End = s1GCache.cmpS2Start; + s1GCache.cmpS2TailSize = 0; + } else { + cmpRevertS2FirstToken = + Clip(cmpRevertS2FirstToken, static_cast(0), static_cast(batchCache.cmpRevertS2Size - 1U)); + cmpRevertS2LastToken = + Clip(cmpRevertS2LastToken, static_cast(0), static_cast(batchCache.cmpRevertS2Size - 1U)); + // 如果压缩后长度为0,则直接返回 + if ((cmpRevertS2LastToken + 1) / cmpRatio_ == 0) { + s1GCache.cmpS2End = s1GCache.cmpS2Start; + s1GCache.cmpS2TailSize = 0; + return; + } + // 获取压缩后的 token 索引 + uint64_t cmpS2FirstToken = + (cmpRevertS2FirstToken + 1) / cmpRatio_ == 0 ? 0 : (cmpRevertS2FirstToken + 1) / cmpRatio_ - 1U; + uint64_t cmpS2LastToken = (cmpRevertS2LastToken + 1) / cmpRatio_ - 1U; + // cmpS2LastToken 与 topk 取最小 + uint32_t s1Idx = GetS1Idx(batchCache.s1Size, s1GCache.s1GIdx); + uint32_t bsStride = GetBsStride(s1GCache.bIdx, s1Idx); + uint32_t cmpTopkSize = GetCmpTopkLength(bsStride); + s1GCache.actCmpS2Size = isSparseCmpKv_ ? + std::min(static_cast(cmpS2LastToken - cmpS2FirstToken + 1), cmpTopkSize) : + static_cast(cmpS2LastToken - cmpS2FirstToken + 1); + s1GCache.cmpS2End = s1GCache.actCmpS2Size == 0 ? + s1GCache.cmpS2Start : + s1GCache.cmpS2Start + (s1GCache.actCmpS2Size - 1U) / s2BaseSize_ + 1U; + s1GCache.cmpS2TailSize = s1GCache.actCmpS2Size % s2BaseSize_; + } +} + +void SparseFlashMlaMetadataCpuKernel::GatherOriAndCmpCache(S1GCache &s1GCache) +{ + s1GCache.s2Start = 0; + if (s1GCache.cmpS1GBlock > 0) { + s1GCache.s1GLastBlockCost = s1GCache.cmpS1GLastBlockCost; + s1GCache.s2End = s1GCache.cmpS2End; + } else { + s1GCache.s1GLastBlockCost = s1GCache.oriS1GLastBlockCost; + s1GCache.s2End = s1GCache.oriS2End; + } + s1GCache.s1GBlock = s1GCache.oriS1GBlock + s1GCache.cmpS1GBlock; + s1GCache.s1GCost = s1GCache.oriS1GCost + s1GCache.cmpS1GCost; +} + +void SparseFlashMlaMetadataCpuKernel::CalcS1GCache(uint32_t s1GIdx, const SplitContext &splitContext, + const BatchCache &batchCache, S1GCache &s1GCache) +{ + const SplitInfo &splitInfo = splitContext.splitInfo; + // 如果s1G是空行,则直接返回 + if (splitInfo.s1GBaseNum[batchCache.bIdx] == 0) { + s1GCache.s1GCost = 0; + s1GCache.s1GLastBlockCost = 0; + s1GCache.oriS1GNormalBlockCost = 0; + s1GCache.oriS1GLastBlockCost = 0; + s1GCache.cmpS1GNormalBlockCost = 0; + s1GCache.cmpS1GLastBlockCost = 0; + s1GCache.s1GBlock = 0; + s1GCache.s2Loop = 0; + s1GCache.s2Start = 0; + s1GCache.cmpS2Start = 0; + s1GCache.s2End = 0; + return; + } + s1GCache.bIdx = batchCache.bIdx; + s1GCache.s1GIdx = s1GIdx; + s1GCache.actOriS2Size = 0U; + s1GCache.actCmpS2Size = 0U; + s1GCache.reductionBlockSize = 0U; + // 计算 ori_kv 有效负载起止 + if (hasOriKv_) { + // 计算 ori_kv 的 s2Token 起止 + auto oriS2TokenRange = CalcS2TokenRange(s1GIdx, batchCache, ORI_KV); + // 计算 ori_kv 的 s2Block 起止 + CalcOriBlockRange(oriS2TokenRange, batchCache, s1GCache); + } else { + // ori_kv s2Token 起止初始化为0 + s1GCache.oriS2Start = 0; + s1GCache.oriS2End = s1GCache.oriS2Start; + s1GCache.oriS2TailSize = 0; + } + // 计算 cmp_kv 有效负载起止 + if (hasCmpKv_) { + // 计算 cmp_kv 的 s2Token 起止 + auto cmpRevertS2TokenRange = CalcS2TokenRange(s1GIdx, batchCache, CMP_KV); + // 计算 cmp_kv 的 s2Block 起止 + CalcCmpBlockRange(cmpRevertS2TokenRange, batchCache, s1GCache); + } else { + // cmp_kv s2Token 起止初始化为0 + s1GCache.cmpS2Start = s1GCache.oriS2End; + s1GCache.cmpS2End = s1GCache.cmpS2Start; + s1GCache.cmpS2TailSize = 0; + } + if (isBatchConsistency_) { + // The reduction block only depends on this row, so its reduction tree is independent of the surrounding batch. + uint64_t actTotalS2Size = static_cast(s1GCache.actOriS2Size) + s1GCache.actCmpS2Size; + if (actTotalS2Size > 0U) { + uint64_t rawReductionBlockSize = actTotalS2Size / BATCH_CONSISTENCY_MAX_REDUCTION_PARTS; + s1GCache.reductionBlockSize = static_cast( + (rawReductionBlockSize + s2BaseSize_ - 1U) / s2BaseSize_ * s2BaseSize_); + s1GCache.reductionBlockSize = std::max(s1GCache.reductionBlockSize, s2BaseSize_); + s1GCache.oriS2End = s1GCache.actOriS2Size == 0U ? + 0U : + (s1GCache.actOriS2Size - 1U) / s1GCache.reductionBlockSize + 1U; + s1GCache.oriS2TailSize = s1GCache.actOriS2Size % s1GCache.reductionBlockSize; + s1GCache.cmpS2Start = s1GCache.oriS2End; + s1GCache.cmpS2End = s1GCache.actCmpS2Size == 0U ? + s1GCache.cmpS2Start : + s1GCache.cmpS2Start + (s1GCache.actCmpS2Size - 1U) / s1GCache.reductionBlockSize + 1U; + s1GCache.cmpS2TailSize = s1GCache.actCmpS2Size % s1GCache.reductionBlockSize; + } + } + // 计算基本块负载 + CalcCostTable(splitInfo.s1GTailSize[s1GCache.bIdx], s1GCache.reductionBlockSize, s1GCache.oriS2TailSize, + s1GCache.cmpS2TailSize); + // 计算 ori 和 cmp 部分的 cost 和 block 信息 + CalcOriS1GCache(s1GCache, splitInfo); + CalcCmpS1GCache(s1GCache, splitInfo); + // 汇总 ori 和 cmp 部分的 cost 和 block 信息 + GatherOriAndCmpCache(s1GCache); + s1GCache.s2Loop = static_cast( + (static_cast(s1GCache.actOriS2Size) + s2BaseSize_ - 1U) / s2BaseSize_ + + (static_cast(s1GCache.actCmpS2Size) + s2BaseSize_ - 1U) / s2BaseSize_); +} + +void SparseFlashMlaMetadataCpuKernel::CalcBatchCost(uint32_t bIdx, const SplitContext &splitContext, CostInfo &costInfo) +{ + const SplitInfo &splitInfo = splitContext.splitInfo; + + costInfo.bN2CostOfEachBatch[bIdx] = 0; + costInfo.bN2BlockOfEachBatch[bIdx] = 0U; + costInfo.bN2S2LoopOfEachBatch[bIdx] = 0U; + costInfo.bN2LastBlockCostOfEachBatch[bIdx] = 0U; + + if (GetS1SeqSize(bIdx) == 0U) { + return; + } + if (!hasOriKv_ && !hasCmpKv_) { + return; + } else if (!hasOriKv_) { + if (GetCmpS2SeqSize(bIdx) == 0U) { + return; + } + } else if (!hasCmpKv_) { + if (GetOriS2SeqSize(bIdx) == 0U) { + return; + } + } else { + if ((GetOriS2SeqSize(bIdx) == 0U) && GetCmpS2SeqSize(bIdx) == 0U) { + return; + } + } + + BatchCache bCache; + S1GCache s1GCache; + CalcBatchCache(bIdx, splitContext, bCache); + for (uint32_t s1GIdx = 0; s1GIdx < splitInfo.s1GBaseNum[bIdx]; s1GIdx++) { + CalcS1GCache(s1GIdx, splitContext, bCache, s1GCache); + costInfo.bN2CostOfEachBatch[bIdx] += s1GCache.s1GCost; + costInfo.bN2BlockOfEachBatch[bIdx] += s1GCache.s1GBlock; + costInfo.bN2S2LoopOfEachBatch[bIdx] += s1GCache.s2Loop; + // 更新最大S1G行开销 + if (s1GCache.s1GCost > costInfo.maxS1GCost) { + costInfo.maxS1GCost = s1GCache.s1GCost; + } + if (s1GCache.s1GBlock > 0) { + costInfo.bN2LastBlockCostOfEachBatch[bIdx] = s1GCache.s1GLastBlockCost; + } + } +} + +void SparseFlashMlaMetadataCpuKernel::CalcCostInfo(SplitContext &splitContext) +{ + const SplitInfo &splitInfo = splitContext.splitInfo; + CostInfo &costInfo = splitContext.costInfo; + + if (splitInfo.isKvSeqAllZero) { + costInfo.totalCost = 0; + costInfo.totalBlockNum = 0U; + return; + } + + // 计算 batch 的负载并记录,用于按batch分配,需要按行计算起止点,统计块数、负载 + for (uint32_t bIdx = 0; bIdx < batchSize_; bIdx++) { + CalcBatchCost(bIdx, splitContext, costInfo); + costInfo.totalCost += costInfo.bN2CostOfEachBatch[bIdx] * numHeadsKv_; + costInfo.totalBlockNum += costInfo.bN2BlockOfEachBatch[bIdx] * numHeadsKv_; + } +} + +void SparseFlashMlaMetadataCpuKernel::UpdateCursor(const SplitContext &splitContext, AssignContext &assignContext) +{ + const SplitInfo &splitInfo = splitContext.splitInfo; + const CostInfo &costInfo = splitContext.costInfo; + + bool UpdateS1G = false; + bool UpdateBatch = false; + + // Update S2 + if (assignContext.curS2Idx >= assignContext.s1GCache.s2End) { // 边界assignInfo.s2End是取不到的开区间 + assignContext.curS2Idx = 0U; + assignContext.curS1GIdx++; + UpdateS1G = true; + } + + // Update S1G + if (assignContext.curS1GIdx >= splitInfo.s1GBaseNum[assignContext.curBIdx]) { + assignContext.curS1GIdx = 0U; + assignContext.curBN2Idx++; + } + + // Update Batch + if (assignContext.curBN2Idx == batchSize_ * numHeadsKv_) { // 所有负载全部分配完,设置最后一个核的右开区间,返回 + assignContext.curS1GIdx = 0U; + assignContext.curS2Idx = 0U; + assignContext.isFinished = true; + return; + } + + if (assignContext.curBN2Idx / numHeadsKv_ != assignContext.curBIdx) { + assignContext.curBIdx = assignContext.curBN2Idx / numHeadsKv_; + assignContext.curS1GIdx = 0U; + UpdateBatch = true; + UpdateS1G = true; + } + + // Update Cache + if (UpdateBatch) { + CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache); + assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx]; + assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx]; + assignContext.bN2S2Loop = costInfo.bN2S2LoopOfEachBatch[assignContext.curBIdx]; + } + if (UpdateS1G) { + CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache); + assignContext.curS2Idx = supportFd_ ? assignContext.s1GCache.oriS2Start : 0U; + } +} + +void SparseFlashMlaMetadataCpuKernel::AssignByBatch(const SplitContext &splitContext, AssignContext &assignContext) +{ + if (assignContext.isFinished) { + return; + } + const CostInfo &costInfo = splitContext.costInfo; + const SplitInfo &splitInfo = splitContext.splitInfo; + while (assignContext.bN2Cost == 0 || + IsWithinTolerance(assignContext.coreCache.costLimit, + costInfo.bN2LastBlockCostOfEachBatch[assignContext.curBIdx] / FA_TOLERANCE_RATIO, + assignContext.coreCache.cost + assignContext.bN2Cost)) { + assignContext.coreCache.cost += assignContext.bN2Cost; + assignContext.coreCache.block += assignContext.bN2Block; + assignContext.coreCache.s2Loop += assignContext.bN2S2Loop; + assignContext.curBN2Idx++; + // to the end + if (assignContext.curBN2Idx == batchSize_ * numHeadsKv_) { + assignContext.curS1GIdx = 0U; + assignContext.curS2Idx = 0U; + assignContext.isFinished = true; + return; + } + + // next batch + if (assignContext.curBN2Idx / numHeadsKv_ != assignContext.curBIdx) { + assignContext.curBIdx = assignContext.curBN2Idx / numHeadsKv_; + CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache); + } + + assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx]; + assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx]; + assignContext.bN2S2Loop = costInfo.bN2S2LoopOfEachBatch[assignContext.curBIdx]; + assignContext.curS1GIdx = 0U; + CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache); + assignContext.curS2Idx = assignContext.s1GCache.s2Start; + } +} + +void SparseFlashMlaMetadataCpuKernel::AssignByRow(const SplitContext &splitContext, AssignContext &assignContext) +{ + if (assignContext.isFinished) { + return; + } + + while (IsWithinTolerance(assignContext.coreCache.costLimit, + assignContext.s1GCache.s1GLastBlockCost / FA_TOLERANCE_RATIO, + assignContext.coreCache.cost + assignContext.s1GCache.s1GCost)) { + assignContext.coreCache.cost += assignContext.s1GCache.s1GCost; + assignContext.coreCache.block += assignContext.s1GCache.s1GBlock; + assignContext.coreCache.s2Loop += assignContext.s1GCache.s2Loop; + // 当前batch被分配一行出去,更新剩余负载 + assignContext.bN2Cost = assignContext.bN2Cost > assignContext.s1GCache.s1GCost ? + assignContext.bN2Cost - assignContext.s1GCache.s1GCost : + 0; + assignContext.bN2Block = assignContext.bN2Block > assignContext.s1GCache.s1GBlock ? + assignContext.bN2Block - assignContext.s1GCache.s1GBlock : + 0U; + assignContext.bN2S2Loop = assignContext.bN2S2Loop > assignContext.s1GCache.s2Loop ? + assignContext.bN2S2Loop - assignContext.s1GCache.s2Loop : + 0U; + // 计算新一行的信息 + do { + assignContext.curS1GIdx++; + CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache); + } while (assignContext.s1GCache.s1GBlock == 0); + assignContext.curS2Idx = assignContext.s1GCache.s2Start; + } +} + +int64_t SparseFlashMlaMetadataCpuKernel::CalcCurBlockCost(const AssignContext &assignContext) +{ + int64_t curCost = 0; + if (assignContext.curS2Idx < assignContext.s1GCache.cmpS2Start) { + curCost = assignContext.s1GCache.oriS1GNormalBlockCost; + if (assignContext.curS2Idx == (assignContext.s1GCache.cmpS2Start - 1U)) { + curCost = assignContext.s1GCache.oriS1GLastBlockCost; + } + } else { + curCost = assignContext.s1GCache.cmpS1GNormalBlockCost; + if (assignContext.curS2Idx == (assignContext.s1GCache.s2End - 1U)) { + curCost = assignContext.s1GCache.cmpS1GLastBlockCost; + } + } + return curCost; +} + +uint32_t SparseFlashMlaMetadataCpuKernel::CalcCurBlockS2Loop(const AssignContext &assignContext) +{ + if (!isBatchConsistency_) { + return 1U; + } + const S1GCache &s1GCache = assignContext.s1GCache; + uint32_t blockSize = s1GCache.reductionBlockSize; + if (assignContext.curS2Idx < s1GCache.cmpS2Start) { + if (assignContext.curS2Idx + 1U == s1GCache.cmpS2Start && s1GCache.oriS2TailSize != 0) { + blockSize = static_cast(s1GCache.oriS2TailSize); + } + } else if (assignContext.curS2Idx + 1U == s1GCache.s2End && s1GCache.cmpS2TailSize != 0) { + blockSize = static_cast(s1GCache.cmpS2TailSize); + } + return (blockSize + s2BaseSize_ - 1U) / s2BaseSize_; +} + +void SparseFlashMlaMetadataCpuKernel::AssignByBlock(const SplitContext &splitContext, AssignContext &assignContext) +{ + if (assignContext.isFinished || !supportFd_) { + return; + } + + int64_t curCost = CalcCurBlockCost(assignContext); + uint32_t curS2Loop = CalcCurBlockS2Loop(assignContext); + + // (costLimit - curCostOnCore) * FA_TOLERANCE_RATIO > curCost;至少分配1块 + while (IsWithinTolerance(assignContext.coreCache.costLimit, curCost / FA_TOLERANCE_RATIO, + assignContext.coreCache.cost + curCost)) { + assignContext.coreCache.cost += curCost; + assignContext.coreCache.block++; + assignContext.coreCache.s2Loop += curS2Loop; + assignContext.curS2Idx++; + // 当前batch被分配一块出去,更新剩余负载 + assignContext.bN2Cost = assignContext.bN2Cost - curCost; + // 当前行被分配一块出去,更新剩余负载 + assignContext.s1GCache.s1GCost = assignContext.s1GCache.s1GCost - curCost; + assignContext.bN2Block--; + assignContext.s1GCache.s1GBlock--; + assignContext.bN2S2Loop = + assignContext.bN2S2Loop > curS2Loop ? assignContext.bN2S2Loop - curS2Loop : 0U; + assignContext.s1GCache.s2Loop = + assignContext.s1GCache.s2Loop > curS2Loop ? assignContext.s1GCache.s2Loop - curS2Loop : 0U; + curCost = CalcCurBlockCost(assignContext); + curS2Loop = CalcCurBlockS2Loop(assignContext); + } +} + +void SparseFlashMlaMetadataCpuKernel::ForceAssign(const SplitContext &splitContext, AssignContext &assignContext) +{ + if (assignContext.isFinished) { + return; + } + + int64_t curCost = CalcCurBlockCost(assignContext); + uint32_t curS2Loop = CalcCurBlockS2Loop(assignContext); + + assignContext.coreCache.cost += curCost; + assignContext.coreCache.block++; + assignContext.coreCache.s2Loop += curS2Loop; + assignContext.curS2Idx++; + // 当前batch被分配一块出去,更新剩余负载 + assignContext.bN2Cost = assignContext.bN2Cost - curCost; + assignContext.bN2Block--; + assignContext.bN2S2Loop = + assignContext.bN2S2Loop > curS2Loop ? assignContext.bN2S2Loop - curS2Loop : 0U; + // 当前行被分配一块出去,更新剩余负载 + assignContext.s1GCache.s1GCost = assignContext.s1GCache.s1GCost - curCost; + assignContext.s1GCache.s1GBlock--; + assignContext.s1GCache.s2Loop = + assignContext.s1GCache.s2Loop > curS2Loop ? assignContext.s1GCache.s2Loop - curS2Loop : 0U; + UpdateCursor(splitContext, assignContext); +} + +bool SparseFlashMlaMetadataCpuKernel::IsNeedRecordFDInfo(const AssignContext &assignContext, + const SplitResult &splitRes) +{ + // 切分点大概率不会刚好在行尾,因此滞后处理归约信息的统计,到下一个切分点再判断是否需要归约 + // 核0无需处理 + if (assignContext.curCoreIdx == 0U) { + return false; + } + // 无跨核行,无需处理 + if (assignContext.curKvSplitPart <= 1U) { + return false; + } + // 需要归约的行还未处理完 + if (assignContext.curBN2Idx == splitRes.bN2End[assignContext.curCoreIdx - 1U] && + assignContext.curS1GIdx == splitRes.gS1End[assignContext.curCoreIdx - 1U]) { + return false; + } + return true; +} + +bool SparseFlashMlaMetadataCpuKernel::IsFirstReductionBlock(const AssignContext &assignContext, + const SplitResult &splitRes) +{ + if (assignContext.curCoreIdx == 0U || splitRes.s2End[assignContext.curCoreIdx - 1U] == 0U) { + return true; + } + return assignContext.curBN2Idx != splitRes.bN2End[assignContext.curCoreIdx - 1U] || + assignContext.curS1GIdx != splitRes.gS1End[assignContext.curCoreIdx - 1U]; +} + +void SparseFlashMlaMetadataCpuKernel::RecordFDInfo(const SplitContext &splitContext, const AssignContext &assignContext, + SplitResult &result) +{ + const SplitInfo &splitInfo = splitContext.splitInfo; + // 需要规约的行是上一个核的切分点所在位置 + uint32_t splitBIdx = result.bN2End[assignContext.curCoreIdx - 1U] / numHeadsKv_; + uint32_t splitS1GIdx = result.gS1End[assignContext.curCoreIdx - 1U]; + uint32_t s1Size = GetS1SeqSize(splitBIdx); + + // 计算归约数据的FD均衡划分信息 + uint32_t curFdS1gSize = + (splitS1GIdx == splitInfo.s1GBaseNum[splitBIdx] - 1U) ? + (static_cast(s1Size) * groupSize_ - static_cast(splitS1GIdx) * mBaseSize_) : + mBaseSize_; + // 记录 + result.maxS2SplitNum = std::max(result.maxS2SplitNum, assignContext.curKvSplitPart); + // 若存在头归约,则切分点一定为上一个核结束的位置 + result.fdRes.fdBN2Idx[result.numOfFdHead] = result.bN2End[assignContext.curCoreIdx - 1U]; + result.fdRes.fdMIdx[result.numOfFdHead] = result.gS1End[assignContext.curCoreIdx - 1U]; + result.fdRes.fdWorkspaceIdx[result.numOfFdHead] = assignContext.preFdDataNum; + result.fdRes.fdS2SplitNum[result.numOfFdHead] = assignContext.curKvSplitPart; + result.fdRes.fdMSize[result.numOfFdHead] = curFdS1gSize; + result.numOfFdHead++; +} + +void SparseFlashMlaMetadataCpuKernel::AssignBlocksToCore(const SplitContext &splitContext, AssignContext &assignContext, + SplitResult &result) +{ + const CostInfo &costInfo = splitContext.costInfo; + result.firstFdDataWorkspaceIdx[assignContext.curCoreIdx] = + assignContext.preFdDataNum + assignContext.curKvSplitPart - 1U; + int64_t avgCost = assignContext.unassignedCost / (aicCoreNum_ - assignContext.curCoreIdx); + assignContext.coreCache = {}; + if (!supportFd_) { + assignContext.coreCache.costLimit = std::max(avgCost, costInfo.maxS1GCost); + } else { + assignContext.coreCache.costLimit = avgCost; + } + // 1、按整batch分配 + AssignByBatch(splitContext, assignContext); + // 2、按行分配 + AssignByRow(splitContext, assignContext); + // 3、按块分配 + AssignByBlock(splitContext, assignContext); + // 4、强制分配 + if (assignContext.coreCache.block == 0 && supportFd_) { + ForceAssign(splitContext, assignContext); + } + result.bN2End[assignContext.curCoreIdx] = assignContext.curBN2Idx; + result.gS1End[assignContext.curCoreIdx] = assignContext.curS1GIdx; + result.s2End[assignContext.curCoreIdx] = assignContext.curS2Idx; + result.maxCost = std::max(result.maxCost, assignContext.coreCache.cost); + assignContext.unassignedCost -= assignContext.coreCache.cost; + result.maxS2LoopNum = std::max(assignContext.coreCache.s2Loop, result.maxS2LoopNum); + // 对之前的归约信息进行记录并清理 + if (IsNeedRecordFDInfo(assignContext, result)) { + if (isBatchConsistency_ && remainedBlockNum_ > 0U) { + // curKvSplitPart already reserves one slot for the core that finishes this row. + assignContext.curKvSplitPart += remainedBlockNum_ - 1U; + } + RecordFDInfo(splitContext, assignContext, result); + assignContext.preFdDataNum += assignContext.curKvSplitPart; + assignContext.curKvSplitPart = 1U; + remainedBlockNum_ = 0U; + } + // 更新S2切分信息 + if (assignContext.curS2Idx > assignContext.s1GCache.s2Start && + assignContext.curS2Idx <= assignContext.s1GCache.s2End) { + if (isBatchConsistency_) { + if (IsFirstReductionBlock(assignContext, result)) { + assignContext.curKvSplitPart++; + } else { + assignContext.curKvSplitPart += + result.s2End[assignContext.curCoreIdx] - result.s2End[assignContext.curCoreIdx - 1U]; + } + remainedBlockNum_ = assignContext.s1GCache.s1GBlock; + } else { + assignContext.curKvSplitPart++; + } + } +} + +void SparseFlashMlaMetadataCpuKernel::CalcSplitPlan(int64_t costLimit, const SplitContext &splitContext, + SplitResult &result) +{ + const CostInfo &costInfo = splitContext.costInfo; + const SplitInfo &splitInfo = splitContext.splitInfo; + if (aicCoreNum_ == 0U) { + return; + } + result.maxCost = 0U; + result.usedCoreNum = 0U; + + AssignContext assignContext{}; + assignContext.curBIdx = 0U; + assignContext.curS1GIdx = 0U; + assignContext.unassignedCost = costInfo.totalCost; + assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx]; + assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx]; + assignContext.bN2S2Loop = costInfo.bN2S2LoopOfEachBatch[assignContext.curBIdx]; + CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache); + CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache); + assignContext.curS2Idx = assignContext.s1GCache.s2Start; + // 负载分配 + for (uint32_t i = 0; i < aicCoreNum_; ++i) { + if (result.maxCost > costLimit) { + return; + } + if (assignContext.isFinished || assignContext.unassignedCost <= 0) { + break; + } + assignContext.curCoreIdx = i; + AssignBlocksToCore(splitContext, assignContext, result); + } + result.usedCoreNum = assignContext.curCoreIdx + 1; +} + +void SparseFlashMlaMetadataCpuKernel::SplitFD(SplitResult &splitRes) +{ + // 计算FD的总数据量 + uint64_t totalFDLoad = 0; + for (uint32_t i = 0; i < splitRes.numOfFdHead; i++) { + totalFDLoad += splitRes.fdRes.fdS2SplitNum[i] * splitRes.fdRes.fdMSize[i]; + } + // 计算当前最大冗余vec核数 + uint32_t emptyVectorNum = aivCoreNum_ - splitRes.numOfFdHead; + // 计算每个核处理的load + uint64_t averageLoad = (totalFDLoad + aivCoreNum_ - 1U) / aivCoreNum_; // 向上取整,避免核负载为0 + uint32_t curCoreIndex = 0; + for (uint32_t i = 0; i < splitRes.numOfFdHead; i++) { + // 冗余vec核数为0,此时规约任务无法进行更小的切分,只能1个vec核计算1个规约任务 + if (emptyVectorNum == 0U) { + splitRes.fdRes.fdIdx[curCoreIndex] = i; + splitRes.fdRes.fdMStart[curCoreIndex] = 0U; + splitRes.fdRes.fdMNum[curCoreIndex] = splitRes.fdRes.fdMSize[i]; + curCoreIndex++; + continue; + } + // 计算当前归约任务所用核数,向下取整,避免使用核数超出总核数 + uint32_t curFDVectorNum = splitRes.fdRes.fdS2SplitNum[i] * splitRes.fdRes.fdMSize[i] / averageLoad; + curFDVectorNum = std::max(1U, curFDVectorNum); + // 计算当前归约任务每个核的行数,向上取整,避免行数为0 + uint32_t curAveMSize = (splitRes.fdRes.fdMSize[i] + curFDVectorNum - 1U) / curFDVectorNum; + curFDVectorNum = (splitRes.fdRes.fdMSize[i] + curAveMSize - 1U) / curAveMSize; + // 需要使用的vec核数与当前剩余可用vec核数取最小 + curFDVectorNum = std::min(curFDVectorNum, emptyVectorNum + 1U); // 1: Fd任务自身带一个核 + // FD负载分配 + for (uint32_t vid = 0; vid < curFDVectorNum; vid++) { + splitRes.fdRes.fdIdx[curCoreIndex] = i; + splitRes.fdRes.fdMStart[curCoreIndex] = vid * curAveMSize; + splitRes.fdRes.fdMNum[curCoreIndex] = + (vid < curFDVectorNum - 1) ? curAveMSize : (splitRes.fdRes.fdMSize[i] - vid * curAveMSize); + curCoreIndex++; + } + // 更新冗余vec核数 + emptyVectorNum -= (curFDVectorNum - 1U); // 1: 空余核不包含FD自身的核,要-1 + } + splitRes.fdRes.fdUsedVecNum = curCoreIndex; +} + +bool SparseFlashMlaMetadataCpuKernel::BalanceSchedule(SplitResult &splitRes) +{ + SplitContext splitContext(batchSize_); + + // 1、划分基本块,统计信息 + CalcSplitInfo(splitContext); + // 全空case + if (splitContext.splitInfo.isKvSeqAllZero) { + splitRes.usedCoreNum = 1U; + splitRes.bN2End[0] = batchSize_ * numHeadsKv_; + splitRes.gS1End[0] = 0U; + splitRes.s2End[0] = 0U; + return true; + } + CalcCostInfo(splitContext); + + // 2、获取每个核的分配方案 + splitRes.maxCost = INT64_MAX; + splitRes.usedCoreNum = 1U; + + CalcSplitPlan(splitRes.maxCost, splitContext, splitRes); + // 3、存在FD任务,对FD进行负载均衡分配 + if (splitRes.numOfFdHead > 0U) { + SplitFD(splitRes); + } + splitRes.usedCoreNum = std::max(splitRes.usedCoreNum, 1U); // 至少使用1个core + return true; +} + +bool SparseFlashMlaMetadataCpuKernel::GenMetadata(SplitResult &splitRes) +{ + optiling::detail::SmlaMetadata *metadataPtr = static_cast(metadata_->GetData()); + *metadataPtr = {}; + // FA Metadata Generate + if (isSplitG_) { + for (size_t i = 0; i < aicCoreNum_; i++) { + // 单核s2计算轮次最大数量 + metadataPtr->faMetadata[2 * i][FA_S2_MAX_NUM] = splitRes.maxS2LoopNum; + metadataPtr->faMetadata[2 * i + 1][FA_S2_MAX_NUM] = splitRes.maxS2LoopNum; + + if (i >= splitRes.usedCoreNum) { + metadataPtr->faMetadata[2 * i][FA_CORE_ENABLE_INDEX] = 0; // AIC disenable + metadataPtr->faMetadata[2 * i + 1][FA_CORE_ENABLE_INDEX] = 0; // AIC disenable + continue; + } + metadataPtr->faMetadata[2 * i][FA_CORE_ENABLE_INDEX] = 1; // AIC enable + metadataPtr->faMetadata[2 * i + 1][FA_CORE_ENABLE_INDEX] = 1; // AIC enable + // FA START + metadataPtr->faMetadata[2 * i][FA_BN2_START_INDEX] = i == 0 ? 0 : splitRes.bN2End[i - 1]; + metadataPtr->faMetadata[2 * i][FA_M_START_INDEX] = i == 0 ? 0 : splitRes.gS1End[i - 1]; + metadataPtr->faMetadata[2 * i][FA_S2_START_INDEX] = i == 0 ? 0 : splitRes.s2End[i - 1]; + + metadataPtr->faMetadata[2 * i + 1][FA_BN2_START_INDEX] = i == 0 ? 0 : splitRes.bN2End[i - 1]; + metadataPtr->faMetadata[2 * i + 1][FA_M_START_INDEX] = i == 0 ? 0 : splitRes.gS1End[i - 1]; + metadataPtr->faMetadata[2 * i + 1][FA_S2_START_INDEX] = i == 0 ? 0 : splitRes.s2End[i - 1]; + // FA END + metadataPtr->faMetadata[2 * i][FA_BN2_END_INDEX] = splitRes.bN2End[i]; + metadataPtr->faMetadata[2 * i][FA_M_END_INDEX] = splitRes.gS1End[i]; + metadataPtr->faMetadata[2 * i][FA_S2_END_INDEX] = splitRes.s2End[i]; + + metadataPtr->faMetadata[2 * i + 1][FA_BN2_END_INDEX] = splitRes.bN2End[i]; + metadataPtr->faMetadata[2 * i + 1][FA_M_END_INDEX] = splitRes.gS1End[i]; + metadataPtr->faMetadata[2 * i + 1][FA_S2_END_INDEX] = splitRes.s2End[i]; + // firstFdDataWorkspace + metadataPtr->faMetadata[2 * i][FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX] = splitRes.firstFdDataWorkspaceIdx[i]; + metadataPtr->faMetadata[2 * i + 1][FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX] = + splitRes.firstFdDataWorkspaceIdx[i]; + } + } else { + for (size_t i = 0; i < aicCoreNum_; ++i) { + if (i >= splitRes.usedCoreNum) { + metadataPtr->faMetadata[i][FA_CORE_ENABLE_INDEX] = 0; // AIC disenable + continue; + } + metadataPtr->faMetadata[i][FA_CORE_ENABLE_INDEX] = 1; // AIC enable + // FA START + metadataPtr->faMetadata[i][FA_BN2_START_INDEX] = i == 0 ? 0 : splitRes.bN2End[i - 1]; + metadataPtr->faMetadata[i][FA_M_START_INDEX] = i == 0 ? 0 : splitRes.gS1End[i - 1]; + metadataPtr->faMetadata[i][FA_S2_START_INDEX] = i == 0 ? 0 : splitRes.s2End[i - 1]; + // FA END + metadataPtr->faMetadata[i][FA_BN2_END_INDEX] = splitRes.bN2End[i]; + metadataPtr->faMetadata[i][FA_M_END_INDEX] = splitRes.gS1End[i]; + metadataPtr->faMetadata[i][FA_S2_END_INDEX] = splitRes.s2End[i]; + // firstFdDataWorkspace + metadataPtr->faMetadata[i][FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX] = splitRes.firstFdDataWorkspaceIdx[i]; + } + } + + // FD Metadata Generate + for (size_t i = 0; i < aivCoreNum_; ++i) { + if (i >= splitRes.fdRes.fdUsedVecNum) { + metadataPtr->fdMetadata[i][FD_CORE_ENABLE_INDEX] = 0; // AIV disenable + continue; + } + metadataPtr->fdMetadata[i][FD_CORE_ENABLE_INDEX] = 1; // AIV enable + uint32_t curFdIdx = splitRes.fdRes.fdIdx[i]; + metadataPtr->fdMetadata[i][FD_BN2_IDX_INDEX] = splitRes.fdRes.fdBN2Idx[curFdIdx]; + metadataPtr->fdMetadata[i][FD_M_IDX_INDEX] = splitRes.fdRes.fdMIdx[curFdIdx]; + metadataPtr->fdMetadata[i][FD_WORKSPACE_IDX_INDEX] = splitRes.fdRes.fdWorkspaceIdx[curFdIdx]; + metadataPtr->fdMetadata[i][FD_WORKSPACE_NUM_INDEX] = splitRes.fdRes.fdS2SplitNum[curFdIdx]; + metadataPtr->fdMetadata[i][FD_M_START_INDEX] = splitRes.fdRes.fdMStart[i]; + metadataPtr->fdMetadata[i][FD_M_NUM_INDEX] = splitRes.fdRes.fdMNum[i]; + } + return true; +} + +namespace { +static const char *kernelType = "SparseFlashMlaMetadata"; +REGISTER_CPU_KERNEL(kernelType, SparseFlashMlaMetadataCpuKernel); +} // namespace + +}; // namespace aicpu diff --git a/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.h b/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.h new file mode 100644 index 000000000000..e95358149a9e --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.h @@ -0,0 +1,398 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file sparse_flash_mla_metadata_aicpu.h + * \brief + */ + +#ifndef SPARSE_FLASH_MLA_METADATA_AICPU_H +#define SPARSE_FLASH_MLA_METADATA_AICPU_H + +#include +#include +#include +#include "cpu_context.h" +#include "cpu_kernel.h" +#include "cpu_tensor.h" +#include "../../sparse_flash_mla/op_kernel/sparse_flash_mla_kernel_metadata.h" +#include "../../common/op_kernel/aicpu_common.h" + +namespace aicpu { +constexpr int64_t FA_TOLERANCE_RATIO = 2; +constexpr uint32_t COST_WEIGHT_M = 6U; +constexpr uint32_t COST_WEIGHT_S2 = 10U; +constexpr uint32_t BATCH_CONSISTENCY_MAX_REDUCTION_PARTS = 32U; +constexpr bool ORI_KV = false; +constexpr bool CMP_KV = true; +constexpr uint32_t NO_MASK = 0; +constexpr uint32_t HAS_MASK = 1; + +enum BlockType : uint32_t { + ORI_NORMAL_BLOCK = 0, + ORI_TAIL_BLOCK, + CMP_NORMAL_BLOCK, + CMP_TAIL_BLOCK, + BLOCK_MAX_TYPE +}; + +enum class SparseMode : uint8_t { + DEFAULT_MASK = 0, + ALL_MASK, + LEFT_UP_CAUSAL, + RIGHT_DOWN_CAUSAL, + BAND, + SPARSE_BUTT, +}; + +enum class ValidSocVersion { + ASCEND910 = 0, + ASCEND950 +}; + +template +using Range = std::pair; + +template +using BlockCost = std::array(BLOCK_MAX_TYPE)>, static_cast(BLOCK_MAX_TYPE)>; + +template +T Clip(T value, T minValue, T maxValue) +{ + if (value < minValue) { + return minValue; + } + if (value > maxValue) { + return maxValue; + } + return value; +} + +template +inline bool IsWithinTolerance(T limit, T tolerance, T value) +{ + return limit + tolerance >= value; +} + +// 分核功能模块输出:FD信息,包含需要归约的数据索引及其分核信息 +struct FlashDecodeResult { + uint32_t fdUsedVecNum{0U}; // 归约过程使用的vector数量 + // 1、归约任务的索引信息 + std::vector fdBN2Idx{}; // 每个归约任务的BN2索引,脚标为归约任务的序号,最大为核数-1 + std::vector fdMIdx{}; // 每个归约任务的GS1索引,脚标为归约任务的序号 + std::vector fdWorkspaceIdx{}; // 每个归约任务在workspace中的存放位置 + std::vector fdS2SplitNum{}; // 每个归约任务的S2核间切分份数,脚标为归约任务的序号 + std::vector fdMSize{}; // 每个归约任务m轴大小,脚标为归约任务的序号 + // 2、FD负载均衡阶段,归约任务的分核(vec)信息 + std::vector fdIdx{}; // FD负载均衡阶段,每个vector处理的归约任务对应ID + std::vector fdMStart{}; // FD负载均衡阶段,每个vector处理的归约任务的m轴起点 + std::vector fdMNum{}; // FD负载均衡阶段,每个vector处理的归约任务的m轴行数 + + FlashDecodeResult(uint32_t aicNum, uint32_t aivNum) + : fdBN2Idx(aicNum), + fdMIdx(aicNum), + fdWorkspaceIdx(aicNum), + fdS2SplitNum(aicNum), + fdMSize(aicNum), + fdIdx(aivNum), + fdMStart(aivNum), + fdMNum(aivNum) + {} +}; + +// 分核功能模块输出:FA阶段的核间分核信息 +struct SplitResult { + uint32_t usedCoreNum{0U}; // 使用的核数量 + std::vector bN2End{}; // 每个核处理数据的BN2结束点 + std::vector gS1End{}; // 每个核处理数据的GS1结束点 + std::vector s2End{}; // 每个核处理数据的S2结束点 + std::vector firstFdDataWorkspaceIdx{}; // 每个核第一份归约任务的存放位置 + int64_t maxCost{0}; // 慢核开销 + uint32_t numOfFdHead{0U}; // 归约任务数量 + uint32_t maxS2SplitNum{0U}; // 单个归约任务最大分核数量 + uint32_t maxS2LoopNum{0U}; // 单个核最大s2计算轮次数量 + FlashDecodeResult fdRes{0U, 0U}; // FD信息 + + SplitResult(uint32_t aicNum, uint32_t aivNum) + : bN2End(aicNum), + gS1End(aicNum), + s2End(aicNum), + firstFdDataWorkspaceIdx(aicNum), + fdRes(aicNum, aivNum) {}; +}; + +// 分核功能模块内部使用:记录切分信息 +struct SplitInfo { + std::vector s1GBaseNum{}; // S1G方向,切了多少个基本块 + std::vector oriS2BaseNum{}; // oriS2方向,切了多少个基本块 + std::vector cmpS2BaseNum{}; // cmpS2方向,切了多少个基本块 + std::vector s1GTailSize{}; // S1G方向,尾块size + std::vector oriS2TailSize{}; // oriS2方向,尾块size + std::vector cmpS2TailSize{}; // cmpS2方向,尾块size + bool isKvSeqAllZero{true}; + + explicit SplitInfo(uint32_t batchSize) + : s1GBaseNum(batchSize), + oriS2BaseNum(batchSize), + cmpS2BaseNum(batchSize), + s1GTailSize(batchSize), + oriS2TailSize(batchSize), + cmpS2TailSize(batchSize) + {} +}; + +// 分核功能模块内部使用:记录batch的开销信息 +struct CostInfo { + std::vector bN2CostOfEachBatch{}; // 整个batch的开销 + std::vector bN2BlockOfEachBatch{}; // 整个batch的开销 + std::vector bN2S2LoopOfEachBatch{}; // 整个batch的s2计算轮次数量 + std::vector bN2LastBlockCostOfEachBatch{}; // batch最后一块的开销 + uint32_t totalBlockNum{0U}; + int64_t totalCost{0}; + int64_t maxS1GCost{0}; // 记录所有S1G行中的最大开销 + + explicit CostInfo(uint32_t batchSize) + : bN2CostOfEachBatch(batchSize), + bN2BlockOfEachBatch(batchSize), + bN2S2LoopOfEachBatch(batchSize), + bN2LastBlockCostOfEachBatch(batchSize) + {} +}; + +// 分核功能模块内部使用:分核过程中,case基本信息的上下文信息,组合以减少接口传参数量 +struct SplitContext { + SplitInfo splitInfo{0U}; + CostInfo costInfo{0U}; + + explicit SplitContext(uint32_t batchSize) + : splitInfo(batchSize), + costInfo(batchSize) + {} +}; + +// 分核功能模块内部使用:记录batch相关的临时信息 +struct BatchCache { + uint32_t bIdx{0U}; + uint32_t s1Size{0U}; + uint32_t oriS2Size{0U}; + uint64_t cmpRevertS2Size{0U}; + int64_t oriPreTokenLeftUp{0}; + int64_t oriNextTokenLeftUp{0}; + int64_t cmpPreTokenLeftUp{0}; + int64_t cmpNextTokenLeftUp{0}; + BlockCost typeCost{}; +}; + +// 分核功能模块内部使用:记录当前行(S1G)的临时信息 +struct S1GCache { + uint32_t bIdx{0U}; + uint32_t s1GIdx{0U}; + uint32_t s2Start{0U}; + uint32_t s2End{0U}; + uint32_t oriS2Start{0U}; + uint32_t oriS2End{0U}; + uint32_t cmpS2Start{0U}; // win部分与cmp部分的切分点 + uint32_t cmpS2End{0U}; + int64_t s1GCost{0}; + int64_t s1GLastBlockCost{0}; + uint32_t s1GBlock{0U}; + uint32_t s2Loop{0U}; + uint32_t oriS1GBlock{0U}; + int64_t oriS1GCost{0}; + int64_t oriS1GLastBlockCost{0}; + int64_t oriS1GNormalBlockCost{0}; + uint32_t cmpS1GBlock{0U}; + int64_t cmpS1GCost{0}; + int64_t cmpS1GLastBlockCost{0}; + int64_t cmpS1GNormalBlockCost{0}; + int64_t oriS2TailSize{0}; + int64_t cmpS2TailSize{0}; + uint32_t actOriS2Size{0U}; + uint32_t actCmpS2Size{0U}; + uint32_t reductionBlockSize{0U}; +}; + +// 分核功能模块内部使用:记录分配过程中,当前核的负载信息 +struct CoreCache { + int64_t costLimit{0}; // 负载上限 + int64_t cost{0}; // 已分配负载 + uint32_t block{0U}; // 已分配块数 + uint32_t s2Loop{0U}; // 已分配s2计算轮次数量 +}; + +// 分核功能模块内部使用:记录分配过程中的上下文信息 +struct AssignContext { + uint32_t curBIdx{0U}; + uint32_t curBN2Idx{0U}; + uint32_t curS1GIdx{0U}; + uint32_t curS2Idx{0U}; + uint32_t curCoreIdx{0U}; + int64_t unassignedCost{0}; + uint32_t curKvSplitPart{1U}; + uint32_t preFdDataNum{0U}; + + int64_t bN2Cost{0}; + uint32_t bN2Block{0U}; + uint32_t bN2S2Loop{0U}; + bool isFinished{false}; + BatchCache batchCache{}; + S1GCache s1GCache{}; + CoreCache coreCache{}; +}; + +class SparseFlashMlaMetadataCpuKernel : public CpuKernel { +public: + SparseFlashMlaMetadataCpuKernel() = default; + ~SparseFlashMlaMetadataCpuKernel() = default; + uint32_t Compute(CpuKernelContext &ctx) override; + +private: + bool Prepare(CpuKernelContext &ctx); + int32_t GetQueryBatchSize(); + void CalcOriMaskMode(); + void CalcCmpMaskMode(); + ValidSocVersion ProcessSocVersion(); + bool ParamsCheck(); + int32_t GetSumOfQuerySeq(); + bool ParamsInit(); + bool BalanceSchedule(SplitResult &splitRes); + bool GenMetadata(SplitResult &splitRes); + // util + uint32_t GetS1SeqSize(uint32_t bIdx); + uint32_t GetOriS2SeqSize(uint32_t bIdx); + uint32_t GetCmpS2SeqSize(uint32_t bIdx); + uint64_t GetRevertS2Size(uint32_t bIdx); + uint32_t GetS1Idx(uint32_t s1Size, uint32_t s1GIdx); + uint32_t GetBsStride(uint32_t bIdx, uint32_t s1Idx); + uint32_t GetOriTopkLength(uint32_t bsStride); + uint32_t GetOriTopkLength(uint32_t s1GIdx, const BatchCache &batchCache) const; + uint32_t ReadOriTopkLengthAtRow(uint32_t bIdx, uint32_t s1Idx) const; + uint32_t GetCmpTopkLength(uint32_t bsStride); + int64_t CalcOriPreTokenLeftUp(uint32_t s1Size, uint32_t s2Size); + int64_t CalcOriNextTokenLeftUp(uint32_t s1Size, uint32_t s2Size); + int64_t CalcCmpPreTokenLeftUp(uint32_t s1Size, uint64_t s2Size); + int64_t CalcCmpNextTokenLeftUp(uint32_t s1Size, uint64_t s2Size); + Range CalcS2TokenRange(uint32_t s1GIdx, const BatchCache &batchCache, bool isCmpKv); + int64_t OriCalcCost(uint32_t basicM, uint32_t basicS2); + int64_t CmpCalcCost(uint32_t basicM, uint32_t basicS2); + void CalcCostTable(uint32_t s1GTailSize, uint32_t reductionBlockSize, uint32_t oriS2TailSize, + uint32_t cmpS2TailSize); + + // cache calculation + void CalcBatchCache(uint32_t bIdx, const SplitContext &splitContext, BatchCache &batchCache); + void CalcOriBlockRange(const Range &oriS2TokenRange, const BatchCache &batchCache, S1GCache &s1GCache); + void CalcCmpBlockRange(const Range &cmpS2TokenRange, const BatchCache &batchCache, S1GCache &s1GCache); + void CalcOriS1GCache(S1GCache &s1GCache, const SplitInfo &splitInfo); + void CalcCmpS1GCache(S1GCache &s1GCache, const SplitInfo &splitInfo); + void GatherOriAndCmpCache(S1GCache &s1GCache); + void CalcS1GCache(uint32_t s1GIdx, const SplitContext &splitContext, const BatchCache &batchCache, + S1GCache &s1GCache); + + // preprocess + void CalcSplitInfo(SplitContext &splitContext); + void CalcBatchCost(uint32_t bIdx, const SplitContext &splitContext, CostInfo &costInfo); + void CalcCostInfo(SplitContext &splitContext); + + // assign + void UpdateCursor(const SplitContext &splitContext, AssignContext &assignContext); + void AssignByBatch(const SplitContext &splitContext, AssignContext &assignContext); + void AssignByRow(const SplitContext &splitContext, AssignContext &assignContext); + int64_t CalcCurBlockCost(const AssignContext &assignContext); + uint32_t CalcCurBlockS2Loop(const AssignContext &assignContext); + void AssignByBlock(const SplitContext &splitContext, AssignContext &assignContext); + void ForceAssign(const SplitContext &splitContext, AssignContext &assignContext); + void AssignBlocksToCore(const SplitContext &splitContext, AssignContext &assignContext, SplitResult &result); + + // FD + bool IsNeedRecordFDInfo(const AssignContext &assignContext, const SplitResult &splitRes); + bool IsFirstReductionBlock(const AssignContext &assignContext, const SplitResult &splitRes); + void RecordFDInfo(const SplitContext &splitContext, const AssignContext &assignContext, SplitResult &result); + + // main + void SplitFD(SplitResult &splitRes); + void CalcSplitPlan(int64_t costLimit, const SplitContext &splitContext, SplitResult &result); + +private: + // input + Tensor *cuSeqlensQ_ = nullptr; + Tensor *cuSeqlensOriKv_ = nullptr; + Tensor *cuSeqlensCmpKv_ = nullptr; + Tensor *sequsedQ_ = nullptr; + Tensor *sequsedOriKv_ = nullptr; + Tensor *sequsedCmpKv_ = nullptr; + Tensor *cmpResidualKv_ = nullptr; + Tensor *oriTopkLength_ = nullptr; + Tensor *cmpTopkLength_ = nullptr; + + // output + Tensor *metadata_ = nullptr; + + // attributes + int32_t batchSize_ = 0; + int32_t maxSeqlenQ_ = 0; + int32_t numHeadsQ_ = 0; + int32_t maxSeqlenOriKv_ = 0; + int32_t maxSeqlenCmpKv_ = 0; + int32_t numHeadsKv_ = 1; + int32_t headDim_ = 0; + int32_t oriTopK_ = 0; + int32_t cmpTopK_ = 0; + int32_t cmpRatio_ = 0; + int32_t oriMaskMode_ = static_cast(SparseMode::BAND); + int32_t cmpMaskMode_ = static_cast(SparseMode::RIGHT_DOWN_CAUSAL); + int64_t oriWinLeft_ = 127; + int64_t oriWinRight_ = 0; + std::string layoutQ_ = "BSND"; + std::string layoutKv_ = "PA_BBND"; + bool hasOriKv_ = true; + bool hasCmpKv_ = true; + uint32_t aicCoreNum_ = optiling::AIC_CORE_MAX_NUM; + uint32_t aivCoreNum_ = optiling::AIV_CORE_MAX_NUM; + bool isBatchConsistency_ = false; + + // attr + std::string socVersion_ = "Ascend950"; + int64_t oriPreToken_ = 0; + int64_t oriNextToken_ = 0; + int64_t cmpPreToken_ = 0; + int64_t cmpNextToken_ = 0; + uint32_t groupSize_ = 0; + uint32_t mBaseSize_ = 0; + uint32_t s2BaseSize_ = 128U; + bool isS1G_ = true; + bool supportFd_ = false; + uint32_t oriAttentionMode_ = HAS_MASK; + uint32_t cmpAttentionMode_ = HAS_MASK; + BlockCost typeCost_ = {}; + bool isSplitG_ = false; + bool isSparseOriKv_ = false; + bool isSparseCmpKv_ = false; + bool hasOriTopkLength_ = false; + uint32_t remainedBlockNum_ = 0U; + +private: + enum class ParamId : uint32_t { + // input + cuSeqlensQ = 0, + cuSeqlensOriKv = 1, + cuSeqlensCmpKv = 2, + sequsedQ = 3, + sequsedOriKv = 4, + sequsedCmpKv = 5, + cmpResidualKv = 6, + oriTopkLength = 7, + cmpTopkLength = 8, + // output + metaData = 0, + }; +}; +} // namespace aicpu + +#endif diff --git a/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.json b/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.json new file mode 100644 index 000000000000..eb1535386361 --- /dev/null +++ b/csrc/attention/sparse_flash_mla_metadata/op_kernel_aicpu/sparse_flash_mla_metadata_aicpu.json @@ -0,0 +1,15 @@ +{ + "SparseFlashMlaMetadata":{ + "opInfo":{ + "computeCost":"100", + "engine":"DNN_VM_AICPU", + "flagAsync":"False", + "flagPartial":"False", + "functionName":"RunCpuKernel", + "kernelSo":"libtransformer_aicpu_kernels.so", + "opKernelLib":"CUSTAICPUKernel", + "userDefined":"True", + "workspaceSize":"100" + } + } +} diff --git a/csrc/build_aclnn.sh b/csrc/build_aclnn.sh index 2be057999439..a12287161320 100755 --- a/csrc/build_aclnn.sh +++ b/csrc/build_aclnn.sh @@ -115,6 +115,8 @@ elif [[ "$SOC_VERSION" =~ ^ascend910b ]]; then "vllm_quant_lightning_indexer_metadata" "quant_lightning_indexer_v2" "quant_lightning_indexer_v2_metadata" + "sparse_flash_mla" + "sparse_flash_mla_metadata" "sparse_attn_sharedkv" "sparse_attn_sharedkv_metadata" "hc_pre" @@ -171,6 +173,8 @@ elif [[ "$SOC_VERSION" =~ ^ascend910_93 ]]; then "vllm_quant_lightning_indexer_metadata" "quant_lightning_indexer_v2" "quant_lightning_indexer_v2_metadata" + "sparse_flash_mla" + "sparse_flash_mla_metadata" "sparse_attn_sharedkv" "sparse_attn_sharedkv_metadata" "hc_pre" diff --git a/csrc/moe/hc_pre/op_host/hc_pre_def.cpp b/csrc/moe/hc_pre/op_host/hc_pre_def.cpp index 12027002706d..569217f34741 100644 --- a/csrc/moe/hc_pre/op_host/hc_pre_def.cpp +++ b/csrc/moe/hc_pre/op_host/hc_pre_def.cpp @@ -41,6 +41,11 @@ class HcPre : public OpDef { .DataType({ge::DT_FLOAT}) .Format({ge::FORMAT_ND}) .UnknownShapeFormat({ge::FORMAT_ND}); + this->Input("pre_mix") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT}) + .Format({ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND}); this->Output("y") .ParamType(REQUIRED) .DataType({ge::DT_BF16}) @@ -56,6 +61,11 @@ class HcPre : public OpDef { .DataType({ge::DT_FLOAT}) .Format({ge::FORMAT_ND}) .UnknownShapeFormat({ge::FORMAT_ND}); + this->Output("pre") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT}) + .Format({ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND}); this->Attr("hc_mult").AttrType(OPTIONAL).Int(4); this->Attr("hc_sinkhorn_iters").AttrType(OPTIONAL).Int(20); diff --git a/csrc/moe/hc_pre/op_host/hc_pre_tiling.cpp b/csrc/moe/hc_pre/op_host/hc_pre_tiling.cpp index e58ed6bc2044..2115791b7496 100644 --- a/csrc/moe/hc_pre/op_host/hc_pre_tiling.cpp +++ b/csrc/moe/hc_pre/op_host/hc_pre_tiling.cpp @@ -134,6 +134,9 @@ ge::graphStatus HcPreTiling::GetShapeAttrsInfoInner() "hc_base size should be equal with mixhc, but is %ld", baseFirstDim), return ge::GRAPH_FAILED); + tilingData_.set_hasPreMix(context_->GetInputShape(4) != nullptr ? 1 : 0); + tilingData_.set_hasPreOut(context_->GetOutputShape(3) != nullptr ? 1 : 0); + OPS_ERR_IF(GetAttr() != ge::GRAPH_SUCCESS, OPS_LOG_E(context_->GetNodeName(), "get attr failed."), return ge::GRAPH_FAILED); @@ -258,7 +261,6 @@ ge::graphStatus HcPreTiling::CalcMKSplitCoreMembasePart2Tiling() tilingData_.set_d(d_); tilingData_.set_hcMultAlign(hcMultAlign_); tilingData_.set_rowOfFormerBlock(rowOfFormerBlock_); - tilingData_.set_rowOfTailBlock(rowOfTailBlock_); tilingData_.set_rowLoopOfFormerBlock(rowLoopOfFormerBlock_); tilingData_.set_rowLoopOfTailBlock(rowLoopOfTailBlock_); tilingData_.set_stage2RowFactor(rowFactor_); diff --git a/csrc/moe/hc_pre/op_host/hc_pre_tiling.h b/csrc/moe/hc_pre/op_host/hc_pre_tiling.h index 83781f5b55d6..b9d2d8d315fd 100644 --- a/csrc/moe/hc_pre/op_host/hc_pre_tiling.h +++ b/csrc/moe/hc_pre/op_host/hc_pre_tiling.h @@ -48,7 +48,6 @@ TILING_DATA_FIELD_DEF(int64_t, hcMult); TILING_DATA_FIELD_DEF(int64_t, d); TILING_DATA_FIELD_DEF(int64_t, hcMultAlign); TILING_DATA_FIELD_DEF(int64_t, rowOfFormerBlock); -TILING_DATA_FIELD_DEF(int64_t, rowOfTailBlock); TILING_DATA_FIELD_DEF(int64_t, rowLoopOfFormerBlock); TILING_DATA_FIELD_DEF(int64_t, rowLoopOfTailBlock); TILING_DATA_FIELD_DEF(int64_t, rowFactor); @@ -99,6 +98,8 @@ TILING_DATA_FIELD_DEF(int64_t, stage1MFactor); TILING_DATA_FIELD_DEF(int64_t, bufferPool0Size); TILING_DATA_FIELD_DEF(int64_t, bufferPool1Size); TILING_DATA_FIELD_DEF(int64_t, mUbSize); +TILING_DATA_FIELD_DEF(int64_t, hasPreMix); +TILING_DATA_FIELD_DEF(int64_t, hasPreOut); END_TILING_DATA_DEF; @@ -158,4 +159,4 @@ class HcPreTiling { }; } // namespace optiling -#endif // HC_PRE_SINKHORN_TILING_H \ No newline at end of file +#endif // HC_PRE_SINKHORN_TILING_H diff --git a/csrc/moe/hc_pre/op_host/hc_pre_tiling_arch35.h b/csrc/moe/hc_pre/op_host/hc_pre_tiling_arch35.h index 4b57cb7b0088..a4519017dca9 100644 --- a/csrc/moe/hc_pre/op_host/hc_pre_tiling_arch35.h +++ b/csrc/moe/hc_pre/op_host/hc_pre_tiling_arch35.h @@ -267,7 +267,6 @@ ge::graphStatus HcPreTilingRegbase::CalcRegbaseOpTiling() tilingData_.set_d(d_); tilingData_.set_hcMultAlign(hcMultAlign_); tilingData_.set_rowOfFormerBlock(rowOfFormerBlock_); - tilingData_.set_rowOfTailBlock(rowOfTailBlock_); tilingData_.set_rowLoopOfFormerBlock(rowLoopOfFormerBlock_); tilingData_.set_rowLoopOfTailBlock(rowLoopOfTailBlock_); tilingData_.set_rowFactor(rowFactor_); @@ -382,7 +381,6 @@ ge::graphStatus HcPreTilingRegbase::CalcMKSplitCorePart2Tiling() tilingData_.set_d(d_); tilingData_.set_hcMultAlign(hcMultAlign_); tilingData_.set_rowOfFormerBlock(rowOfFormerBlock_); - tilingData_.set_rowOfTailBlock(rowOfTailBlock_); tilingData_.set_rowLoopOfFormerBlock(rowLoopOfFormerBlock_); tilingData_.set_rowLoopOfTailBlock(rowLoopOfTailBlock_); tilingData_.set_stage2RowFactor(rowFactor_); diff --git a/csrc/moe/hc_pre/op_kernel/hc_pre.cpp b/csrc/moe/hc_pre/op_kernel/hc_pre.cpp index 7ccc1986c802..cdf8ddb17ec6 100644 --- a/csrc/moe/hc_pre/op_kernel/hc_pre.cpp +++ b/csrc/moe/hc_pre/op_kernel/hc_pre.cpp @@ -31,8 +31,8 @@ using namespace AscendC; extern "C" __global__ __aicore__ void hc_pre(GM_ADDR x, GM_ADDR hc_fn, GM_ADDR hc_scale, GM_ADDR hc_base, - GM_ADDR y, GM_ADDR post, GM_ADDR comb_frag, GM_ADDR workspace, - GM_ADDR tiling) + GM_ADDR pre_mix, GM_ADDR y, GM_ADDR post, GM_ADDR comb_frag, + GM_ADDR pre, GM_ADDR workspace, GM_ADDR tiling) { KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2); if (workspace == nullptr) { @@ -80,10 +80,10 @@ extern "C" __global__ __aicore__ void hc_pre(GM_ADDR x, GM_ADDR hc_fn, GM_ADDR h TPipe pipeStage2; HcPre::HcPreMembaseKSplitCorePart2 op2; - op2.Init(x, hc_scale, hc_base, y, post, comb_frag, userWs, tilingData, &pipeStage2); + op2.Init(x, hc_scale, hc_base, y, post, comb_frag, pre_mix, pre, userWs, tilingData, &pipeStage2); op2.Process(); pipeStage2.Destroy(); } #endif -} \ No newline at end of file +} diff --git a/csrc/moe/hc_pre/op_kernel/hc_pre_m_k_split_core.h b/csrc/moe/hc_pre/op_kernel/hc_pre_m_k_split_core.h index 412b06f4ec21..82b4e7fe7097 100644 --- a/csrc/moe/hc_pre/op_kernel/hc_pre_m_k_split_core.h +++ b/csrc/moe/hc_pre/op_kernel/hc_pre_m_k_split_core.h @@ -214,7 +214,7 @@ class HcPreMembaseKSplitCorePart2 { __aicore__ inline void InitGlobalBuffers(GM_ADDR x, GM_ADDR hcScale, GM_ADDR hcBase, GM_ADDR y, GM_ADDR post, GM_ADDR combFrag, - GM_ADDR workspace) + GM_ADDR preMix, GM_ADDR pre, GM_ADDR workspace) { xGm.SetGlobalBuffer((__gm__ T *)x); hcScaleGm.SetGlobalBuffer((__gm__ float *)hcScale); @@ -222,6 +222,8 @@ class HcPreMembaseKSplitCorePart2 { yGm.SetGlobalBuffer((__gm__ T *)y); postGm.SetGlobalBuffer((__gm__ float *)post); combFragGm.SetGlobalBuffer((__gm__ float *)combFrag); + preMixGm.SetGlobalBuffer((__gm__ float *)preMix); + preGm.SetGlobalBuffer((__gm__ float *)pre); workspaceGm.SetGlobalBuffer((__gm__ float *)workspace); } @@ -275,6 +277,8 @@ class HcPreMembaseKSplitCorePart2 { pipe->InitBuffer(maskPatternBuf, RoundUp(MASK_PATTERN_BASE_SIZE * MASK_PATTERN_REPEAT_SIZE) * sizeof(uint32_t)); + pipe->InitBuffer(preMixBuf, + tilingData->stage2RowFactor * tilingData->hcMultAlign * sizeof(float)); } __aicore__ inline void GetLocalTensors() @@ -292,16 +296,17 @@ class HcPreMembaseKSplitCorePart2 { yCastLocal = yCastBuf.Get(); rsqrtLocal = rsqrtBuf.Get(); maskPatternLocal = maskPatternBuf.Get(); + preMixLocal = preMixBuf.Get(); SetGatherMaskPattern(maskPatternLocal); } __aicore__ inline void Init(GM_ADDR x, GM_ADDR hcScale, GM_ADDR hcBase, - GM_ADDR y, GM_ADDR post, GM_ADDR combFrag, GM_ADDR workspace, - const HcPreTilingData *tilingDataPtr, TPipe *pipePtr) + GM_ADDR y, GM_ADDR post, GM_ADDR combFrag, GM_ADDR preMix, GM_ADDR pre, + GM_ADDR workspace, const HcPreTilingData *tilingDataPtr, TPipe *pipePtr) { pipe = pipePtr; tilingData = tilingDataPtr; - InitGlobalBuffers(x, hcScale, hcBase, y, post, combFrag, workspace); + InitGlobalBuffers(x, hcScale, hcBase, y, post, combFrag, preMix, pre, workspace); int64_t stage1UsedCoreNum = tilingData->cubeBlockDimK; int64_t xQueNum2 = tilingData->stage2RowFactor * tilingData->hcMult * RoundUp(tilingData->dFactor); @@ -393,6 +398,23 @@ int64_t curBsIdxForAll = (stage2BlockIdx * tilingData->rowLoopOfFormerBlock + ProcessPre(mixes01ReduceLocal, mixes01ReduceLocal, hcBase0Local, rsqrtLocal, rowBrcbLocal0, hcBrcbLocal1, hcScaleGm.GetValue(0), tilingData->hcEps, curRowFactor, tilingData->hcMult); + if (tilingData->hasPreOut != 0) { + event_t eventIdPreOut = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3)); + SetFlag(eventIdPreOut); + WaitFlag(eventIdPreOut); + CopyOut(mixes01ReduceLocal, + preGm[stage2BlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult + + rowOuterIdx * tilingData->stage2RowFactor * tilingData->hcMult], + 1, curRowFactor * tilingData->hcMult); + } + if (tilingData->hasPreMix != 0) { + CopyIn(preMixGm[stage2BlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult + + rowOuterIdx * tilingData->stage2RowFactor * tilingData->hcMult], + preMixLocal, 1, curRowFactor * tilingData->hcMult); + event_t eventIdPreMix = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V)); + SetFlag(eventIdPreMix); + WaitFlag(eventIdPreMix); + } for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++) { int64_t curDFactor = (dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor; @@ -404,7 +426,8 @@ int64_t curBsIdxForAll = (stage2BlockIdx * tilingData->rowLoopOfFormerBlock + xQue.template EnQue(xLocal); xLocal = xQue.template DeQue(); yLocal = yQue.template AllocTensor(); - ProcessY(yLocal, xLocal, mixes01ReduceLocal, hcBrcbLocal1, xCastLocal, yCastLocal, curRowFactor, + ProcessY(yLocal, xLocal, tilingData->hasPreMix != 0 ? preMixLocal : mixes01ReduceLocal, + hcBrcbLocal1, xCastLocal, yCastLocal, curRowFactor, tilingData->hcMult, curDFactor); xQue.template FreeTensor(xLocal); yQue.template EnQue(yLocal); @@ -510,6 +533,8 @@ int64_t curBsIdxForAll = (stage2BlockIdx * tilingData->rowLoopOfFormerBlock + GlobalTensor yGm; GlobalTensor postGm; GlobalTensor combFragGm; + GlobalTensor preMixGm; + GlobalTensor preGm; TQue mixesQue01; TQue mixesQue2; @@ -536,6 +561,7 @@ int64_t curBsIdxForAll = (stage2BlockIdx * tilingData->rowLoopOfFormerBlock + TBuf xCastBuf; TBuf yCastBuf; TBuf maskPatternBuf; + TBuf preMixBuf; LocalTensor mixes01Local; LocalTensor mixes2Local; @@ -557,8 +583,9 @@ int64_t curBsIdxForAll = (stage2BlockIdx * tilingData->rowLoopOfFormerBlock + LocalTensor yCastLocal; LocalTensor squareSumOutLocal; LocalTensor maskPatternLocal; + LocalTensor preMixLocal; }; } // namespace HcPreSinkhorn -#endif \ No newline at end of file +#endif diff --git a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_def.cpp b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_def.cpp index 00136a478bf3..4ecc33bcf35e 100644 --- a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_def.cpp +++ b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_def.cpp @@ -79,6 +79,21 @@ class MoeGatingTopKHash : public OpDef { ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) .AutoContiguous(); + this->Input("bias_vl") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, + ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, + ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, + ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); this->Output("y") .ParamType(REQUIRED) .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, @@ -130,6 +145,8 @@ class MoeGatingTopKHash : public OpDef { this->Attr("out_flag").AttrType(OPTIONAL).Bool(false); this->Attr("routed_scaling_factor").AttrType(OPTIONAL).Float(1.0); this->Attr("eps").AttrType(OPTIONAL).Float(1e-20f); + this->Attr("image_sentinel_lo").AttrType(OPTIONAL).Int(129257); + this->Attr("image_sentinel_count").AttrType(OPTIONAL).Int(5); this->AICore().AddConfig("ascend910b"); this->AICore().AddConfig("ascend910_93"); @@ -143,4 +160,4 @@ class MoeGatingTopKHash : public OpDef { }; OP_ADD(MoeGatingTopKHash); -} // namespace ops \ No newline at end of file +} // namespace ops diff --git a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_proto.cpp b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_proto.cpp index f1b1d26ec062..c20506f48ae8 100644 --- a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_proto.cpp +++ b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_proto.cpp @@ -25,6 +25,9 @@ namespace ge { * @par Inputs: * @li x: A 2D tensor which moe gating topk is applied, The shape is: (B*S, E), format supports ND, and data type must be float16, float or bfloat16. E(Expert num) can not be greater than 2048. E(Expert num) should be divisible by group_count. * @li bias: A 1D tensor which is "bias" in moe gating topk. The shape is: (E), format supports ND, and data type must be the same as that of x. + * @li bias_vl: An optional 1D vision-token correction bias. When present, input_ids in + * [image_sentinel_lo, image_sentinel_lo + image_sentinel_count) use this bias; + * all other rows use bias. * * @par Outputs: * @li y: A 2D tensor which is the topk value result of moe gating topk, format supports ND, and data type must be the same as that of x. @@ -49,6 +52,7 @@ REG_OP(MoeGatingTopKHash) .OPTIONAL_INPUT(bias, TensorType({DT_FLOAT, DT_FLOAT16, DT_BF16})) .OPTIONAL_INPUT(input_ids, TensorType({DT_INT64, DT_INT32})) .OPTIONAL_INPUT(tid2eid, TensorType({DT_INT64, DT_INT32})) + .OPTIONAL_INPUT(bias_vl, TensorType({DT_FLOAT, DT_FLOAT16, DT_BF16})) .OUTPUT(y, TensorType({DT_FLOAT, DT_FLOAT16, DT_BF16})) .OUTPUT(expert_idx, TensorType({DT_INT32})) .OUTPUT(out, TensorType({DT_FLOAT})) @@ -61,8 +65,10 @@ REG_OP(MoeGatingTopKHash) .ATTR(out_flag, Bool, false) .ATTR(routed_scaling_factor, Float, 1.0) .ATTR(eps, Float, 1e-20f) + .ATTR(image_sentinel_lo, Int, 129257) + .ATTR(image_sentinel_count, Int, 5) .OP_END_FACTORY_REG(MoeGatingTopKHash) } // namespace ge -#endif // OPS_OP_PROTO_INC_MOEGATINGTOPK_H_ \ No newline at end of file +#endif // OPS_OP_PROTO_INC_MOEGATINGTOPK_H_ diff --git a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling.cpp b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling.cpp index b2f4ae5be0bd..ab4e185fd6d5 100644 --- a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling.cpp +++ b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling.cpp @@ -37,6 +37,7 @@ const static int64_t X_INPUT_INDEX = 0; const static int64_t BIAS_INPUT_INDEX = 1; const static int64_t INPUT_IDS_INPUT_INDEX = 2; const static int64_t TID_TO_EID_INPUT_INDEX = 3; +const static int64_t BIAS_VL_INPUT_INDEX = 4; const static int64_t Y_OUTPUT_INDEX = 0; const static int64_t EXPERT_IDX_OUTPUT_INDEX = 1; const static int64_t OUT_OUTPUT_INDEX = 2; @@ -49,6 +50,8 @@ const static int64_t NORM_TYPE_ATTR_INDEX = 5; const static int64_t OUT_FLAG_ATTR_INDEX = 6; const static int64_t ROUTED_SCALING_FACTOR_ATTR_INDEX = 7; const static int64_t EPS_ATTR_INDEX = 8; +const static int64_t IMAGE_SENTINEL_LO_ATTR_INDEX = 9; +const static int64_t IMAGE_SENTINEL_COUNT_ATTR_INDEX = 10; const static int64_t DEFAULT_WORKSPACE_SIZE = 16777216; // 预留16M空间 const static uint32_t DATATYPESIZE_FLOAT = 4; const static bool IS_LARGEST = true; @@ -124,16 +127,18 @@ class MoeGatingTopKHashTilingBase { const gert::Shape *biasShape_ = nullptr; const gert::Shape *inputIdsShape_ = nullptr; const gert::Shape *tid2eidShape_ = nullptr; + const gert::Shape *biasVlShape_ = nullptr; const gert::Shape *yShape_ = nullptr; const gert::Shape *expertIdxShape_ = nullptr; const gert::Shape *outShape_ = nullptr; - ge::DataType inputIdsDtype; - ge::DataType tid2eidDtype; + ge::DataType inputIdsDtype = ge::DataType::DT_INT32; + ge::DataType tid2eidDtype = ge::DataType::DT_INT32; uint64_t coreNum_ = 0; int64_t rows_ = 0; int64_t expertCount_ = 0; int64_t addBias_ = 0; + int64_t addBiasVl_ = 0; int64_t k_ = 0; int64_t kGroup_ = 0; @@ -144,6 +149,8 @@ class MoeGatingTopKHashTilingBase { int64_t normType_ = NORM_TYPE_SOFTMAX; int64_t outFlag_ = OUT_FLAG_FALSE; int64_t hashFlag_ = 0; + int64_t imageSentinelLo_ = 129257; + int64_t imageSentinelCount_ = 5; float routedScalingFactor_ = 1.0; float eps_ = 1e-20f; @@ -179,12 +186,19 @@ ge::graphStatus MoeGatingTopKHashTilingBase::CheckInputShape() } moeGatingTopKTilingData_.set_addBias(addBias_); - if (inputIdsShape_ != nullptr) { + if (biasVlShape_ != nullptr) { + addBiasVl_ = 1; + size_t biasVlDimNum = biasVlShape_->GetDimNum(); OPS_ERR_IF( - tid2eidShape_ == nullptr, - OPS_LOG_E(context_, "The tid2eid should not be empty when inputIds has value."), + biasVlDimNum != BIAS_INPUT_DIMS || biasVlShape_->GetDim(0) != expertCount_, + OPS_LOG_E(context_, "bias_vl must be a 1D tensor with expertCount elements."), + return ge::GRAPH_FAILED); + OPS_ERR_IF(inputIdsShape_ == nullptr, + OPS_LOG_E(context_, "input_ids is required when bias_vl is present."), return ge::GRAPH_FAILED); } + moeGatingTopKTilingData_.set_addBiasVl(addBiasVl_); + if (tid2eidShape_ != nullptr) { OPS_ERR_IF( inputIdsShape_ == nullptr, @@ -231,6 +245,13 @@ ge::graphStatus MoeGatingTopKHashTilingBase::CheckAttr() OPS_LOG_E(context_, "norm type softplus only supported when groupCount equals 1, but got %ld.", groupCount_), return ge::GRAPH_FAILED); + OPS_ERR_IF(addBiasVl_ && groupCount_ != 1, + OPS_LOG_E(context_, "bias_vl routing currently requires groupCount=1, but got %ld.", groupCount_), + return ge::GRAPH_FAILED); + OPS_ERR_IF(addBiasVl_ && imageSentinelCount_ <= 0, + OPS_LOG_E(context_, "image_sentinel_count must be positive when bias_vl is present."), + return ge::GRAPH_FAILED); + OPS_ERR_IF(groupSelectMode_ != GROUP_SELECT_MODE_SUM && groupSelectMode_ != GROUP_SELECT_MODE_MAX, OPS_LOG_E(context_, "group select mode is: %ld, but currently only support %ld and %ld.", groupSelectMode_, GROUP_SELECT_MODE_SUM, GROUP_SELECT_MODE_MAX), @@ -288,6 +309,9 @@ ge::graphStatus MoeGatingTopKHashTilingBase::GetShapeAttrsInfo() auto tid2eidShapePtr = context_->GetOptionalInputShape(TID_TO_EID_INPUT_INDEX); tid2eidShape_ = tid2eidShapePtr == nullptr ? nullptr : &tid2eidShapePtr->GetStorageShape(); + auto biasVlShapePtr = context_->GetOptionalInputShape(BIAS_VL_INPUT_INDEX); + biasVlShape_ = biasVlShapePtr == nullptr ? nullptr : &biasVlShapePtr->GetStorageShape(); + // 获取输出shape auto yShapePtr = context_->GetOutputShape(Y_OUTPUT_INDEX); OPS_LOG_E_IF_NULL(context_, yShapePtr, return ge::GRAPH_FAILED); @@ -316,6 +340,14 @@ ge::graphStatus MoeGatingTopKHashTilingBase::GetShapeAttrsInfo() ge::TypeUtils::DataTypeToSerialString(xDtype).c_str()), return ge::GRAPH_FAILED); } + if (biasVlShapePtr != nullptr) { + auto biasVlDtype = context_->GetOptionalInputDesc(BIAS_VL_INPUT_INDEX)->GetDataType(); + OPS_ERR_IF((biasVlDtype != xDtype), + OPS_LOG_E(context_, "bias_vl dtype %s not equal x dtype %s, please check.", + ge::TypeUtils::DataTypeToSerialString(biasVlDtype).c_str(), + ge::TypeUtils::DataTypeToSerialString(xDtype).c_str()), + return ge::GRAPH_FAILED); + } if (inputIdsShapePtr != nullptr) { inputIdsDtype = context_->GetOptionalInputDesc(INPUT_IDS_INPUT_INDEX)->GetDataType(); OPS_ERR_IF((inputIdsDtype != ge::DataType::DT_INT32 && inputIdsDtype != ge::DataType::DT_INT64), @@ -422,6 +454,18 @@ ge::graphStatus MoeGatingTopKHashTilingBase::GetShapeAttrsInfo() } OPS_LOG_I(context_, "Attr eps is: %f ", eps_); + const int64_t *imageSentinelLoPtr = attrs->GetAttrPointer(IMAGE_SENTINEL_LO_ATTR_INDEX); + if (imageSentinelLoPtr != nullptr) { + imageSentinelLo_ = *imageSentinelLoPtr; + } + moeGatingTopKTilingData_.set_imageSentinelLo(imageSentinelLo_); + + const int64_t *imageSentinelCountPtr = attrs->GetAttrPointer(IMAGE_SENTINEL_COUNT_ATTR_INDEX); + if (imageSentinelCountPtr != nullptr) { + imageSentinelCount_ = *imageSentinelCountPtr; + } + moeGatingTopKTilingData_.set_imageSentinelCount(imageSentinelCount_); + inputDtypeSize_ = static_cast(ge::GetSizeByDataType(context_->GetInputDesc(0)->GetDataType())); return ge::GRAPH_SUCCESS; } @@ -637,4 +681,4 @@ static ge::graphStatus TilingPrepareForMoeGatingTopKHash(gert::TilingParseContex IMPL_OP_OPTILING(MoeGatingTopKHash) .Tiling(TilingForMoeGatingTopKHash) .TilingParse(TilingPrepareForMoeGatingTopKHash); -} // namespace optiling \ No newline at end of file +} // namespace optiling diff --git a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling.h b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling.h index 8c47b37cff3a..ed1225c39631 100644 --- a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling.h +++ b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling.h @@ -39,6 +39,7 @@ TILING_DATA_FIELD_DEF(int64_t, perCoreRowCount); TILING_DATA_FIELD_DEF(int64_t, lastCoreRowCount); TILING_DATA_FIELD_DEF(int64_t, expertCount); TILING_DATA_FIELD_DEF(int64_t, addBias); +TILING_DATA_FIELD_DEF(int64_t, addBiasVl); TILING_DATA_FIELD_DEF(int64_t, k); TILING_DATA_FIELD_DEF(int64_t, kGroup); TILING_DATA_FIELD_DEF(int64_t, groupCount); @@ -49,6 +50,8 @@ TILING_DATA_FIELD_DEF(int64_t, renorm); TILING_DATA_FIELD_DEF(int64_t, normType); TILING_DATA_FIELD_DEF(int64_t, outFlag); TILING_DATA_FIELD_DEF(int64_t, hashFlag); +TILING_DATA_FIELD_DEF(int64_t, imageSentinelLo); +TILING_DATA_FIELD_DEF(int64_t, imageSentinelCount); TILING_DATA_FIELD_DEF(int64_t, vmsCount); TILING_DATA_FIELD_DEF(float, routedScalingFactor); TILING_DATA_FIELD_DEF(float, eps); @@ -63,6 +66,7 @@ TILING_DATA_FIELD_DEF(int64_t, perCoreRowCount); TILING_DATA_FIELD_DEF(int64_t, lastCoreRowCount); TILING_DATA_FIELD_DEF(int64_t, expertCount); TILING_DATA_FIELD_DEF(int64_t, addBias); +TILING_DATA_FIELD_DEF(int64_t, addBiasVl); TILING_DATA_FIELD_DEF(int64_t, k); TILING_DATA_FIELD_DEF(int64_t, kGroup); TILING_DATA_FIELD_DEF(int64_t, groupCount); @@ -73,6 +77,8 @@ TILING_DATA_FIELD_DEF(int64_t, renorm); TILING_DATA_FIELD_DEF(int64_t, normType); TILING_DATA_FIELD_DEF(int64_t, outFlag); TILING_DATA_FIELD_DEF(int64_t, hashFlag); +TILING_DATA_FIELD_DEF(int64_t, imageSentinelLo); +TILING_DATA_FIELD_DEF(int64_t, imageSentinelCount); TILING_DATA_FIELD_DEF(int64_t, vmsCount); TILING_DATA_FIELD_DEF(float, routedScalingFactor); TILING_DATA_FIELD_DEF(float, eps); diff --git a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling_arch35.h b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling_arch35.h index 3c13c08566f8..e56d931bda71 100644 --- a/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling_arch35.h +++ b/csrc/moe/moe_gating_top_k_hash/op_host/moe_gating_top_k_hash_tiling_arch35.h @@ -44,6 +44,7 @@ const static int64_t X_INPUT_INDEX = 0; const static int64_t BIAS_INPUT_INDEX = 1; const static int64_t INPUT_IDS_INPUT_INDEX = 2; const static int64_t TID_TO_EID_INPUT_INDEX = 3; +const static int64_t BIAS_VL_INPUT_INDEX = 4; const static int64_t Y_OUTPUT_INDEX = 0; const static int64_t EXPERT_IDX_OUTPUT_INDEX = 1; const static int64_t OUT_OUTPUT_INDEX = 2; @@ -117,6 +118,7 @@ class MoeGatingTopKHashTilingRegbase { const gert::Shape *outShape_ = nullptr; const gert::Shape *inputIdsShape_ = nullptr; const gert::Shape *tid2eidShape_ = nullptr; + const gert::Shape *biasVlShape_ = nullptr; ge::DataType inputIdsDtype; ge::DataType tid2eidDtype; @@ -174,6 +176,13 @@ ge::graphStatus MoeGatingTopKHashTilingRegbase::CheckInputShape() return ge::GRAPH_FAILED); } moeGatingTopKTilingData_.set_addBias(addBias_); + moeGatingTopKTilingData_.set_addBiasVl(0); + moeGatingTopKTilingData_.set_imageSentinelLo(129257); + moeGatingTopKTilingData_.set_imageSentinelCount(5); + + OPS_ERR_IF(biasVlShape_ != nullptr, + OPS_LOG_E(context_, "bias_vl routing is not implemented for Ascend 950."), + return ge::GRAPH_FAILED); if (inputIdsShape_ != nullptr) { OPS_ERR_IF( @@ -277,6 +286,8 @@ ge::graphStatus MoeGatingTopKHashTilingRegbase::GetShapeAttrsInfo() inputIdsShape_ = inputIdsShapePtr == nullptr ? nullptr : &inputIdsShapePtr->GetStorageShape(); auto tid2eidShapePtr = context_->GetOptionalInputShape(TID_TO_EID_INPUT_INDEX); tid2eidShape_ = tid2eidShapePtr == nullptr ? nullptr : &tid2eidShapePtr->GetStorageShape(); + auto biasVlShapePtr = context_->GetOptionalInputShape(BIAS_VL_INPUT_INDEX); + biasVlShape_ = biasVlShapePtr == nullptr ? nullptr : &biasVlShapePtr->GetStorageShape(); // 获取输出shape auto yShapePtr = context_->GetOutputShape(Y_OUTPUT_INDEX); diff --git a/csrc/moe/moe_gating_top_k_hash/op_kernel/moe_gating_top_k_hash.cpp b/csrc/moe/moe_gating_top_k_hash/op_kernel/moe_gating_top_k_hash.cpp index 807713b1c81c..6052615fa2f4 100644 --- a/csrc/moe/moe_gating_top_k_hash/op_kernel/moe_gating_top_k_hash.cpp +++ b/csrc/moe/moe_gating_top_k_hash/op_kernel/moe_gating_top_k_hash.cpp @@ -35,8 +35,10 @@ using namespace AscendC; using namespace MoeGatingTopKHash; -extern "C" __global__ __aicore__ void moe_gating_top_k_hash(GM_ADDR x, GM_ADDR bias, GM_ADDR inputIds, GM_ADDR tid2eid, GM_ADDR y, GM_ADDR expertIdx, - GM_ADDR out, GM_ADDR workspace, GM_ADDR tiling) +extern "C" __global__ __aicore__ void moe_gating_top_k_hash(GM_ADDR x, GM_ADDR bias, GM_ADDR inputIds, + GM_ADDR tid2eid, GM_ADDR biasVl, GM_ADDR y, + GM_ADDR expertIdx, GM_ADDR out, GM_ADDR workspace, + GM_ADDR tiling) { KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY); if (g_coreType == AIC) { @@ -63,31 +65,31 @@ extern "C" __global__ __aicore__ void moe_gating_top_k_hash(GM_ADDR x, GM_ADDR b GET_TILING_DATA_WITH_STRUCT(MoeGatingTopKHashTilingData, tilingData, tiling); const MoeGatingTopKHashTilingData *__restrict t = &tilingData; MoeGatingTopKHashWithoutGroup op; - op.Init(x, bias, inputIds, tid2eid, y, expertIdx, out, userWS, t, &tPipe); + op.Init(x, bias, inputIds, tid2eid, biasVl, y, expertIdx, out, userWS, t, &tPipe); op.Process(); } else if (TILING_KEY_IS(TILING_KEY_WITHOUT_GROUP_1)) { GET_TILING_DATA_WITH_STRUCT(MoeGatingTopKHashTilingData, tilingData, tiling); const MoeGatingTopKHashTilingData *__restrict t = &tilingData; MoeGatingTopKHashWithoutGroup op; - op.Init(x, bias, inputIds, tid2eid, y, expertIdx, out, userWS, t, &tPipe); + op.Init(x, bias, inputIds, tid2eid, biasVl, y, expertIdx, out, userWS, t, &tPipe); op.Process(); } else if (TILING_KEY_IS(TILING_KEY_WITHOUT_GROUP_2)) { GET_TILING_DATA_WITH_STRUCT(MoeGatingTopKHashTilingData, tilingData, tiling); const MoeGatingTopKHashTilingData *__restrict t = &tilingData; MoeGatingTopKHashWithoutGroup op; - op.Init(x, bias, inputIds, tid2eid, y, expertIdx, out, userWS, t, &tPipe); + op.Init(x, bias, inputIds, tid2eid, biasVl, y, expertIdx, out, userWS, t, &tPipe); op.Process(); } else if (TILING_KEY_IS(TILING_KEY_WITHOUT_GROUP_3)) { GET_TILING_DATA_WITH_STRUCT(MoeGatingTopKHashTilingData, tilingData, tiling); const MoeGatingTopKHashTilingData *__restrict t = &tilingData; MoeGatingTopKHashWithoutGroup op; - op.Init(x, bias, inputIds, tid2eid, y, expertIdx, out, userWS, t, &tPipe); + op.Init(x, bias, inputIds, tid2eid, biasVl, y, expertIdx, out, userWS, t, &tPipe); op.Process(); } else if (TILING_KEY_IS(TILING_KEY_WITHOUT_GROUP_4)) { GET_TILING_DATA_WITH_STRUCT(MoeGatingTopKHashTilingData, tilingData, tiling); const MoeGatingTopKHashTilingData *__restrict t = &tilingData; MoeGatingTopKHashWithoutGroup op; - op.Init(x, bias, inputIds, tid2eid, y, expertIdx, out, userWS, t, &tPipe); + op.Init(x, bias, inputIds, tid2eid, biasVl, y, expertIdx, out, userWS, t, &tPipe); op.Process(); } else if (TILING_KEY_IS(TILING_KEY_GENERALIZED)) { GET_TILING_DATA_WITH_STRUCT(MoeGatingTopKHashTilingData, tilingData, tiling); diff --git a/csrc/moe/moe_gating_top_k_hash/op_kernel/moe_gating_top_k_hash_without_group.h b/csrc/moe/moe_gating_top_k_hash/op_kernel/moe_gating_top_k_hash_without_group.h index 03b5e1303118..fd1fc336ea27 100644 --- a/csrc/moe/moe_gating_top_k_hash/op_kernel/moe_gating_top_k_hash_without_group.h +++ b/csrc/moe/moe_gating_top_k_hash/op_kernel/moe_gating_top_k_hash_without_group.h @@ -24,17 +24,18 @@ template class MoeGatingTopKHashWithoutGroup { public: __aicore__ inline MoeGatingTopKHashWithoutGroup(){}; - __aicore__ inline void Init(GM_ADDR x, GM_ADDR bias, GM_ADDR inputIds, GM_ADDR tid2eid, GM_ADDR y, GM_ADDR expertIdx, GM_ADDR out, GM_ADDR workspace, + __aicore__ inline void Init(GM_ADDR x, GM_ADDR bias, GM_ADDR inputIds, GM_ADDR tid2eid, GM_ADDR biasVl, GM_ADDR y, GM_ADDR expertIdx, GM_ADDR out, GM_ADDR workspace, const MoeGatingTopKHashTilingData *tilingData, TPipe *tPipe); __aicore__ inline void Process(); private: __aicore__ inline void CopyInBiasAndInitExpertId(); __aicore__ inline void CopyInX(int64_t progress); - __aicore__ inline void ComputeX(); + __aicore__ inline void ComputeX(bool useVisionBias); __aicore__ inline void CopuOutXNorm(int64_t row); __aicore__ inline void SelectTopKExpertIdx(); __aicore__ inline void SelectExpertIdxByHash(int64_t row); + __aicore__ inline bool IsImageRow(int64_t row); __aicore__ inline void SelectTopKExpertScore(); __aicore__ inline void CopyOut(int64_t row); @@ -46,6 +47,7 @@ class MoeGatingTopKHashWithoutGroup { TQue outOutQueue_; TBuf biasBuf_; // 存放输入bias + TBuf biasVlBuf_; // vision-token correction bias TBuf expertIdBuf_; // 专家编号 TBuf xNormWithBiasBuf_; // 存放加了bias之后的值 TBuf xNormBuf_; // 存放计算sigmoid或softmax的值 @@ -54,6 +56,7 @@ class MoeGatingTopKHashWithoutGroup { GlobalTensor xGm_; GlobalTensor biasGm_; + GlobalTensor biasVlGm_; GlobalTensor inputIdsGm_; GlobalTensor tid2eidGm_; GlobalTensor yGm_; @@ -65,6 +68,7 @@ class MoeGatingTopKHashWithoutGroup { int64_t curCoreRowCount_ = 0; int64_t expertCount_ = 0; bool addBias_ = false; + bool addBiasVl_ = false; bool outFlag_ = false; bool hashFlag_ = false; int64_t k_ = 0; @@ -104,6 +108,20 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::CopyInBiasAndIn PipeBarrier(); } } + if (addBiasVl_) { + LocalTensor biasVlTensor = biasVlBuf_.Get(); + if constexpr (IsSameType::value) { + DataCopyPad(biasVlTensor, biasVlGm_, dataCopyParams, dataCopyPadParams); + SetWaitFlag(HardEvent::MTE2_V); + } else { + DataCopyPad(biasVlTensor[expertCountAlign_].ReinterpretCast(), biasVlGm_, dataCopyParams, + dataCopyPadParams); + SetWaitFlag(HardEvent::MTE2_V); + Cast(biasVlTensor, biasVlTensor[expertCountAlign_].ReinterpretCast(), RoundMode::CAST_NONE, + expertCountAlign_); + PipeBarrier(); + } + } ArithProgression(expertIdTensor, static_cast(0), static_cast(1), expertCount_); } @@ -123,12 +141,11 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::CopyInX(int64_t } template -__aicore__ inline void MoeGatingTopKHashWithoutGroup::ComputeX() +__aicore__ inline void MoeGatingTopKHashWithoutGroup::ComputeX(bool useVisionBias) { LocalTensor xNormTensor = xNormBuf_.Get(); LocalTensor xInLocalTensor = xInQueue_.DeQue(); LocalTensor xNormWithBiasTensor = xNormWithBiasBuf_.Get(); - LocalTensor biasTensor = biasBuf_.Get(); if constexpr (!IsSameType::value) { Cast(xInLocalTensor, xInLocalTensor[expertCountAlign_].ReinterpretCast(), RoundMode::CAST_NONE, @@ -176,8 +193,10 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::ComputeX() Sqrt(xNormTensor, calcNormTmpTensor, expertCount_); PipeBarrier(); } - if (addBias_) { - Add(xNormWithBiasTensor, xNormTensor, biasTensor, expertCount_); + if (useVisionBias && addBiasVl_) { + Add(xNormWithBiasTensor, xNormTensor, biasVlBuf_.Get(), expertCount_); + } else if (addBias_) { + Add(xNormWithBiasTensor, xNormTensor, biasBuf_.Get(), expertCount_); } else { DataCopy(xNormWithBiasTensor, xNormTensor, expertCountAlign_); } @@ -195,6 +214,17 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::ComputeX() xInQueue_.FreeTensor(xInLocalTensor); } +template +__aicore__ inline bool MoeGatingTopKHashWithoutGroup::IsImageRow(int64_t row) +{ + if (!addBiasVl_) { + return false; + } + U1 tokenId = inputIdsGm_.GetValue(row); + return tokenId >= static_cast(tilingData_->imageSentinelLo) && + tokenId < static_cast(tilingData_->imageSentinelLo + tilingData_->imageSentinelCount); +} + template __aicore__ inline void MoeGatingTopKHashWithoutGroup::CopuOutXNorm(int64_t row) { @@ -313,7 +343,7 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::SelectExpertIdx } template -__aicore__ inline void MoeGatingTopKHashWithoutGroup::Init(GM_ADDR x, GM_ADDR bias, GM_ADDR inputIds, GM_ADDR tid2eid, +__aicore__ inline void MoeGatingTopKHashWithoutGroup::Init(GM_ADDR x, GM_ADDR bias, GM_ADDR inputIds, GM_ADDR tid2eid, GM_ADDR biasVl, GM_ADDR y, GM_ADDR expertIdx, GM_ADDR out, GM_ADDR workspace, const MoeGatingTopKHashTilingData *tilingData, TPipe *tPipe) { @@ -328,6 +358,7 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::Init(GM_ADDR x, } expertCount_ = tilingData_->expertCount; addBias_ = tilingData_->addBias == 1; + addBiasVl_ = tilingData_->addBiasVl == 1; outFlag_ = tilingData_->outFlag == 1; hashFlag_ = tilingData_->hashFlag == 1; k_ = tilingData_->k; @@ -337,6 +368,7 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::Init(GM_ADDR x, // init input gm buf xGm_.SetGlobalBuffer((__gm__ T *)x + perCoreRowCount_ * expertCount_ * blockIdx_, expertCount_); biasGm_.SetGlobalBuffer((__gm__ T *)bias, expertCount_); + biasVlGm_.SetGlobalBuffer((__gm__ T *)biasVl, expertCount_); inputIdsGm_.SetGlobalBuffer((__gm__ U1 *)inputIds); tid2eidGm_.SetGlobalBuffer((__gm__ U2 *)tid2eid); @@ -353,6 +385,7 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::Init(GM_ADDR x, // init calc buf pipe_->InitBuffer(biasBuf_, expertCountAlign_ * sizeof(float) * (sizeof(float) / sizeof(T))); + pipe_->InitBuffer(biasVlBuf_, expertCountAlign_ * sizeof(float) * (sizeof(float) / sizeof(T))); pipe_->InitBuffer(expertIdBuf_, expertCountAlign_ * sizeof(int32_t)); pipe_->InitBuffer(xNormBuf_, expertCountAlign_ * sizeof(float)); pipe_->InitBuffer(xNormWithBiasBuf_, expertCountAlign_ * sizeof(float)); @@ -367,13 +400,15 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::Process() { CopyInBiasAndInitExpertId(); for (int64_t row = 0; row < curCoreRowCount_; row++) { + int64_t globalRow = row + perCoreRowCount_ * blockIdx_; + bool useVisionBias = IsImageRow(globalRow); CopyInX(row); - ComputeX(); + ComputeX(useVisionBias); if (outFlag_) { CopuOutXNorm(row); } - if (hashFlag_) { - SelectExpertIdxByHash(row + perCoreRowCount_ * blockIdx_); + if (hashFlag_ && !useVisionBias) { + SelectExpertIdxByHash(globalRow); } else { SelectTopKExpertIdx(); } @@ -382,4 +417,4 @@ __aicore__ inline void MoeGatingTopKHashWithoutGroup::Process() } } } // namespace MoeGatingTopKHash -#endif // MOE_GATING_TOP_K_E_K_WITHOUT_GROUP_H \ No newline at end of file +#endif // MOE_GATING_TOP_K_E_K_WITHOUT_GROUP_H diff --git a/csrc/torch_binding.cpp b/csrc/torch_binding.cpp index c1d5ec88ceda..7610eb092048 100644 --- a/csrc/torch_binding.cpp +++ b/csrc/torch_binding.cpp @@ -41,6 +41,8 @@ #include "attention/lightning_indexer/lightning_indexer_torch_adpt.h" #include "moe/moe_gating_top_k/moe_gating_top_k_torch_adpt.h" #include "attention/sparse_flash_attention/sparse_flash_attention_torch_adpt.h" +#include "attention/sparse_flash_mla/sparse_flash_mla_torch_adpt.h" +#include "attention/quant_lightning_indexer_v2/quant_lightning_indexer_v2_torch_adpt.h" #include "attention/kv_quant_sparse_flash_attention/kv_quant_sparse_flash_attention_torch_adpt.h" #include "attention/fused_sparse_attention_overlap/fused_sparse_attention_overlap_torch_adpt.h" #include "attention/lightning_indexer_quant/lightning_indexer_quant_torch_adpt.h" @@ -690,7 +692,10 @@ std::tuple moe_gating_top_k_hash( int64_t group_select_mode, int64_t renorm, int64_t norm_type, - bool out_flag) + bool out_flag, + const c10::optional& bias_vl_opt, + int64_t image_sentinel_lo, + int64_t image_sentinel_count) { TORCH_CHECK(x.dim() == 2, "x must be 2D, but got dim=", x.dim()); @@ -748,9 +753,25 @@ std::tuple moe_gating_top_k_hash( TORCH_CHECK(tid2eid.dim() >= 1, "tid2eid must have dim>=1, but got dim=", tid2eid.dim()); } + if (bias_vl_opt.has_value() && bias_vl_opt->defined()) { + const auto& bias_vl = *bias_vl_opt; + TORCH_CHECK(input_ids_opt.has_value() && input_ids_opt->defined(), + "input_ids is required when bias_vl is present"); + TORCH_CHECK(bias_vl.dim() == 1, "bias_vl must be 1D, but got dim=", bias_vl.dim()); + TORCH_CHECK(bias_vl.size(0) == expert_num, + "bias_vl.size(0) must equal expert_num. bias_vl.size(0)=", + bias_vl.size(0), ", expert_num=", expert_num); + TORCH_CHECK(bias_vl.scalar_type() == x.scalar_type(), + "bias_vl dtype must equal x dtype. x=", x.scalar_type(), + ", bias_vl=", bias_vl.scalar_type()); + TORCH_CHECK(image_sentinel_count > 0, + "image_sentinel_count must be > 0, but got ", image_sentinel_count); + } + const at::Tensor& bias = c10::value_or_else(bias_opt, [] { return at::Tensor(); }); const at::Tensor& input_ids = c10::value_or_else(input_ids_opt, [] { return at::Tensor(); }); const at::Tensor& tid2eid = c10::value_or_else(tid2eid_opt, [] { return at::Tensor(); }); + const at::Tensor& bias_vl = c10::value_or_else(bias_vl_opt, [] { return at::Tensor(); }); at::Tensor y = at::empty({rows, k}, x.options()); at::Tensor expert_idx = at::empty({rows, k}, x.options().dtype(at::kInt)); @@ -761,15 +782,18 @@ std::tuple moe_gating_top_k_hash( bias, input_ids, tid2eid, + bias_vl, k, k_group, group_count, - routed_scaling_factor, - eps, group_select_mode, renorm, norm_type, out_flag, + routed_scaling_factor, + eps, + image_sentinel_lo, + image_sentinel_count, y, expert_idx, out); @@ -1136,6 +1160,30 @@ std::tuple npu_quant_lightning_indexer_v2_npu( return std::tuple(sparse_indices_out, sparse_values_out); } +std::tuple npu_quant_lightning_indexer_v2_compat_npu( + const at::Tensor &query, const at::Tensor &key, const at::Tensor &weights, + const at::Tensor &query_dequant_scale, const at::Tensor &key_dequant_scale, + int64_t topk, int64_t quant_mode, + const c10::optional &cu_seqlens_q, + const c10::optional &cu_seqlens_k, + const c10::optional &seqused_q, + const c10::optional &seqused_k, + const c10::optional &cmp_residual_k, + const c10::optional &block_table, + const c10::optional &output_idx_offset, + const c10::optional &metadata, + int64_t max_seqlen_q, c10::string_view layout_q, c10::string_view layout_k, + int64_t mask_mode, int64_t cmp_ratio, int64_t return_value) +{ + TORCH_CHECK(return_value == 0, "npu_quant_lightning_indexer_v2 only supports return_value=0"); + auto outputs = qli_v2::QuantLightningIndexerCandidate( + query, key, weights, query_dequant_scale, key_dequant_scale, topk, quant_mode, + c10::nullopt, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k, + block_table, output_idx_offset, metadata, max_seqlen_q, layout_q, layout_k, + mask_mode, cmp_ratio, 3, 2048, 8); + return {std::get<0>(outputs), std::get<1>(outputs)}; +} + std::tuple construct_output_tensor(const at::Tensor &q, std::string layout, bool return_softmax_lse) { @@ -1391,8 +1439,6 @@ at::Tensor npu_hc_post_npu( } constexpr int64_t HC_PRE_HC_LIMIT = 4; -constexpr int64_t HC_PRE_D_LIMIT = 4096; -constexpr int64_t HC_PRE_D_LIMIT_EXTEND = 7168; constexpr int64_t HC_PRE_MIX_HC_LIMIT = 24; std::tuple construct_hc_pre_output_tensor(const at::Tensor& x, int64_t hc_mult) @@ -1423,11 +1469,23 @@ std::tuple construct_hc_pre_output_tensor(co return std::tuple(y, post, comb_frag); } +at::Tensor construct_hc_pre_pre_output_tensor(const at::Tensor& x, int64_t hc_mult) +{ + at::SmallVector pre_size; + if (x.dim() == 4) { + pre_size = {x.size(0), x.size(1), hc_mult}; + } else if (x.dim() == 3) { + pre_size = {x.size(0), hc_mult}; + } + return at::empty(pre_size, x.options().dtype(at::kFloat)); +} + void check_hc_pre_shape_and_dtype( const at::Tensor& x, const at::Tensor& hc_fn, const at::Tensor& hc_scale, const at::Tensor& hc_base, + const c10::optional& pre_mix, int64_t hc_mult) { constexpr int64_t HC_SCALE_SIZE = 3; @@ -1442,8 +1500,6 @@ void check_hc_pre_shape_and_dtype( auto d = x_dims == 4 ? x.size(3) : x.size(2); TORCH_CHECK(hc_mult == HC_PRE_HC_LIMIT, "hc_mult only supports ", HC_PRE_HC_LIMIT, ", actual ", hc_mult, "."); TORCH_CHECK(hc == HC_PRE_HC_LIMIT, "The hc of x only supports ", HC_PRE_HC_LIMIT, ", actual ", hc, "."); - TORCH_CHECK(d == HC_PRE_D_LIMIT || d == HC_PRE_D_LIMIT_EXTEND, "The d of x only supports ", HC_PRE_D_LIMIT, - " or ", HC_PRE_D_LIMIT_EXTEND, ", actual ", d, "."); TORCH_CHECK(hc_fn.dim() == 2, "Input tensor hc_fn's dim num should be 2, actual ", hc_fn.dim(), "."); TORCH_CHECK(hc_fn.size(0) == HC_PRE_MIX_HC_LIMIT, "The hc_fn.shape[0] only supports ", HC_PRE_MIX_HC_LIMIT, ", actual ", hc_fn.size(0), "."); @@ -1460,28 +1516,51 @@ void check_hc_pre_shape_and_dtype( TORCH_CHECK(hc_fn.dtype() == at::kFloat, "hc_fn's dtype should be FLOAT32."); TORCH_CHECK(hc_scale.dtype() == at::kFloat, "hc_scale's dtype should be FLOAT32."); TORCH_CHECK(hc_base.dtype() == at::kFloat, "hc_base's dtype should be FLOAT32."); + if (pre_mix.has_value() && pre_mix->defined()) { + TORCH_CHECK(pre_mix->dtype() == at::kFloat, "pre_mix's dtype should be FLOAT32."); + TORCH_CHECK(pre_mix->dim() == x_dims - 1, "pre_mix's dim num should be ", x_dims - 1, ", actual ", + pre_mix->dim(), "."); + for (auto i = 0; i < x_dims - 1; i++) { + TORCH_CHECK(pre_mix->size(i) == x.size(i), "pre_mix.shape[", i, "] should equal x.shape[", i, + "], actual ", pre_mix->size(i), " vs ", x.size(i), "."); + } + } } -std::tuple run_hc_pre_fusion( +std::tuple run_hc_pre_fusion( const at::Tensor& x, const at::Tensor& hc_fn, const at::Tensor& hc_scale, const at::Tensor& hc_base, - int64_t hc_mult, int64_t hc_sinkhorn_iters, double norm_eps, double hc_eps) + const c10::optional& pre_mix, int64_t hc_mult, int64_t hc_sinkhorn_iters, double norm_eps, + double hc_eps) { auto output_tensors = construct_hc_pre_output_tensor(x, hc_mult); at::Tensor y = std::get<0>(output_tensors); at::Tensor post = std::get<1>(output_tensors); at::Tensor comb_frag = std::get<2>(output_tensors); - EXEC_NPU_CMD(aclnnHcPre, x, hc_fn, hc_scale, hc_base, hc_mult, hc_sinkhorn_iters, hc_eps, norm_eps, - y, post, comb_frag); + at::Tensor pre = construct_hc_pre_pre_output_tensor(x, hc_mult); + EXEC_NPU_CMD(aclnnHcPre, x, hc_fn, hc_scale, hc_base, pre_mix, hc_mult, hc_sinkhorn_iters, hc_eps, norm_eps, + y, post, comb_frag, pre); - return std::tuple(y, post, comb_frag); + return std::tuple(y, post, comb_frag, pre); } std::tuple npu_hc_pre_v2_npu( const at::Tensor& x, const at::Tensor& hc_fn, const at::Tensor& hc_scale, const at::Tensor& hc_base, int64_t hc_mult, int64_t hc_sinkhorn_iters, double norm_eps, double hc_eps) { - check_hc_pre_shape_and_dtype(x, hc_fn, hc_scale, hc_base, hc_mult); - return run_hc_pre_fusion(x, hc_fn, hc_scale, hc_base, hc_mult, hc_sinkhorn_iters, norm_eps, hc_eps); + const c10::optional pre_mix = c10::nullopt; + check_hc_pre_shape_and_dtype(x, hc_fn, hc_scale, hc_base, pre_mix, hc_mult); + auto outputs = run_hc_pre_fusion(x, hc_fn, hc_scale, hc_base, pre_mix, hc_mult, hc_sinkhorn_iters, norm_eps, + hc_eps); + return {std::get<0>(outputs), std::get<1>(outputs), std::get<2>(outputs)}; +} + +std::tuple npu_hc_pre_v3_npu( + const at::Tensor& x, const at::Tensor& hc_fn, const at::Tensor& hc_scale, const at::Tensor& hc_base, + const c10::optional& pre_mix, int64_t hc_mult, int64_t hc_sinkhorn_iters, double norm_eps, + double hc_eps) +{ + check_hc_pre_shape_and_dtype(x, hc_fn, hc_scale, hc_base, pre_mix, hc_mult); + return run_hc_pre_fusion(x, hc_fn, hc_scale, hc_base, pre_mix, hc_mult, hc_sinkhorn_iters, norm_eps, hc_eps); } void inplace_partial_rotary_mul_npu(at::Tensor & x, const at::Tensor &r1, const at::Tensor &r2, c10::string_view rotary_mode, at::IntArrayRef partial_slice, bool negate_sin) @@ -3020,6 +3099,58 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) ); ops.impl("npu_sparse_flash_attention", torch::kPrivateUse1, &vllm_ascend::npu_sparse_flash_attention); + ops.def( + "npu_quant_lightning_indexer_v2_metadata(int num_heads_q, int num_heads_k, int head_dim, int topk, " + "int quant_mode, *, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None, Tensor? seqused_q=None, " + "Tensor? seqused_k=None, Tensor? cmp_residual_k=None, int batch_size=0, int max_seqlen_q=0, " + "int max_seqlen_k=0, str layout_q='TND', str layout_k='PA_BBND', int mask_mode=3, int cmp_ratio=1, " + "str device='npu') -> Tensor" + ); + ops.impl("npu_quant_lightning_indexer_v2_metadata", torch::kPrivateUse1, + &vllm_ascend::npu_quant_lightning_indexer_v2_metadata_npu); + ops.def( + "npu_quant_lightning_indexer_v3(Tensor query, Tensor key, Tensor weights, " + "Tensor query_dequant_scale, Tensor key_dequant_scale, int topk, int quant_mode, *, " + "Tensor? candidate_topk_index=None, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None, " + "Tensor? seqused_q=None, Tensor? seqused_k=None, Tensor? cmp_residual_k=None, " + "Tensor? block_table=None, Tensor? output_idx_offset=None, Tensor? metadata=None, " + "int max_seqlen_q=-1, str layout_q='TND', str layout_k='PA_BBND', int mask_mode=3, " + "int cmp_ratio=1, int candidate_mode=3, int candidate_topk_blocks=2048, " + "int candidate_block_size=8) -> (Tensor, Tensor, Tensor)" + ); + ops.impl("npu_quant_lightning_indexer_v3", torch::kPrivateUse1, + &vllm_ascend::qli_v2::QuantLightningIndexerCandidate); + ops.impl("npu_quant_lightning_indexer_v3", torch::kMeta, + &vllm_ascend::qli_v2::QuantLightningIndexerCandidate); + + ops.def( + "npu_sparse_flash_mla_metadata(int num_heads_q, int num_heads_kv, int head_dim, *, " + "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_ori_kv=None, Tensor? cu_seqlens_cmp_kv=None, " + "Tensor? seqused_q=None, Tensor? seqused_ori_kv=None, Tensor? seqused_cmp_kv=None, " + "Tensor? cmp_residual_kv=None, Tensor? ori_topk_length=None, Tensor? cmp_topk_length=None, " + "int batch_size=0, int max_seqlen_q=0, int max_seqlen_ori_kv=0, int max_seqlen_cmp_kv=0, " + "int ori_topk=0, int cmp_topk=0, int cmp_ratio=0, int ori_mask_mode=0, int cmp_mask_mode=0, " + "int ori_win_left=-1, int ori_win_right=-1, str layout_q='BSND', str layout_kv='BSND', " + "bool has_ori_kv=True, bool has_cmp_kv=True) -> Tensor" + ); + ops.impl("npu_sparse_flash_mla_metadata", torch::kPrivateUse1, + &vllm_ascend::npu_sparse_flash_mla_metadata); + + ops.def( + "npu_sparse_flash_mla(Tensor q, *, Tensor? ori_kv=None, Tensor? cmp_kv=None, " + "Tensor? ori_sparse_indices=None, Tensor? cmp_sparse_indices=None, " + "Tensor? ori_block_table=None, Tensor? cmp_block_table=None, " + "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_ori_kv=None, Tensor? cu_seqlens_cmp_kv=None, " + "Tensor? seqused_q=None, Tensor? seqused_ori_kv=None, Tensor? seqused_cmp_kv=None, " + "Tensor? cmp_residual_kv=None, Tensor? ori_topk_length=None, Tensor? cmp_topk_length=None, " + "Tensor? sinks=None, Tensor? metadata=None, float softmax_scale=1.0, int cmp_ratio=0, " + "int ori_mask_mode=0, int cmp_mask_mode=0, int ori_win_left=-1, int ori_win_right=-1, " + "str layout_q='BSND', str layout_kv='BSND', int topk_value_mode=1, " + "bool return_softmax_lse=False) -> (Tensor, Tensor)" + ); + ops.impl("npu_sparse_flash_mla", torch::kPrivateUse1, + &vllm_ascend::npu_sparse_flash_mla); + ops.def( "npu_kv_quant_sparse_flash_attention(Tensor query, Tensor key, Tensor value," " Tensor sparse_indices, float scale_value, *," @@ -3131,7 +3262,10 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) "int group_select_mode=0, " "int renorm=0, " "int norm_type=0, " - "bool out_flag=False" + "bool out_flag=False, " + "Tensor? bias_vl=None, " + "int image_sentinel_lo=129257, " + "int image_sentinel_count=5" ") -> (Tensor y, Tensor expert_idx, Tensor out)" ); ops.impl("moe_gating_top_k_hash", torch::kPrivateUse1,&vllm_ascend::moe_gating_top_k_hash); @@ -3205,7 +3339,8 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) "int mask_mode=3, int cmp_ratio=4, int return_value=0" ") -> (Tensor sparse_indices, Tensor sparse_values)" ); - ops.impl("npu_quant_lightning_indexer_v2", torch::kPrivateUse1, &vllm_ascend::npu_quant_lightning_indexer_v2_npu); + ops.impl("npu_quant_lightning_indexer_v2", torch::kPrivateUse1, + &vllm_ascend::npu_quant_lightning_indexer_v2_compat_npu); ops.def( "npu_sparse_attn_sharedkv(" @@ -3289,30 +3424,6 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) ); ops.impl("npu_vllm_quant_lightning_indexer_metadata", torch::kPrivateUse1, &vllm_ascend::npu_vllm_quant_lightning_indexer_metadata_npu); - ops.def( - "npu_quant_lightning_indexer_v2_metadata(" - "int num_heads_q, " - "int num_heads_k, " - "int head_dim, " - "int topk, " - "int quant_mode, *, " - "Tensor? cu_seqlens_q=None, " - "Tensor? cu_seqlens_k=None, " - "Tensor? seqused_q=None, " - "Tensor? seqused_k=None, " - "Tensor? cmp_residual_k=None, " - "int batch_size=0, " - "int max_seqlen_q=-1, " - "int max_seqlen_k=-1, " - "str layout_q=\"TND\", " - "str layout_k=\"PA_BBND\", " - "int mask_mode=3, " - "int cmp_ratio=4, " - "str device=\"npu\"" - ") -> (Tensor metadata)" - ); - ops.impl("npu_quant_lightning_indexer_v2_metadata", torch::kPrivateUse1, &vllm_ascend::npu_quant_lightning_indexer_v2_metadata_npu); - ops.def( "npu_hc_post(" "Tensor x, " @@ -3328,10 +3439,19 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) "Tensor x, Tensor hc_fn, Tensor hc_scale, Tensor hc_base, " "int hc_mult, int hc_sinkhorn_iters, " "float norm_eps, float hc_eps" - ") -> (Tensor out0, Tensor out1, Tensor out2)" + ") -> (Tensor y, Tensor post, Tensor comb_frag)" ); ops.impl("npu_hc_pre_v2", torch::kPrivateUse1, &vllm_ascend::npu_hc_pre_v2_npu); + ops.def( + "npu_hc_pre_v3(" + "Tensor x, Tensor hc_fn, Tensor hc_scale, Tensor hc_base, Tensor? pre_mix=None, *, " + "int hc_mult=4, int hc_sinkhorn_iters=20, " + "float norm_eps=1e-6, float hc_eps=1e-6" + ") -> (Tensor y, Tensor post, Tensor comb_frag, Tensor pre)" + ); + ops.impl("npu_hc_pre_v3", torch::kPrivateUse1, &vllm_ascend::npu_hc_pre_v3_npu); + ops.def( "inplace_partial_rotary_mul(" "Tensor(a!) x, Tensor r1, Tensor r2, str rotary_mode, int[] partial_slice, bool negate_sin=False" diff --git a/csrc/torch_binding_meta.cpp b/csrc/torch_binding_meta.cpp index 5c65a4facdae..aa7a269413eb 100644 --- a/csrc/torch_binding_meta.cpp +++ b/csrc/torch_binding_meta.cpp @@ -306,6 +306,78 @@ std::tuple npu_sparse_flash_attention_meta( return std::tuple(output, softmax_max, softmax_sum); } +at::Tensor npu_sparse_flash_mla_metadata_meta( + int64_t num_heads_q, int64_t num_heads_kv, int64_t head_dim, + const c10::optional &cu_seqlens_q, + const c10::optional &cu_seqlens_ori_kv, + const c10::optional &cu_seqlens_cmp_kv, + const c10::optional &seqused_q, + const c10::optional &seqused_ori_kv, + const c10::optional &seqused_cmp_kv, + const c10::optional &cmp_residual_kv, + const c10::optional &ori_topk_length, + const c10::optional &cmp_topk_length, + int64_t batch_size, int64_t max_seqlen_q, + int64_t max_seqlen_ori_kv, int64_t max_seqlen_cmp_kv, + int64_t ori_topk, int64_t cmp_topk, int64_t cmp_ratio, + int64_t ori_mask_mode, int64_t cmp_mask_mode, + int64_t ori_win_left, int64_t ori_win_right, + c10::string_view layout_q, c10::string_view layout_kv, + bool has_ori_kv, bool has_cmp_kv) +{ + return at::empty_symint(c10::SymDimVector{c10::SymInt(1024)}, + at::TensorOptions().dtype(at::kInt).device(c10::kMeta)); +} + +std::tuple npu_sparse_flash_mla_meta( + const at::Tensor &q, + const c10::optional &ori_kv, + const c10::optional &cmp_kv, + const c10::optional &ori_sparse_indices, + const c10::optional &cmp_sparse_indices, + const c10::optional &ori_block_table, + const c10::optional &cmp_block_table, + const c10::optional &cu_seqlens_q, + const c10::optional &cu_seqlens_ori_kv, + const c10::optional &cu_seqlens_cmp_kv, + const c10::optional &seqused_q, + const c10::optional &seqused_ori_kv, + const c10::optional &seqused_cmp_kv, + const c10::optional &cmp_residual_kv, + const c10::optional &ori_topk_length, + const c10::optional &cmp_topk_length, + const c10::optional &sinks, + const c10::optional &metadata, + double softmax_scale, int64_t cmp_ratio, int64_t ori_mask_mode, + int64_t cmp_mask_mode, int64_t ori_win_left, int64_t ori_win_right, + c10::string_view layout_q, c10::string_view layout_kv, + int64_t topk_value_mode, bool return_softmax_lse) +{ + auto output = at::empty_symint(q.sym_sizes(), q.options().device(c10::kMeta)); + if (!return_softmax_lse) { + auto lse = at::empty_symint(c10::SymDimVector{c10::SymInt(0)}, + q.options().dtype(at::kFloat).device(c10::kMeta)); + return {output, lse}; + } + + TORCH_CHECK(ori_kv.has_value() || cmp_kv.has_value(), + "ori_kv or cmp_kv is required when return_softmax_lse is true"); + const auto &kv = ori_kv.has_value() ? ori_kv.value() : cmp_kv.value(); + const auto layout_q_str = std::string(layout_q); + const auto layout_kv_str = std::string(layout_kv); + const auto kv_heads = layout_kv_str == "TND" ? kv.sym_size(1) : kv.sym_size(2); + c10::SymDimVector lse_shape; + if (layout_q_str == "BSND") { + lse_shape = {q.sym_size(0), kv_heads, q.sym_size(1), q.sym_size(2) / kv_heads}; + } else { + TORCH_CHECK(layout_q_str == "TND", "layout_q must be BSND or TND"); + lse_shape = {kv_heads, q.sym_size(0), q.sym_size(1) / kv_heads}; + } + auto lse = at::empty_symint(lse_shape, + q.options().dtype(at::kFloat).device(c10::kMeta)); + return {output, lse}; +} + at::Tensor npu_sparse_attention_score_meta( const at::Tensor &query, const at::Tensor &key, const at::Tensor &value, const at::Tensor &select_idx, const at::Tensor &block_table, @@ -737,7 +809,10 @@ std::tuple moe_gating_top_k_hash_meta( int64_t group_select_mode, int64_t renorm, int64_t norm_type, - bool out_flag) + bool out_flag, + const c10::optional& bias_vl_opt, + int64_t image_sentinel_lo, + int64_t image_sentinel_count) { TORCH_CHECK(x.dim() == 2, "x must be 2D, but got dim=", x.dim()); TORCH_CHECK( @@ -770,6 +845,18 @@ std::tuple moe_gating_top_k_hash_meta( ", bias=", bias.scalar_type()); } + if (bias_vl_opt.has_value() && bias_vl_opt->defined()) { + const auto& bias_vl = *bias_vl_opt; + TORCH_CHECK(input_ids_opt.has_value() && input_ids_opt->defined(), + "input_ids is required when bias_vl is present"); + TORCH_CHECK(bias_vl.dim() == 1, "bias_vl must be 1D, but got dim=", bias_vl.dim()); + TORCH_CHECK(bias_vl.scalar_type() == x.scalar_type(), + "bias_vl dtype must equal x dtype. x=", x.scalar_type(), + ", bias_vl=", bias_vl.scalar_type()); + TORCH_CHECK(image_sentinel_count > 0, + "image_sentinel_count must be > 0, but got ", image_sentinel_count); + } + if (input_ids_opt.has_value() && input_ids_opt->defined()) { const auto& input_ids = *input_ids_opt; TORCH_CHECK(input_ids.scalar_type() == at::kInt || input_ids.scalar_type() == at::kLong, @@ -1168,16 +1255,37 @@ std::tuple construct_hc_pre_output_tensor(co return std::tuple(y, post, comb_frag); } +at::Tensor construct_hc_pre_pre_output_tensor(const at::Tensor& x, int64_t hc_mult) +{ + at::SmallVector pre_size; + if (x.dim() == 4) { + pre_size = {x.sym_size(0), x.sym_size(1), hc_mult}; + } else if (x.dim() == 3) { + pre_size = {x.sym_size(0), hc_mult}; + } + return at::empty_symint(c10::SymIntArrayRef(pre_size), x.options().dtype(at::kFloat)); +} + std::tuple npu_hc_pre_meta( const at::Tensor& x, const at::Tensor& hc_fn, const at::Tensor& hc_scale, const at::Tensor& hc_base, int64_t hc_mult, int64_t hc_sinkhorn_iters, double norm_eps, double hc_eps) +{ + auto output_tensors = construct_hc_pre_output_tensor(x, hc_mult); + return {std::get<0>(output_tensors), std::get<1>(output_tensors), std::get<2>(output_tensors)}; +} + +std::tuple npu_hc_pre_v3_meta( + const at::Tensor& x, const at::Tensor& hc_fn, const at::Tensor& hc_scale, const at::Tensor& hc_base, + const c10::optional& pre_mix, int64_t hc_mult, int64_t hc_sinkhorn_iters, double norm_eps, + double hc_eps) { auto output_tensors = construct_hc_pre_output_tensor(x, hc_mult); at::Tensor y = std::get<0>(output_tensors); at::Tensor post = std::get<1>(output_tensors); at::Tensor comb_frag = std::get<2>(output_tensors); + at::Tensor pre = construct_hc_pre_pre_output_tensor(x, hc_mult); - return std::tuple(y, post, comb_frag); + return std::tuple(y, post, comb_frag, pre); } void inplace_partial_rotary_mul_meta( @@ -2078,6 +2186,8 @@ TORCH_LIBRARY_IMPL_EXPAND(CONCAT(_C, _ascend), Meta, ops) { ops.impl("npu_lightning_indexer", &vllm_ascend::meta::npu_lightning_indexer_meta); // Sparse flash attention ops.impl("npu_sparse_flash_attention", &vllm_ascend::meta::npu_sparse_flash_attention_meta); + ops.impl("npu_sparse_flash_mla_metadata", &vllm_ascend::meta::npu_sparse_flash_mla_metadata_meta); + ops.impl("npu_sparse_flash_mla", &vllm_ascend::meta::npu_sparse_flash_mla_meta); ops.impl("npu_sparse_attention_score", &vllm_ascend::meta::npu_sparse_attention_score_meta); ops.impl("npu_k2q_csr", &vllm_ascend::meta::npu_k2q_csr_meta); ops.impl("npu_sparse_attention_score_prefill", @@ -2113,6 +2223,7 @@ TORCH_LIBRARY_IMPL_EXPAND(CONCAT(_C, _ascend), Meta, ops) { ops.impl("npu_sparse_attn_sharedkv_metadata", &vllm_ascend::meta::npu_sparse_attn_sharedkv_metadata_meta); ops.impl("npu_hc_post", &vllm_ascend::meta::npu_hc_post_meta); ops.impl("npu_hc_pre_v2", &vllm_ascend::meta::npu_hc_pre_meta); + ops.impl("npu_hc_pre_v3", &vllm_ascend::meta::npu_hc_pre_v3_meta); ops.impl("inplace_partial_rotary_mul", &vllm_ascend::meta::inplace_partial_rotary_mul_meta); ops.impl("npu_rms_norm_dynamic_quant", &vllm_ascend::meta::npu_rms_norm_dynamic_quant_meta); ops.impl("kv_compress_epilog", &vllm_ascend::meta::kv_compress_epilog_meta); diff --git a/docs/source/developer_guide/aurora_qli_candidate.md b/docs/source/developer_guide/aurora_qli_candidate.md new file mode 100644 index 000000000000..0bc637f05c66 --- /dev/null +++ b/docs/source/developer_guide/aurora_qli_candidate.md @@ -0,0 +1,155 @@ +# Aurora QLI V2 and candidate integration + +Aurora's indexer uses `QuantLightningIndexerV2` with paged INT8 index K. +The imported operator source is from `ops-transformer-qli_candidate.zip`, +SHA256 `771f0c16b9119c676c10cebef713b168f127ff194f3f2f98edcbcd0979c6f966`. +Its companion `QuantLightningIndexerV2Metadata` is built and registered too. + +## Data flow + +`DeepseekV41Indexer.select` projects Q and head weights and applies RoPE. +`select_projected` quantizes each Q head to INT8 with a representable FP16 +scale. Head weights and K scales are FP16. Existing source-owned K-cache +updates remain unchanged. + +The operator consumes `TND` Q and `PA_BBND` index K directly. Aurora's +four-slot layer-outermost allocator packs index K and FP16 scales after KV +inside each shared page. C2 uses 64-row views with a 131072-byte block stride; +C1 uses 128-row views with a 147712-byte block stride. The Torch adapter passes +their actual leading strides to ACLNN. +Neither layout requires gathering the whole context into a dense tensor. + +Metadata receives original query boundaries, compressed K lengths, and the +original sequence length modulo the compression ratio. The residual tensor +must be absent for ratio 1. On A2/A3, this repository compiles the model's +TND/PA_BBND layout only; BSND and nonpaged calls are rejected by host validation. +A5 retains its general QLI dtype, layout and quantization support. + +The three candidate modes share one native operator: + +| Mode | Meaning | Result | +| --- | --- | --- | +| 1 | Candidate source | Unfiltered position TopK and candidate block IDs | +| 2 | Candidate consumer | Rerank using this layer's Q and weights within source blocks | +| 3 | Candidate disabled | Ordinary position TopK | + +Candidates are INT32 block IDs of shape `[tokens, 1, candidate_topk_blocks]`, +with `-1` padding. They are neither index-K vectors nor position TopK. A block +contains 8 compressed positions. Block scores use the maximum position score; +the last visible block is pinned. Shared attention state retains this tensor +within one forward and resets it on the next forward. + +The model sorts returned position indices chronologically and moves `-1` +padding to the end before attention. Empty compressed contexts return empty +indices and, for a source, all-invalid candidates. A consumer without a source +raises an error. + +## Current A3 contract + +- INT8 Q/K, FP16 head weights and per-head Q/per-token K scales; quant mode 2. +- 32 or 64 replicated index heads, head dimension 128, one index-K head. +- Aurora compression ratios 1 and 2, causal mask mode 3. +- Position TopK in `[1, 2048]`; candidate blocks a multiple of 64 in `[64, 2048]`. +- Candidate block size is exactly 8 in this kernel implementation. +- TND query and PA_BBND key layouts only in the compiled package. +- The A3 operator returns indices and candidate IDs, not score values. + +The numerical reference follows the supplied INT8 golden: INT32 QK divided by +1024, FP16 ReLU and `weight * Q_scale`, FP32 head reduction, then K scaling. +Query quantization and FP16 weight rounding differ from the earlier floating-Q +small-operator path. Operator agreement with this quantized reference does not +establish full-model or dataset accuracy. + +## Build and regression coverage + +Build with `pip install -v -e . --no-deps --no-build-isolation` on the paired +CANN/NPU environment. Both new symbols are registered on PrivateUse1 and Meta. + +`tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_qli.py` +checks candidate generation/consumption, different consumer queries, ratio +boundaries, paged views, mixed requests, 2048 candidate blocks, 64 heads, empty +contexts and Meta shapes against an independent CPU reference. Ties at the +TopK cutoff use score validity, uniqueness and count rather than arbitrary +index ordering. + +The mixed-request model-indexer cases allocate the actual four-slot cache +configuration, including the null ID, and test source/consumer selection on +its strided index K/scale views. The SparseFlashMla suite also compares the +slot-backed BF16 views against the earlier block-outermost layout for ratios +0/1/2 and decode, prefill and mixed requests. These updated tests have not yet +been executed; remote torch/NPU verification is deferred. The end-to-end +results below describe the original layout. + +The imported host tiling needed one semantic fix: TND candidate size is +`T * N_k * blocks`, because T already includes every request. Multiplying by +batch size again incorrectly rejects mixed-request consumers. Kernel offsets +already use TND query prefixes and require no corresponding change. + +Validation includes single-chip eager and Meta contracts, plus the full +40-layer W8A8 checkpoint on one A3 server with TP4/DP4/EP16 (Engram and +DSpark disabled). Four end-to-end cases were each run twice at temperature 0, +seed 7: arithmetic, capital-city lookup, multi-turn recall, and a 20,032-token +retrieval prompt. Both runs returned the expected answers and identical output +tokens. The long prompt exceeds the 16,384-position candidate block budget. +It exercises chunked prefill and decode with actual candidate filtering. + +The checkpoint-provided `encoding.encode_messages(..., thinking_mode="chat")` +was used with `/v1/completions`; the checkpoint does not provide a standard +chat template. Repeated output agreement validates these fixed functional +cases, not a baseline-versus-candidate dataset accuracy comparison. Graph, +64K/128K requests, multi-node execution, Engram and DSpark remain unvalidated. + +## Compilation scope + +Aurora's A2/A3 path uses INT8 Q/K, FP16 weights/scales and quant mode 2. Its +compiled QLI V2 template matrix has one key, down from 4. A5 retains the full +16-key matrix, including paged BSND/TND queries and matching nonpaged BSND/TND +layouts. A5 dtype registration, host validation and kernel dispatch retain +FP8, MXFP8, HiFloat8, MXFP4 and INT8 (quant modes 1/3/4/5/2 respectively). +The Aurora INT8 call sites do not establish that other A5 paths are unused. +Both architectures retain their original template argument encodings. + +Candidate modes 1/2/3, compression ratio, TopK and sequence lengths remain +runtime parameters. The A2/A3 candidate implementation and the A5 implementation +remain in separate architecture branches. This pruning does not add A5 candidate +support or change which operators the default A5 package builds: QLI V2 and +SparseFlashMla are currently included by the A2/A3 package lists. Explicit A5 +builds of these operators use the A5 template selections. + +SparseFlashMla is likewise restricted to the model's BF16 TND/PA_BBND SWA/CSA +calls. It compiles 6 keys on A2/A3 and 12 on A5; shared keys remain on both, +while single-head CSA specialization is A2/A3-only and split-G/vectorized +addressing is A5-only. HCA, independent original-KV sparse templates, other +layouts and FP16 are excluded on both architectures. Host validation rejects +pruned contracts before kernel lookup. The full compilation matrix is in +`csrc/attention/sparse_flash_mla/docs/ratio2_a2a3.md`. +HcPre already isolates A2/A3 key 0 from A5 keys 1000/1001; both A5 paths can +be selected by runtime token counts, so neither is removed. + +The fused Compressor serves DeepSeek V4; V4.1 currently uses small operators +for compression. Compressor already selects four keys per architecture: +TH/BF16, interleaved RoPE, continuous cache, and `coff=1/2`, with FP32 RoPE. +A2/A3 selects EMPTY_X/PERF; A5 selects NORMAL/EMPTY_X. The two `coff` values +cover V4 compression ratios 128 and 4 respectively, and empty input remains +supported. A5 FULL_LOAD requires BSH, so it is already excluded. Host and kernel +sources are selected separately for arch32 (A2/A3) and arch35 (A5). + +Compressor dtype registration now matches those existing selections: one +signature per architecture, down from four on A2/A3 and two on A5. Norm weights +remain BF16 on A2/A3 and FP32 on A5. Host validation likewise rejects uncompiled +layouts, dtypes and modes. The empty-input entry uses a discarded `else` branch +so the compiler does not instantiate a computation kernel after an unconditional +return. The four keys themselves and their encodings remain unchanged. + +The CPU-only template regression can be run without importing torch or CANN: + +```bash +python3 -m unittest discover -s tests/ut/ops -p test_aurora_tiling_keys.py -v +``` + +This checks preprocessor selections, architecture isolation, model layout +calls, dtype registration and unchanged argument encodings. It also compiles +the real Compressor entry against stubs that reject unwanted template +instantiations, including any computation for EMPTY_X. It does not measure CANN build time +or execute NPU kernels. Rebuild the operator package before the numerical +regressions; retained key counts alone do not establish a wall-clock speedup. diff --git a/setup.py b/setup.py index 1d036c14b996..d9b81252ecdd 100644 --- a/setup.py +++ b/setup.py @@ -509,6 +509,7 @@ def _read_requirements(filename: str) -> list[str]: packages=find_packages(exclude=("docs", "examples", "tests*", "csrc")), package_data={ "vllm_ascend.observability": ["config/*.yaml"], + "vllm_ascend.patch.platform.patch_deepseek_v41_frontend": ["LICENSE"], }, python_requires=">=3.10", install_requires=get_requirements(), diff --git a/tests/check_aurora_ring_static.py b/tests/check_aurora_ring_static.py new file mode 100644 index 000000000000..6fd64e62ea52 --- /dev/null +++ b/tests/check_aurora_ring_static.py @@ -0,0 +1,359 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# mypy: ignore-errors +"""Static placement arithmetic only: never import torch, vllm, or their modules.""" + +from __future__ import annotations + +import ast +import dataclasses +import json +from collections import defaultdict +from pathlib import Path +from types import SimpleNamespace + +ROOT = Path(__file__).resolve().parents[1] +SOURCE = ROOT / "vllm_ascend/core/deepseek_v41.py" + + +@dataclasses.dataclass(frozen=True) +class ItemSize: + itemsize: int + + +@dataclasses.dataclass(frozen=True) +class Resource: + block_size: int + head_size: int + dtype: ItemSize + compress_ratio: int = 1 + num_kv_heads: int = 1 + scale_dim: int = 0 + scale_dtype: ItemSize = ItemSize(2) + page_size_padded: int | None = None + sliding_window: int = 128 + + @property + def storage_block_size(self): + return self.block_size // self.compress_ratio + + @property + def page_size_bytes(self): + return self.page_size_padded or sum(planes(self)) + + +class Full(Resource): + pass + + +class Index(Resource): + pass + + +class State(Resource): + pass + + +class SWA(Resource): + pass + + +class Draft(Resource): + pass + + +@dataclasses.dataclass(frozen=True) +class Uniform: + kv_cache_specs: dict + + @classmethod + def from_specs(cls, specs): + return cls(specs) + + +symbols = { + "dataclass": dataclasses.dataclass, + "DeepseekV41FullSpec": Full, + "DeepseekV41IndexerSpec": Index, + "DeepseekV41CompressorStateSpec": State, + "DeepseekV41SWASpec": SWA, + "DeepseekV41DraftSWASpec": Draft, + "is_v41_spec": lambda s: isinstance(s, (Full, Index, State, SWA, Draft)), + "replace": dataclasses.replace, + "UniformTypeKVCacheSpecs": Uniform, + "KVCacheGroupSpec": lambda **kw: SimpleNamespace(**kw), + "KVCacheTensor": lambda **kw: SimpleNamespace(**kw), + "may_override_num_blocks": lambda cfg, count: cfg.override if cfg.override is not None else count, + "STATE_RING_ROWS": 32, + "CUDAGraphMode": SimpleNamespace(NONE="none", FULL="full", FULL_DECODE_ONLY="decode"), +} +names = { + "CachePlacement", + "CacheSlot", + "_layer_number", + "_draft_layer_number", + "_cache_plane_sizes", + "plan_cache_slots", + "_uniform", + "group_cache_specs", + "make_cache_groups", + "cache_slots_from_groups", + "pool_bytes_per_block", + "allocate_cache_config", + "validate_cache_runtime", +} +module = ast.parse(SOURCE.read_text()) +selected = ast.Module( + body=[node for node in module.body if isinstance(node, (ast.ClassDef, ast.FunctionDef)) and node.name in names], + type_ignores=[], +) +assert len(selected.body) == len(names) +exec(compile(selected, str(SOURCE), "exec"), symbols) +plan = symbols["plan_cache_slots"] +planes = symbols["_cache_plane_sizes"] + + +def specs_for(block, width, index_width): + specs = {} + for layer in range(40): + prefix = f"language_model.model.layers.{layer}.self_attn" + specs[prefix + ".swa_cache"] = SWA(block, width, ItemSize(2)) + if layer in (2, 8, 14, 20): + ratio = 1 if layer == 20 else 2 + specs[prefix + ".long_kv_cache"] = Full(block, width, ItemSize(2), ratio) + specs[prefix + ".indexer.k_cache"] = Index(block, index_width, ItemSize(1), ratio, scale_dim=1) + if ratio == 2: + specs[prefix + ".compressor.state_cache"] = State(32, 2 * width, ItemSize(4)) + return specs + + +specs = specs_for(128, 512, 128) +slots = plan(specs) +assert len(specs) == 51 +assert [s.page_size_bytes for s in slots] == [131072, 131072, 131072, 147712] +assert sum(s.page_size_bytes for s in slots) == 540928 +assert [len(s.placements) for s in slots] == [13, 13, 13, 12] +assert plan(dict(reversed(list(specs.items())))) == slots +padded = { + p.name: dataclasses.replace(specs[p.name], page_size_padded=p.page_size_bytes) for s in slots for p in s.placements +} +assert plan(padded) == slots +assert all(s.page_size_padded is None for s in specs.values()) +swa_padding = [ + p.page_size_bytes - sum(planes(specs[p.name])) + for s in slots + for p in s.placements + if isinstance(specs[p.name], SWA) +] +assert swa_padding.count(0) == 30 and swa_padding.count(16640) == 10 +for slot in slots: + for p in slot.placements: + assert p.offset >= 0 and p.offset + p.page_size_bytes <= slot.page_size_bytes + assert sum(planes(specs[p.name])) <= p.page_size_bytes +for idx, slot in enumerate(slots): + kv, index = slot.placements[:2] + assert index.offset == sum(planes(specs[kv.name])) + assert index.offset + index.page_size_bytes == slot.page_size_bytes + assert index.offset + planes(specs[index.name])[0] == (73728 if idx < 3 else 147456) + +# Group membership reference: 8 full resources, 3 states and ten SWA quartets. +groups = [ + [p.name for s in slots for p in s.placements if isinstance(specs[p.name], (Full, Index))], + [p.name for s in slots for p in s.placements if isinstance(specs[p.name], State)], +] +for start in range(0, 40, 4): + groups.append([f"language_model.model.layers.{i}.self_attn.swa_cache" for i in range(start, start + 4)]) +assert [len(g) for g in groups] == [8, 3] + [4] * 10 +placement = {p.name: (idx, p) for idx, s in enumerate(slots) for p in s.placements} +N = 17 +bases = [0] +for s in slots[:-1]: + bases.append(bases[-1] + N * s.page_size_bytes) +intervals = [] +for gid, group in enumerate(groups): + bid = gid + 1 + for name in group: + idx, p = placement[name] + start = bases[idx] + bid * slots[idx].page_size_bytes + p.offset + for size in planes(specs[name]): + intervals.append((start, start + size, name)) + start += size +for i, (a0, a1, _) in enumerate(intervals): + for b0, b1, _ in intervals[i + 1 :]: + assert a1 <= b0 or b1 <= a0 +for start in (127, 128, 129, 255, 256, 257): + ids = [7, 19, 3] + for pos in range(start - 3, start): + original = ids[pos // 128] * 128 + pos % 128 + if pos % 2: + assert original // 2 == ids[pos // 128] * 64 + (pos % 128) // 2 +for block, width, index_width in [(64, 8, 4), (128, 256, 64), (256, 512, 128)]: + small = specs_for(block, width, index_width) + for slot in plan(small): + assert all(p.offset + sum(planes(small[p.name])) <= slot.page_size_bytes for p in slot.placements) +assert "torch" not in __import__("sys").modules and "vllm" not in __import__("sys").modules +print( + json.dumps( + { + "status": "passed", + "scope": "extracted placement arithmetic; no torch/vllm imports or runtime tests", + "groups": len(groups), + "cache_specs": len(specs), + "slots": [s.page_size_bytes for s in slots], + "bytes_per_global_id": sum(s.page_size_bytes for s in slots), + "swa_unpadded": 30, + "swa_padded": 10, + "disjoint_payload_intervals": len(intervals), + "rank_shrink_offsets": "component offsets independent of N", + }, + indent=2, + ) +) + +# Execute the actual grouping and allocator for the optional G12 overlay. +draft_specs = dict(specs) +for stage in range(3): + draft_specs[f"mtp.{stage}.self_attn.swa_cache"] = Draft(128, 512, ItemSize(2)) +draft_slots = plan(draft_specs) +assert len(draft_slots) == 4 and sum(s.page_size_bytes for s in draft_slots) == 540928 +assert [len(s.placements) for s in draft_slots] == [14, 14, 14, 12] +target_groups = symbols["make_cache_groups"](symbols["group_cache_specs"](specs)) +draft_groups = symbols["make_cache_groups"](symbols["group_cache_specs"](draft_specs)) +assert len(draft_groups) == 13 and sum(len(g.layer_names) for g in draft_groups) == 54 +assert draft_groups[:12] == target_groups +assert draft_groups[12].layer_names == [f"mtp.{i}.self_attn.swa_cache" for i in range(3)] +draft_padded = {n: s for g in draft_groups for n, s in g.kv_cache_spec.kv_cache_specs.items()} +assert plan(draft_padded) == plan(dict(reversed(list(draft_padded.items())))) == draft_slots +for override in (None, 5): + count, allocations = symbols["allocate_cache_config"](SimpleNamespace(override=override), draft_groups, 17 * 540928) + assert count == (17 if override is None else override) + assert len(allocations) == 4 and sum(a.size for a in allocations) == count * 540928 + for stage in range(3): + assert f"mtp.{stage}.self_attn.swa_cache" in allocations[stage].shared_by + # Mirrors upstream rank-capacity shrinking without changing offsets/stride. + for allocation, slot in zip(allocations, draft_slots): + assert allocation.size // count * 3 == 3 * slot.page_size_bytes +for override in (1, 18): + try: + symbols["allocate_cache_config"](SimpleNamespace(override=override), draft_groups, 17 * 540928) + except ValueError: + pass + else: + raise AssertionError("Unsafe block override accepted") +draft_intervals = list(intervals) +for stage in range(3): + # G12 owns ID 13, independent of all target group IDs 1..12. + begin = bases[stage] + 13 * draft_slots[stage].page_size_bytes + end = begin + sum(planes(draft_specs[f"mtp.{stage}.self_attn.swa_cache"])) + assert all(end <= lo or hi <= begin for lo, hi, _ in draft_intervals) + draft_intervals.append((begin, end, f"mtp.{stage}")) +assert len(draft_intervals) == 58 +print("PASS: DSpark has 13 groups, 54 specs, four buffers, 540928 bytes/ID and 58 disjoint payload planes.") + +# A verifier writes anchor P and S speculative input rows. After A acceptances +# the next forward starts at P+A+1. Its previous row must survive if that start +# is odd. Test every acceptance count, including complete rejection. +rollback_cases = 0 +for start in (0, 1, 15, 16, 17, 31, 32, 33, 127, 128, 129): + for proposed in range(1, 32): + ring = {p % 32: p for p in range(max(0, start - 32), start)} + end = start + proposed + 1 + for p in range(max(start, end - 32), end): + ring[p % 32] = p + for accepted in range(proposed + 1): + next_start = start + accepted + 1 + if next_start % 2: + assert ring[(next_start - 1) % 32] == next_start - 1 + rollback_cases += 1 +# S=32 can overwrite the anchor needed after complete rejection. +unsafe = {p % 32: p for p in range(1, 33)} +assert unsafe[0] != 0 +print(f"PASS: {rollback_cases} speculative rejection cases; S=32 is correctly outside the safe bound.") + +for mode in ("none", "decode"): + for proposed in (1, 15, 31, 32): + config = SimpleNamespace( + use_v2_model_runner=False, + model_config=SimpleNamespace(enforce_eager=mode == "none"), + compilation_config=SimpleNamespace(cudagraph_mode=mode), + cache_config=SimpleNamespace(enable_prefix_caching=False, cache_dtype="auto"), + scheduler_config=SimpleNamespace(disable_hybrid_kv_cache_manager=False), + parallel_config=SimpleNamespace(), + kv_transfer_config=None, + speculative_config=SimpleNamespace(use_dspark=lambda: True, num_speculative_tokens=proposed), + ) + try: + symbols["validate_cache_runtime"](config) + except ValueError: + assert proposed == 32 + else: + assert proposed < 32 and config.cache_config.cache_dtype == "bfloat16" +config.speculative_config.num_speculative_tokens = 31 +config.speculative_config.num_speculative_tokens_per_batch_size = [(1, 8, 32)] +try: + symbols["validate_cache_runtime"](config) +except ValueError: + pass +else: + raise AssertionError("Per-batch draft length escaped the ring retention limit") +print("PASS: actual runtime guards enforce the retention bound and BF16 draft backend in both target modes.") + +# Model the read-before-write schedule with token identities, including ring wraps. +cases = 0 +for start in range(130): + for length in (0, 1, 2, 15, 16, 17, 31, 32, 33, 129): + cache = {p % 32: p for p in range(max(0, start - 32), start)} + groups = (start + length) // 2 - start // 2 + outputs = {} + for group in range(start // 2, start // 2 + groups): + pair = [2 * group, 2 * group + 1] + actual = [p if p >= start else cache[p % 32] for p in pair] + assert actual == pair + output_row = 2 * group + 1 - start + assert 0 <= output_row < length + outputs[output_row] = tuple(actual) + writes = list(range(start + length - min(length, 32), start + length)) + assert len({p % 32 for p in writes}) == len(writes) + for p in writes: + cache[p % 32] = p + if (start + length) % 2: + assert cache[(start + length - 1) % 32] == start + length - 1 + cases += 1 +for block_id in (1, 3, 7, 16): + for pos in (0, 1, 31, 32, 33, 129): + offset = (block_id * 32 + pos % 32) * 1024 * 4 + assert offset == block_id * 131072 + pos % 32 * 4096 +print(f"PASS: {cases} ring schedules, exact 128-KiB page addressing; state demand is one ID per request.") + +# Check actual manager allocation methods without loading vLLM or torch. + + +class BaseManager: + def __init__(self): + self.req_to_blocks = defaultdict(list) + self.claimed = 0 + + def claim(count): + self.claimed += count + return [SimpleNamespace(block_id=self.claimed)] + + self.block_pool = SimpleNamespace(get_new_blocks=claim) + + +manager_source = ast.parse((ROOT / "vllm_ascend/core/circular_buffer.py").read_text()) +manager_class = next( + n for n in manager_source.body if isinstance(n, ast.ClassDef) and n.name == "AscendCircularBufferManager" +) +manager_env = {"FullAttentionManager": BaseManager} +exec(compile(ast.Module(body=[manager_class], type_ignores=[]), "", "exec"), manager_env) +manager = manager_env["AscendCircularBufferManager"]() +for request in ("a", "b"): + assert manager.get_num_blocks_to_allocate(request, 129, [], 0, 0, 129) == 1 + assert len(manager.allocate_new_blocks(request, 129, 129)) == 1 + for tokens in (130, 1024, 65536): + assert manager.get_num_blocks_to_allocate(request, tokens, [], 0, 0, tokens) == 0 + assert manager.allocate_new_blocks(request, tokens, tokens) == [] + manager.allocate_external_computed_blocks(request, 0, tokens) + manager.remove_skipped_blocks(request, tokens) +assert manager.claimed == 2 and not manager._record_new_block_ids +print("PASS: extracted manager retains exactly one block per request.") diff --git a/tests/deepseek_v41_cache_utils.py b/tests/deepseek_v41_cache_utils.py new file mode 100644 index 000000000000..3c0607d139ad --- /dev/null +++ b/tests/deepseek_v41_cache_utils.py @@ -0,0 +1,68 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Shared cache construction for Aurora worker and native-operator tests.""" + +from types import SimpleNamespace + +import torch +from vllm.v1.kv_cache_interface import KVCacheConfig + +from tests.deepseek_v41_reference import build_v41_cache_specs +from vllm_ascend.core.deepseek_v41 import ( + DeepseekV41DraftSWASpec, + allocate_cache_config, + cache_slots_from_groups, + group_cache_specs, + make_cache_groups, + pool_bytes_per_block, + reshape_cache, +) + + +def make_cache_config(num_blocks, *, block_size=128, head_size=512, index_size=128, draft_layers=0): + config = dict( + num_hidden_layers=40, + compress_ratios=[0, 0] + [2] * 18 + [1] * 20, + kv_source_layers=[2, 8, 14, 20], + index_source_layers=[2, 8, 14, 20, 24, 28, 32, 36], + candidate_source_layer=20, + candidate_topk_blocks=64, + candidate_block_size=8, + index_topk=512, + engram_layer_ids=[1, 14], + sliding_window=128, + head_dim=head_size, + index_head_dim=index_size, + ) + runtime = SimpleNamespace(cache_config=SimpleNamespace(block_size=block_size, num_gpu_blocks_override=None)) + specs = build_v41_cache_specs(config, runtime) + for stage in range(draft_layers): + specs[f"mtp.{stage}.self_attn.swa_cache"] = DeepseekV41DraftSWASpec( + block_size=block_size, + num_kv_heads=1, + head_size=head_size, + dtype=torch.bfloat16, + sliding_window=config["sliding_window"], + cache_dtype_str="bfloat16", + model_version="deepseek_v4", + ) + groups = make_cache_groups(group_cache_specs(specs)) + blocks, tensors = allocate_cache_config(runtime, groups, num_blocks * pool_bytes_per_block(groups)) + return KVCacheConfig(num_blocks=blocks, kv_cache_tensors=tensors, kv_cache_groups=groups) + + +def allocate_cache_views(config, device="cpu"): + specs = {n: s for g in config.kv_cache_groups for n, s in g.kv_cache_spec.kv_cache_specs.items()} + backings, caches = [], {} + for allocation, slot in zip(config.kv_cache_tensors, cache_slots_from_groups(config.kv_cache_groups)): + raw = torch.zeros(allocation.size, dtype=torch.uint8, device=device) + backings.append(raw) + for placement in slot.placements: + caches[placement.name] = reshape_cache( + raw, + specs[placement.name], + num_blocks=config.num_blocks, + offset=placement.offset, + block_stride=slot.page_size_bytes, + ) + return backings, caches diff --git a/tests/deepseek_v41_reference.py b/tests/deepseek_v41_reference.py new file mode 100644 index 000000000000..3667cde76edc --- /dev/null +++ b/tests/deepseek_v41_reference.py @@ -0,0 +1,261 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""PyTorch references and static fixtures for DeepSeek V4.1 tests.""" + +from typing import Any + +import torch +import torch.nn.functional as F + +from vllm_ascend.core.deepseek_v41 import ( + STATE_RING_ROWS, + DeepseekV41CompressorStateSpec, + DeepseekV41FullSpec, + DeepseekV41IndexerSpec, + DeepseekV41SWASpec, +) +from vllm_ascend.models.deepseek_v41.compressor import _read, text_config_of +from vllm_ascend.models.deepseek_v41.model import build_layer_plan + + +def build_v41_cache_specs(config: Any, vllm_config: Any, prefix: str = "model"): + """Describe the source-shared cache graph for allocation tests.""" + config = text_config_of(config) + block_size = vllm_config.cache_config.block_size + if block_size <= 0 or block_size % 2: + raise ValueError("V4.1 logical block_size must be a positive multiple of two") + width = _read(config, "head_dim") + index_width = _read(config, "index_head_dim") + window = _read(config, "sliding_window") + specs = {} + for role in build_layer_plan(config).layers: + attn_prefix = f"{prefix}.layers.{role.layer_idx}.self_attn" + specs[f"{attn_prefix}.swa_cache"] = DeepseekV41SWASpec( + block_size=block_size, + num_kv_heads=1, + head_size=width, + dtype=torch.bfloat16, + sliding_window=window, + ) + if not role.is_kv_source: + continue + specs[f"{attn_prefix}.long_kv_cache"] = DeepseekV41FullSpec( + block_size=block_size, + num_kv_heads=1, + head_size=width, + dtype=torch.bfloat16, + tokens_per_state=role.compress_ratio, + storage_block_size=block_size // role.compress_ratio, + ) + specs[f"{attn_prefix}.indexer.k_cache"] = DeepseekV41IndexerSpec( + block_size=block_size, + num_kv_heads=1, + head_size=index_width, + dtype=torch.int8, + tokens_per_state=role.compress_ratio, + storage_block_size=block_size // role.compress_ratio, + scale_dim=1, + scale_dtype=torch.float16, + ) + if role.compress_ratio == 2: + specs[f"{attn_prefix}.compressor.state_cache"] = DeepseekV41CompressorStateSpec( + block_size=STATE_RING_ROWS, + num_kv_heads=1, + head_size=2 * width, + dtype=torch.float32, + ) + return specs + + +def scatter_cache(cache: torch.Tensor, slots: torch.Tensor, values: torch.Tensor) -> None: + """PyTorch reference for writing flat slots into a paged cache.""" + cache = cache.squeeze(-2) + slots = slots[: values.shape[0]].long() + valid = slots >= 0 + physical = slots.clamp_min(0) + pages = torch.div(physical, cache.shape[1], rounding_mode="floor") + rows = physical.remainder(cache.shape[1]) + write_values = torch.where( + valid.view((-1,) + (1,) * (values.ndim - 1)), + values, + torch.zeros_like(values), + ) + cache[pages, rows] = write_values.to(cache.dtype) + + +def gather_cache_rows(cache: torch.Tensor, slots: torch.Tensor) -> torch.Tensor: + """PyTorch reference for reading rows from a block-strided cache view.""" + cache = cache.squeeze(-2) + slots = slots.long() + pages = torch.div(slots, cache.shape[1], rounding_mode="floor") + rows = slots.remainder(cache.shape[1]) + return cache[pages, rows] + + +def paged_prefix(cache, block_table, length, block_size): + """Materialize one request's logical prefix from a paged cache.""" + if length <= 0: + return cache.new_empty((0, cache.shape[-1])) + blocks = (length + block_size - 1) // block_size + page_ids = block_table[:blocks].long() + return cache.squeeze(-2).index_select(0, page_ids).flatten(0, 1)[:length] + + +def select_candidate_blocks(logits, compress_lens, topk_blocks, block_size): + """PyTorch reference for level-one candidate block selection.""" + width = logits.shape[-1] + if width == 0: + return torch.zeros_like(logits, dtype=torch.bool) + scores = F.pad(logits, (0, -width % block_size), value=-torch.inf) + scores = scores.unflatten(-1, (-1, block_size)).amax(-1) + num_blocks = scores.shape[-1] + if not torch.is_tensor(compress_lens): + compress_lens = torch.tensor(compress_lens, device=logits.device) + last = (compress_lens - 1) // block_size + scores = scores.masked_fill( + torch.arange(num_blocks, device=logits.device) == last, + torch.inf, + ) + top = scores.topk(min(topk_blocks, num_blocks), dim=-1) + keep = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, top.indices, top.values > -torch.inf) + return keep.repeat_interleave(block_size, dim=-1)[..., :width] + + +def select_index_topk(logits, compress_lens, index_topk): + """PyTorch reference for chronological level-two position TopK.""" + width = logits.shape[-1] + if width == 0: + return torch.empty((*logits.shape[:-1], 0), dtype=torch.int32, device=logits.device) + topk = min(index_topk, width) + indices = logits.topk(topk, dim=-1, sorted=False).indices.sort(-1).values + return torch.where(indices < compress_lens, indices, -1).int() + + +def small_op_attention( + q, + positions, + swa_cache, + swa_metadata, + *, + source_cache=None, + source_metadata=None, + compress_ratio=0, + window_size=128, + index_topk=512, + compressed_indices=None, + sinks=None, + softmax_scale=1.0, +): + """PyTorch reference for SWA plus selected compressed KV attention.""" + query_starts = swa_metadata.query_start_loc.tolist() + seq_lens = swa_metadata.seq_lens.tolist() + outputs = [] + for req_idx, (q_start, q_end) in enumerate(zip(query_starts[:-1], query_starts[1:])): + seq_len = int(seq_lens[req_idx]) + local = paged_prefix( + swa_cache, + swa_metadata.block_table[req_idx], + seq_len, + swa_metadata.storage_block_size, + ) + compressed = None + if source_cache is not None: + compressed_len = int(source_metadata.cache_seq_lens[req_idx]) + compressed = paged_prefix( + source_cache, + source_metadata.block_table[req_idx], + compressed_len, + source_metadata.storage_block_size, + ) + for token_idx in range(q_start, q_end): + position = int(positions[token_idx]) + local_start = max(0, position - window_size + 1) + keys = local[local_start : position + 1] + if compressed is not None: + visible = (position + 1) // compress_ratio + if compressed_indices is None: + selected = compressed[max(0, visible - index_topk) : visible] + else: + indices = compressed_indices[token_idx].long() + indices = indices[(indices >= 0) & (indices < visible)] + selected = compressed.index_select(0, indices) + keys = torch.cat((keys, selected)) + logits = torch.einsum("hd,kd->hk", q[token_idx].float(), keys.float()) + logits *= softmax_scale + if sinks is not None: + logits = torch.cat((logits, sinks.float().unsqueeze(-1)), -1) + probs = logits.softmax(-1)[..., :-1] + else: + probs = logits.softmax(-1) + outputs.append(torch.einsum("hk,kd->hd", probs, keys.float())) + return torch.stack(outputs).to(q.dtype) + + +def compressor_ratio2_reference(compressor, x, start_pos: int, state_cache, state_block_table): + """Reference ratio-2 compressor over one request's private FP32 ring.""" + if start_pos < 0 or x.ndim != 2: + raise ValueError("Expected nonnegative start_pos and [tokens, hidden] input") + if ( + state_cache is None + or state_cache.ndim != 3 + or state_cache.shape[-1] != 2 * compressor.width + or state_cache.dtype != torch.float32 + ): + raise ValueError("Ratio2 requires paged FP32 [pages, block_size, 2*head_dim] state") + if not isinstance(state_block_table, (list, tuple)): + raise ValueError("Reference compressor requires a host list/tuple state_block_table") + block_size = state_cache.shape[1] + if block_size != STATE_RING_ROWS or len(state_block_table) != 1: + raise ValueError("State requires one 32-row ring block per request") + + def state_row(position): + offset = position % block_size + physical_block = state_block_table[0] + if not isinstance(physical_block, int) or not 0 < physical_block < state_cache.shape[0]: + raise ValueError("Compressor state refers to an absent/null/out-of-range page") + return state_cache[physical_block, offset] + + if x.shape[0]: + first = start_pos - start_pos % compressor.ratio + for position in range(first, start_pos + x.shape[0]): + state_row(position) + kv = compressor.wkv(x.float()) + score = compressor.wgate(x.float()) + completed = [] + for token in range(x.shape[0]): + position = start_pos + token + row = state_row(position) + row[: compressor.width] = kv[token] + row[compressor.width :] = score[token] + if (position + 1) % compressor.ratio == 0: + group = torch.stack([state_row(position - 1), row]) + pooled = (group[:, : compressor.width] * group[:, compressor.width :].softmax(dim=0)).sum(dim=0) + completed.append(pooled) + latent = torch.stack(completed).to(x.dtype) if completed else x.new_empty((0, compressor.width)) + return compressor.norm(latent) + + +def hc_mixes_reference(layer, x, hc_fn, hc_scale, hc_base): + """Reference mHC coefficient construction.""" + x_float = x.float() + flat = x_float.flatten(-2) + mixes = F.linear(flat, hc_fn) + mixes *= torch.rsqrt(flat.square().mean(-1, keepdim=True) + layer.norm_eps) + pre, post, comb = mixes.split([layer.hc_mult, layer.hc_mult, layer.hc_mult * layer.hc_mult], -1) + pre = torch.sigmoid(pre * hc_scale[0] + hc_base[: layer.hc_mult]) + layer.hc_eps + post = 2 * torch.sigmoid(post * hc_scale[1] + hc_base[layer.hc_mult : 2 * layer.hc_mult]) + comb = comb.unflatten(-1, (layer.hc_mult, layer.hc_mult)) + comb = comb * hc_scale[2] + hc_base[2 * layer.hc_mult :].view(layer.hc_mult, layer.hc_mult) + comb = comb.softmax(-1) + layer.hc_eps + comb = comb / (comb.sum(-2, keepdim=True) + layer.hc_eps) + for _ in range(layer.hc_sinkhorn_iters - 1): + comb = comb / (comb.sum(-1, keepdim=True) + layer.hc_eps) + comb = comb / (comb.sum(-2, keepdim=True) + layer.hc_eps) + return pre, post, comb + + +def hc_post_reference(x, residual, post, comb): + """Reference mHC residual expansion.""" + y = post.unsqueeze(-1) * x.unsqueeze(-2) + y += (comb.unsqueeze(-1) * residual.unsqueeze(-2)).sum(-3) + return y.to(x.dtype) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_cache_store.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_cache_store.py new file mode 100644 index 000000000000..0c397c9ba9b9 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_cache_store.py @@ -0,0 +1,172 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Validate V4.1 fused stores against the PyTorch cache-write reference.""" + +from types import SimpleNamespace + +import pytest +import torch +import torch_npu # noqa: F401 + +from tests.deepseek_v41_cache_utils import allocate_cache_views, make_cache_config +from tests.deepseek_v41_reference import scatter_cache +from vllm_ascend.attention.dsa_v41 import scatter_cache_sk +from vllm_ascend.models.deepseek_v4.dspark import DeepseekV4DSparkModel +from vllm_ascend.utils import enable_custom_op + +enable_custom_op() + + +def _slot_mapping_2d(slots, block_size): + valid = slots >= 0 + physical = slots.clamp_min(0) + indices = torch.stack( + ( + torch.div(physical, block_size, rounding_mode="floor"), + physical.remainder(block_size), + ), + dim=-1, + ).to(torch.int32) + indices[~valid] = -1 + return indices + + +def _reference_scatter(cache, slots, values): + valid = slots >= 0 + scatter_cache(cache, slots[valid], values[valid]) + + +@pytest.mark.parametrize("stage", [0, 1, 2]) +def test_dspark_context_store_uses_slot_backed_view_without_target_corruption(stage): + config = make_cache_config(17, draft_layers=3) + expected_raw, expected = allocate_cache_views(config, "npu") + actual_raw, actual = allocate_cache_views(config, "npu") + name = f"mtp.{stage}.self_attn.swa_cache" + target = f"model.layers.{(2, 8, 14)[stage]}.self_attn.long_kv_cache" + actual[target][2].fill_(7) + expected[target][2].fill_(7) + slots = torch.tensor([-1, 7 * 128 + 127, 13 * 128], dtype=torch.int64, device="npu") + values = torch.randn(3, 512, dtype=torch.bfloat16, device="npu") + model = SimpleNamespace(vllm_config=SimpleNamespace(cache_config=SimpleNamespace(cache_dtype="bfloat16"))) + attn = SimpleNamespace( + dsa_attn=SimpleNamespace( + swa_cache_layer=SimpleNamespace( + kv_cache=[actual[name]], + block_size=128, + ) + ) + ) + DeepseekV4DSparkModel._store_standard_swa_kv(model, values.unsqueeze(1), slots, attn) + _reference_scatter(expected[name], slots, values) + torch.npu.synchronize() + for actual_buffer, expected_buffer in zip(actual_raw, expected_raw): + torch.testing.assert_close(actual_buffer.cpu(), expected_buffer.cpu(), rtol=0, atol=0) + + +@pytest.mark.parametrize( + "name,rows,width,dtype", + [ + ("model.layers.3.self_attn.swa_cache", 128, 512, torch.bfloat16), + ("model.layers.2.self_attn.long_kv_cache", 64, 512, torch.bfloat16), + ("model.layers.20.self_attn.long_kv_cache", 128, 512, torch.bfloat16), + ], +) +def test_fused_store_matches_reference_in_layer_slots(name, rows, width, dtype): + torch.manual_seed(47) + config = make_cache_config(7) + expected_backing, expected = allocate_cache_views(config, "npu") + actual_backing, actual = allocate_cache_views(config, "npu") + slots = torch.tensor( + [-1, rows + 3, 5 * rows + rows - 1], + dtype=torch.int64, + device="npu", + ) + values = torch.randn(3, width, dtype=dtype, device="npu") + + _reference_scatter(expected[name], slots, values) + indices = _slot_mapping_2d(slots, rows) + scatter_cache_sk(actual[name], indices, values) + torch.npu.synchronize() + + for expected_raw, actual_raw in zip(expected_backing, actual_backing): + torch.testing.assert_close(actual_raw.cpu(), expected_raw.cpu(), rtol=0, atol=0) + + +@pytest.mark.parametrize("kind", ["random", "zero", "tiny"]) +def test_indexer_dynamic_quant_and_fused_store_match_reference(kind): + torch.manual_seed(53) + config = make_cache_config(7) + expected_backing, expected = allocate_cache_views(config, "npu") + actual_backing, actual = allocate_cache_views(config, "npu") + name = "model.layers.2.self_attn.indexer.k_cache" + expected_key, expected_scale = expected[name] + actual_key, actual_scale = actual[name] + rows = expected_key.shape[1] + slots = torch.tensor( + [-1, rows + 1, 3 * rows + rows - 1], + dtype=torch.int64, + device="npu", + ) + key = torch.randn(3, 128, dtype=torch.bfloat16, device="npu") + if kind == "zero": + key.zero_() + elif kind == "tiny": + key.mul_(1e-7) + + reference_scale = key.float().abs().amax(-1, keepdim=True).clamp_min_(1e-12) / 127.0 + reference_key = (key.float() / reference_scale).round_().clamp_(-127, 127).to(torch.int8) + actual_quant, actual_quant_scale = torch_npu.npu_dynamic_quant(key, dst_type=torch.int8) + torch.testing.assert_close(actual_quant.cpu(), reference_key.cpu(), rtol=0, atol=0) + torch.testing.assert_close( + actual_quant_scale.float().cpu(), + reference_scale.squeeze(-1).cpu(), + rtol=1e-5, + atol=1e-8, + ) + + _reference_scatter(expected_key, slots, reference_key) + _reference_scatter(expected_scale, slots, reference_scale.to(torch.float16)) + indices = _slot_mapping_2d(slots, rows) + scatter_cache_sk(actual_key, indices, actual_quant) + scatter_cache_sk( + actual_scale, + indices, + actual_quant_scale.unsqueeze(-1).to(torch.float16), + ) + torch.npu.synchronize() + + for expected_raw, actual_raw in zip(expected_backing, actual_backing): + torch.testing.assert_close(actual_raw.cpu(), expected_raw.cpu(), rtol=0, atol=0) + + +@pytest.mark.parametrize( + "name,plane,width,dtype", + [ + ("model.layers.3.self_attn.swa_cache", None, 512, torch.bfloat16), + ("model.layers.2.self_attn.long_kv_cache", None, 512, torch.bfloat16), + ("model.layers.20.self_attn.long_kv_cache", None, 512, torch.bfloat16), + ("model.layers.2.self_attn.indexer.k_cache", 0, 128, torch.int8), + ("model.layers.2.self_attn.indexer.k_cache", 1, 1, torch.float16), + ], +) +def test_negative_coordinates_do_not_modify_packed_backing(name, plane, width, dtype): + torch.manual_seed(59) + config = make_cache_config(7) + backing, caches = allocate_cache_views(config, "npu") + cache = caches[name] if plane is None else caches[name][plane] + before = [tensor.clone() for tensor in backing] + indices = torch.full((3, 2), -1, dtype=torch.int32, device="npu") + if dtype == torch.int8: + values = torch.randint(-127, 128, (3, width), dtype=dtype, device="npu") + else: + values = torch.randn(3, width, dtype=dtype, device="npu") + + # The view is packed into a shared allocation: its page stride is larger + # than the contiguous stride implied by the visible plane shape. + squeezed = cache.squeeze(-2) + assert squeezed.stride(0) > squeezed.shape[1] * squeezed.stride(1) + scatter_cache_sk(cache, indices, values) + torch.npu.synchronize() + + for expected_raw, actual_raw in zip(before, backing): + torch.testing.assert_close(actual_raw.cpu(), expected_raw.cpu(), rtol=0, atol=0) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_hc_fallback.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_hc_fallback.py new file mode 100644 index 000000000000..bf257effe59b --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_hc_fallback.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 + +import torch +import torch.nn.functional as F +import torch_npu # noqa: F401 + +from tests.deepseek_v41_reference import hc_post_reference +from vllm_ascend.models.deepseek_v41.model import DeepseekV41DecoderLayer + +HC_MULT = 4 +HIDDEN_SIZE = 5120 +SINKHORN_ITERS = 3 +NORM_EPS = 1e-6 +HC_EPS = 1e-6 + + +def _layer() -> DeepseekV41DecoderLayer: + layer = DeepseekV41DecoderLayer.__new__(DeepseekV41DecoderLayer) + layer.hc_mult = HC_MULT + layer.hc_sinkhorn_iters = SINKHORN_ITERS + layer.norm_eps = NORM_EPS + layer.hc_eps = HC_EPS + return layer + + +def _reference(x, hc_fn, hc_scale, hc_base, pre_mix): + x_float = x.float() + x_flat = x_float.flatten(-2) + mixes = F.linear(x_flat, hc_fn) * torch.rsqrt(x_flat.square().mean(-1, keepdim=True) + NORM_EPS) + pre, post, comb = mixes.split([HC_MULT, HC_MULT, HC_MULT * HC_MULT], dim=-1) + comb = comb.unflatten(-1, (HC_MULT, HC_MULT)) + pre = torch.sigmoid(pre * hc_scale[0] + hc_base[:HC_MULT]) + HC_EPS + post = 2 * torch.sigmoid(post * hc_scale[1] + hc_base[HC_MULT : 2 * HC_MULT]) + comb = comb * hc_scale[2] + hc_base[2 * HC_MULT :].view(HC_MULT, HC_MULT) + comb = comb.softmax(-1) + HC_EPS + comb = comb / (comb.sum(-2, keepdim=True) + HC_EPS) + for _ in range(SINKHORN_ITERS - 1): + comb = comb / (comb.sum(-1, keepdim=True) + HC_EPS) + comb = comb / (comb.sum(-2, keepdim=True) + HC_EPS) + y = (pre_mix.unsqueeze(-1) * x_float).sum(dim=-2).to(x.dtype) + return y, post, comb, pre + + +def test_v41_hc_pre_handoff_5120_on_npu(): + torch.manual_seed(19) + x = torch.randn(2, HC_MULT, HIDDEN_SIZE, dtype=torch.bfloat16) + hc_fn = torch.randn(24, HC_MULT * HIDDEN_SIZE, dtype=torch.float32) / HIDDEN_SIZE + hc_scale = torch.randn(3, dtype=torch.float32) + hc_base = torch.randn(24, dtype=torch.float32) + pre_mix = torch.rand(2, HC_MULT, dtype=torch.float32) + expected = _reference(x, hc_fn, hc_scale, hc_base, pre_mix) + + actual = _layer().hc_pre(x.npu(), hc_fn.npu(), hc_scale.npu(), hc_base.npu(), pre_mix.npu()) + + for actual_tensor, expected_tensor in zip(actual, expected): + torch.testing.assert_close( + actual_tensor.cpu().float(), + expected_tensor.float(), + atol=5e-3, + rtol=5e-3, + ) + + expected_post = hc_post_reference(expected[0], x, expected[1], expected[2]) + actual_post = _layer().hc_post(actual[0], x.npu(), actual[1], actual[2]) + torch.testing.assert_close(actual_post.cpu().float(), expected_post.float(), atol=2e-2, rtol=2e-2) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_metadata.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_metadata.py new file mode 100644 index 000000000000..c00614fcc124 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_metadata.py @@ -0,0 +1,209 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Exercise shared V4.1 metadata with native operators and changed-input replay.""" + +from types import SimpleNamespace +from typing import Any + +import pytest +import torch +import torch_npu # noqa: F401 +from vllm.forward_context import BatchDescriptor + +from tests.deepseek_v41_cache_utils import make_cache_config +from vllm_ascend.attention.dsa_v41 import DeepseekV41MetadataBuilder +from vllm_ascend.ops import rope_dsv4 +from vllm_ascend.utils import enable_custom_op +from vllm_ascend.worker.device_metadata import DeviceMetadataExecutor + +# SmlaMetadata and QliV2Metadata define the initialized payload. The remainder +# of each 1024-word allocation is reserved and is not read by the kernels. +SMLA_METADATA_WORDS = 36 * 9 + 72 * 8 +QLI_METADATA_WORDS = 36 * 8 + 72 * 8 + +METADATA_TENSORS = ( + "seq_lens", + "cache_seq_lens", + "cmp_residual", + "slot_mapping", + "smla_metadata", + "qli_metadata", + "c2_ring_metadata", + "c2_complete_mask", + "c2_source_positions", + "c2_source_cos", + "c2_source_sin", +) + + +def _builders(runtime, device, deferred): + result = [] + for group in make_cache_config(17).kv_cache_groups: + layers_by_spec: dict[Any, list[str]] = {} + for name, spec in group.kv_cache_spec.kv_cache_specs.items(): + layers_by_spec.setdefault(spec, []).append(name) + group_builders = [] + for spec, names in layers_by_spec.items(): + builder = DeepseekV41MetadataBuilder(spec, names, runtime, device) + # Resolve C2's exact source RoPE table for either execution mode. + builder.enable_device_metadata() + builder._device_metadata_enabled = deferred + group_builders.append(builder) + result.append(group_builders) + return result + + +def _build(builders, lengths, query_len, *, shared, block_offset=0, idle=False): + device = builders[0][0]._seq_lens.device + positions = torch.tensor( + [*range(lengths[0] - query_len, lengths[0]), *range(lengths[1] - query_len, lengths[1]), 0], + device=device, + dtype=torch.int64, + ) + query_start_loc_cpu = torch.tensor([0, query_len, 2 * query_len, 2 * query_len + 1], dtype=torch.int32) + seq_lens_cpu = torch.tensor([*lengths, 0], dtype=torch.int32) + query_start_loc = query_start_loc_cpu.to(device) + seq_lens = seq_lens_cpu.to(device) + batch_shared: dict[str, Any] = {} + metadata, tasks = [], [] + for gid, group_builders in enumerate(builders): + group_shared: dict[str, Any] = {} + blocks = torch.tensor([gid + 1 + block_offset, gid + 2 + block_offset, 0], device=device, dtype=torch.int32) + slots = torch.cat( + ( + blocks[0] * 128 + positions[:query_len] % 128, + blocks[1] * 128 + positions[query_len : 2 * query_len] % 128, + torch.full((1,), -1, device=device, dtype=torch.int64), + ) + ) + common = SimpleNamespace( + query_start_loc=query_start_loc, + query_start_loc_cpu=query_start_loc_cpu, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + positions=positions, + slot_mapping=slots, + block_table_tensor=blocks[:, None], + num_reqs=3, + num_input_tokens=positions.numel(), + num_actual_tokens=2 * query_len, + max_query_len=query_len, + max_seq_len=max(lengths), + is_prefilling=torch.tensor([query_len > 1, query_len > 1, False]), + ) + for builder in group_builders: + metadata.append( + builder.build( + 0, + common, + num_actual_reqs=2, + skip_ring_state_update=idle, + common_v41_metadata=group_shared, + common_v41_batch_metadata=batch_shared if shared else None, + ) + ) + tasks.extend(builder.take_device_metadata_tasks()) + return metadata, tasks + + +def _consume(metadata): + outputs = {} + for index, entry in enumerate(metadata): + for field in METADATA_TENSORS: + value = getattr(entry, field) + if value is not None: + if field == "smla_metadata": + value = value[:SMLA_METADATA_WORDS] + elif field == "qli_metadata": + value = value[:QLI_METADATA_WORDS] + outputs[index, field] = value.clone() + if entry.cos is not None: + for layer in (0, 2): + name = f"model.layers.{layer}.self_attn.attn" + outputs[index, f"cos:{layer}"] = entry.cos[name].clone() + outputs[index, f"sin:{layer}"] = entry.sin[name].clone() + return outputs + + +@pytest.mark.parametrize("deferred", [False, True]) +@pytest.mark.parametrize("query_len", [1, 3]) +@torch.inference_mode() +def test_grouped_metadata_matches_independent_builds_and_replays(monkeypatch, deferred, query_len): + assert enable_custom_op(), "V4.1 native metadata operators are required" + device = torch.device("npu:0") + torch.npu.set_device(device) + monkeypatch.setattr(rope_dsv4, "_ROPE_STATE", rope_dsv4.RopeGlobalState()) + runtime = SimpleNamespace( + model_config=SimpleNamespace( + hf_text_config=dict( + sliding_window=128, + num_attention_heads=32, + head_dim=512, + qk_rope_head_dim=64, + index_topk=512, + index_n_heads=64, + index_head_dim=128, + ) + ), + parallel_config=SimpleNamespace(tensor_parallel_size=1), + scheduler_config=SimpleNamespace(max_num_batched_tokens=8, max_num_seqs=4), + speculative_config=None, + ) + with device: + for layer in (0, 2, 8, 14): + rope_dsv4.ComplexExpRotaryEmbedding( + runtime, + f"model.layers.{layer}.self_attn.attn", + head_size=512, + rotary_dim=64, + max_position_embeddings=512, + base=10000 if layer == 0 else 1000000, + scaling_factor=1, + ) + reference_builders = _builders(runtime, device, False) + shared_builders = _builders(runtime, device, deferred) + executor = DeviceMetadataExecutor() if deferred else None + descriptor = BatchDescriptor(num_tokens=3, num_reqs=3) if query_len == 1 else None + graph = None + pointers = None + for iteration, (lengths, idle) in enumerate([([127, 128], False), ([130, 131], False), ([130, 131], True)]): + reference, _ = _build(reference_builders, lengths, query_len, shared=False, block_offset=iteration, idle=idle) + expected = {key: value.cpu() for key, value in _consume(reference).items()} + metadata, tasks = _build(shared_builders, lengths, query_len, shared=True, block_offset=iteration, idle=idle) + if executor is not None: + assert len(tasks) == 6 + executor.submit(tasks, descriptor) + current_pointers = tuple( + tuple(getattr(entry, field).data_ptr() for field in METADATA_TENSORS if getattr(entry, field) is not None) + for entry in metadata + ) + if pointers is not None: + assert current_pointers == pointers + pointers = current_pointers + if query_len == 1 and graph is not None: + graph.replay() + elif query_len == 1: + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True): + if executor is not None: + for task in tasks: + executor.wait(task.stage, task.group_id) + outputs = _consume(metadata) + graph.replay() + else: + if executor is not None: + for task in tasks: + executor.wait(task.stage, task.group_id) + outputs = _consume(metadata) + if executor is not None: + executor.release() + torch.npu.synchronize() + assert outputs.keys() == expected.keys() + for key, value in outputs.items(): + torch.testing.assert_close( + value.cpu(), + expected[key], + rtol=0, + atol=0, + msg=lambda detail, iteration=iteration, key=key: f"iteration={iteration}, {key}: {detail}", + ) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_qli.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_qli.py new file mode 100644 index 000000000000..65c81771865d --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_qli.py @@ -0,0 +1,448 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Aurora QLI/candidate correctness on A3, including real paged cache views.""" + +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F +import torch_npu # noqa: F401 + +from tests.deepseek_v41_cache_utils import allocate_cache_views, make_cache_config +from vllm_ascend.utils import enable_custom_op, is_950 + +enable_custom_op() + +HEADS = 32 +WIDTH = 128 +PAGE = 64 +# Fixed before NPU validation. Values at the TopK cutoff may tie or differ by +# FP16 intermediate rounding; every selected score must meet this cutoff. +RTOL = 1e-3 +ATOL = 1e-5 + + +def _paged(value, *, strided=True, page_size=PAGE): + pages = (value.shape[0] + page_size - 1) // page_size + shape = (pages, page_size, *value.shape[1:]) + dense = torch.zeros(shape, dtype=value.dtype) + order = torch.arange(pages - 1, -1, -1) + padded = F.pad(value.flatten(1), (0, 0, 0, pages * page_size - value.shape[0])).reshape(shape) + dense[order] = padded + dense = dense.npu() + if strided: + stride = (dense.stride(0) * 2, *dense.stride()[1:]) + backing = torch.zeros(pages * stride[0] + 128, dtype=value.dtype, device="npu") + view = backing.as_strided(shape, stride, 128) + view.copy_(dense) + dense = view + return dense, order.to(torch.int32).unsqueeze(0).npu() + + +def _data(length, qlen=5, heads=HEADS, seed=41, strided=True, page_size=PAGE): + gen = torch.Generator().manual_seed(seed) + q = torch.randint(-8, 9, (qlen, heads, WIDTH), dtype=torch.int8, generator=gen) + k = torch.randint(-8, 9, (length, 1, WIDTH), dtype=torch.int8, generator=gen) + w = (torch.rand(qlen, heads, generator=gen) * 2 - 0.5).half() + qs = (torch.rand(qlen, heads, generator=gen) / 32 + 0.01).half() + ks = (torch.rand(length, 1, generator=gen) / 16 + 0.01).half() + key, table = _paged(k, strided=strided, page_size=page_size) + scale, _ = _paged(ks, strided=strided, page_size=page_size) + return q, k, w, qs, ks, key, scale, table + + +def _scores(q, k, w, qs, ks, ratio, residual, mask=3): + # Independent CPU form of the vendor INT8 golden's two matmuls. Preserve + # FP16 QK/1024 and w*qscale rounding before the FP32 head reduction. + qk = torch.einsum("qhd,kd->qhk", q.float(), k[:, 0].float()) + relu = (qk / 1024).clamp_min(0).half().float() + score = (relu * (w * qs).float().unsqueeze(-1)).sum(1) * ks[:, 0].float() + visible = torch.full((q.shape[0],), k.shape[0], dtype=torch.long) + if mask == 3: + visible = (k.shape[0] * ratio + residual - q.shape[0] + torch.arange(1, q.shape[0] + 1)) // ratio + visible.clamp_min_(0) + score.masked_fill_(torch.arange(k.shape[0])[None, :] >= visible[:, None], -torch.inf) + return score, visible + + +def _mask_non_candidates(score, membership): + # A2/A3 retain reachable non-candidates as low-priority TopK fillers. + # Causal padding must remain -inf; the A5 kernel is unchanged. + penalty = -torch.inf if is_950() else -1e30 + score.masked_fill_(~membership & torch.isfinite(score), penalty) + + +def _check_topk(score, indices, count): + indices = indices.cpu().reshape(score.shape[0], -1) + for row, idx in zip(score, indices): + valid = idx[idx >= 0].long() + expected_count = min(count, int(torch.isfinite(row).sum())) + assert valid.numel() == expected_count + assert valid.unique().numel() == valid.numel() + assert bool((valid < row.numel()).all()) + if expected_count: + selected = row[valid] + assert bool(torch.isfinite(selected).all()) + cutoff = row.topk(expected_count).values[-1] + assert selected.min() >= cutoff - (ATOL + RTOL * cutoff.abs()) + + +def _check_candidates(score, visible, candidates, blocks): + block_score = F.pad(score, (0, -score.shape[-1] % 8), value=-torch.inf).unflatten(-1, (-1, 8)).amax(-1) + for i, n in enumerate(visible): + if n > 0: + block_score[i, (n - 1) // 8] = torch.inf + indices = candidates.cpu().reshape(score.shape[0], -1) + for row, n, idx in zip(block_score, visible, indices): + valid = idx[idx >= 0].long() + reachable = int((n + 7) // 8) + assert valid.numel() == min(blocks, reachable) + assert valid.unique().numel() == valid.numel() + if not reachable: + continue + assert int((n - 1) // 8) in valid.tolist() + cutoff = row.topk(min(blocks, reachable)).values[-1] + if not torch.isinf(cutoff): + assert row[valid].min() >= cutoff - (ATOL + RTOL * cutoff.abs()) + return indices + + +def _invoke(data, layout, ratio, mode, *, candidates=None, blocks=64, mask=3, residual=None, topk=None): + q, k, w, qs, ks, key, scale, table = data + if residual is None: + residual = ratio - 1 + topk = min(128, k.shape[0]) if topk is None else topk + cu = torch.tensor([0, q.shape[0]], dtype=torch.int32, device="npu") + used = torch.tensor([k.shape[0]], dtype=torch.int32, device="npu") + res = torch.tensor([residual], dtype=torch.int32, device="npu") + common = dict( + cu_seqlens_q=cu if layout == "TND" else None, + seqused_k=used, + cmp_residual_k=res if ratio != 1 and mask != 0 else None, + max_seqlen_q=q.shape[0], + layout_q=layout, + layout_k="PA_BBND", + mask_mode=mask, + cmp_ratio=ratio, + ) + metadata = torch.ops._C_ascend.npu_quant_lightning_indexer_v2_metadata( + q.shape[1], 1, WIDTH, topk, 2, batch_size=1, max_seqlen_k=k.shape[0], **common + ) + query, weights, query_scale = q.npu(), w.npu(), qs.npu() + if layout == "BSND": + query, weights, query_scale = (x.unsqueeze(0) for x in (query, weights, query_scale)) + print( + { + "layout": layout, + "ratio": ratio, + "candidate_mode": mode, + "q_shape": list(query.shape), + "k_shape": list(key.shape), + "k_stride": list(key.stride()), + "scale_stride": list(scale.stride()), + "k_offset": key.storage_offset(), + "topk": topk, + "candidate_blocks": blocks, + }, + flush=True, + ) + output, values, candidate_out = torch.ops._C_ascend.npu_quant_lightning_indexer_v3( + query, + key, + weights, + query_scale, + scale, + topk, + 2, + candidate_topk_index=candidates, + block_table=table, + metadata=metadata, + candidate_mode=mode, + candidate_topk_blocks=blocks, + candidate_block_size=8, + **common, + ) + torch.npu.synchronize() + assert values.numel() == 0 + return output, candidate_out + + +@pytest.mark.parametrize("ratio", [1, 2]) +@pytest.mark.parametrize("length", [7, 513, 1025]) +@pytest.mark.parametrize("mode", [1, 2, 3]) +def test_native_qli_candidate(ratio, length, mode): + layout = "TND" + data = _data(length) + score, visible = _scores(data[0], data[1], data[2], data[3], data[4], ratio, ratio - 1) + candidate_in = None + if mode == 2: + _, candidate_in = _invoke(data, layout, ratio, 1) + _check_candidates(score, visible, candidate_in, 64) + # A different consumer query must rerank within source blocks. + other = _data(length, seed=17) + data = (other[0], data[1], other[2], other[3], *data[4:]) + score, visible = _scores(data[0], data[1], data[2], data[3], data[4], ratio, ratio - 1) + block_ids = candidate_in.cpu().reshape(data[0].shape[0], -1) + membership = (torch.arange(length)[None, None, :] // 8 == block_ids[:, :, None]).any(1) + _mask_non_candidates(score, membership) + output, candidates = _invoke(data, layout, ratio, mode, candidates=candidate_in) + _check_topk(score, output, min(128, length)) + if mode == 1: + _check_candidates(score, visible, candidates, 64) + else: + assert candidates.numel() == 0 + + +def test_native_qli_supports_bsnd_query_layout(): + data = _data(513) + score, _ = _scores(data[0], data[1], data[2], data[3], data[4], 1, 0) + output, _ = _invoke(data, "BSND", 1, 3) + _check_topk(score, output, 128) + + +@pytest.mark.parametrize( + "ratio,residual,qlen,length,blocks", + [ + (2, 0, 2, 1, 64), # First query has no completed compressed key. + (2, 0, 1, 513, 64), + (1, 0, 2, 17001, 2048), # Production candidate width, actual filtering. + ], +) +def test_candidate_boundaries(ratio, residual, qlen, length, blocks): + data = _data(length, qlen=qlen) + score, visible = _scores(data[0], data[1], data[2], data[3], data[4], ratio, residual) + output, candidates = _invoke(data, "TND", ratio, 1, residual=residual, blocks=blocks) + _check_topk(score, output, min(128, length)) + _check_candidates(score, visible, candidates, blocks) + + +@pytest.mark.parametrize("mode", [1, 3]) +def test_fake_qli_shapes(mode): + q = torch.empty(3, 32, 128, device="meta", dtype=torch.int8) + k = torch.empty(16, 64, 1, 128, device="meta", dtype=torch.int8) + w = torch.empty(3, 32, device="meta", dtype=torch.float16) + ks = torch.empty(16, 64, 1, device="meta", dtype=torch.float16) + out, values, candidates = torch.ops._C_ascend.npu_quant_lightning_indexer_v3( + q, k, w, w, ks, 128, 2, candidate_mode=mode, candidate_topk_blocks=64 + ) + assert out.shape == (3, 1, 128) and out.dtype == torch.int32 + assert candidates.shape == ((3, 1, 64) if mode == 1 else (0,)) + assert values.shape == (0,) + + +def _indexer(ratio): + from vllm_ascend.models.deepseek_v41.indexer import DeepseekV41Indexer + + obj = DeepseekV41Indexer.__new__(DeepseekV41Indexer) + torch.nn.Module.__init__(obj) + obj.width, obj.n_heads, obj.index_topk, obj.compress_ratio = WIDTH, HEADS, 128, ratio + return obj + + +def _quant_query_reference(query): + scale = (query.float().abs().amax(-1) / 127).half().clamp_min(2.0**-24) + quant = (query.float() / scale.float().unsqueeze(-1)).round().clamp(-127, 127).to(torch.int8) + return quant, scale + + +@pytest.mark.parametrize("ratio,zero_first", [(1, False), (2, False), (2, True)]) +def test_model_indexer_mixed_batch(ratio, zero_first): + page_size = 128 // ratio + parts = [_data(7, qlen=2, seed=11, page_size=page_size), _data(1025, qlen=3, seed=29, page_size=page_size)] + key = torch.cat([d[5] for d in parts]) + scale = torch.cat([d[6] for d in parts]).unsqueeze(-1) + config = make_cache_config(key.shape[0] + 1) + _, caches = allocate_cache_views(config, "npu") + source = 2 if ratio == 2 else 20 + slot_key, slot_scale = caches[f"model.layers.{source}.self_attn.indexer.k_cache"] + slot_key[1:].copy_(key) + slot_scale[1:].copy_(scale) + key, scale = slot_key, slot_scale + assert key.shape[1] == page_size + assert not key.is_contiguous() and not scale.is_contiguous() + assert not key[0].any() and not scale[0].any() + table = torch.zeros(2, parts[1][7].shape[1], dtype=torch.int32, device="npu") + table[0, : parts[0][7].shape[1]] = parts[0][7][0] + 1 + table[1] = parts[1][7][0] + parts[0][5].shape[0] + 1 + lengths = [0 if zero_first else 7, 1025] + original = [lengths[0] * ratio + (ratio - 1), 1025 * ratio + (ratio - 1)] + # A zero-key request must still represent a valid original token range. + if zero_first: + parts[0] = _data(7, qlen=1, seed=11) + original[0] = 1 + qlens = [d[0].shape[0] for d in parts] + starts = [0, qlens[0], sum(qlens)] + query = torch.cat([d[0].float() * 0.03125 for d in parts]).bfloat16() + weights = torch.cat([d[2] for d in parts]) + positions = torch.cat([torch.arange(n - q, n) for n, q in zip(original, qlens)]) + metadata = SimpleNamespace( + max_cache_seq_len=1025, + max_query_len=max(qlens), + query_start_loc=torch.tensor(starts, dtype=torch.int32, device="npu"), + cache_seq_lens=torch.tensor(lengths, dtype=torch.int32, device="npu"), + seq_lens=torch.tensor(original, dtype=torch.int32, device="npu"), + block_table=table, + ) + metadata.cmp_residual = ( + torch.tensor([value % ratio for value in original], dtype=torch.int32, device="npu") if ratio != 1 else None + ) + metadata.qli_metadata = torch.ops._C_ascend.npu_quant_lightning_indexer_v2_metadata( + HEADS, + 1, + WIDTH, + 128, + 2, + cu_seqlens_q=metadata.query_start_loc, + seqused_k=metadata.cache_seq_lens, + cmp_residual_k=metadata.cmp_residual, + batch_size=2, + max_seqlen_q=metadata.max_query_len, + max_seqlen_k=metadata.max_cache_seq_len, + layout_q="TND", + layout_k="PA_BBND", + mask_mode=3, + cmp_ratio=ratio, + ) + obj = _indexer(ratio) + options = dict(candidate_topk_blocks=64, candidate_block_size=8) + out, candidates = obj.select_projected( + query.npu(), + weights.npu(), + positions.npu(), + (key, scale), + metadata, + is_candidate_source=True, + uses_candidate_filter=False, + candidates=None, + **options, + ) + assert out.shape == (sum(qlens), 128) + assert candidates.shape == (sum(qlens), 1, 64) + for i, (start, end) in enumerate(zip(starts[:-1], starts[1:])): + q, qs = _quant_query_reference(query[start:end]) + k = parts[i][1][: lengths[i]] + ks = parts[i][4][: lengths[i]] + score, visible = _scores(q, k, weights[start:end], qs, ks, ratio, original[i] % ratio) + _check_topk(score, out[start:end], 128) + if lengths[i]: + _check_candidates(score, visible, candidates[start:end], 64) + else: + assert bool((candidates[start:end] == -1).all()) + # Same shared candidate tensor, different consumer projection. + consumer_query = -query + consumer_out, retained = obj.select_projected( + consumer_query.npu(), + weights.npu(), + positions.npu(), + (key, scale), + metadata, + is_candidate_source=False, + uses_candidate_filter=True, + candidates=candidates, + **options, + ) + assert retained is candidates + for i, (start, end) in enumerate(zip(starts[:-1], starts[1:])): + q, qs = _quant_query_reference(consumer_query[start:end]) + score, _ = _scores( + q, parts[i][1][: lengths[i]], weights[start:end], qs, parts[i][4][: lengths[i]], ratio, original[i] % ratio + ) + block_ids = candidates[start:end].cpu().reshape(end - start, -1) + keep = (torch.arange(lengths[i])[None, None, :] // 8 == block_ids[:, :, None]).any(1) + _mask_non_candidates(score, keep) + _check_topk(score, consumer_out[start:end], 128) + valid = consumer_out[start:end].cpu() + valid = torch.where(valid < 0, torch.iinfo(torch.int32).max, valid) + assert bool((valid[:, 1:] >= valid[:, :-1]).all()) + + +def test_missing_candidate_rejected(): + obj = _indexer(1) + with pytest.raises(RuntimeError, match="before its source"): + obj.select_projected( + None, + None, + None, + None, + None, + is_candidate_source=False, + uses_candidate_filter=True, + candidate_topk_blocks=64, + candidate_block_size=8, + candidates=None, + ) + + +def test_empty_compressed_cache(): + obj = _indexer(2) + q = torch.empty(1, 32, 128, device="npu", dtype=torch.bfloat16) + out, candidates = obj.select_projected( + q, + None, + None, + None, + SimpleNamespace(max_cache_seq_len=0), + is_candidate_source=True, + uses_candidate_filter=False, + candidate_topk_blocks=64, + candidate_block_size=8, + candidates=None, + ) + assert out.shape == (1, 0) and candidates.shape == (1, 1, 64) + assert bool((candidates == -1).all()) + + +@pytest.mark.parametrize("ratio", [1, 2]) +def test_64_head_candidate_consumer(ratio): + data = _data(1025, qlen=1, heads=64) + score, visible = _scores(data[0], data[1], data[2], data[3], data[4], ratio, ratio - 1) + output, candidates = _invoke(data, "TND", ratio, 1) + _check_topk(score, output, 128) + _check_candidates(score, visible, candidates, 64) + block_ids = candidates.cpu().reshape(1, -1) + keep = (torch.arange(1025)[None, None, :] // 8 == block_ids[:, :, None]).any(1) + _mask_non_candidates(score, keep) + output, _ = _invoke(data, "TND", ratio, 2, candidates=candidates) + _check_topk(score, output, 128) + + +@pytest.mark.parametrize("topk", [512, 2048]) +def test_position_topk_width(topk): + data = _data(4097, qlen=1) + score, visible = _scores(data[0], data[1], data[2], data[3], data[4], 2, 1) + output, candidates = _invoke(data, "TND", 2, 1, topk=topk) + _check_topk(score, output, topk) + _check_candidates(score, visible, candidates, 64) + block_ids = candidates.cpu().reshape(1, -1) + keep = (torch.arange(4097)[None, None, :] // 8 == block_ids[:, :, None]).any(1) + _mask_non_candidates(score, keep) + output, _ = _invoke(data, "TND", 2, 2, candidates=candidates, topk=topk) + _check_topk(score, output, topk) + + +@pytest.mark.skipif(is_950(), reason="Candidate fill semantics changed only on A2/A3") +@pytest.mark.parametrize("strided", [False, True]) +@pytest.mark.parametrize("length,qlen", [(91, 5), (4096, 1)]) +def test_candidate_shortfall_keeps_reachable_indices(length, qlen, strided): + topk = 2048 + data = _data(length, qlen=qlen, strided=strided) + score, visible = _scores(data[0], data[1], data[2], data[3], data[4], 1, 0) + block_ids = torch.arange(64, dtype=torch.int32).repeat(qlen, 1) + block_ids[:, 9:11] = -1 + block_ids[block_ids >= (length + 7) // 8] = -1 + candidates = block_ids.unsqueeze(1).npu() + membership = (torch.arange(length)[None, None, :] // 8 == block_ids[:, :, None]).any(1) + _mask_non_candidates(score, membership) + + output, _ = _invoke(data, "TND", 1, 2, candidates=candidates, topk=topk) + _check_topk(score, output, topk) + # Missing candidate blocks must not turn reachable tokens into -1 slots. + for row, indices in enumerate(output.cpu().reshape(qlen, topk)): + valid = indices[indices >= 0].long() + reachable = torch.arange(int(visible[row])) + candidate_positions = reachable[membership[row, : reachable.numel()]] + assert torch.isin(candidate_positions, valid).all() + outside_count = (~membership[row, valid]).sum().item() + assert outside_count == min(topk, reachable.numel()) - candidate_positions.numel() + if reachable.numel() <= topk: + assert torch.equal(valid.sort().values, reachable) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_ring_compressor.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_ring_compressor.py new file mode 100644 index 000000000000..249b4009296e --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_ring_compressor.py @@ -0,0 +1,207 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Deferred NPU checks for the actual four-slot FP32 circular state views.""" + +import pytest +import torch +import torch_npu # noqa: F401 + +from tests.deepseek_v41_cache_utils import allocate_cache_views, make_cache_config +from vllm_ascend.models.deepseek_v41.compressor import DeepseekV41RMSNorm +from vllm_ascend.ops.triton.compressor.compressor_triton import _cube_core_num, compressor_from_projected + + +def _reference(kv, scores, state, metadata): + out = torch.zeros_like(kv, dtype=torch.bfloat16) + # Snapshot residuals before updating: chunks may span several ring wraps. + for start, used, _, base, block in metadata.T.tolist(): + if not used or not 0 < block < state.shape[0]: + continue + for group in range(start // 2, (start + used) // 2): + rows = [] + for position in (2 * group, 2 * group + 1): + if position < start: + rows.append(state[block, position % 32].clone()) + else: + index = base + position - start + rows.append(torch.cat((kv[index], scores[index]))) + pair = torch.stack(rows) + width = kv.shape[1] + out[base + 2 * group + 1 - start] = ( + (pair[:, :width] * pair[:, width:].softmax(dim=0)).sum(dim=0).to(torch.bfloat16) + ) + for index in range(max(0, used - 32), used): + state[block, (start + index) % 32] = torch.cat((kv[base + index], scores[base + index])) + return out + + +def _inputs(length, start): + torch.manual_seed(41) + _, caches = allocate_cache_views(make_cache_config(13), "npu") + state = caches["model.layers.2.self_attn.compressor.state_cache"].squeeze(-2) + assert state.shape == (13, 32, 1024) and state.stride() == (32768, 1024, 1) + # Nonzero stale pages detect unexpected null/inactive writes and BF16 casts. + initial = torch.randn(state.shape, dtype=torch.float32) + initial[:, :, 0] = 1.0001 + state.copy_(initial) + kv_cpu = torch.randn(length + 2, 512, dtype=torch.float32) + kv_cpu[:, 0] = 1.0001 + scores_cpu = torch.randn_like(kv_cpu) + meta_cpu = torch.tensor( + [[start, start + 1, 0], [length, 1, 0], [0, length, length + 1], [0, length, length + 1], [3, 7, 0]], + dtype=torch.int32, + ) + return state, initial, kv_cpu, scores_cpu, meta_cpu + + +@pytest.mark.parametrize("length", [0, 1, 2, 15, 16, 17, 31, 32, 33, 129]) +@pytest.mark.parametrize("start", [0, 1, 31, 32, 33]) +def test_projected_ring_matches_fp32_reference(length, start): + state, initial, kv_cpu, scores_cpu, meta_cpu = _inputs(length, start) + expected_state = initial.clone() + expected = _reference(kv_cpu, scores_cpu, expected_state, meta_cpu) + kv, scores, meta = kv_cpu.npu(), scores_cpu.npu(), meta_cpu.npu() + out = torch.empty_like(kv, dtype=torch.bfloat16) + compressor_from_projected(kv, scores, state, meta, out, max_query_len=max(length, 1), num_cores=_cube_core_num()) + torch.testing.assert_close(out.cpu(), expected, rtol=0.016, atol=1e-5) + torch.testing.assert_close(state.cpu(), expected_state, rtol=0, atol=0) + norm = DeepseekV41RMSNorm(512, 1e-6) + expected_norm = expected.float() * torch.rsqrt(expected.float().square().mean(-1, keepdim=True) + norm.eps) + expected_norm = expected_norm.to(expected.dtype) * norm.weight + torch.testing.assert_close(norm.npu()(out).cpu(), expected_norm, rtol=0.016, atol=1e-5) + + +def test_projected_ring_graph_replay_uses_new_metadata_and_request_ids(): + state, initial, kv_cpu, scores_cpu, meta_cpu = _inputs(1, 0) + kv, scores, meta = kv_cpu.npu(), scores_cpu.npu(), meta_cpu.npu() + out = torch.empty_like(kv, dtype=torch.bfloat16) + cores = _cube_core_num() + + def run(): + compressor_from_projected(kv, scores, state, meta, out, max_query_len=1, num_cores=cores) + + run() # Compile outside capture. + torch.npu.synchronize() + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True): + run() + pointers = (meta.data_ptr(), state.data_ptr(), out.data_ptr()) + for start in (0, 1, 31, 32, 33, 127, 128, 129): + controls = meta_cpu.clone() + controls[0] = torch.tensor([start, start + 1, start + 2]) + controls[1] = torch.tensor([1, 0, 1]) if start % 2 else torch.tensor([1, 1, 0]) + controls[4] = torch.tensor([7, 0, 3]) if start % 2 else torch.tensor([3, 7, 0]) + state.copy_(initial) + meta.copy_(controls) + expected_state = initial.clone() + expected = _reference(kv_cpu, scores_cpu, expected_state, controls) + graph.replay() + torch.npu.synchronize() + assert pointers == (meta.data_ptr(), state.data_ptr(), out.data_ptr()) + torch.testing.assert_close(out.cpu(), expected, rtol=0.016, atol=1e-5) + torch.testing.assert_close(state.cpu(), expected_state, rtol=0, atol=0) + + +def test_empty_projected_batch_keeps_ring_untouched(): + state, initial, _, _, _ = _inputs(0, 0) + kv = torch.empty(0, 512, dtype=torch.float32, device="npu") + out = torch.empty_like(kv, dtype=torch.bfloat16) + controls = torch.empty(5, 0, dtype=torch.int32, device="npu") + compressor_from_projected(kv, kv, state, controls, out, max_query_len=0, num_cores=1) + torch.testing.assert_close(state.cpu(), initial, rtol=0, atol=0) + + +@pytest.mark.parametrize("start", [0, 1, 31, 32, 127, 128]) +@pytest.mark.parametrize("accepted", [0, 1, 15, 30, 31]) +def test_ring_residual_after_dspark_rejection_matches_verified_prefix(start, accepted): + _, caches = allocate_cache_views(make_cache_config(17, draft_layers=3), "npu") + state = caches["model.layers.2.self_attn.compressor.state_cache"].squeeze(-2) + state[7].fill_(17) + kv = torch.randn(32, 512, dtype=torch.float32, device="npu") + scores = torch.randn_like(kv) + metadata = torch.tensor([[start], [32], [0], [0], [7]], dtype=torch.int32, device="npu") + out = torch.empty_like(kv, dtype=torch.bfloat16) + cores = _cube_core_num() + compressor_from_projected(kv, scores, state, metadata, out, max_query_len=32, num_cores=cores) + next_start = start + accepted + 1 + next_kv = torch.randn(2, 512, dtype=torch.float32, device="npu") + next_scores = torch.randn_like(next_kv) + controls = torch.tensor([[next_start], [2], [0], [0], [7]], dtype=torch.int32, device="npu") + next_out = torch.empty_like(next_kv, dtype=torch.bfloat16) + if next_start % 2: + pair_kv = torch.stack((kv[accepted], next_kv[0])) + pair_scores = torch.stack((scores[accepted], next_scores[0])) + output_index = 0 + else: + pair_kv, pair_scores = next_kv, next_scores + output_index = 1 + expected = (pair_kv * pair_scores.softmax(0)).sum(0).bfloat16() + compressor_from_projected(next_kv, next_scores, state, controls, next_out, max_query_len=2, num_cores=cores) + torch.testing.assert_close(next_out[output_index].cpu(), expected.cpu(), rtol=0.016, atol=1e-5) + assert not next_out[1 - output_index].any() + # G12's separate live ID is untouched by target state updates. + assert not caches["mtp.0.self_attn.swa_cache"][13].any() + + +@pytest.mark.parametrize("invalid", ["state_dtype", "projection_dtype", "strided_state", "ring_rows", "output_dtype"]) +def test_projected_ring_rejects_incompatible_views(invalid): + state, _, kv_cpu, scores_cpu, controls = _inputs(1, 0) + kv, scores, meta = kv_cpu.npu(), scores_cpu.npu(), controls.npu() + out = torch.empty_like(kv, dtype=torch.bfloat16) + if invalid == "state_dtype": + state = state.bfloat16() + elif invalid == "projection_dtype": + kv = kv.bfloat16() + elif invalid == "strided_state": + state = torch.empty(13, 64, 1024, dtype=torch.float32, device="npu")[:, ::2] + elif invalid == "ring_rows": + state = state[:, :16] + else: + out = out.float() + with pytest.raises(ValueError): + compressor_from_projected(kv, scores, state, meta, out, max_query_len=1, num_cores=1) + + +@pytest.mark.parametrize("graph_mode", [False, True]) +def test_mixed_prefill_first_residual_uses_ring_at_projection_boundary(graph_mode): + """First request history must not read row -1 of packed projections.""" + torch.manual_seed(1301) + kv_cpu = torch.randn(3968, 512) + scores_cpu = torch.randn_like(kv_cpu) + # Allocate projections before other large tensors to exercise their boundary. + kv, scores = kv_cpu.npu(), scores_cpu.npu() + initial = torch.zeros(5954, 32, 1024, dtype=torch.float32) + blocks = [12, 145, 278, 411, 544, 677] + for block in blocks: + initial[block] = torch.randn(32, 1024) + state = initial.npu() + controls = torch.tensor( + [ + [1301, 1301, 1366, 0, 0, 0], + [6, 6, 15, 1338, 1300, 1303], + [0, 6, 12, 27, 1365, 2665], + [0, 6, 12, 27, 1365, 2665], + blocks, + ], + dtype=torch.int32, + ) + meta = controls.npu() + out = torch.empty_like(kv, dtype=torch.bfloat16) + cores = _cube_core_num() + expected_state = initial.clone() + expected = _reference(kv_cpu, scores_cpu, expected_state, controls) + + def run(): + compressor_from_projected(kv, scores, state, meta, out, max_query_len=1338, num_cores=cores) + + run() + if graph_mode: + torch.npu.synchronize() + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True): + run() + state.copy_(initial) + graph.replay() + torch.npu.synchronize() + torch.testing.assert_close(out.cpu(), expected, rtol=0.016, atol=1e-5) + torch.testing.assert_close(state.cpu(), expected_state, rtol=0, atol=0) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_rmsnorm.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_rmsnorm.py new file mode 100644 index 000000000000..70680ef24669 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_rmsnorm.py @@ -0,0 +1,57 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""NPU coverage for the shared compressor and indexer K RMSNorm.""" + +import pytest +import torch +import torch_npu # noqa: F401 + +from vllm_ascend.models.deepseek_v41.compressor import DeepseekV41RMSNorm + + +def _reference(x, weight, eps): + normalized = x.float() * torch.rsqrt(x.float().square().mean(-1, keepdim=True) + eps) + return normalized.to(x.dtype) * weight + + +@pytest.mark.parametrize("width", [128, 512]) +@pytest.mark.parametrize("tokens", [0, 1, 32, 4096]) +@pytest.mark.parametrize("eps", [1e-6, 1e-3]) +@pytest.mark.parametrize("scale", [0.0, 1e-4, 1.0]) +@torch.inference_mode() +def test_rmsnorm_matches_reference(width, tokens, eps, scale): + torch.manual_seed(41) + x = (torch.randn(tokens, width) * scale).bfloat16() + norm = DeepseekV41RMSNorm(width, eps).npu() + norm.weight.copy_(torch.randn_like(norm.weight)) + expected = _reference(x, norm.weight.cpu(), eps) + x_npu = x.npu() + actual = norm(x_npu) + assert actual.shape == x.shape and actual.dtype == x.dtype + torch.testing.assert_close(x_npu.cpu(), x, rtol=0, atol=0) + # The old equation rounds to BF16 before multiplying by the weight. + torch.testing.assert_close(actual.cpu(), expected, rtol=0.016, atol=1e-5) + + +@pytest.mark.parametrize("width", [128, 512]) +@torch.inference_mode() +def test_rmsnorm_graph_replay_uses_new_input(width): + torch.manual_seed(42) + norm = DeepseekV41RMSNorm(width, 1e-6).npu() + norm.weight.copy_(torch.randn_like(norm.weight)) + weight = norm.weight.cpu() + x = torch.randn(32, width, dtype=torch.bfloat16, device="npu") + norm(x) + torch.npu.synchronize() + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True): + actual = norm(x) + pointers = (x.data_ptr(), actual.data_ptr()) + for scale in (1.0, 1e-4, 0.0): + updated = (torch.randn(32, width) * scale).bfloat16() + x.copy_(updated) + graph.replay() + torch.npu.synchronize() + assert (x.data_ptr(), actual.data_ptr()) == pointers + expected = _reference(updated, weight, norm.eps) + torch.testing.assert_close(actual.cpu(), expected, rtol=0.016, atol=1e-5) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_sparse_mla.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_sparse_mla.py new file mode 100644 index 000000000000..844233a0fb9f --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_deepseek_v41_sparse_mla.py @@ -0,0 +1,144 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Compare Aurora's slot-backed BF16 cache with the prior block-outermost views.""" + +from types import SimpleNamespace + +import pytest +import torch +import torch_npu # noqa: F401 + +from tests.deepseek_v41_cache_utils import allocate_cache_views, make_cache_config +from vllm_ascend.attention.dsa_v41 import DeepseekV41EagerAttentionImpl +from vllm_ascend.utils import enable_custom_op + +enable_custom_op() + +WIDTH = 512 +HEADS = 32 +BLOCK = 128 +OLD_BLOCK_STRIDE = 393216 + + +@pytest.mark.parametrize("ratio", [0, 1, 2]) +@pytest.mark.parametrize("query_lens", [(1,), (3,), (1, 3)]) +def test_slot_backed_attention_matches_block_outermost(ratio, query_lens): + torch.manual_seed(41) + config = make_cache_config(13) + specs = {n: s for g in config.kv_cache_groups for n, s in g.kv_cache_spec.kv_cache_specs.items()} + _, candidate = allocate_cache_views(config, "npu") + # Original token lengths include odd compression tails and page boundaries. + lengths = [257, 383][: len(query_lens)] + ori_ids = [[3, 1, 5], [2, 4, 6]][: len(lengths)] + cmp_ids = [[9, 7, 11], [8, 10, 12]][: len(lengths)] + swa_name = "model.layers.3.self_attn.swa_cache" # Padded fourth slot. + cmp_name = f"model.layers.{2 if ratio == 2 else 20}.self_attn.long_kv_cache" + names = [swa_name] + ([cmp_name] if ratio else []) + baseline = {} + for name in names: + spec = specs[name] + raw = torch.zeros(config.num_blocks * OLD_BLOCK_STRIDE, dtype=torch.uint8, device="npu") + baseline[name] = raw.view(spec.dtype).as_strided( + (config.num_blocks, spec.storage_block_size, 1, WIDTH), + (OLD_BLOCK_STRIDE // spec.dtype.itemsize, WIDTH, WIDTH, 1), + ) + ids = ori_ids if name == swa_name else cmp_ids + cache_lengths = lengths if name == swa_name else [n // ratio for n in lengths] + for physical_ids, count in zip(ids, cache_lengths): + data = torch.randn(count, 1, WIDTH, dtype=torch.bfloat16, device="npu") + for logical, physical in enumerate(physical_ids): + start = logical * spec.storage_block_size + end = min(start + spec.storage_block_size, count) + if end > start: + baseline[name][physical, : end - start].copy_(data[start:end]) + candidate[name][physical, : end - start].copy_(data[start:end]) + + query = torch.randn(sum(query_lens), HEADS, WIDTH, dtype=torch.bfloat16, device="npu") + cu = torch.tensor([0, *torch.tensor(query_lens).cumsum(0).tolist()], dtype=torch.int32, device="npu") + seq_lens = torch.tensor(lengths, dtype=torch.int32, device="npu") + cmp_lens = seq_lens // ratio if ratio else None + metadata = SimpleNamespace( + swa=SimpleNamespace( + num_reqs=len(lengths), + query_start_loc=cu, + seq_lens=seq_lens, + block_table=torch.tensor(ori_ids, dtype=torch.int32, device="npu"), + max_query_len=max(query_lens), + max_seq_len=max(lengths), + ori_sparse_indices=None, + ori_topk_length=None, + ori_mask_mode=4, + ori_win_left=127, + ori_win_right=0, + ), + attention=SimpleNamespace( + block_table=torch.tensor(cmp_ids, dtype=torch.int32, device="npu"), + cache_seq_lens=cmp_lens, + max_cache_seq_len=max(lengths) // ratio, + ) + if ratio + else None, + ) + indices = None + if ratio: + indices = torch.full((sum(query_lens), 512), -1, dtype=torch.int32, device="npu") + row = 0 + for length, qlen in zip(lengths, query_lens): + for position in range(length - qlen, length): + visible = (position + 1) // ratio + indices[row, :visible] = torch.arange(visible, dtype=torch.int32, device="npu") + row += 1 + operator_metadata = metadata.attention if ratio else metadata.swa + cmp_residual = torch.remainder(seq_lens, ratio) if ratio else None + if ratio: + operator_metadata.cmp_residual = cmp_residual + operator_metadata.smla_metadata = torch.ops._C_ascend.npu_sparse_flash_mla_metadata( + HEADS, + 1, + WIDTH, + cu_seqlens_q=cu, + seqused_ori_kv=seq_lens, + seqused_cmp_kv=cmp_lens, + cmp_residual_kv=cmp_residual, + batch_size=len(lengths), + max_seqlen_q=max(query_lens), + max_seqlen_ori_kv=max(lengths), + max_seqlen_cmp_kv=max(lengths) // ratio if ratio else 0, + ori_topk=0, + ori_topk_length=None, + cmp_topk=512 if ratio else 0, + cmp_ratio=ratio, + ori_mask_mode=4, + cmp_mask_mode=3 if ratio else 0, + ori_win_left=127, + ori_win_right=0, + layout_q="TND", + layout_kv="PA_BBND", + has_ori_kv=True, + has_cmp_kv=bool(ratio), + ) + impl = DeepseekV41EagerAttentionImpl.__new__(DeepseekV41EagerAttentionImpl) + impl.role = SimpleNamespace(compress_ratio=ratio) + impl.topology = SimpleNamespace(index_topk=512) + outputs = [] + for caches in (baseline, candidate): + attn = SimpleNamespace( + head_dim=WIDTH, + window_size=128, + n_local_heads=HEADS, + shared_state=SimpleNamespace(smla_metadata={}), + dsa_attn=SimpleNamespace(swa_cache_layer=SimpleNamespace(kv_cache=[caches[swa_name]])), + attn_sink=torch.zeros(HEADS, dtype=torch.float32, device="npu"), + softmax_scale=WIDTH**-0.5, + ) + outputs.append( + impl._native_attention( + attn, + query, + metadata, + source_cache=caches[cmp_name] if ratio else None, + compressed_indices=indices, + ).cpu() + ) + assert torch.isfinite(outputs[0]).all() and torch.isfinite(outputs[1]).all() + torch.testing.assert_close(outputs[0], outputs[1], rtol=0, atol=0) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_npu_hc_pre.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_npu_hc_pre.py index c0180053dd69..a4b50858396a 100644 --- a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_npu_hc_pre.py +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_npu_hc_pre.py @@ -55,7 +55,8 @@ def _hc_pre_cpu( hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + pre_mix: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: x_float = x.float() x_flat = x_float.flatten(-2) inv_rms = torch.rsqrt(x_flat.square().mean(-1, keepdim=True) + NORM_EPS) @@ -78,8 +79,9 @@ def _hc_pre_cpu( comb_frag = comb_frag / (comb_frag.sum(-1, keepdim=True) + HC_EPS) comb_frag = comb_frag / (comb_frag.sum(-2, keepdim=True) + HC_EPS) - y = (pre.unsqueeze(-1) * x_float).sum(dim=-2).to(x.dtype) - return y, post, comb_frag + mix_for_y = pre_mix.float() if pre_mix is not None else pre + y = (mix_for_y.unsqueeze(-1) * x_float).sum(dim=-2).to(x.dtype) + return y, post, comb_frag, pre def _assert_close_with_pass_rate( @@ -105,7 +107,7 @@ def _assert_close_with_pass_rate( def _compare_hc_pre_with_cpu(shape: tuple[int, ...]): x, hc_fn, hc_scale, hc_base = _make_hc_pre_inputs(shape) - expected_y, expected_post, expected_comb_frag = _hc_pre_cpu( + expected_y, expected_post, expected_comb_frag, _ = _hc_pre_cpu( x, hc_fn, hc_scale, @@ -179,3 +181,45 @@ def test_npu_hc_pre_v2_bf16_extended_hidden_size(): gc.collect() torch.npu.empty_cache() torch.npu.reset_peak_memory_stats() + + +@torch.inference_mode() +def test_npu_hc_pre_v3_uses_external_pre_mix(): + shape = (2, HC_MULT, HIDDEN_SIZE) + x, hc_fn, hc_scale, hc_base = _make_hc_pre_inputs(shape) + pre_mix = torch.rand(shape[:-1], dtype=torch.float32) + expected_y, expected_post, expected_comb_frag, expected_pre = _hc_pre_cpu(x, hc_fn, hc_scale, hc_base, pre_mix) + y, post, comb_frag, pre = torch.ops._C_ascend.npu_hc_pre_v3( + x.npu(), + hc_fn.npu(), + hc_scale.npu(), + hc_base.npu(), + pre_mix.npu(), + hc_mult=HC_MULT, + hc_sinkhorn_iters=HC_SINKHORN_ITERS, + norm_eps=NORM_EPS, + hc_eps=HC_EPS, + ) + assert y.shape == (shape[0], shape[-1]) + assert post.shape == (shape[0], HC_MULT) + assert comb_frag.shape == (shape[0], HC_MULT, HC_MULT) + assert pre.shape == (shape[0], HC_MULT) + assert y.dtype == torch.bfloat16 + assert post.dtype == torch.float32 + assert comb_frag.dtype == torch.float32 + assert pre.dtype == torch.float32 + for actual, expected, diff_threshold, required_pass_rate in ( + (y, expected_y, Y_DIFF_THRESHOLD, Y_REQUIRED_PASS_RATE), + (post, expected_post, AUX_DIFF_THRESHOLD, AUX_REQUIRED_PASS_RATE), + (comb_frag, expected_comb_frag, AUX_DIFF_THRESHOLD, AUX_REQUIRED_PASS_RATE), + (pre, expected_pre, AUX_DIFF_THRESHOLD, AUX_REQUIRED_PASS_RATE), + ): + _assert_close_with_pass_rate( + actual, + expected, + diff_threshold=diff_threshold, + required_pass_rate=required_pass_rate, + ) + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_npu_moe_gating_top_k_hash.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_npu_moe_gating_top_k_hash.py new file mode 100644 index 000000000000..3627d2efc194 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_npu_moe_gating_top_k_hash.py @@ -0,0 +1,115 @@ +import pytest +import torch + +from vllm_ascend.utils import enable_custom_op + +enable_custom_op() + +IMAGE_SENTINEL_LO = 129257 +IMAGE_SENTINEL_COUNT = 5 + + +def _reference( + logits: torch.Tensor, + text_bias: torch.Tensor, + bias_vl: torch.Tensor, + input_ids: torch.Tensor, + tid2eid: torch.Tensor | None, + top_k: int, + routed_scaling_factor: float, +) -> tuple[torch.Tensor, torch.Tensor]: + scores = torch.nn.functional.softplus(logits.float()).sqrt() + image_mask = (input_ids >= IMAGE_SENTINEL_LO) & (input_ids < IMAGE_SENTINEL_LO + IMAGE_SENTINEL_COUNT) + row_bias = torch.where( + image_mask[:, None], + bias_vl.float()[None, :], + text_bias.float()[None, :], + ) + dynamic_ids = torch.topk(scores + row_bias, top_k, dim=-1, sorted=True).indices + if tid2eid is None: + expert_ids = dynamic_ids + else: + text_ids = tid2eid[input_ids.clamp_max(tid2eid.shape[0] - 1)].long() + expert_ids = torch.where(image_mask[:, None], dynamic_ids, text_ids) + weights = scores.gather(1, expert_ids) + weights = weights / weights.sum(dim=-1, keepdim=True) + return (weights * routed_scaling_factor).to(logits.dtype), expert_ids.int() + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("with_hash", [False, True]) +@pytest.mark.parametrize("execution_mode", ["eager", "graph"]) +def test_v41_vision_bias_and_image_sentinel( + dtype: torch.dtype, + with_hash: bool, + execution_mode: str, +): + torch.manual_seed(20260909) + rows, experts, top_k = 8, 384, 6 + logits = torch.randn(rows, experts, dtype=dtype) + text_bias = torch.randn(experts, dtype=dtype) * 0.2 + bias_vl = torch.randn(experts, dtype=dtype) * 0.2 + input_ids = torch.tensor( + [11, IMAGE_SENTINEL_LO, 22, IMAGE_SENTINEL_LO + 2, 33, IMAGE_SENTINEL_LO + 4, 44, 55], + dtype=torch.int64, + ) + tid2eid = None + if with_hash: + tid2eid = torch.empty(64, top_k, dtype=torch.int32) + for token_id in range(tid2eid.shape[0]): + tid2eid[token_id] = torch.randperm(experts)[:top_k] + + expected_weights, expected_ids = _reference( + logits, + text_bias, + bias_vl, + input_ids, + tid2eid, + top_k, + routed_scaling_factor=1.5, + ) + npu_logits = logits.npu() + npu_text_bias = text_bias.npu() + npu_bias_vl = bias_vl.npu() + npu_input_ids = input_ids.npu() + npu_tid2eid = tid2eid.npu() if tid2eid is not None else None + + def run_op(): + return torch.ops._C_ascend.moe_gating_top_k_hash( + x=npu_logits, + k=top_k, + bias=npu_text_bias, + input_ids=npu_input_ids, + tid2eid=npu_tid2eid, + k_group=1, + group_count=1, + routed_scaling_factor=1.5, + eps=1e-20, + group_select_mode=1, + renorm=0, + norm_type=2, + out_flag=False, + bias_vl=npu_bias_vl, + image_sentinel_lo=IMAGE_SENTINEL_LO, + image_sentinel_count=IMAGE_SENTINEL_COUNT, + ) + + if execution_mode == "graph": + run_op() + torch.npu.synchronize() + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph): + actual_weights, actual_ids, _ = run_op() + graph.replay() + torch.npu.synchronize() + else: + actual_weights, actual_ids, _ = run_op() + + torch.testing.assert_close(actual_ids.cpu(), expected_ids, rtol=0, atol=0) + tolerance = 1e-5 if dtype == torch.float32 else 1e-2 + torch.testing.assert_close( + actual_weights.cpu(), + expected_weights, + rtol=tolerance, + atol=tolerance, + ) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_prepare_indexer_indices.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_prepare_indexer_indices.py new file mode 100644 index 000000000000..a26657cd8660 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_prepare_indexer_indices.py @@ -0,0 +1,84 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch +import torch_npu # noqa: F401 + +from vllm_ascend.ops.triton.prepare_indexer_indices import prepare_indexer_indices + + +def reference(selected, positions, compress_ratio): + visible = ((positions + 1) // compress_ratio).unsqueeze(-1) + valid = (selected >= 0) & (selected < visible) + sentinel = torch.iinfo(torch.int32).max + selected = torch.where(valid, selected, sentinel).sort(dim=-1).values + return torch.where(selected == sentinel, -1, selected) + + +@pytest.mark.parametrize("topk", [1, 7, 8, 33, 128, 512, 2047, 2048]) +@pytest.mark.parametrize("tokens", [0, 1, 3, 41, 129]) +@pytest.mark.parametrize("compress_ratio", [1, 2]) +@torch.inference_mode() +def test_prepare_indexer_indices(topk, tokens, compress_ratio): + torch.manual_seed(41) + selected = torch.randint(-3, 1000, (tokens, topk), dtype=torch.int32, device="npu") + positions = torch.randint(-1, 2000, (tokens,), dtype=torch.int64, device="npu") + original = selected.clone() + actual = prepare_indexer_indices(selected, positions, compress_ratio) + assert actual.is_contiguous() + torch.testing.assert_close(actual.cpu(), reference(selected.cpu(), positions.cpu(), compress_ratio), rtol=0, atol=0) + torch.testing.assert_close(selected, original, rtol=0, atol=0) + + +@pytest.mark.parametrize("compress_ratio", [1, 2]) +@pytest.mark.parametrize("position_dtype", [torch.int32, torch.int64]) +@torch.inference_mode() +def test_prepare_indexer_indices_boundaries(compress_ratio, position_dtype): + # Preserve distinct INT32 indices above FP32's exact-integer range, ties, + # negative sentinels, causal boundaries and all-invalid rows. + row = torch.tensor([0, -1, -3, 9, 9, 10, 11, 2**24 - 1, 2**24, 2**24 + 1, 2**24 + 2, 2**31 - 2, 2**31 - 1]) + selected = row.int().repeat(5, 1).npu() + positions = torch.tensor([-1, 0, 19, 2**25 + 3, 2**31 - 1], dtype=position_dtype, device="npu") + actual = prepare_indexer_indices(selected, positions, compress_ratio) + torch.testing.assert_close(actual.cpu(), reference(selected.cpu(), positions.cpu(), compress_ratio), rtol=0, atol=0) + + +@torch.inference_mode() +def test_prepare_indexer_indices_full_int32_range(): + torch.manual_seed(42) + selected = torch.randint(0, 2**31 - 1, (41, 2048), dtype=torch.int32, device="npu") + positions = torch.full((41,), 2**32, dtype=torch.int64, device="npu") + actual = prepare_indexer_indices(selected, positions, 2) + torch.testing.assert_close(actual.cpu(), reference(selected.cpu(), positions.cpu(), 2), rtol=0, atol=0) + + +@torch.inference_mode() +def test_prepare_indexer_indices_noncontiguous(): + selected = torch.randint(-1, 2000, (41, 256), dtype=torch.int32, device="npu")[:, ::2] + positions = torch.arange(82, dtype=torch.int64, device="npu")[::2] + actual = prepare_indexer_indices(selected, positions, 2) + torch.testing.assert_close(actual.cpu(), reference(selected.cpu(), positions.cpu(), 2), rtol=0, atol=0) + + +@pytest.mark.parametrize("compress_ratio", [1, 2]) +@pytest.mark.parametrize("tokens", [41, 129]) +@torch.inference_mode() +def test_prepare_indexer_indices_graph_replay(compress_ratio, tokens): + selected = torch.randint(-1, 4096, (tokens, 2048), dtype=torch.int32, device="npu") + positions = torch.full((tokens,), 4095, dtype=torch.int64, device="npu") + prepare_indexer_indices(selected, positions, compress_ratio) + torch.npu.synchronize() + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True): + actual = prepare_indexer_indices(selected, positions, compress_ratio) + pointer = actual.data_ptr() + for last_position in (0, 127, 8191): + selected.copy_(torch.randint_like(selected, -1, 4096)) + positions.fill_(last_position) + graph.replay() + torch.npu.synchronize() + assert actual.data_ptr() == pointer + torch.testing.assert_close( + actual.cpu(), reference(selected.cpu(), positions.cpu(), compress_ratio), rtol=0, atol=0 + ) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_qli_candidate_rows.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_qli_candidate_rows.py new file mode 100644 index 000000000000..d5f128afe347 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_qli_candidate_rows.py @@ -0,0 +1,76 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch +import torch_npu # noqa: F401 + +from vllm_ascend.utils import bootstrap_custom_op_env + +bootstrap_custom_op_env(include_vendor_lib=True) +import vllm_ascend.vllm_ascend_C # type: ignore[import-untyped] # noqa: E402,F401 + + +@pytest.mark.parametrize("ratio", [1, 2]) +@pytest.mark.parametrize("query_count", [1, 2, 5]) +def test_candidate_consumer_preserves_each_query_mask(ratio, query_count): + """A multi-query tile must reload each row's candidates across KV tiles.""" + torch.manual_seed(71) + device = "npu" + heads, width, topk, block_size, kv_len = 64, 128, 1024, 128, 6144 + query = torch.randint(-90, 90, (query_count, heads, width), dtype=torch.int8, device=device) + key = torch.randint(-90, 90, (kv_len // block_size, block_size, 1, width), dtype=torch.int8, device=device) + weights = torch.rand(query_count, heads, dtype=torch.float16, device=device) + query_scale = torch.full((query_count, heads), 0.01, dtype=torch.float16, device=device) + key_scale = torch.full(key.shape[:-1], 0.01, dtype=torch.float16, device=device) + qsl = torch.tensor([0, query_count], dtype=torch.int32, device=device) + lengths = torch.tensor([kv_len], dtype=torch.int32, device=device) + residual = torch.zeros_like(lengths) if ratio == 2 else None + table = torch.arange(kv_len // block_size, dtype=torch.int32, device=device).unsqueeze(0) + # Alternate masks: row 0 excludes the middle KV tile, row 1 excludes the + # first. A stale mask remains plausible but selects forbidden positions. + masks = [torch.cat((torch.arange(256), torch.arange(512, 768))), torch.arange(256, 768)] + candidates_cpu = torch.stack([masks[row % 2] for row in range(query_count)]).int() + candidates = candidates_cpu.unsqueeze(1).to(device) + common = dict( + cu_seqlens_q=qsl, + seqused_k=lengths, + cmp_residual_k=residual, + max_seqlen_q=query_count, + layout_q="TND", + layout_k="PA_BBND", + mask_mode=3, + cmp_ratio=ratio, + ) + metadata = torch.ops._C_ascend.npu_quant_lightning_indexer_v2_metadata( + heads, + 1, + width, + topk, + 2, + batch_size=1, + max_seqlen_k=kv_len, + device=str(query.device), + **common, + ) + selected, _, _ = torch.ops._C_ascend.npu_quant_lightning_indexer_v3( + query, + key, + weights, + query_scale, + key_scale, + topk, + 2, + block_table=table, + metadata=metadata, + candidate_topk_index=candidates, + candidate_mode=2, + candidate_topk_blocks=512, + candidate_block_size=8, + **common, + ) + for row, indices in enumerate(selected.cpu().reshape(query_count, topk)): + assert indices.unique().numel() == topk + visible = (kv_len * ratio - query_count + row + 1) // ratio + assert ((indices >= 0) & (indices < visible)).all() + assert torch.isin(indices // 8, candidates_cpu[row]).all(), f"wrong candidate row at query {row}" diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_quantize_indexer_query.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_quantize_indexer_query.py new file mode 100644 index 000000000000..d1f63a0e75ca --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_quantize_indexer_query.py @@ -0,0 +1,81 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch +import torch_npu # noqa: F401 + +from vllm_ascend.ops.triton.quantize_indexer_query import quantize_indexer_query + + +def reference(query): + scale = (query.float().abs().amax(-1) / 127.0).half().clamp_min_(2.0**-24) + quantized = (query.float() / scale.float().unsqueeze(-1)).round().clamp(-127, 127).to(torch.int8) + return quantized, scale + + +def assert_quantized_equal(query, expected): + actual = quantize_indexer_query(query) + for output, ref in zip(actual, expected): + assert output.is_contiguous() + torch.testing.assert_close(output.cpu(), ref.cpu(), rtol=0, atol=0) + + +@pytest.mark.parametrize("heads", [32, 64]) +@pytest.mark.parametrize("tokens", [0, 1, 3, 32, 129, 4096]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) +@torch.inference_mode() +def test_quantize_indexer_query(heads, tokens, dtype): + torch.manual_seed(41) + query = torch.randn(tokens, heads, 128, dtype=dtype, device="npu") + original = query.clone() + assert_quantized_equal(query, reference(query)) + torch.testing.assert_close(query, original, rtol=0, atol=0) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) +@torch.inference_mode() +def test_quantize_indexer_query_rounding_and_scale(dtype): + # Scale=1 exposes positive and negative ties. Other rows cover zero, + # FP16 subnormal scales, scale rounding and FP16 scale overflow. + query = torch.zeros(1, 32, 128, dtype=torch.float32) + query[0, 0, :10] = torch.tensor([127, -127, 0.5, -0.5, 1.5, -1.5, 2.5, -2.5, 3.5, -3.5]) + query[0, 2] = query[0, 0] * 2.0**-24 + query[0, 3] = query[0, 0] * 2.0**-25 + query[0, 4] = torch.linspace(-1.001, 1.001, 128) + query[0, 5] = torch.finfo(dtype).max + query = query.to(dtype).npu() + expected = reference(query) + torch.testing.assert_close( + expected[0][0, 0, :10].cpu(), torch.tensor([127, -127, 0, 0, 2, -2, 2, -2, 4, -4], dtype=torch.int8) + ) + assert_quantized_equal(query, expected) + + +@pytest.mark.parametrize("heads", [32, 64]) +@torch.inference_mode() +def test_quantize_indexer_query_noncontiguous(heads): + query = torch.randn(3, heads, 256, dtype=torch.bfloat16, device="npu")[..., ::2] + assert not query.is_contiguous() + assert_quantized_equal(query, reference(query)) + + +@pytest.mark.parametrize("heads", [32, 64]) +@pytest.mark.parametrize("tokens", [3, 129]) +@torch.inference_mode() +def test_quantize_indexer_query_graph_replay(heads, tokens): + query = torch.randn(tokens, heads, 128, dtype=torch.bfloat16, device="npu") + quantize_indexer_query(query) + torch.npu.synchronize() + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True): + actual = quantize_indexer_query(query) + pointers = tuple(output.data_ptr() for output in actual) + for magnitude in (2.0, 1e-6, 0.0): + query.copy_(torch.randn_like(query) * magnitude) + expected = reference(query) + graph.replay() + torch.npu.synchronize() + assert tuple(output.data_ptr() for output in actual) == pointers + for output, ref in zip(actual, expected): + torch.testing.assert_close(output.cpu(), ref.cpu(), rtol=0, atol=0) diff --git a/tests/e2e/pull_request/one_card/test_deepseek_v4_vision_precision.py b/tests/e2e/pull_request/one_card/test_deepseek_v4_vision_precision.py new file mode 100644 index 000000000000..5401fb1162ab --- /dev/null +++ b/tests/e2e/pull_request/one_card/test_deepseek_v4_vision_precision.py @@ -0,0 +1,218 @@ +import json +import os +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F +import torch_npu # noqa: F401 +from safetensors.torch import load_file +from vllm.config import VllmConfig, set_current_vllm_config + +from vllm_ascend.models.deepseek_v4.vision import ( + DeepseekV4VisionAttention, + DeepseekV4VisionBlock, + DeepseekV4ViT, + get_vision_cos_sin, +) +from vllm_ascend.utils import register_ascend_customop + +VISION_DIM = 1024 +VISION_HEADS = 16 +VISION_INTER_DIM = 2816 +RMS_NORM_EPS = 1e-6 + + +@pytest.fixture(scope="module", autouse=True) +def register_ascend_vision_ops(): + register_ascend_customop() + with set_current_vllm_config(VllmConfig()): + yield + + +@pytest.fixture +def vision_config(): + return SimpleNamespace( + vision_dim=VISION_DIM, + vision_n_heads=VISION_HEADS, + vision_inter_dim=VISION_INTER_DIM, + ) + + +def _new_bf16_module(module_cls, config): + original_dtype = torch.get_default_dtype() + try: + torch.set_default_dtype(torch.bfloat16) + module = module_cls(config) + finally: + torch.set_default_dtype(original_dtype) + return module.eval().npu() + + +def _reference_rms_norm(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + x_fp32 = x.float() + normalized = x_fp32 * torch.rsqrt(x_fp32.square().mean(-1, keepdim=True) + RMS_NORM_EPS) + return (weight * normalized).to(dtype) + + +def _reference_apply_rotary(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + x1, x2 = x.float().chunk(2, dim=-1) + return torch.cat((x1 * cos - x2 * sin, x2 * cos + x1 * sin), dim=-1).to(dtype) + + +def _reference_attention( + attention: DeepseekV4VisionAttention, + x: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, +) -> torch.Tensor: + num_tokens = x.size(0) + q, k, v = ( + tensor.view(num_tokens, attention.n_heads, attention.head_dim) for tensor in attention.wqkv(x).chunk(3, dim=-1) + ) + q = _reference_apply_rotary(q, cos, sin) + k = _reference_apply_rotary(k, cos, sin) + output = F.scaled_dot_product_attention( + q.transpose(0, 1), + k.transpose(0, 1), + v.transpose(0, 1), + ) + return attention.wo(output.transpose(0, 1).reshape(num_tokens, -1)) + + +def _reference_mlp(block: DeepseekV4VisionBlock, x: torch.Tensor) -> torch.Tensor: + gate, up = block.mlp.w1(x).chunk(2, dim=-1) + return block.mlp.w2(F.silu(gate) * up) + + +def _reference_block( + block: DeepseekV4VisionBlock, + x: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, +) -> torch.Tensor: + attention_input = _reference_rms_norm(x, block.norm1.weight) + residual = x + _reference_attention(block.attn, attention_input, cos, sin) + mlp_input = _reference_rms_norm(residual, block.norm2.weight) + return residual + _reference_mlp(block, mlp_input) + + +def _assert_vision_close( + actual: torch.Tensor, + expected: torch.Tensor, + *, + max_abs_error: float, + mean_abs_error: float, + min_cosine_similarity: float, +) -> None: + actual_fp32 = actual.float().reshape(-1) + expected_fp32 = expected.float().reshape(-1) + error = (actual_fp32 - expected_fp32).abs() + cosine_similarity = F.cosine_similarity(actual_fp32, expected_fp32, dim=0) + + max_error = error.max().item() + mean_error = error.mean().item() + cosine = cosine_similarity.item() + metrics = f"max_abs_error={max_error}, mean_abs_error={mean_error}, cosine_similarity={cosine}" + print(metrics) + + assert max_error <= max_abs_error, metrics + assert mean_error <= mean_abs_error, metrics + assert cosine >= min_cosine_similarity, metrics + + +@pytest.mark.parametrize("grid", [(4, 4), (16, 24), (48, 72)]) +def test_fused_attention_matches_original_formula(vision_config, grid): + torch.manual_seed(7) + attention = _new_bf16_module(DeepseekV4VisionAttention, vision_config) + num_tokens = grid[0] * grid[1] + x = torch.randn(num_tokens, VISION_DIM, dtype=torch.bfloat16, device="npu") + cos, sin = get_vision_cos_sin(grid[0], grid[1], attention.head_dim // 2, 10000.0) + cos = cos.npu() + sin = sin.npu() + + with torch.inference_mode(): + expected = _reference_attention(attention, x, cos, sin) + actual = attention(x, cos, sin) + torch.npu.synchronize() + + _assert_vision_close( + actual, + expected, + max_abs_error=0.03125, + mean_abs_error=0.002, + min_cosine_similarity=0.9999, + ) + + +@pytest.mark.parametrize("grid", [(4, 4), (16, 24), (48, 72)]) +def test_fused_vision_block_matches_original_formula(vision_config, grid): + torch.manual_seed(17) + block = _new_bf16_module(DeepseekV4VisionBlock, vision_config) + num_tokens = grid[0] * grid[1] + x = torch.randn(num_tokens, VISION_DIM, dtype=torch.bfloat16, device="npu") + cos, sin = get_vision_cos_sin(grid[0], grid[1], VISION_DIM // VISION_HEADS // 2, 10000.0) + cos = cos.npu() + sin = sin.npu() + + with torch.inference_mode(): + expected = _reference_block(block, x, cos, sin) + actual = block(x.clone(), cos, sin) + torch.npu.synchronize() + + _assert_vision_close( + actual, + expected, + max_abs_error=0.0625, + mean_abs_error=0.004, + min_cosine_similarity=0.9999, + ) + + +def test_real_weights_full_vit_matches_original_formula(): + model_path = os.getenv("DEEPSEEK_V4_VISION_MODEL_PATH") + if not model_path: + pytest.skip("DEEPSEEK_V4_VISION_MODEL_PATH is not set") + assert model_path is not None + + model_root = Path(model_path) + config = SimpleNamespace(**json.loads((model_root / "config.json").read_text())) + checkpoint = load_file(model_root / "quant_model_weights-00078-of-00078.safetensors") + vision_weights = { + name.removeprefix("vision."): tensor for name, tensor in checkpoint.items() if name.startswith("vision.") + } + model = _new_bf16_module(DeepseekV4ViT, config) + model.load_state_dict(vision_weights, strict=True) + + torch.manual_seed(29) + grid = (16, 24) + patches = torch.randn( + grid[0] * grid[1], + 3, + config.vision_patch_size, + config.vision_patch_size, + dtype=torch.bfloat16, + device="npu", + ) + + with torch.inference_mode(): + expected = model.patch_embed(patches) + cos, sin = get_vision_cos_sin(grid[0], grid[1], model.rope_dim, model.rope_theta) + cos = cos.npu() + sin = sin.npu() + for block in model.blocks: + expected = _reference_block(block, expected, cos, sin) + expected = _reference_rms_norm(expected, model.norm.weight) + actual = model(patches, *grid) + torch.npu.synchronize() + + _assert_vision_close( + actual, + expected, + max_abs_error=0.125, + mean_abs_error=0.005, + min_cosine_similarity=0.9999, + ) diff --git a/tests/ut/attention/test_dsa_cp_o_proj_weight_switch.py b/tests/ut/attention/test_dsa_cp_o_proj_weight_switch.py index e8d4657503d6..97a62dd74006 100644 --- a/tests/ut/attention/test_dsa_cp_o_proj_weight_switch.py +++ b/tests/ut/attention/test_dsa_cp_o_proj_weight_switch.py @@ -63,6 +63,14 @@ def test_enablement_is_not_gated_by_hardware_family(self): tp_group = SimpleNamespace(world_size=2, rank_in_group=0) layer = self._OProj() with ( + patch( + "vllm_ascend.attention.context_parallel.dsa_cp.get_ascend_config", + return_value=SimpleNamespace(multistream_dsv4_dsa_overlap=True), + ), + patch( + "vllm_ascend.attention.context_parallel.dsa_cp.is_a5_bf16_kv_enabled", + return_value=False, + ), patch( "vllm_ascend.attention.context_parallel.dsa_cp.enable_dsa_cp_full_o_proj", return_value=True, @@ -90,9 +98,9 @@ def test_enablement_is_not_gated_by_hardware_family(self): n_local_groups=1, window_size=1, compress_ratio=1, - wq_a=object(), - wq_b=object(), - wkv=object(), + wq_a=layer, + wq_b=layer, + wkv=layer, q_norm=object(), kv_norm=object(), swa_cache_layer=SimpleNamespace(prefix="swa"), @@ -102,6 +110,10 @@ def test_enablement_is_not_gated_by_hardware_family(self): attn_sink=torch.empty(2), ) + self.assertTrue(impl.multistream_dsv4_dsa_overlap) + self.assertIs(impl.cv_wq_a.linear, layer) + self.assertIs(impl.cv_wkv.linear, layer) + self.assertIs(impl.cv_wq_b.linear, layer) self.assertTrue(impl.enable_dsa_cp_full_o_proj) profile.supports.assert_called_once_with(HardwareCapability.FP8_ATTENTION) diff --git a/tests/ut/attention/test_dsa_v1.py b/tests/ut/attention/test_dsa_v1.py index 1b7fdd9d689e..c2289a7bfe07 100644 --- a/tests/ut/attention/test_dsa_v1.py +++ b/tests/ut/attention/test_dsa_v1.py @@ -797,6 +797,10 @@ def test_dsa_cp_attention_waits_before_sas_consumer(compress_ratio: int, monkeyp "vllm_ascend.attention.context_parallel.dsa_cp.get_current_vllm_config", _make_vllm_config, ) + monkeypatch.setattr( + "vllm_ascend.attention.context_parallel.dsa_cp.get_ascend_config", + lambda: SimpleNamespace(multistream_dsv4_dsa_overlap=False), + ) impl = cast(AscendDSACPImpl, _make_impl(AscendDSACPImpl)) impl.compress_ratio = compress_ratio impl.compressor_overlap = False @@ -2119,3 +2123,98 @@ def fake_o_proj(o_proj_input: torch.Tensor, output_tensor: torch.Tensor) -> torc output[:local_num_actual_tokens], attention_output.view(local_num_actual_tokens, 2), ) + + +def test_v4_backend_keeps_post_projection_q_norm_by_default(): + assert _make_impl().apply_q_norm is True + + +@pytest.mark.parametrize("apply_q_norm", [False, True]) +@pytest.mark.parametrize("multistream", [False, True]) +@pytest.mark.parametrize("w8a8", [False, True]) +@torch.inference_mode() +def test_post_projection_q_norm_switch_preserves_lora_norm(monkeypatch, apply_q_norm, multistream, w8a8): + from contextlib import nullcontext + + from vllm_ascend.attention import dsa_v1 + + impl = _make_impl() + impl.apply_q_norm = apply_q_norm + impl.n_local_heads = 2 + impl.head_dim = 2 + impl.nope_head_dim = 1 + impl.rope_head_dim = 1 + impl.eps = 1e-6 + + class Norm(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor([1.0, 2.0, 0.5, 1.5])) + + def forward(self, x): + return x * torch.rsqrt(x.square().mean(-1, keepdim=True) + impl.eps) * self.weight + + class Wrapper: + _quant_method = None + _has_communication = False + + def __init__(self, linear): + self.linear = linear + + def quantize(self, x): + return x, None + + def matmul(self, x, scale): + return self.linear(x) + + def linear(weight): + layer = torch.nn.Linear(4, 4, bias=False) + layer.weight.copy_(torch.diag(torch.tensor(weight))) + layer.weight_scale = torch.ones(4) + return layer + + impl.wq_a = linear([1.0, 2.0, 3.0, 4.0]) + impl.wq_b = linear([4.0, 0.5, 2.0, 3.0]) + impl.wkv = torch.nn.Linear(4, 2, bias=False) + impl.wkv.weight_scale = torch.ones(2) + impl.q_norm = Norm() + impl.kv_norm = lambda x: x + impl.cv_wq_a, impl.cv_wq_b, impl.cv_wkv = map(Wrapper, (impl.wq_a, impl.wq_b, impl.wkv)) + stream = MagicMock() + monkeypatch.setattr(torch.npu, "current_stream", lambda: stream) + monkeypatch.setattr(dsa_v1, "dsv4_dsa_overlap_stream", lambda: stream) + monkeypatch.setattr(dsa_v1, "npu_stream_switch", lambda *args, **kwargs: nullcontext()) + monkeypatch.setattr(dsa_v1, "_is_w8a8_dynamic", lambda layer: w8a8) + monkeypatch.setattr(dsa_v1, "get_dsa_attn_kv_plan", lambda config: MagicMock()) + monkeypatch.setattr(torch.ops._C_ascend, "inplace_partial_rotary_mul", lambda *args, **kwargs: None, raising=False) + monkeypatch.setattr(dsa_v1.torch_npu, "npu_dynamic_quant", lambda x: (x, torch.ones(x.shape[0])), raising=False) + monkeypatch.setattr( + torch.ops._C_ascend, + "npu_rms_norm_dynamic_quant", + lambda x, weight, epsilon: (impl.q_norm(x), torch.ones(x.shape[0])), + raising=False, + ) + monkeypatch.setattr( + dsa_v1.torch_npu, + "npu_quant_matmul", + lambda x, weight, scale, **kwargs: torch.nn.functional.linear(x, weight, kwargs.get("bias")), + raising=False, + ) + rms = MagicMock(side_effect=lambda q, eps, norm: q * torch.rsqrt(q.square().mean(-1, keepdim=True) + eps)) + monkeypatch.setattr(dsa_v1.DeviceOperator, "apply_dsa_q_rms", rms) + + hidden = torch.tensor([[1.0, 2.0, 3.0, 4.0], [-2.0, 1.0, 4.0, 0.5]]) + expected_qr = impl.q_norm(impl.wq_a(hidden)) + projected = impl.wq_b(expected_qr).unflatten(-1, (2, 2)) + expected = ( + projected * torch.rsqrt(projected.square().mean(-1, keepdim=True) + impl.eps) if apply_q_norm else projected + ) + assert not torch.allclose(projected.square().mean(-1), torch.ones(2, 2)) + args = (hidden, None, None, torch.empty(1), torch.arange(2)) + if multistream: + q, qr, _, _ = impl._mla_prolog_multistream(*args) + else: + q, qr, _ = impl._mla_prolog_single_stream(*args, write_swa_cache=False) + torch.testing.assert_close(qr, expected_qr) + torch.testing.assert_close(q, expected) + assert rms.call_count == int(apply_q_norm) diff --git a/tests/ut/core/test_circular_buffer.py b/tests/ut/core/test_circular_buffer.py new file mode 100644 index 000000000000..5dfb5639958c --- /dev/null +++ b/tests/ut/core/test_circular_buffer.py @@ -0,0 +1,110 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch +from vllm.v1.core.block_pool import BlockPool +from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheGroupSpec, UniformTypeKVCacheSpecs + +from vllm_ascend.core.circular_buffer import ( + AscendCircularBufferManager, + AscendCircularBufferSpec, + is_circular_spec, + prefix_cacheable, +) +from vllm_ascend.patch.platform.patch_circular_buffer import _truncate_computed_blocks +from vllm_ascend.patch.platform.patch_kv_cache_coordinator import AscendHybridKVCacheCoordinator +from vllm_ascend.worker.block_table import BlockTable + + +def ring_spec(): + return AscendCircularBufferSpec(block_size=32, num_kv_heads=1, head_size=1024, dtype=torch.float32) + + +def test_ring_lifetime_reuse_and_external_tokens(): + pool = BlockPool(num_gpu_blocks=5, enable_caching=True, hash_block_size=32) + manager = AscendCircularBufferManager( + ring_spec(), + pool, + enable_caching=True, + kv_cache_group_id=0, + scheduler_block_size=128, + needs_kv_cache_zeroing=True, + ) + free = pool.get_num_free_blocks() + assert manager.get_num_blocks_to_allocate("a", 129, [], 0, 0, 129) == 1 + a = manager.allocate_new_blocks("a", 129, 129) + b = manager.allocate_new_blocks("b", 1, 1) + assert len(a) == len(b) == 1 and a[0].block_id != b[0].block_id + assert not a[0].is_null and pool.get_num_free_blocks() == free - 2 + for tokens in (130, 1024, 65536): + assert manager.get_num_blocks_to_allocate("a", tokens, [], tokens - 1, 0, tokens) == 0 + assert manager.allocate_new_blocks("a", tokens, tokens) == [] + manager.allocate_external_computed_blocks("a", 0, tokens) + manager.remove_skipped_blocks("a", tokens) + manager.cache_blocks( + SimpleNamespace(request_id="a"), + tokens, + replay_boundary=tokens - 1, + ) + assert manager.req_to_blocks["a"] == a + assert manager.take_new_block_ids() == [] + assert manager.get_num_common_prefix_blocks("a") == manager.get_num_skipped_tokens(65536) == 0 + manager.free("a") # Finish or preempt, then resume under a new allocation. + manager.free("b") + assert pool.get_num_free_blocks() == free + manager.allocate_external_computed_blocks("a", 0, 65536) + assert len(manager.req_to_blocks["a"]) == 1 + manager.free("a") + assert pool.get_num_free_blocks() == free + + +def test_uniform_properties_and_single_plane_size(): + spec = ring_spec() + uniform = UniformTypeKVCacheSpecs(block_size=32, kv_cache_specs={"s0": spec, "s1": spec}) + assert spec.page_size_bytes == 131072 + assert spec.max_memory_usage_bytes(None) == 131072 + assert uniform.max_num_blocks_per_req(None, 65536) == 1 + assert is_circular_spec(uniform) and not prefix_cacheable(uniform) + assert not uniform.prefix_cacheable + assert not prefix_cacheable(SimpleNamespace(participates_in_prefix_caching=False)) + + +def test_scratch_groups_do_not_reduce_prefix_hits_or_truncation(): + full = FullAttentionSpec(block_size=128, num_kv_heads=1, head_size=8, dtype=torch.float32) + scratch = ring_spec() + groups = [KVCacheGroupSpec(["kv"], full), KVCacheGroupSpec(["state"], scratch)] + coordinator = SimpleNamespace( + kv_cache_config=SimpleNamespace(kv_cache_groups=groups), + single_type_managers=[SimpleNamespace(), SimpleNamespace()], + eagle_group_ids=set(), + scheduler_block_size=128, + _get_effective_block_size=lambda spec: spec.block_size, + ) + AscendHybridKVCacheCoordinator.verify_and_split_kv_cache_groups(coordinator) + assert len(coordinator.attention_groups) == 1 + assert coordinator.attention_groups[0].group_ids == [0] + host = SimpleNamespace( + kv_cache_config=SimpleNamespace(kv_cache_groups=groups), + coordinator=SimpleNamespace( + single_type_managers=[SimpleNamespace(block_size=128), SimpleNamespace(block_size=32)] + ), + create_kv_cache_blocks=lambda blocks: blocks, + ) + assert _truncate_computed_blocks(host, SimpleNamespace(blocks=([1, 2], [])), 128) == ([1], []) + coordinator.kv_cache_config.kv_cache_groups = [groups[1]] + AscendHybridKVCacheCoordinator.verify_and_split_kv_cache_groups(coordinator) + assert AscendHybridKVCacheCoordinator.find_longest_cache_hit(coordinator, [], 128) == (([],), 0) + + +@pytest.mark.parametrize("draft", [False, True]) +def test_ring_bypasses_position_to_page_mapping(draft): + mapping = torch.zeros(8, dtype=torch.int64) + table = SimpleNamespace(is_circular_group=True, slot_mapping=SimpleNamespace(gpu=mapping)) + if draft: + BlockTable.compute_slot_mapping_draft(table, Mock(), Mock()) + else: + BlockTable.compute_slot_mapping(table, 1, Mock(), Mock()) + assert mapping.tolist() == [-1] * 8 diff --git a/tests/ut/distributed/eplb/test_state.py b/tests/ut/distributed/eplb/test_state.py index f38bfde993e3..79d17a637a24 100644 --- a/tests/ut/distributed/eplb/test_state.py +++ b/tests/ut/distributed/eplb/test_state.py @@ -32,8 +32,8 @@ def test_layer_state_builds_routing_table_and_preserves_captured_tensor( lambda: SimpleNamespace(rank_in_group=1), ) monkeypatch.setattr( - eplb_state._eplb_ops, - "build_expert_replica_routing_table", + eplb_state, + "_build_expert_replica_routing_table", build_routing_table, ) layer_state = AscendEplbLayerState() diff --git a/tests/ut/model_executor/warmup/test_deepseek_v41_triton_warmup.py b/tests/ut/model_executor/warmup/test_deepseek_v41_triton_warmup.py new file mode 100644 index 000000000000..890eca1c2a20 --- /dev/null +++ b/tests/ut/model_executor/warmup/test_deepseek_v41_triton_warmup.py @@ -0,0 +1,76 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch + +from vllm_ascend.model_executor.warmup import deepseek_v41_triton_warmup as warmup + + +@pytest.mark.parametrize( + "topk,cores,max_tokens,expected", + [ + (2048, 40, 4096, [1, 41]), + (2048, 40, 40, [1]), + (2048, 48, 4096, [1, 49]), + (128, 40, 4096, [1, 41, 81, 161, 321, 641]), + (128, 40, 128, [1, 41, 81]), + (7, 40, 1, [1]), + ], +) +def test_collect_indexer_warmup_token_counts(topk, cores, max_tokens, expected): + assert warmup.collect_indexer_warmup_token_counts(topk, cores, max_tokens) == expected + + +@pytest.mark.parametrize("model_type", ["deepseek_v4.1", "deepseek_v41", "deepseek_v4.1_text", "deepseek_v41_text"]) +def test_warmup_covers_tiles_and_active_compression_ratios(monkeypatch, model_type): + config = SimpleNamespace( + model_type=model_type, + num_hidden_layers=4, + compress_ratios=[0, 1, 2, 2, 8], + index_n_heads=64, + index_head_dim=128, + index_topk=2048, + ) + worker = SimpleNamespace( + model_config=SimpleNamespace(hf_text_config=config, dtype=torch.bfloat16), + scheduler_config=SimpleNamespace(max_num_batched_tokens=4096), + device=torch.device("cpu"), + ) + quantize, prepare = Mock(), Mock() + monkeypatch.setattr(warmup, "HAS_TRITON", True) + monkeypatch.setattr(warmup, "get_vectorcore_num", lambda: 40) + monkeypatch.setattr(warmup, "quantize_indexer_query", quantize) + monkeypatch.setattr(warmup, "prepare_indexer_indices", prepare) + warmup.deepseek_v41_triton_warmup(worker) + query = quantize.call_args.args[0] + assert query.shape == (1, 64, 128) + assert query.dtype == torch.bfloat16 + assert [(call.args[0].shape[0], call.args[2]) for call in prepare.call_args_list] == [ + (1, 1), + (1, 2), + (41, 1), + (41, 2), + ] + assert all(call.args[1].dtype == torch.int64 for call in prepare.call_args_list) + + +@pytest.mark.parametrize( + "has_triton,model_type,ratios", [(False, "deepseek_v41", [1]), (True, "other", [1]), (True, "deepseek_v41", [0])] +) +def test_warmup_skips_unused_indexer(monkeypatch, has_triton, model_type, ratios): + worker = SimpleNamespace( + model_config=SimpleNamespace( + hf_text_config=SimpleNamespace(model_type=model_type, num_hidden_layers=1, compress_ratios=ratios) + ) + ) + quantize, prepare = Mock(), Mock() + monkeypatch.setattr(warmup, "HAS_TRITON", has_triton) + monkeypatch.setattr(warmup, "quantize_indexer_query", quantize) + monkeypatch.setattr(warmup, "prepare_indexer_indices", prepare) + warmup.deepseek_v41_triton_warmup(worker) + quantize.assert_not_called() + prepare.assert_not_called() diff --git a/tests/ut/models/test_deepseek_v41_cache.py b/tests/ut/models/test_deepseek_v41_cache.py new file mode 100644 index 000000000000..5ce7033e6cbb --- /dev/null +++ b/tests/ut/models/test_deepseek_v41_cache.py @@ -0,0 +1,1667 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from dataclasses import replace +from types import SimpleNamespace +from typing import Any +from unittest.mock import Mock + +import pytest +import torch +import torch_npu +from vllm.config import CUDAGraphMode +from vllm.v1.core import kv_cache_utils +from vllm.v1.kv_cache_interface import KVCacheConfig + +from tests.deepseek_v41_cache_utils import allocate_cache_views, make_cache_config +from tests.deepseek_v41_reference import ( + build_v41_cache_specs, + compressor_ratio2_reference, + gather_cache_rows, + scatter_cache, + select_candidate_blocks, + select_index_topk, +) +from vllm_ascend.attention import dsa_v41 +from vllm_ascend.attention.dsa_v41 import ( + DeepseekV41CacheLayer, + DeepseekV41EagerAttentionImpl, + DeepseekV41MetadataBuilder, + compressed_slot_mapping, + pad_sparse_indices, + scatter_cache_sk, +) +from vllm_ascend.core.deepseek_v41 import ( + DeepseekV41DraftSWASpec, + DeepseekV41FullSpec, + DeepseekV41IndexerSpec, + DeepseekV41SWASpec, + allocate_cache_config, + cache_slots_from_groups, + group_cache_specs, + make_cache_groups, + plan_cache_slots, + pool_bytes_per_block, + request_blocks, + reshape_cache, +) +from vllm_ascend.models.deepseek_v41.compressor import DeepseekV41Compressor +from vllm_ascend.models.deepseek_v41.model import build_layer_plan +from vllm_ascend.worker.device_metadata import DeviceMetadataStage + + +@pytest.fixture(autouse=True) +def mock_npu_rms_norm(monkeypatch): + # Keep cache/state tests on CPU; operator accuracy is covered on NPU. + def rms_norm(x, gamma, epsilon=1e-6): + rstd = torch.rsqrt(x.float().square().mean(dim=-1, keepdim=True) + epsilon) + return (x.float() * rstd).to(x.dtype) * gamma, rstd + + monkeypatch.setattr(torch_npu, "npu_rms_norm", rms_norm) + + +@pytest.fixture +def config(): + # Deliberately small parameter dimensions; source topology matches the backbone. + return dict( + num_hidden_layers=40, + compress_ratios=[0, 0] + [2] * 18 + [1] * 20 + [0] * 3, + kv_source_layers=[2, 8, 14, 20], + index_source_layers=[2, 8, 14, 20, 24, 28, 32, 36], + candidate_source_layer=20, + candidate_topk_blocks=16, + candidate_block_size=8, + index_topk=8, + engram_layer_ids=[1, 14], + sliding_window=128, + head_dim=8, + index_head_dim=4, + hidden_size=16, + num_attention_heads=4, + index_n_heads=2, + q_lora_rank=8, + o_lora_rank=4, + o_groups=2, + rms_norm_eps=1e-6, + ) + + +@pytest.fixture +def runtime(config): + return SimpleNamespace( + model_config=SimpleNamespace(hf_text_config=config, enforce_eager=True), + cache_config=SimpleNamespace( + block_size=64, + enable_prefix_caching=False, + cache_dtype="auto", + num_gpu_blocks_override=None, + prefix_cache_retention_interval=None, + ), + compilation_config=SimpleNamespace(static_forward_context={}), + scheduler_config=SimpleNamespace(disable_hybrid_kv_cache_manager=False), + parallel_config=SimpleNamespace( + pipeline_parallel_size=1, + decode_context_parallel_size=1, + prefill_context_parallel_size=1, + tensor_parallel_size=1, + ), + speculative_config=None, + kv_transfer_config=None, + use_v2_model_runner=False, + ) + + +def collect_specs(runtime, prefix="model"): + return build_v41_cache_specs(runtime.model_config.hf_text_config, runtime, prefix) + + +def test_owner_counts_nested_config_and_source_resolution(config, runtime): + topology = build_layer_plan({"text_config": config}) + specs = collect_specs(runtime) + assert len(specs) == 51 + assert topology.kv_consumers(2) == tuple(range(2, 8)) + assert topology.kv_consumers(20) == tuple(range(20, 40)) + assert topology.layer(26).kv_source_layer == 20 + assert topology.layer(26).index_source_layer == 24 + assert specs["model.layers.20.self_attn.long_kv_cache"].storage_block_size == 64 + assert specs["model.layers.2.self_attn.long_kv_cache"].storage_block_size == 32 + assert "model.layers.20.self_attn.compressor.state_cache" not in specs + + +@pytest.mark.parametrize("block_size", [0, -2, 3, 63]) +def test_invalid_block_sizes(config, runtime, block_size): + runtime.cache_config.block_size = block_size + with pytest.raises(ValueError, match="multiple of two"): + build_v41_cache_specs(config, runtime) + + +def test_twelve_groups_share_four_layer_slots(config, runtime): + original = collect_specs(runtime) + uniform = group_cache_specs(original) + assert [len(g.kv_cache_specs) for g in uniform] == [8, 3] + [4] * 10 + assert [g.block_size for g in uniform] == [64, 32] + [64] * 10 + groups = make_cache_groups(uniform) + specs = {n: s for g in uniform for n, s in g.kv_cache_specs.items()} + assert all(s.page_size_padded is None for s in original.values()) + assert group_cache_specs(specs) == uniform # Replanning cannot accumulate padding. + assert group_cache_specs(dict(reversed(list(original.items())))) == uniform + slots = cache_slots_from_groups(groups) + blocks, allocations = allocate_cache_config(runtime, groups, pool_bytes_per_block(groups) * 10 + 1) + assert blocks == 10 and len(allocations) == 4 + cache_config = KVCacheConfig(num_blocks=blocks, kv_cache_tensors=allocations, kv_cache_groups=groups) + raw, caches = allocate_cache_views(cache_config) + assert len({t.data_ptr() for t in raw}) == 4 + assert sum(t.numel() for t in raw) == blocks * pool_bytes_per_block(groups) + assert set(caches) == set(original) + for backing, allocation, slot in zip(raw, allocations, slots): + assert allocation.offset == 0 and allocation.block_stride == slot.page_size_bytes + assert allocation.size == blocks * slot.page_size_bytes + assert allocation.layers == [p.name for p in slot.placements] + for placement in slot.placements: + spec = specs[placement.name] + cache = caches[placement.name] + views = cache if isinstance(cache, tuple) else (cache,) + assert views[0].shape == (blocks, spec.storage_block_size, 1, spec.head_size) + assert views[0].data_ptr() == backing.data_ptr() + placement.offset + assert all(v.stride(0) * v.element_size() == slot.page_size_bytes for v in views) + assert spec.page_size_bytes == placement.page_size_bytes + if isinstance(spec, DeepseekV41IndexerSpec): + key, scale = cache + assert key.dtype == torch.int8 and scale.dtype == torch.float16 + assert scale.data_ptr() - key.data_ptr() == spec.storage_block_size * spec.head_size + assert scale.shape == (blocks, spec.storage_block_size, 1, 1) + + +def test_production_layout_matches_design(config, runtime): + runtime.cache_config.block_size = 128 + specs = build_v41_cache_specs(dict(config, head_dim=512, index_head_dim=128), runtime) + groups = make_cache_groups(group_cache_specs(specs)) + assert len(groups) == 12 + assert [g.kv_cache_spec.page_size_bytes for g in groups] == [540928, 393216] + [540928] * 10 + assert pool_bytes_per_block(groups) == 540928 + slots = cache_slots_from_groups(groups) + assert [slot.page_size_bytes for slot in slots] == [131072] * 3 + [147712] + assert [len(slot.placements) for slot in slots] == [13, 13, 13, 12] + for i, slot in enumerate(slots): + assert slot.placements[1].offset == (65536 if i < 3 else 131072) + assert slot.placements[1].page_size_bytes == (65536 if i < 3 else 16640) + padded = {n: s for g in groups for n, s in g.kv_cache_spec.kv_cache_specs.items()} + swa_padding = [ + s.page_size_bytes - s.real_page_size_bytes for s in padded.values() if isinstance(s, DeepseekV41SWASpec) + ] + assert swa_padding.count(0) == 30 and swa_padding.count(16640) == 10 + blocks, tensors = allocate_cache_config(runtime, groups, 540928 * 3) + assert blocks == 3 + cache_config = KVCacheConfig(num_blocks=blocks, kv_cache_tensors=tensors, kv_cache_groups=groups) + _, caches = allocate_cache_views(cache_config) + assert sum(caches[n].is_contiguous() for n, s in padded.items() if isinstance(s, DeepseekV41SWASpec)) == 30 + + +def test_shared_slots_isolate_groups_and_recycled_ids(config, runtime): + groups = make_cache_groups(group_cache_specs(collect_specs(runtime))) + count = len(groups) + 1 + blocks, tensors = allocate_cache_config(runtime, groups, pool_bytes_per_block(groups) * count) + cfg = KVCacheConfig(num_blocks=blocks, kv_cache_tensors=tensors, kv_cache_groups=groups) + _, caches = allocate_cache_views(cfg) + expected = [] + for group_idx, group in enumerate(groups): + block_id = group_idx + 1 + for resource_idx, name in enumerate(group.layer_names): + cache = caches[name] + for plane_idx, view in enumerate(cache if isinstance(cache, tuple) else (cache,)): + slots = block_id * view.shape[1] + torch.arange(view.shape[1]) + value = torch.full( + (view.shape[1], view.shape[-1]), 1 + group_idx + resource_idx + plane_idx, dtype=view.dtype + ) + scatter_cache(view, slots, value) + expected.append((group_idx, view, slots, value)) + for _, view, slots, value in expected: + torch.testing.assert_close(gather_cache_rows(view, slots), value) + assert not view[0].any() + # Simulate release of group 0's ID and reassignment to a SWA group. + # The released full-context views are no longer valid; all other IDs remain intact. + for name in groups[2].layer_names: + caches[name][1].fill_(99) + for group_idx, view, slots, value in expected: + if group_idx != 0: + torch.testing.assert_close(gather_cache_rows(view, slots), value) + + +def test_slot_planner_rejects_missing_and_mismatched_pairs(runtime): + specs = collect_specs(runtime) + index_name = "model.layers.2.self_attn.indexer.k_cache" + with pytest.raises(ValueError, match="incompatible KV/index"): + plan_cache_slots({n: s for n, s in specs.items() if n != index_name}) + specs[index_name] = replace(specs[index_name], tokens_per_state=1) + with pytest.raises(ValueError, match="incompatible KV/index"): + plan_cache_slots(specs) + + +def test_merged_group_requires_common_logical_block_size(runtime): + specs = collect_specs(runtime) + for suffix in ("long_kv_cache", "indexer.k_cache"): + name = f"model.layers.20.self_attn.{suffix}" + specs[name] = replace(specs[name], block_size=128) + with pytest.raises(ValueError, match="Incompatible V4.1 resource layouts"): + group_cache_specs(specs) + + +@pytest.mark.parametrize("offset,stride,match", [(1, 257, "aligned"), (250, 256, "exceeds")]) +def test_invalid_view_layout_rejected(offset, stride, match): + spec = DeepseekV41FullSpec(block_size=16, num_kv_heads=1, head_size=4, dtype=torch.bfloat16) + with pytest.raises(ValueError, match=match): + reshape_cache( + torch.zeros(2 * stride, dtype=torch.uint8), spec, num_blocks=2, offset=offset, block_stride=stride + ) + + +def test_view_with_nonzero_backing_storage_offset(): + spec = DeepseekV41FullSpec(block_size=16, num_kv_heads=1, head_size=4, dtype=torch.bfloat16) + backing = torch.zeros(16 + 2 * 256, dtype=torch.uint8) + raw = backing[16:] + cache = reshape_cache(raw, spec, num_blocks=2, offset=32, block_stride=256) + cache[1].fill_(7) + assert cache.data_ptr() == backing.data_ptr() + 48 + torch.testing.assert_close(backing[304:432].view(torch.bfloat16), torch.full((64,), 7, dtype=torch.bfloat16)) + assert not backing[:48].any() + + +def test_view_accepts_latest_vllm_int8_backing_storage(): + spec = DeepseekV41FullSpec(block_size=16, num_kv_heads=1, head_size=4, dtype=torch.bfloat16) + raw = torch.zeros(2 * 256, dtype=torch.int8) + cache = reshape_cache(raw, spec, num_blocks=2, offset=32, block_stride=256) + assert cache.shape == (2, 16, 1, 4) + + +def test_request_accounting_counts_merged_full_context_once(runtime): + runtime.model_config.max_model_len = 1024 + runtime.max_in_flight_tokens = 128 + groups = make_cache_groups(group_cache_specs(collect_specs(runtime))) + bounded = sum( + max(s.max_memory_usage_bytes(runtime) // s.page_size_bytes for s in g.kv_cache_spec.kv_cache_specs.values()) + for g in groups[1:] + ) + assert request_blocks(runtime, groups) == 1024 // 64 + bounded + + +def test_mixed_layouts_rejected(config, runtime): + specs = collect_specs(runtime) + specs["foreign"] = object() + with pytest.raises(ValueError, match="foreign resources"): + group_cache_specs(specs) + + +def test_unsafe_override_rejected(config, runtime): + groups = make_cache_groups(group_cache_specs(collect_specs(runtime))) + runtime.cache_config.num_gpu_blocks_override = 100 + with pytest.raises(ValueError, match="unsafe block override"): + allocate_cache_config(runtime, groups, 1) + + +def test_safe_override_and_reserved_null_capacity(runtime): + groups = make_cache_groups(group_cache_specs(collect_specs(runtime))) + page = pool_bytes_per_block(groups) + runtime.cache_config.num_gpu_blocks_override = 3 + blocks, tensors = allocate_cache_config(runtime, groups, 5 * page + 1) + assert blocks == 3 and sum(t.size for t in tensors) == 3 * page + runtime.cache_config.num_gpu_blocks_override = None + with pytest.raises(ValueError, match="reserved null block"): + allocate_cache_config(runtime, groups, page) + + +def test_v0271_entrypoint_and_admission_use_slot_reservation(runtime): + runtime.model_config.max_model_len = 1024 + runtime.max_in_flight_tokens = 128 + groups = make_cache_groups(group_cache_specs(collect_specs(runtime))) + page = pool_bytes_per_block(groups) + config = kv_cache_utils.get_kv_cache_config_from_groups(runtime, groups, 100 * page) + assert config.num_blocks == 100 and len(config.kv_cache_tensors) == 4 + assert sum(t.size for t in config.kv_cache_tensors) == 100 * page + demand = request_blocks(runtime, groups) + assert kv_cache_utils._pool_bytes_per_block(groups) == page + assert kv_cache_utils._max_memory_usage_bytes_from_groups(runtime, groups) == (demand + 1) * page + assert kv_cache_utils.get_max_concurrency_for_kv_cache_config(runtime, config) == 99 / demand + + +def test_model_registration_and_binding(runtime): + specs = collect_specs(runtime, "language_model.model") + context = runtime.compilation_config.static_forward_context + modules = torch.nn.ModuleDict() + for index, (name, spec) in enumerate(specs.items()): + modules[str(index)] = DeepseekV41CacheLayer(runtime, name, spec) + assert len(context) == 51 + assert all(module.kv_cache[0].numel() == 0 for module in context.values()) + state = context["language_model.model.layers.2.self_attn.compressor.state_cache"] + assert state is context["language_model.model.layers.2.self_attn.compressor.state_cache"] + assert not state.spec.prefix_cacheable + assert state.spec.storage_block_size == 32 + owned_names = [name for name, module in modules.named_modules() if hasattr(module, "kv_cache")] + assert len(owned_names) == 51 + + +@pytest.mark.parametrize("feature", ["spec", "pp", "graph"]) +def test_unsupported_runtime_fails_before_registration(runtime, feature): + if feature == "spec": + runtime.speculative_config = object() + elif feature == "pp": + runtime.parallel_config.pipeline_parallel_size = 2 + else: + runtime.model_config.enforce_eager = False + with pytest.raises(NotImplementedError): + from vllm_ascend.core.deepseek_v41 import validate_cache_runtime + + validate_cache_runtime(runtime) + assert not runtime.compilation_config.static_forward_context + + +def test_prefix_cache_runtime_is_supported(runtime): + from vllm_ascend.core.deepseek_v41 import validate_cache_runtime + + runtime.cache_config.enable_prefix_caching = True + validate_cache_runtime(runtime) + + +def test_full_decode_only_runtime_is_supported(runtime): + from vllm_ascend.core.deepseek_v41 import validate_cache_runtime + + runtime.model_config.enforce_eager = False + runtime.compilation_config.cudagraph_mode = CUDAGraphMode.FULL_DECODE_ONLY + validate_cache_runtime(runtime) + + +@pytest.mark.parametrize("mode", [CUDAGraphMode.NONE, CUDAGraphMode.FULL_DECODE_ONLY]) +@pytest.mark.parametrize("count", [1, 15, 31]) +def test_dspark_runtime_preserves_ring_retention_limit(runtime, mode, count): + from vllm_ascend.core.deepseek_v41 import validate_cache_runtime + + runtime.speculative_config = SimpleNamespace(use_dspark=lambda: True, num_speculative_tokens=count) + runtime.compilation_config.cudagraph_mode = mode + validate_cache_runtime(runtime) + assert runtime.cache_config.cache_dtype == "bfloat16" + + +@pytest.mark.parametrize("count", [0, 32, 63]) +def test_dspark_rejects_verification_tail_that_cannot_fit_ring(runtime, count): + from vllm_ascend.core.deepseek_v41 import validate_cache_runtime + + runtime.speculative_config = SimpleNamespace(use_dspark=lambda: True, num_speculative_tokens=count) + with pytest.raises(ValueError, match="1..31"): + validate_cache_runtime(runtime) + + +def test_dspark_is_one_additional_group_in_existing_slots(runtime): + runtime.model_config.max_model_len = 4096 + runtime.max_in_flight_tokens = 256 + target = make_cache_config(17) + draft = make_cache_config(17, draft_layers=3) + groups = draft.kv_cache_groups + assert len(groups) == 13 and sum(len(g.layer_names) for g in groups) == 54 + assert groups[:12] == target.kv_cache_groups + assert groups[12].layer_names == [f"mtp.{i}.self_attn.swa_cache" for i in range(3)] + assert all(isinstance(s, DeepseekV41DraftSWASpec) for s in groups[12].kv_cache_spec.kv_cache_specs.values()) + assert len(draft.kv_cache_tensors) == 4 + assert [t.size for t in draft.kv_cache_tensors] == [t.size for t in target.kv_cache_tensors] + assert pool_bytes_per_block(groups) == 540928 + added = request_blocks(runtime, groups) - request_blocks(runtime, target.kv_cache_groups) + spec = next(iter(groups[12].kv_cache_spec.kv_cache_specs.values())) + assert added == (spec.max_memory_usage_bytes(runtime) + spec.page_size_bytes - 1) // spec.page_size_bytes + padded = {n: s for g in groups for n, s in g.kv_cache_spec.kv_cache_specs.items()} + assert plan_cache_slots(padded) == plan_cache_slots(dict(reversed(list(padded.items())))) + backings, views = allocate_cache_views(draft) + assert sum(b.numel() for b in backings) == 17 * 540928 + for stage in range(3): + name = f"mtp.{stage}.self_attn.swa_cache" + assert name in draft.kv_cache_tensors[stage].layers + assert views[name].data_ptr() == backings[stage].data_ptr() + assert views[name].shape == (17, 128, 1, 512) + assert views[name].stride() == (65536, 512, 512, 1) + + +def test_dspark_slots_isolate_groups_and_reuse_released_ids(): + cfg = make_cache_config(17, draft_layers=3) + _, views = allocate_cache_views(cfg) + # Each group owns a different live global ID, including the draft group. + for gid, group in enumerate(cfg.kv_cache_groups): + for name in group.layer_names: + planes = views[name] if isinstance(views[name], tuple) else (views[name],) + for plane in planes: + plane[gid + 1].fill_(gid + 1) + for gid, group in enumerate(cfg.kv_cache_groups): + for name in group.layer_names: + planes = views[name] if isinstance(views[name], tuple) else (views[name],) + for plane in planes: + assert (plane[gid + 1] == gid + 1).all() + assert (plane[0] == 0).all() + # After target G0 releases ID 1, G12 may use it without touching live IDs. + for stage in range(3): + view = views[f"mtp.{stage}.self_attn.swa_cache"] + view[1].fill_(7) + assert (view[13] == 13).all() + assert (view[0] == 0).all() + + +@pytest.mark.parametrize("change", [{"head_size": 1024}, {"block_size": 256}, {"sliding_window": 256}]) +def test_dspark_geometry_cannot_expand_existing_slots(change): + cfg = make_cache_config(3, draft_layers=3) + specs = {n: s for g in cfg.kv_cache_groups for n, s in g.kv_cache_spec.kv_cache_specs.items()} + specs["mtp.0.self_attn.swa_cache"] = replace(specs["mtp.0.self_attn.swa_cache"], **change) + with pytest.raises(ValueError, match="geometry"): + plan_cache_slots(specs) + + +@pytest.mark.parametrize("count", [1, 2, 4]) +def test_dspark_requires_three_slot_owners(count): + with pytest.raises(ValueError, match="three ordered"): + make_cache_config(3, draft_layers=count) + + +@pytest.mark.parametrize("full_graph_mode", [False, True]) +def test_c2_builder_keeps_fixed_rows_for_mixed_parity_and_padding(config, runtime, full_graph_mode): + spec = collect_specs(runtime)["model.layers.2.self_attn.compressor.state_cache"] + builder = DeepseekV41MetadataBuilder(spec, [], runtime, torch.device("cpu")) + common = SimpleNamespace( + slot_mapping=torch.tensor([10, 11, -1]), + block_table_tensor=torch.tensor([[1], [2], [0]]), + query_start_loc=torch.tensor([0, 1, 2, 3]), + query_start_loc_cpu=torch.tensor([0, 1, 2, 3]), + seq_lens=torch.tensor([3, 4, 9]), + seq_lens_cpu=torch.tensor([3, 4, 9]), + positions=torch.tensor([2, 3, 0]), + num_reqs=3, + num_actual_tokens=2, + num_input_tokens=3, + max_query_len=1, + max_seq_len=9, + is_prefilling=torch.tensor([False, False, False]), + ) + + metadata = builder.build(0, common, num_actual_reqs=2, full_graph_mode=full_graph_mode) + + assert metadata.num_actual_reqs == 2 + assert metadata.seq_lens.tolist() == [3, 4, 0] + assert metadata.c2_complete_mask.tolist() == [False, True, False] + assert metadata.c2_ring_metadata.tolist() == [[2, 3, 0], [1, 1, 0], [0, 1, 2], [0, 1, 2], [1, 2, 0]] + assert metadata.slot_mapping.tolist() == [-1, -1, -1] + assert metadata.c2_source_positions.tolist() == [0, 2, 0] + assert metadata.c2_metadata_group_id == id(builder._c2_complete_mask) + pointer = metadata.c2_ring_metadata.data_ptr() + common.seq_lens = torch.tensor([4, 5, 8]) + common.seq_lens_cpu = common.seq_lens + common.positions = torch.tensor([3, 4, 0]) + common.block_table_tensor = torch.tensor([[7], [3], [0]]) + replay = builder.build(0, common, num_actual_reqs=2, full_graph_mode=full_graph_mode) + assert replay.c2_ring_metadata.data_ptr() == pointer + assert replay.c2_complete_mask.tolist() == [True, False, False] + assert replay.c2_ring_metadata[4].tolist() == [7, 3, 0] + idle = builder.build(0, common, num_actual_reqs=2, skip_ring_state_update=True) + assert idle.c2_ring_metadata[1].tolist() == [0, 0, 0] + assert idle.c2_ring_metadata[4].tolist() == [0, 0, 0] + assert not idle.c2_complete_mask.any() + + +@pytest.mark.parametrize("full_graph_mode", [False, True]) +@pytest.mark.parametrize("index_first", [False, True]) +@pytest.mark.parametrize( + "num_actual_reqs,num_actual_tokens,stored_rows", + [(3, 5, [1, 2, 3, 4]), (2, 5, [1, 2]), (3, 3, [1, 2]), (0, 5, []), (3, 0, [])], +) +def test_c2_builder_prepares_shared_store_mask( + runtime, full_graph_mode, index_first, num_actual_reqs, num_actual_tokens, stored_rows +): + specs = collect_specs(runtime) + names = ["model.layers.2.self_attn.long_kv_cache", "model.layers.2.self_attn.indexer.k_cache"] + if index_first: + names.reverse() + builders = [DeepseekV41MetadataBuilder(specs[name], [name], runtime, torch.device("cpu")) for name in names] + common = SimpleNamespace( + slot_mapping=torch.tensor([10, 11, 15, 129, 131]), + block_table_tensor=torch.tensor([[1], [2], [4]]), + query_start_loc=torch.tensor([0, 2, 3, 5]), + query_start_loc_cpu=torch.tensor([0, 2, 3, 5]), + seq_lens=torch.tensor([4, 6, 10]), + seq_lens_cpu=torch.tensor([4, 6, 10]), + positions=torch.tensor([2, 3, 5, 7, 9]), + num_reqs=3, + num_actual_tokens=num_actual_tokens, + num_input_tokens=5, + max_query_len=2, + max_seq_len=10, + is_prefilling=torch.tensor([True, False, True]), + ) + original_slots = common.slot_mapping.clone() + shared: dict[str, Any] = {} + metadata = [ + builder.build( + 0, common, num_actual_reqs=num_actual_reqs, full_graph_mode=full_graph_mode, common_v41_metadata=shared + ) + for builder in builders + ] + expected = torch.full((5, 2), -1, dtype=torch.int32) + coordinates = torch.tensor([[-1, -1], [0, 5], [0, 7], [2, 0], [2, 1]], dtype=torch.int32) + expected[stored_rows] = coordinates[stored_rows] + pointer = metadata[0].slot_mapping.data_ptr() + for result in metadata: + assert result.slot_mapping.data_ptr() == pointer + torch.testing.assert_close(result.slot_mapping, expected) + torch.testing.assert_close(common.slot_mapping, original_slots) + + # New metadata must update the captured address and retain position parity + # even when a padded slot happens to contain a valid physical coordinate. + common.positions = torch.tensor([3, 4, 6, 8, 10]) + common.slot_mapping = torch.tensor([11, 11, 15, 129, 131]) + common.num_actual_tokens = 5 + replay = builders[0].build(0, common, full_graph_mode=full_graph_mode) + assert replay.slot_mapping.data_ptr() == pointer + assert replay.slot_mapping.tolist() == [[0, 5], [-1, -1], [-1, -1], [-1, -1], [-1, -1]] + idle = builders[0].build(0, common, skip_ring_state_update=True) + assert idle.slot_mapping.data_ptr() == pointer + assert idle.slot_mapping.tolist() == [[-1, -1]] * 5 + + +def test_scatter_cache_redirects_invalid_rows_to_null_row(): + cache = torch.full((1, 8, 1, 2), -3.0) + values = torch.tensor([[9.0, 9.0], [7.0, 8.0]]) + + scatter_cache(cache, torch.tensor([-1, 3]), values) + + assert cache[0, 0, 0].tolist() == [0.0, 0.0] + assert cache[0, 3, 0].tolist() == [7.0, 8.0] + + +def test_scatter_cache_sk_consumes_prepared_coordinates_and_preserves_stride( + monkeypatch, +): + backing = torch.zeros(3 * 128, dtype=torch.uint8) + cache = torch.as_strided( + backing.view(torch.float32), + size=(3, 4, 1, 2), + stride=(32, 2, 2, 1), + ) + values = torch.tensor([[9.0, 9.0], [7.0, 8.0]]) + indices = torch.tensor([[-1, -1], [1, 3]], dtype=torch.int32) + calls = [] + + def scatter(var, indices, updates): + calls.append((var, indices, updates)) + + monkeypatch.setattr( + torch.ops._C_ascend, + "npu_scatter_nd_update_sk", + scatter, + raising=False, + ) + scatter_cache_sk(cache, indices, values) + + var, actual_indices, updates = calls[0] + assert var.shape == (3, 4, 2) + assert var.stride() == (32, 2, 1) + assert actual_indices.data_ptr() == indices.data_ptr() + torch.testing.assert_close(actual_indices, indices) + assert updates.tolist() == [[9.0, 9.0], [7.0, 8.0]] + + +def test_compression_slot_mapping(): + slots = torch.tensor([-1, 0, 1, 62, 63, 320, 321, 383]) + assert compressed_slot_mapping(slots, 2).tolist() == [-1, -1, 0, -1, 31, -1, 160, 191] + assert torch.equal(compressed_slot_mapping(slots, 1), slots) + + +@pytest.mark.parametrize("compress_ratio", [0, 1, 2]) +def test_supported_ratios_route_to_native_sparse_flash_mla(monkeypatch, compress_ratio): + impl = DeepseekV41EagerAttentionImpl.__new__(DeepseekV41EagerAttentionImpl) + impl.role = SimpleNamespace( + compress_ratio=compress_ratio, + has_long_context=compress_ratio > 0, + ) + impl.long_kv_source_prefix = "source" + impl.topology = SimpleNamespace(index_topk=512) + source_cache = object() + monkeypatch.setattr( + "vllm_ascend.attention.dsa_v41.get_forward_context", + lambda: SimpleNamespace(no_compile_layers={"source": SimpleNamespace(kv_cache=[source_cache])}), + ) + expected = object() + + def native(*args, **kwargs): + return expected + + monkeypatch.setattr(impl, "_native_attention", native) + + actual = impl._attention( + SimpleNamespace(), + object(), + SimpleNamespace(swa=object(), attention=object()), + object() if compress_ratio else None, + ) + + assert actual is expected + + +def test_candidate_blocks_pin_partial_tail_and_drop_unreachable_blocks(): + scores = torch.tensor([[9.0, 8.0, 7.0, 6.0, 5.0, 4.0, -torch.inf, -torch.inf]]) + # With two candidate blocks, the best old block and the partially filled + # newest block must survive. The unreachable final block must not. + mask = select_candidate_blocks(scores, torch.tensor([[6]]), topk_blocks=2, block_size=2) + assert mask.tolist() == [[True, True, False, False, True, True, False, False]] + + +def test_index_topk_is_chronological_and_marks_unreachable_slots(): + scores = torch.tensor([[1.0, 7.0, 3.0, -torch.inf, -torch.inf]]) + selected = select_index_topk(scores, torch.tensor([[3]]), index_topk=4) + assert selected.tolist() == [[0, 1, 2, -1]] + + +def test_sparse_indices_are_padded_for_native_mla(): + selected = torch.tensor([[2, 7], [1, -1]], dtype=torch.int32) + padded = pad_sparse_indices(selected, 4) + assert padded.shape == (2, 1, 4) + assert padded.tolist() == [[[2, 7, -1, -1]], [[1, -1, -1, -1]]] + + +def test_state_metadata_disables_ordinary_token_slots(config, runtime): + specs = collect_specs(runtime) + spec = specs["model.layers.2.self_attn.compressor.state_cache"] + builder = DeepseekV41MetadataBuilder(spec, [], runtime, torch.device("cpu")) + slots = torch.tensor([7 * 16 + 15, 3 * 16, -1]) + common = SimpleNamespace( + slot_mapping=slots, + positions=None, + block_table_tensor=torch.tensor([[7, 3]]), + query_start_loc=torch.tensor([0, 2]), + query_start_loc_cpu=torch.tensor([0, 2]), + seq_lens=torch.tensor([17]), + seq_lens_cpu=torch.tensor([17]), + num_reqs=1, + num_actual_tokens=2, + num_input_tokens=2, + max_query_len=2, + max_seq_len=17, + is_prefilling=torch.tensor([True]), + ) + metadata = builder.build(0, common) + assert metadata.is_compressor_state + assert (metadata.slot_mapping == -1).all() + assert metadata.compress_ratio == 1 + assert metadata.storage_block_size == 32 + assert metadata.max_query_len == 2 + assert metadata.max_seq_len == 17 + assert metadata.query_start_loc.tolist() == [0, 2] + assert metadata.cache_seq_lens.tolist() == [17] + assert metadata.cache_seq_lens is metadata.seq_lens + assert metadata.num_prefills == 1 + assert metadata.num_prefill_tokens == 2 + + +def test_slot_mapping_is_shared_per_compatible_cache_group(config, runtime): + specs = collect_specs(runtime) + common = SimpleNamespace( + slot_mapping=torch.tensor([1, 2, 65, -1]), + positions=None, + block_table_tensor=torch.tensor([[5, 7]]), + query_start_loc=torch.tensor([0, 4]), + query_start_loc_cpu=torch.tensor([0, 4]), + seq_lens=torch.tensor([4]), + seq_lens_cpu=torch.tensor([4]), + num_reqs=1, + num_actual_tokens=3, + num_input_tokens=4, + max_query_len=4, + max_seq_len=4, + is_prefilling=torch.tensor([True]), + ) + full_group_metadata: dict[str, Any] = {} + long_metadata = DeepseekV41MetadataBuilder( + specs["model.layers.2.self_attn.long_kv_cache"], + ["model.layers.2.self_attn.long_kv_cache"], + runtime, + torch.device("cpu"), + ).build(0, common, common_v41_metadata=full_group_metadata) + index_metadata = DeepseekV41MetadataBuilder( + specs["model.layers.2.self_attn.indexer.k_cache"], + ["model.layers.2.self_attn.indexer.k_cache"], + runtime, + torch.device("cpu"), + ).build(0, common, common_v41_metadata=full_group_metadata) + + assert long_metadata.slot_mapping.data_ptr() == index_metadata.slot_mapping.data_ptr() + assert long_metadata.slot_mapping.tolist() == [ + [0, 0], + [-1, -1], + [1, 0], + [-1, -1], + ] + + # The SWA builder receives a different per-group publication dictionary, + # so it owns an independent mapping computed from that group's flat slots. + swa_metadata = DeepseekV41MetadataBuilder( + specs["model.layers.3.self_attn.swa_cache"], + ["model.layers.3.self_attn.swa_cache"], + runtime, + torch.device("cpu"), + ).build(0, common, common_v41_metadata={}) + assert swa_metadata.slot_mapping.data_ptr() != long_metadata.slot_mapping.data_ptr() + assert swa_metadata.slot_mapping.tolist() == [ + [0, 1], + [0, 2], + [1, 1], + [-1, -1], + ] + + +def test_compressed_metadata_exposes_original_and_cache_coordinates(config, runtime): + specs = collect_specs(runtime) + spec = specs["model.layers.2.self_attn.long_kv_cache"] + builder = DeepseekV41MetadataBuilder(spec, [], runtime, torch.device("cpu")) + # Request 0 starts halfway through a compression pair; request 1 ends + # with an incomplete pair. Only completed pairs become cache rows. + common = SimpleNamespace( + slot_mapping=torch.tensor([1, 2, 3, 65, 66]), + positions=torch.tensor([1, 2, 3, 1, 2]), + block_table_tensor=torch.tensor([[5, 7], [9, 0]]), + query_start_loc=torch.tensor([0, 3, 5]), + query_start_loc_cpu=torch.tensor([0, 3, 5]), + seq_lens=torch.tensor([4, 3]), + seq_lens_cpu=torch.tensor([4, 3]), + num_reqs=2, + num_actual_tokens=5, + num_input_tokens=5, + max_query_len=3, + max_seq_len=4, + is_prefilling=torch.tensor([True, False]), + ) + metadata = builder.build(0, common) + assert metadata.seq_lens.tolist() == [4, 3] + assert metadata.query_start_loc.tolist() == [0, 3, 5] + assert metadata.cache_seq_lens.tolist() == [2, 1] + assert metadata.cmp_residual.tolist() == [0, 1] + assert metadata.max_cache_seq_len == 2 + assert metadata.slot_mapping.tolist() == [ + [0, 0], + [-1, -1], + [0, 1], + [1, 0], + [-1, -1], + ] + assert metadata.num_prefills == 1 + assert metadata.num_prefill_tokens == 3 + assert metadata.num_decodes == 1 + assert metadata.num_decode_tokens == 2 + + +@pytest.mark.parametrize("deferred", [False, True]) +@pytest.mark.parametrize("query_len", [1, 3]) +def test_batch_metadata_reuses_work_and_keeps_group_slots_separate(runtime, monkeypatch, deferred, query_len): + groups = make_cache_config(17).kv_cache_groups + builders = [] + for group in groups: + layers_by_spec: dict[Any, list[str]] = {} + for name, spec in group.kv_cache_spec.kv_cache_specs.items(): + layers_by_spec.setdefault(spec, []).append(name) + builders.append( + [ + DeepseekV41MetadataBuilder(spec, names, runtime, torch.device("cpu")) + for spec, names in layers_by_spec.items() + ] + ) + assert sum(map(len, builders)) == 25 + counts = Mock(wraps=dsa_v41._request_counts) + compressed_slots = Mock(wraps=dsa_v41.compressed_slot_mapping) + rope = Mock(side_effect=lambda positions, **kwargs: (positions.float().clone(), -positions.float())) + monkeypatch.setattr(dsa_v41, "_request_counts", counts) + monkeypatch.setattr(dsa_v41, "compressed_slot_mapping", compressed_slots) + monkeypatch.setattr(dsa_v41, "get_cos_and_sin_dsa", rope) + + def native_metadata(*args, **kwargs): + lengths = kwargs.get("seqused_ori_kv", kwargs.get("seqused_k")) + return torch.full( + (dsa_v41.V41_METADATA_BUFFER_SIZE,), int(lengths.sum()) + kwargs["cmp_ratio"], dtype=torch.int32 + ) + + smla = Mock(side_effect=native_metadata) + qli = Mock(side_effect=native_metadata) + monkeypatch.setattr(torch.ops._C_ascend, "npu_sparse_flash_mla_metadata", smla, raising=False) + monkeypatch.setattr(torch.ops._C_ascend, "npu_quant_lightning_indexer_v2_metadata", qli, raising=False) + for group_builders in builders: + for builder in group_builders: + # Only native operator dispatch is mocked; all coordinates use CPU torch. + builder._supports_device_ops = not isinstance(builder.kv_cache_spec, dsa_v41.DeepseekV41CompressorStateSpec) + builder._device_metadata_enabled = deferred + + def build_batch(lengths, block_offset=0, idle=False): + batch_shared: dict[str, Any] = {} + results, tasks = [], [] + positions = torch.tensor( + [*range(lengths[0] - query_len, lengths[0]), *range(lengths[1] - query_len, lengths[1]), 0] + ) + query_start_loc = torch.tensor([0, query_len, 2 * query_len, 2 * query_len + 1], dtype=torch.int32) + for gid, group_builders in enumerate(builders): + block_ids = torch.tensor([gid + 1 + block_offset, gid + 2 + block_offset, 0], dtype=torch.int32) + flat_slots = torch.cat( + ( + block_ids[0] * 128 + positions[:query_len], + block_ids[1] * 128 + positions[query_len : 2 * query_len], + torch.tensor([-1]), + ) + ) + common = SimpleNamespace( + query_start_loc=query_start_loc, + query_start_loc_cpu=query_start_loc, + seq_lens=torch.tensor([*lengths, 999], dtype=torch.int32), + seq_lens_cpu=None, + _seq_lens_cpu=torch.tensor([*lengths, 999], dtype=torch.int32), + positions=positions, + slot_mapping=flat_slots, + block_table_tensor=block_ids[:, None], + num_reqs=3, + num_input_tokens=len(positions), + num_actual_tokens=2 * query_len, + max_query_len=query_len, + max_seq_len=max(lengths), + is_prefilling=torch.tensor([query_len > 1, query_len > 1, False]), + ) + group_shared: dict[str, Any] = {} + group_results = [] + for builder in group_builders: + metadata = builder.build( + 0, + common, + num_actual_reqs=2, + skip_ring_state_update=idle, + full_graph_mode=query_len == 1, + common_v41_metadata=group_shared, + common_v41_batch_metadata=batch_shared, + ) + group_results.append(metadata) + tasks.extend(builder.take_device_metadata_tasks()) + results.append(group_results) + # Building later groups must never modify an earlier group's slots. + if gid >= 2: + expected = torch.stack((flat_slots.clamp_min(0) // 128, flat_slots.clamp_min(0) % 128), dim=1).int() + expected[-1] = -1 + torch.testing.assert_close(group_results[0].slot_mapping, expected) + for task in sorted(tasks, key=lambda task: task.stage): + task.run() + if deferred: + assert [task.stage for task in tasks].count(DeviceMetadataStage.ATTENTION) == 3 + assert [task.stage for task in tasks].count(DeviceMetadataStage.INDEXER) == 2 + assert [task.stage for task in tasks].count(DeviceMetadataStage.COMPRESSOR) == 1 + else: + assert not tasks + return results, tuple((task.stage, task.group_id) for task in tasks) + + previous_pointers = previous_frontiers = None + for iteration, (lengths, idle) in enumerate([([7, 8], False), ([10, 11], False), ([10, 11], True)]): + results, frontiers = build_batch(lengths, block_offset=iteration, idle=idle) + all_metadata = [metadata for group_results in results for metadata in group_results] + assert counts.call_count == iteration + 1 + assert rope.call_count == iteration + 1 + assert compressed_slots.call_count == iteration + 1 + assert smla.call_count == 3 * (iteration + 1) + assert qli.call_count == 2 * (iteration + 1) + for metadata in all_metadata: + assert metadata.seq_lens.tolist() == [*lengths, 0] + assert metadata.seq_lens is all_metadata[0].seq_lens + c2 = [metadata for metadata in results[0] if metadata.compress_ratio == 2] + assert c2[0].cache_seq_lens is c2[1].cache_seq_lens + assert c2[0].cmp_residual is c2[1].cmp_residual + assert c2[0].cache_seq_lens.tolist() == [n // 2 for n in lengths] + [0] + assert c2[0].cmp_residual.tolist() == [n % 2 for n in lengths] + [0] + assert c2[0].max_cache_seq_len == max(lengths) // 2 + for metadata in results[0]: + if metadata.compress_ratio == 1: + assert metadata.cache_seq_lens is metadata.seq_lens + swa = [metadata for group_results in results[2:] for metadata in group_results] + assert all(metadata.cos is swa[0].cos and metadata.sin is swa[0].sin for metadata in swa) + assert all(metadata.smla_metadata is swa[0].smla_metadata for metadata in swa) + assert int(swa[0].smla_metadata[0]) == sum(lengths) + assert len({group_results[0].slot_mapping.data_ptr() for group_results in results[2:]}) == 10 + for group_results in results[2:]: + assert group_results[0].slot_mapping is group_results[1].slot_mapping + if idle: + assert (c2[0].slot_mapping == -1).all() + assert (results[1][0].c2_ring_metadata[1] == 0).all() + pointers = tuple( + ( + metadata.seq_lens.data_ptr(), + metadata.cache_seq_lens.data_ptr(), + metadata.slot_mapping.data_ptr(), + None if metadata.smla_metadata is None else metadata.smla_metadata.data_ptr(), + ) + for metadata in all_metadata + ) + if previous_pointers is not None: + assert pointers == previous_pointers + assert frontiers == previous_frontiers + previous_pointers, previous_frontiers = pointers, frontiers + + +@pytest.mark.parametrize("end", [127, 128, 129, 255, 256, 257]) +def test_merged_metadata_preserves_nonconsecutive_block_ids(runtime, end): + runtime.cache_config.block_size = 128 + group = group_cache_specs(collect_specs(runtime))[0] + table = torch.tensor([[7, 19, 3]], dtype=torch.int32) + positions = torch.arange(end - 3, end) + original_slots = table[0, positions // 128] * 128 + positions % 128 + common = SimpleNamespace( + slot_mapping=original_slots, + positions=positions, + block_table_tensor=table, + query_start_loc=torch.tensor([0, 3]), + query_start_loc_cpu=torch.tensor([0, 3]), + seq_lens=torch.tensor([end]), + seq_lens_cpu=torch.tensor([end]), + num_reqs=1, + num_actual_tokens=3, + num_input_tokens=3, + max_query_len=3, + max_seq_len=end, + is_prefilling=torch.tensor([True]), + ) + for name, spec in group.kv_cache_specs.items(): + metadata = DeepseekV41MetadataBuilder(spec, [name], runtime, torch.device("cpu")).build(0, common) + ratio = spec.tokens_per_state + rows = 128 // ratio + expected = table[0, positions // 128] * rows + (positions % 128) // ratio + expected = torch.where((positions + 1) % ratio == 0, expected, -1) + valid = expected >= 0 + physical = expected.clamp_min(0) + expected_2d = torch.stack( + (physical // spec.storage_block_size, physical % spec.storage_block_size), + dim=-1, + ).to(torch.int32) + expected_2d[~valid] = -1 + torch.testing.assert_close(metadata.slot_mapping, expected_2d) + assert metadata.logical_block_size == 128 + assert metadata.storage_block_size == rows + assert metadata.cache_seq_lens.tolist() == [end // ratio] + torch.testing.assert_close(metadata.block_table, table) + torch.testing.assert_close(common.slot_mapping, original_slots) + + +@pytest.mark.parametrize("end", [15, 16, 17, 31, 32, 33, 127, 128, 129, 255, 256, 257]) +def test_state_boundary_mapping_with_padded_pages(runtime, end): + group = group_cache_specs(collect_specs(runtime))[1] + positions = torch.arange(end - 2, end) + common = SimpleNamespace( + slot_mapping=torch.full((2,), -1), + block_table_tensor=torch.tensor([[7]], dtype=torch.int32), + positions=positions, + query_start_loc=torch.tensor([0, 2]), + query_start_loc_cpu=torch.tensor([0, 2]), + seq_lens=torch.tensor([end]), + seq_lens_cpu=torch.tensor([end]), + num_reqs=1, + num_actual_tokens=2, + num_input_tokens=2, + max_query_len=2, + max_seq_len=end, + is_prefilling=torch.tensor([True]), + ) + spec = next(iter(group.kv_cache_specs.values())) + metadata = DeepseekV41MetadataBuilder(spec, [], runtime, torch.device("cpu")).build(0, common) + assert metadata.slot_mapping.tolist() == [-1, -1] + assert metadata.storage_block_size == metadata.logical_block_size == 32 + assert metadata.c2_ring_metadata[:, 0].tolist() == [end - 2, 2, 0, 0, 7] + assert metadata.c2_source_positions.tolist() == [int(p - 1) if p % 2 else 0 for p in positions] + + +def test_actual_attention_parameter_ownership(config, runtime): + topology = build_layer_plan(config) + assert not topology.layer(0).has_long_context + assert topology.layer(2).is_kv_source and topology.layer(2).is_index_source + assert topology.layer(20).is_kv_source and topology.layer(20).compress_ratio == 1 + assert topology.layer(24).is_index_source and not topology.layer(24).is_kv_source + assert not topology.layer(26).is_index_source + assert topology.layer(26).kv_source_layer == 20 + + +@pytest.mark.parametrize("chunks", [(1, 1, 1, 2, 2), (3, 4), (2, 2, 3), (7,)]) +@torch.inference_mode() +def test_compressor_chunk_boundary_matches_vector_reference(config, chunks): + torch.manual_seed(7) + compressor = DeepseekV41Compressor(config, 2) + x = torch.randn(7, 16, dtype=torch.bfloat16) + kv = compressor.wkv(x.float())[:6].reshape(3, 2, 8) + gate = compressor.wgate(x.float())[:6].reshape(3, 2, 8) + expected = compressor.norm((kv * gate.softmax(dim=1)).sum(dim=1).to(x.dtype)) + state = torch.full((6, 32, 16), float("nan"), dtype=torch.float32) + block_table = [4] + actual = [] + start = 0 + for size in chunks: + actual.append(compressor_ratio2_reference(compressor, x[start : start + size], start, state, block_table)) + start += size + torch.testing.assert_close(torch.cat(actual), expected) + torch.testing.assert_close(state[4, 6, :8], compressor.wkv(x[-1:].float())[0]) + + +@pytest.mark.parametrize("num_tokens", [1, 2, 3, 5]) +@pytest.mark.parametrize("start", [0, 1]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) +def test_ring_source_reuses_prepared_store_coordinates(monkeypatch, num_tokens, start, dtype): + from vllm_ascend.attention import dsa_v41 + + positions = torch.arange(start, start + num_tokens) + completed = positions.remainder(2) == 1 + slots = torch.tensor([[7, 63], [19, 0], [19, 1], [3, 0], [3, 1]], dtype=torch.int32)[:num_tokens] + slots[~completed] = -1 + original_slots = slots.clone() + rope = torch.zeros(num_tokens, 1, 2) + state = SimpleNamespace( + c2_ring_metadata=torch.zeros(5, 1, dtype=torch.int32), + c2_metadata_group_id="ring", + c2_source_cos=rope, + c2_source_sin=rope, + ) + events = [] + hidden_states = torch.randn(num_tokens, 8, dtype=dtype) + original_hidden_states = hidden_states.clone() + + def pool(kv, score, metadata): + assert kv.dtype == score.dtype == torch.float32 + # Identity projections must receive the same FP32 conversion. + assert kv is score + torch.testing.assert_close(kv, hidden_states.float(), rtol=0, atol=0) + if dtype == torch.float32: + assert kv is hidden_states + assert metadata is state + events.append("pool") + return kv.to(torch.bfloat16) + + expected = slots.clone() + + def update_keys(latent, coordinates, cos, sin): + events.append("index") + assert coordinates.data_ptr() == slots.data_ptr() + torch.testing.assert_close(coordinates, expected) + + def store(cache, coordinates, values): + events.append("kv") + assert coordinates.data_ptr() == slots.data_ptr() + torch.testing.assert_close(coordinates, expected) + + monkeypatch.setattr(dsa_v41, "wait_for_device_metadata", lambda *args: events.append("wait")) + monkeypatch.setattr(dsa_v41, "scatter_cache_sk", store) + monkeypatch.setattr(torch.ops._C_ascend, "inplace_partial_rotary_mul", lambda *args, **kwargs: None, raising=False) + attn = SimpleNamespace( + compressor=SimpleNamespace(wkv=lambda x: x, wgate=lambda x: x, pool_projected=pool), + indexer=SimpleNamespace(update_keys=update_keys), + long_kv_cache=SimpleNamespace(kv_cache=[torch.empty(0)]), + head_dim=8, + nope_head_dim=6, + ) + cache = SimpleNamespace(slot_mapping=slots) + metadata = SimpleNamespace( + compressor=SimpleNamespace(cache=cache, state=state), + indexer=SimpleNamespace(cache=cache), + ) + DeepseekV41EagerAttentionImpl._write_compressed_source( + SimpleNamespace(role=SimpleNamespace(compress_ratio=2)), + attn, + hidden_states, + positions, + rope, + rope, + metadata, + ) + assert events == ["wait", "pool", "index", "kv"] + torch.testing.assert_close(slots, original_slots) + torch.testing.assert_close(hidden_states, original_hidden_states, rtol=0, atol=0) + + +def test_state_uses_one_ring_page_and_block_table_entry(config, runtime): + from vllm_ascend.core.circular_buffer import AscendCircularBufferSpec + + spec = collect_specs(runtime)["model.layers.2.self_attn.compressor.state_cache"] + assert isinstance(spec, AscendCircularBufferSpec) + assert spec.compress_ratio == 1 and not spec.prefix_cacheable + assert spec.storage_block_size == 32 + assert spec.page_size_bytes == 32 * 16 * 4 + assert spec.max_num_blocks_per_req(runtime, 1024) == 1 + assert spec.max_memory_usage_bytes(runtime) == spec.page_size_bytes + + +@pytest.mark.parametrize("change", [{"dtype": torch.bfloat16}, {"block_size": 16}, {"compress_ratio": 2}]) +def test_state_spec_rejects_precision_or_capacity_changes(runtime, change): + spec = collect_specs(runtime)["model.layers.2.self_attn.compressor.state_cache"] + with pytest.raises(ValueError, match="32-row FP32"): + replace(spec, **change) + + +def test_ring_view_rejects_unrepresented_page_padding(runtime): + spec = collect_specs(runtime)["model.layers.2.self_attn.compressor.state_cache"] + stride = 2 * spec.page_size_bytes + with pytest.raises(ValueError, match="fill its slot"): + reshape_cache(torch.zeros(3 * stride, dtype=torch.uint8), spec, num_blocks=3, offset=0, block_stride=stride) + + +def test_projected_model_entry_keeps_fp32_state_and_existing_norm(config, monkeypatch): + compressor = DeepseekV41Compressor(config, 2) + compressor.register_buffer("_ring_pooled", torch.empty(4, 8, dtype=torch.bfloat16), persistent=False) + compressor._ring_num_cores = 1 + state = torch.zeros(3, 32, 1, 16, dtype=torch.float32) + compressor.state_cache = SimpleNamespace(kv_cache=[state]) + metadata = SimpleNamespace(c2_ring_metadata=torch.zeros(5, 1, dtype=torch.int32), max_query_len=2) + pooled = torch.randn(2, 8, dtype=torch.bfloat16) + expected = compressor.norm(pooled).clone() + pointer = compressor._ring_pooled.data_ptr() + + def kernel(kv, scores, state_view, controls, out, **kwargs): + assert kv.dtype == scores.dtype == state_view.dtype == torch.float32 + assert state_view.data_ptr() == state.data_ptr() + assert controls is metadata.c2_ring_metadata + assert out.data_ptr() == pointer and out.dtype == torch.bfloat16 + out.copy_(pooled) + return out + + monkeypatch.setattr("vllm_ascend.ops.triton.compressor.compressor_triton.compressor_from_projected", kernel) + hidden = torch.randn(2, 16, dtype=torch.bfloat16) + actual = compressor.pool_projected(compressor.wkv(hidden.float()), compressor.wgate(hidden.float()), metadata) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + assert compressor.wkv.weight.dtype == compressor.wgate.weight.dtype == torch.float32 + + +@torch.inference_mode() +def test_compressor_rejects_missing_previous_state_page(config): + compressor = DeepseekV41Compressor(config, 2) + state = torch.full((3, 32, 16), float("nan"), dtype=torch.float32) + with pytest.raises(ValueError, match="absent/null"): + compressor_ratio2_reference( + compressor, + torch.zeros(1, 16, dtype=torch.bfloat16), + 1, + state, + [0], + ) + + +@torch.inference_mode() +def test_state_page_reuse_does_not_require_request_reset(config): + compressor = DeepseekV41Compressor(config, 2) + state = torch.full((3, 32, 16), float("nan"), dtype=torch.float32) + x = torch.randn(2, 16, dtype=torch.bfloat16) + expected = compressor_ratio2_reference(compressor, x, 0, state, [1]).clone() + state[1].fill_(12345) + actual = compressor_ratio2_reference(compressor, x, 0, state, [1]) + torch.testing.assert_close(actual, expected) + + +def test_state_registers_circular_manager(monkeypatch): + from vllm.v1.kv_cache_spec_registry import KVCacheSpecRegistry + + from vllm_ascend.core.circular_buffer import AscendCircularBufferManager + from vllm_ascend.core.deepseek_v41 import DeepseekV41CompressorStateSpec + from vllm_ascend.core.kv_cache_interface import register_ascend_kv_cache_specs + + registrations = {} + + def record(kvcache_spec_cls, manager_class, uniform_type_base_spec): + registrations[kvcache_spec_cls] = manager_class + + monkeypatch.setattr(KVCacheSpecRegistry, "register", record) + register_ascend_kv_cache_specs() + assert registrations[DeepseekV41CompressorStateSpec] is AscendCircularBufferManager + + +@torch.inference_mode() +def test_interleaved_request_state_isolation(config): + compressor = DeepseekV41Compressor(config, 2) + state = torch.full((3, 32, 16), float("nan"), dtype=torch.float32) + first = torch.randn(2, 16, dtype=torch.bfloat16) + second = torch.randn(2, 16, dtype=torch.bfloat16) + compressor_ratio2_reference(compressor, first[:1], 0, state, [1]) + saved = state[1, 0].clone() + compressor_ratio2_reference(compressor, second, 0, state, [2]) + torch.testing.assert_close(state[1, 0], saved) + actual = compressor_ratio2_reference(compressor, first[1:], 1, state, [1]) + expected = compressor_ratio2_reference(compressor, first, 0, state, [1]) + torch.testing.assert_close(actual, expected) + + +class _CPCommon(SimpleNamespace): + def replace(self, **kwargs): + return type(self)(**(vars(self) | kwargs)) + + +def _cp_common(): + # The second request resumes in the middle of a ratio-2 pair. + return _CPCommon( + slot_mapping=torch.tensor([0, 1, 2, 68]), + block_table_tensor=torch.tensor([[0, 1], [1, 2]]), + query_start_loc=torch.tensor([0, 3, 4], dtype=torch.int32), + query_start_loc_cpu=torch.tensor([0, 3, 4], dtype=torch.int32), + seq_lens=torch.tensor([3, 5], dtype=torch.int32), + seq_lens_cpu=torch.tensor([3, 5], dtype=torch.int32), + num_reqs=2, + num_actual_tokens=4, + num_input_tokens=4, + max_query_len=3, + max_seq_len=5, + positions=torch.tensor([0, 1, 2, 4]), + is_prefilling=torch.tensor([True, True]), + causal=True, + ) + + +@pytest.mark.parametrize( + "rank,size,query_offsets,seq_lens,positions", + [ + (0, 2, [0, 2, 2], [2, 0], [0, 1]), + (1, 2, [0, 1, 2], [3, 5], [2, 4]), + (5, 8, [0, 0, 0], [0, 0], []), + ], +) +def test_v41_cp_metadata_preserves_global_compression_and_local_causality( + runtime, monkeypatch, rank, size, query_offsets, seq_lens, positions +): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPMetadataBuilder + + monkeypatch.setattr( + "vllm_ascend.attention.context_parallel.dsa_cp.get_tp_group", + lambda: SimpleNamespace(world_size=size, rank_in_group=rank), + ) + spec = collect_specs(runtime)["model.layers.2.self_attn.long_kv_cache"] + builder = DeepseekV41CPMetadataBuilder(spec, [], runtime, torch.device("cpu")) + metadata = builder.build(0, _cp_common(), common_v41_metadata={}, common_v41_batch_metadata={}) + assert metadata.query_start_loc.tolist() == query_offsets + assert metadata.query_start_loc.dtype == torch.int32 + pointer = metadata.query_start_loc.data_ptr() + assert builder.build(0, _cp_common()).query_start_loc.data_ptr() == pointer + assert metadata.seq_lens.tolist() == seq_lens + assert metadata.positions.tolist() == positions + assert metadata.global_metadata.seq_lens.tolist() == [3, 5] + assert metadata.global_metadata.cache_seq_lens.tolist() == [1, 2] + assert metadata.global_metadata.slot_mapping.tolist() == [[-1, -1], [0, 0], [-1, -1], [-1, -1]] + assert metadata.num_actual_tokens == len(positions) + + +@pytest.mark.parametrize("rank", [0, 1]) +def test_v41_cp_uses_device_seq_lens_when_cpu_mirror_is_upper_bound(runtime, monkeypatch, rank): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPMetadataBuilder + + monkeypatch.setattr( + "vllm_ascend.attention.context_parallel.dsa_cp.get_tp_group", + lambda: SimpleNamespace(world_size=2, rank_in_group=rank), + ) + # A speculative rejection has corrected device lengths and positions; + # the host mirror still describes the optimistic upper bound. + common = _cp_common().replace( + seq_lens=torch.tensor([7, 9], dtype=torch.int32), + seq_lens_cpu=None, + _seq_lens_cpu=torch.tensor([9, 11], dtype=torch.int32), + positions=torch.tensor([4, 5, 6, 8]), + max_seq_len=11, + ) + spec = collect_specs(runtime)["model.layers.2.self_attn.long_kv_cache"] + builder = DeepseekV41CPMetadataBuilder(spec, [], runtime, torch.device("cpu")) + metadata = builder.build(0, common) + expected = [6, 0] if rank == 0 else [7, 9] + assert metadata.global_metadata.seq_lens.tolist() == [7, 9] + assert metadata.seq_lens.tolist() == expected + assert metadata.cache_seq_lens.tolist() == [n // 2 for n in expected] + assert metadata.cmp_residual.tolist() == [n % 2 for n in expected] + + +@pytest.mark.parametrize("rank", [0, 1, 3, 7]) +@pytest.mark.parametrize("prefill", [False, True]) +def test_v41_cp_rope_preserves_global_rows_across_builds(runtime, monkeypatch, rank, prefill): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPMetadataBuilder + from vllm_ascend.ops import rope_dsv4 as rope + + monkeypatch.setattr( + "vllm_ascend.attention.context_parallel.dsa_cp.get_tp_group", + lambda: SimpleNamespace(world_size=8, rank_in_group=rank), + ) + state = rope.RopeGlobalState() + full = torch.arange(128, dtype=torch.float32).reshape(128, 1, 1, 1) + state.full_rope_cache["test"] = (full, full + 1000) + state.runtime_buffer["test"] = {"default": (torch.zeros(8, 1, 1, 1), torch.zeros(8, 1, 1, 1))} + state.registry_summary["test"] = {"default"} + state.layer_info["test.layer"] = ("test", ["default"]) + monkeypatch.setattr(rope, "_ROPE_STATE", state) + calls = [] + + def gather(positions, **kwargs): + calls.append(positions.clone()) + return rope.get_cos_and_sin_dsa(positions, **kwargs) + + monkeypatch.setattr("vllm_ascend.attention.dsa_v41.get_cos_and_sin_dsa", gather) + spec = collect_specs(runtime)["model.layers.0.self_attn.swa_cache"] + builder = DeepseekV41CPMetadataBuilder(spec, [], runtime, torch.device("cpu")) + pointer = None + for step in (0, 3): + positions = torch.tensor([10, 20, 30, 40]) + step + offsets = torch.arange(5, dtype=torch.int32) + common = _cp_common().replace( + positions=positions, + num_reqs=4, + query_start_loc=offsets, + query_start_loc_cpu=offsets, + seq_lens=(positions + 1).int(), + seq_lens_cpu=(positions + 1).int(), + max_query_len=1, + max_seq_len=int(positions.max()) + 1, + block_table_tensor=torch.zeros(4, 2, dtype=torch.int32), + is_prefilling=torch.full((4,), prefill), + ) + metadata = builder.build(0, common) + global_cos = metadata.global_metadata.cos["test.layer"] + local_cos = metadata.cos["test.layer"] + local_sin = metadata.sin["test.layer"] + expected = positions[rank : rank + 1].float() + torch.testing.assert_close(global_cos.flatten(), positions.float()) + torch.testing.assert_close(local_cos.flatten(), expected) + torch.testing.assert_close(local_sin.flatten(), expected + 1000) + if not prefill: + if pointer is not None: + assert global_cos.data_ptr() == pointer + pointer = global_cos.data_ptr() + if expected.numel(): + assert local_cos.data_ptr() == global_cos.data_ptr() + rank * global_cos.element_size() + assert len(calls) == 2 # One global gather per build, including empty local ranks. + + +@pytest.mark.parametrize( + "cache_name,field,stage", + [ + ("model.layers.0.self_attn.swa_cache", "smla_metadata", DeviceMetadataStage.ATTENTION), + ("model.layers.2.self_attn.indexer.k_cache", "qli_metadata", DeviceMetadataStage.INDEXER), + ("model.layers.2.self_attn.compressor.state_cache", "c2_ring_metadata", DeviceMetadataStage.COMPRESSOR), + ], +) +def test_v41_cp_builds_device_controls_only_on_consuming_side(runtime, monkeypatch, cache_name, field, stage): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPMetadataBuilder + + monkeypatch.setattr( + "vllm_ascend.attention.context_parallel.dsa_cp.get_tp_group", + lambda: SimpleNamespace(world_size=2, rank_in_group=0), + ) + builder = DeepseekV41CPMetadataBuilder(collect_specs(runtime)[cache_name], [], runtime, torch.device("cpu")) + global_builder = builder._global_builder + compressor = stage == DeviceMetadataStage.COMPRESSOR + for side in (builder, global_builder): + # Queue native metadata operations without invoking NPU kernels on CPU. + side._device_metadata_enabled = True + side._supports_device_ops = not compressor + assert global_builder._smla_metadata.numel() == 0 + assert global_builder._qli_metadata.numel() == 0 + assert builder._c2_ring_metadata.numel() == 0 + assert builder._c2_complete_mask.numel() == 0 + assert builder._c2_source_positions.numel() == 0 + assert builder._c2_source_cos.numel() == 0 + assert builder._c2_source_sin.numel() == 0 + for _ in range(2): + metadata = builder.build(0, _cp_common(), common_v41_metadata={}, common_v41_batch_metadata={}) + owner, unused = (metadata.global_metadata, metadata) if compressor else (metadata, metadata.global_metadata) + assert getattr(owner, field) is not None + assert getattr(unused, field) is None + tasks = builder.take_device_metadata_tasks() + assert len(tasks) == 1 + assert tasks[0].stage == stage + assert builder.take_device_metadata_tasks() == () + + +@pytest.mark.parametrize("local_tokens", [0, 1, 2]) +def test_v41_cp_output_exchange_only_pads_partial_ranks(monkeypatch, local_tokens): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPImpl + + impl = DeepseekV41CPImpl("layer", SimpleNamespace(is_kv_source=False), None, None, None) + calls = [] + monkeypatch.setattr("vllm_ascend.attention.context_parallel.dsa_v41_cp.get_tp_group", lambda: None) + + def exchange(tensor, group): + calls.append(tensor) + return torch.ones((4, 2, 3)) + + monkeypatch.setattr("vllm_ascend.attention.context_parallel.dsa_v41_cp.restore_tp_heads", exchange) + projection = SimpleNamespace(_forward_o_proj=lambda tensor: tensor.flatten(1)) + attn = SimpleNamespace(dsa_attn=SimpleNamespace(dsa_attn=SimpleNamespace(impl=projection))) + destination = torch.empty((3, 6)) + local_output = torch.ones((local_tokens, 4, 3)) + output = impl._project_output( + attn, + local_output, + torch.empty((3, 6)), + SimpleNamespace(swa=SimpleNamespace(cp_token_range=(0, 2, 2, 4))), + projected=destination, + ) + assert output.data_ptr() == destination.data_ptr() + assert len(calls) == 1 + assert calls[0].shape == (2, 4, 3) + assert (calls[0] is local_output) == (local_tokens == 2) + torch.testing.assert_close(calls[0][:local_tokens], local_output) + assert torch.count_nonzero(calls[0][local_tokens:]) == 0 + assert output.shape == (3, 6) + + +def test_v41_cp_consumers_reuse_local_topk_and_candidates(): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPImpl + + impl = DeepseekV41CPImpl( + "layer", SimpleNamespace(is_kv_source=False, has_long_context=True, is_index_source=False), None, None, None + ) + indices = torch.tensor([[0, 2], [1, 3]]) + candidates = torch.tensor([[True, False]]) + shared = SimpleNamespace(topk_indices=indices, candidates=candidates) + actual = impl._select_sparse_indices( + SimpleNamespace(shared_state=shared), + torch.empty(16, 1), + torch.empty(2, 1), + None, + None, + None, + None, + ) + assert actual.data_ptr() == indices.data_ptr() + torch.testing.assert_close(actual, indices) + assert shared.candidates is candidates + + +@pytest.mark.parametrize("pcp,cp", [(False, False), (True, False), (False, True), (True, True)]) +def test_v41_backend_routes_metadata_and_execution_together(monkeypatch, pcp, cp): + from vllm_ascend.attention.context_parallel import dsa_v41_cp + from vllm_ascend.attention.dsa_v41 import DeepseekV41CacheBackend + + monkeypatch.setattr(dsa_v41_cp, "enable_pcp", lambda: pcp) + monkeypatch.setattr(dsa_v41_cp, "enable_dsa_cp", lambda: cp) + if pcp: + with pytest.raises(NotImplementedError, match="PCP is not supported"): + DeepseekV41CacheBackend.get_builder_cls() + return + builder, impl = dsa_v41_cp.get_v41_cp_classes() + assert DeepseekV41CacheBackend.get_builder_cls() is builder + assert not DeepseekV41CacheBackend.supports_pcp() + if cp: + assert issubclass(impl, dsa_v41_cp.DeepseekV41CPImpl) + + +def test_v41_cp_accepts_async_seq_lens_mirror(runtime, monkeypatch): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPMetadataBuilder + + monkeypatch.setattr( + "vllm_ascend.attention.context_parallel.dsa_cp.get_tp_group", + lambda: SimpleNamespace(world_size=2, rank_in_group=1), + ) + common = _cp_common() + common._seq_lens_cpu = common.seq_lens_cpu + common.seq_lens_cpu = None + spec = collect_specs(runtime)["model.layers.2.self_attn.long_kv_cache"] + metadata = DeepseekV41CPMetadataBuilder(spec, [], runtime, torch.device("cpu")).build(0, common) + assert metadata.seq_lens.tolist() == [3, 5] + assert metadata.global_metadata.cache_seq_lens.tolist() == [1, 2] + + +def test_v41_cp_resolves_own_planes_with_native_draft_metadata_present(): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPImpl + + impl = DeepseekV41CPImpl( + prefix="model.layers.0.self_attn", + role=SimpleNamespace(is_kv_source=False, compress_ratio=0), + topology=None, + long_kv_source_prefix=None, + index_k_source_prefix=None, + ) + global_swa = object() + metadata = { + impl.swa_prefix: SimpleNamespace(global_metadata=global_swa), + "mtp.0.self_attn.swa_cache": SimpleNamespace(seq_lens=torch.tensor([4])), + } + assert impl._global_layer_metadata(metadata).swa is global_swa + + +@pytest.mark.parametrize("v2,pcp", [(False, 2), (True, 1)]) +def test_v41_runtime_rejects_pcp_and_mrv2(runtime, v2, pcp): + from vllm_ascend.core.deepseek_v41 import validate_cache_runtime + + runtime.use_v2_model_runner = v2 + runtime.parallel_config.prefill_context_parallel_size = pcp + with pytest.raises(NotImplementedError, match="runner V1" if v2 else "PCP=1"): + validate_cache_runtime(runtime) + + +@pytest.mark.parametrize("overlap", [False, True]) +def test_v41_query_preparation_keeps_mainline_preprocess(overlap): + from unittest.mock import Mock + + from vllm_ascend.attention.dsa_v41 import DeepseekV41EagerAttentionImpl + + impl = DeepseekV41EagerAttentionImpl.__new__(DeepseekV41EagerAttentionImpl) + impl.role = SimpleNamespace(is_kv_source=True) + impl.preprocess = Mock(return_value=("q", "qr")) + impl.multistream_preprocess = Mock(return_value=("q", "qr")) + impl._write_compressed_source = Mock() + attn = SimpleNamespace( + dsa_attn=SimpleNamespace(dsa_attn=SimpleNamespace(impl=SimpleNamespace(multistream_dsv4_dsa_overlap=overlap))) + ) + metadata = SimpleNamespace(swa=SimpleNamespace(num_actual_tokens=6)) + assert impl._prepare_queries(attn, "hidden", "positions", "cos", "sin", metadata) == ("q", "qr") + selected = impl.multistream_preprocess if overlap else impl.preprocess + other = impl.preprocess if overlap else impl.multistream_preprocess + selected.assert_called_once_with(attn, "hidden", "cos", "sin", metadata.swa) + other.assert_not_called() + impl._write_compressed_source.assert_called_once_with(attn, "hidden", "positions", "cos", "sin", metadata) + + +@pytest.mark.parametrize("overlap", [False, True]) +def test_v41_cp_query_preparation_uses_full_inputs_only_for_overlap(overlap): + from unittest.mock import Mock + + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPImpl + + impl = DeepseekV41CPImpl.__new__(DeepseekV41CPImpl) + impl._project_q = Mock(return_value=("q", "qr")) + impl.multistream_preprocess = Mock(return_value=("q", "qr")) + impl._write_compressed_source = Mock() + attn = SimpleNamespace( + dsa_attn=SimpleNamespace(dsa_attn=SimpleNamespace(impl=SimpleNamespace(multistream_dsv4_dsa_overlap=overlap))) + ) + metadata = SimpleNamespace(swa=SimpleNamespace(num_actual_tokens=2, cp_token_range=(2, 4, 2, 6))) + assert impl._prepare_queries(attn, "abcdef", "positions", "cos", "sin", metadata) == ("q", "qr") + if overlap: + impl.multistream_preprocess.assert_called_once_with(attn, "abcdef", "cos", "sin", metadata.swa) + impl._project_q.assert_not_called() + else: + impl._project_q.assert_called_once_with(attn, "cd", "cos", "sin") + impl.multistream_preprocess.assert_not_called() + impl._write_compressed_source.assert_not_called() + + +@pytest.mark.parametrize("overlap", [False, True]) +@pytest.mark.parametrize("local_tokens", [0, 2]) +def test_v41_cp_input_preparation_updates_empty_rank_cache(overlap, local_tokens): + from unittest.mock import Mock + + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPImpl + + impl = DeepseekV41CPImpl.__new__(DeepseekV41CPImpl) + full = torch.arange(24).reshape(6, 4) + global_metadata = SimpleNamespace(swa=SimpleNamespace(num_actual_tokens=5)) + impl._global_layer_metadata = Mock(return_value=global_metadata) + impl._update_caches = Mock() + attn = SimpleNamespace( + dsa_attn=SimpleNamespace(dsa_attn=SimpleNamespace(impl=SimpleNamespace(multistream_dsv4_dsa_overlap=overlap))) + ) + metadata = SimpleNamespace(swa=SimpleNamespace(cp_token_range=(3, 6, 3, 6), num_actual_tokens=local_tokens)) + assert impl._prepare_inputs_and_caches(attn, full, metadata, {}) is None + if not overlap or local_tokens == 0: + impl._update_caches.assert_called_once() + assert torch.equal(impl._update_caches.call_args.args[1], full[:5]) + assert impl._update_caches.call_args.args[2] is global_metadata + else: + impl._update_caches.assert_not_called() + + +def test_v41_cp_inherits_forward(): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPImpl + from vllm_ascend.attention.dsa_v41 import DeepseekV41EagerAttentionImpl + + assert DeepseekV41CPImpl.forward is DeepseekV41EagerAttentionImpl.forward + + +@pytest.mark.parametrize("rank", [None, 0, 1, 7]) +@pytest.mark.parametrize("first_seq_len", [3, 260]) +def test_dspark_v41_noncausal_metadata_preserves_full_visible_block(runtime, monkeypatch, rank, first_seq_len): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPMetadataBuilder + from vllm_ascend.core.deepseek_v41 import DeepseekV41DraftSWASpec + + runtime.speculative_config = SimpleNamespace(num_speculative_tokens=3) + spec = DeepseekV41DraftSWASpec( + block_size=128, + num_kv_heads=1, + head_size=8, + dtype=torch.bfloat16, + sliding_window=128, + cache_dtype_str="bfloat16", + model_version="deepseek_v4", + ) + common = _cp_common().replace( + causal=False, + # Deliberately non-identity pages: logical positions must not be mapped + # here because SparseFlashMLA performs the physical lookup itself. + block_table_tensor=torch.tensor([[7, 3, 9], [5, 2, 8]], dtype=torch.int32), + seq_lens=torch.tensor([first_seq_len, 5], dtype=torch.int32), + seq_lens_cpu=torch.tensor([first_seq_len, 5], dtype=torch.int32), + max_seq_len=max(first_seq_len, 5), + positions=torch.tensor([first_seq_len - 3, first_seq_len - 2, first_seq_len - 1, 4]), + ) + native = Mock(return_value=torch.zeros(dsa_v41.V41_METADATA_BUFFER_SIZE, dtype=torch.int32)) + monkeypatch.setattr(torch.ops._C_ascend, "npu_sparse_flash_mla_metadata", native, raising=False) + builder = DeepseekV41MetadataBuilder(spec, [], runtime, torch.device("cpu")) + builder._supports_device_ops = True + full = builder.build_for_drafting(common, 1) + torch.testing.assert_close( + native.call_args.kwargs["ori_topk_length"], + (full.ori_sparse_indices >= 0).sum(-1, dtype=torch.int32), + ) + assert full.ori_topk_length is native.call_args.kwargs["ori_topk_length"] + assert full.ori_mask_mode == 0 + assert full.ori_sparse_indices.shape[0] == 4 + # Every query of the first request can see its complete draft block. + torch.testing.assert_close(full.ori_sparse_indices[0], full.ori_sparse_indices[2]) + expected = list(range(max(0, first_seq_len - 3 - 128), first_seq_len)) + assert full.ori_sparse_indices[0, 0, : len(expected)].tolist() == expected + assert torch.all(full.ori_sparse_indices[0, 0, len(expected) :] == -1) + assert full.ori_sparse_indices[3, 0, :5].tolist() == list(range(5)) + assert full.ori_topk_length[:, 0].tolist() == [len(expected)] * 3 + [5] + if rank is None: + return + monkeypatch.setattr( + "vllm_ascend.attention.context_parallel.dsa_cp.get_tp_group", + lambda: SimpleNamespace(world_size=8, rank_in_group=rank), + ) + native.reset_mock() + local_builder = DeepseekV41CPMetadataBuilder(spec, [], runtime, torch.device("cpu")) + local_builder._supports_device_ops = True + local = local_builder.build_for_drafting(common, 1) + if rank >= full.num_actual_tokens: + native.assert_not_called() + assert local.smla_metadata is local_builder._smla_metadata + assert torch.count_nonzero(local.smla_metadata) == 0 + else: + torch.testing.assert_close( + native.call_args.kwargs["ori_topk_length"], + (local.ori_sparse_indices >= 0).sum(-1, dtype=torch.int32), + ) + torch.testing.assert_close(local.ori_sparse_indices, full.ori_sparse_indices[rank : rank + 1]) + assert local.seq_lens.tolist() == ([first_seq_len, 0] if rank < 3 else [0, 0]) + assert local.ori_mask_mode == 0 diff --git a/tests/ut/models/test_deepseek_v41_dspark.py b/tests/ut/models/test_deepseek_v41_dspark.py new file mode 100644 index 000000000000..71447ae49705 --- /dev/null +++ b/tests/ut/models/test_deepseek_v41_dspark.py @@ -0,0 +1,188 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Deferred torch checks for Aurora DSpark model/cache integration.""" + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest +import torch +from vllm.v1.core.single_type_kv_cache_manager import SlidingWindowManager +from vllm.v1.kv_cache_spec_registry import KVCacheSpecRegistry +from vllm.v1.worker.gpu_model_runner import GPUModelRunner + +from vllm_ascend.core.deepseek_v41 import DeepseekV41DraftSWASpec +from vllm_ascend.core.kv_cache_interface import AscendSlidingWindowMLASpec, register_ascend_kv_cache_specs +from vllm_ascend.models.deepseek_v4.model import AscendDeepseekV4SWACache +from vllm_ascend.models.deepseek_v41 import dspark as deepseek_v41_dspark_module +from vllm_ascend.models.deepseek_v41.dspark import ( + DeepseekV41DSparkAttention, + DeepseekV41DSparkDecoderLayer, + DeepseekV41DSparkModel, + DeepseekV41DSparkSWACache, +) +from vllm_ascend.models.deepseek_v41.model import DeepseekV41Model +from vllm_ascend.worker.model_runner_v1 import NPUModelRunner + + +def test_draft_cache_uses_v41_backend_and_explicit_aurora_spec(monkeypatch): + spec = AscendSlidingWindowMLASpec( + block_size=128, + num_kv_heads=1, + head_size=512, + dtype=torch.bfloat16, + sliding_window=128, + cache_dtype_str="bfloat16", + model_version="deepseek_v4", + ) + with patch.object(AscendDeepseekV4SWACache, "get_kv_cache_spec", return_value=spec): + cache = DeepseekV41DSparkSWACache.__new__(DeepseekV41DSparkSWACache) + draft = cache.get_kv_cache_spec(None) + from vllm_ascend.attention.dsa_v41 import DeepseekV41CacheBackend + + assert DeepseekV41DSparkSWACache.get_attn_backend(None) is DeepseekV41CacheBackend + assert type(draft) is DeepseekV41DraftSWASpec + assert draft.page_size_bytes == 131072 + assert DeepseekV41DSparkDecoderLayer.attention_cls is DeepseekV41DSparkAttention + assert DeepseekV41DSparkAttention.swa_cache_cls is DeepseekV41DSparkSWACache + registrations = {} + + def record(kvcache_spec_cls, manager_class, uniform_type_base_spec): + registrations[kvcache_spec_cls] = manager_class + + monkeypatch.setattr(KVCacheSpecRegistry, "register", record) + register_ascend_kv_cache_specs() + assert registrations[type(draft)] is SlidingWindowManager + + +def test_composite_config_selects_checkpoint_aux_layers(): + runner = NPUModelRunner.__new__(NPUModelRunner) + text = SimpleNamespace(dspark_target_layer_ids=[37, 38, 39]) + runner.speculative_config = SimpleNamespace( + use_dspark=lambda: True, + draft_model_config=SimpleNamespace(hf_config=SimpleNamespace(text_config=text)), + ) + with patch.object(GPUModelRunner, "_get_eagle3_aux_layers_from_config", return_value=None): + assert runner._get_eagle3_aux_layers_from_config() == (38, 39, 40) + + +def test_target_exports_residual_entering_selected_layers(monkeypatch): + monkeypatch.setattr( + "vllm_ascend.models.deepseek_v41.model.get_pp_group", + lambda: SimpleNamespace(is_first_rank=True, is_last_rank=True), + ) + + class Layer: + def __init__(self, index): + self.layer_idx = index + self.engram = None + + def __call__(self, positions, hidden, pre_mix, unused, input_ids): + return hidden + self.layer_idx + 1, pre_mix + + @staticmethod + def hc_collapse(hidden, pre_mix): + return hidden.mean(dim=1) + + model = SimpleNamespace( + hc_mult=4, + needs_moe_input_ids=False, + prepare_engram=lambda input_ids, positions: ({}, torch.empty(0, dtype=torch.bool)), + aux_hidden_state_layers=(1, 3), + shared_attention_state=SimpleNamespace(reset=lambda: None), + layers=[Layer(i) for i in range(3)], + norm=lambda hidden: hidden, + ) + hidden = torch.arange(12, dtype=torch.float32).reshape(3, 4) + output, aux = DeepseekV41Model.forward(model, torch.arange(3), torch.arange(3), None, inputs_embeds=hidden) + torch.testing.assert_close(aux[0], hidden) + torch.testing.assert_close(aux[1], hidden + 3) + torch.testing.assert_close(output, hidden + 6) + + +@pytest.mark.parametrize("cp", [False, True]) +def test_v41_draft_routes_to_v41_and_disables_post_projection_q_norm(cp): + from vllm_ascend.models.deepseek_v4.model import DeepseekV4Attention + + ordinary_backend = SimpleNamespace(apply_q_norm=True) + draft_backend = SimpleNamespace(apply_q_norm=True) + + def initialize_base(instance, **kwargs): + torch.nn.Module.__init__(instance) + instance.compress_ratio = 0 + instance.scale = 512**-0.5 + instance.dsa_attn = SimpleNamespace(dsa_attn=SimpleNamespace(impl=draft_backend)) + + from vllm_ascend.attention.context_parallel.dsa_v41_cp import DeepseekV41CPImpl + from vllm_ascend.attention.dsa_v41 import DeepseekV41EagerAttentionImpl + + config = SimpleNamespace(compilation_config=SimpleNamespace(static_forward_context={})) + with ( + patch.object(DeepseekV4Attention, "__init__", initialize_base), + patch("vllm_ascend.attention.context_parallel.dsa_v41_cp.enable_dsa_cp", return_value=cp), + patch("vllm_ascend.attention.context_parallel.dsa_v41_cp.enable_pcp", return_value=False), + ): + draft = DeepseekV41DSparkAttention(vllm_config=config, prefix="mtp.0.self_attn") + assert type(draft.v41_impl) is (DeepseekV41CPImpl if cp else DeepseekV41EagerAttentionImpl) + assert config.compilation_config.static_forward_context[draft.v41_layer_name] is draft + assert draft.softmax_scale == 512**-0.5 + assert draft.dsa_attn.dsa_attn.impl.apply_q_norm is False + assert ordinary_backend.apply_q_norm is True + + +def test_v41_draft_sequence_parallel_shards_inputs_and_restores_output(monkeypatch): + class Layer: + @staticmethod + def hc_collapse(hidden, pre_mix): + return hidden.mean(dim=1) + + def __call__(self, positions, hidden, pre_mix, llama_4_scaling, input_ids): + return hidden, pre_mix + + hidden = torch.arange(16, dtype=torch.float32).reshape(4, 4) + input_ids = torch.tensor([11, 12, 13, 14]) + padding = torch.tensor([False, True, False, False]) + forward_context = SimpleNamespace(is_padding=padding) + sharded_hidden = hidden[:2].unsqueeze(1).repeat(1, 4, 1) + sharded_ids = input_ids[:2] + gathered = torch.arange(20, dtype=torch.float32).reshape(5, 4) + sp_shard = MagicMock(side_effect=[sharded_hidden, sharded_ids]) + sp_all_gather = MagicMock(return_value=gathered) + padding_mask = MagicMock(return_value=torch.tensor([False, True])) + monkeypatch.setattr(deepseek_v41_dspark_module, "sp_shard", sp_shard) + monkeypatch.setattr(deepseek_v41_dspark_module, "sp_all_gather", sp_all_gather) + monkeypatch.setattr(deepseek_v41_dspark_module, "sp_padding_mask", padding_mask) + monkeypatch.setattr(deepseek_v41_dspark_module, "is_forward_context_available", lambda: True) + monkeypatch.setattr(deepseek_v41_dspark_module, "get_forward_context", lambda: forward_context) + monkeypatch.setattr(deepseek_v41_dspark_module.envs, "VLLM_MOE_SKIP_PADDING", True) + + model = SimpleNamespace( + embed_tokens=MagicMock(return_value=hidden), + hc_mult=4, + use_sequence_parallel=True, + needs_moe_input_ids=False, + layers={"40": Layer()}, + ) + + output = DeepseekV41DSparkModel.forward(model, input_ids, torch.arange(4)) + + padding_mask.assert_called_once() + assert forward_context.is_padding.tolist() == [False, True] + assert sp_shard.call_args_list[0].args[0].shape == (4, 4, 4) + assert sp_shard.call_args_list[1].args[0] is input_ids + sp_all_gather.assert_called_once() + torch.testing.assert_close(output, gathered[:4]) + + +def test_v41_draft_context_store_uses_physical_pairs_and_preserves_padding(): + from vllm_ascend.models.deepseek_v41.dspark import DeepseekV41DSparkModel + + cache = torch.empty(3, 128, 1, 8) + attn = SimpleNamespace(dsa_attn=SimpleNamespace(swa_cache_layer=SimpleNamespace(block_size=128, kv_cache=[cache]))) + values = torch.randn(3, 1, 8) + with patch("vllm_ascend.models.deepseek_v41.dspark.scatter_cache_sk") as store: + DeepseekV41DSparkModel._store_standard_swa_kv(None, values, torch.tensor([129, -1, 258]), attn) + actual_cache, slots, updates = store.call_args.args + assert actual_cache is cache + assert slots.tolist() == [[1, 1], [-1, -1], [2, 2]] + torch.testing.assert_close(updates, values.squeeze(1)) diff --git a/tests/ut/models/test_deepseek_v41_hyper_connection.py b/tests/ut/models/test_deepseek_v41_hyper_connection.py new file mode 100644 index 000000000000..04e7e9a1ec70 --- /dev/null +++ b/tests/ut/models/test_deepseek_v41_hyper_connection.py @@ -0,0 +1,352 @@ +# SPDX-License-Identifier: Apache-2.0 + +from unittest.mock import MagicMock, patch + +import pytest +import torch + +from tests.deepseek_v41_reference import hc_mixes_reference, hc_post_reference +from vllm_ascend.models.deepseek_v41 import model as deepseek_v41_module +from vllm_ascend.models.deepseek_v41.model import DeepseekV41DecoderLayer + + +def _layer() -> DeepseekV41DecoderLayer: + layer = DeepseekV41DecoderLayer.__new__(DeepseekV41DecoderLayer) + torch.nn.Module.__init__(layer) + layer.hc_mult = 4 + layer.hc_sinkhorn_iters = 3 + layer.norm_eps = 1e-6 + layer.hc_eps = 1e-6 + return layer + + +def test_v41_hc_pre_dispatches_fused_operator_with_pre_mix(): + layer = _layer() + x = torch.randn(2, 4, 8, dtype=torch.bfloat16) + hc_fn = torch.randn(24, 32, dtype=torch.float32) + hc_scale = torch.randn(3, dtype=torch.float32) + hc_base = torch.randn(24, dtype=torch.float32) + pre_mix = torch.randn(2, 4, dtype=torch.float32) + expected = ( + torch.randn(2, 8, dtype=torch.bfloat16), + torch.randn(2, 4), + torch.randn(2, 4, 4), + torch.randn(2, 4), + ) + + with patch.object( + torch.ops._C_ascend, + "npu_hc_pre_v3", + create=True, + return_value=expected, + ) as op: + actual = layer.hc_pre(x, hc_fn, hc_scale, hc_base, pre_mix) + + assert actual is expected + op.assert_called_once_with( + x, + hc_fn, + hc_scale, + hc_base, + pre_mix, + hc_mult=4, + hc_sinkhorn_iters=3, + norm_eps=1e-6, + hc_eps=1e-6, + ) + + +def test_v41_forward_threads_pre_mix_through_fused_hc_pre(): + layer = _layer() + hidden_states = torch.randn(2, 4, 8, dtype=torch.bfloat16) + incoming_pre = torch.randn(2, 4, dtype=torch.float32) + attn_pre = torch.randn(2, 4, dtype=torch.float32) + ffn_pre = torch.randn(2, 4, dtype=torch.float32) + post = torch.randn(2, 4, dtype=torch.float32) + comb = torch.randn(2, 4, 4, dtype=torch.float32) + collapsed = torch.randn(2, 8, dtype=torch.bfloat16) + layer.hc_attn_fn = torch.nn.Parameter(torch.empty(24, 32)) + layer.hc_attn_scale = torch.nn.Parameter(torch.empty(3)) + layer.hc_attn_base = torch.nn.Parameter(torch.empty(24)) + layer.hc_ffn_fn = torch.nn.Parameter(torch.empty(24, 32)) + layer.hc_ffn_scale = torch.nn.Parameter(torch.empty(3)) + layer.hc_ffn_base = torch.nn.Parameter(torch.empty(24)) + layer.hc_pre = MagicMock( + side_effect=[ + (collapsed, post, comb, attn_pre), + (collapsed, post, comb, ffn_pre), + ] + ) + layer.input_layernorm = MagicMock(side_effect=lambda value: value) + normalized = torch.randn_like(collapsed) + normalized_fp32 = normalized.float() + layer.rms_norm_cast = MagicMock(return_value=(normalized, normalized_fp32)) + layer.self_attn = MagicMock(side_effect=lambda _positions, value, _scaling: value) + layer.mlp = MagicMock(side_effect=lambda value, **_kwargs: value) + layer.hc_post = MagicMock(side_effect=lambda _x, residual, _post, _comb: residual) + + input_ids = torch.tensor([11, 22]) + output, next_pre = layer.forward(torch.arange(2), hidden_states, incoming_pre, input_ids=input_ids) + + assert output is hidden_states + assert next_pre is ffn_pre + assert layer.hc_pre.call_args_list[0].args[-1] is incoming_pre + assert layer.hc_pre.call_args_list[1].args[-1] is attn_pre + layer.rms_norm_cast.assert_called_once_with(collapsed) + layer.mlp.assert_called_once_with( + normalized, + input_ids=input_ids, + hidden_states_fp32=normalized_fp32, + already_sequence_parallel=False, + ) + + +def test_v41_forward_gathers_attention_and_keeps_moe_sharded(monkeypatch): + layer = _layer() + layer.use_sequence_parallel = True + hidden_states = torch.randn(2, 4, 8, dtype=torch.bfloat16) + collapsed = torch.randn(2, 8, dtype=torch.bfloat16) + post = torch.randn(2, 4, dtype=torch.float32) + comb = torch.randn(2, 4, 4, dtype=torch.float32) + pre = torch.randn(2, 4, dtype=torch.float32) + layer.hc_attn_fn = torch.nn.Parameter(torch.empty(24, 32)) + layer.hc_attn_scale = torch.nn.Parameter(torch.empty(3)) + layer.hc_attn_base = torch.nn.Parameter(torch.empty(24)) + layer.hc_ffn_fn = torch.nn.Parameter(torch.empty(24, 32)) + layer.hc_ffn_scale = torch.nn.Parameter(torch.empty(3)) + layer.hc_ffn_base = torch.nn.Parameter(torch.empty(24)) + layer.hc_pre = MagicMock(side_effect=[(collapsed, post, comb, pre)] * 2) + layer.input_layernorm = MagicMock(side_effect=lambda value: value) + layer.rms_norm_cast = MagicMock(return_value=(collapsed, collapsed.float())) + layer.self_attn = MagicMock(side_effect=lambda _positions, value, _scaling: value) + layer.mlp = MagicMock(side_effect=lambda value, **_kwargs: value) + layer.hc_post = MagicMock(side_effect=lambda _x, residual, _post, _comb: residual) + all_gather = MagicMock(return_value=collapsed) + reduce_scatter = MagicMock(return_value=collapsed) + monkeypatch.setattr(deepseek_v41_module, "sp_all_gather", all_gather) + monkeypatch.setattr(deepseek_v41_module, "sp_reduce_scatter", reduce_scatter) + + layer.forward(torch.arange(2), hidden_states, pre, input_ids=torch.tensor([1, 2])) + + all_gather.assert_called_once_with(collapsed) + reduce_scatter.assert_called_once() + torch.testing.assert_close(reduce_scatter.call_args.args[0], collapsed) + assert layer.mlp.call_args.kwargs["already_sequence_parallel"] is True + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +def test_v41_rms_norm_cast_preserves_rounded_routing_input(dtype): + layer = _layer() + x = torch.randn(2, 8, dtype=dtype) + normalized = torch.randn_like(x) + normalized_fp32 = normalized.float() + norm = MagicMock(return_value=normalized) + norm.weight = torch.ones(8, dtype=dtype) + norm.variance_epsilon = 1e-6 + layer.post_attention_layernorm = norm + + with ( + patch("vllm_ascend.models.deepseek_v4.model.enable_custom_op", return_value=True), + patch.object( + torch.ops._C_ascend, + "npu_rms_norm_cast", + create=True, + return_value=(normalized, normalized_fp32), + ) as op, + ): + actual, actual_fp32 = layer.rms_norm_cast(x) + + assert actual is normalized + torch.testing.assert_close(actual_fp32, normalized.float(), rtol=0, atol=0) + op.assert_called_once_with(x, norm.weight, norm.variance_epsilon) + assert actual_fp32 is normalized_fp32 + norm.assert_not_called() + + +def test_v41_hc_reference_supports_hidden_size_5120(): + torch.manual_seed(7) + layer = _layer() + x = torch.randn(2, 4, 5120, dtype=torch.bfloat16) + hc_fn = torch.randn(24, 4 * 5120, dtype=torch.float32) / 5120 + hc_scale = torch.randn(3, dtype=torch.float32) + hc_base = torch.randn(24, dtype=torch.float32) + + pre, post, comb = hc_mixes_reference(layer, x, hc_fn, hc_scale, hc_base) + y = layer.hc_collapse(x, pre) + + assert y.shape == (2, 5120) + assert y.dtype == torch.bfloat16 + assert post.shape == (2, 4) + assert post.dtype == torch.float32 + assert comb.shape == (2, 4, 4) + assert comb.dtype == torch.float32 + torch.testing.assert_close(comb.sum(-2), torch.ones(2, 4), atol=2e-5, rtol=2e-5) + + restored = hc_post_reference(y, x, post, comb) + assert restored.shape == x.shape + assert restored.dtype == x.dtype + + +def test_v41_hc_post_matches_reference_equation(): + torch.manual_seed(11) + x = torch.randn(3, 5, dtype=torch.bfloat16) + residual = torch.randn(3, 4, 5, dtype=torch.bfloat16) + post = torch.randn(3, 4, dtype=torch.float32) + comb = torch.randn(3, 4, 4, dtype=torch.float32) + + actual = hc_post_reference(x, residual, post, comb) + expected = (post.unsqueeze(-1) * x.unsqueeze(-2) + (comb.unsqueeze(-1) * residual.unsqueeze(-2)).sum(dim=-3)).to( + x.dtype + ) + torch.testing.assert_close(actual, expected) + + +def test_v41_hc_post_dispatches_fused_operator_with_batch_dimension(): + layer = _layer() + x = torch.randn(3, 5, dtype=torch.bfloat16) + residual = torch.randn(3, 4, 5, dtype=torch.bfloat16) + post = torch.randn(3, 4, dtype=torch.float32) + comb = torch.randn(3, 4, 4, dtype=torch.float32) + expected = torch.randn_like(residual).unsqueeze(0) + + with patch.object( + torch.ops._C_ascend, + "npu_hc_post", + create=True, + return_value=expected, + ) as op: + actual = layer.hc_post(x, residual, post, comb) + + torch.testing.assert_close(actual, expected.squeeze(0)) + op.assert_called_once() + for actual_arg, expected_arg in zip( + op.call_args.args, + ( + x.unsqueeze(0), + residual.unsqueeze(0), + post.unsqueeze(0), + comb.unsqueeze(0), + ), + ): + torch.testing.assert_close(actual_arg, expected_arg) + + +def test_v41_dspark_propagates_delayed_mix_and_collapses_final_stream(): + from vllm_ascend.models.deepseek_v41.dspark import DeepseekV41DSparkModel + + model = DeepseekV41DSparkModel.__new__(DeepseekV41DSparkModel) + torch.nn.Module.__init__(model) + model.hc_mult = 2 + model.needs_moe_input_ids = False + model.use_sequence_parallel = False + model.embed_tokens = torch.nn.Embedding(4, 3) + seen = [] + + class Layer(torch.nn.Module): + hc_collapse = staticmethod(DeepseekV41DecoderLayer.hc_collapse) + + def forward(self, positions, hidden, pre_mix, llama_4_scaling=None, input_ids=None): + seen.append(pre_mix.clone()) + return hidden + 1, pre_mix.flip(-1) + + model.layers = torch.nn.ModuleDict({"40": Layer(), "41": Layer(), "42": Layer()}) + ids = torch.tensor([0, 1]) + result = model(ids, torch.tensor([2, 3])) + torch.testing.assert_close(result, model.embed_tokens(ids) + 3) + assert [x[0].tolist() for x in seen] == [[1, 0], [0, 1], [1, 0]] + + +def test_v41_target_emits_input_residual_for_selected_aux_layers(): + from vllm_ascend.models.deepseek_v41.model import DeepseekV41Model + + model = DeepseekV41Model.__new__(DeepseekV41Model) + torch.nn.Module.__init__(model) + model.hc_mult = 2 + model.needs_moe_input_ids = False + model.embed_tokens = torch.nn.Embedding(4, 3) + model.norm = torch.nn.Identity() + model.shared_attention_state = MagicMock() + model._set_aux_hidden_state_layers((1, 3)) + + class Layer(torch.nn.Module): + hc_collapse = staticmethod(DeepseekV41DecoderLayer.hc_collapse) + + def __init__(self, idx): + super().__init__() + self.layer_idx = idx + self.engram = None + + def forward(self, positions, hidden, pre_mix, scaling, input_ids=None): + return hidden + 1, pre_mix + + model.layers = torch.nn.ModuleList([Layer(i) for i in range(3)]) + ids = torch.tensor([0, 1]) + with patch( + "vllm_ascend.models.deepseek_v41.model.get_pp_group", + return_value=MagicMock(is_first_rank=True, is_last_rank=True), + ): + output, aux = model.forward( + ids, + torch.tensor([0, 1]), + None, + engram_lookups={}, + engram_mask=torch.empty(0, dtype=torch.bool), + ) + embedded = model.embed_tokens(ids) + torch.testing.assert_close(output, embedded + 3) + assert len(aux) == 2 + torch.testing.assert_close(aux[0], embedded) + torch.testing.assert_close(aux[1], embedded + 2) + + +def test_v41_dspark_decoder_uses_draft_experts_instead_of_target_config(): + from contextlib import ExitStack + from types import SimpleNamespace + + import vllm_ascend.models.deepseek_v41.dspark as shared + from vllm_ascend.models.deepseek_v41.dspark import DeepseekV41DSparkModel + + draft = SimpleNamespace( + hc_mult=4, + hidden_size=8, + dspark_block_size=5, + num_nextn_predict_layers=3, + dspark_target_layer_ids=[37, 38, 39], + num_hidden_layers=40, + vocab_size=16, + rms_norm_eps=1e-6, + hc_eps=1e-6, + n_routed_experts=128, + num_experts_per_tok=3, + ) + config = SimpleNamespace( + model_config=SimpleNamespace(hf_config=SimpleNamespace(n_routed_experts=384)), + parallel_config=SimpleNamespace(use_sequence_parallel_moe=False), + speculative_config=SimpleNamespace(draft_model_config=SimpleNamespace(hf_text_config=draft)), + quant_config=None, + ) + + def make_layer(*args, **kwargs): + layer = torch.nn.Module() + layer.mlp = SimpleNamespace(gate=SimpleNamespace(tid2eid=None, bias_vl=None)) + return layer + + factory = MagicMock(side_effect=make_layer) + with ExitStack() as stack: + for name in ( + "VocabParallelEmbedding", + "ColumnParallelLinear", + "RMSNorm", + "DSparkMarkovHead", + "DSparkConfidenceHead", + ): + stack.enter_context(patch.object(shared, name, side_effect=lambda *args, **kwargs: torch.nn.Identity())) + stack.enter_context(patch.object(shared, "DeepseekV41DSparkDecoderLayer", factory)) + stack.enter_context(patch.object(shared, "validate_cache_runtime")) + model = DeepseekV41DSparkModel(vllm_config=config) + assert len(model.layers) == factory.call_count == 3 + for call in factory.call_args_list: + assert call.kwargs["config"] is draft + assert call.kwargs["config"].n_routed_experts == 128 + assert call.kwargs["is_draft_layer"] + assert config.model_config.hf_config.n_routed_experts == 384 diff --git a/tests/ut/models/test_deepseek_v41_layer_plan.py b/tests/ut/models/test_deepseek_v41_layer_plan.py new file mode 100644 index 000000000000..62b615a290eb --- /dev/null +++ b/tests/ut/models/test_deepseek_v41_layer_plan.py @@ -0,0 +1,78 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM Ascend project + +import pytest +import torch + +from vllm_ascend.models.deepseek_v41.model import ( + DeepseekV41SharedAttentionState, + build_layer_plan, +) + + +@pytest.fixture +def text_config() -> dict: + return { + "num_hidden_layers": 40, + "compress_ratios": [0, 0] + [2] * 18 + [1] * 20 + [0] * 3, + "kv_source_layers": [2, 8, 14, 20], + "index_source_layers": [2, 8, 14, 20, 24, 28, 32, 36], + "candidate_source_layer": 20, + "candidate_topk_blocks": 2048, + "candidate_block_size": 8, + "index_topk": 512, + "engram_layer_ids": [1, 14], + } + + +def test_builds_expected_source_groups(text_config: dict): + topology = build_layer_plan(text_config) + + assert topology.kv_consumers(2) == tuple(range(2, 8)) + assert topology.kv_consumers(8) == tuple(range(8, 14)) + assert topology.kv_consumers(14) == tuple(range(14, 20)) + assert topology.kv_consumers(20) == tuple(range(20, 40)) + + assert topology.index_consumers(20) == tuple(range(20, 24)) + assert topology.index_consumers(24) == tuple(range(24, 28)) + assert topology.index_consumers(36) == tuple(range(36, 40)) + + +def test_layer_26_resolves_layer_20_kv_and_layer_24_index(text_config: dict): + topology = build_layer_plan(text_config) + role = topology.layer(26) + assert role.kv_source_layer == 20 + assert role.index_source_layer == 24 + assert topology.candidate_source_layer == 20 + assert role.compress_ratio == 1 + assert role.uses_candidate_filter + + +def test_source_roles_and_engram_slots(text_config: dict): + topology = build_layer_plan(text_config) + + assert topology.layer(2).is_kv_source + assert topology.layer(2).is_index_source + assert topology.layer(20).is_candidate_source + assert topology.layer(1).engram_slot == 0 + assert topology.layer(14).engram_slot == 1 + assert topology.layer(0).kv_source_layer is None + + +def test_shared_state_resets_sparse_attention_metadata(): + topk_indices = torch.zeros((4, 1, 512), dtype=torch.int32) + candidates = torch.zeros((4, 1, 16), dtype=torch.int32) + state = DeepseekV41SharedAttentionState(topk_indices, candidates) + + state.reset() + + assert state.topk_indices is topk_indices + assert state.candidates is candidates + + +def test_rejects_ratio_mismatch(text_config: dict): + broken = dict(text_config) + broken["compress_ratios"] = list(text_config["compress_ratios"]) + broken["compress_ratios"][8] = 1 + with pytest.raises(ValueError, match="KV source"): + build_layer_plan(broken) diff --git a/tests/ut/models/test_deepseek_v41_operator_contracts.py b/tests/ut/models/test_deepseek_v41_operator_contracts.py new file mode 100644 index 000000000000..98cf9dfd8425 --- /dev/null +++ b/tests/ut/models/test_deepseek_v41_operator_contracts.py @@ -0,0 +1,53 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import ast +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[3] + + +def test_deepseek_v41_model_call_sites_use_compiled_operator_layouts(): + """Keep framework call sites aligned with the operator PR contracts.""" + source = ast.parse((REPO_ROOT / "vllm_ascend/attention/dsa_v41.py").read_text()) + calls = [ + node + for node in ast.walk(source) + if isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr in ("npu_sparse_flash_mla", "npu_sparse_flash_mla_metadata") + ] + assert len(calls) == 2 + for call in calls: + layouts = {kw.arg: ast.literal_eval(kw.value) for kw in call.keywords if kw.arg in ("layout_q", "layout_kv")} + assert layouts == {"layout_q": "TND", "layout_kv": "PA_BBND"} + + source = ast.parse((REPO_ROOT / "vllm_ascend/models/deepseek_v41/indexer.py").read_text()) + common = next( + node.value + for node in ast.walk(source) + if isinstance(node, ast.Assign) + and any(isinstance(target, ast.Name) and target.id == "common" for target in node.targets) + ) + assert isinstance(common, ast.Call) + layouts = {kw.arg: ast.literal_eval(kw.value) for kw in common.keywords if kw.arg in ("layout_q", "layout_k")} + assert layouts == {"layout_q": "TND", "layout_k": "PA_BBND"} + + +def test_compressor_call_sites_use_compiled_modes(): + for filename, count in ( + ("models/deepseek_v4/compressor.py", 1), + ("attention/context_parallel/dsa_cp.py", 2), + ): + source = ast.parse((REPO_ROOT / "vllm_ascend" / filename).read_text()) + calls = [ + node + for node in ast.walk(source) + if isinstance(node, ast.Call) and ast.unparse(node.func) == "torch.ops._C_ascend.compressor" + ] + assert len(calls) == count + for call in calls: + modes = { + kw.arg: ast.literal_eval(kw.value) for kw in call.keywords if kw.arg in ("rotary_mode", "cache_mode") + } + assert modes == {"rotary_mode": 2, "cache_mode": 1} diff --git a/tests/ut/models/test_deepseek_v41_preprocess.py b/tests/ut/models/test_deepseek_v41_preprocess.py new file mode 100644 index 000000000000..ad97a9369602 --- /dev/null +++ b/tests/ut/models/test_deepseek_v41_preprocess.py @@ -0,0 +1,292 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from contextlib import contextmanager +from types import SimpleNamespace +from typing import Any +from unittest.mock import Mock + +import pytest +import torch + +from vllm_ascend.attention import dsa_v41 +from vllm_ascend.attention.context_parallel import dsa_v41_cp + + +def test_dsa_v41_custom_op_forwards_its_output_buffer(monkeypatch): + hidden = torch.zeros(1, 8) + output = torch.empty_like(hidden) + impl = Mock() + attn = SimpleNamespace(v41_impl=impl) + monkeypatch.setattr( + dsa_v41, + "get_forward_context", + lambda: SimpleNamespace(no_compile_layers={"layer": attn}), + ) + + dsa_v41.dsa_v41_forward(hidden, output, "layer") + + impl.forward.assert_called_once_with(attn, None, hidden, output) + + +@pytest.mark.parametrize("share_quant", [False, True]) +@pytest.mark.parametrize("num_tokens", [1, 5]) +@pytest.mark.parametrize("cp", [False, True]) +@torch.inference_mode() +def test_preprocess_equivalence_and_stream_dependencies(monkeypatch, share_quant, num_tokens, cp): + """Check Q/qr/cache parity and the cross-stream producer/consumer ordering.""" + trace: list[tuple[str, str, str]] = [] + active = "main" + + class Stream: + def __init__(self, name): + self.name = name + + def record_event(self): + event = f"event{len(trace)}" + trace.append((self.name, "record", event)) + return event + + def wait_event(self, event): + trace.append((self.name, "wait", event)) + + def wait_stream(self, stream): + trace.append((self.name, "join", stream.name)) + + main, aux = Stream("main"), Stream("aux") + + @contextmanager + def switch(stream, *, enabled): + nonlocal active + previous = active + active = stream.name + try: + yield + finally: + active = previous + + class Wrapper: + _has_communication = False + + def __init__(self, name, linear, quant_method): + self.name, self.linear = name, linear + self._quant_method = quant_method + + def quantize(self, x): + trace.append((active, self.name, "quantize")) + return x, None + + def matmul(self, x, scale, bias=None): + trace.append((active, self.name, "matmul")) + assert bias is self.linear.bias + return self.linear(x) + + def norm(name): + def apply(x): + trace.append((active, name, "norm")) + return x * torch.rsqrt(x.square().mean(-1, keepdim=True) + 1e-6) + + return apply + + def rope(x, cos, sin, **kwargs): + trace.append((active, "rope", "apply")) + # Exercise the in-place write on a view, including Q/KV partial slices. + start, end = kwargs["partial_slice"] + assert x.shape[0] == cos.shape[0] == sin.shape[0] + x[..., start:end].mul_(cos).add_(sin) + + def scatter(cache, slots, values): + trace.append((active, "cache", "scatter")) + for slot, value in zip(slots, values): + if slot[0] >= 0: + cache[slot[0], slot[1]].copy_(value) + + monkeypatch.setattr(torch.npu, "current_stream", lambda: main) + monkeypatch.setattr(dsa_v41, "dsv4_dsa_overlap_stream", lambda: aux) + monkeypatch.setattr(dsa_v41, "npu_stream_switch", switch) + monkeypatch.setattr(dsa_v41, "scatter_cache_sk", scatter) + monkeypatch.setattr(dsa_v41_cp, "dsv4_dsa_overlap_stream", lambda: aux) + monkeypatch.setattr(dsa_v41_cp, "npu_stream_switch", switch) + monkeypatch.setattr(dsa_v41_cp, "scatter_cache_sk", scatter) + monkeypatch.setattr(torch.ops._C_ascend, "inplace_partial_rotary_mul", rope, raising=False) + torch.manual_seed(7) + cache = torch.zeros(2, num_tokens, 4) + q_a, q_b, kv = torch.nn.Linear(8, 6), torch.nn.Linear(6, 8), torch.nn.Linear(8, 4) + wrappers = SimpleNamespace( + cv_wq_a=Wrapper("qa", q_a, object()), + cv_wq_b=Wrapper("qb", q_b, object()), + cv_wkv=Wrapper("kv", kv, object() if share_quant else SimpleNamespace()), + ) + attn = SimpleNamespace( + wq_a=q_a, + wq_b=q_b, + wkv=kv, + q_norm=norm("q"), + kv_norm=norm("kv"), + n_local_heads=1 if cp else 2, + n_heads=2, + head_dim=4, + nope_head_dim=2, + dsa_attn=SimpleNamespace( + swa_cache_layer=SimpleNamespace(kv_cache=[cache]), + dsa_attn=SimpleNamespace(impl=wrappers), + ), + ) + slots = torch.tensor([[1, i] for i in range(num_tokens)]) + if num_tokens > 1: + slots[-1] = -1 + metadata = SimpleNamespace(slot_mapping=slots) + hidden = torch.randn(num_tokens, 8) + cos, sin = torch.randn(num_tokens, 1, 1, 2), torch.randn(num_tokens, 1, 1, 2) + start = num_tokens // 2 if cp else 0 + local_hidden, local_cos, local_sin = hidden[start:], cos[start:], sin[start:] + impl = object.__new__(dsa_v41_cp.DeepseekV41CPImpl if cp else dsa_v41.DeepseekV41EagerAttentionImpl) + if cp: + expected_q, expected_qr = impl._project_q(attn, local_hidden, local_cos, local_sin) + scatter(cache, slots, impl._project_kv(attn, hidden, cos, sin)) + else: + expected_q, expected_qr = impl.preprocess(attn, hidden, cos, sin, metadata) + expected_cache = cache.clone() + cache.zero_() + trace.clear() + kwargs: dict[str, Any] = {} + if cp: + impl.role = SimpleNamespace(is_kv_source=False) + attn.rotary_emb = SimpleNamespace(layername="layer") + metadata.num_actual_tokens = num_tokens + global_metadata = SimpleNamespace(swa=metadata, rope=lambda *args: (cos, sin)) + metadata = SimpleNamespace( + swa=SimpleNamespace( + num_actual_tokens=num_tokens - start, cp_token_range=(start, num_tokens, num_tokens - start, num_tokens) + ) + ) + impl._global_layer_metadata = Mock(return_value=global_metadata) + monkeypatch.setattr(dsa_v41_cp, "get_forward_context", lambda: SimpleNamespace(attn_metadata={})) + metadata = metadata.swa + q, qr = impl.multistream_preprocess(attn, hidden, local_cos, local_sin, metadata, **kwargs) + torch.testing.assert_close(q, expected_q) + torch.testing.assert_close(qr, expected_qr) + torch.testing.assert_close(cache, expected_cache) + assert qr.is_floating_point() + assert (("aux", "kv", "quantize") in trace) == (cp or not share_quant) + assert trace.count(("aux", "cache", "scatter")) == 1 + kv_mm = trace.index(("aux", "kv", "matmul")) + kv_done = trace[kv_mm + 1] + assert kv_done[:2] == ("aux", "record") + assert trace.index(("main", "wait", kv_done[2])) < trace.index(("main", "qb", "matmul")) + part3 = trace[trace.index(("main", "wait", kv_done[2])) - 1] + assert part3[:2] == ("main", "record") + assert trace.index(("aux", "wait", part3[2])) < trace.index(("aux", "kv", "norm")) + assert trace.index(("aux", "cache", "scatter")) < trace.index(("main", "join", "aux")) + assert trace.index(("main", "join", "aux")) < trace.index(("main", "rope", "apply")) + + +@pytest.mark.parametrize("enabled", [False, True]) +def test_forward_selects_preprocess_from_v1_switch(monkeypatch, enabled): + impl = object.__new__(dsa_v41.DeepseekV41EagerAttentionImpl) + impl.role = SimpleNamespace(is_kv_source=False) + hidden = torch.zeros(1, 8) + q, qr = torch.zeros(1, 2, 4), torch.zeros(1, 6) + metadata = SimpleNamespace( + positions=torch.zeros(1), swa=SimpleNamespace(num_actual_tokens=1), rope=lambda *args: (None, None) + ) + impl._get_layer_metadata = Mock(return_value=metadata) + impl.preprocess = Mock(return_value=(q, qr)) + impl.multistream_preprocess = Mock(return_value=(q, qr)) + impl._select_sparse_indices = Mock(return_value=None) + impl._attention = Mock(return_value=q) + v1_impl = SimpleNamespace( + multistream_dsv4_dsa_overlap=enabled, + _forward_o_proj=lambda q, output: output.zero_(), + ) + attn = SimpleNamespace( + rotary_emb=SimpleNamespace(layername="layer"), + dsa_attn=SimpleNamespace(dsa_attn=SimpleNamespace(impl=v1_impl)), + nope_head_dim=2, + head_dim=4, + ) + metadata.rope = lambda *args: (torch.zeros(1), torch.zeros(1)) + monkeypatch.setattr(dsa_v41, "get_forward_context", lambda: SimpleNamespace(attn_metadata={})) + monkeypatch.setattr(torch.ops._C_ascend, "inplace_partial_rotary_mul", lambda *args, **kwargs: None, raising=False) + output = torch.full_like(hidden, 1) + result = impl.forward(attn, None, hidden, output) + selected = impl.multistream_preprocess if enabled else impl.preprocess + unused = impl.preprocess if enabled else impl.multistream_preprocess + selected.assert_called_once() + unused.assert_not_called() + assert result is output + assert torch.count_nonzero(output) == 0 + + +@pytest.mark.parametrize("local_tokens", [0, 2]) +@pytest.mark.parametrize("is_source", [False, True]) +def test_cp_multistream_forward_preserves_full_cache_updates(monkeypatch, local_tokens, is_source): + impl = object.__new__(dsa_v41_cp.DeepseekV41CPImpl) + impl.role = SimpleNamespace(is_kv_source=is_source) + hidden = torch.arange(40, dtype=torch.float32).reshape(5, 8) + global_cos, global_sin = torch.ones(4), torch.ones(4) + local_cos, local_sin = global_cos[2 : 2 + local_tokens], global_sin[2 : 2 + local_tokens] + metadata = SimpleNamespace( + swa=SimpleNamespace(num_actual_tokens=local_tokens, cp_token_range=(2, 4, 2, 4)), + positions=torch.arange(2, 2 + local_tokens), + rope=lambda *args: (local_cos, local_sin), + ) + global_metadata = SimpleNamespace( + swa=SimpleNamespace(num_actual_tokens=4), + positions=torch.arange(4), + rope=lambda *args: (global_cos, global_sin), + ) + impl._get_layer_metadata = Mock(return_value=metadata) + impl._global_layer_metadata = Mock(return_value=global_metadata) + q, qr = torch.zeros(local_tokens, 2, 4), torch.zeros(local_tokens, 6) + impl.multistream_preprocess = Mock(return_value=(q, qr)) + impl._update_caches = Mock() + impl._write_compressed_source = Mock() + impl._select_sparse_indices = Mock(return_value=None) + impl._attention = Mock(return_value=q) + impl._project_output = Mock(side_effect=lambda *args, projected: projected.zero_()) + attn = SimpleNamespace( + rotary_emb=SimpleNamespace(layername="layer"), + n_heads=2, + enable_dsa_cp=True, + head_dim=4, + nope_head_dim=2, + dsa_attn=SimpleNamespace(dsa_attn=SimpleNamespace(impl=SimpleNamespace(multistream_dsv4_dsa_overlap=True))), + ) + monkeypatch.setattr(dsa_v41, "get_forward_context", lambda: SimpleNamespace(attn_metadata={})) + monkeypatch.setattr(torch.ops._C_ascend, "inplace_partial_rotary_mul", lambda *args, **kwargs: None, raising=False) + output = torch.ones_like(hidden) + assert impl.forward(attn, None, hidden, output) is output + impl._project_output.assert_called_once() + assert torch.count_nonzero(output) == 0 + if local_tokens: + impl._update_caches.assert_not_called() + impl.multistream_preprocess.assert_called_once() + args, kwargs = impl.multistream_preprocess.call_args + torch.testing.assert_close(args[1], hidden) + assert args[2] is local_cos and args[3] is local_sin + assert args[4] is metadata.swa + assert not kwargs + # The mocked preprocessor owns compressed-cache writes. + impl._write_compressed_source.assert_not_called() + assert impl._select_sparse_indices.call_args.args[-1] is metadata + else: + impl.multistream_preprocess.assert_not_called() + impl._update_caches.assert_called_once() + args = impl._update_caches.call_args.args + torch.testing.assert_close(args[1], hidden[:4]) + assert args[2] is global_metadata + impl._write_compressed_source.assert_not_called() + impl._attention.assert_not_called() + + +def test_cp_disabled_overlap_delegates_to_serial_forward(monkeypatch): + impl = object.__new__(dsa_v41_cp.DeepseekV41CPImpl) + serial = Mock(return_value=object()) + monkeypatch.setattr(dsa_v41.DeepseekV41EagerAttentionImpl, "forward", serial) + attn = SimpleNamespace( + dsa_attn=SimpleNamespace(dsa_attn=SimpleNamespace(impl=SimpleNamespace(multistream_dsv4_dsa_overlap=False))) + ) + hidden, output = torch.zeros(2, 8), torch.empty(2, 8) + assert impl.forward(attn, None, hidden, output) is serial.return_value + serial.assert_called_once_with(attn, None, hidden, output) diff --git a/tests/ut/models/test_deepseek_v41_release_config.py b/tests/ut/models/test_deepseek_v41_release_config.py new file mode 100644 index 000000000000..31033e920282 --- /dev/null +++ b/tests/ut/models/test_deepseek_v41_release_config.py @@ -0,0 +1,119 @@ +import pytest +from vllm import ModelRegistry + +from vllm_ascend.deepseek_v41_config import DeepseekV41Config +from vllm_ascend.models import register_model + + +def _released_text_config(): + return { + "model_type": "deepseek_v41_text", + "num_hidden_layers": 40, + "kv_source_layer_ids": [2, 8, 14, 20], + "index_source_layer_ids": [2, 8, 14, 20, 24, 28, 32, 36], + "candidate_source_layer_id": 20, + "engram_pad_token_id": 2, + "dspark_n_routed_experts": 128, + "dspark_num_experts_per_tok": 3, + } + + +def _rotation_config(): + return { + "value_projection_rotated": True, + "value_basis": "quarot_global", + "key_and_gate_basis": "original", + "runtime_delta_rotation": False, + } + + +def test_released_config_names_are_available_to_existing_runtime(): + config = DeepseekV41Config( + architectures=["DeepseekV41ForCausalLM"], + text_config=_released_text_config(), + vision_config={ + "model_type": "deepseek_v41_vision", + "num_hidden_layers": 32, + "max_image_tokens": 1024, + }, + engram_rotation_config=_rotation_config(), + ) + + assert config.model_type == "deepseek_v41" + assert config.text_config.model_type == "deepseek_v41_text" + assert config.kv_source_layers == config.kv_source_layer_ids + assert config.index_source_layers == config.index_source_layer_ids + assert config.candidate_source_layer == config.candidate_source_layer_id == 20 + assert config.engram_pad_id == config.engram_pad_token_id == 2 + assert config.dspark_n_activated_experts == 3 + assert config.dspark_num_experts_per_tok == 3 + assert config.vision_max_n_token == config.vision_config.max_image_tokens == 1024 + assert config.vision_config.max_num_tokens == 1024 + assert config.engram_rotation_config == _rotation_config() + # The released checkpoint renamed the architecture to CausalLM but still + # carries and serves the complete vision path. + assert config.is_mm_prefix_lm + assert config.mm_prefix_clamp_sliding_window + assert config.mm_prefix_span_leading_pad_modulus == 2 + + +def test_released_causal_architecture_uses_multimodal_wrapper(monkeypatch): + calls = [] + monkeypatch.setattr( + ModelRegistry, + "register_model", + lambda *args, **kwargs: calls.append((args, kwargs)), + ) + + register_model() + + assert any( + args + == ( + "DeepseekV41ForCausalLM", + "vllm_ascend.models.deepseek_v41.vl_model:AscendDeepseekV41ForConditionalGeneration", + ) + for args, _kwargs in calls + ) + + +def test_legacy_conditional_config_remains_supported(): + config = DeepseekV41Config( + architectures=["DeepseekV41ForConditionalGeneration"], + text_config={ + "model_type": "deepseek_v4.1_text", + "kv_source_layers": [2], + "index_source_layers": [2], + "candidate_source_layer": 2, + "engram_pad_id": 2, + "dspark_n_activated_experts": 3, + }, + vision_config={ + "model_type": "deepseek_v4.1_vision", + "num_hidden_layers": 32, + "max_num_tokens": 1024, + }, + ) + + assert config.text_config.kv_source_layer_ids == [2] + assert config.text_config.engram_pad_token_id == 2 + assert config.text_config.dspark_num_experts_per_tok == 3 + assert config.is_mm_prefix_lm + assert config.mm_prefix_span_leading_pad_modulus == 2 + + +def test_release_and_legacy_aliases_cannot_disagree(): + text = _released_text_config() + text["kv_source_layers"] = [3] + with pytest.raises(ValueError, match="Conflicting DeepSeek V4.1 config fields"): + DeepseekV41Config(text_config=text) + + +def test_unknown_engram_rotation_contract_fails_closed(): + rotation = _rotation_config() + rotation["runtime_delta_rotation"] = True + with pytest.raises(ValueError, match="Unsupported DeepSeek V4.1 Engram rotation"): + DeepseekV41Config( + text_config=_released_text_config(), + engram_rotation_config=rotation, + ) diff --git a/tests/ut/models/test_deepseek_v41_vision.py b/tests/ut/models/test_deepseek_v41_vision.py new file mode 100644 index 000000000000..8e595360ba93 --- /dev/null +++ b/tests/ut/models/test_deepseek_v41_vision.py @@ -0,0 +1,156 @@ +from types import SimpleNamespace + +import torch +from PIL import Image +from torch import nn +from vllm.model_executor.models.interfaces import supports_multimodal +from vllm.multimodal.processing import InputProcessingContext + +from vllm_ascend.deepseek_v41_config import DeepseekV41Config +from vllm_ascend.models.deepseek_v41.engram_hash import ( + valid_engram_token_mask, +) +from vllm_ascend.models.deepseek_v41.mm_preprocess import ( + COMPRESS_PAD_TO, + IMAGE, + IMAGE_END, + IMAGE_NEW_LINE, + IMAGE_PAD_ID, + IMAGE_START, + IMAGE_TOKEN_ID, + DeepseekV41VLProcessingInfo, + DeepseekV41VLProcessor, + image_sentinel_mask, + image_token_types, + leading_compressor_pad, +) +from vllm_ascend.models.deepseek_v41.model import AscendDeepseekV41ForCausalLM +from vllm_ascend.models.deepseek_v41.vl_model import ( + AscendDeepseekV41ForConditionalGeneration, +) + + +def test_v41_vision_wrapper_uses_v41_language_backbone(): + assert supports_multimodal(AscendDeepseekV41ForConditionalGeneration) + assert AscendDeepseekV41ForConditionalGeneration.language_model_cls is AscendDeepseekV41ForCausalLM + assert "_processor_factory" in AscendDeepseekV41ForConditionalGeneration.__dict__ + + +def test_v41_processing_info_accepts_v41_config(): + config = DeepseekV41Config( + text_config={}, + vision_config={ + "num_hidden_layers": 1, + "patch_size": 14, + "downsample_ratio": 3, + "max_num_tokens": 1024, + }, + ) + model_config = SimpleNamespace(hf_config=config) + ctx = InputProcessingContext(model_config=model_config, tokenizer=None) + + assert DeepseekV41VLProcessingInfo(ctx).get_hf_config() is config + assert config.image_sentinel_base_id == IMAGE_TOKEN_ID + assert config.image_pad_token_id == IMAGE_PAD_ID + assert config.is_mm_prefix_lm + assert config.mm_prefix_clamp_sliding_window + assert config.mm_prefix_span_leading_pad_modulus == COMPRESS_PAD_TO == 2 + + +def test_v41_image_roles_use_reference_reading_order(): + assert image_token_types(2, 3).tolist() == [ + IMAGE_START, + IMAGE, + IMAGE, + IMAGE, + IMAGE_NEW_LINE, + IMAGE, + IMAGE, + IMAGE, + IMAGE_NEW_LINE, + IMAGE_END, + ] + + assert [leading_compressor_pad(i) for i in range(4)] == [1, 0, 1, 0] + + +def test_v41_processor_emits_types_without_v4_perm(): + config = DeepseekV41Config( + text_config={}, + vision_config={ + "num_hidden_layers": 1, + "hidden_size": 16, + "num_attention_heads": 2, + "intermediate_size": 32, + "patch_size": 14, + "rope_theta": 10000.0, + "downsample_ratio": 3, + "max_num_tokens": 64, + "min_pixels": 0, + "max_wh_ratio": None, + }, + ) + result = DeepseekV41VLProcessor(config)(images=[Image.new("RGB", (84, 42))]) + + assert result["vit_grid"].tolist() == [[3, 6]] + assert result["llm_grid"].tolist() == [[1, 2]] + assert result["types"].tolist() == [ + IMAGE_START, + IMAGE, + IMAGE, + IMAGE_NEW_LINE, + IMAGE_END, + ] + assert "perm" not in result + + +def test_v41_image_and_alignment_pad_are_dead_to_engram(): + token_ids = torch.tensor([17, IMAGE_TOKEN_ID, IMAGE_PAD_ID, 18]) + expected = torch.tensor([True, False, False, True]) + + torch.testing.assert_close(image_sentinel_mask(token_ids), ~expected) + torch.testing.assert_close( + valid_engram_token_mask(token_ids, IMAGE_TOKEN_ID, IMAGE_PAD_ID), + expected, + ) + + +def test_v41_span_has_three_delimiters_and_no_image_pad_parameter(): + wrapper = object.__new__(AscendDeepseekV41ForConditionalGeneration) + nn.Module.__init__(wrapper) + wrapper.image_start = nn.Parameter(torch.tensor([1.0, 1.0])) + wrapper.image_newline = nn.Parameter(torch.tensor([2.0, 2.0])) + wrapper.image_end = nn.Parameter(torch.tensor([3.0, 3.0])) + image_embeds = torch.tensor([[10.0, 10.0], [20.0, 20.0]]) + + span = wrapper._build_image_span( + image_embeds, + image_token_types(1, 2), + ) + + torch.testing.assert_close( + span, + torch.tensor( + [ + [1.0, 1.0], + [10.0, 10.0], + [20.0, 20.0], + [2.0, 2.0], + [3.0, 3.0], + ] + ), + ) + assert "image_pad" not in dict(wrapper.named_parameters()) + + +def test_v41_alignment_pad_uses_plain_image_token_embedding(): + class LanguageModel(nn.Module): + def embed_input_ids(self, input_ids): + return input_ids.unsqueeze(-1) + + wrapper = object.__new__(AscendDeepseekV41ForConditionalGeneration) + nn.Module.__init__(wrapper) + wrapper.language_model = LanguageModel() + + embeddings = wrapper.embed_input_ids(torch.tensor([7, IMAGE_PAD_ID, IMAGE_TOKEN_ID])) + assert embeddings.squeeze(-1).tolist() == [7, IMAGE_TOKEN_ID, IMAGE_TOKEN_ID] diff --git a/tests/ut/models/test_deepseek_v4_moe.py b/tests/ut/models/test_deepseek_v4_moe.py index a09811fbe19f..3148f0e30e36 100644 --- a/tests/ut/models/test_deepseek_v4_moe.py +++ b/tests/ut/models/test_deepseek_v4_moe.py @@ -36,6 +36,80 @@ def forward(self, hidden_states, router_logits, input_ids=None): return hidden_states +@pytest.mark.parametrize("is_internal_router", [False, True]) +@pytest.mark.parametrize("is_sequence_parallel", [False, True]) +@pytest.mark.parametrize("has_fp32_input", [False, True]) +def test_deepseek_v4_moe_reuses_fp32_input_on_matching_token_shard( + monkeypatch, is_internal_router, is_sequence_parallel, has_fp32_input +): + hidden_states = torch.randn(4, 8, dtype=torch.bfloat16) + fp32_input = hidden_states.float() if has_fp32_input else None + input_ids = torch.tensor([11, 22, 13, 24]) + experts = MagicMock(side_effect=lambda **kwargs: kwargs["hidden_states"]) + experts.is_internal_router = is_internal_router + gate = SimpleNamespace(tid2eid=torch.zeros(32, 2), weight=torch.randn(3, 8)) + moe = SimpleNamespace(gate=gate, experts=experts, is_sequence_parallel=is_sequence_parallel, tp_size=1) + monkeypatch.setattr(deepseek_v4_module, "sequence_parallel_chunk", lambda x: x[2:]) + monkeypatch.setattr(deepseek_v4_module, "sp_all_gather", lambda x: torch.cat([x, x])) + linear = MagicMock(wraps=torch.nn.functional.linear) + monkeypatch.setattr(deepseek_v4_module.F, "linear", linear) + + deepseek_v4_module.DeepseekV4MoE.forward(moe, hidden_states, input_ids, fp32_input) + + kwargs = experts.call_args.kwargs + expected_hidden = hidden_states[2:] if is_sequence_parallel else hidden_states + torch.testing.assert_close(kwargs["hidden_states"], expected_hidden, rtol=0, atol=0) + assert kwargs["input_ids"] is input_ids + if is_internal_router: + router_input = kwargs["router_logits"] + linear.assert_not_called() + else: + linear.assert_called_once() + router_input = linear.call_args.args[0] + torch.testing.assert_close(kwargs["router_logits"], expected_hidden.float() @ gate.weight.T) + torch.testing.assert_close(router_input.float(), expected_hidden.float(), rtol=0, atol=0) + if has_fp32_input: + expected_fp32 = fp32_input[2:] if is_sequence_parallel else fp32_input + assert router_input.dtype == torch.float32 + assert router_input.data_ptr() == expected_fp32.data_ptr() + + +def test_deepseek_v4_dsa_cp_keeps_moe_input_sequence_parallel(): + layer = deepseek_v4_module.DeepseekV2DecoderLayer.__new__(deepseek_v4_module.DeepseekV2DecoderLayer) + nn.Module.__init__(layer) + layer.use_sequence_parallel_moe = True + layer.enable_dsa_cp = True + layer.hc_attn_fn = nn.Parameter(torch.empty(1)) + layer.hc_attn_scale = nn.Parameter(torch.empty(1)) + layer.hc_attn_base = nn.Parameter(torch.empty(1)) + layer.hc_ffn_fn = nn.Parameter(torch.empty(1)) + layer.hc_ffn_scale = nn.Parameter(torch.empty(1)) + layer.hc_ffn_base = nn.Parameter(torch.empty(1)) + + hidden_states = torch.randn(2, 4, 8) + collapsed = torch.randn(2, 8) + layer.hc_pre = MagicMock( + side_effect=[ + (collapsed, torch.empty(0), torch.empty(0)), + (collapsed, torch.empty(0), torch.empty(0)), + ] + ) + layer.input_layernorm = MagicMock(side_effect=lambda value: value) + layer.self_attn = MagicMock(side_effect=lambda **kwargs: kwargs["hidden_states"]) + layer.hc_post = MagicMock(side_effect=lambda value, *_args: value) + layer.rms_norm_cast = MagicMock(return_value=(collapsed, collapsed.float())) + layer.mlp = MagicMock(return_value=collapsed) + + layer.forward( + torch.arange(2), + hidden_states, + None, + input_ids=torch.tensor([11, 22]), + ) + + assert layer.mlp.call_args.kwargs["already_sequence_parallel"] is True + + def test_deepseek_v4_hash_layer_uses_upstream_hash_router(monkeypatch): gate = _FakeGate() diff --git a/tests/ut/models/test_deepseek_v4_routing_inputs.py b/tests/ut/models/test_deepseek_v4_routing_inputs.py new file mode 100644 index 000000000000..75246a04c13f --- /dev/null +++ b/tests/ut/models/test_deepseek_v4_routing_inputs.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM Ascend project + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest +import torch + +from vllm_ascend.models.deepseek_v4 import model as deepseek_v4_module +from vllm_ascend.models.deepseek_v4.dspark import DeepseekV4DSparkModel +from vllm_ascend.models.deepseek_v4.model import DeepseekV4Model +from vllm_ascend.models.deepseek_v41 import model as deepseek_v41_module +from vllm_ascend.models.deepseek_v41.dspark import DeepseekV41DSparkModel +from vllm_ascend.models.deepseek_v41.model import DeepseekV41Model + + +class _CaptureRoutingLayer: + def __init__(self, index): + self.layer_idx = index + self.engram = None + self.input_ids = [] + + def __call__(self, positions, hidden_states, residual, llama_4_scaling=None, input_ids=None): + self.input_ids.append(input_ids) + return hidden_states, residual + + @staticmethod + def hc_collapse(hidden_states, pre_mix): + return hidden_states.mean(dim=1) + + +@pytest.mark.parametrize( + "model_cls", + [DeepseekV4Model, DeepseekV41Model, DeepseekV4DSparkModel, DeepseekV41DSparkModel], +) +@pytest.mark.parametrize("input_dtype", [torch.int32, torch.int64]) +@pytest.mark.parametrize("needs_moe_input_ids", [False, True]) +def test_model_prepares_routing_ids_once_per_forward(monkeypatch, model_cls, input_dtype, needs_moe_input_ids): + pp_group = SimpleNamespace(is_first_rank=True, is_last_rank=True) + monkeypatch.setattr(deepseek_v4_module, "get_pp_group", lambda: pp_group) + monkeypatch.setattr(deepseek_v41_module, "get_pp_group", lambda: pp_group) + monkeypatch.setattr(deepseek_v4_module, "get_pp_transport_tensors", lambda *_args: []) + layers = [_CaptureRoutingLayer(i) for i in range(3)] + is_draft = model_cls in (DeepseekV4DSparkModel, DeepseekV41DSparkModel) + hidden_states = torch.randn(4, 8) + embed = MagicMock(return_value=hidden_states) + prepare_engram = MagicMock(return_value=({}, torch.empty(0, dtype=torch.bool))) + model = SimpleNamespace( + needs_moe_input_ids=needs_moe_input_ids, + use_sequence_parallel_moe=False, + use_sequence_parallel=False, + hc_mult=4, + layers={str(i): layer for i, layer in enumerate(layers)} if is_draft else layers, + start_layer=0, + end_layer=len(layers), + aux_hidden_state_layers=(), + _mtp_hidden_buffer=None, + embed_input_ids=embed, + embed_tokens=embed, + prepare_engram=prepare_engram, + shared_attention_state=SimpleNamespace(reset=MagicMock()), + hc_head=lambda hidden, *_args: hidden.mean(dim=1), + hc_head_fn=None, + hc_head_scale=None, + hc_head_base=None, + norm=lambda hidden: hidden, + ) + positions = torch.arange(4) + # A new draft substep must use the new IDs, including new placeholder rows. + for step, values in enumerate(([-1, 0, 129259, 15], [13, -1, 129257, -1])): + input_ids = torch.tensor(values, dtype=input_dtype) + original = input_ids.clone() + with patch.object(torch, "where", wraps=torch.where) as where: + if is_draft: + model_cls.forward(model, input_ids, positions) + else: + model_cls.forward(model, input_ids, positions, None) + + assert where.call_count == int(needs_moe_input_ids) + routed_ids = layers[0].input_ids[step] + for layer in layers: + assert layer.input_ids[step] is routed_ids + expected = torch.tensor([0 if value == -1 else value for value in values], dtype=input_dtype) + if needs_moe_input_ids: + torch.testing.assert_close(routed_ids, expected, rtol=0, atol=0) + assert routed_ids is not input_ids + else: + assert routed_ids is input_ids + torch.testing.assert_close(input_ids, original, rtol=0, atol=0) + assert embed.call_args.args[0] is input_ids + if model_cls is DeepseekV41Model: + assert prepare_engram.call_args.args[0] is input_ids diff --git a/tests/ut/models/test_deepseek_v4_vision_preprocess.py b/tests/ut/models/test_deepseek_v4_vision_preprocess.py index d71a03a4f617..4bea78bf3da0 100644 --- a/tests/ut/models/test_deepseek_v4_vision_preprocess.py +++ b/tests/ut/models/test_deepseek_v4_vision_preprocess.py @@ -165,3 +165,13 @@ def process(): competing_call.result() assert all(torch.equal(output, torch.tensor([[1]])) for output in outputs) + + +def test_hf_processor_accepts_base_class_call_signature(): + info = _ConcurrentStubInfo() + info.tokenizer = _NonThreadSafeTokenizer() + processor = DeepseekV4VLMultiModalProcessor(info, None) + + output = processor._call_hf_processor("prompt", {"images": []}, {}) + + assert torch.equal(output["input_ids"], torch.tensor([[1]])) diff --git a/tests/ut/models/test_engram_hbm.py b/tests/ut/models/test_engram_hbm.py new file mode 100644 index 000000000000..ecb74712a0cc --- /dev/null +++ b/tests/ut/models/test_engram_hbm.py @@ -0,0 +1,300 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Run standalone: pytest --confcutdir=tests/ut/models test_engram_hbm.py.""" + +import importlib.util +import json +import sys +from datetime import timedelta +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +from safetensors.torch import save_file + +ROOT = Path(__file__).resolve().parents[3] + + +def load_module(name): + spec = importlib.util.spec_from_file_location(name, ROOT / f"vllm_ascend/models/deepseek_v41/{name}.py") + if spec is None or spec.loader is None: + raise ImportError(f"Unable to load {name}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +hbm = load_module("engram_hbm") +gate_mod = load_module("engram_gate") +hash_mod = load_module("engram_hash") +gate = gate_mod.engram_gate + + +@pytest.mark.parametrize("cp", ["none", "v41", "v41_empty_rank"]) +def test_engram_history_metadata_uses_full_requests(cp): + boundaries = torch.tensor([0, 3, 9], dtype=torch.int32) + pages = torch.tensor([[7, 8, 9], [12, 13, 14]], dtype=torch.int32) + # Device tensors are intentionally unusable: history must use host mirrors. + fields = dict( + query_start_loc=None, + block_table=None, + storage_block_size=4, + query_start_loc_cpu=boundaries, + block_table_cpu=pages, + ) + request = SimpleNamespace(**fields) + if cp.startswith("v41"): + local = torch.tensor([0, 0, 0] if cp == "v41_empty_rank" else [0, 0, 2]) + metadata = SimpleNamespace( + global_metadata=request, + query_start_loc=local, + query_start_loc_cpu=local, + block_table=None, + storage_block_size=4, + ) + else: + metadata = request + actual_boundaries, actual_pages, block_size = hash_mod.engram_history_metadata(metadata) + assert torch.equal(actual_boundaries, boundaries.long()) + assert torch.equal(actual_pages, pages) + assert block_size == 4 + + +@pytest.mark.parametrize("missing", ["query_start_loc_cpu", "block_table_cpu"]) +def test_engram_history_metadata_requires_cpu_mirrors(missing): + metadata = SimpleNamespace( + query_start_loc_cpu=torch.tensor([0, 1]), + block_table_cpu=torch.tensor([[7]]), + query_start_loc=None, + block_table=None, + storage_block_size=4, + ) + setattr(metadata, missing, None) + with pytest.raises(ValueError, match="requires query_start_loc_cpu and block_table_cpu"): + hash_mod.engram_history_metadata(metadata) + + +def _worker(rank, rendezvous): + torch.set_num_threads(1) + dist.init_process_group("gloo", init_method=rendezvous, rank=rank, world_size=4, timeout=timedelta(seconds=60)) + tp = None + for ranks in ([0, 1], [2, 3]): + group = dist.new_group(ranks) + if rank in ranks: + tp = group + q = hbm.EngramQueryGroup(dist.group.WORLD, dist.group.WORLD, tp, (rank // 2) * 2) + # Deliberately not divisible by group size; includes signed zero and NaN bits. + weights = torch.arange(29 * 8).reshape(29, 8).to(torch.bfloat16) + weights.view(torch.int16)[0, 0] = -32768 + weights.view(torch.int16)[28, 7] = 32705 + table = hbm.NodeShardedEngram(29, 8, q, device="cpu") + table.weight.data.copy_(weights[table.start : table.end]) + for a, b in [(0, 0), (1, 0), (0, 17), (3, 7), (31, 1), (1, 1), (0, 0)]: + n = (a, b)[rank // 2] + ids = (torch.arange(n * 3).reshape(n, 3) * 7 + rank // 2) % 29 + actual = table(ids) + expected = weights[ids] + assert torch.equal(actual.view(torch.int16), expected.view(torch.int16)), (rank, a, b) + bad = torch.tensor([[29 if rank // 2 == 1 else 0]]) + with pytest.raises(IndexError): + table(bad) + assert torch.equal( + table(torch.tensor([[1, 8, 28]])).view(torch.int16), weights[torch.tensor([[1, 8, 28]])].view(torch.int16) + ) + ids = torch.tensor([[1, 8, 28]]) + packed = table.forward_many([ids, ids]) + assert torch.equal(packed[0].view(torch.int16), weights[ids].view(torch.int16)) + assert torch.equal(packed[1].view(torch.int16), weights[ids].view(torch.int16)) + single = table.route_many([table], [ids])[0] + assert torch.equal(single.view(torch.int16), weights[ids].view(torch.int16)) + # Distinct row counts exercise different shard boundaries in one exchange. + other_weights = -torch.arange(37 * 8).reshape(37, 8).bfloat16() + other = hbm.NodeShardedEngram(37, 8, q, device="cpu") + other.weight.data.copy_(other_weights[other.start : other.end]) + for a, b in [(0, 0), (1, 0), (0, 17), (31, 1), (1, 31), (3, 7)]: + n = (a, b)[rank // 2] + first_ids = (torch.arange(n * 3).reshape(n, 3) * 7 + rank // 2) % 29 + second_ids = (torch.arange(n * 3).reshape(n, 3) * 11 + rank // 2) % 37 + first, second = table.route_many([table, other], [first_ids, second_ids]) + assert torch.equal(first.view(torch.int16), weights[first_ids].view(torch.int16)) + assert torch.equal(second.view(torch.int16), other_weights[second_ids].view(torch.int16)) + with pytest.raises(IndexError): + table.route_many([table, other], [ids, torch.tensor([[37 if rank // 2 else 0]])]) + recovered = table.route_many([table, other], [ids, ids]) + assert torch.equal(recovered[1].view(torch.int16), other_weights[ids].view(torch.int16)) + dist.destroy_process_group() + + +def _compressed_wire_worker(rank, rendezvous): + torch.set_num_threads(1) + dist.init_process_group("gloo", init_method=rendezvous, rank=rank, world_size=4, timeout=timedelta(seconds=60)) + for ranks in ([0, 1], [2, 3]): + group = dist.new_group(ranks) + if rank in ranks: + tp = group + q = hbm.EngramQueryGroup(dist.group.WORLD, dist.group.WORLD, tp, rank // 2 * 2) + table = hbm.NodeShardedEngram(37, 256, q, device="cpu", storage_format="int8") + codes = (torch.arange(37 * 256).reshape(37, 256) % 256 - 128).to(torch.int8) + scales = torch.linspace(0.001, 0.9, 37 * 8).reshape(37, 8) + codes[0] = 0 + table.weight.data.copy_(codes[table.start : table.end]) + table.weight_scale.copy_(scales[table.start : table.end]) + reference = (codes.float().reshape(37, 8, 32) * scales[..., None]).reshape(37, 256).bfloat16() + table.compressed_int8_wire = True + for a, b in ((0, 0), (1, 0), (0, 17), (3, 7), (128, 1), (1, 128)): + n = (a, b)[rank // 2] + ids = (torch.arange(n * 24).reshape(n, 24) * 7 + rank // 2) % 37 + if n: + ids[0, :3] = torch.tensor([0, 36, 0]) + actual = table(ids) + assert torch.equal(actual.view(torch.int16), reference[ids].view(torch.int16)), (rank, a, b) + with pytest.raises(IndexError): + table(torch.tensor([[37 if rank // 2 else 0]])) + ids = torch.tensor([[0, 36, 1]]) + assert torch.equal(table(ids).view(torch.int16), reference[ids].view(torch.int16)) + dist.destroy_process_group() + + +def _mixed_storage_worker(rank, rendezvous): + torch.set_num_threads(1) + dist.init_process_group("gloo", init_method=rendezvous, rank=rank, world_size=4, timeout=timedelta(seconds=60)) + for ranks in ([0, 1], [2, 3]): + group = dist.new_group(ranks) + if rank in ranks: + tp = group + q = hbm.EngramQueryGroup(dist.group.WORLD, dist.group.WORLD, tp, rank // 2 * 2) + bf16 = hbm.NodeShardedEngram(41, 256, q, device="cpu", storage_format="bf16") + int8 = hbm.NodeShardedEngram(43, 256, q, device="cpu", storage_format="int8") + bf16_reference = torch.arange(41 * 256, dtype=torch.float32).reshape(41, 256).bfloat16() + int8_codes = (torch.arange(43 * 256).reshape(43, 256) % 255 - 127).to(torch.int8) + int8_scales = torch.linspace(0.25, 1.0, 43 * 8).reshape(43, 8) + int8_reference = (int8_codes.float().reshape(43, 8, 32) * int8_scales[..., None]).reshape(43, 256).bfloat16() + bf16.weight.data.copy_(bf16_reference[bf16.start : bf16.end]) + int8.weight.data.copy_(int8_codes[int8.start : int8.end]) + int8.weight_scale.copy_(int8_scales[int8.start : int8.end]) + count = (0, 5)[rank // 2] + first_ids = (torch.arange(count * 7).reshape(count, 7) * 3 + rank // 2) % 41 + second_ids = (torch.arange(count * 7).reshape(count, 7) * 5 + rank // 2) % 43 + first, second = bf16.route_many([bf16, int8], [first_ids, second_ids]) + assert torch.equal(first.view(torch.int16), bf16_reference[first_ids].view(torch.int16)) + assert torch.equal(second.view(torch.int16), int8_reference[second_ids].view(torch.int16)) + dist.destroy_process_group() + + +def _fp8_offload_worker(rank, rendezvous): + torch.set_num_threads(1) + dist.init_process_group("gloo", init_method=rendezvous, rank=rank, world_size=4, timeout=timedelta(seconds=60)) + for ranks in ([0, 1], [2, 3]): + group = dist.new_group(ranks) + if rank in ranks: + tp = group + q = hbm.EngramQueryGroup(dist.group.WORLD, dist.group.WORLD, tp, rank // 2 * 2) + for storage in ("fp8", "mxfp8"): + table = hbm.NodeShardedEngram(37, 32, q, device="cpu", storage_format=storage) + codes = (torch.arange(37 * 32).reshape(37, 32) % 240 - 120).to(torch.float8_e4m3fn) + scales = torch.tensor([[1.0]] * 37, dtype=torch.float8_e8m0fnu) + table.weight.data.copy_(codes[table.start : table.end]) + table.weight_scale.copy_(scales[table.start : table.end]) + ids = torch.tensor([[0, 36, 1], [7, 7, 19]]) if rank // 2 == 0 else torch.empty((0, 3), dtype=torch.long) + actual = table(ids) + reference = codes.float().reshape(37, 1, 32).mul(scales.float().reshape(37, 1, 1)).reshape(37, 32).bfloat16() + expected = reference[ids] if rank // 2 == 0 else torch.empty((0, 3, 32), dtype=torch.bfloat16) + assert actual.is_pinned() is False + assert torch.equal(actual.view(torch.int16), expected.view(torch.int16)), (rank, storage) + dist.destroy_process_group() + + +def test_cross_dp_variable_idle_and_exact_bf16(tmp_path): + mp.spawn(_worker, args=(f"file://{tmp_path / 'rendezvous'}",), nprocs=4, join=True) + + +def test_int8_compressed_wire_cross_dp(tmp_path): + mp.spawn(_compressed_wire_worker, args=(f"file://{tmp_path / 'compressed'}",), nprocs=4, join=True) + + +def test_route_many_mixed_storage_and_idle(tmp_path): + mp.spawn(_mixed_storage_worker, args=(f"file://{tmp_path / 'mixed'}",), nprocs=4, join=True) + + +def test_fp8_offload_cross_dp_and_idle(tmp_path): + mp.spawn(_fp8_offload_worker, args=(f"file://{tmp_path / 'fp8'}",), nprocs=4, join=True) + + +def test_shard_loader(tmp_path): + weights = torch.arange(29 * 8).reshape(29, 8).to(torch.bfloat16) + key = "layers.1.engram.embed.weight" + save_file({key: weights}, tmp_path / "weights.safetensors") + (tmp_path / "quant_model_weights.safetensors.index.json").write_text( + json.dumps({"weight_map": {key: "weights.safetensors"}}) + ) + pieces = [] + for rank in range(16): + q = type("QueryGroup", (), {"size": 16, "rank": rank})() + if rank * 2 >= 29: + with pytest.raises(ValueError): + hbm.NodeShardedEngram(29, 8, q, device="cpu") + continue + table = hbm.NodeShardedEngram(29, 8, q, device="cpu") + table.load_checkpoint(tmp_path, key, chunk_rows=3) + pieces.append(table.weight) + assert torch.equal(torch.cat(pieces), weights) + + +def test_cached_metadata_stays_on_cpu_with_device_context(): + q = type("QueryGroup", (), {"size": 4, "rank": 1, "is_source": False})() + with torch.device("meta"): + table = hbm.NodeShardedEngram(17, 32, q, device="cpu") + for ids in (torch.empty((0, 3), dtype=torch.int64), torch.tensor([[1, 2, 3]])): + flat, order, metadata = table._metadata(ids) + assert flat.device.type == order.device.type == metadata.device.type == "cpu" + assert torch.equal(metadata, torch.zeros(5, dtype=torch.int64)) + + +def test_loader_rejects_wrong_dtype(tmp_path): + key = "layers.1.engram.embed.weight" + save_file({key: torch.zeros(17, 8)}, tmp_path / "weights.safetensors") + (tmp_path / "quant_model_weights.safetensors.index.json").write_text( + json.dumps({"weight_map": {key: "weights.safetensors"}}) + ) + q = type("QueryGroup", (), {"size": 2, "rank": 0})() + table = hbm.NodeShardedEngram(17, 8, q, device="cpu") + with pytest.raises(ValueError, match="expected BF16"): + table.load_checkpoint(tmp_path, key) + + +def test_gate_preserves_masked_rows(): + torch.manual_seed(7) + hidden = torch.randn(3, 4, 32).bfloat16() + key = torch.randn(3, 4, 32).bfloat16() + value = torch.randn(3, 32).bfloat16() + out = gate(hidden, key, value, torch.randn(4, 32), torch.eye(32), torch.tensor([True, False, True]), 1e-5) + assert torch.equal(out[1], hidden[1]) + assert torch.isfinite(out.float()).all() + + +@pytest.mark.parametrize("barrier_token", [98, 99]) +def test_hash_causal_barrier(barrier_token): + h = hash_mod.PagedNgramHistory.__new__(hash_mod.PagedNgramHistory) + h.token_map = torch.arange(100) + h.pad_id = 2 + h.image_token_id = 99 + h.lookback = 2 + h.image_pad_token_id = 98 + h.primes = torch.tensor([[[101, 103]]]) + h.offsets = torch.tensor([[0, 101]]) + h.multipliers = torch.tensor([[3, 5]]) + h.pages = {} + values, mask = h.update( + torch.tensor([0, 5, 9, barrier_token, 13, 17]), + torch.arange(6), + torch.zeros(6, dtype=torch.long), + torch.tensor([[5, 1]]), + 4, + ) + assert values.shape == (6, 1, 2) and not mask[3] + # The first token on the next page must hash against padding, not the image. + assert values[4, 0, 0].item() == ((13 * 3) ^ (h.pad_id * 5)) % 101 diff --git a/tests/ut/ops/test_aurora_tiling_keys.py b/tests/ut/ops/test_aurora_tiling_keys.py new file mode 100644 index 000000000000..70d1a1b2bb81 --- /dev/null +++ b/tests/ut/ops/test_aurora_tiling_keys.py @@ -0,0 +1,315 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# mypy: ignore-errors +"""Check the compiled template matrix without requiring CANN or an NPU. + +The C++ preprocessor selects the real architecture branch and expands the +real headers. The stub only exposes macro arguments for enumeration; this is +not a CANN compilation or a numerical operator test. +""" + +import ast +import itertools +import shutil +import subprocess +import tempfile +import unittest +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[3] +TILING_HEADERS = ( + REPO_ROOT / "csrc/attention/sparse_flash_mla/op_kernel/sparse_flash_mla_template_tiling_key.h", + REPO_ROOT / "csrc/attention/quant_lightning_indexer_v2/op_kernel/quant_lightning_indexer_v2_template_tiling_key.h", + REPO_ROOT / "csrc/attention/compressor/op_kernel/arch32/compressor_template_tiling_key.h", + REPO_ROOT / "csrc/attention/compressor/op_kernel/arch35/compressor_template_tiling_key.h", +) +TEMPLATE_ARGUMENT_STUB = """ +#define ASCENDC_TPL_ARGS_DECL(op, ...) declaration = [__VA_ARGS__] +#define ASCENDC_TPL_BOOL_DECL(name, ...) [#name, 1, [__VA_ARGS__]] +#define ASCENDC_TPL_DTYPE_DECL(name, ...) [#name, "dtype", [__VA_ARGS__]] +#define ASCENDC_TPL_UINT_DECL(name, bits, kind, ...) [#name, bits, [__VA_ARGS__]] +#define ASCENDC_TPL_SEL(...) selection = [__VA_ARGS__] +#define ASCENDC_TPL_ARGS_SEL(...) [__VA_ARGS__] +#define ASCENDC_TPL_BOOL_SEL(name, ...) [__VA_ARGS__] +#define ASCENDC_TPL_DTYPE_SEL(name, ...) [__VA_ARGS__] +#define ASCENDC_TPL_UINT_SEL(name, kind, ...) [__VA_ARGS__] +#define ASCENDC_TPL_TILING_STRUCT_SEL(...) +""" + +COMPRESSOR_KERNEL_STUB = """ +#pragma once +#include +#define __global__ +#define __aicore__ +#define __gm__ +#define REGISTER_TILING_DEFAULT(...) +#define KERNEL_TASK_TYPE_DEFAULT(...) +#define GET_TILING_DATA_WITH_STRUCT(type, name, ...) type name; +namespace optiling { struct CompressorTilingData {}; } +namespace Compressor { +struct TPipe {}; +enum class X_LAYOUT : uint8_t { BSH = 0, TH = 1 }; +enum class X_DTYPE : uint8_t { BF16 = 0, FP16 = 1 }; +enum class COFF : uint8_t { DISABLE = 1, OVERLAP = 2 }; +enum class ROTARY_MODE : uint8_t { HALF = 1, INTERLEAVE = 2 }; +enum class CACHE_MODE : uint8_t { CONTINUOUS = 1, CYCLE = 2 }; +enum class ROPE_DTYPE : uint8_t { SAME_AS_X = 0, FP32 = 1 }; +enum class TEMPLATE_ID : uint8_t { + NORMAL = 0, EMPTY_X = 1, FULL_LOAD = 2 +}; +template struct COMPType {}; +template struct Kernel { + Kernel(...) { + static_assert(Tag == EXPECTED_KERNEL, "Unexpected computation template instantiated"); + } + void Init(...) {} + void Process() {} +}; +template using CompressorKernel = Kernel; +#if __CCE_AICORE__ == 220 +template using CompressorKernelPerf = Kernel; +#endif +template using CompressorKernelFullLoad = Kernel; +} +""" + + +class AuroraTilingKeysTest(unittest.TestCase): + @classmethod + def setUpClass(cls): + compiler = shutil.which("c++") + if compiler is None: + raise unittest.SkipTest("A C++ preprocessor is required to inspect template selections") + cls.matrices = {} + cls.declarations = {} + with tempfile.TemporaryDirectory() as directory: + stub = Path(directory) / "ascendc/host_api/tiling/template_argument.h" + stub.parent.mkdir(parents=True) + stub.write_text(TEMPLATE_ARGUMENT_STUB) + for header, arch in itertools.product(TILING_HEADERS, (None, 220, 310)): + compressor_arch = {"arch32": 220, "arch35": 310}.get(header.parent.name) + if compressor_arch is not None and arch != compressor_arch: + continue + args = [compiler, "-E", "-P", "-x", "c++", "-I", directory] + if arch is not None: + args.append(f"-D__CCE_AICORE__={arch}") + output = subprocess.check_output([*args, str(header)], text=True, timeout=30) + assignments = {node.targets[0].id: ast.literal_eval(node.value) for node in ast.parse(output).body} + op = "compressor" if compressor_arch is not None else header.parents[1].name + cls.declarations[op, arch] = assignments["declaration"] + cls.matrices[op, arch] = [ + key for selection in assignments["selection"] for key in itertools.product(*selection) + ] + + def test_compressor_retains_model_dispatch_and_key_encoding(self): + declaration = [ + ["X_LAYOUT", 1, [0, 1]], + ["X_DTYPE", 4, [0, 1]], + ["COFF", 2, [1, 2]], + ["ROTARY_MODE", 2, [1, 2]], + ["CACHE_MODE", 2, [1, 2]], + ["TEMPLATE_ID", 2, [0, 1, 2]], + ] + for arch in (220, 310): + # V4 ratio 4 overlaps (coff=2), ratio 128 does not (coff=1). + # TH selects PERF on A2/A3 and NORMAL on A5; EMPTY_X is required + # on both. FULL_LOAD requires BSH and is unreachable in the model. + template_ids = (0, 1) + expected = { + (1, 0, coff, 2, 1, template_id) for coff, template_id in itertools.product((1, 2), template_ids) + } + with self.subTest(arch=arch): + self.assertEqual(set(self.matrices["compressor", arch]), expected) + self.assertEqual(len(self.matrices["compressor", arch]), 4) + self.assertEqual( + self.declarations["compressor", arch], + declaration, + ) + + def test_compressor_dtype_registration_matches_selected_templates(self): + root = REPO_ROOT / "csrc/attention/compressor/op_host" + source = (root / "compressor_def.cpp").read_text() + for arch in ("arch32", "arch35"): + with self.subTest(arch=arch): + rows = [] + for line in source.splitlines(): + if 'Input("' in line or 'Output("' in line: + name = line.split('"')[1] + elif ".DataType" in line: + rows.append((name, line.split("{", 1)[1].split("}", 1)[0])) + self.assertEqual(len(rows), 14) + expected = { + "x": "ge::DT_BF16, ge::DT_FLOAT16", + "wkv": "ge::DT_BF16, ge::DT_FLOAT16", + "wgate": "ge::DT_BF16, ge::DT_FLOAT16", + "cmp_kv": "ge::DT_BF16, ge::DT_FLOAT16", + "state_cache": "ge::DT_FLOAT", + "ape": "ge::DT_FLOAT", + "norm_weight": "ge::DT_FLOAT", + "rope_sin": "ge::DT_FLOAT", + "rope_cos": "ge::DT_FLOAT", + "state_block_table": "ge::DT_INT32", + "cu_seqlens": "ge::DT_INT32", + "seqused": "ge::DT_INT32", + "start_pos": "ge::DT_INT32", + } + for name, dtypes in rows: + self.assertEqual(dtypes.strip(), expected[name]) + host = (root / arch / "compressor_tiling.h").read_text() + host_expected = expected.copy() + if arch == "arch35": + for name in ("x", "wkv", "wgate", "cmp_kv"): + host_expected[name] = "ge::DT_BF16" + for name, dtype in host_expected.items(): + self.assertRegex(host, rf"\{{{name.upper()}_NAME,\s*\{{{dtype}\}}\}}") + self.assertRegex( + host, + r"\{X_NAME,\s*\{COMPRESSOR_DIM_NUM_2" + + (r",\s*COMPRESSOR_DIM_NUM_3" if arch == "arch32" else "") + + r"\}\}", + ) + expected_modes = { + "ROTARY_MODE": "1, 2" if arch == "arch32" else "2", + "CACHE_MODE": "1", + } + for mode, values in expected_modes.items(): + definitions = [ + line for line in host.splitlines() if line.startswith(f"const std::vector {mode} ") + ] + self.assertGreaterEqual(len(definitions), 1) + # arch32 carries a DAY0 branch first; the final definition + # is the normal production branch selected without DAY0_SCOPE. + self.assertRegex(definitions[-1], rf"\{{\s*{values}\s*\}}") + + def test_compressor_empty_input_does_not_instantiate_computation(self): + compiler = shutil.which("c++") + source = (REPO_ROOT / "csrc/attention/compressor/op_kernel/compressor.cpp").read_text() + for arch in (220, 310): + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + arch_dir = root / ("arch32" if arch == 220 else "arch35") + arch_dir.mkdir() + (arch_dir / "stub.h").write_text(COMPRESSOR_KERNEL_STUB) + # Only the target architecture's headers exist: crossing the + # architecture branch fails compilation instead of going unnoticed. + for name in ( + "compressor_kernel.h", + "compressor_kernel_perf.h", + "compressor_kernel_full_load.h", + ): + (arch_dir / name).write_text('#include "stub.h"\n') + for key in self.matrices["compressor", arch]: + with self.subTest(arch=arch, key=key): + expected_kernel = -1 if key[5] == 1 else int(arch == 220) + invocation = ", ".join(map(str, key)) + arguments = ", ".join(["nullptr"] * 16) + (root / "entry.cpp").write_text( + source + f"\nvoid instantiate() {{ compressor<{invocation}>({arguments}); }}\n" + ) + result = subprocess.run( + [ + compiler, + "-std=c++17", + "-fsyntax-only", + f"-D__CCE_AICORE__={arch}", + f"-DEXPECTED_KERNEL={expected_kernel}", + str(root / "entry.cpp"), + ], + capture_output=True, + text=True, + timeout=30, + ) + self.assertEqual(result.returncode, 0, result.stderr) + + def test_sparse_mla_covers_model_dispatch_on_each_architecture(self): + # DoOpTiling derives flags from hardware, template mode and local head + # count. A5 CSA can enable address vectorization depending on the KV + # block shape and UB capacity; both paths must stay compiled. + for arch in (220, 310): + expected = set() + for ratio, heads, batch_consistency in itertools.product((0, 1, 2), (1, 2, 4, 8, 16, 32, 64, 128), (0, 1)): + mode = 0 if ratio == 0 else 2 + split_g = int(arch == 310 and heads > 64) + head_ratio_one = int(arch == 220 and mode == 2 and heads == 1) + vectorize_flags = (0, 1) if arch == 310 and mode == 2 else (0,) + for vectorize in vectorize_flags: + expected.add((0, 1, 2, mode, split_g, head_ratio_one, batch_consistency, vectorize)) + with self.subTest(arch=arch): + selected = self.matrices["sparse_flash_mla", arch] + self.assertEqual(set(selected), expected) + self.assertEqual(len(selected), 6 if arch == 220 else 12) + self.assertEqual(len(selected), len(set(selected))) + + def test_sparse_mla_architectures_exclude_each_others_specializations(self): + a2a3 = set(self.matrices["sparse_flash_mla", 220]) + a5 = set(self.matrices["sparse_flash_mla", 310]) + self.assertTrue(all(key[4] == 0 and key[7] == 0 for key in a2a3)) + self.assertTrue(all(key[5] == 0 for key in a5)) + self.assertEqual(len(a2a3 & a5), 4) + host = self.matrices["sparse_flash_mla", None] + self.assertEqual(set(host), a2a3 | a5) + self.assertEqual(len(host), 14) + + def test_sparse_mla_scope_and_a2a3_qli_quantization(self): + for arch in (None, 220, 310): + with self.subTest(arch=arch): + for key in self.matrices["sparse_flash_mla", arch]: + self.assertEqual(key[1:3], (1, 2)) # TND / PA_BBND + self.assertIn(key[3], (0, 2)) # SWA / CSA + if arch != 310: + # A2/A3: INT8 Q/K, INT32 output, paged attention, TND / PA_BBND. + self.assertEqual( + self.matrices["quant_lightning_indexer_v2", arch], + [ + (2, 2, 3, 1, 0, 2), + (2, 2, 3, 1, 1, 2), + (2, 2, 3, 0, 0, 0), + (2, 2, 3, 0, 1, 1), + ], + ) + + def test_qli_a5_retains_full_dtype_and_layout_matrix(self): + # A5's general QLI supports FP8/MXFP8, HiFloat8, MXFP4 and INT8. + # FP8 and MXFP8 share a dtype key and dispatch by runtime quant_mode. + expected = { + (dtype, dtype, 3, paged, q_layout, k_layout) + for dtype, (paged, q_layout, k_layout) in itertools.product( + (36, 34, 40, 2), ((1, 0, 2), (1, 1, 2), (0, 0, 0), (0, 1, 1)) + ) + } + selected = self.matrices["quant_lightning_indexer_v2", 310] + self.assertEqual(set(selected), expected) + self.assertEqual(len(selected), 16) + + def test_key_argument_order_widths_and_values_stay_unchanged(self): + # Keep the declaration, including unused values: changing its encoding + # can change the numeric keys shared by host tiling and kernel lookup. + expected = [ + ["FLASH_DECODE", 1, [0, 1]], + ["LAYOUT_T", 4, [0, 1]], + ["KV_LAYOUT_T", 4, [0, 1, 2]], + ["TEMPLATE_MODE", 4, [0, 1, 2, 3, 4]], + ["SPLIT_G", 1, [0, 1]], + ["HEAD_RATIO_ONE", 1, [0, 1]], + ["BATCH_CONSISTENCY", 1, [0, 1]], + ["IS_VEC_S2PHYADDR", 1, [0, 1]], + ] + for arch in (None, 220, 310): + with self.subTest(arch=arch): + self.assertEqual(self.declarations["sparse_flash_mla", arch], expected) + qk_types = [36, 34, 40, 2] if arch == 310 else [2] + self.assertEqual( + self.declarations["quant_lightning_indexer_v2", arch], + [ + ["DT_Q", "dtype", qk_types], + ["DT_K", "dtype", qk_types], + ["DT_OUT", "dtype", [3]], + ["PAGE_ATTENTION", 1, [1, 0]], + ["Q_LAYOUT_T", 4, [0, 1]], + ["K_LAYOUT_T", 4, [0, 1, 2]], + ], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/ut/ops/test_comm_utils.py b/tests/ut/ops/test_comm_utils.py index 892f58974e01..5fc6d514d5a5 100644 --- a/tests/ut/ops/test_comm_utils.py +++ b/tests/ut/ops/test_comm_utils.py @@ -59,6 +59,20 @@ def test_async_all_to_all(self, input_tensor, output_split_sizes, input_split_si assert handle is not None assert isinstance(handle, mocker.MagicMock) + def test_async_all_to_all_normalizes_numpy_splits(self, mocker: MockerFixture): + input_tensor = torch.randn(8, 16) + splits = torch.tensor([2, 2, 2, 2], dtype=torch.float32).numpy() + all_to_all = mocker.patch("torch.distributed.all_to_all_single", return_value=mocker.MagicMock()) + + async_all_to_all(input_tensor, splits, splits, mocker.MagicMock()) + + assert all_to_all.call_args.kwargs["input_split_sizes"] == [2, 2, 2, 2] + assert all_to_all.call_args.kwargs["output_split_sizes"] == [2, 2, 2, 2] + + def test_async_all_to_all_rejects_mismatched_input_splits(self): + with pytest.raises(RuntimeError, match="MoE all-to-all input split mismatch"): + async_all_to_all(torch.randn(8, 16), [2, 2, 2, 2], [1, 1, 1, 1], object()) + @pytest.mark.parametrize( "world_size, test_tensor, expected", [(1, torch.randn(8, 16), (8, 16)), (4, torch.randn(8, 16), (32, 16))] ) diff --git a/tests/ut/ops/test_fused_moe.py b/tests/ut/ops/test_fused_moe.py index 9f8714fe1c03..c5868f97996e 100644 --- a/tests/ut/ops/test_fused_moe.py +++ b/tests/ut/ops/test_fused_moe.py @@ -515,6 +515,7 @@ def test_routed_experts_select_experts_validates_router_logits(monkeypatch): monkeypatch.setattr(routed_experts_module, "get_forward_context", lambda: SimpleNamespace(input_ids=None)) monkeypatch.setattr(routed_experts_module, "get_current_vllm_config", lambda: None) monkeypatch.setattr(routed_experts_module, "get_moe_num_logical_experts", lambda *args, **kwargs: 3) + monkeypatch.setattr(routed_experts_module, "get_ascend_config", lambda: SimpleNamespace(enable_force_eplb=False)) result_weights, result_ids = routed_experts._select_experts( hidden_states=hidden_states, @@ -575,6 +576,7 @@ def test_routing_replay_captures_logical_ids_before_ascend_mapping(monkeypatch): "get_moe_num_logical_experts", lambda *args, **kwargs: 4, ) + monkeypatch.setattr(routed_experts_module, "get_ascend_config", lambda: SimpleNamespace(enable_force_eplb=False)) hidden_states = torch.randn(2, 4) router_logits = torch.tensor( [[0.1, 0.9, 0.2, 0.8], [0.7, 0.2, 0.6, 0.1]], @@ -607,6 +609,7 @@ def test_routing_replay_disabled_keeps_ascend_routing_unchanged(monkeypatch): "get_moe_num_logical_experts", lambda *args, **kwargs: 4, ) + monkeypatch.setattr(routed_experts_module, "get_ascend_config", lambda: SimpleNamespace(enable_force_eplb=False)) hidden_states = torch.randn(2, 4) router_logits = torch.tensor( [[0.1, 0.9, 0.2, 0.8], [0.7, 0.2, 0.6, 0.1]], @@ -624,9 +627,11 @@ def test_routing_replay_disabled_keeps_ascend_routing_unchanged(monkeypatch): torch.testing.assert_close(physical_ids, log2phy[expected_logical_ids]) -def test_hash_router_uses_explicit_input_ids(monkeypatch): - input_ids = torch.tensor([11, 22], dtype=torch.int32) - hidden_states = torch.randn(2, 4) +@pytest.mark.parametrize("hidden_dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("input_dtype", [torch.int32, torch.int64]) +def test_hash_router_preserves_fp32_weights_and_explicit_input_ids(monkeypatch, hidden_dtype, input_dtype): + input_ids = torch.tensor([0, 22], dtype=input_dtype) + hidden_states = torch.randn(2, 4, dtype=hidden_dtype) router_logits = torch.randn(2, 4) topk_weights = torch.randn(2, 2) topk_ids = torch.zeros(2, 2, dtype=torch.int32) @@ -655,22 +660,81 @@ def test_hash_router_uses_explicit_input_ids(monkeypatch): tid2eid=torch.ones(32, 4, dtype=torch.int32), ) - weights, ids = router._compute_routing( + routed_experts = _build_routing_replay_experts(router, None) + routed_experts.n_shared_experts = 0 + monkeypatch.setattr(routed_experts_module, "get_moe_num_logical_experts", lambda *args, **kwargs: 4) + monkeypatch.setattr(routed_experts_module, "get_ascend_config", lambda: SimpleNamespace(enable_force_eplb=False)) + weights, ids = routed_experts._select_experts( hidden_states, router_logits, - torch.int32, + enable_force_load_balance=False, input_ids=input_ids, ) assert weights is topk_weights + assert weights.dtype == torch.float32 assert ids is topk_ids torch.testing.assert_close(hash_op.call_args.kwargs["input_ids"], input_ids.to(torch.int64)) + if input_dtype == torch.int64: + assert hash_op.call_args.kwargs["input_ids"] is input_ids prepare_finalize.all_gather_input_id_with_dp_group.assert_called_once() with pytest.raises(ValueError, match="hash MoE routing requires input_ids"): router._compute_routing(hidden_states, router_logits, torch.int32) +def test_vision_router_fuses_bias_and_image_sentinel(monkeypatch): + input_ids = torch.tensor([11, 129259], dtype=torch.int32) + hidden_states = torch.randn(2, 4) + router_logits = torch.randn(2, 4, dtype=torch.float32) + text_bias = torch.randn(4, dtype=torch.float32) + bias_vl = torch.randn(4, dtype=torch.bfloat16) + topk_weights = torch.randn(2, 2) + topk_ids = torch.zeros(2, 2, dtype=torch.int32) + prepare_finalize = SimpleNamespace(all_gather_input_id_with_dp_group=MagicMock(side_effect=lambda value: value)) + monkeypatch.setattr( + fused_topk_router_module, + "_EXTRA_CTX", + SimpleNamespace( + moe_comm_type=MoECommType.ALLGATHER, + moe_comm_method=SimpleNamespace(prepare_finalize=prepare_finalize), + ), + ) + hash_op = MagicMock(return_value=(topk_weights, topk_ids, None)) + monkeypatch.setattr( + fused_topk_router_module.torch.ops._C_ascend, + "moe_gating_top_k_hash", + hash_op, + raising=False, + ) + router = AscendFusedTopKRouter( + top_k=2, + global_num_experts=4, + num_expert_group=1, + topk_group=1, + scoring_func="sqrtsoftplus", + e_score_correction_bias=text_bias, + bias_vl=bias_vl, + image_sentinel_lo=129257, + ) + + weights, ids = router._compute_routing( + hidden_states, + router_logits, + torch.int64, + input_ids=input_ids, + ) + + kwargs = hash_op.call_args.kwargs + assert weights is topk_weights + assert ids.dtype == torch.int64 + assert kwargs["bias"] is text_bias + assert kwargs["bias_vl"].dtype == router_logits.dtype + torch.testing.assert_close(kwargs["input_ids"], input_ids.to(torch.int64)) + assert kwargs["image_sentinel_lo"] == 129257 + assert kwargs["image_sentinel_count"] == 5 + + def test_hash_router_chunks_unaligned_input_ids_for_sequence_parallel(monkeypatch): input_ids = torch.tensor([11, 22, 33, 44], dtype=torch.int32) hidden_states = torch.randn(2, 4) @@ -793,13 +857,16 @@ def test_routed_experts_forward_impl_runs_current_flow(monkeypatch, return_with_ "_EXTRA_CTX", SimpleNamespace( in_profile_run=False, + moe_comm_type=MoECommType.MC2, moe_comm_method=moe_comm_method, eplb_heat_collection_status=False, ), ) + monkeypatch.setattr(routed_experts_module, "activate_moe_comm_method", lambda *args: None) monkeypatch.setattr(routed_experts_module, "get_forward_context", lambda: SimpleNamespace(all_moe_layers=None)) monkeypatch.setattr(routed_experts_module, "get_current_vllm_config", lambda: None) monkeypatch.setattr(routed_experts_module, "get_moe_num_logical_experts", lambda *args, **kwargs: 3) + monkeypatch.setattr(routed_experts_module, "get_ascend_config", lambda: SimpleNamespace(enable_force_eplb=False)) result = routed_experts.forward_impl( hidden_states=hidden_states, @@ -1735,6 +1802,54 @@ def test_forward_impl_returns_current_runner_contract(monkeypatch, has_shared_ex ascend_shared_experts.forward.assert_not_called() +@pytest.mark.parametrize("has_shared_experts", [False, True]) +@pytest.mark.parametrize("has_fp32_input", [False, True]) +def test_internal_router_reuses_fused_fp32_input(monkeypatch, has_shared_experts, has_fp32_input): + runner = AscendMoERunner.__new__(AscendMoERunner) + nn.Module.__init__(runner) + hidden_states = torch.randn(2, 4, dtype=torch.bfloat16) + router_input = hidden_states.float() if has_fp32_input else hidden_states + input_ids = torch.tensor([11, 22]) + weight = torch.randn(3, 4) + routed_out = torch.randn_like(hidden_states) + shared_out = torch.randn_like(hidden_states) + events = FusedMoEEvents(None, None, None, None, None) + runner.routed_experts = SimpleNamespace( + forward_impl=MagicMock(return_value=(routed_out, events) if has_shared_experts else routed_out) + ) + runner.ascend_shared_experts = ( + SimpleNamespace( + prepare_input_before_routed_experts=MagicMock(return_value=(hidden_states, None)), + forward=MagicMock(return_value=shared_out), + ) + if has_shared_experts + else None + ) + runner._sequence_parallel_context = MagicMock(return_value=nullcontext()) + runner.gate = SimpleNamespace(weight_fp32=weight) + runner.routed_input_transform = None + runner.routed_output_transform = None + monkeypatch.setattr(AscendMoERunner, "is_internal_router", property(lambda _: True)) + monkeypatch.setattr(fused_moe_module.torch.npu, "current_stream", MagicMock()) + linear = MagicMock(wraps=F.linear) + monkeypatch.setattr(fused_moe_module.F, "linear", linear) + + runner._forward_impl(hidden_states, router_input, shared_experts_input=None, input_ids=input_ids) + + linear.assert_called_once() + assert linear.call_args.args[0].dtype == torch.float32 + if has_fp32_input: + assert linear.call_args.args[0] is router_input + routed_kwargs = runner.routed_experts.forward_impl.call_args.kwargs + assert routed_kwargs["hidden_states"] is hidden_states + assert routed_kwargs["input_ids"] is input_ids + torch.testing.assert_close(routed_kwargs["router_logits"], hidden_states.float() @ weight.T) + if has_shared_experts: + assert runner.ascend_shared_experts is not None + runner.ascend_shared_experts.prepare_input_before_routed_experts.assert_not_called() + runner.ascend_shared_experts.forward.assert_called_once() + + def test_forward_impl_keeps_full_width_input_for_shared_experts(monkeypatch): runner = AscendMoERunner.__new__(AscendMoERunner) nn.Module.__init__(runner) diff --git a/tests/ut/ops/test_moe_comm_isolation.py b/tests/ut/ops/test_moe_comm_isolation.py new file mode 100644 index 000000000000..8984f8a5bf5a --- /dev/null +++ b/tests/ut/ops/test_moe_comm_isolation.py @@ -0,0 +1,70 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Target and DSpark experts must not overwrite each other's dispatch state.""" + +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from vllm_ascend.ascend_forward_context import MoECommType +from vllm_ascend.ops.fused_moe import moe_comm_method as comm + + +@pytest.mark.parametrize("ep_size", [1, 8]) +def test_expert_shapes_select_distinct_dispatchers(monkeypatch, ep_size): + monkeypatch.setattr(comm, "_MoECommMethods", {}) + monkeypatch.setattr(comm, "_MoECommMethodsByConfig", {}) + monkeypatch.setattr(comm, "_EXTRA_CTX", SimpleNamespace()) + constructors = [] + for name in ("AllGatherCommImpl", "AlltoAllCommImpl", "MC2CommImpl", "FusedMC2CommImpl"): + constructor = MagicMock(side_effect=lambda config: SimpleNamespace(owner=config)) + monkeypatch.setattr(comm, name, constructor) + constructors.append(constructor) + common = dict( + hidden_dim=5120, intermediate_size_per_partition=2048, ep_size=ep_size, tp_size=1, dp_size=1, pcp_size=1 + ) + target = SimpleNamespace(num_experts=384, num_local_experts=384 // ep_size, experts_per_token=6, **common) + draft = SimpleNamespace(num_experts=128, num_local_experts=128 // ep_size, experts_per_token=3, **common) + comm.setup_moe_comm_method(target) + comm.setup_moe_comm_method(draft) + kinds = [MoECommType.ALLGATHER] + if ep_size > 1: + kinds += [MoECommType.ALLTOALL, MoECommType.MC2, MoECommType.FUSED_MC2] + for kind in kinds: + target_method = comm.get_moe_comm_method(kind, target) + draft_method = comm.get_moe_comm_method(kind, draft) + assert target_method is not draft_method + assert target_method.owner is target and draft_method.owner is draft + for config, expected in ((target, target_method), (draft, draft_method), (target, target_method)): + assert comm.activate_moe_comm_method(kind, config) is expected + assert comm._EXTRA_CTX.moe_comm_method is expected + assert sum(c.call_count for c in constructors) == 2 * len(kinds) + + +def test_single_shape_keeps_forward_context_binding(monkeypatch): + original = SimpleNamespace(owner="forward-context") + cached = SimpleNamespace(owner="cached") + config = SimpleNamespace( + num_experts=128, + num_local_experts=128, + experts_per_token=8, + hidden_dim=4096, + intermediate_size_per_partition=1024, + ep_size=1, + tp_size=1, + dp_size=1, + pcp_size=1, + ) + key = comm._moe_config_key(config) + monkeypatch.setattr(comm, "_MoECommMethods", {MoECommType.ALLGATHER: cached}) + monkeypatch.setattr( + comm, + "_MoECommMethodsByConfig", + {(MoECommType.ALLGATHER, key): cached}, + ) + extra_ctx = SimpleNamespace(moe_comm_method=original) + monkeypatch.setattr(comm, "_EXTRA_CTX", extra_ctx) + + assert comm.activate_moe_comm_method(MoECommType.ALLGATHER, config) is cached + assert extra_ctx.moe_comm_method is original diff --git a/tests/ut/ops/test_rope_proxy.py b/tests/ut/ops/test_rope_proxy.py index 9920eea2aa42..39c4f833346a 100644 --- a/tests/ut/ops/test_rope_proxy.py +++ b/tests/ut/ops/test_rope_proxy.py @@ -2,9 +2,57 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import pytest import torch -from vllm_ascend.ops.rope_dsv4 import RopeDataProxy +from vllm_ascend.ops import rope_dsv4 +from vllm_ascend.ops.rope_dsv4 import ( + ComplexExpRotaryEmbedding, + RopeDataProxy, + get_full_cos_and_sin_dsa_for_layer, +) + + +def test_full_rope_lookup_resolves_exact_layer_config(monkeypatch): + first = (torch.randn(4, 1, 1, 8), torch.randn(4, 1, 1, 8)) + second = (torch.randn(4, 1, 1, 8), torch.randn(4, 1, 1, 8)) + monkeypatch.setattr( + rope_dsv4._ROPE_STATE, + "layer_info", + { + "model.layers.0.self_attn.attn": ("base", ["default"]), + "model.layers.2.self_attn.attn": ("compressed", ["default"]), + }, + ) + monkeypatch.setattr( + rope_dsv4._ROPE_STATE, + "full_rope_cache", + {"base": first, "compressed": second}, + ) + + actual = get_full_cos_and_sin_dsa_for_layer("model.layers.2.self_attn.attn") + + assert actual[0] is second[0] + assert actual[1] is second[1] + with pytest.raises(KeyError, match="not registered"): + get_full_cos_and_sin_dsa_for_layer("missing") + + +def test_zero_original_length_disables_yarn(): + dim = 8 + base = 10000 + actual = ComplexExpRotaryEmbedding.precompute_freqs_cis( + dim, + seqlen=65536, + original_seq_len=0, + base=base, + factor=16, + beta_fast=32, + beta_slow=1, + ) + expected = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) + torch.testing.assert_close(actual, expected) + # ────────────────────────────────────────────── # Equivalence: pad_to + slice vs pad-positions + gather + slice diff --git a/tests/ut/ops/test_token_dispatcher.py b/tests/ut/ops/test_token_dispatcher.py index 1a1e1311d35f..e938e2f4150d 100644 --- a/tests/ut/ops/test_token_dispatcher.py +++ b/tests/ut/ops/test_token_dispatcher.py @@ -310,7 +310,7 @@ def test_w4a8_group_dispatch_keeps_prefix_sum_group_list(self): self.assertEqual(kwargs["expert_token_nums_type"], EXPERT_TOKEN_NUMS_TYPE_CUMSUM) def test_get_combine_mc_kwargs_with_quant(self): - hidden_states = torch.randn(10, 128) + hidden_states = torch.randn(10, 128, dtype=torch.bfloat16) topk_ids = torch.randint(0, 8, (10, 1)) topk_weights = torch.randn(10, 1) expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7]) @@ -341,6 +341,8 @@ def test_get_combine_mc_kwargs_with_quant(self): self.dispatcher.moe_expert_num = len(expert_map) kwargs = self.dispatcher.get_combine_mc_kwargs(hidden_states, combine_metadata) self.assertIn("tp_send_counts", kwargs) + self.assertIs(kwargs["expert_scales"], topk_weights) + self.assertEqual(kwargs["expert_scales"].dtype, torch.float32) def test_get_combine_mc_kwargs_combine_quant_mode_forces_quant_mode(self): # When additional_config.combine_quant_mode is non-zero (here 4), the diff --git a/tests/ut/patch/platform/deepseek_v41/__init__.py b/tests/ut/patch/platform/deepseek_v41/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/examples/example.txt b/tests/ut/patch/platform/deepseek_v41/fixtures/examples/example.txt new file mode 100644 index 000000000000..9200d36c90f6 --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/examples/example.txt @@ -0,0 +1,7 @@ +中国的首都是哪里? + +列出100以内的所有素数。 + +DeepSeek是做什么的公司? + +请按“第一张、第二张”的顺序回答:第一张图examples/images/carrots.jpeg和第二张图examples/images/corn.jpeg中分别是什么食材?它们通常食用的部位分别是什么? diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/examples/example_harmony.json b/tests/ut/patch/platform/deepseek_v41/fixtures/examples/example_harmony.json new file mode 100644 index 000000000000..fae1c2608f8e --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/examples/example_harmony.json @@ -0,0 +1,96 @@ +[ + { + "messages": [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "请按“第一张、第二张”的顺序回答:第一张图" + }, + { + "type": "image_url", + "image_url": { + "url": "examples/images/carrots.jpeg" + } + }, + { + "type": "text", + "text": "和第二张图" + }, + { + "type": "image_url", + "image_url": { + "url": "examples/images/corn.jpeg" + } + }, + { + "type": "text", + "text": "中分别是什么食材?它们通常食用的部位分别是什么?" + } + ] + } + ] + }, + { + "messages": [ + { + "role": "system", + "content": "You are a helpful assistant." + }, + { + "role": "user", + "content": "中国的首都是哪里?" + } + ] + }, + { + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the weather for a specific location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string", "description": "The city name"}, + "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]} + }, + "required": ["location"] + } + } + } + ], + "messages": [ + { + "role": "system", + "content": "You are a helpful assistant." + }, + { + "role": "user", + "content": "What's the weather like in Beijing?" + } + ] + }, + { + "messages": [ + { + "role": "system", + "content": "You are a helpful assistant." + }, + { + "role": "user", + "content": "Hello" + }, + { + "role": "assistant", + "content": "Hi there! How can I help you?" + }, + { + "role": "system", + "content": "Mid-conversation instruction update: reply in Chinese only. (deepseek_v41 only)" + } + ] + } +] diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_1.json b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_1.json new file mode 100644 index 000000000000..c7435727e24c --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_1.json @@ -0,0 +1,82 @@ +{ + "thinking_mode": "thinking", + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the weather for a specific location", + "parameters": { + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "The city name" + }, + "unit": { + "type": "string", + "enum": ["celsius", "fahrenheit"], + "description": "Temperature unit" + } + }, + "required": ["location"] + } + } + }, + { + "type": "function", + "function": { + "name": "search", + "description": "Search the web for information", + "parameters": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "Search query" + }, + "num_results": { + "type": "integer", + "description": "Number of results to return" + } + }, + "required": ["query"] + } + } + } + ], + "messages": [ + { + "role": "system", + "content": "You are a helpful assistant." + }, + { + "role": "user", + "content": "What's the weather like in Beijing?" + }, + { + "role": "assistant", + "reasoning_content": "The user wants the weather in Beijing. I should call get_weather.", + "content": "", + "tool_calls": [ + { + "type": "function", + "function": { + "name": "get_weather", + "arguments": "{\"location\": \"Beijing\", \"unit\": \"celsius\"}" + } + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_0", + "content": "{\"temperature\": 22, \"condition\": \"sunny\", \"humidity\": 45}" + }, + { + "role": "assistant", + "reasoning_content": "Got the weather data. Let me format a nice response.", + "content": "The weather in Beijing is currently sunny with a temperature of 22\u00b0C and 45% humidity." + } + ] +} diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_2.json b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_2.json new file mode 100644 index 000000000000..132bb05a51e7 --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_2.json @@ -0,0 +1,24 @@ +[ + { + "role": "system", + "content": "You are a helpful assistant." + }, + { + "role": "user", + "content": "Hello" + }, + { + "role": "assistant", + "reasoning_content": "The user said hello, I should greet back.", + "content": "Hi there! How can I help you?" + }, + { + "role": "user", + "content": "What is the capital of France?" + }, + { + "role": "assistant", + "reasoning_content": "The user asks about the capital of France. It is Paris.", + "content": "The capital of France is Paris." + } +] diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_3.json b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_3.json new file mode 100644 index 000000000000..ff82f1c947e6 --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_3.json @@ -0,0 +1,89 @@ +[ + { + "role": "system", + "content": "该助手为DeepSeek,由深度求索公司创造。" + }, + { + "role": "latest_reminder", + "content": "2026-02-21,星期六,广州,App,中文" + }, + { + "role": "developer", + "content": "小柴胡冲剂和布洛芬能一起吃吗?\n\nCITATION FORMAT: 【{cursor_id}†L{start_line_id}(-L{end_line_id})?】", + "tools": [ + { + "type": "function", + "function": { + "name": "search", + "description": "Web search. Split multiple queries with '||'.", + "parameters": { + "type": "object", + "properties": { + "queries": { + "type": "string", + "description": "query1||query2" + } + }, + "required": ["queries"], + "additionalProperties": false + } + } + }, + { + "type": "function", + "function": { + "name": "open", + "description": "Batch open IDs (format 【{id}†...】) or URLs.", + "parameters": { + "type": "object", + "properties": { + "open_list": { + "type": "array", + "items": { + "type": "object", + "properties": { + "id": { + "description": "ID or URL", + "anyOf": [{"type": "integer"}, {"type": "string"}], + "default": -1 + }, + "loc": {"type": "integer", "description": "Start line", "default": -1}, + "num_lines": {"type": "integer", "description": "", "default": -1} + }, + "additionalProperties": false + }, + "description": "" + } + }, + "required": ["open_list"], + "additionalProperties": false + } + } + } + ] + }, + { + "role": "assistant", + "content": "", + "reasoning_content": "用户想知道小柴胡冲剂和布洛芬能否一起服用。", + "tool_calls": [ + { + "type": "function", + "function": { + "name": "search", + "arguments": "{\"queries\": \"小柴胡冲剂 布洛芬 相互作用 一起吃\"}" + } + } + ] + }, + { + "role": "tool", + "content": "[0]" + }, + { + "role": "assistant", + "content": "请及时就医。", + "reasoning_content": "现在开始组织回答。", + "tool_calls": [] + } +] diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_4.json b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_4.json new file mode 100644 index 000000000000..86cd1f0a6807 --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_4.json @@ -0,0 +1,28 @@ +[ + { + "role": "system", + "content": "该助手为DeepSeek-V3,由深度求索公司创造。\n今天是2025年10月17日,星期五。" + }, + { + "role": "latest_reminder", + "content": "2024-11-15,上海市,App,中文" + }, + { + "role": "user", + "content": "热海大滚锅是世界著名温泉吗" + }, + { + "role": "assistant", + "content": "热海大滚锅在中国乃至全球的地热奇观中占有重要地位,但“世界著名”的称号更侧重于它作为独特的地质现象和旅游景点。", + "mask": 1 + }, + { + "role": "user", + "content": "世界著名温泉有哪些", + "task": "action" + }, + { + "role": "assistant", + "content": "Search" + } +] diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_5.json b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_5.json new file mode 100644 index 000000000000..2a2fa6a89e81 --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_input_5.json @@ -0,0 +1,39 @@ +{ + "thinking_mode": "thinking", + "reasoning_effort": "max", + "messages": [ + { + "role": "system", + "content": "You are a helpful vision assistant." + }, + { + "role": "user", + "content": [ + { + "type": "text", + "text": "请按“第一张、第二张”的顺序回答:第一张图" + }, + { + "type": "image_url", + "image_url": { + "url": "examples/images/carrots.jpeg" + } + }, + { + "type": "text", + "text": "和第二张图" + }, + { + "type": "image_url", + "image_url": { + "url": "examples/images/corn.jpeg" + } + }, + { + "type": "text", + "text": "中分别是什么食材?它们通常食用的部位分别是什么?" + } + ] + } + ] +} diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_1.txt b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_1.txt new file mode 100644 index 000000000000..2960285848ae --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_1.txt @@ -0,0 +1,38 @@ +<|begin▁of▁sentence|><|System|>Reasoning Effort: 50 (range 1-100, the higher the value, the more thorough the reasoning) + +You are a helpful assistant. + +## Tools + +You have access to a set of tools to help answer the user's question. You can invoke tools by writing a "<|DSML| calls>" block like the following: + +<|DSML| calls> +<|DSML| invoke name="$TOOL_NAME"> +<|DSML| parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE +... + +<|DSML| invoke name="$TOOL_NAME2"> +... + + + +String parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`. + +If thinking_mode is enabled (triggered by ), you MUST output your complete reasoning inside ... BEFORE any tool calls or final response. + +Otherwise, output directly after with tool calls or final response. + +### Available Tool Schemas + +{"name": "get_weather", "description": "Get the weather for a specific location", "parameters": {"type": "object", "properties": {"location": {"type": "string", "description": "The city name"}, "unit": {"type": "string", "enum": ["celsius", "fahrenheit"], "description": "Temperature unit"}}, "required": ["location"]}} +{"name": "search", "description": "Search the web for information", "parameters": {"type": "object", "properties": {"query": {"type": "string", "description": "Search query"}, "num_results": {"type": "integer", "description": "Number of results to return"}}, "required": ["query"]}} + +You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls. +<|User|>What's the weather like in Beijing?<|Assistant|>The user wants the weather in Beijing. I should call get_weather. + +<|DSML| calls> +<|DSML| invoke name="get_weather"> +<|DSML| parameter name="location" string="true">Beijing +<|DSML| parameter name="unit" string="true">celsius + +<|end▁of▁sentence|><|User|>{"temperature": 22, "condition": "sunny", "humidity": 45}<|Assistant|>Got the weather data. Let me format a nice response.The weather in Beijing is currently sunny with a temperature of 22°C and 45% humidity.<|end▁of▁sentence|> \ No newline at end of file diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_2.txt b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_2.txt new file mode 100644 index 000000000000..78eb6be5de6f --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_2.txt @@ -0,0 +1 @@ +<|begin▁of▁sentence|><|System|>You are a helpful assistant.<|User|>Hello<|Assistant|>Hi there! How can I help you?<|end▁of▁sentence|><|User|>What is the capital of France?<|Assistant|>The capital of France is Paris.<|end▁of▁sentence|> \ No newline at end of file diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_3.txt b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_3.txt new file mode 100644 index 000000000000..c3a9508790ff --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_3.txt @@ -0,0 +1,37 @@ +<|begin▁of▁sentence|><|System|>该助手为DeepSeek,由深度求索公司创造。<|latest_reminder|>2026-02-21,星期六,广州,App,中文<|User|>小柴胡冲剂和布洛芬能一起吃吗? + +CITATION FORMAT: 【{cursor_id}†L{start_line_id}(-L{end_line_id})?】 + +## Tools + +You have access to a set of tools to help answer the user's question. You can invoke tools by writing a "<|DSML| calls>" block like the following: + +<|DSML| calls> +<|DSML| invoke name="$TOOL_NAME"> +<|DSML| parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE +... + +<|DSML| invoke name="$TOOL_NAME2"> +... + + + +String parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`. + +If thinking_mode is enabled (triggered by ), you MUST output your complete reasoning inside ... BEFORE any tool calls or final response. + +Otherwise, output directly after with tool calls or final response. + +### Available Tool Schemas + +{"name": "search", "description": "Web search. Split multiple queries with '||'.", "parameters": {"type": "object", "properties": {"queries": {"type": "string", "description": "query1||query2"}}, "required": ["queries"], "additionalProperties": false}} +{"name": "open", "description": "Batch open IDs (format 【{id}†...】) or URLs.", "parameters": {"type": "object", "properties": {"open_list": {"type": "array", "items": {"type": "object", "properties": {"id": {"description": "ID or URL", "anyOf": [{"type": "integer"}, {"type": "string"}], "default": -1}, "loc": {"type": "integer", "description": "Start line", "default": -1}, "num_lines": {"type": "integer", "description": "", "default": -1}}, "additionalProperties": false}, "description": ""}}, "required": ["open_list"], "additionalProperties": false}} + +You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls. +<|Assistant|> + +<|DSML| calls> +<|DSML| invoke name="search"> +<|DSML| parameter name="queries" string="true">小柴胡冲剂 布洛芬 相互作用 一起吃 + +<|end▁of▁sentence|><|User|>[0]<|Assistant|>请及时就医。<|end▁of▁sentence|> \ No newline at end of file diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_4.txt b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_4.txt new file mode 100644 index 000000000000..efad296eb4ff --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_4.txt @@ -0,0 +1,2 @@ +<|begin▁of▁sentence|><|System|>该助手为DeepSeek-V3,由深度求索公司创造。 +今天是2025年10月17日,星期五。<|latest_reminder|>2024-11-15,上海市,App,中文<|User|>热海大滚锅是世界著名温泉吗<|Assistant|>热海大滚锅在中国乃至全球的地热奇观中占有重要地位,但“世界著名”的称号更侧重于它作为独特的地质现象和旅游景点。<|end▁of▁sentence|><|User|>世界著名温泉有哪些<|Assistant|><|action|>Search<|end▁of▁sentence|> \ No newline at end of file diff --git a/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_5.txt b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_5.txt new file mode 100644 index 000000000000..a130bcc2c568 --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/fixtures/test_output_5.txt @@ -0,0 +1,11 @@ +<|begin▁of▁sentence|><|System|>Reasoning Effort: 100 (range 1-100, the higher the value, the more thorough the reasoning) + +You are a helpful vision assistant.<|User|>请按“第一张、第二张”的顺序回答:第一张图 + +<|deepseek_image|> + +和第二张图 + +<|deepseek_image|> + +中分别是什么食材?它们通常食用的部位分别是什么?<|Assistant|> \ No newline at end of file diff --git a/tests/ut/patch/platform/deepseek_v41/test_encoding.py b/tests/ut/patch/platform/deepseek_v41/test_encoding.py new file mode 100644 index 000000000000..ef0dc77f11d7 --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/test_encoding.py @@ -0,0 +1,429 @@ +# SPDX-License-Identifier: MIT +# Copyright (c) 2023 DeepSeek +# Reference tests; see vllm_ascend/patch/platform/patch_deepseek_v41_frontend/LICENSE. +""" +Tests for encoding.py (DeepSeek-V4.1 encoding). + +Adapted from dsv41-master/deepseek_harmony/tests/test_deepseek_v41.py for the +self-contained dict-based API in this repo. +""" + +import json +from pathlib import Path +from typing import Any + +import pytest + +from vllm_ascend.patch.platform.patch_deepseek_v41_frontend import encoding as enc +from vllm_ascend.patch.platform.patch_deepseek_v41_frontend.encoding import ( + IMAGE_PLACEHOLDER, + SYSTEM_SP_TOKEN, + encode_messages, + merge_tool_messages, + parse_message_from_completion_text, + render_message, +) + +REASONING_EFFORT_TEMPLATE = ( + SYSTEM_SP_TOKEN + "Reasoning Effort: {budget} " + "(range 1-100, the higher the value, the more thorough the reasoning)\n\n" +) + +V41_TOOL_CALL_OUTPUT = ( + " reason summary\n\n" + "<|DSML| calls>\n" + '<|DSML| invoke name="lookup">\n' + '<|DSML| parameter name="query" string="true">value' + "\n" + '<|DSML| parameter name="limit" string="false">2' + "\n" + "\n" + "<|end▁of▁sentence|>" +) + + +def make_tool() -> dict: + return { + "type": "function", + "function": { + "name": "lookup", + "description": "Look up a value", + "parameters": { + "type": "object", + "properties": { + "query": {"type": "string"}, + "limit": {"type": "integer"}, + }, + }, + }, + } + + +def make_tool_call_messages() -> list: + return [ + {"role": "user", "content": "question"}, + { + "role": "assistant", + "reasoning_content": " reason ", + "content": "summary", + "tool_calls": [ + { + "type": "function", + "function": { + "name": "lookup", + "arguments": '{"query":"value","limit":2}', + }, + } + ], + }, + ] + + +# ============================================================ +# Vision +# ============================================================ + + +def test_v41_renders_images() -> None: + prompt, media = encode_messages( + [ + { + "role": "user", + "content": [ + {"type": "text", "text": "inspect"}, + {"type": "image_url", "image_url": {"url": "/unused/image.png"}}, + ], + } + ], + thinking_mode="chat", + return_multi_modal_data=True, + ) + + assert prompt == (f"<|begin▁of▁sentence|><|User|>inspect\n\n{IMAGE_PLACEHOLDER}<|Assistant|>") + assert media == {"images": [{"type": "image", "url": "/unused/image.png"}]} + + +def test_v41_rejects_image_placeholder_in_text() -> None: + with pytest.raises(ValueError): + encode_messages( + [{"role": "user", "content": f"hi {IMAGE_PLACEHOLDER}"}], + thinking_mode="chat", + ) + + +# ============================================================ +# Reasoning Effort +# ============================================================ + + +@pytest.mark.parametrize( + ("effort", "budget"), + [ + (None, 50), + ("low", 25), + ("high", 50), + ("xhigh", 75), + ("max", 100), + (1, 1), + (42, 42), + (100, 100), + ], +) +def test_v41_maps_reasoning_effort_to_1_100_budget( + effort: Any, + budget: int, +) -> None: + prompt = encode_messages( + [{"role": "user", "content": "question"}], + thinking_mode="thinking", + reasoning_effort=effort, + ) + + assert prompt == ( + "<|begin▁of▁sentence|>" + f"{REASONING_EFFORT_TEMPLATE.format(budget=budget)}" + "<|User|>question<|Assistant|>" + ) + + +def test_v41_only_adds_reasoning_effort_to_first_thinking_message() -> None: + messages = [ + {"role": "system", "content": "system"}, + {"role": "user", "content": "question"}, + ] + + later_message = render_message(1, messages, thinking_mode="thinking", reasoning_effort=100) + chat_message = render_message(0, messages, thinking_mode="chat", reasoning_effort=100) + + assert "Reasoning Effort:" not in later_message + assert "Reasoning Effort:" not in chat_message + + +def test_v41_chat_mode_has_no_reasoning_effort_or_system_token() -> None: + prompt = encode_messages( + [{"role": "user", "content": "hello"}], + thinking_mode="chat", + reasoning_effort="max", + ) + assert prompt == "<|begin▁of▁sentence|><|User|>hello<|Assistant|>" + + +@pytest.mark.parametrize("effort", [-1, 0, 101, "medium"]) +def test_v41_rejects_out_of_range_or_unknown_reasoning_effort( + effort: Any, +) -> None: + with pytest.raises(AssertionError, match=r"int within \[1,100\]"): + encode_messages( + [{"role": "user", "content": "question"}], + thinking_mode="thinking", + reasoning_effort=effort, + ) + + +@pytest.mark.parametrize("effort", [True, False, 1.5]) +def test_v41_rejects_non_string_non_integer_effort_types(effort: Any) -> None: + # bool is not `type(...) is int`; float is invalid too + with pytest.raises(AssertionError): + encode_messages( + [{"role": "user", "content": "question"}], + thinking_mode="thinking", + reasoning_effort=effort, + ) + + +# ============================================================ +# System token +# ============================================================ + + +def test_v41_leading_system_message_uses_system_token() -> None: + prompt = encode_messages( + [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "hello"}, + ], + thinking_mode="chat", + ) + assert prompt == ( + "<|begin▁of▁sentence|><|System|>You are a helpful assistant.<|User|>hello<|Assistant|>" + ) + + +def test_v41_mid_conversation_system_message() -> None: + prompt = encode_messages( + [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": "q1"}, + {"role": "assistant", "content": "a1", "reasoning_content": "r1"}, + {"role": "system", "content": "mid sys"}, + ], + thinking_mode="thinking", + reasoning_effort=88, + ) + # Mid-conversation system gets its own <|System|> token and triggers + # the assistant generation header afterwards. + assert prompt == ( + "<|begin▁of▁sentence|>" + f"{REASONING_EFFORT_TEMPLATE.format(budget=88)}" + "sys<|User|>q1<|Assistant|>a1<|end▁of▁sentence|>" + "<|System|>mid sys<|Assistant|>" + ) + + +# ============================================================ +# DSML tool tags +# ============================================================ + + +def test_v41_tool_instructions_use_spaced_dsml_tags_in_chat_mode() -> None: + prompt = encode_messages( + [ + {"role": "system", "content": "system", "tools": [make_tool()]}, + {"role": "user", "content": "question"}, + ], + thinking_mode="chat", + ) + + assert ( + "<|DSML| calls>\n" + '<|DSML| invoke name="$TOOL_NAME">\n' + '<|DSML| parameter name="$PARAMETER_NAME" ' + 'string="true|false">$PARAMETER_VALUE\n' + "...\n" + "" + ) in prompt + assert "<|DSML|tool_calls>" not in prompt + assert "<|DSML|invoke" not in prompt + assert "<|DSML|parameter" not in prompt + + +def test_v41_renders_spaced_dsml_with_v4_assistant_semantics() -> None: + messages = make_tool_call_messages() + + prompt = render_message(1, messages, thinking_mode="thinking") + + assert prompt == V41_TOOL_CALL_OUTPUT + + +def test_v41_parses_spaced_dsml_roundtrip() -> None: + messages = make_tool_call_messages() + + parsed = parse_message_from_completion_text(V41_TOOL_CALL_OUTPUT, thinking_mode="thinking") + + assert parsed["role"] == "assistant" + assert parsed["reasoning_content"] == " reason " + assert parsed["content"] == "summary" + assert parsed["tool_calls"] + assert parsed["tool_calls"][0]["function"]["name"] == "lookup" + assert json.loads(parsed["tool_calls"][0]["function"]["arguments"]) == { + "query": "value", + "limit": 2, + } + + # Re-encoding the parsed message reproduces the original completion text + assert ( + encode_messages( + [parsed], + thinking_mode="thinking", + context=messages[:1], + ) + == V41_TOOL_CALL_OUTPUT + ) + + +def test_v41_parse_rejects_unspaced_v4_dsml() -> None: + v4_output = ( + V41_TOOL_CALL_OUTPUT.replace("|DSML| calls", "|DSML|tool_calls") + .replace("|DSML| invoke", "|DSML|invoke") + .replace("|DSML| parameter", "|DSML|parameter") + ) + with pytest.raises(AssertionError): + parse_message_from_completion_text(v4_output, thinking_mode="thinking") + + +# ============================================================ +# Multi-turn flow +# ============================================================ + + +def test_v41_drop_thinking_without_tools() -> None: + prompt = encode_messages( + [ + {"role": "user", "content": "q1"}, + {"role": "assistant", "content": "a1", "reasoning_content": "r1"}, + {"role": "user", "content": "q2"}, + ], + thinking_mode="thinking", + drop_thinking=True, + ) + # Earlier turn reasoning dropped, form; new turn opens + assert "<|User|>q1<|Assistant|>a1<|end▁of▁sentence|>" in prompt + assert "r1" not in prompt + assert prompt.endswith("<|User|>q2<|Assistant|>") + + +# ============================================================ +# Preprocessing +# ============================================================ + + +def test_merge_tool_messages_creates_tool_result_blocks() -> None: + merged = merge_tool_messages( + [ + {"role": "assistant", "content": "", "tool_calls": []}, + {"role": "tool", "tool_call_id": "a", "content": "r1"}, + {"role": "tool", "tool_call_id": "b", "content": "r2"}, + ] + ) + assert len(merged) == 2 + assert merged[1]["role"] == "user" + assert [b["type"] for b in merged[1]["content_blocks"]] == ["tool_result", "tool_result"] + + +def test_v41_task_sp_token() -> None: + prompt = encode_messages( + [{"role": "user", "content": "classify me", "task": "query"}], + thinking_mode="chat", + ) + assert prompt.endswith("classify me<|query|>") + assert "<|Assistant|>" not in prompt + + +# ============================================================ +# Golden fixtures from encoding/tests +# ============================================================ + +ENCODING_DIR = Path(__file__).resolve().parent +ENCODING_FIXTURES_DIR = ENCODING_DIR / "fixtures" +INFERENCE_EXAMPLES_DIR = ENCODING_FIXTURES_DIR / "examples" + +FIXTURE_CASE_IDS = sorted(int(p.stem.split("_")[-1]) for p in ENCODING_FIXTURES_DIR.glob("test_input_*.json")) + + +@pytest.mark.parametrize("case_id", FIXTURE_CASE_IDS) +def test_examples_encoding_golden_outputs(case_id: int) -> None: + """Each tests/encoding input must encode to its checked-in golden output.""" + input_file = ENCODING_FIXTURES_DIR / f"test_input_{case_id}.json" + output_file = ENCODING_FIXTURES_DIR / f"test_output_{case_id}.txt" + assert output_file.exists(), f"missing golden output: {output_file.name} (run tests/encoding/regen_outputs.py)" + + case = enc.load_cases(str(input_file))[0] + prompt, _ = enc.encode_case(case, thinking_mode="chat") + + assert prompt == output_file.read_text(), ( + f"{output_file.name} is stale; regenerate with tests/encoding/regen_outputs.py" + ) + + +def test_examples_v41_output_uses_v41_format_markers() -> None: + """Sanity-check the V4.1 goldens actually exercise V4.1-specific format.""" + # case 1: tool calls with spaced DSML tags + out1 = (ENCODING_FIXTURES_DIR / "test_output_1.txt").read_text() + assert "<|DSML| calls>" in out1 and '<|DSML| invoke name="get_weather">' in out1 + assert "<|DSML|tool_calls>" not in out1 + + # case 5: numeric reasoning effort behind the system token + out5 = (ENCODING_FIXTURES_DIR / "test_output_5.txt").read_text() + assert out5.startswith("<|begin▁of▁sentence|>" + REASONING_EFFORT_TEMPLATE.format(budget=100)) + assert out5.count(IMAGE_PLACEHOLDER) == 2 + + +def test_examples_vl_txt_and_json_encode_identically() -> None: + """The TXT (last block of example.txt) and JSON vision examples must encode identically.""" + txt = (INFERENCE_EXAMPLES_DIR / "example.txt").read_text().rstrip("\n").split("\n\n")[-1] + messages = [{"role": "user", "content": enc.parse_tagged_text(txt)}] + p1, m1 = encode_messages(messages, thinking_mode="chat", return_multi_modal_data=True) + + case = enc.load_cases(str(INFERENCE_EXAMPLES_DIR / "example_harmony.json"))[0] + p2, m2 = enc.encode_case(case, thinking_mode="chat") + + assert p1 == p2 + assert m1["images"] == m2 + assert len(m2) == 2 + + +def test_examples_harmony_cases_encode() -> None: + """All example_harmony.json cases encode without error.""" + cases = enc.load_cases(str(INFERENCE_EXAMPLES_DIR / "example_harmony.json")) + assert len(cases) == 4 + + # case 1 (vision) is covered by test_examples_vl_txt_and_json_encode_identically + + # cases are pure OpenAI format: mode/effort are passed at call time + prompt = encode_messages(cases[1]["messages"], thinking_mode="thinking", reasoning_effort=75) + assert REASONING_EFFORT_TEMPLATE.format(budget=75) in prompt + + # case 3: tools with spaced DSML tags + prompt, _ = enc.encode_case(cases[2], thinking_mode="chat") + assert "<|DSML| calls>" in prompt + + # case 4: mid-conversation system message triggers assistant header + prompt, _ = enc.encode_case(cases[3], thinking_mode="chat") + assert "<|System|>Mid-conversation instruction update" in prompt + assert prompt.endswith("<|Assistant|>") + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v"])) diff --git a/tests/ut/patch/platform/deepseek_v41/test_frontend.py b/tests/ut/patch/platform/deepseek_v41/test_frontend.py new file mode 100644 index 000000000000..e0243209c24b --- /dev/null +++ b/tests/ut/patch/platform/deepseek_v41/test_frontend.py @@ -0,0 +1,410 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +import base64 +import copy +import io +import json +import subprocess +import sys +from pathlib import Path +from types import SimpleNamespace +from typing import Any + +import pytest +from PIL import Image +from tokenizers import Tokenizer # type: ignore[import-untyped] +from tokenizers.models import WordLevel # type: ignore[import-untyped] +from transformers import PreTrainedTokenizerFast +from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest +from vllm.parser.parser_manager import ParserManager +from vllm.renderers.params import ChatParams +from vllm.utils.async_utils import make_async +from xgrammar import Grammar, GrammarCompiler, GrammarMatcher, TokenizerInfo + +from vllm_ascend.patch.platform.patch_deepseek_v41_frontend import register_frontend +from vllm_ascend.patch.platform.patch_deepseek_v41_frontend.encoding import ( + encode_messages, + load_cases, + parse_message_from_completion_text, +) +from vllm_ascend.patch.platform.patch_deepseek_v41_frontend.parser import DeepseekV41Parser, DeepseekV41ToolParser +from vllm_ascend.patch.platform.patch_deepseek_v41_frontend.renderer import DeepseekV41Renderer +from vllm_ascend.patch.platform.patch_deepseek_v41_frontend.tokenizer import get_deepseek_v41_tokenizer + +FIXTURES = Path(__file__).parent / "fixtures" +TOOL: dict[str, Any] = { + "type": "function", + "function": { + "name": "lookup", + "parameters": { + "type": "object", + "properties": {"query": {"type": "string"}, "limit": {"type": "integer"}}, + "required": ["query", "limit"], + }, + }, +} + + +@pytest.fixture +def tokenizer(): + tokens = ["[UNK]", "", "", "<|end▁of▁sentence|>", "|DSML|"] + backend = Tokenizer(WordLevel({token: i for i, token in enumerate(tokens)}, unk_token="[UNK]")) + return get_deepseek_v41_tokenizer( + PreTrainedTokenizerFast(tokenizer_object=backend, additional_special_tokens=tokens[1:]) + ) + + +def request(**kwargs): + return ChatCompletionRequest(model="v41", messages=[{"role": "user", "content": "hi"}], **kwargs) + + +def call_text(arguments, name="lookup"): + params = "\n".join( + f'<|DSML| parameter name="{key}" string="{str(isinstance(value, str)).lower()}">' + f"{value if isinstance(value, str) else json.dumps(value, ensure_ascii=False)}" + for key, value in arguments.items() + ) + return f'<|DSML| invoke name="{name}">\n{params}\n' + + +def completion(arguments, thinking=True): + return ( + (" reason " if thinking else "") + + "summary\n\n<|DSML| calls>\n" + + call_text(arguments) + + "\n" + ) + + +@pytest.mark.parametrize("case_id", range(1, 6)) +def test_tokenizer_matches_checkpoint_goldens(tokenizer, case_id): + case = load_cases(str(FIXTURES / f"test_input_{case_id}.json"))[0] + before = copy.deepcopy(case) + actual = tokenizer.apply_chat_template( + case["messages"], + tokenize=False, + thinking=case.get("thinking_mode", "chat") == "thinking", + reasoning_effort=case.get("reasoning_effort"), + context=case.get("context"), + drop_thinking=case.get("drop_thinking", True), + ) + assert actual == (FIXTURES / f"test_output_{case_id}.txt").read_text() + assert case == before + + +@pytest.mark.parametrize( + "effort,budget", [(None, 50), ("low", 25), ("high", 50), ("xhigh", 75), ("max", 100), (1, 1), (42, 42), (100, 100)] +) +def test_tokenizer_preserves_numeric_effort(tokenizer, effort, budget): + prompt = tokenizer.apply_chat_template(request().messages, reasoning_effort=effort, tokenize=False) + assert f"<|System|>Reasoning Effort: {budget} (range 1-100," in prompt + assert prompt.endswith("") + + +@pytest.mark.parametrize("effort", [0, 101, True, 1.5, "medium", "minimal", "invalid"]) +def test_tokenizer_rejects_invalid_effort(tokenizer, effort): + with pytest.raises((AssertionError, ValueError)): + tokenizer.apply_chat_template(request().messages, reasoning_effort=effort) + + +@pytest.mark.parametrize("kwargs", [{"thinking": False}, {"enable_thinking": False}, {"reasoning_effort": "none"}]) +def test_chat_mode_matches_parser_initial_state(tokenizer, kwargs): + prompt = tokenizer.apply_chat_template(request().messages, tokenize=False, **kwargs) + assert prompt.endswith("") and "Reasoning Effort:" not in prompt + parser = DeepseekV41Parser(tokenizer, chat_template_kwargs=kwargs) + assert parser.extract_reasoning("answer", request()) == (None, "answer") + + +def test_tools_follow_existing_system_and_preserve_history(tokenizer): + messages: list[dict[str, Any]] = [ + {"role": "system", "content": "system"}, + {"role": "user", "content": "q"}, + {"role": "assistant", "content": "a", "reasoning": "old thought"}, + {"role": "user", "content": "next"}, + ] + before = copy.deepcopy(messages) + expected = copy.deepcopy(messages) + expected[0]["tools"] = [TOOL] + expected[2]["reasoning_content"] = expected[2].pop("reasoning") + actual = tokenizer.apply_chat_template(messages, tools=[TOOL], tokenize=False) + assert actual == encode_messages(expected, thinking_mode="thinking") + assert actual.index("system") < actual.index("## Tools") + assert "old thought" in actual + assert messages == before + + +def test_openai_tool_defaults_do_not_change_reference_prompt(tokenizer): + req = request(tools=[TOOL], tool_choice="auto") + tools = [tool.model_dump() for tool in req.tools] + assert tools[0]["function"]["description"] is None + actual = tokenizer.apply_chat_template(req.messages, tools=tools, tokenize=False) + expected = [{"role": "system", "content": "", "tools": [TOOL]}, *req.messages] + assert actual == encode_messages(expected, thinking_mode="thinking") + + +@pytest.mark.parametrize("size", [1, 2, 7, 19, 10000]) +@pytest.mark.parametrize("thinking", [False, True]) +def test_streaming_parallel_calls_and_types(tokenizer, size, thinking): + arguments = { + "text": '中文\\"\n<|DSML|parameter>literal', + "number_string": "42", + "boolean": False, + "null": None, + "array": [1, "x"], + "object": {"a": 2}, + } + text = completion(arguments, thinking).replace( + "\n", "\n" + call_text({}, "empty") + "\n" + ) + req = request(tools=[TOOL], tool_choice="auto") + parser = DeepseekV41Parser(tokenizer, chat_template_kwargs={"thinking": thinking}) + reasoning, content, calls = parser.parse(text, req, enable_auto_tools=True) + assert reasoning == (" reason " if thinking else None) + assert content == "summary" + assert [call.name for call in calls] == ["lookup", "empty"] + assert json.loads(calls[0].arguments) == arguments + assert json.loads(calls[1].arguments) == {} + parser = DeepseekV41Parser(tokenizer, chat_template_kwargs={"thinking": thinking}) + deltas = [] + for offset in range(0, len(text), size): + chunk = text[offset : offset + size] + delta = parser.parse_delta(chunk, [], req, finished=offset + size >= len(text)) + if delta: + deltas.append(delta) + assert "".join(delta.reasoning or "" for delta in deltas) == (" reason " if thinking else "") + assert "".join(delta.content or "" for delta in deltas) == "summary" + for index, expected in enumerate([arguments, {}]): + fragments = [ + tc.function.arguments or "" + for delta in deltas + for tc in delta.tool_calls or [] + if tc.index == index and tc.function + ] + assert json.loads("".join(fragments)) == expected + + +def test_registry_composes_reasoning_and_tool_parser(tokenizer): + cls = ParserManager.get_parser("deepseek_v41", "deepseek_v41", enable_auto_tools=True) + parser = cls(tokenizer) + reason, content, calls = parser.parse( + completion({"query": "value", "limit": 2}), request(tools=[TOOL], tool_choice="auto"), enable_auto_tools=True + ) + assert reason == " reason " and content == "summary" + assert json.loads(calls[0].arguments) == {"query": "value", "limit": 2} + + +def test_parser_matches_reference_completion_fields(tokenizer): + text = completion({"query": "value", "limit": 2}) + "<|end▁of▁sentence|>" + expected = parse_message_from_completion_text(text, "thinking") + parser = DeepseekV41Parser(tokenizer) + reasoning, content, calls = parser.parse(text, request(tools=[TOOL], tool_choice="auto"), enable_auto_tools=True) + assert reasoning == expected["reasoning_content"] + assert content == expected["content"] + assert calls[0].name == expected["tool_calls"][0]["function"]["name"] + assert json.loads(calls[0].arguments) == json.loads(expected["tool_calls"][0]["function"]["arguments"]) + + +def test_dsml_string_flag_overrides_schema(tokenizer): + parser = DeepseekV41Parser(tokenizer, tools=request(tools=[TOOL]).tools) + _, _, calls = parser.parse( + completion({"query": "value", "limit": "42"}), request(tools=[TOOL], tool_choice="auto"), enable_auto_tools=True + ) + assert json.loads(calls[0].arguments)["limit"] == "42" + + +@pytest.mark.parametrize("choice", ["required", {"type": "function", "function": {"name": "lookup"}}, "auto"]) +def test_structural_tag_accepts_v41_and_rejects_v4(tokenizer, choice): + tool = copy.deepcopy(TOOL) + tool["function"]["strict"] = True + req = request(tools=[tool], tool_choice=choice) + tag = DeepseekV41ToolParser(tokenizer).get_structural_tag(req) + compiler = GrammarCompiler(TokenizerInfo.from_huggingface(tokenizer)) + grammar = compiler.compile_grammar(Grammar.from_structural_tag(tag)) + text = "\n\n<|DSML| calls>\n" + call_text({"query": "value", "limit": 2}) + "\n" + matcher = GrammarMatcher(grammar) + assert matcher.accept_string(text) + assert matcher.is_completed() + if choice != "auto": + assert not GrammarMatcher(grammar).accept_string( + text.replace("|DSML| calls", "|DSML|tool_calls").replace("|DSML| ", "|DSML|") + ) + assert not GrammarMatcher(grammar).accept_string(text.replace('string="false">2', 'string="false">"wrong"')) + + +def test_renderer_uses_raw_blocks_and_keeps_media(tokenizer, monkeypatch): + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "a"}, + {"type": "image_url", "image_url": {"url": "unused.png"}}, + {"type": "text", "text": "b"}, + ], + } + ] + media = {"image": [object()]} + uuids = {"image": ["test-image"]} + monkeypatch.setattr( + "vllm_ascend.patch.platform.patch_deepseek_v41_frontend.renderer.parse_chat_messages", + lambda *a, **kw: ([{"role": "user", "content": "flattened"}], media, uuids), + ) + renderer = object.__new__(DeepseekV41Renderer) + renderer.model_config = SimpleNamespace() + renderer.get_tokenizer = lambda: tokenizer + _, prompt = renderer.render_messages( + messages, ChatParams(chat_template_kwargs={"thinking": False, "tokenize": False}) + ) + assert prompt["prompt"] == encode_messages(messages, thinking_mode="chat") + assert prompt["multi_modal_data"] is media and prompt["multi_modal_uuids"] is uuids + + +def test_renderer_async_and_request_response_format(tokenizer, monkeypatch): + req = request( + reasoning_effort="xhigh", + tools=[TOOL], + tool_choice="auto", + response_format={"type": "json_schema", "json_schema": {"name": "result", "schema": {"type": "object"}}}, + ) + conversation = [{"role": "user", "content": "hi"}] + + async def parse(*args, **kwargs): + return conversation, None, None + + monkeypatch.setattr( + "vllm_ascend.patch.platform.patch_deepseek_v41_frontend.renderer.parse_chat_messages_async", parse + ) + renderer = object.__new__(DeepseekV41Renderer) + renderer.model_config = SimpleNamespace() + renderer.get_tokenizer = lambda: tokenizer + renderer._apply_chat_template_async = make_async(renderer._apply_chat_template) + params = req.build_chat_params(None, "auto") + # OnlineRenderer adds normalized tools after building ChatParams. + params.chat_template_kwargs["tools"] = [TOOL] + params.chat_template_kwargs["tokenize"] = False + _, prompt = asyncio.run(renderer.render_messages_async(req.messages, params)) + expected = [ + {"role": "system", "content": "", "tools": [TOOL], "response_format": {"type": "object"}}, + *req.messages, + ] + assert prompt["prompt"] == encode_messages(expected, thinking_mode="thinking", reasoning_effort="xhigh") + + +@pytest.mark.parametrize( + "schema,value", + [ + ({"type": "string", "enum": ["red", "blue"]}, "red"), + ({"type": "object", "properties": {"nested": {"type": "integer"}}, "required": ["nested"]}, {"nested": 2}), + ({"type": "array", "items": {"type": "integer"}}, [1, 2]), + ({"type": ["string", "null"]}, None), + ({"type": "boolean"}, True), + ], +) +def test_grammar_parameter_schemas(tokenizer, schema, value): + tool = copy.deepcopy(TOOL) + tool["function"]["parameters"] = { + "type": "object", + "properties": {"value": schema, "optional": {"type": "integer"}}, + "required": ["value"], + } + tag = DeepseekV41ToolParser(tokenizer).get_structural_tag(request(tools=[tool], tool_choice="required")) + compiler = GrammarCompiler(TokenizerInfo.from_huggingface(tokenizer)) + grammar = compiler.compile_grammar(Grammar.from_structural_tag(tag)) + text = "\n\n<|DSML| calls>\n" + call_text({"value": value}) + "\n" + matcher = GrammarMatcher(grammar) + assert matcher.accept_string(text) and matcher.is_completed() + assert not GrammarMatcher(grammar).accept_string(text.replace('name="value"', 'name="unknown"')) + if schema.get("enum"): + assert not GrammarMatcher(grammar).accept_string(text.replace(">red<", ">green<")) + + +def test_grammar_empty_call_and_named_tool_filter(tokenizer): + tool = {"type": "function", "function": {"name": "empty", "parameters": {"type": "object", "properties": {}}}} + req = request(tools=[TOOL, tool], tool_choice={"type": "function", "function": {"name": "empty"}}) + tag = DeepseekV41ToolParser(tokenizer).get_structural_tag(req) + compiler = GrammarCompiler(TokenizerInfo.from_huggingface(tokenizer)) + grammar = compiler.compile_grammar(Grammar.from_structural_tag(tag)) + text = "\n\n<|DSML| calls>\n" + call_text({}, "empty") + "\n" + matcher = GrammarMatcher(grammar) + assert matcher.accept_string(text) and matcher.is_completed() + assert not GrammarMatcher(grammar).accept_string(text.replace('name="empty"', 'name="lookup"')) + + +def test_v4_registration_is_unchanged(): + from vllm.renderers.registry import RENDERER_REGISTRY + from vllm.tokenizers.registry import TokenizerRegistry + + old = (TokenizerRegistry.load_tokenizer_cls("deepseek_v4"), RENDERER_REGISTRY.load_renderer_cls("deepseek_v4")) + register_frontend() + assert old == ( + TokenizerRegistry.load_tokenizer_cls("deepseek_v4"), + RENDERER_REGISTRY.load_renderer_cls("deepseek_v4"), + ) + assert RENDERER_REGISTRY.load_renderer_cls("deepseek_v41") is DeepseekV41Renderer + + +@pytest.mark.parametrize("asynchronous", [False, True]) +def test_real_media_loader_preserves_interleaved_order(tokenizer, monkeypatch, asynchronous): + from vllm.entrypoints.chat_utils import BaseMultiModalItemTracker + + # Only the model processor is a stand-in; vLLM loads both actual PNGs. + processor = SimpleNamespace(info=SimpleNamespace(validate_num_items=lambda *a: None)) + monkeypatch.setattr(BaseMultiModalItemTracker, "mm_processor", processor) + monkeypatch.setattr( + BaseMultiModalItemTracker, "model_cls", SimpleNamespace(get_placeholder_str=lambda *a: "<|deepseek_image|>") + ) + blocks = [] + for color in ("red", "blue"): + buf = io.BytesIO() + Image.new("RGB", (2, 2), color).save(buf, format="PNG") + blocks.extend( + [ + {"type": "text", "text": color}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode()}, + }, + ] + ) + messages = [{"role": "user", "content": blocks}] + renderer = object.__new__(DeepseekV41Renderer) + renderer.model_config = SimpleNamespace( + hf_config=SimpleNamespace(), + multimodal_config=None, + allowed_local_media_path="", + allowed_media_domains=None, + is_multimodal_model=True, + enable_prompt_embeds=False, + ) + renderer.get_tokenizer = lambda: tokenizer + renderer._apply_chat_template_async = make_async(renderer._apply_chat_template) + params = ChatParams(chat_template_kwargs={"thinking": False, "tokenize": False}) + if asynchronous: + _, prompt = asyncio.run(renderer.render_messages_async(messages, params)) + else: + _, prompt = renderer.render_messages(messages, params) + assert prompt["prompt"] == encode_messages(messages, thinking_mode="chat") + images = prompt["multi_modal_data"]["image"] + assert [image.media.getpixel((0, 0)) for image in images] == [(255, 0, 0), (0, 0, 255)] + + +def test_global_patch_registers_frontend_in_fresh_process(): + subprocess.run( + [ + sys.executable, + "-c", + "import vllm_ascend; " + "from vllm_ascend.utils import adapt_patch; " + "adapt_patch(is_global_patch=True); " + "from vllm.tokenizers.registry import TokenizerRegistry; " + "from vllm.renderers.registry import RENDERER_REGISTRY; " + "from vllm.parser.parser_manager import ParserManager; " + "assert TokenizerRegistry.load_tokenizer_cls('deepseek_v41').__name__ == 'DeepseekV41Tokenizer'; " + "assert RENDERER_REGISTRY.load_renderer_cls('deepseek_v41').__name__ == 'DeepseekV41Renderer'; " + "assert ParserManager.get_parser('deepseek_v41', 'deepseek_v41', enable_auto_tools=True)", + ], + check=True, + capture_output=True, + text=True, + timeout=120, + ) diff --git a/tests/ut/patch/platform/test_patch_speculative_config_dspark.py b/tests/ut/patch/platform/test_patch_speculative_config_dspark.py index 7e4a1a5617be..b5d5f1cae34c 100644 --- a/tests/ut/patch/platform/test_patch_speculative_config_dspark.py +++ b/tests/ut/patch/platform/test_patch_speculative_config_dspark.py @@ -174,3 +174,115 @@ def validate_config(validated_config): patch_speculative_config._dspark_post_init(config) assert config.target_parallel_config is original_parallel_config + + +def test_deepseek_v41_dspark_selects_v41_drafter_and_expert_shape(): + text_config = SimpleNamespace( + model_type="deepseek_v4.1_text", + dspark_target_layer_ids=[37, 38, 39], + dspark_n_routed_experts=128, + dspark_n_activated_experts=3, + num_nextn_predict_layers=3, + ) + text_config.update = lambda values: text_config.__dict__.update(values) + # vLLM's generic DeepSeek-V4 dSPark normalization has already rewritten + # the composite root before the Ascend post-init hook runs. + hf_config = SimpleNamespace( + model_type="deepseek_v4", + architectures=["DSparkDraftModel"], + text_config=text_config, + ) + hf_config.update = lambda values: hf_config.__dict__.update(values) + model_arch_config = ModelArchitectureConfig( + architectures=["DeepseekV41ForConditionalGeneration"], + model_type="deepseek_v4.1", + text_model_type="deepseek_v4.1_text", + hidden_size=5120, + total_num_hidden_layers=43, + total_num_attention_heads=64, + head_size=512, + vocab_size=129280, + total_num_kv_heads=1, + num_experts=384, + num_experts_per_token=6, + quantization_config=None, + is_deepseek_mla=True, + is_mm_prefix_lm=True, + rswa_window=128, + derived_max_model_len_and_key=(1048576, "max_position_embeddings"), + ) + registry = MagicMock() + registry.inspect_model_cls.return_value = ( + "model-info", + "DeepseekV41DSparkDraftModel", + ) + draft_model_config = SimpleNamespace( + hf_config=hf_config, + model_arch_config=model_arch_config, + registry=registry, + ) + + _normalize_deepseek_v4_dspark_draft(draft_model_config) + + assert hf_config.architectures == ["DeepseekV41DSparkDraftModel"] + assert hf_config.model_type == "deepseek_v4.1" + assert text_config.n_routed_experts == 128 + assert text_config.num_experts_per_tok == 3 + assert text_config.n_mtp_layers == 3 + assert draft_model_config.model_arch_config.num_experts == 128 + if "num_experts_per_token" in ModelArchitectureConfig.__dataclass_fields__: + assert draft_model_config.model_arch_config.num_experts_per_token == 3 + registry.inspect_model_cls.assert_called_once_with(["DeepseekV41DSparkDraftModel"], draft_model_config) + + +def test_released_deepseek_v41_dspark_names_select_v41_drafter(): + text_config = SimpleNamespace( + model_type="deepseek_v41_text", + dspark_target_layer_ids=[37, 38, 39], + dspark_n_routed_experts=128, + dspark_num_experts_per_tok=3, + num_nextn_predict_layers=3, + ) + text_config.update = lambda values: text_config.__dict__.update(values) + hf_config = SimpleNamespace( + model_type="deepseek_v41", + architectures=["DeepseekV41ForCausalLM"], + text_config=text_config, + ) + hf_config.update = lambda values: hf_config.__dict__.update(values) + model_arch_config = ModelArchitectureConfig( + architectures=["DeepseekV41ForCausalLM"], + model_type="deepseek_v41", + text_model_type="deepseek_v41_text", + hidden_size=5120, + total_num_hidden_layers=43, + total_num_attention_heads=64, + head_size=512, + vocab_size=129280, + total_num_kv_heads=1, + num_experts=384, + num_experts_per_token=6, + quantization_config=None, + is_deepseek_mla=True, + is_mm_prefix_lm=False, + rswa_window=128, + derived_max_model_len_and_key=(1048576, "max_position_embeddings"), + ) + registry = MagicMock() + registry.inspect_model_cls.return_value = ( + "model-info", + "DeepseekV41DSparkDraftModel", + ) + draft_model_config = SimpleNamespace( + hf_config=hf_config, + model_arch_config=model_arch_config, + registry=registry, + ) + + _normalize_deepseek_v4_dspark_draft(draft_model_config) + + assert hf_config.architectures == ["DeepseekV41DSparkDraftModel"] + assert hf_config.model_type == "deepseek_v41" + assert text_config.n_routed_experts == 128 + assert text_config.num_experts_per_tok == 3 + assert draft_model_config.model_arch_config.num_experts_per_token == 3 diff --git a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py index 6971b3f0c57d..64cb19c80138 100644 --- a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py +++ b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py @@ -900,6 +900,26 @@ def __init__(self) -> None: self.use_eagle = False +def test_verify_and_split_accepts_one_unique_spec_across_groups() -> None: + kv_cache_config = _make_deepseek_v4_kv_cache_config() + repeated_group = kv_cache_config.kv_cache_groups[0] + coordinator = AscendHybridKVCacheCoordinator.__new__(AscendHybridKVCacheCoordinator) + coordinator.kv_cache_config = replace( + kv_cache_config, + kv_cache_groups=[repeated_group, repeated_group], + ) + coordinator.single_type_managers = (_FakeEagleManager(), _FakeEagleManager()) + coordinator.eagle_group_ids = set() + coordinator.enable_partial_hash_hits = False + coordinator.dcp_world_size = 1 + coordinator.enable_caching = False + + coordinator.verify_and_split_kv_cache_groups() + + assert len(coordinator.attention_groups) == 1 + assert coordinator.attention_groups[0].group_ids == [0, 1] + + def test_verify_and_split_propagates_eagle_to_managers() -> None: """Regression for DeepSeek-V4 prefix-cache hit rate 0% with MTP/EAGLE. diff --git a/tests/ut/quantization/configs/test_modelslim_config.py b/tests/ut/quantization/configs/test_modelslim_config.py index 5d6a9295deb9..cae36b045db4 100644 --- a/tests/ut/quantization/configs/test_modelslim_config.py +++ b/tests/ut/quantization/configs/test_modelslim_config.py @@ -67,6 +67,27 @@ def test_from_config(self): self.assertIsInstance(config, AscendModelSlimConfig) self.assertEqual(config.quant_description, self.sample_config) + def test_from_metadata_only_config_defers_description_load(self): + config = AscendModelSlimConfig.from_config({"quant_method": "ascend", "model_quant_type": "W8A8_DYNAMIC"}) + self.assertEqual(config.quant_description, {}) + + def test_deepseek_v41_packed_mapping_uses_checkpoint_shard_names(self): + self.ascend_config._update_packed_modules_mapping("deepseek_v4.1") + self.assertEqual( + self.ascend_config.packed_modules_mapping["gate_up_proj"], + ["w1", "w3"], + ) + self.assertEqual( + self.ascend_config.packed_modules_mapping["experts"], + ["experts.0.w1", "experts.0.w2", "experts.0.w3"], + ) + + def test_deepseek_v41_quant_prefix_maps_terminal_projection(self): + self.assertEqual( + self.ascend_config.quant_prefix_mapper("deepseek_v4.1", "model.layers.0.mlp.shared_experts.down_proj"), + "layers.0.ffn.shared_experts.w2", + ) + @patch("vllm_ascend.quantization.configs.modelslim_config.torch.npu.is_available") def test_override_quantization_method(self, mock_is_available): # Test when NPU is available @@ -592,6 +613,27 @@ def test_deepseek_v4_vision_maps_wrapped_language_model_prefixes(self): expected, ) + def test_deepseek_v41_vision_maps_wrapped_language_model_prefixes(self): + config = AscendModelSlimConfig( + { + "embed.weight": "W8A8_DYNAMIC", + "layers.0.attn.q_proj.weight": "W8A8_DYNAMIC", + "head.weight": "FLOAT", + } + ) + + cases = { + "language_model.model.embed_tokens": "embed", + "language_model.model.layers.0.self_attn.q_proj": ("layers.0.attn.q_proj"), + "language_model.lm_head": "head", + } + for prefix, expected in cases.items(): + with self.subTest(prefix=prefix): + self.assertEqual( + config.quant_prefix_mapper("deepseek_v4.1", prefix), + expected, + ) + def test_lm_head_maps_to_language_model_lm_head_when_quant_key_exists(self): config = AscendModelSlimConfig({"language_model.lm_head.weight": "FLOAT"}) diff --git a/tests/ut/spec_decode/test_dspark_proposer.py b/tests/ut/spec_decode/test_dspark_proposer.py index 7e0498e56db8..8e1024f6bdc9 100644 --- a/tests/ut/spec_decode/test_dspark_proposer.py +++ b/tests/ut/spec_decode/test_dspark_proposer.py @@ -820,8 +820,38 @@ def _make_proposer_for_init(): proposer.device = torch.device("cpu") proposer.runner = SimpleNamespace(device_metadata_executor=None) proposer.dcp_size = 1 + proposer._per_group_block_tables = {} + proposer._per_group_slot_mappings = {} return proposer + def test_aurora_draft_uses_only_group_twelve(self, monkeypatch): + from tests.deepseek_v41_cache_utils import make_cache_config + + config = make_cache_config(17, draft_layers=3) + draft_names = config.kv_cache_groups[12].layer_names + backend = MagicMock() + backend.full_cls_name.return_value = "AscendDSASWABackend" + modules = {name: SimpleNamespace(get_attn_backend=lambda: backend) for name in draft_names} + monkeypatch.setattr( + "vllm_ascend.spec_decode.dspark_proposer.get_layers_from_vllm_config", lambda *args, **kw: modules + ) + proposer = self._make_proposer_for_init() + proposer.model = SimpleNamespace(get_draft_kv_cache_layer_names=lambda: draft_names) + proposer.max_query_tokens = 16 + proposer.max_num_tokens = 32 + with patch.object(AttentionGroup, "create_metadata_builders"): + proposer.initialize_attn_backend(config, [128, 32] + [128] * 11) + assert proposer.kv_cache_gid == 12 + assert len(proposer.draft_attn_groups) == 1 + assert set(proposer.draft_attn_groups[0].layer_names) == set(draft_names) + assert proposer._layer_group_idx == [12, 12, 12] + target_table = torch.tensor([[2]], dtype=torch.int32) + draft_table = torch.tensor([[7]], dtype=torch.int32) + proposer.set_per_group_attn_metadata(2, target_table, torch.tensor([256])) + proposer.set_per_group_attn_metadata(12, draft_table, torch.tensor([896])) + assert proposer._per_group_block_tables[12] is draft_table + assert proposer._per_group_block_tables[12] is not target_table + @pytest.mark.parametrize( ("dcp_size", "pcp_enabled", "has_executor", "expected_tokens"), [ diff --git a/tests/ut/worker/test_model_runner_v1.py b/tests/ut/worker/test_model_runner_v1.py index 79c38eccfd6d..73e77523ee0b 100644 --- a/tests/ut/worker/test_model_runner_v1.py +++ b/tests/ut/worker/test_model_runner_v1.py @@ -6,11 +6,12 @@ import numpy as np import torch -from vllm.config import CUDAGraphMode +from vllm.config import CompilationConfig, CUDAGraphMode from vllm.model_executor.layers.attention import MLAAttention from vllm.model_executor.models.deepseek_v2 import DeepseekV32IndexerCache from vllm.sampling_params import SamplingParams from vllm.v1.attention.backends.utils import reorder_batch_to_split_decodes_and_prefills +from vllm.v1.cudagraph_dispatcher import CudagraphDispatcher from vllm.v1.kv_cache_interface import ( FullAttentionSpec, HiddenStateCacheSpec, @@ -25,6 +26,7 @@ from vllm.v1.worker.gpu_input_batch import CachedRequestState, InputBatch from vllm.v1.worker.gpu_model_runner import GPUModelRunner +from tests.deepseek_v41_cache_utils import make_cache_config from vllm_ascend.attention.mla_v1 import AscendMLABackend from vllm_ascend.attention.utils import get_sfa_qsfa_packed_head_dim from vllm_ascend.core.kv_cache_interface import ( @@ -97,6 +99,107 @@ def test_partial_prompt_with_state_dispatches_speculative_decode_graph(self): class TestDummyRunSlotInvalidation(unittest.TestCase): + def test_padded_speculative_dummy_preserves_logical_query_lengths(self): + # DSA CP rounds 186 tokens to 192 without adding a logical request. + # Also cover dispatchers that pad the request count itself. + for num_tokens, padded_tokens, padded_reqs in ((186, 192, 31), (12, 24, 4), (6, 24, 4), (12, 12, 2)): + with self.subTest(num_tokens=num_tokens): + runner = NPUModelRunner.__new__(NPUModelRunner) + runner.uniform_decode_query_len = 6 + runner.scheduler_config = SimpleNamespace(max_num_batched_tokens=4096, max_num_seqs=32) + runner.dynamic_eplb = False + runner.dcp_size = 1 + runner.speculative_config = None + runner._has_gdn = True + runner.vllm_config = MagicMock() + agreed_counts = torch.tensor([padded_tokens, 192], dtype=torch.int32) + runner._determine_batch_execution_and_padding = MagicMock( + return_value=( + CUDAGraphMode.FULL, + SimpleNamespace(num_tokens=padded_tokens, num_reqs=padded_reqs), + False, + agreed_counts, + None, + ) + ) + runner.synchronize_input_prep = nullcontext + runner._should_build_dummy_attn_metadata = MagicMock(return_value=True) + runner.optimistic_seq_lens_cpu = torch.zeros(32, dtype=torch.int32) + runner.seq_lens = MagicMock() + runner.arange_np = np.arange(4096, dtype=np.int32) + runner.query_pos = SimpleNamespace(np=np.zeros(4096, dtype=np.int32)) + runner.query_start_loc = SimpleNamespace(np=np.zeros(33, dtype=np.int32), copy_to_gpu=MagicMock()) + runner.gdn_query_start_loc = SimpleNamespace(np=np.zeros(33, dtype=np.int32), copy_to_gpu=MagicMock()) + + def check_offsets( + *args, + runner=runner, + num_tokens=num_tokens, + padded_tokens=padded_tokens, + padded_reqs=padded_reqs, + agreed_counts=agreed_counts, + ): + num_reqs = num_tokens // 6 + expected = np.concatenate( + (np.arange(num_reqs + 1) * 6, np.full(padded_reqs - num_reqs, num_tokens)) + ) + np.testing.assert_array_equal(runner.query_start_loc.np[: padded_reqs + 1], expected) + np.testing.assert_array_equal(runner.gdn_query_start_loc.np[: padded_reqs + 1], expected) + np.testing.assert_array_equal(runner.query_pos.np[:num_tokens], np.tile(np.arange(6), num_reqs)) + torch.testing.assert_close(agreed_counts, torch.tensor([padded_tokens, 192], dtype=torch.int32)) + raise RuntimeError("logical query lengths checked") + + get_cumsum = runner._get_cumsum_and_arange + + def check_padded_schedule( + schedule, + arange, + num_tokens=num_tokens, + padded_reqs=padded_reqs, + get_cumsum=get_cumsum, + ): + num_reqs = num_tokens // 6 + expected = np.concatenate((np.full(num_reqs, 6), np.zeros(padded_reqs - num_reqs))) + np.testing.assert_array_equal(schedule, expected) + self.assertEqual(schedule.dtype, np.int32) + self.assertEqual(int(schedule.sum()), num_tokens) + return get_cumsum(schedule, arange) + + runner._get_cumsum_and_arange = check_padded_schedule + runner._pad_query_start_loc_for_fia = check_offsets + with ( + patch("vllm_ascend.worker.model_runner_v1.using_paged_attention", return_value=False), + self.assertRaisesRegex(RuntimeError, "logical query lengths checked"), + ): + runner._dummy_run(num_tokens, uniform_decode=True, cudagraph_runtime_mode=CUDAGraphMode.FULL) + + def test_padded_dummy_preserves_other_dp_token_counts(self): + for token_counts in ([24, 8], [8, 24], [8, 8]): + with self.subTest(token_counts=token_counts): + runner = NPUModelRunner.__new__(NPUModelRunner) + runner.uniform_decode_query_len = 1 + runner.scheduler_config = SimpleNamespace(max_num_batched_tokens=32, max_num_seqs=4) + runner.dynamic_eplb = False + runner.dcp_size = 1 + agreed_counts = torch.tensor(token_counts, dtype=torch.int32) + runner._determine_batch_execution_and_padding = MagicMock( + return_value=( + CUDAGraphMode.NONE, + SimpleNamespace(num_tokens=8, num_reqs=1), + False, + agreed_counts, + None, + ) + ) + + def check_agreed_counts(agreed_counts=agreed_counts, token_counts=token_counts): + torch.testing.assert_close(agreed_counts, torch.tensor(token_counts, dtype=torch.int32)) + raise RuntimeError("DP token counts checked") + + runner.synchronize_input_prep = check_agreed_counts + with self.assertRaisesRegex(RuntimeError, "DP token counts checked"): + runner._dummy_run(1, uniform_decode=True) + def test_backend_metadata_sees_invalidated_dummy_slots(self): runner = NPUModelRunner.__new__(NPUModelRunner) runner.kvpp = KVPPRuntime() @@ -130,7 +233,12 @@ def test_backend_metadata_sees_invalidated_dummy_slots(self): slot_mapping=SimpleNamespace(gpu=slot_mappings[index]) ) runner.input_batch = SimpleNamespace(block_table=block_tables) - runner.kv_cache_config = SimpleNamespace(kv_cache_groups=[object(), object()]) + runner.kv_cache_config = SimpleNamespace( + kv_cache_groups=[ + SimpleNamespace(kv_cache_spec=object()), + SimpleNamespace(kv_cache_spec=object()), + ] + ) def check_slots_before_build(**_kwargs): for slot_mapping in slot_mappings: @@ -142,6 +250,64 @@ def check_slots_before_build(**_kwargs): with self.assertRaisesRegex(RuntimeError, "metadata checked"): runner._dummy_run(1) + def test_graph_capture_invalidates_only_v41_active_slots(self): + runner = NPUModelRunner.__new__(NPUModelRunner) + runner.uniform_decode_query_len = 1 + runner.scheduler_config = SimpleNamespace(max_num_batched_tokens=8, max_num_seqs=8) + runner.dynamic_eplb = False + runner.dcp_size = 1 + runner.speculative_config = None + runner.use_compress = True + runner._has_gdn = False + runner.vllm_config = MagicMock() + + runner._determine_batch_execution_and_padding = MagicMock( + return_value=( + CUDAGraphMode.FULL, + SimpleNamespace(num_tokens=2, num_reqs=2), + None, + None, + None, + ) + ) + runner._should_build_dummy_attn_metadata = MagicMock(return_value=True) + runner.synchronize_input_prep = MagicMock(return_value=nullcontext()) + runner._get_cumsum_and_arange = MagicMock(return_value=np.array([1, 2], dtype=np.int32)) + runner._pad_query_start_loc_for_fia = MagicMock(return_value=2) + + runner.optimistic_seq_lens_cpu = torch.zeros(8, dtype=torch.int32) + runner.seq_lens = MagicMock() + runner.query_pos = SimpleNamespace(np=np.zeros(8, dtype=np.int32)) + runner.query_start_loc = SimpleNamespace(np=np.zeros(9, dtype=np.int32), copy_to_gpu=MagicMock()) + runner.positions = MagicMock() + runner._dsa_positions_cpu_buf = MagicMock() + + v41_group = make_cache_config(17).kv_cache_groups[0] + other_group = SimpleNamespace(kv_cache_spec=object()) + slot_mappings = [torch.tensor([3, 4]), torch.tensor([7, 8])] + block_tables = MagicMock() + block_tables.__getitem__.side_effect = lambda index: SimpleNamespace( + slot_mapping=SimpleNamespace(gpu=slot_mappings[index]) + ) + runner.input_batch = SimpleNamespace(block_table=block_tables) + runner.kv_cache_config = SimpleNamespace( + kv_cache_groups=[v41_group, other_group], + num_blocks=17, + ) + + def check_slots_before_build(**_kwargs): + torch.testing.assert_close(slot_mappings[0], torch.full_like(slot_mappings[0], -1)) + torch.testing.assert_close(slot_mappings[1], torch.tensor([7, 8])) + raise RuntimeError("metadata checked") + + runner._build_attention_metadata = check_slots_before_build + + with ( + patch("vllm_ascend.worker.model_runner_v1.using_paged_attention", return_value=False), + self.assertRaisesRegex(RuntimeError, "metadata checked"), + ): + runner._dummy_run(2, cudagraph_runtime_mode=CUDAGraphMode.FULL, is_graph_capturing=True) + class TestDeviceMetadataFullGraphEvents(unittest.TestCase): def test_full_mode_requires_external_events(self): @@ -513,6 +679,74 @@ def test_kvpp_allocate_and_reshape_views(self): caches = runner._reshape_kv_cache_tensors(cache_config, raw) assert_attention_cache_views(caches, raw, packed) + def test_v41_layer_outer_buffers_allocate_and_reshape(self): + runner = self._build_runner() + config = make_cache_config(4) + # vLLM shrinks each tensor proportionally when another rank has less + # capacity. Component offsets must not depend on the old block count. + for allocation in config.kv_cache_tensors: + allocation.size = allocation.size // config.num_blocks * 3 + config.num_blocks = 3 + raw = runner._allocate_kv_cache_tensors(config) + prefix = "model.layers." + long_name = prefix + "2.self_attn.long_kv_cache" + index_name = prefix + "2.self_attn.indexer.k_cache" + assert raw[long_name] is raw[index_name] + assert raw[long_name] is raw[prefix + "0.self_attn.swa_cache"] + assert raw[long_name] is not raw[prefix + "8.self_attn.long_kv_cache"] + unique = {id(value): value for value in raw.values()} + assert len(unique) == 4 + assert sum(buffer.numel() for buffer in unique.values()) == 3 * 540928 + runner._kv_cache_spec_attn_group_iterator = lambda: [ + SimpleNamespace(kv_cache_spec=spec, backend=runner.attn_backend, layer_names=[name]) + for group in config.kv_cache_groups + for name, spec in group.kv_cache_spec.kv_cache_specs.items() + ] + caches = runner._reshape_kv_cache_tensors(config, raw) + key, scale = caches[index_name] + assert key.shape == (3, 64, 1, 128) + assert scale.shape == (3, 64, 1, 1) + assert key.data_ptr() - caches[long_name].data_ptr() == 65536 + assert scale.data_ptr() - key.data_ptr() == 8192 + assert key.stride(0) == 131072 + assert scale.stride(0) == 65536 + assert not key.is_contiguous() + assert caches[prefix + "0.self_attn.swa_cache"].is_contiguous() + assert not caches[prefix + "3.self_attn.swa_cache"].is_contiguous() + + def test_v41_rejects_obsolete_allocation_descriptors(self): + runner = self._build_runner() + config = make_cache_config(3) + config.kv_cache_tensors[0].block_stride = 0 + with self.assertRaisesRegex(ValueError, "allocation disagrees"): + runner._allocate_kv_cache_tensors(config) + + def test_v41_dspark_shares_four_backings_after_rank_shrink(self): + runner = self._build_runner() + config = make_cache_config(5, draft_layers=3) + for allocation in config.kv_cache_tensors: + allocation.size = allocation.size // config.num_blocks * 3 + config.num_blocks = 3 + raw = runner._allocate_kv_cache_tensors(config) + assert len({id(value) for value in raw.values()}) == 4 + runner._kv_cache_spec_attn_group_iterator = lambda: [ + SimpleNamespace(kv_cache_spec=spec, backend=runner.attn_backend, layer_names=[name]) + for group in config.kv_cache_groups + for name, spec in group.kv_cache_spec.kv_cache_specs.items() + ] + caches = runner._reshape_kv_cache_tensors(config, raw) + for stage, source in enumerate((2, 8, 14)): + draft = f"mtp.{stage}.self_attn.swa_cache" + target = f"model.layers.{source}.self_attn.long_kv_cache" + assert raw[draft] is raw[target] + assert caches[draft].shape == (3, 128, 1, 512) + assert caches[draft].stride(0) * 2 == 131072 + assert caches[draft].data_ptr() == caches[target].data_ptr() + caches[target][1].fill_(2) + caches[draft][2].fill_(3) + assert (caches[target][1] == 2).all() + assert (caches[draft][0] == 0).all() + def test_allocate_kv_cache_uses_layer_spec_for_draft_gqa(self): runner = self._build_runner() runner.sparse_kv_offload_enabled = False @@ -726,34 +960,64 @@ def get_builder_cls(cls): {(target_layer,), (draft_layer,), (cache_layer,)}, ) - def test_explicit_capture_sizes_must_align_spec_decode_and_sp(self): - for capture_sizes, expected_tp_size in (([48, 96], 1), ([16, 32], 16)): - with self.subTest(capture_sizes=capture_sizes): + def test_cp_or_sp_decode_dispatch_keys_are_actually_captured(self): + cases = ( + (True, 8, 6, False), + (True, 16, 6, False), + (True, 8, 6, True), + (True, 8, 1, False), + (False, 8, 6, True), + (False, 16, 6, True), + (False, 8, 1, True), + ) + for cp_enabled, tp_size, query_len, enable_sp in cases: + with self.subTest(cp=cp_enabled, tp_size=tp_size, query_len=query_len, enable_sp=enable_sp): runner = self._build_runner() - compilation_config = SimpleNamespace( - pass_config=SimpleNamespace(enable_sp=True), - cudagraph_capture_sizes=capture_sizes, - resolve_cudagraph_mode_and_sizes=MagicMock(return_value=CUDAGraphMode.FULL_DECODE_ONLY), + max_tokens = 32 * query_len + config = CompilationConfig( + cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY, + cudagraph_capture_sizes=list(range(tp_size, max_tokens + 1, tp_size)), + max_cudagraph_capture_size=max_tokens, ) - runner.compilation_config = compilation_config - runner.vllm_config.compilation_config = compilation_config - runner.parallel_config = SimpleNamespace(tensor_parallel_size=16) - runner.uniform_decode_query_len = 6 - runner.kv_cache_config = SimpleNamespace() - runner.max_num_reqs = 16 - runner.cudagraph_dispatcher = MagicMock() - runner.cudagraph_dispatcher.get_capture_descs.return_value = [] + config.pass_config.enable_sp = enable_sp + runner.compilation_config = config + runner.parallel_config = SimpleNamespace(tensor_parallel_size=tp_size) + runner.vllm_config = SimpleNamespace( + compilation_config=config, + parallel_config=runner.parallel_config, + num_speculative_tokens=query_len - 1, + lora_config=None, + scheduler_config=SimpleNamespace(max_num_seqs=32), + ) + runner.uniform_decode_query_len = query_len + runner.kv_cache_config = SimpleNamespace(has_mamba_layers=False) + runner.max_num_reqs = 32 + runner.cudagraph_dispatcher = CudagraphDispatcher(runner.vllm_config) runner.speculative_config = None runner.drafter = None runner.use_aclgraph = False - - runner._check_and_update_cudagraph_mode([], []) - - call_kwargs = compilation_config.resolve_cudagraph_mode_and_sizes.call_args.kwargs - self.assertEqual( - call_kwargs["tensor_parallel_size"], - expected_tp_size, - ) + with ( + patch("vllm_ascend.worker.model_runner_v1.enable_dsa_cp", return_value=cp_enabled), + patch("vllm_ascend.worker.model_runner_v1.enable_sp", return_value=enable_sp), + patch("vllm_ascend.worker.model_runner_v1.update_pass_config", return_value=nullcontext()), + ): + runner._check_and_update_cudagraph_mode([], []) + dispatcher = runner.cudagraph_dispatcher + captured = set() + for mode, descriptors in dispatcher.get_capture_descs(): + for descriptor in descriptors: + actual_mode, actual_key = dispatcher.dispatch( + runner._pad_for_sequence_parallelism(descriptor.num_tokens), uniform_decode=True + ) + self.assertEqual((actual_mode, actual_key), (mode, descriptor)) + captured.add(actual_key) + # Includes the one-request dummy used by idle DP ranks. + for num_reqs in range(1, 33): + mode, key = dispatcher.dispatch( + runner._pad_for_sequence_parallelism(num_reqs * query_len), uniform_decode=True + ) + self.assertEqual(mode, CUDAGraphMode.FULL) + self.assertIn(key, captured) def test_sparse_c8_indexer_reuses_raw_cache_from_shared_descriptor(self): runner = self._build_runner() @@ -2178,7 +2442,11 @@ def test_history_gate_uses_only_actual_requests(self): runner.model = object() runner.vllm_config = SimpleNamespace() runner.parallel_config = SimpleNamespace(num_ubatches=1) - runner.model_config = SimpleNamespace(enforce_eager=True) + runner.model_config = SimpleNamespace( + enforce_eager=True, + hf_config=SimpleNamespace(model_type="test"), + hf_text_config=SimpleNamespace(model_type="test"), + ) runner.cache_config = SimpleNamespace(kv_sharing_fast_prefill=False, mamba_cache_mode=None) runner.input_batch = SimpleNamespace( num_reqs=2, req_ids=["a", "b"], num_computed_tokens_cpu=np.array(computed) diff --git a/tests/ut/worker/test_worker_v1.py b/tests/ut/worker/test_worker_v1.py index 93ace797b15f..a232ca577a2a 100644 --- a/tests/ut/worker/test_worker_v1.py +++ b/tests/ut/worker/test_worker_v1.py @@ -930,7 +930,11 @@ def test_execute_dummy_batch(self): worker.execute_dummy_batch() # Verify call - mock_model_runner._dummy_run.assert_called_once_with(mock_uniform_decode_query_len, uniform_decode=True) + mock_model_runner._dummy_run.assert_called_once_with( + mock_uniform_decode_query_len, + uniform_decode=True, + skip_gdn_state_update=True, + ) @patch("vllm_ascend.worker.worker.plan_sparse_kv_offload_memory") @patch("vllm_ascend.worker.worker.get_ascend_config") diff --git a/vllm_ascend/ascend_config.py b/vllm_ascend/ascend_config.py index a82543c7b814..e2973cff444b 100644 --- a/vllm_ascend/ascend_config.py +++ b/vllm_ascend/ascend_config.py @@ -446,6 +446,10 @@ class AscendConfig: # ---- user-input switches: bool/int/list/str, auto type validation ---- enable_cpu_binding: bool = True + # Enable the V4.1 node-sharded Engram path. + enable_engram: bool = True + # V4.1 node-sharded Engram storage; BF16 output and projections are unchanged. + engram_storage: Literal["bf16", "int8", "fp8", "mxfp8"] = "bf16" multistream_dsv4_dsa_overlap: bool = True enable_prefill_mc2: bool = False multistream_overlap_shared_expert: bool = False diff --git a/vllm_ascend/attention/context_parallel/dsa_common.py b/vllm_ascend/attention/context_parallel/dsa_common.py new file mode 100644 index 000000000000..ba0f7f7a6f85 --- /dev/null +++ b/vllm_ascend/attention/context_parallel/dsa_common.py @@ -0,0 +1,23 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Shared DSA token layouts for replicated-cache context parallelism.""" + +import torch +import torch.distributed as dist + + +def restore_tp_heads(output, tp_group): + """Exchange [local tokens, all heads] for [all tokens, local heads].""" + if tp_group.world_size == 1: + return output + tokens, heads, width = output.shape + local_heads = heads // tp_group.world_size + send = ( + output.view(tokens, tp_group.world_size, local_heads, width) + .permute(1, 0, 2, 3) + .contiguous() + .view(-1, local_heads, width) + ) + recv = torch.empty_like(send) + dist.all_to_all_single(recv, send, group=tp_group.device_group) + return recv diff --git a/vllm_ascend/attention/context_parallel/dsa_cp.py b/vllm_ascend/attention/context_parallel/dsa_cp.py index 5c22db5a5041..bf2f39184762 100644 --- a/vllm_ascend/attention/context_parallel/dsa_cp.py +++ b/vllm_ascend/attention/context_parallel/dsa_cp.py @@ -3,7 +3,6 @@ from typing import TYPE_CHECKING, Any, ClassVar, TypeAlias import torch -import torch.distributed as dist import torch.nn.functional as F import torch_npu from vllm.config import CUDAGraphMode, VllmConfig, get_current_vllm_config @@ -13,8 +12,10 @@ from vllm.v1.attention.backend import AttentionCGSupport, AttentionImplBase, AttentionMetadataBuilder from vllm.v1.kv_cache_interface import AttentionSpec +from vllm_ascend.ascend_config import get_ascend_config from vllm_ascend.attention import dsa_v1 from vllm_ascend.attention.attention_v1 import AscendAttentionState +from vllm_ascend.attention.context_parallel.dsa_common import restore_tp_heads from vllm_ascend.attention.dsa_attn_kv_plan import ( get_dsa_attn_kv_plan, is_a5_bf16_kv_enabled, @@ -41,6 +42,7 @@ from vllm_ascend.models.common.ops.sequence_parallel import sp_reduce_scatter from vllm_ascend.models.deepseek_v4.compressor import AscendCompressorMetadata from vllm_ascend.models.deepseek_v4.indexer import AscendIndexerMetadata +from vllm_ascend.ops.cv_linear import CVLinearWrapper from vllm_ascend.ops.linear import AscendUnquantizedLinearMethod from vllm_ascend.ops.rope_dsv4 import RopeDataProxy, get_cos_and_sin_dsa, get_full_cos_and_sin_dsa from vllm_ascend.ops.triton.dsa_cp import build_local_metadata_triton @@ -1403,6 +1405,14 @@ def __init__( self.vllm_config = kwargs.get("vllm_config", get_current_vllm_config()) + # V4.1 CP preprocessing uses the same split projections as ordinary DSA. + self.cv_wq_a = CVLinearWrapper(self.wq_a) + self.cv_wkv = CVLinearWrapper(self.wkv) + self.cv_wq_b = CVLinearWrapper(self.wq_b) + self.multistream_dsv4_dsa_overlap = get_ascend_config().multistream_dsv4_dsa_overlap + if self.multistream_dsv4_dsa_overlap and is_a5_bf16_kv_enabled(self.vllm_config): + self.multistream_dsv4_dsa_overlap = False + # indexer param if self.indexer is not None: self.indexer_heads: int = self.indexer.n_heads @@ -1658,7 +1668,9 @@ def forward( # type: ignore[override] ) num_tokens = o_proj_input.shape[0] - # o + # Keep gathered projection weights alive until all asynchronous NPU + # consumers, including reduce-scatter and the output copy, are queued. + # V4.1 calls _forward_o_proj only without temporary gathered weights. if full_gather_wo_a_enabled: self._switch_o_proj_to_full_weight() o_proj_groups = self.n_group if full_gather_wo_a_enabled else self.n_local_groups @@ -1686,9 +1698,6 @@ def forward( # type: ignore[override] o_proj_input = o.reshape(num_tokens, -1) else: o_proj_input = o_proj_input.view(num_tokens, o_proj_groups, -1) - # wo_a = self.wo_a.weight.view(o_proj_groups, self.o_lora_rank, -1) - # o = torch.einsum("tgd,grd->tgr", o, wo_a) - # A5 BF16 uses the same 3D [groups, hidden, rank] layout. o_proj_input = torch_npu.npu_transpose_batchmatmul( o_proj_input, self._get_batched_wo_a_weight(o_proj_groups), @@ -1712,6 +1721,56 @@ def forward( # type: ignore[override] return output + def _forward_o_proj(self, o_proj_input, full_gather_wo_a_enabled=False): + """Project CP attention output with TP or temporarily gathered weights.""" + num_tokens = o_proj_input.shape[0] + # o + if full_gather_wo_a_enabled: + self._switch_o_proj_to_full_weight() + o_proj_groups = self.n_group if full_gather_wo_a_enabled else self.n_local_groups + try: + use_a5_quant_o_proj = self.support_fp8_attention and _has_weight_scale(self.wo_a) + if use_a5_quant_o_proj: + o = o_proj_input.view(num_tokens, o_proj_groups, -1) + wo_a_method = getattr(self.wo_a.quant_method, "quant_method", self.wo_a.quant_method) + if isinstance(wo_a_method, AscendUnquantizedLinearMethod): + o = torch.bmm(o.transpose(0, 1), self._get_batched_wo_a_weight(o_proj_groups)).transpose(0, 1) + else: + o, swiglu_out_scale = torch_npu.npu_dynamic_mx_quant(o, dst_type=torch.float8_e4m3fn) + o = torch_npu.npu_transpose_quant_batchmatmul( + o, + self._get_batched_wo_a_weight(o_proj_groups), + dtype=torch.bfloat16, + bias=None, + group_sizes=(0, 0, 32), + x1_scale=swiglu_out_scale.view(torch.float8_e8m0fnu), + x2_scale=self._get_batched_wo_a_scale(o_proj_groups).view(torch.float8_e8m0fnu), + perm_x1=(1, 0, 2), + perm_x2=(0, 1, 2), + perm_y=(1, 0, 2), + ) + o_proj_input = o.reshape(num_tokens, -1) + else: + o_proj_input = o_proj_input.view(num_tokens, o_proj_groups, -1) + # wo_a = self.wo_a.weight.view(o_proj_groups, self.o_lora_rank, -1) + # o = torch.einsum("tgd,grd->tgr", o, wo_a) + # A5 BF16 uses the same 3D [groups, hidden, rank] layout. + o_proj_input = torch_npu.npu_transpose_batchmatmul( + o_proj_input, + self._get_batched_wo_a_weight(o_proj_groups), + bias=None, + scale=None, + perm_x1=(1, 0, 2), + perm_x2=(0, 1, 2), + perm_y=(1, 0, 2), + batch_split_factor=1, + ) + o_proj_input = o_proj_input.reshape(num_tokens, -1) + return self._apply_wo_b(o_proj_input, full_gather_wo_a_enabled) + finally: + if full_gather_wo_a_enabled: + self._switch_o_proj_to_local_weight() + def _forward( self, layer_name, @@ -1949,7 +2008,6 @@ def _restore_tp_head_layout( assert attn_metadata.req_metadata is not None req_metadata = attn_metadata.req_metadata cp_metadata = req_metadata.cp_metadata - num_tokens = local_attn_output.shape[0] torch.ops._C_ascend.inplace_partial_rotary_mul( local_attn_output.unsqueeze(1), cp_metadata.local_cos[layer_name], @@ -1962,15 +2020,7 @@ def _restore_tp_head_layout( if self.tp_size == 1 or skip_all_to_all: return local_attn_output - send = ( - local_attn_output.view(num_tokens, self.tp_size, self.n_local_heads, self.head_dim) - .permute(1, 0, 2, 3) - .contiguous() - .view(-1, self.n_local_heads, self.head_dim) - ) - recv = torch.empty_like(send) - dist.all_to_all_single(recv, send, group=self.tp_group.device_group) - return recv + return restore_tp_heads(local_attn_output, self.tp_group) def _update_indexer_cache( self, diff --git a/vllm_ascend/attention/context_parallel/dsa_v41_cp.py b/vllm_ascend/attention/context_parallel/dsa_v41_cp.py new file mode 100644 index 000000000000..adae91e1bc63 --- /dev/null +++ b/vllm_ascend/attention/context_parallel/dsa_v41_cp.py @@ -0,0 +1,252 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""V4.1 replicated-cache TP-token DSA CP adapter.""" + +from dataclasses import replace + +import torch +from vllm.distributed import get_tp_group +from vllm.forward_context import get_forward_context + +from vllm_ascend.attention.context_parallel.dsa_common import restore_tp_heads +from vllm_ascend.attention.context_parallel.dsa_cp import AscendDSACPMetadataBuilder +from vllm_ascend.attention.dsa_v1 import dsv4_dsa_overlap_stream +from vllm_ascend.attention.dsa_v41 import ( + DeepseekV41EagerAttentionImpl, + DeepseekV41MetadataBuilder, + _config_value, + scatter_cache_sk, +) +from vllm_ascend.attention.utils import enable_pcp +from vllm_ascend.utils import enable_dsa_cp, npu_stream_switch + + +def get_v41_cp_classes(): + if enable_pcp(): + raise NotImplementedError("V4.1 PCP is not supported") + if enable_dsa_cp(): + return DeepseekV41CPMetadataBuilder, DeepseekV41CPImpl + return DeepseekV41MetadataBuilder, DeepseekV41EagerAttentionImpl + + +class _ReplicatedCacheMetadataBuilder(DeepseekV41MetadataBuilder): + """Keep global cache metadata independent from local query buffers.""" + + def __init__(self, kv_cache_spec, layer_names, vllm_config, device): + super().__init__(kv_cache_spec, layer_names, vllm_config, device, build_compressor_metadata=False) + self._global_builder = DeepseekV41MetadataBuilder( + kv_cache_spec, layer_names, vllm_config, device, build_query_metadata=False + ) + + def enable_device_metadata(self): + super().enable_device_metadata() + self._global_builder.enable_device_metadata() + + def take_device_metadata_tasks(self): + return ( + *self._global_builder.take_device_metadata_tasks(), + *super().take_device_metadata_tasks(), + ) + + def _build_global_metadata(self, common_prefix_len, common, fast_build, kwargs): + global_kwargs = dict(kwargs) + shared = kwargs.get("common_v41_metadata") + if shared is not None: + global_kwargs["common_v41_metadata"] = shared.setdefault("cp_global", {}) + batch_shared = kwargs.get("common_v41_batch_metadata") + if batch_shared is not None: + global_kwargs["common_v41_batch_metadata"] = batch_shared.setdefault("cp_global", {}) + return self._global_builder.build(common_prefix_len, common, fast_build, **global_kwargs) + + +class DeepseekV41CPMetadataBuilder(_ReplicatedCacheMetadataBuilder): + def __init__(self, kv_cache_spec, layer_names, vllm_config, device): + super().__init__(kv_cache_spec, layer_names, vllm_config, device) + # SMLA consumes INT32 offsets at a fixed address during graph replay. + self._cp_query_start_loc = self._seq_lens.new_zeros(self._seq_lens.numel() + 1) + + # Reuse Legacy DSACP's request intersection and causal-prefix calculation. + _local_token_range = staticmethod(AscendDSACPMetadataBuilder._local_token_range) + + def build(self, common_prefix_len, common_attn_metadata, fast_build=False, **kwargs): + common = common_attn_metadata + global_metadata = self._build_global_metadata(common_prefix_len, common, fast_build, kwargs) + seq_lens_cpu = ( + common._seq_lens_cpu if getattr(common, "_seq_lens_cpu", None) is not None else common.seq_lens_cpu + ) + start, end, per_rank, padded, qsl, seq_lens = AscendDSACPMetadataBuilder._build_local_token_metadata( + self, + common.num_reqs, + common.num_input_tokens, + common.query_start_loc_cpu, + seq_lens_cpu, + is_noncausal=not bool(getattr(common, "causal", True)), + ) + actual_end = min(end, common.num_actual_tokens) + actual_start = min(start, actual_end) + # Padding participates in the output exchange, not in cache reads. + qsl = qsl.clamp_max(actual_end - actual_start).to(self._cp_query_start_loc.dtype) + query_start_loc = self._cp_query_start_loc[: qsl.numel()] + query_start_loc.copy_(qsl) + # Device lengths are authoritative after speculative rejection; the + # CPU mirror may still be an upper bound. Remove only the query suffix + # beyond this rank's token interval from each request's device length. + query_ends = common.query_start_loc_cpu[1 : common.num_reqs + 1] + suffix = query_ends - query_ends.clamp(min=actual_start, max=actual_end) + if not bool(getattr(common, "causal", True)): + suffix = torch.zeros_like(suffix) + local_seq_lens = (common.seq_lens[: common.num_reqs] - suffix.to(common.seq_lens.device)).clamp_min(0) + local_seq_lens = torch.where(query_start_loc[1:] > query_start_loc[:-1], local_seq_lens, 0) + local_common = common.replace( + query_start_loc=query_start_loc, + query_start_loc_cpu=qsl, + seq_lens=local_seq_lens, + seq_lens_cpu=seq_lens, + num_actual_tokens=actual_end - actual_start, + num_input_tokens=actual_end - actual_start, + positions=common.positions[actual_start:actual_end], + slot_mapping=common.slot_mapping[actual_start:actual_end], + max_query_len=int((qsl[1:] - qsl[:-1]).max()) if common.num_reqs else 0, + max_seq_len=int(seq_lens.max()) if common.num_reqs else 0, + ) + kwargs["num_query_heads"] = _config_value(self.vllm_config.model_config.hf_text_config, "num_attention_heads") + if global_metadata.cos is not None and global_metadata.sin is not None: + # Q owns a contiguous token slice of the global KV batch. Reuse + # that slice: a second cached RoPE gather would overwrite the + # process-wide buffer still referenced by global KV metadata. + kwargs["rope_views"] = ( + global_metadata.cos[actual_start:actual_end], + global_metadata.sin[actual_start:actual_end], + ) + if global_metadata.ori_sparse_indices is not None: + kwargs["ori_sparse_indices"] = global_metadata.ori_sparse_indices[actual_start:actual_end] + local = super().build(common_prefix_len, local_common, fast_build, **kwargs) + return replace(local, global_metadata=global_metadata, cp_token_range=(start, end, per_rank, padded)) + + +class DeepseekV41CPImpl(DeepseekV41EagerAttentionImpl): + def multistream_preprocess(self, attn, hidden_states, cos, sin, swa_metadata): + """Slice local Q from full inputs and overlap replicated KV preprocessing.""" + global_metadata = self._global_layer_metadata(get_forward_context().attn_metadata) + kv_hidden_states = hidden_states[: global_metadata.swa.num_actual_tokens] + start, _, _, _ = swa_metadata.cp_token_range + hidden_states = hidden_states[start : start + swa_metadata.num_actual_tokens] + kv_cos, kv_sin = global_metadata.rope(attn.rotary_emb.layername, kv_hidden_states.shape[0]) + swa_metadata = global_metadata.swa + main_stream = torch.npu.current_stream() + aux_stream = dsv4_dsa_overlap_stream() + v1_impl = attn.dsa_attn.dsa_attn.impl + wq_a, wkv, wq_b = v1_impl.cv_wq_a, v1_impl.cv_wkv, v1_impl.cv_wq_b + + # Q and KV own different token ranges, even with identical quantizers. + q_quant, q_scale = wq_a.quantize(hidden_states) + q_quant_done = main_stream.record_event() + with npu_stream_switch(aux_stream, enabled=True): + aux_stream.wait_event(q_quant_done) + kv_quant, kv_scale = wkv.quantize(kv_hidden_states) + kv_quant_done = aux_stream.record_event() + q_a = wq_a.matmul(q_quant, q_scale, bias=attn.wq_a.bias) + + # Serialize Cube matmuls while overlapping Q Vector work with KV Cube. + part2_start = main_stream.record_event() + main_stream.wait_event(kv_quant_done) + with npu_stream_switch(aux_stream, enabled=True): + aux_stream.wait_event(part2_start) + kv = wkv.matmul(kv_quant, kv_scale, bias=attn.wkv.bias) + kv_matmul_done = aux_stream.record_event() + qr = attn.q_norm(q_a) + q_b_quant, q_b_scale = wq_b.quantize(qr) + + # KV Vector work uses global RoPE and global cache slots. + part3_start = main_stream.record_event() + main_stream.wait_event(kv_matmul_done) + with npu_stream_switch(aux_stream, enabled=True): + aux_stream.wait_event(part3_start) + kv = attn.kv_norm(kv).view(-1, 1, attn.head_dim) + torch.ops._C_ascend.inplace_partial_rotary_mul( + kv.unsqueeze(1), + kv_cos, + kv_sin, + rotary_mode="interleave", + partial_slice=[attn.nope_head_dim, attn.head_dim], + ) + scatter_cache_sk(attn.dsa_attn.swa_cache_layer.kv_cache[0], swa_metadata.slot_mapping, kv.squeeze(1)) + q = wq_b.matmul(q_b_quant, q_b_scale, bias=attn.wq_b.bias).unflatten(-1, (attn.n_heads, attn.head_dim)) + main_stream.wait_stream(aux_stream) + torch.ops._C_ascend.inplace_partial_rotary_mul( + q.unsqueeze(1), + cos, + sin, + rotary_mode="interleave", + partial_slice=[attn.nope_head_dim, attn.head_dim], + ) + # Both streams have joined before compressor/indexer cache reads. + if self.role.is_kv_source: + self._write_compressed_source( + attn, + kv_hidden_states, + global_metadata.positions[: kv_hidden_states.shape[0]], + kv_cos, + kv_sin, + global_metadata, + ) + return q.to(hidden_states.dtype), qr + + def _global_layer_metadata(self, metadata_by_prefix): + global_by_prefix = {} + # The runner also includes DSpark's native DSA metadata in this map. + # Resolve only the cache planes consumed by this target layer. + for prefix in ( + self.swa_prefix, + self.long_kv_source_prefix, + self.index_k_source_prefix, + self.compressor_state_prefix, + ): + if prefix is None: + continue + metadata = metadata_by_prefix[prefix] + if metadata.global_metadata is None: + raise ValueError(f"V4.1 CP is missing global cache metadata for {prefix}") + global_by_prefix[prefix] = metadata.global_metadata + return self._get_layer_metadata(global_by_prefix) + + def _prepare_inputs_and_caches(self, attn, hidden_states, metadata, metadata_by_prefix): + if not attn.dsa_attn.dsa_attn.impl.multistream_dsv4_dsa_overlap or metadata.swa.num_actual_tokens == 0: + # Empty query ranks still update replicated caches before exchange. + global_metadata = self._global_layer_metadata(metadata_by_prefix) + self._update_caches(attn, hidden_states[: global_metadata.swa.num_actual_tokens], global_metadata) + + def _prepare_queries(self, attn, hidden_states, positions, cos, sin, metadata): + if attn.dsa_attn.dsa_attn.impl.multistream_dsv4_dsa_overlap: + return self.multistream_preprocess(attn, hidden_states, cos, sin, metadata.swa) + # Replicated caches were updated before the TP token slice. + start, _, _, _ = metadata.swa.cp_token_range + hidden_states = hidden_states[start : start + metadata.swa.num_actual_tokens] + return self._project_q(attn, hidden_states, cos, sin) + + def _select_sparse_indices(self, attn, hidden_states, qr, positions, cos, sin, metadata): + if not self.role.has_long_context: + return None + if not self.role.is_index_source: + shared = attn.shared_state + if shared is None: + raise RuntimeError("V4.1 shared attention state is not initialized") + # ``hidden_states`` still owns the full pre-CP token batch here, + # while ``qr`` was projected from this rank's local query slice. + # SparseFlashMla requires cmp_sparse_indices.T to match q.T. + return shared.topk_indices[: qr.shape[0]] + start, _, _, _ = metadata.swa.cp_token_range + hidden_states = hidden_states[start : start + metadata.swa.num_actual_tokens] + return super()._select_sparse_indices(attn, hidden_states, qr, positions, cos, sin, metadata) + + def _project_output(self, attn, output, hidden_states, metadata, *, projected): + _, _, per_rank, _ = metadata.swa.cp_token_range + padded = output + if output.shape[0] != per_rank: + padded = output.new_zeros((per_rank, output.shape[1], output.shape[2])) + padded[: output.shape[0]] = output + exchanged = restore_tp_heads(padded, get_tp_group()) + # The inherited V4 module owns quantized weights and TP projection logic. + local_output = attn.dsa_attn.dsa_attn.impl._forward_o_proj(exchanged) + projected.copy_(local_output[: hidden_states.shape[0]]) + return projected diff --git a/vllm_ascend/attention/dsa_v1.py b/vllm_ascend/attention/dsa_v1.py index c2527bcd1bef..6efe6a7a73b4 100644 --- a/vllm_ascend/attention/dsa_v1.py +++ b/vllm_ascend/attention/dsa_v1.py @@ -424,11 +424,16 @@ def build_dspark_swa_indices( index_width: int | None = None, indices_output: torch.Tensor | None = None, buffer: torch.Tensor | None = None, + *, + use_logical_indices: bool = False, ) -> tuple[torch.Tensor, torch.Tensor]: """Build DSpark non-causal visible slot ids for a paged SWA cache. Each token in a draft block sees the trailing context window plus the whole current draft block. Invalid/padded rows get lens=0 and -1 slots. + ``use_logical_indices`` returns positions within each sequence for + SparseFlashMLA, which applies its own block-table lookup. The default + preserves physical slots for existing DSA callers. When ``buffer`` is given, the per-token slots are copied into its leading rows and the returned tensor is a slice view of ``buffer``. This keeps the @@ -456,13 +461,15 @@ def build_dspark_swa_indices( cols = torch.arange(index_width, device=start_pos.device) col_mask = cols[None, :] < visible_lens[:, None] pos = start_pos[:, None] + cols[None, :] - block_nums = pos // block_size - # Clamp to valid block-table columns so gather never goes OOB on the - # out-of-range columns (their results are discarded by col_mask anyway). - safe_nums = block_nums.clamp(min=0, max=int(block_table.shape[1]) - 1) - block_offsets = pos % block_size - block_ids = torch.gather(block_table, 1, safe_nums) - slot_ids = (block_ids * block_size + block_offsets).to(torch.int32) + if use_logical_indices: + slot_ids = pos.to(torch.int32) + else: + block_nums = pos // block_size + # Clamp out-of-range columns before gathering; col_mask discards them. + safe_nums = block_nums.clamp(min=0, max=int(block_table.shape[1]) - 1) + block_offsets = pos % block_size + block_ids = torch.gather(block_table, 1, safe_nums) + slot_ids = (block_ids * block_size + block_offsets).to(torch.int32) slot_ids = slot_ids.where(col_mask, torch.full_like(slot_ids, -1)) per_token_slots = torch.repeat_interleave(slot_ids, query_lens, dim=0, output_size=num_decode_tokens).unsqueeze(1) @@ -1505,6 +1512,7 @@ def __init__( self.wkv = kwargs["wkv"] self.q_norm = kwargs["q_norm"] self.q_norm_without_weight = kwargs["q_norm_without_weight"] + self.apply_q_norm = True self.kv_norm = kwargs["kv_norm"] # CV wrapper: split wq_a/wkv/wq_b into quantize(Vector) + matmul(Cube) @@ -1821,7 +1829,8 @@ def _mla_prolog_single_stream( qr = self.q_norm(q_a) q = self.wq_b(qr).unflatten(-1, (self.n_local_heads, self.head_dim)) qr_pertoken_scale = None - q = DeviceOperator.apply_dsa_q_rms(q, self.eps, self.q_norm_without_weight) + if self.apply_q_norm: + q = DeviceOperator.apply_dsa_q_rms(q, self.eps, self.q_norm_without_weight) torch.ops._C_ascend.inplace_partial_rotary_mul( q.unsqueeze(1), @@ -1995,7 +2004,8 @@ def _mla_prolog_multistream( e_tail_overlap_done = torch.npu.current_stream().record_event() tail_overlap_output = overlap_result, e_tail_overlap_done - q = DeviceOperator.apply_dsa_q_rms(q, self.eps, self.q_norm_without_weight) + if self.apply_q_norm: + q = DeviceOperator.apply_dsa_q_rms(q, self.eps, self.q_norm_without_weight) torch.ops._C_ascend.inplace_partial_rotary_mul( q.unsqueeze(1), cos, diff --git a/vllm_ascend/attention/dsa_v41.py b/vllm_ascend/attention/dsa_v41.py new file mode 100644 index 000000000000..94f8de8d6f1c --- /dev/null +++ b/vllm_ascend/attention/dsa_v41.py @@ -0,0 +1,1205 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""DeepSeek V4.1 DSA metadata and fused attention execution. + +The model file owns the network topology and projection modules. This module +owns the attention execution boundary: it gathers every cache plane's metadata +before running the compressor, indexer and sparse-attention operators without +moving cache or scheduler knowledge back into the model. +""" + +from dataclasses import dataclass +from typing import Any + +import torch +import torch.nn.functional as F +from torch import nn +from vllm.compilation.breakable_cudagraph import eager_break_during_capture +from vllm.config import VllmConfig +from vllm.forward_context import get_forward_context +from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase +from vllm.utils.torch_utils import direct_register_custom_op +from vllm.v1.attention.backend import ( + AttentionBackend, + AttentionCGSupport, + AttentionMetadata, + AttentionMetadataBuilder, +) + +from vllm_ascend.attention.dsa_v1 import build_dspark_swa_indices, dsv4_dsa_overlap_stream +from vllm_ascend.core.deepseek_v41 import ( + DeepseekV41CompressorStateSpec, + DeepseekV41DraftSWASpec, + DeepseekV41FullSpec, + DeepseekV41IndexerSpec, + DeepseekV41SWASpec, +) +from vllm_ascend.core.kv_cache_interface import get_kv_cache_compression_ratio +from vllm_ascend.ops.rope_dsv4 import ( + get_cos_and_sin_dsa, + get_full_cos_and_sin_dsa_for_layer, +) +from vllm_ascend.utils import npu_stream_switch +from vllm_ascend.worker.device_metadata import ( + DeviceMetadataStage, + DeviceMetadataTask, + wait_for_device_metadata, +) + +V41_METADATA_BUFFER_SIZE = 1024 + + +@eager_break_during_capture +def dsa_v41_forward( + hidden_states: torch.Tensor, + output: torch.Tensor, + layer_name: str, +) -> None: + """Execute V4.1 attention behind an explicit graph side-effect boundary.""" + forward_context = get_forward_context() + attn = forward_context.no_compile_layers[layer_name] + attn.v41_impl.forward(attn, None, hidden_states, output) + + +def dsa_v41_forward_fake( + hidden_states: torch.Tensor, + output: torch.Tensor, + layer_name: str, +) -> None: + return None + + +direct_register_custom_op( + op_name="dsa_v41_forward", + op_func=dsa_v41_forward, + mutates_args=["output"], + fake_impl=dsa_v41_forward_fake, + dispatch_key="PrivateUse1", +) + + +def _config_value(config: Any, name: str, default: Any = None) -> Any: + """Read one field from either an HF config object or a raw config dict.""" + if isinstance(config, dict): + return config.get(name, default) + return getattr(config, name, default) + + +@dataclass +class DeepseekV41Metadata(AttentionMetadata): + """Scheduler and cache-plane contract for one V4.1 cache resource. + + ``seq_lens``/``query_start_loc`` always stay in original-token + coordinates, matching the common vLLM metadata. The ``cache_*`` fields + describe the rows visible to the concrete cache plane. Keeping both + coordinate systems here lets future fused kernels replace the eager path + without rebuilding scheduling metadata in the model. + """ + + block_table: torch.Tensor + query_start_loc: torch.Tensor + seq_lens: torch.Tensor + slot_mapping: torch.Tensor + compress_ratio: int + storage_block_size: int + is_compressor_state: bool + cache_kind: str = "unknown" + positions: torch.Tensor | None = None + cos: Any = None + sin: Any = None + num_actual_tokens: int = 0 + num_input_tokens: int = 0 + num_reqs: int = 0 + num_actual_reqs: int = 0 + num_decodes: int = 0 + num_decode_tokens: int = 0 + num_prefills: int = 0 + num_prefill_tokens: int = 0 + logical_block_size: int = 0 + query_start_loc_cpu: torch.Tensor | None = None + block_table_cpu: torch.Tensor | None = None + seq_lens_cpu: torch.Tensor | None = None + cache_seq_lens: torch.Tensor | None = None + max_query_len: int = 0 + max_seq_len: int = 0 + max_cache_seq_len: int = 0 + attn_state: Any = None + is_prefilling: torch.Tensor | None = None + causal: bool | torch.Tensor = True + ori_sparse_indices: torch.Tensor | None = None + ori_topk_length: torch.Tensor | None = None + ori_mask_mode: int = 4 + ori_win_left: int = 0 + ori_win_right: int = 0 + smla_metadata: torch.Tensor | None = None + qli_metadata: torch.Tensor | None = None + cmp_residual: torch.Tensor | None = None + c2_ring_metadata: torch.Tensor | None = None + c2_complete_mask: torch.Tensor | None = None + c2_source_positions: torch.Tensor | None = None + c2_source_cos: torch.Tensor | None = None + c2_source_sin: torch.Tensor | None = None + c2_metadata_group_id: int | None = None + global_metadata: "DeepseekV41Metadata | None" = None + cp_token_range: tuple[int, int, int, int] | None = None + + +@dataclass(frozen=True) +class DeepseekV41CompressorMetadata: + """V4-shaped cache/state bundle consumed by the compressor stage.""" + + cache: DeepseekV41Metadata + state: DeepseekV41Metadata | None = None + + +@dataclass(frozen=True) +class DeepseekV41IndexerMetadata: + """V4-shaped source cache bundle consumed by the indexer stage.""" + + cache: DeepseekV41Metadata + + +@dataclass(frozen=True) +class DeepseekV41LayerMetadata: + """All metadata consumed by one V4.1 attention layer invocation.""" + + attention: DeepseekV41Metadata | None + swa: DeepseekV41Metadata + compressor: DeepseekV41CompressorMetadata | None + indexer: DeepseekV41IndexerMetadata | None + + @property + def positions(self) -> torch.Tensor: + if self.swa.positions is None: + raise RuntimeError("V4.1 SWA metadata does not contain input positions") + return self.swa.positions + + def rope(self, layer_name: str, num_tokens: int): + if self.swa.cos is None or self.swa.sin is None: + raise RuntimeError("V4.1 SWA metadata does not contain RoPE tensors") + return self.swa.cos[layer_name][:num_tokens], self.swa.sin[layer_name][:num_tokens] + + +def compressed_slot_mapping(slot_mapping: torch.Tensor, ratio: int) -> torch.Tensor: + """Convert original-token physical slots to completed compressed slots. + + Logical block sizes must be divisible by ratio. Negative/padded slots and + incomplete compression groups never produce a write. + """ + if ratio not in (1, 2): + raise ValueError("V4.1 only supports ratio 1 or 2") + valid = (slot_mapping >= 0) & ((slot_mapping + 1) % ratio == 0) + return torch.where(valid, slot_mapping // ratio, -1) + + +def _request_counts(common: Any, num_reqs: int): + """Return V4-shaped request counters without synchronizing the NPU.""" + is_prefilling = getattr(common, "is_prefilling", None) + query_start_loc_cpu = getattr(common, "query_start_loc_cpu", None) + if ( + is_prefilling is None + or query_start_loc_cpu is None + or getattr(is_prefilling, "device", None) is None + or is_prefilling.device.type != "cpu" + ): + return 0, 0, 0, 0 + flags = is_prefilling[:num_reqs].bool() + query_lens_cpu = query_start_loc_cpu[1 : num_reqs + 1] - query_start_loc_cpu[:num_reqs] + num_prefills = int(flags.sum().item()) + num_decodes = num_reqs - num_prefills + num_prefill_tokens = int(query_lens_cpu[flags].sum().item()) + num_decode_tokens = int(query_lens_cpu[~flags].sum().item()) + return num_decodes, num_decode_tokens, num_prefills, num_prefill_tokens + + +def scatter_cache_sk( + cache: torch.Tensor, + slot_mapping: torch.Tensor, + values: torch.Tensor, +) -> None: + """Store rows using builder-prepared coordinates and V4's Ascend op. + + V4.1 cache planes can be views into a larger layer-outermost slot, so the + physical page stride is not necessarily the contiguous stride implied by + the plane shape. ``npu_scatter_nd_update_sk`` preserves that stride and + treats the builder's ``[-1, -1]`` coordinates as skipped rows, matching V4. + """ + if slot_mapping.ndim != 2 or slot_mapping.shape[-1] != 2: + raise ValueError( + f"V4.1 fused cache store requires builder-prepared [T, 2] slot_mapping, got {tuple(slot_mapping.shape)}" + ) + cache = cache.squeeze(-2) + indices = slot_mapping[: values.shape[0]] + updates = values.to(cache.dtype).contiguous() + torch.ops._C_ascend.npu_scatter_nd_update_sk(cache, indices, updates) + + +def pad_sparse_indices(indices: torch.Tensor, topk: int) -> torch.Tensor: + """Convert V4.1's compact [T, K] selection into SMLA [T, 1, topk].""" + if indices.ndim != 2: + raise ValueError(f"V4.1 sparse indices must be rank 2, got {indices.shape}") + if indices.shape[-1] > topk: + raise ValueError(f"V4.1 sparse indices width {indices.shape[-1]} exceeds operator topk {topk}") + if indices.shape[-1] < topk: + indices = F.pad(indices, (0, topk - indices.shape[-1]), value=-1) + return indices.unsqueeze(1).contiguous().int() + + +class DeepseekV41EagerAttentionImpl: + """V4-shaped execution boundary backed by fused Ascend operators. + + Projection, compressor and indexer modules remain registered by the model, + while this object resolves the complete per-layer metadata bundle and owns + their invocation order. That is the same separation used by ``dsa_v1``: + model construction is independent from cache-aware attention execution. + """ + + def __init__(self, prefix, role, topology, long_kv_source_prefix, index_k_source_prefix): + self.prefix = prefix + self.layer_name = f"{prefix}.attn" + self.role = role + self.topology = topology + self.swa_prefix = f"{prefix}.swa_cache" + self.long_kv_source_prefix = long_kv_source_prefix + self.index_k_source_prefix = index_k_source_prefix + self.compressor_state_prefix = ( + f"{prefix}.compressor.state_cache" if role.is_kv_source and role.compress_ratio == 2 else None + ) + + def _get_layer_metadata(self, metadata) -> DeepseekV41LayerMetadata: + try: + swa = metadata[self.swa_prefix] + long_kv = metadata[self.long_kv_source_prefix] if self.long_kv_source_prefix is not None else None + index_k = metadata[self.index_k_source_prefix] if self.index_k_source_prefix is not None else None + compressor_state = ( + metadata[self.compressor_state_prefix] if self.compressor_state_prefix is not None else None + ) + except KeyError as exc: + raise RuntimeError(f"Missing V4.1 cache metadata for {exc.args[0]}") from exc + return DeepseekV41LayerMetadata( + attention=long_kv, + swa=swa, + compressor=( + DeepseekV41CompressorMetadata(long_kv, compressor_state) + if self.role.is_kv_source and long_kv is not None + else None + ), + indexer=(DeepseekV41IndexerMetadata(index_k) if index_k is not None else None), + ) + + @staticmethod + def _project_q(attn, hidden_states, cos, sin): + qr = attn.q_norm(attn.wq_a(hidden_states)) + q = attn.wq_b(qr).unflatten(-1, (-1, attn.head_dim)) + torch.ops._C_ascend.inplace_partial_rotary_mul( + q.unsqueeze(1), + cos, + sin, + rotary_mode="interleave", + partial_slice=[attn.nope_head_dim, attn.head_dim], + ) + return q.to(hidden_states.dtype), qr + + @staticmethod + def _project_kv(attn, hidden_states, cos, sin): + kv = attn.kv_norm(attn.wkv(hidden_states)).view(-1, 1, attn.head_dim) + torch.ops._C_ascend.inplace_partial_rotary_mul( + kv.unsqueeze(1), + cos, + sin, + rotary_mode="interleave", + partial_slice=[attn.nope_head_dim, attn.head_dim], + ) + return kv.squeeze(1) + + @classmethod + def _project_q_kv(cls, attn, hidden_states, cos, sin): + q, qr = cls._project_q(attn, hidden_states, cos, sin) + return q, qr, cls._project_kv(attn, hidden_states, cos, sin) + + def _update_caches(self, attn, hidden_states, metadata): + if hidden_states.shape[0] == 0: + return + positions = metadata.positions[: hidden_states.shape[0]] + cos, sin = metadata.rope(attn.rotary_emb.layername, hidden_states.shape[0]) + kv = self._project_kv(attn, hidden_states, cos, sin) + scatter_cache_sk(attn.dsa_attn.swa_cache_layer.kv_cache[0], metadata.swa.slot_mapping, kv) + if self.role.is_kv_source: + self._write_compressed_source(attn, hidden_states, positions, cos, sin, metadata) + + def _prepare_inputs_and_caches(self, attn, hidden_states, metadata, metadata_by_prefix): + """Prepare caches before query work; ordinary preprocessing writes them.""" + pass + + def _prepare_queries(self, attn, hidden_states, positions, cos, sin, metadata): + hidden_states = hidden_states[: metadata.swa.num_actual_tokens] + v1_impl = attn.dsa_attn.dsa_attn.impl + preprocess = self.multistream_preprocess if v1_impl.multistream_dsv4_dsa_overlap else self.preprocess + q, qr = preprocess(attn, hidden_states, cos, sin, metadata.swa) + if self.role.is_kv_source: + self._write_compressed_source(attn, hidden_states, positions, cos, sin, metadata) + return q, qr + + def _project_output(self, attn, output, hidden_states, metadata, *, projected): + padded = output + if output.shape[0] != hidden_states.shape[0]: + padded = output.new_zeros((hidden_states.shape[0], output.shape[1], output.shape[2])) + padded[: output.shape[0]] = output + attn.dsa_attn.dsa_attn.impl._forward_o_proj(padded, projected) + return projected + + def preprocess(self, attn, hidden_states, cos, sin, swa_metadata): + """Project Q/KV and populate this layer's SWA cache on the current stream.""" + q, qr, kv = self._project_q_kv(attn, hidden_states, cos, sin) + scatter_cache_sk( + attn.dsa_attn.swa_cache_layer.kv_cache[0], + swa_metadata.slot_mapping, + kv, + ) + return q, qr + + def multistream_preprocess(self, attn, hidden_states, cos, sin, swa_metadata): + """Overlap Q Vector work with KV Cube work, then reverse their roles. + + Reuse V1's stream and projection wrappers. V4.1 keeps floating-point + qr for its indexer and has no post-Wq_b Q RMSNorm. Stage events serialize + the Cube matmuls; the final join makes SWA writes visible to attention. + """ + main_stream = torch.npu.current_stream() + aux_stream = dsv4_dsa_overlap_stream() + v1_impl = attn.dsa_attn.dsa_attn.impl + wq_a, wkv, wq_b = v1_impl.cv_wq_a, v1_impl.cv_wkv, v1_impl.cv_wq_b + share_quant = ( + type(wq_a._quant_method) is type(wkv._quant_method) and wq_a._has_communication == wkv._has_communication + ) + + # Part 1: Q_a matmul (Cube) overlaps independent KV quantization (Vector). + q_quant, q_scale = wq_a.quantize(hidden_states) + kv_quant_done = None + if share_quant: + kv_quant, kv_scale = q_quant, q_scale + else: + q_quant_done = main_stream.record_event() + with npu_stream_switch(aux_stream, enabled=True): + aux_stream.wait_event(q_quant_done) + kv_quant, kv_scale = wkv.quantize(hidden_states) + kv_quant_done = aux_stream.record_event() + q_a = wq_a.matmul(q_quant, q_scale, bias=attn.wq_a.bias) + + # Part 2: Q normalization/quantization (Vector) overlaps KV matmul (Cube). + part2_start = main_stream.record_event() + if kv_quant_done is not None: + main_stream.wait_event(kv_quant_done) + with npu_stream_switch(aux_stream, enabled=True): + aux_stream.wait_event(part2_start) + kv = wkv.matmul(kv_quant, kv_scale, bias=attn.wkv.bias) + kv_matmul_done = aux_stream.record_event() + qr = attn.q_norm(q_a) + q_b_quant, q_b_scale = wq_b.quantize(qr) + + # Part 3: Q_b matmul (Cube) overlaps KV norm, RoPE and cache store (Vector). + part3_start = main_stream.record_event() + main_stream.wait_event(kv_matmul_done) + with npu_stream_switch(aux_stream, enabled=True): + aux_stream.wait_event(part3_start) + kv = attn.kv_norm(kv).view(-1, 1, attn.head_dim) + torch.ops._C_ascend.inplace_partial_rotary_mul( + kv.unsqueeze(1), + cos, + sin, + rotary_mode="interleave", + partial_slice=[attn.nope_head_dim, attn.head_dim], + ) + scatter_cache_sk( + attn.dsa_attn.swa_cache_layer.kv_cache[0], + swa_metadata.slot_mapping, + kv.squeeze(1), + ) + q = wq_b.matmul(q_b_quant, q_b_scale, bias=attn.wq_b.bias).unflatten(-1, (attn.n_local_heads, attn.head_dim)) + main_stream.wait_stream(aux_stream) + torch.ops._C_ascend.inplace_partial_rotary_mul( + q.unsqueeze(1), + cos, + sin, + rotary_mode="interleave", + partial_slice=[attn.nope_head_dim, attn.head_dim], + ) + return q.to(hidden_states.dtype), qr + + def _write_compressed_source( + self, + attn, + hidden_states, + positions, + cos, + sin, + metadata, + ): + compressor = attn.compressor + if compressor is None or metadata.compressor is None or metadata.indexer is None: + raise RuntimeError("V4.1 KV source is missing compressor or source metadata") + compressor_metadata = metadata.compressor + indexer_metadata = metadata.indexer + ratio = self.role.compress_ratio + if ratio == 1: + latent = compressor(hidden_states) + # C1 source positions are the current token positions. Reuse the + # query RoPE selected by the SWA metadata builder instead of + # indexing the global table a second time. + source_cos = cos + source_sin = sin + index_slots = indexer_metadata.cache.slot_mapping[: positions.shape[0]] + long_slots = compressor_metadata.cache.slot_mapping[: positions.shape[0]] + else: + if compressor_metadata.state is None: + raise RuntimeError("V4.1 ratio-2 source is missing compressor-state metadata") + state_metadata = compressor_metadata.state + if state_metadata.c2_ring_metadata is None or state_metadata.c2_metadata_group_id is None: + raise RuntimeError("V4.1 ring compressor metadata is missing") + wait_for_device_metadata(DeviceMetadataStage.COMPRESSOR, state_metadata.c2_metadata_group_id) + hidden_states_fp32 = hidden_states.float() + kv = compressor.wkv(hidden_states_fp32) + score = compressor.wgate(hidden_states_fp32) + latent = compressor.pool_projected(kv, score, state_metadata) + source_cos = state_metadata.c2_source_cos + source_sin = state_metadata.c2_source_sin + if source_cos is None or source_sin is None: + fallback_cos, fallback_sin = get_cos_and_sin_dsa(state_metadata.c2_source_positions) + source_cos = fallback_cos[attn.rotary_emb.layername] + source_sin = fallback_sin[attn.rotary_emb.layername] + source_cos = source_cos[: positions.shape[0]] + source_sin = source_sin[: positions.shape[0]] + index_slots = indexer_metadata.cache.slot_mapping[: positions.shape[0]] + long_slots = compressor_metadata.cache.slot_mapping[: positions.shape[0]] + + if attn.indexer is None: + raise RuntimeError("V4.1 KV source is missing its indexer") + attn.indexer.update_keys( + latent, + index_slots, + source_cos, + source_sin, + ) + latent = latent.view(-1, 1, attn.head_dim) + torch.ops._C_ascend.inplace_partial_rotary_mul( + latent.unsqueeze(1), + source_cos, + source_sin, + rotary_mode="interleave", + partial_slice=[attn.nope_head_dim, attn.head_dim], + ) + scatter_cache_sk( + attn.long_kv_cache.kv_cache[0], + long_slots, + latent.squeeze(1), + ) + + def _select_sparse_indices(self, attn, hidden_states, qr, positions, cos, sin, metadata): + hidden_states = hidden_states[: metadata.swa.num_actual_tokens] + if not self.role.has_long_context: + return None + shared = attn.shared_state + if shared is None: + raise RuntimeError("V4.1 shared attention state is not initialized") + if not self.role.is_index_source: + return shared.topk_indices[: hidden_states.shape[0]] + if attn.indexer is None or metadata.indexer is None: + raise RuntimeError("V4.1 index source is missing indexer metadata") + + context = get_forward_context().no_compile_layers + source_layer = context[self.index_k_source_prefix] + selected, candidates = attn.indexer.select( + hidden_states, + qr, + positions, + cos, + sin, + source_layer.kv_cache[0], + metadata.indexer.cache, + is_candidate_source=self.role.is_candidate_source, + uses_candidate_filter=self.role.uses_candidate_filter, + candidate_topk_blocks=self.topology.candidate_topk_blocks, + candidate_block_size=self.topology.candidate_block_size, + candidates=shared.candidates[: hidden_states.shape[0]], + ) + shared.topk_indices[: selected.shape[0]].copy_(selected) + if self.role.is_candidate_source: + shared.candidates[: candidates.shape[0]].copy_(candidates) + return shared.topk_indices[: selected.shape[0]] + + def _attention(self, attn, q, metadata, compressed_indices): + source_cache = None + if self.role.has_long_context: + source_cache = get_forward_context().no_compile_layers[self.long_kv_source_prefix].kv_cache[0] + return self._native_attention( + attn, + q, + metadata, + source_cache=source_cache, + compressed_indices=compressed_indices, + ) + + def _native_attention( + self, + attn, + q, + metadata, + *, + source_cache, + compressed_indices, + ): + """Run SparseFlashMla with the same PA metadata for both operator stages.""" + if attn.head_dim != 512: + raise ValueError(f"SparseFlashMla requires head_dim 512, got {attn.head_dim}") + if attn.window_size != 128: + raise ValueError(f"A2/A3 SparseFlashMla requires sliding_window 128, got {attn.window_size}") + num_heads = q.shape[1] + if not 1 <= num_heads <= 128 or num_heads & (num_heads - 1): + raise ValueError( + "A2/A3 SparseFlashMla requires the local query-head count to be " + f"a power of two in [1, 128], got {num_heads}" + ) + has_compressed = self.role.compress_ratio in (1, 2) + ratio = self.role.compress_ratio if has_compressed else 0 + num_reqs = metadata.swa.num_reqs + query_start_loc = metadata.swa.query_start_loc[: num_reqs + 1] + seq_lens = metadata.swa.seq_lens[:num_reqs] + ori_block_table = metadata.swa.block_table[:num_reqs] + cmp_block_table = None + cmp_seq_lens = None + cmp_residual = None + cmp_indices = None + cmp_topk = 0 + if has_compressed: + if source_cache is None or metadata.attention is None or compressed_indices is None: + raise RuntimeError("V4.1 compressed attention is missing KV or TopK metadata") + cmp_block_table = metadata.attention.block_table[:num_reqs] + cmp_seq_lens = metadata.attention.cache_seq_lens[:num_reqs] + cmp_residual = metadata.attention.cmp_residual + cmp_topk = self.topology.index_topk + if cmp_topk not in (512, 1024): + raise ValueError(f"SparseFlashMla only supports TopK 512 or 1024, got {cmp_topk}") + cmp_indices = pad_sparse_indices(compressed_indices, cmp_topk) + + operator_metadata = metadata.attention if has_compressed else metadata.swa + op_metadata = operator_metadata.smla_metadata + if op_metadata is None: + raise RuntimeError(f"V4.1 ratio-{ratio} SMLA metadata was not built") + wait_for_device_metadata( + DeviceMetadataStage.ATTENTION, + id(op_metadata), + ) + output, _ = torch.ops._C_ascend.npu_sparse_flash_mla( + q, + ori_kv=attn.dsa_attn.swa_cache_layer.kv_cache[0], + cmp_kv=source_cache, + ori_sparse_indices=metadata.swa.ori_sparse_indices, + ori_topk_length=metadata.swa.ori_topk_length, + cmp_sparse_indices=cmp_indices, + ori_block_table=ori_block_table, + cmp_block_table=cmp_block_table, + cu_seqlens_q=query_start_loc, + seqused_ori_kv=seq_lens, + seqused_cmp_kv=cmp_seq_lens, + cmp_residual_kv=cmp_residual, + sinks=attn.attn_sink, + metadata=op_metadata, + softmax_scale=attn.softmax_scale, + cmp_ratio=ratio, + ori_mask_mode=metadata.swa.ori_mask_mode, + cmp_mask_mode=3 if has_compressed else 0, + ori_win_left=metadata.swa.ori_win_left, + ori_win_right=metadata.swa.ori_win_right, + layout_q="TND", + layout_kv="PA_BBND", + topk_value_mode=1, + return_softmax_lse=False, + ) + return output + + @staticmethod + def update_graph_params(*args, **kwargs): + """V4.1 owns stable metadata buffers; no backend pointer patch is needed.""" + return None + + def forward(self, attn, positions, hidden_states, output: torch.Tensor | None = None): + # The custom-op caller provides a graph-stable output buffer. Write + # O-projection results into it directly instead of materializing a + # second full hidden-state tensor and copying it at the graph boundary. + if output is None: + output = torch.empty_like(hidden_states) + forward_context = get_forward_context() + if forward_context.attn_metadata is None: + output.zero_() + return output + metadata = self._get_layer_metadata(forward_context.attn_metadata) + self._prepare_inputs_and_caches(attn, hidden_states, metadata, forward_context.attn_metadata) + num_tokens = metadata.swa.num_actual_tokens + if num_tokens: + positions = metadata.positions[:num_tokens] + cos, sin = metadata.rope(attn.rotary_emb.layername, num_tokens) + q, qr = self._prepare_queries(attn, hidden_states, positions, cos, sin, metadata) + compressed_indices = self._select_sparse_indices(attn, hidden_states, qr, positions, cos, sin, metadata) + attention_output = self._attention(attn, q, metadata, compressed_indices) + torch.ops._C_ascend.inplace_partial_rotary_mul( + attention_output.unsqueeze(1), + cos, + -sin, + rotary_mode="interleave", + partial_slice=[attn.nope_head_dim, attn.head_dim], + ) + else: + heads = attn.n_heads if getattr(attn, "enable_dsa_cp", False) else attn.n_local_heads + attention_output = hidden_states.new_empty((0, heads, attn.head_dim)) + self._project_output(attn, attention_output, hidden_states, metadata, projected=output) + return output + + +class DeepseekV41MetadataBuilder(AttentionMetadataBuilder[DeepseekV41Metadata]): + def __init__( + self, + kv_cache_spec, + layer_names, + vllm_config, + device, + *, + build_query_metadata=True, + build_compressor_metadata=True, + ): + super().__init__(kv_cache_spec, layer_names, vllm_config, device) + max_tokens = getattr(vllm_config.scheduler_config, "max_num_batched_tokens", 4096) + max_reqs = getattr(vllm_config.scheduler_config, "max_num_seqs", 256) + self._supports_device_ops = getattr(device, "type", "cpu") != "cpu" + # CP uses global cache-write controls and local query controls. These + # roles are fixed before allocation and graph capture. + self._build_query_metadata = build_query_metadata + self._build_compressor_metadata = build_compressor_metadata + query_metadata_size = V41_METADATA_BUFFER_SIZE if build_query_metadata else 0 + compressor_tokens = max_tokens if build_compressor_metadata else 0 + compressor_reqs = max_reqs if build_compressor_metadata else 0 + self._slot_mapping = torch.full((max_tokens,), -1, dtype=torch.int64, device=device) + self._slot_mapping_2d = torch.full((max_tokens, 2), -1, dtype=torch.int32, device=device) + self._seq_lens = torch.zeros(max_reqs, dtype=torch.int32, device=device) + self._cache_seq_lens = torch.zeros(max_reqs, dtype=torch.int32, device=device) + self._cmp_residual = torch.zeros(max_reqs, dtype=torch.int32, device=device) + self._smla_metadata = torch.zeros(query_metadata_size, dtype=torch.int32, device=device) + self._qli_metadata = torch.zeros(query_metadata_size, dtype=torch.int32, device=device) + self._c2_ring_metadata = torch.zeros(5 * compressor_reqs, dtype=torch.int32, device=device) + self._c2_complete_mask = torch.zeros(compressor_tokens, dtype=torch.bool, device=device) + self._c2_source_positions = torch.zeros(compressor_tokens, dtype=torch.int64, device=device) + text_config = vllm_config.model_config.hf_text_config + rope_dim = int( + _config_value( + text_config, + "qk_rope_head_dim", + _config_value(text_config, "head_dim"), + ) + ) + c2_rope_rows = ( + compressor_tokens + if self._supports_device_ops and isinstance(kv_cache_spec, DeepseekV41CompressorStateSpec) + else 0 + ) + self._c2_source_cos = torch.ones( + (c2_rope_rows, 1, 1, rope_dim), + dtype=torch.float32, + device=device, + ) + self._c2_source_sin = torch.zeros_like(self._c2_source_cos) + self._c2_rope_layer_names = tuple( + name.removesuffix(".compressor.state_cache") + ".attn" + for name in layer_names + if name.endswith(".compressor.state_cache") + ) + self._c2_full_source_rope: tuple[torch.Tensor, torch.Tensor] | None = None + self._device_metadata_enabled = False + self._device_metadata_tasks: tuple[DeviceMetadataTask, ...] = () + + @classmethod + def get_cudagraph_support( + cls, + vllm_config: VllmConfig, + kv_cache_spec, + ) -> AttentionCGSupport: + return AttentionCGSupport.UNIFORM_BATCH + + def build_for_cudagraph_capture( + self, + common_attn_metadata, + **kwargs, + ) -> DeepseekV41Metadata: + return self.build( + common_prefix_len=0, + common_attn_metadata=common_attn_metadata, + **kwargs, + ) + + def build_for_drafting(self, common_attn_metadata, draft_index, **kwargs): + if not isinstance(self.kv_cache_spec, DeepseekV41DraftSWASpec): + raise TypeError("V4.1 drafting requires a draft SWA cache") + # DSpark issues one eager block per step. Group-local tables and slots + # remain independent; the builder owns the operator metadata buffers. + return self.build(0, common_attn_metadata) + + def enable_device_metadata(self) -> None: + self._device_metadata_enabled = True + if self._build_compressor_metadata and isinstance(self.kv_cache_spec, DeepseekV41CompressorStateSpec): + if not self._c2_rope_layer_names: + raise RuntimeError("V4.1 compressor-state builder has no source RoPE layer") + source_rope = get_full_cos_and_sin_dsa_for_layer(self._c2_rope_layer_names[0]) + for rope_layer_name in self._c2_rope_layer_names[1:]: + other_rope = get_full_cos_and_sin_dsa_for_layer(rope_layer_name) + if any(other.data_ptr() != source.data_ptr() for other, source in zip(other_rope, source_rope)): + raise RuntimeError("V4.1 ratio-2 source layers must share one RoPE table") + self._c2_full_source_rope = source_rope + + def take_device_metadata_tasks(self) -> tuple[DeviceMetadataTask, ...]: + tasks = self._device_metadata_tasks + self._device_metadata_tasks = () + return tasks + + def _publish_task( + self, + shared: dict[str, Any], + key: str, + buffer: torch.Tensor, + stage: DeviceMetadataStage, + run, + ) -> torch.Tensor: + existing = shared.get(key) + if existing is not None: + return existing + shared[key] = buffer + if self._device_metadata_enabled: + self._device_metadata_tasks = ( + *self._device_metadata_tasks, + DeviceMetadataTask(stage, run, id(buffer)), + ) + else: + run() + return buffer + + def _build_batch_metadata(self, common, num_reqs, num_actual_reqs, num_input_tokens): + self._seq_lens[:num_reqs].copy_(common.seq_lens[:num_reqs]) + if num_actual_reqs < num_reqs: + self._seq_lens[num_actual_reqs:num_reqs].zero_() + seq_lens_cpu = getattr(common, "seq_lens_cpu", None) + if seq_lens_cpu is None: + seq_lens_cpu = getattr(common, "_seq_lens_cpu", None) + max_seq_len = int(getattr(common, "max_seq_len", 0)) + if seq_lens_cpu is not None: + max_seq_len = int(seq_lens_cpu[:num_actual_reqs].max().item()) if num_actual_reqs else 0 + num_decodes, num_decode_tokens, num_prefills, num_prefill_tokens = _request_counts(common, num_reqs) + positions = common.positions + if positions is not None: + positions = positions[:num_input_tokens].long() + return dict( + query_start_loc=common.query_start_loc[: num_reqs + 1], + query_start_loc_cpu=getattr(common, "query_start_loc_cpu", None), + seq_lens=self._seq_lens[:num_reqs], + seq_lens_cpu=seq_lens_cpu, + positions=positions, + max_cache_seq_len=max_seq_len, + num_decodes=num_decodes, + num_decode_tokens=num_decode_tokens, + num_prefills=num_prefills, + num_prefill_tokens=num_prefill_tokens, + ) + + def build( + self, + common_prefix_len, + common_attn_metadata, + fast_build=False, + **kwargs, + ): + if common_prefix_len: + raise NotImplementedError("V4.1 prefix caching is not implemented") + self._device_metadata_tasks = () + spec = self.kv_cache_spec + common = common_attn_metadata + is_compressor_state = isinstance(spec, DeepseekV41CompressorStateSpec) + ratio = spec.compress_ratio if is_compressor_state else get_kv_cache_compression_ratio(spec) + if isinstance(spec, (DeepseekV41SWASpec, DeepseekV41DraftSWASpec)): + cache_kind = "swa" + elif isinstance(spec, DeepseekV41FullSpec): + cache_kind = "long_kv" + elif isinstance(spec, DeepseekV41IndexerSpec): + cache_kind = "index_k" + elif is_compressor_state: + cache_kind = "compressor_state" + else: + raise TypeError(f"Unsupported V4.1 cache spec: {type(spec).__name__}") + + num_reqs = int(getattr(common, "num_reqs", common.seq_lens.shape[0])) + num_actual_reqs = int(kwargs.get("num_actual_reqs", num_reqs)) + num_actual_reqs = min(num_actual_reqs, num_reqs) + num_input_tokens = int(getattr(common, "num_input_tokens", common.slot_mapping.shape[0])) + num_actual_tokens = int(getattr(common, "num_actual_tokens", num_input_tokens)) + shared = kwargs.get("common_v41_metadata") + if shared is None: + shared = {} + batch_shared = kwargs.get("common_v41_batch_metadata") + if batch_shared is None: + batch_shared = shared + + # The runner resets both dictionaries on each build. Batch values do + # not depend on physical block IDs; slot mappings remain group-local. + batch_metadata = batch_shared.get("batch") + if batch_metadata is None: + batch_metadata = self._build_batch_metadata(common, num_reqs, num_actual_reqs, num_input_tokens) + batch_shared["batch"] = batch_metadata + coordinates = dict(batch_metadata) + seq_lens = coordinates["seq_lens"] + positions = coordinates["positions"] + + # SWA uses original-token coordinates; circular state has no token slots. + # Long KV and index K are addressed in completed compression groups. + compressed = cache_kind in {"long_kv", "index_k"} + if is_compressor_state: + # State writes use ring ownership metadata; this buffer stays PAD. + slots = self._slot_mapping[:num_input_tokens] + else: + # Scope ``shared`` to one framework KV cache group in the model + # runner. Long KV and Indexer builders with the same physical + # layout then share one persistent [T, 2] mapping, while every SWA + # group owns a distinct mapping buffer. + slot_key = f"slot:c{ratio}:b{spec.storage_block_size}" + prepared_slots = shared.get(slot_key) + if prepared_slots is None: + active_slots = common.slot_mapping[:num_input_tokens] + if compressed and ratio != 1: + active_slots = compressed_slot_mapping(active_slots, ratio) + valid = active_slots >= 0 + if compressed and ratio == 2: + # Prepare the C2 store mask once per cache group, before + # forward. Match the ring compressor's completion policy. + if kwargs.get("skip_ring_state_update", False): + valid.zero_() + else: + valid_end = common.query_start_loc[num_actual_reqs].clamp_max(num_actual_tokens) + valid &= torch.arange(num_input_tokens, device=active_slots.device) < valid_end + if positions is not None: + valid &= positions.remainder(2) == 1 + physical = active_slots.clamp_min(0) + self._slot_mapping_2d[:num_input_tokens, 0].copy_( + torch.where( + valid, + torch.div( + physical, + spec.storage_block_size, + rounding_mode="floor", + ), + -1, + ) + ) + self._slot_mapping_2d[:num_input_tokens, 1].copy_( + torch.where( + valid, + physical.remainder(spec.storage_block_size), + -1, + ) + ) + prepared_slots = self._slot_mapping_2d[:num_input_tokens] + shared[slot_key] = prepared_slots + slots = prepared_slots + plane_ratio = ratio if compressed else 1 + coordinates["cache_seq_lens"] = seq_lens + cmp_residual_buffer = None + if compressed and ratio == 2: + compressed_lengths = batch_shared.get("lengths:c2") + if compressed_lengths is None: + torch.div(seq_lens, ratio, rounding_mode="floor", out=self._cache_seq_lens[:num_reqs]) + torch.remainder(seq_lens, ratio, out=self._cmp_residual[:num_reqs]) + compressed_lengths = (self._cache_seq_lens[:num_reqs], self._cmp_residual[:num_reqs]) + batch_shared["lengths:c2"] = compressed_lengths + coordinates["cache_seq_lens"], cmp_residual_buffer = compressed_lengths + coordinates["max_cache_seq_len"] //= plane_ratio + cos = sin = None + if cache_kind == "swa" and positions is not None: + rope = kwargs.get("rope_views") + if rope is None: + rope = batch_shared.get("rope") + if rope is None: + rope = get_cos_and_sin_dsa(positions, use_cache=coordinates["num_prefills"] == 0) + batch_shared["rope"] = rope + cos, sin = rope + text_config = self.vllm_config.model_config.hf_text_config + window_size = int(_config_value(text_config, "sliding_window", 0)) + n_local_heads = ( + int(_config_value(text_config, "num_attention_heads")) + // self.vllm_config.parallel_config.tensor_parallel_size + ) + n_local_heads = int(kwargs.get("num_query_heads", n_local_heads)) + head_dim = int(_config_value(text_config, "head_dim")) + index_topk = int(_config_value(text_config, "index_topk")) + ori_sparse_indices = kwargs.get("ori_sparse_indices") + noncausal = not bool(getattr(common, "causal", True)) + if noncausal and not isinstance(spec, DeepseekV41DraftSWASpec): + raise ValueError("V4.1 noncausal attention requires a DSpark draft SWA cache") + if noncausal and ori_sparse_indices is None: + ori_sparse_indices, _ = build_dspark_swa_indices( + common.block_table_tensor[:num_reqs], + self.vllm_config.speculative_config.num_speculative_tokens, + window_size, + spec.storage_block_size, + common.query_start_loc[: num_reqs + 1], + seq_lens, + num_actual_tokens, + use_logical_indices=True, + ) + ori_topk_length = ( + (ori_sparse_indices >= 0).sum(dim=-1, dtype=torch.int32) + if ori_sparse_indices is not None and noncausal + else None + ) + ori_mask_mode = 0 if noncausal else 4 + ori_win_left = max(0, window_size - 1) + ori_win_right = 0 + operator_ratio = 0 if cache_kind == "swa" else ratio + smla_metadata = None + qli_metadata = None + + if self._build_query_metadata and self._supports_device_ops and cache_kind in {"swa", "long_kv"}: + has_compressed = operator_ratio in (1, 2) + cmp_seq_lens = coordinates["cache_seq_lens"] if has_compressed else None + cmp_residual = cmp_residual_buffer + + def build_smla_metadata() -> None: + # Keep graph event frontiers stable even when this CP rank has no query. + if num_actual_tokens == 0: + self._smla_metadata.zero_() + return + value = torch.ops._C_ascend.npu_sparse_flash_mla_metadata( + n_local_heads, + 1, + head_dim, + cu_seqlens_q=common.query_start_loc[: num_reqs + 1].int(), + seqused_ori_kv=seq_lens, + seqused_cmp_kv=cmp_seq_lens, + cmp_residual_kv=cmp_residual, + batch_size=num_reqs, + max_seqlen_q=int(getattr(common, "max_query_len", 0)), + max_seqlen_ori_kv=int(getattr(common, "max_seq_len", 0)), + max_seqlen_cmp_kv=(coordinates["max_cache_seq_len"] if has_compressed else 0), + ori_topk=ori_sparse_indices.shape[-1] if ori_sparse_indices is not None else 0, + ori_topk_length=ori_topk_length, + cmp_topk=index_topk if has_compressed else 0, + cmp_ratio=operator_ratio, + ori_mask_mode=ori_mask_mode, + cmp_mask_mode=3 if has_compressed else 0, + ori_win_left=ori_win_left, + ori_win_right=ori_win_right, + layout_q="TND", + layout_kv="PA_BBND", + has_ori_kv=True, + has_cmp_kv=has_compressed, + ) + self._smla_metadata.copy_(value) + + smla_metadata = self._publish_task( + batch_shared, + f"smla:c{operator_ratio}", + self._smla_metadata, + DeviceMetadataStage.ATTENTION, + build_smla_metadata, + ) + + if self._build_query_metadata and self._supports_device_ops and cache_kind == "index_k": + residual = cmp_residual_buffer + + def build_qli_metadata() -> None: + value = torch.ops._C_ascend.npu_quant_lightning_indexer_v2_metadata( + int(_config_value(text_config, "index_n_heads")), + 1, + int(_config_value(text_config, "index_head_dim")), + index_topk, + 2, + cu_seqlens_q=common.query_start_loc[: num_reqs + 1].int(), + seqused_k=coordinates["cache_seq_lens"], + cmp_residual_k=residual, + batch_size=num_reqs, + max_seqlen_q=int(getattr(common, "max_query_len", 0)), + max_seqlen_k=coordinates["max_cache_seq_len"], + layout_q="TND", + layout_k="PA_BBND", + mask_mode=3, + cmp_ratio=ratio, + ) + self._qli_metadata.copy_(value) + + qli_metadata = self._publish_task( + batch_shared, + f"qli:c{ratio}", + self._qli_metadata, + DeviceMetadataStage.INDEXER, + build_qli_metadata, + ) + + c2_ring_metadata = None + c2_complete_mask = None + c2_source_positions = None + c2_source_cos = None + c2_source_sin = None + c2_metadata_group_id = None + if self._build_compressor_metadata and cache_kind == "compressor_state" and positions is not None: + ring_meta = self._c2_ring_metadata[: 5 * num_reqs].view(5, num_reqs) + input_positions = positions + if self._supports_device_ops: + if self._c2_full_source_rope is None: + raise RuntimeError("V4.1 source RoPE buffers were not initialized") + full_source_cos, full_source_sin = self._c2_full_source_rope + else: + full_source_cos = full_source_sin = None + + def build_c2_metadata() -> None: + starts = common.query_start_loc[:num_reqs].int() + ends = common.query_start_loc[1 : num_reqs + 1].int() + query_lens = ends - starts + live = torch.arange(num_reqs, device=starts.device) < num_actual_reqs + used = (ends.clamp_max(num_actual_tokens) - starts).clamp_min(0) + used = torch.where(live, used, 0) + if kwargs.get("skip_ring_state_update", False): + used = torch.zeros_like(used) + ring_meta[0].copy_((seq_lens - query_lens).clamp_min(0)) + ring_meta[1].copy_(used) + ring_meta[2].copy_(starts) + ring_meta[3].copy_(starts) + ring_meta[4].copy_(torch.where(used > 0, common.block_table_tensor[:num_reqs, 0], 0)) + valid_end = common.query_start_loc[num_actual_reqs].clamp_max(num_actual_tokens) + valid = torch.arange(num_input_tokens, device=input_positions.device) < valid_end + complete = (input_positions.remainder(2) == 1) & valid + if kwargs.get("skip_ring_state_update", False): + complete = torch.zeros_like(complete) + self._c2_complete_mask[:num_input_tokens].copy_(complete) + self._c2_source_positions[:num_input_tokens].copy_( + torch.where( + complete, + input_positions - 1, + torch.zeros_like(input_positions), + ) + ) + if full_source_cos is not None and full_source_sin is not None: + gather_idx = ( + self._c2_source_positions[:num_input_tokens] + .reshape(-1, 1, 1, 1) + .expand( + num_input_tokens, + 1, + 1, + full_source_cos.shape[-1], + ) + ) + torch.gather( + full_source_cos, + 0, + gather_idx, + out=self._c2_source_cos[:num_input_tokens], + ) + torch.gather( + full_source_sin, + 0, + gather_idx, + out=self._c2_source_sin[:num_input_tokens], + ) + + compressor_group = self._publish_task( + shared, + "c2:compressor", + self._c2_complete_mask, + DeviceMetadataStage.COMPRESSOR, + build_c2_metadata, + ) + if compressor_group is not self._c2_complete_mask: + raise RuntimeError("V4.1 compressor metadata must have one owner") + c2_complete_mask = self._c2_complete_mask[:num_input_tokens] + c2_ring_metadata = ring_meta + c2_source_positions = self._c2_source_positions[:num_input_tokens] + if self._supports_device_ops: + c2_source_cos = self._c2_source_cos[:num_input_tokens] + c2_source_sin = self._c2_source_sin[:num_input_tokens] + c2_metadata_group_id = id(self._c2_complete_mask) + return DeepseekV41Metadata( + block_table=common.block_table_tensor[:num_reqs], + block_table_cpu=( + common.block_table_cpu[:num_reqs] if getattr(common, "block_table_cpu", None) is not None else None + ), + slot_mapping=slots, + compress_ratio=ratio, + storage_block_size=spec.storage_block_size, + is_compressor_state=is_compressor_state, + cache_kind=cache_kind, + cos=cos, + sin=sin, + num_actual_tokens=num_actual_tokens, + num_input_tokens=num_input_tokens, + num_reqs=num_reqs, + num_actual_reqs=num_actual_reqs, + logical_block_size=spec.block_size, + max_query_len=int(getattr(common, "max_query_len", 0)), + max_seq_len=int(getattr(common, "max_seq_len", 0)), + attn_state=getattr(common, "attn_state", None), + is_prefilling=getattr(common, "is_prefilling", None), + causal=getattr(common, "causal", True), + ori_sparse_indices=ori_sparse_indices, + ori_topk_length=ori_topk_length, + ori_mask_mode=ori_mask_mode, + ori_win_left=ori_win_left, + ori_win_right=ori_win_right, + smla_metadata=smla_metadata, + qli_metadata=qli_metadata, + cmp_residual=cmp_residual_buffer, + c2_ring_metadata=c2_ring_metadata, + c2_complete_mask=c2_complete_mask, + c2_source_positions=c2_source_positions, + c2_source_cos=c2_source_cos, + c2_source_sin=c2_source_sin, + c2_metadata_group_id=c2_metadata_group_id, + **coordinates, + ) + + +class DeepseekV41CacheBackend(AttentionBackend): + """Cache-only backend: supplies layout and metadata, not an AttentionImpl.""" + + @staticmethod + def get_name(): + return "ASCEND_DSA_V41_CACHE" + + @staticmethod + def get_impl_cls(): + return DeepseekV41EagerAttentionImpl + + @staticmethod + def get_builder_cls(): + from vllm_ascend.attention.context_parallel.dsa_v41_cp import get_v41_cp_classes + + return get_v41_cp_classes()[0] + + @classmethod + def supports_pcp(cls) -> bool: + return False + + @staticmethod + def get_kv_cache_shape(num_blocks, block_size, num_kv_heads, head_size, cache_dtype_str="auto"): + return num_blocks, block_size, num_kv_heads, head_size + + +class DeepseekV41CacheLayer(nn.Module, AttentionLayerBase): + supports_dcp = False + + def __init__(self, vllm_config, prefix, spec): + super().__init__() + self.prefix = prefix + self.spec = spec + self.kv_cache = [torch.empty(0)] + context = vllm_config.compilation_config.static_forward_context + if prefix in context: + raise ValueError(f"Duplicate V4.1 cache prefix: {prefix}") + context[prefix] = self + + def get_kv_cache_spec(self, vllm_config): + return self.spec + + def get_attn_backend(self): + return DeepseekV41CacheBackend diff --git a/vllm_ascend/attention/utils.py b/vllm_ascend/attention/utils.py index d1809fe560c0..191049691b9a 100644 --- a/vllm_ascend/attention/utils.py +++ b/vllm_ascend/attention/utils.py @@ -289,6 +289,9 @@ class AscendCommonAttentionMetadata(CommonAttentionMetadata): # E.g., tensor([128, 256, 64]) for 3 requests with different seq lengths. seq_lens_cpu: torch.Tensor = None + # Host mirror of this cache group's block table, including padded rows. + block_table_cpu: torch.Tensor | None = None + # CPU tensor of already computed tokens count per request. # E.g., tensor([100, 200, 50]) means req0 has 100 tokens already computed. num_computed_tokens_cpu: torch.Tensor = None @@ -346,6 +349,7 @@ def _slice_reqs(x): # there will be error about shape mismatch during reshape and cache. # This is really strange since vLLM slices them as well block_table_tensor=self.block_table_tensor, + block_table_cpu=self.block_table_cpu, slot_mapping=self.slot_mapping, causal=self.causal, actual_seq_lengths_q=self.actual_seq_lengths_q[:num_actual_tokens], diff --git a/vllm_ascend/core/circular_buffer.py b/vllm_ascend/core/circular_buffer.py new file mode 100644 index 000000000000..7d75e4ac61c5 --- /dev/null +++ b/vllm_ascend/core/circular_buffer.py @@ -0,0 +1,140 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""v0.27-compatible, single-page circular scratch ownership.""" + +from dataclasses import dataclass + +from vllm.v1.core.single_type_kv_cache_manager import FullAttentionManager +from vllm.v1.kv_cache_interface import AttentionSpec, KVCacheSpec, UniformTypeKVCacheSpecs + + +def prefix_cacheable(spec): + if isinstance(spec, UniformTypeKVCacheSpecs): + return all(prefix_cacheable(member) for member in spec.kv_cache_specs.values()) + return bool(getattr(spec, "prefix_cacheable", True)) and bool(getattr(spec, "participates_in_prefix_caching", True)) + + +# Upstream gained these properties after v0.27.1. Honor the existing GLM +# opt-out as well, without changing its manager or overwriting newer APIs. +if not hasattr(KVCacheSpec, "prefix_cacheable"): + KVCacheSpec.prefix_cacheable = property(lambda self: getattr(self, "participates_in_prefix_caching", True)) +if "prefix_cacheable" not in UniformTypeKVCacheSpecs.__dict__: + UniformTypeKVCacheSpecs.prefix_cacheable = property( + lambda self: all(prefix_cacheable(s) for s in self.kv_cache_specs.values()) + ) + + +@dataclass(frozen=True, kw_only=True) +class AscendCircularBufferSpec(AttentionSpec): + """A single packed plane, rather than AttentionSpec's default K+V pair.""" + + @property + def storage_block_size(self): + return self.block_size + + @property + def real_page_size_bytes(self): + return self.block_size * self.num_kv_heads * self.head_size * self.dtype.itemsize + + @property + def unpadded_page_size_bytes(self): + """Expose the circular buffer's single-plane size to latest vLLM.""" + return self.real_page_size_bytes + + @property + def prefix_cacheable(self): + return False + + def max_memory_usage_bytes(self, vllm_config): + return self.page_size_bytes + + def max_num_blocks_per_req(self, vllm_config, max_len): + return 1 + + def is_uniform_with_collection(self, specs): + return all(type(s) is type(self) and s.block_size == self.block_size for s in specs.values()) + + +def is_circular_spec(spec): + if isinstance(spec, UniformTypeKVCacheSpecs): + return bool(spec.kv_cache_specs) and all(is_circular_spec(s) for s in spec.kv_cache_specs.values()) + return isinstance(spec, AscendCircularBufferSpec) + + +class AscendCircularBufferManager(FullAttentionManager): + """One private page until free/preemption; no prefix hits or pruning.""" + + supports_fine_grained_hash_lookup = False + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self._record_new_block_ids = False + + def _claim_ring_block(self, request_id): + blocks = self.req_to_blocks[request_id] + if blocks: + return [] + new_blocks = self.block_pool.get_new_blocks(1) + blocks.extend(new_blocks) + return new_blocks + + def get_num_blocks_to_allocate( + self, + request_id, + num_tokens, + new_computed_blocks, + total_computed_tokens, + num_local_computed_tokens, + num_tokens_main_model, + apply_admission_cap=False, + ): + return 0 if self.req_to_blocks.get(request_id) else 1 + + def allocate_new_blocks(self, request_id, num_tokens, num_tokens_main_model): + return self._claim_ring_block(request_id) + + def allocate_external_computed_blocks(self, request_id, num_local_computed_tokens, num_external_computed_tokens): + self._claim_ring_block(request_id) + + @classmethod + def find_longest_cache_hit( + cls, + block_hashes, + max_length, + kv_cache_group_ids, + block_pool, + kv_cache_spec, + drop_eagle_block, + alignment_tokens, + dcp_world_size=1, + pcp_world_size=1, + ): + return tuple([] for _ in kv_cache_group_ids), 0 + + def cache_blocks( + self, + request, + num_tokens, + retention_interval=None, + *, + replay_boundary=None, + ): + pass + + def add_local_computed_blocks( + self, + request_id, + new_computed_blocks, + num_local_computed_tokens, + num_external_computed_tokens, + ): + pass + + def remove_skipped_blocks(self, request_id, processed_computed_tokens, num_prompt_tokens=None): + pass + + def get_num_common_prefix_blocks(self, running_request_id): + return 0 + + def get_num_skipped_tokens(self, num_computed_tokens): + return 0 diff --git a/vllm_ascend/core/deepseek_v41.py b/vllm_ascend/core/deepseek_v41.py new file mode 100644 index 000000000000..47bb61cb5aea --- /dev/null +++ b/vllm_ascend/core/deepseek_v41.py @@ -0,0 +1,380 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Framework-side V4.1 cache specs and layer-outermost hybrid allocation.""" + +from dataclasses import dataclass, replace + +import torch +from vllm.config import CUDAGraphMode +from vllm.v1.core.kv_cache_utils import may_override_num_blocks +from vllm.v1.kv_cache_interface import KVCacheGroupSpec, KVCacheTensor, UniformTypeKVCacheSpecs + +from vllm_ascend.core.circular_buffer import AscendCircularBufferSpec +from vllm_ascend.core.kv_cache_interface import ( + AscendMLAAttentionSpec, + AscendSlidingWindowMLASpec, + get_kv_cache_compression_ratio, + get_storage_block_size, +) +from vllm_ascend.utils import vllm_version_is + +STATE_RING_ROWS = 32 + + +@dataclass(frozen=True, kw_only=True) +class DeepseekV41FullSpec(AscendMLAAttentionSpec): + def is_uniform_with_collection(self, specs): + return all( + isinstance(s, (DeepseekV41FullSpec, DeepseekV41IndexerSpec)) + and s.block_size == self.block_size + and get_kv_cache_compression_ratio(s) in (1, 2) + for s in specs.values() + ) + + +@dataclass(frozen=True, kw_only=True) +class DeepseekV41IndexerSpec(AscendMLAAttentionSpec): + """INT8 index keys followed by FP16 scales inside each shared slot page.""" + + def is_uniform_with_collection(self, specs): + return all( + isinstance(s, (DeepseekV41FullSpec, DeepseekV41IndexerSpec)) + and s.block_size == self.block_size + and get_kv_cache_compression_ratio(s) in (1, 2) + for s in specs.values() + ) + + +@dataclass(frozen=True, kw_only=True) +class DeepseekV41SWASpec(AscendSlidingWindowMLASpec): + def is_uniform_with_collection(self, specs): + return all( + isinstance(s, DeepseekV41SWASpec) and s.sliding_window == self.sliding_window for s in specs.values() + ) + + +@dataclass(frozen=True, kw_only=True) +class DeepseekV41DraftSWASpec(AscendSlidingWindowMLASpec): + """DSpark SWA owned by G12, aliasing target slots at distinct block IDs.""" + + def __post_init__(self): + if self.dtype != torch.bfloat16 or self.num_kv_heads != 1 or self.compress_ratio != 1: + raise ValueError("Aurora DSpark requires one uncompressed BF16 KV plane") + + def is_uniform_with_collection(self, specs): + return all( + isinstance(s, DeepseekV41DraftSWASpec) + and s.block_size == self.block_size + and s.sliding_window == self.sliding_window + for s in specs.values() + ) + + +@dataclass(frozen=True, kw_only=True) +class DeepseekV41CompressorStateSpec(AscendCircularBufferSpec): + """One private FP32 KV/score ring page for each active request.""" + + compress_ratio: int = 1 + + def __post_init__(self): + if self.dtype != torch.float32 or self.block_size != STATE_RING_ROWS or self.compress_ratio != 1: + raise ValueError("Aurora state requires a 32-row FP32 uncompressed ring") + if self.num_kv_heads != 1: + raise ValueError("Aurora state requires one packed KV/score plane") + + +def is_v41_spec(spec): + return isinstance( + spec, + ( + DeepseekV41FullSpec, + DeepseekV41IndexerSpec, + DeepseekV41SWASpec, + DeepseekV41DraftSWASpec, + DeepseekV41CompressorStateSpec, + ), + ) + + +def _uniform(members, label): + if not members: + raise ValueError(f"V4.1 cache group {label} is empty") + uniform = UniformTypeKVCacheSpecs.from_specs(members) + if uniform is None: + raise ValueError(f"Incompatible V4.1 resource layouts in {label}") + return uniform + + +@dataclass(frozen=True) +class CachePlacement: + name: str + offset: int + page_size_bytes: int + + +@dataclass(frozen=True) +class CacheSlot: + page_size_bytes: int + placements: tuple[CachePlacement, ...] + + +def _layer_number(name): + try: + return int(name.rsplit(".layers.", 1)[1].split(".", 1)[0]) + except (IndexError, ValueError) as exc: + raise ValueError(f"Invalid V4.1 cache resource name: {name}") from exc + + +def _cache_plane_sizes(spec): + rows = get_storage_block_size(spec) * spec.num_kv_heads + key_bytes = rows * spec.head_size * spec.dtype.itemsize + if isinstance(spec, DeepseekV41IndexerSpec): + return key_bytes, rows * spec.scale_dim * spec.scale_dtype.itemsize + return (key_bytes,) + + +def _draft_layer_number(name): + try: + return int(("." + name).rsplit(".mtp.", 1)[1].split(".", 1)[0]) + except (IndexError, ValueError) as exc: + raise ValueError(f"Invalid Aurora DSpark cache resource name: {name}") from exc + + +def plan_cache_slots(specs): + """Place source KV/index tuples, state and SWA in four shared layer slots. + + Sizes come from payloads, never previously padded specs. Different groups + overlay a slot at distinct live block IDs; a source's KV and index share + the same ID at disjoint offsets within its page. + """ + if not all(is_v41_spec(spec) for spec in specs.values()): + raise ValueError( + "V4.1 requires explicit target or Aurora DSpark cache specs; foreign resources are unsupported" + ) + full = sorted((n for n, s in specs.items() if isinstance(s, DeepseekV41FullSpec)), key=_layer_number) + state = sorted((n for n, s in specs.items() if isinstance(s, DeepseekV41CompressorStateSpec)), key=_layer_number) + swa = sorted((n for n, s in specs.items() if isinstance(s, DeepseekV41SWASpec)), key=_layer_number) + draft = sorted((n for n, s in specs.items() if isinstance(s, DeepseekV41DraftSWASpec)), key=_draft_layer_number) + if draft and list(map(_draft_layer_number, draft)) != [0, 1, 2]: + raise ValueError("Aurora DSpark requires exactly three ordered draft layers: mtp.0, mtp.1, mtp.2") + if list(map(_layer_number, full)) != [2, 8, 14, 20]: + raise ValueError("V4.1 requires KV source layers 2, 8, 14, 20") + if list(map(_layer_number, state)) != [2, 8, 14]: + raise ValueError("V4.1 requires state source layers 2, 8, 14") + if list(map(_layer_number, swa)) != list(range(40)): + raise ValueError("V4.1 requires exactly 40 ordered SWA resources") + + slots = [] + for slot_idx, kv_name in enumerate(full): + prefix, suffix = kv_name.rsplit(".", 1) + index_name = prefix + ".indexer.k_cache" + index_spec = specs.get(index_name) + kv_spec = specs[kv_name] + ratio = 2 if slot_idx < len(state) else 1 + if ( + suffix != "long_kv_cache" + or not isinstance(index_spec, DeepseekV41IndexerSpec) + or get_kv_cache_compression_ratio(kv_spec) != ratio + or get_kv_cache_compression_ratio(index_spec) != ratio + or kv_spec.block_size != index_spec.block_size + ): + raise ValueError(f"V4.1 source {prefix} has incompatible KV/index specs") + aliases = ([state[slot_idx]] if slot_idx < len(state) else []) + swa[slot_idx :: len(full)] + kv_bytes = sum(_cache_plane_sizes(kv_spec)) + index_bytes = sum(_cache_plane_sizes(index_spec)) + capacity = max(kv_bytes + index_bytes, *(sum(_cache_plane_sizes(specs[n])) for n in aliases)) + if slot_idx < len(draft): + draft_name = draft[slot_idx] + draft_spec = specs[draft_name] + swa_spec = specs[swa[slot_idx]] + if ( + draft_spec.block_size != swa_spec.block_size + or draft_spec.head_size != swa_spec.head_size + or draft_spec.sliding_window != swa_spec.sliding_window + or sum(_cache_plane_sizes(draft_spec)) > capacity + ): + raise ValueError("Aurora DSpark geometry must match target SWA and fit its existing slot") + aliases.append(draft_name) + placements = [ + CachePlacement(kv_name, 0, kv_bytes), + CachePlacement(index_name, kv_bytes, capacity - kv_bytes), + *(CachePlacement(name, 0, capacity) for name in aliases), + ] + slots.append(CacheSlot(capacity, tuple(placements))) + names = [p.name for slot in slots for p in slot.placements] + if len(names) != len(set(names)) or set(names) != set(specs): + raise ValueError("V4.1 slot placement must cover each resource exactly once") + return tuple(slots) + + +def group_cache_specs(specs): + """Merge full-context resources and pad layer tuples without mutating inputs.""" + if not any(is_v41_spec(s) for s in specs.values()): + return None + slots = plan_cache_slots(specs) + padded = { + p.name: replace(specs[p.name], page_size_padded=p.page_size_bytes) for slot in slots for p in slot.placements + } + full = {n: s for n, s in padded.items() if isinstance(s, (DeepseekV41FullSpec, DeepseekV41IndexerSpec))} + state = {n: s for n, s in padded.items() if isinstance(s, DeepseekV41CompressorStateSpec)} + groups = [_uniform(full, "full"), _uniform(state, "state")] + swa = sorted((n for n, s in padded.items() if isinstance(s, DeepseekV41SWASpec)), key=_layer_number) + groups.extend( + _uniform({n: padded[n] for n in swa[start : start + len(slots)]}, f"swa{start}") + for start in range(0, len(swa), len(slots)) + ) + draft = sorted((n for n, s in padded.items() if isinstance(s, DeepseekV41DraftSWASpec)), key=_draft_layer_number) + if draft: + groups.append(_uniform({n: padded[n] for n in draft}, "dspark")) + return groups + + +def make_cache_groups(grouped_specs): + return [KVCacheGroupSpec(layer_names=list(s.kv_cache_specs), kv_cache_spec=s) for s in grouped_specs] + + +def has_v41_groups(groups): + return any( + is_v41_spec(s) + for g in groups + if isinstance(g.kv_cache_spec, UniformTypeKVCacheSpecs) + for s in g.kv_cache_spec.kv_cache_specs.values() + ) + + +def cache_slots_from_groups(groups): + specs = {} + for group in groups: + if not isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs): + raise ValueError("V4.1 requires uniform-type cache groups") + for name in group.layer_names: + if name in specs: + raise ValueError(f"V4.1 resource belongs to multiple cache groups: {name}") + specs[name] = group.kv_cache_spec.kv_cache_specs[name] + return plan_cache_slots(specs) + + +def pool_bytes_per_block(groups): + return sum(slot.page_size_bytes for slot in cache_slots_from_groups(groups)) + + +def request_blocks(vllm_config, groups): + # Different logical groups consume different IDs in one global block pool. + return sum( + max( + (s.max_memory_usage_bytes(vllm_config) + s.page_size_bytes - 1) // s.page_size_bytes + for s in g.kv_cache_spec.kv_cache_specs.values() + ) + for g in groups + ) + + +def allocate_cache_config(vllm_config, groups, available_memory): + """Allocate four independent layer slots backed by one global block-ID pool.""" + slots = cache_slots_from_groups(groups) + capacity = available_memory // sum(slot.page_size_bytes for slot in slots) + num_blocks = may_override_num_blocks(vllm_config, capacity) + if num_blocks <= 1 or num_blocks > capacity: + raise ValueError("Insufficient V4.1 cache memory (including reserved null block), or unsafe block override") + tensors = [] + for slot in slots: + layer_names = [placement.name for placement in slot.placements] + size = num_blocks * slot.page_size_bytes + if vllm_version_is("0.28.0"): + tensors.append( + KVCacheTensor( + size=size, + shared_by=layer_names, + block_stride=slot.page_size_bytes, + ) + ) + else: + tensors.append( + KVCacheTensor( + size=size, + layers=layer_names, + offset=0, + layer_stride=0, + block_stride=slot.page_size_bytes, + ) + ) + return num_blocks, tensors + + +def reshape_cache(raw: torch.Tensor, spec, *, num_blocks, offset, block_stride): + """Create typed per-page views using the containing slot's physical stride.""" + if raw.dtype not in (torch.int8, torch.uint8) or raw.ndim != 1 or not raw.is_contiguous(): + raise ValueError("V4.1 cache requires contiguous one-dimensional byte storage") + if num_blocks <= 0 or block_stride <= 0 or raw.numel() != num_blocks * block_stride: + raise ValueError("V4.1 cache backing does not match its declared layout") + plane_sizes = _cache_plane_sizes(spec) + if offset < 0 or offset + sum(plane_sizes) > block_stride: + raise ValueError("V4.1 cache component exceeds its slot page") + if isinstance(spec, DeepseekV41CompressorStateSpec) and sum(plane_sizes) != block_stride: + raise ValueError("Aurora circular state must fill its slot with 32 contiguous FP32 rows") + storage_block_size = get_storage_block_size(spec) + + def view(dtype, width, byte_offset): + dtype_size = dtype.itemsize + storage_offset = raw.storage_offset() + byte_offset + if storage_offset % dtype_size or block_stride % dtype_size or raw.numel() % dtype_size: + raise ValueError("V4.1 cache offset/stride is not dtype aligned") + return torch.as_strided( + raw.view(dtype), + size=(num_blocks, storage_block_size, spec.num_kv_heads, width), + stride=(block_stride // dtype_size, spec.num_kv_heads * width, width, 1), + storage_offset=storage_offset // dtype_size, + ) + + key = view(spec.dtype, spec.head_size, offset) + if isinstance(spec, DeepseekV41IndexerSpec): + return key, view(spec.scale_dtype, spec.scale_dim, offset + plane_sizes[0]) + return key + + +def validate_cache_runtime(vllm_config): + if vllm_config.use_v2_model_runner: + raise NotImplementedError("V4.1 cache initialization currently requires model runner V1") + if getattr(vllm_config, "kv_transfer_config", None) is not None: + raise NotImplementedError("V4.1 cache initialization does not support KV transfer") + cudagraph_mode = getattr( + vllm_config.compilation_config, + "cudagraph_mode", + CUDAGraphMode.NONE if vllm_config.model_config.enforce_eager else CUDAGraphMode.FULL, + ) + if cudagraph_mode not in ( + CUDAGraphMode.NONE, + CUDAGraphMode.FULL_DECODE_ONLY, + ): + raise NotImplementedError("V4.1 currently supports only eager or FULL_DECODE_ONLY graph mode") + speculative = vllm_config.speculative_config + if speculative is not None: + use_dspark = getattr(speculative, "use_dspark", None) + if not callable(use_dspark) or not use_dspark(): + raise NotImplementedError("Aurora supports only DSpark speculative decoding") + # Verification writes the anchor and up to S speculative rows. After + # rejection, the earliest needed residual is the verified anchor. + # It must survive the final 32-row write: S must be strictly below 32. + if not 0 < speculative.num_speculative_tokens < STATE_RING_ROWS: + raise ValueError("Aurora DSpark requires 1..31 speculative tokens to preserve FP32 ring residuals") + per_batch = getattr(speculative, "num_speculative_tokens_per_batch_size", None) or () + if any(not 0 <= count <= speculative.num_speculative_tokens for _, _, count in per_batch): + raise ValueError("Aurora DSpark per-batch speculation must stay within the configured ring-safe maximum") + parallel = vllm_config.parallel_config + if any( + getattr(parallel, name, 1) != 1 + for name in ( + "pipeline_parallel_size", + "decode_context_parallel_size", + "prefill_context_parallel_size", + ) + ): + raise NotImplementedError("V4.1 initial runtime requires PP=DCP=PCP=1") + if vllm_config.scheduler_config.disable_hybrid_kv_cache_manager: + raise ValueError("V4.1 requires the hybrid KV cache manager") + if vllm_config.cache_config.cache_dtype not in ("auto", "bfloat16"): + raise NotImplementedError("V4.1 initial cache layout requires BF16") + if speculative is not None: + # Aurora's planes are always BF16. Pin the inherited DSV4 draft + # backend to the same layout, including on hardware where auto is FP8. + vllm_config.cache_config.cache_dtype = "bfloat16" diff --git a/vllm_ascend/core/kv_cache_interface.py b/vllm_ascend/core/kv_cache_interface.py index 8516e7ad4b87..f77fa35782e8 100644 --- a/vllm_ascend/core/kv_cache_interface.py +++ b/vllm_ascend/core/kv_cache_interface.py @@ -19,6 +19,7 @@ ) from vllm.v1.kv_cache_spec_registry import KVCacheSpecRegistry +from vllm_ascend.core.circular_buffer import AscendCircularBufferManager, AscendCircularBufferSpec from vllm_ascend.utils import vllm_version_is @@ -287,6 +288,27 @@ def max_memory_usage_bytes(self, vllm_config: VllmConfig) -> int: def register_ascend_kv_cache_specs() -> None: + from vllm_ascend.core.deepseek_v41 import ( + DeepseekV41CompressorStateSpec, + DeepseekV41DraftSWASpec, + DeepseekV41FullSpec, + DeepseekV41IndexerSpec, + DeepseekV41SWASpec, + ) + + KVCacheSpecRegistry.register( + kvcache_spec_cls=AscendCircularBufferSpec, + manager_class=AscendCircularBufferManager, + uniform_type_base_spec=AscendCircularBufferSpec, + ) + for spec, manager in ( + (DeepseekV41FullSpec, FullAttentionManager), + (DeepseekV41IndexerSpec, FullAttentionManager), + (DeepseekV41SWASpec, SlidingWindowManager), + (DeepseekV41DraftSWASpec, SlidingWindowManager), + (DeepseekV41CompressorStateSpec, AscendCircularBufferManager), + ): + KVCacheSpecRegistry.register(kvcache_spec_cls=spec, manager_class=manager, uniform_type_base_spec=spec) KVCacheSpecRegistry.register( kvcache_spec_cls=AscendMLAAttentionSpec, manager_class=FullAttentionManager, diff --git a/vllm_ascend/deepseek_v41_config.py b/vllm_ascend/deepseek_v41_config.py new file mode 100644 index 000000000000..fa6b28de78ec --- /dev/null +++ b/vllm_ascend/deepseek_v41_config.py @@ -0,0 +1,178 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Transformers config classes for the downstream DeepSeek V4.1 port.""" + +from typing import Any + +from transformers.configuration_utils import PretrainedConfig + + +def _mirror_config_aliases( + values: dict[str, Any], + aliases: dict[str, str], +) -> dict[str, Any]: + """Expose both released and pre-release field names, rejecting conflicts.""" + values = dict(values) + for legacy_name, released_name in aliases.items(): + legacy = values.get(legacy_name) + released = values.get(released_name) + if legacy is not None and released is not None and legacy != released: + raise ValueError( + f"Conflicting DeepSeek V4.1 config fields: {legacy_name}={legacy!r}, {released_name}={released!r}" + ) + value = released if released is not None else legacy + if value is not None: + values[legacy_name] = value + values[released_name] = value + return values + + +class DeepseekV41TextConfig(PretrainedConfig): + model_type = "deepseek_v41_text" + base_config_key = "text_config" + + def __init__(self, model_type: str = "deepseek_v41_text", **kwargs: Any) -> None: + kwargs = _mirror_config_aliases( + kwargs, + { + "kv_source_layers": "kv_source_layer_ids", + "index_source_layers": "index_source_layer_ids", + "candidate_source_layer": "candidate_source_layer_id", + "engram_pad_id": "engram_pad_token_id", + "dspark_n_activated_experts": "dspark_num_experts_per_tok", + }, + ) + for name, value in kwargs.items(): + setattr(self, name, value) + + rope = dict(kwargs.get("rope_scaling") or kwargs.get("rope_parameters") or {}) + rope.setdefault("factor", 1.0) + rope.setdefault("beta_fast", 32) + rope.setdefault("beta_slow", 1) + rope.setdefault( + "original_max_position_embeddings", + kwargs.get("max_position_embeddings", 1048576), + ) + rope.setdefault("rope_theta", kwargs.get("rope_theta", 10000.0)) + self.rope_parameters = rope + + base_kwargs = dict(kwargs) + base_kwargs.pop("rope_scaling", None) + base_kwargs.pop("rope_parameters", None) + super().__init__(**base_kwargs) + self.model_type = model_type + + self.num_hash_layers = int(kwargs.get("num_hash_layers", 0)) + self.n_group = int(kwargs.get("n_group", 1)) + self.topk_group = int(kwargs.get("topk_group", 1)) + self.first_k_dense_replace = int(kwargs.get("first_k_dense_replace", 0)) + self.moe_layer_freq = int(kwargs.get("moe_layer_freq", 1)) + + +class DeepseekV41VisionConfig(PretrainedConfig): + model_type = "deepseek_v41_vision" + base_config_key = "vision_config" + + def __init__(self, model_type: str = "deepseek_v41_vision", **kwargs: Any) -> None: + kwargs = _mirror_config_aliases( + kwargs, + {"max_num_tokens": "max_image_tokens"}, + ) + super().__init__(**kwargs) + self.model_type = model_type + for name, value in kwargs.items(): + setattr(self, name, value) + + +class DeepseekV41Config(PretrainedConfig): + model_type = "deepseek_v41" + sub_configs = { + "text_config": DeepseekV41TextConfig, + "vision_config": DeepseekV41VisionConfig, + } + + def __init__( + self, + text_config: dict[str, Any] | DeepseekV41TextConfig | None = None, + vision_config: dict[str, Any] | DeepseekV41VisionConfig | None = None, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + text_config = text_config or {} + vision_config = vision_config or {} + text_field_names = set(text_config) if isinstance(text_config, dict) else set(vars(text_config)) + self.text_config = ( + text_config if isinstance(text_config, DeepseekV41TextConfig) else DeepseekV41TextConfig(**text_config) + ) + self.vision_config = ( + vision_config + if isinstance(vision_config, DeepseekV41VisionConfig) + else DeepseekV41VisionConfig(**vision_config) + ) + # Include mirrored release/pre-release aliases when flattening the + # text config for model implementations that consume the root config. + text_field_names.update(vars(self.text_config)) + + text_field_names.update( + { + "rope_parameters", + "num_hash_layers", + "n_group", + "topk_group", + "first_k_dense_replace", + "moe_layer_freq", + } + ) + for name in text_field_names: + if not name.startswith("_") and name not in {"architectures", "model_type"}: + setattr(self, name, getattr(self.text_config, name)) + + vision_aliases = { + "vision_n_layers": "num_hidden_layers", + "vision_dim": "hidden_size", + "vision_n_heads": "num_attention_heads", + "vision_inter_dim": "intermediate_size", + "vision_patch_size": "patch_size", + "vision_rope_theta": "rope_theta", + "vision_downsample_ratio": "downsample_ratio", + "vision_max_n_token": "max_image_tokens", + "vision_min_pixels": "min_pixels", + "vision_max_wh_ratio": "max_wh_ratio", + } + for alias, source in vision_aliases.items(): + setattr(self, alias, getattr(self.vision_config, source, None)) + + # The released checkpoint is still multimodal even though its + # architecture was renamed from *ForConditionalGeneration to + # DeepseekV41ForCausalLM. Presence of the populated vision config, + # rather than the architecture suffix, is the capability signal. + vision_enabled = bool(self.vision_n_layers) + self.image_token_id = int(getattr(self, "image_token_id", 129264)) + # V4.1 uses one image token for every span role. The adjacent reserved + # token is used only for vLLM's ratio-2 compressor-alignment row. + self.image_sentinel_base_id = self.image_token_id + self.image_pad_token_id = self.image_token_id + 1 + self.is_mm_prefix_lm = vision_enabled + self.mm_prefix_clamp_sliding_window = vision_enabled + self.mm_prefix_span_leading_pad_modulus = 2 if vision_enabled else 0 + + # The released W8A8 checkpoint states the Engram basis contract + # explicitly. The current gate restores the hidden stream to the + # original basis and adds a value that is already globally rotated. + # Fail closed for future layouts that require different runtime math. + rotation = getattr(self, "engram_rotation_config", None) + if rotation is None: + rotation = { + "value_projection_rotated": True, + "value_basis": "quarot_global", + "key_and_gate_basis": "original", + "runtime_delta_rotation": False, + } + supported_rotation = { + "value_projection_rotated": True, + "value_basis": "quarot_global", + "key_and_gate_basis": "original", + "runtime_delta_rotation": False, + } + if any(rotation.get(name) != value for name, value in supported_rotation.items()): + raise ValueError(f"Unsupported DeepSeek V4.1 Engram rotation contract: {rotation!r}") + self.engram_rotation_config = dict(rotation) diff --git a/vllm_ascend/distributed/eplb/state.py b/vllm_ascend/distributed/eplb/state.py index 605e07f4a150..4fc275e6b16c 100644 --- a/vllm_ascend/distributed/eplb/state.py +++ b/vllm_ascend/distributed/eplb/state.py @@ -12,11 +12,16 @@ from vllm.distributed import get_ep_group from vllm.distributed.eplb import eplb_state as _eplb_state -from vllm_ascend.ops.fused_moe import eplb as _eplb_ops - ASYNC_EPLB_CYCLE_COMMITTED_LOG = "Ascend async EPLB cycle committed" +def _build_expert_replica_routing_table(*args, **kwargs): + """Load Ascend EPLB ops only when a model refreshes its routing table.""" + from vllm_ascend.ops.fused_moe.eplb import build_expert_replica_routing_table + + return build_expert_replica_routing_table(*args, **kwargs) + + def _upstream_from_mapping_accepts_valid_expert_count() -> bool: """Return whether the selected vLLM uses the release mapping contract.""" return "num_valid_physical_experts" in inspect.signature(_eplb_state.EplbState.from_mapping).parameters @@ -60,7 +65,7 @@ def refresh_expert_replica_routing_table(self) -> None: if logical_to_physical_map is None or logical_replica_count is None: raise RuntimeError("Cannot build the replica routing table before EPLB layer state is initialized.") - new_routing_table = _eplb_ops.build_expert_replica_routing_table( + new_routing_table = _build_expert_replica_routing_table( logical_to_physical_map, logical_replica_count, get_ep_group().rank_in_group, diff --git a/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_hybrid_connector.py b/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_hybrid_connector.py index 00d13923d917..f7aee9bcb193 100644 --- a/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_hybrid_connector.py +++ b/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_hybrid_connector.py @@ -53,6 +53,7 @@ from vllm.v1.request import RequestStatus from vllm_ascend.ascend_config import get_ascend_config, init_ascend_config +from vllm_ascend.core.circular_buffer import is_circular_spec from vllm_ascend.distributed.kv_transfer.utils.mooncake_transfer_engine import global_te from vllm_ascend.distributed.kv_transfer.utils.utils import PD_QOS_DEFAULT, get_transfer_timeout_value, inject_qos from vllm_ascend.utils import enable_custom_op, get_kv_cache_tensor_layers, is_vl_model @@ -1361,7 +1362,11 @@ def _truncate_request_for_prefill(self, request: "Request") -> None: def _compute_transfer_block_ids(self, block_ids: BlockIds, prompt_len: int) -> BlockIds: transfer_block_ids = [] + kv_cache_specs = getattr(self, "kv_cache_specs", ()) for i, blocks in enumerate(block_ids): + if i < len(kv_cache_specs) and all(is_circular_spec(spec) for spec in kv_cache_specs[i]): + transfer_block_ids.append(blocks) + continue group_token_len = prompt_len group_block_len = math.ceil(group_token_len / self.group_block_size[i]) if group_block_len > 0: diff --git a/vllm_ascend/model_executor/warmup/deepseek_v41_triton_warmup.py b/vllm_ascend/model_executor/warmup/deepseek_v41_triton_warmup.py new file mode 100644 index 000000000000..92f0bc72a822 --- /dev/null +++ b/vllm_ascend/model_executor/warmup/deepseek_v41_triton_warmup.py @@ -0,0 +1,63 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Warm the finite tile variants used by the V4.1 indexer.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch +from vllm.triton_utils import HAS_TRITON + +from vllm_ascend.ops.triton.prepare_indexer_indices import prepare_indexer_indices +from vllm_ascend.ops.triton.quantize_indexer_query import quantize_indexer_query +from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num + +if TYPE_CHECKING: + from vllm_ascend.worker.worker import NPUWorker + + +def collect_indexer_warmup_token_counts(topk: int, num_cores: int, max_tokens: int) -> list[int]: + """One token count per reachable ``BLOCK_ROWS`` in index postprocessing.""" + # Match the 128 KiB, eight-buffer sort budget in prepare_indexer_indices. + padded_topk = 1 << (topk - 1).bit_length() + max_block_rows = 128 * 1024 // (padded_topk * 4 * 8) + token_counts = [1] + block_rows = 1 + while block_rows < max_block_rows: + tokens = block_rows * num_cores + 1 + if tokens > max_tokens: + break + token_counts.append(tokens) + block_rows *= 2 + return token_counts + + +@torch.inference_mode() +def deepseek_v41_triton_warmup(worker: NPUWorker) -> None: + """Precompile indexer tiles before serving arbitrary eager token counts.""" + if not HAS_TRITON: + return + config = worker.model_config.hf_text_config + if getattr(config, "model_type", None) not in ( + "deepseek_v4.1", + "deepseek_v41", + "deepseek_v4.1_text", + "deepseek_v41_text", + ): + return + ratios = sorted(set(config.compress_ratios[: config.num_hidden_layers]) - {0}) + if not ratios: + return + + device = worker.device + query = torch.zeros(1, config.index_n_heads, config.index_head_dim, dtype=worker.model_config.dtype, device=device) + quantize_indexer_query(query) + token_counts = collect_indexer_warmup_token_counts( + config.index_topk, get_vectorcore_num(), worker.scheduler_config.max_num_batched_tokens + ) + for tokens in token_counts: + selected = torch.zeros(tokens, config.index_topk, dtype=torch.int32, device=device) + positions = torch.zeros(tokens, dtype=torch.int64, device=device) + for ratio in ratios: + prepare_indexer_indices(selected, positions, ratio) diff --git a/vllm_ascend/model_executor/warmup/kernel_warmup.py b/vllm_ascend/model_executor/warmup/kernel_warmup.py index a094ed36847a..772b1a031889 100644 --- a/vllm_ascend/model_executor/warmup/kernel_warmup.py +++ b/vllm_ascend/model_executor/warmup/kernel_warmup.py @@ -10,6 +10,9 @@ from vllm.logger import logger from vllm.triton_utils import HAS_TRITON +from vllm_ascend.model_executor.warmup.deepseek_v41_triton_warmup import ( + deepseek_v41_triton_warmup, +) from vllm_ascend.model_executor.warmup.penalties_triton_warmup import ( penalties_triton_warmup, ) @@ -44,6 +47,7 @@ def kernel_warmup(worker: NPUWorker) -> None: _run_warmup("rejection_sampler", rejection_sampler_triton_warmup, worker) _run_warmup("penalties", penalties_triton_warmup, worker) _run_warmup("rms", triton_rms_warmup, worker) + _run_warmup("deepseek_v41_indexer", deepseek_v41_triton_warmup, worker) elapsed = time.perf_counter() - start logger.info("Triton kernel warmup finished in %.3fs.", elapsed) diff --git a/vllm_ascend/models/__init__.py b/vllm_ascend/models/__init__.py index 0d67c3184a3c..2eb835a89cf4 100644 --- a/vllm_ascend/models/__init__.py +++ b/vllm_ascend/models/__init__.py @@ -35,6 +35,14 @@ def register_model(): "DeepseekV4ForConditionalGeneration", "vllm_ascend.models.deepseek_v4.vl_model:AscendDeepseekV4ForConditionalGeneration", ) + ModelRegistry.register_model( + "DeepseekV41ForConditionalGeneration", + "vllm_ascend.models.deepseek_v41.vl_model:AscendDeepseekV41ForConditionalGeneration", + ) + ModelRegistry.register_model( + "DeepseekV41ForCausalLM", + "vllm_ascend.models.deepseek_v41.vl_model:AscendDeepseekV41ForConditionalGeneration", + ) ModelRegistry.register_model( "MiniMaxM3SparseForCausalLM", "vllm_ascend.models.minimax_m3:MiniMaxM3SparseForCausalLM", @@ -48,6 +56,10 @@ def register_model(): "DSparkDraftModel", "vllm_ascend.models.deepseek_v4.dspark:DSparkDeepseekV4ForCausalLM", ) + ModelRegistry.register_model( + "DeepseekV41DSparkDraftModel", + "vllm_ascend.models.deepseek_v41.dspark:DSparkDeepseekV41ForCausalLM", + ) ModelRegistry.register_model( "LlamaForCausalLMVwnEagle3", "vllm_ascend.models.llama_eagle3_vwn:Eagle3VwnLlamaForCausalLM" ) diff --git a/vllm_ascend/models/deepseek_v4/dspark.py b/vllm_ascend/models/deepseek_v4/dspark.py index 6b58ac16c41f..31a554cf1a25 100644 --- a/vllm_ascend/models/deepseek_v4/dspark.py +++ b/vllm_ascend/models/deepseek_v4/dspark.py @@ -145,6 +145,9 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: } ) + self.needs_moe_input_ids = any( + layer.mlp.gate.tid2eid is not None or layer.mlp.gate.bias_vl is not None for layer in self.layers.values() + ) first_layer = self.layers[str(self.mtp_start_layer_idx)] self.use_sequence_parallel_moe = first_layer.use_sequence_parallel_moe @@ -275,13 +278,16 @@ def forward( input_ids = sp_shard(input_ids) residual = None + moe_input_ids = input_ids + if self.needs_moe_input_ids: + moe_input_ids = torch.where(input_ids == -1, 0, input_ids) for layer in self.layers.values(): hidden_states, residual = layer( positions, hidden_states, residual, llama_4_scaling=None, - input_ids=input_ids, + input_ids=moe_input_ids, ) if use_sp: hidden_states = tensor_model_parallel_all_gather(hidden_states, 0) diff --git a/vllm_ascend/models/deepseek_v4/mm_preprocess.py b/vllm_ascend/models/deepseek_v4/mm_preprocess.py index 84da0a16fa41..e0b00b7e0d57 100644 --- a/vllm_ascend/models/deepseek_v4/mm_preprocess.py +++ b/vllm_ascend/models/deepseek_v4/mm_preprocess.py @@ -390,9 +390,11 @@ def _call_hf_processor( prompt: str, mm_data: Mapping[str, object], mm_kwargs: Mapping[str, object], - tok_kwargs: Mapping[str, object], + tok_kwargs: Mapping[str, object] | None = None, ) -> BatchFeature: - """Combine the local image transform with v0.27 tokenization.""" + """Combine the local image transform with vLLM tokenization.""" + if tok_kwargs is None: + tok_kwargs = {} processor = self.info.get_hf_processor(**mm_kwargs) processed = processor( text=prompt, diff --git a/vllm_ascend/models/deepseek_v4/model.py b/vllm_ascend/models/deepseek_v4/model.py index 04ba9db64b53..3c7dd0c9449d 100644 --- a/vllm_ascend/models/deepseek_v4/model.py +++ b/vllm_ascend/models/deepseek_v4/model.py @@ -372,7 +372,7 @@ def __init__( swiglu_limit=self.swiglu_limit, e_score_correction_bias=self.gate.e_score_correction_bias, bias_vl=self.gate.bias_vl, - image_sentinel_lo=129257, + image_sentinel_lo=getattr(config, "image_sentinel_base_id", 129257), enable_eplb=self.enable_eplb, num_redundant_experts=self.n_redundant_experts, is_sequence_parallel=self.is_sequence_parallel, @@ -385,6 +385,7 @@ def forward( hidden_states: torch.Tensor, input_ids: torch.Tensor | None = None, hidden_states_fp32: torch.Tensor | None = None, + already_sequence_parallel: bool = False, ) -> torch.Tensor: if self.gate.tid2eid is not None and input_ids is None: raise ValueError("DeepSeek V4 hash MoE routing requires input_ids.") @@ -393,6 +394,14 @@ def forward( hidden_states = hidden_states.view(-1, hidden_dim) if hidden_states_fp32 is not None: hidden_states_fp32 = hidden_states_fp32.view(-1, hidden_dim) + # Chunk the hidden states so they aren't replicated across TP ranks. + # This avoids duplicate computation in self.experts. + # TODO: We can replace the all_reduce at the end of attn with a + # reduce_scatter instead of chunking here. + if self.is_sequence_parallel and not already_sequence_parallel: + hidden_states = sequence_parallel_chunk(hidden_states) + if hidden_states_fp32 is not None: + hidden_states_fp32 = sequence_parallel_chunk(hidden_states_fp32) if self.experts.is_internal_router: # In this case, the gate/router runs inside the FusedMoEFactory class @@ -435,7 +444,10 @@ def forward( else: final_hidden_states = fused_moe_out - if not self.is_sequence_parallel and self.tp_size > 1 and fused_moe_out_is_tuple: + if self.is_sequence_parallel and not already_sequence_parallel: + final_hidden_states = sp_all_gather(final_hidden_states) + final_hidden_states = final_hidden_states[:num_tokens] + elif self.tp_size > 1 and fused_moe_out_is_tuple: # Legacy tuple outputs are reduced here. Tensor outputs from the # upstream MoERunner have already gone through its final reduction. final_hidden_states = self.experts.maybe_all_reduce_tensor_model_parallel(final_hidden_states) @@ -460,6 +472,8 @@ def _get_llama_4_scaling( class DeepseekV4Attention(nn.Module): + swa_cache_cls = AscendDeepseekV4SWACache + def __init__( self, vllm_config: VllmConfig, @@ -616,7 +630,7 @@ def __init__( ) k_dtype = get_dsv4_attn_kv_dtype(vllm_config) - swa_cache_layer = AscendDeepseekV4SWACache( + swa_cache_layer = self.swa_cache_cls( head_dim=self.head_dim, window_size=self.window_size, dtype=k_dtype, @@ -671,7 +685,9 @@ def forward( return self.dsa_attn(positions, hidden_states, llama_4_scaling) -class DeepseekV4DecoderLayer(nn.Module): +class DeepseekV2DecoderLayer(nn.Module): + attention_cls = DeepseekV4Attention + def __init__( self, vllm_config: VllmConfig, @@ -698,7 +714,7 @@ def __init__( self.use_sequence_parallel_moe = parallel_config.use_sequence_parallel_moe self.enable_dsa_cp = enable_dsa_cp() # TODO: delete this when enable_dsa_cp is sunset. - attn_cls = DeepseekV4Attention + attn_cls = self.attention_cls self.self_attn = attn_cls( vllm_config=vllm_config, @@ -746,10 +762,10 @@ def rms_norm_cast(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torc return hidden_states, hidden_states.float() def hc_pre(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor): - y = torch.ops._C_ascend.npu_hc_pre_v2( + y, post, comb = torch.ops._C_ascend.npu_hc_pre_v2( x, hc_fn, hc_scale, hc_base, self.hc_mult, self.hc_sinkhorn_iters, self.norm_eps, self.hc_eps ) - return y + return y, post, comb def hc_post(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, comb: torch.Tensor): y = torch.ops._C_ascend.npu_hc_post( @@ -788,18 +804,23 @@ def forward( hidden_states, input_ids=input_ids, hidden_states_fp32=hidden_states_fp32, + already_sequence_parallel=(self.use_sequence_parallel_moe and self.enable_dsa_cp), ) hidden_states = self.hc_post(hidden_states, residual, post, comb) return hidden_states, residual +DeepseekV4DecoderLayer = DeepseekV2DecoderLayer + + @support_torch_compile class DeepseekV4Model(nn.Module, EagleModelMixin): fall_back_to_pt_during_load = False # vLLM #50514 validates and relays the model's existing PP aux payload. supports_aux_hidden_states_over_pp = True AUX_HIDDEN_STATE_KEY = "pp_transport_aux_hidden_states_" + decoder_layer_cls: type[DeepseekV2DecoderLayer] = DeepseekV2DecoderLayer def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() @@ -838,9 +859,13 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): self.embed_tokens = PPMissingLayer() self.start_layer, self.end_layer, self.layers = make_layers( config.num_hidden_layers, - lambda prefix: DeepseekV4DecoderLayer(vllm_config, prefix, topk_indices_buffer=topk_indices_buffer), + lambda prefix: self.decoder_layer_cls(vllm_config, prefix, topk_indices_buffer=topk_indices_buffer), prefix=f"{prefix}.layers", ) + self.needs_moe_input_ids = any( + layer.mlp.gate.tid2eid is not None or layer.mlp.gate.bias_vl is not None + for layer in islice(self.layers, self.start_layer, self.end_layer) + ) if get_pp_group().is_last_rank: self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) @@ -953,13 +978,16 @@ def forward( if pp_group.is_first_rank: hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1) # (b, s, h) -> (b, s, c, h) + moe_input_ids = input_ids + if getattr(self, "needs_moe_input_ids", False): + moe_input_ids = torch.where(input_ids == -1, 0, input_ids) for layer in islice(self.layers, self.start_layer, self.end_layer): hidden_states, residual = layer( positions, hidden_states, residual, llama_4_scaling, - input_ids=input_ids, + input_ids=moe_input_ids, ) if layer.layer_idx + 1 in self.aux_hidden_state_layers: aux_hidden_state = hidden_states.mean(dim=1) diff --git a/vllm_ascend/models/deepseek_v4/vision.py b/vllm_ascend/models/deepseek_v4/vision.py index 2e93646a9ff2..a7fefe2660c8 100644 --- a/vllm_ascend/models/deepseek_v4/vision.py +++ b/vllm_ascend/models/deepseek_v4/vision.py @@ -12,6 +12,10 @@ import torch import torch.nn.functional as F from torch import nn +from vllm.model_executor.layers.activation import SiluAndMul +from vllm.model_executor.layers.attention import MMEncoderAttention +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.rotary_embedding.common import ApplyRotaryEmb @lru_cache(8) @@ -24,25 +28,6 @@ def get_vision_cos_sin(n_h: int, n_w: int, dim: int, theta: float) -> tuple[torc return freqs.cos().unsqueeze(1), freqs.sin().unsqueeze(1) -def apply_rotary(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: - dtype = x.dtype - x1, x2 = x.float().chunk(2, dim=-1) - return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1).to(dtype) - - -class DeepseekV4RMSNorm(nn.Module): - def __init__(self, dim: int, eps: float = 1e-6): - super().__init__() - self.eps = eps - self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32)) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - dtype = x.dtype - x = x.float() - x = x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps) - return (self.weight * x).to(dtype) - - class DeepseekV4PatchEmbed(nn.Module): def __init__(self, config): super().__init__() @@ -59,14 +44,24 @@ def __init__(self, config): self.head_dim = config.vision_dim // config.vision_n_heads self.wqkv = nn.Linear(config.vision_dim, 3 * config.vision_dim) self.wo = nn.Linear(config.vision_dim, config.vision_dim) + self.apply_rotary_emb = ApplyRotaryEmb( + enforce_enable=True, + is_neox_style=True, + enable_fp32_compute=True, + ) + self.attn = MMEncoderAttention( + num_heads=self.n_heads, + head_size=self.head_dim, + scale=self.head_dim**-0.5, + ) def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: n = x.size(0) q, k, v = (t.view(n, self.n_heads, self.head_dim) for t in self.wqkv(x).chunk(3, dim=-1)) - q = apply_rotary(q, cos, sin) - k = apply_rotary(k, cos, sin) - o = F.scaled_dot_product_attention(q.transpose(0, 1), k.transpose(0, 1), v.transpose(0, 1)) - return self.wo(o.transpose(0, 1).reshape(n, -1)) + qk = self.apply_rotary_emb(torch.stack((q, k)), cos.squeeze(1), sin.squeeze(1)) + q, k = qk.unbind() + o = self.attn(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0)) + return self.wo(o.squeeze(0).reshape(n, -1)) class DeepseekV4VisionMLP(nn.Module): @@ -74,23 +69,23 @@ def __init__(self, config): super().__init__() self.w1 = nn.Linear(config.vision_dim, 2 * config.vision_inter_dim, bias=False) self.w2 = nn.Linear(config.vision_inter_dim, config.vision_dim, bias=False) + self.act_fn = SiluAndMul() def forward(self, x: torch.Tensor) -> torch.Tensor: - gate, up = self.w1(x).chunk(2, dim=-1) - return self.w2(F.silu(gate) * up) + return self.w2(self.act_fn(self.w1(x))) class DeepseekV4VisionBlock(nn.Module): def __init__(self, config): super().__init__() - self.norm1 = DeepseekV4RMSNorm(config.vision_dim) + self.norm1 = RMSNorm(config.vision_dim, eps=1e-6, dtype=torch.float32) self.attn = DeepseekV4VisionAttention(config) - self.norm2 = DeepseekV4RMSNorm(config.vision_dim) + self.norm2 = RMSNorm(config.vision_dim, eps=1e-6, dtype=torch.float32) self.mlp = DeepseekV4VisionMLP(config) def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: - x = x + self.attn(self.norm1(x), cos, sin) - return x + self.mlp(self.norm2(x)) + residual = x + self.attn(self.norm1(x), cos, sin) + return residual + self.mlp(self.norm2(residual)) class DeepseekV4ViT(nn.Module): @@ -102,7 +97,7 @@ def __init__(self, config): self.rope_theta = config.vision_rope_theta self.patch_embed = DeepseekV4PatchEmbed(config) self.blocks = nn.ModuleList([DeepseekV4VisionBlock(config) for _ in range(config.vision_n_layers)]) - self.norm = DeepseekV4RMSNorm(config.vision_dim) + self.norm = RMSNorm(config.vision_dim, eps=1e-6, dtype=torch.float32) def forward(self, patches: torch.Tensor, n_vit_h: int, n_vit_w: int) -> torch.Tensor: x = self.patch_embed(patches) diff --git a/vllm_ascend/models/deepseek_v4/vl_model.py b/vllm_ascend/models/deepseek_v4/vl_model.py index b9844dd51aae..f963fe06ade2 100644 --- a/vllm_ascend/models/deepseek_v4/vl_model.py +++ b/vllm_ascend/models/deepseek_v4/vl_model.py @@ -59,6 +59,7 @@ class AscendDeepseekV4ForConditionalGeneration( """DeepSeek-V4 vision entry point using the Ascend text backbone.""" requires_raw_input_tokens = True + language_model_cls = AscendDeepseekV4ForCausalLM @classmethod def get_placeholder_str(cls, modality: str, i: int) -> str | None: @@ -104,7 +105,7 @@ def __init__(self, *, vllm_config, prefix: str = "") -> None: self.aligner.to(dtype=model_config.dtype) with self._mark_language_model(vllm_config): - self.language_model = AscendDeepseekV4ForCausalLM( + self.language_model = self.language_model_cls( vllm_config=vllm_config, prefix=maybe_prefix(prefix, "language_model"), ) diff --git a/vllm_ascend/models/deepseek_v41/README.md b/vllm_ascend/models/deepseek_v41/README.md new file mode 100644 index 000000000000..854c50a9b556 --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/README.md @@ -0,0 +1,218 @@ +# DeepSeek V4.1 eager bring-up + +The V4.1 backbone is registered through `model.py`, matching the package +structure used by DeepSeek V4. There is no separate `modeling.py` and no +model-side KV-cache owner. + +## Current execution path + +- `model.py` owns the V4.1 backbone, delayed mHC handoff, attention projections, + source references, compressor invocation and the correctness-first eager + attention path. +- `attention/dsa_v41.py` owns paged-cache registration and metadata plus the + unfused SWA/long-KV gather, attention and cache-scatter operations. +- `core/deepseek_v41.py` owns cache specs, hybrid grouping, sizing, allocation + and reshape. +- `compressor.py` owns ratio1/ratio2 compressor parameters and the ratio2 + FP32 state cache. `indexer.py` owns Index-K cache updates, Indexer scoring, + candidate-block filtering and chronological TopK selection. + +The model reuses DeepSeek V4's quantization-aware projection, MoE and output +projection implementations. Small operators replace the fused DSA kernel for +the initial eager milestone. + +## Hybrid cache layout + +The allocator follows the DeepSeek V4 layer-tuple pattern with one global +block-ID lifecycle and four separate layer-outermost buffers. Each buffer +contains all blocks for one shared slot; each block contains the owning +group's resource tuple. Different groups own distinct simultaneous live IDs. +Within the merged group, KV and index share an ID at disjoint byte offsets. + +| Group | Resources | Logical block size | Slots | +| --- | --- | --- | --- | +| G0 | C2 KV/index at layers 2, 8, 14; C1 KV/index at layer 20 | 128 | 0-3 | +| G1 | FP32 circular compressor state at layers 2, 8, 14 | 32 | 0-2; slot 3 unused | +| G2-G11 | SWA layers 0-39, four consecutive layers per group, window 128 | 128 | 0-3 | + +Without DSpark there are 12 groups and 51 cache specs. The base attention block size is +128 at production dimensions. C2 stores 64 compressed rows per logical block; +C1 and SWA store 128 rows. State stores 32 uncompressed FP32 rows with width +1024 in one private ring page per request. +C2 and C1 share the original-token block table, with compression applied by +per-layer metadata builders. No paired-block mapping is needed. + +For MLA width 512, index width 128 and one KV head: + +| Slot | Full-context tuple | Page bytes | Component offsets (bytes) | +| --- | --- | --- | --- | +| 0-2 | C2 KV + INT8 index K + FP16 scales | 131072 | KV 0; K 65536; scale 73728 | +| 3 | C1 KV + INT8 index K + FP16 scales | 147712 | KV 0; K 131072; scale 147456 | + +C2 KV has 65536 payload bytes. Its index spec is padded from 8320 to 65536 +bytes, placing 57216 unused bytes after the scales. SWA/state occupy offset +zero and are padded to their assigned slot capacity. Thirty SWA layers have +no padding; the ten layers in slot 3 have 16640 padding bytes per page. Each +state spec uses all 131072 bytes as 32 ring rows; G1 leaves slot 3 +reserved but unused. Padding is applied to cloned specs and is idempotent. + +With `N` global IDs (including the reserved null ID), the raw uint8 buffers +have sizes `N * [131072, 131072, 131072, 147712]`. Allocation, startup admission, +maximum-length sizing and concurrency use their sum: **540928 bytes = +528.25 KiB per global ID**. Merged C1/C2 block demand is counted once. + +Typed zero-copy views retain the slot's physical page stride, not the +component's page size. C2 KV is `[N,64,1,512]` BF16; index K/scales are +`[N,64,1,128]` INT8 and `[N,64,1,1]` FP16. C1 uses the same widths with 128 +rows. SWA is `[N,128,1,512]` BF16; state is `[N,32,1,1024]` FP32. Components +other than state need not be contiguous. State must fill its slot contiguously; +no whole-context gather is introduced by allocation. + +### DSpark in the same four slots + +Aurora DSpark uses `DeepseekV41CacheBackend` and the V4.1 execution path for +both DSA_CP settings. Noncausal draft queries pass explicit physical SWA +indices to SparseFlashMla with mask mode 0. CP slices these global indices +and preserves the full visible KV length for each local request. Context KV +writes use the same stride-aware cache scatter as the target model. The V1 +proposer remains eager; this routing does not enable draft graph capture. + +The optional Aurora DSpark model adds one group, G12, containing exactly three +`DeepseekV41DraftSWASpec` resources: `mtp.0.self_attn.swa_cache`, +`mtp.1.self_attn.swa_cache`, and `mtp.2.self_attn.swa_cache`. They occupy offset +zero in slots 0, 1, and 2 respectively; G12 leaves slot 3 unused. The target +groups and their padding remain unchanged. This is **13 groups, 54 specs and +four physical buffers**, still **540928 bytes per global ID**. + +Each draft view is `[N,128,1,512]` BF16 with 131072-byte block stride and no +padding. The planner validates draft geometry against target SWA and rejects +foreign resources, compressed drafts, extra draft layers, and any geometry +that would enlarge the existing slots. An explicit draft spec keeps G12 +separate from target SWA while reusing its `SlidingWindowManager` semantics. + +G12 owns its own block table and live global IDs. All three draft layers use +that table, accessing different physical slots at the same ID. Target groups +use other live IDs, so sharing the backing does not share live target KV data. +Release/preemption returns IDs to the common pool. G12 adds one group's SWA +page demand, not three groups' demand; available-memory sizing still divides +by 540928, and rank shrinking changes only N. + +DSpark context KV is projected independently for each draft layer using the +inherited DSV4 SWA backend and the group's own slot mappings. The target exports +the incoming residual streams from the checkpoint-selected auxiliary layers. +Composite-config selection reads Aurora's text config. Target and draft MoE +dispatchers are selected by expert/execution shape to avoid sharing mutable +dispatch state across incompatible expert counts. With DSpark, `auto` cache +dtype is resolved to BF16 before constructing the inherited DSV4 draft backend. + +The 32-row FP32 target ring requires **1..31 speculative tokens**. A verifier +writes the anchor plus S speculative input rows. After accepting A drafts, the +next forward starts at `P+A+1`; if that position is odd it needs row `P+A`. +At most S newer rows follow it, so S below 32 preserves it through the tail +write. S=32 can overwrite the anchor after complete rejection and is rejected +at initialization. Per-batch limits cannot exceed the configured maximum. +Rejected compressed KV/index rows remain outside the accepted sequence length +and are overwritten when those positions are recomputed. + +Target eager and `FULL_DECODE_ONLY` modes retain their existing dispatch; +the V1 DSpark proposer runs eagerly. Draft graph capture is not enabled. + +### Earlier design comparisons + +The original block-outermost implementation reserved 393216 bytes per ID +across 17 groups with separate C1/C2 groups, both using logical block 128. +The unimplemented comparison design kept C2 logical block 256 (128 stored +rows), C1 block 128 and 17 groups sharing three 147712-byte slots: 432.75 KiB +per ID. The new design reserves 528.25 KiB per ID but reduces the number of +full-context and SWA-group IDs. Compare memory for the same request workload, +including state retention and free/null IDs, rather than comparing the +per-ID divisor alone. Operator and serving performance remain unmeasured. + +Index source layers 24, 28, 32 and 36 compute new selections in the reference +architecture but do not own another copy of the long KV or Index K. Candidate +blocks originate at layer 20. Consumers retain the source prefix and retrieve +the source cache from `static_forward_context`; shared modules are never +re-registered under consumer layers. + +## Supported milestone and remaining accuracy work + +The runtime contract is model runner V1, eager or `FULL_DECODE_ONLY` mode, +BF16 cache, hybrid KV management, PP/DCP/PCP equal to one, and +tensor/data/expert parallel serving. `FULL_DECODE_ONLY` retains Aurora main's +eager prefill and full-graph decode dispatch. DSpark is the only supported +speculative method, subject to the retention bound above. Prefix caching is +supported for cacheable attention groups; the circular compressor state remains +request-local and is excluded from prefix-cache hits. KV transfer and other +graph modes fail closed. + +The fallback attends over local SWA plus the compressed rows selected by the +Indexer/Candidate path. Engram execution is intentionally disabled: its two +roughly 196 GB embedding tables require a distributed HBM layout, while the +temporary CPU/NFS mmap implementation was both prohibitively slow and +numerically unverified. Full accuracy still requires HBM-sharded Engram at +layers 1 and 14 and reference FP8/FP4 rounding. These omissions must not be +interpreted as full model accuracy. + +## Circular compressor integration + +The v0.27.1 compatibility layer registers a circular spec and manager inside +vLLM-Ascend. G1 owns one global ID per request, retained until finish or +preemption. Its block table has one column; ordinary position-to-page slot +mapping is disabled. C2 uses the original global ID and `position % 32`. +The ring's 128-KiB contiguous page matches the shared-slot stride without +block-ID expansion or copying the cache. Full C1/C2 and SWA groups keep their +existing address calculations. Thirty SWA layers remain unpadded and ten +retain 16.25 KiB padding per page. + +C2 retains the existing FP32 projection weights and computation. The projected +Triton entry point pools current-chunk rows plus prior ring residuals before +updating the last 32 rows of the ring. State remains FP32, including values +that BF16 would round away. The pooled output is BF16 and passes through the +existing model RMSNorm, preserving its epsilon and rounding order. C1 and the +standalone Triton compressor API retain their paths. + +Completed pairs occupy their completion-token output rows, matching existing +long-KV/index slot mappings and source RoPE metadata. Per-source output buffers +are allocated before memory profiling. Device metadata uses persistent builder +buffers for graph replay; inactive requests have zero lengths and IDs. Dummy +capture requests receive distinct non-null ring IDs; idle DP synchronization +runs skip ring state updates. Model runner V1 supports eager prefill and +`FULL_DECODE_ONLY` dispatch, with runtime correctness still unverified. + +G1 admission now reserves one ID regardless of sequence length. Slot capacity +and proportional rank shrinking are unchanged: total backing bytes remain +`N * 540928`. Prefix caching reuses cacheable attention groups while keeping the +circular scratch state request-local; it does not enable KV transfer, V2, or +unsupported parallel modes. + +## Validation + +For the earlier block-outermost implementation, on the A3 remote container +the W8A8 checkpoint loads under TP4/DP4/EP, the +service becomes healthy, and greedy eager smoke requests return coherent +answers (`2+2 -> 4`, Chinese capital question -> Beijing). The focused cache and +mHC suite passes 30 tests. Formal task-level and long-context accuracy remain +follow-up gates. + +The merged four-slot implementation has local static validation only; +neither eager nor full-graph decode correctness has been verified for it. +Allocator/worker/metadata tests and slot-backed QLI/SparseFlashMla tests are +provided, but their torch/NPU execution is deferred. The earlier serving +results above do not validate this layout. Remote synchronization, builds, +correctness verification and performance measurements require a later run. + +Circular-ring changes have **local static checks only**. Run the dependency-free +`python3 tests/check_aurora_ring_static.py` to check placement and ring arithmetic. +The circular manager, metadata, and `test_deepseek_v41_ring_compressor.py` tests +are authored but have not been executed with torch/NPU. Numerical correctness, +full-decode graph replay, model serving, and performance remain **not verified**. +No remote synchronization, builds, or device execution were performed for this +change. Remote Run Manifest evidence belongs to a later requested phase. + +DSpark coverage adds dependency-free grouping/allocation checks and 5797 +acceptance/rejection schedules. Torch tests cover shared storage, group-ID +isolation, rank shrinking, configuration/auxiliary-state wiring and per-group +draft metadata. NPU tests cover draft context stores in the actual shared +backings and ring residuals after rejection. These torch/NPU tests are authored +but **not executed**; DSpark correctness, target graph replay and performance +remain **not verified**. diff --git a/vllm_ascend/models/deepseek_v41/__init__.py b/vllm_ascend/models/deepseek_v41/__init__.py new file mode 100644 index 000000000000..19452237b2e5 --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/__init__.py @@ -0,0 +1,3 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""DeepSeek V4.1 construction components (not yet a runnable model).""" diff --git a/vllm_ascend/models/deepseek_v41/compressor.py b/vllm_ascend/models/deepseek_v41/compressor.py new file mode 100644 index 000000000000..66465bd7ce6b --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/compressor.py @@ -0,0 +1,130 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""FP32 C2 ring compressor, ratio-1 path, and fused RMS normalization.""" + +from typing import Any + +import torch +import torch_npu +from torch import nn + +from vllm_ascend.attention.dsa_v41 import DeepseekV41CacheLayer +from vllm_ascend.core.deepseek_v41 import STATE_RING_ROWS, DeepseekV41CompressorStateSpec + + +def _read(config: Any, name: str) -> Any: + if isinstance(config, dict): + try: + return config[name] + except KeyError as exc: + raise ValueError(f"DeepSeek V4.1 config is missing {name!r}") from exc + try: + return getattr(config, name) + except AttributeError as exc: + raise ValueError(f"DeepSeek V4.1 config is missing {name!r}") from exc + + +def text_config_of(config: Any) -> Any: + if isinstance(config, dict): + return config.get("text_config", config) + return getattr(config, "text_config", config) + + +class DeepseekV41CompressorStateCache(DeepseekV41CacheLayer): + """State-cache module owning one packed FP32 circular page per request. + + Pass kv_cache[0].squeeze(-2) and the state's block table to the compressor. + The V4 constructor itself cannot be reused: it asserts ratio in (4, 128). + """ + + def __init__(self, vllm_config, prefix, spec): + if spec.dtype != torch.float32 or spec.compress_ratio != 1 or spec.block_size != STATE_RING_ROWS: + raise ValueError("V4.1 compressor state requires a 32-row FP32 ring") + super().__init__(vllm_config, prefix, spec) + self.state_dim = spec.head_size + self.dtype = spec.dtype + self.compress_ratio = 2 # Pooling ratio; spec storage ratio remains one. + self.block_size = spec.block_size + + +class DeepseekV41RMSNorm(nn.Module): + def __init__(self, width, eps): + super().__init__() + self.weight = nn.Parameter(torch.ones(width, dtype=torch.bfloat16)) + self.eps = eps + + def forward(self, x): + return torch_npu.npu_rms_norm(x, self.weight, epsilon=self.eps)[0] + + +class DeepseekV41Compressor(nn.Module): + def __init__(self, config, ratio, vllm_config=None, prefix="compressor"): + super().__init__() + if ratio not in (1, 2): + raise ValueError("V4.1 compressor requires ratio 1 or 2") + self.ratio = ratio + self.width = _read(config, "head_dim") + dim = _read(config, "hidden_size") + self.wkv = nn.Linear(dim, self.width, bias=False, dtype=torch.float32 if ratio == 2 else torch.bfloat16) + self.norm = DeepseekV41RMSNorm(self.width, _read(config, "rms_norm_eps")) + if ratio == 2: + self.wgate = nn.Linear(dim, self.width, bias=False, dtype=torch.float32) + # Allocate persistent output before memory profiling, so its footprint + # is included in the cache budget rather than added after allocation. + if vllm_config is not None: + capacity = getattr(vllm_config.scheduler_config, "max_num_batched_tokens", 4096) + self.register_buffer( + "_ring_pooled", + torch.empty(capacity, self.width, dtype=torch.bfloat16, device=self.wkv.weight.device), + persistent=False, + ) + # Standalone unfused-reference tests may supply pages explicitly. + if vllm_config is not None: + self.state_cache = DeepseekV41CompressorStateCache( + vllm_config, + f"{prefix}.state_cache", + DeepseekV41CompressorStateSpec( + block_size=STATE_RING_ROWS, + num_kv_heads=1, + head_size=2 * self.width, + dtype=torch.float32, + ), + ) + + def prepare_ring_compressor(self, max_tokens, device): + """Check the profiled per-source buffer and resolve hardware before capture.""" + from vllm_ascend.ops.triton.compressor.compressor_triton import _cube_core_num + + actual_device = self._ring_pooled.device + compatible_device = actual_device.type == device.type and ( + device.index is None or actual_device.index == device.index + ) + if self._ring_pooled.shape[0] < max_tokens or not compatible_device: + raise ValueError("Ring output capacity/device must be established before memory profiling") + self._ring_num_cores = _cube_core_num() + + def pool_projected(self, kv, scores, metadata): + from vllm_ascend.ops.triton.compressor.compressor_triton import compressor_from_projected + + if not hasattr(self, "_ring_pooled") or not hasattr(self, "_ring_num_cores"): + raise RuntimeError("Ring compressor must be initialized before graph capture") + if kv.shape[0] > self._ring_pooled.shape[0]: + raise ValueError("Compressor batch exceeds its prepared output capacity") + pooled = compressor_from_projected( + kv, + scores, + self.state_cache.kv_cache[0].squeeze(-2), + metadata.c2_ring_metadata, + self._ring_pooled[: kv.shape[0]], + max_query_len=metadata.max_query_len, + num_cores=self._ring_num_cores, + ) + return self.norm(pooled) + + def forward(self, x): + """Project an uncompressed source; ratio-2 uses ``pool_projected``.""" + if x.ndim != 2: + raise ValueError("Expected [tokens, hidden] input") + if self.ratio != 1: + raise RuntimeError("Ratio-2 compression must use pool_projected") + return self.norm(self.wkv(x)) diff --git a/vllm_ascend/models/deepseek_v41/dspark.py b/vllm_ascend/models/deepseek_v41/dspark.py new file mode 100644 index 000000000000..2410ecf57d3c --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/dspark.py @@ -0,0 +1,242 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Aurora / DeepSeek-V4.1 dSPark draft model for Ascend.""" + +import torch +import vllm.envs as envs +from vllm.compilation.decorators import support_torch_compile +from vllm.forward_context import get_forward_context, is_forward_context_available +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ColumnParallelLinear +from vllm.model_executor.layers.vocab_parallel_embedding import VocabParallelEmbedding +from vllm.model_executor.models.utils import maybe_prefix + +from vllm_ascend.attention.context_parallel.dsa_v41_cp import get_v41_cp_classes +from vllm_ascend.attention.dsa_v41 import DeepseekV41CacheBackend, scatter_cache_sk +from vllm_ascend.core.deepseek_v41 import DeepseekV41DraftSWASpec, validate_cache_runtime +from vllm_ascend.models.common.ops.sequence_parallel import ( + sp_all_gather, + sp_padding_mask, + sp_shard, +) +from vllm_ascend.models.deepseek_v4.dspark import ( + DeepseekV4DSparkModel, + DSparkConfidenceHead, + DSparkDeepseekV4ForCausalLM, + DSparkMarkovHead, + _get_dspark_num_mtp_layers, +) +from vllm_ascend.models.deepseek_v4.model import AscendDeepseekV4SWACache, DeepseekV4Attention +from vllm_ascend.models.deepseek_v41.model import ( + DeepseekV41Attention, + DeepseekV41DecoderLayer, + DeepseekV41LayerRole, +) + + +class DeepseekV41DSparkSWACache(AscendDeepseekV4SWACache): + def get_kv_cache_spec(self, vllm_config): + spec = super().get_kv_cache_spec(vllm_config) + return DeepseekV41DraftSWASpec( + block_size=spec.block_size, + num_kv_heads=spec.num_kv_heads, + head_size=spec.head_size, + dtype=spec.dtype, + sliding_window=spec.sliding_window, + cache_dtype_str=spec.cache_dtype_str, + model_version=spec.model_version, + ) + + def get_attn_backend(self): + return DeepseekV41CacheBackend + + +class DeepseekV41DSparkAttention(DeepseekV4Attention): + swa_cache_cls = DeepseekV41DSparkSWACache + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if self.compress_ratio != 0: + raise ValueError("Aurora DSpark supports only uncompressed draft SWA layers") + # V4.1 applies Q LoRA RMSNorm only, without a second per-head Q norm. + self.dsa_attn.dsa_attn.impl.apply_q_norm = False + self.softmax_scale = self.scale + self.shared_state = None + prefix = kwargs["prefix"] + self.v41_impl = get_v41_cp_classes()[1]( + prefix=prefix, + role=DeepseekV41LayerRole( + layer_idx=int(prefix.split(".")[-2]), + compress_ratio=0, + kv_source_layer=None, + index_source_layer=None, + is_kv_source=False, + is_index_source=False, + is_candidate_source=False, + uses_candidate_filter=False, + engram_slot=None, + ), + topology=None, + long_kv_source_prefix=None, + index_k_source_prefix=None, + ) + self.v41_layer_name = f"{prefix}.v41_attn" + context = kwargs["vllm_config"].compilation_config.static_forward_context + if self.v41_layer_name in context: + raise ValueError(f"Duplicate V4.1 attention layer: {self.v41_layer_name}") + context[self.v41_layer_name] = self + + forward = DeepseekV41Attention.forward + + +class DeepseekV41DSparkDecoderLayer(DeepseekV41DecoderLayer): + """V4.1 delayed-mHC block with a draft-only SWA attention backend.""" + + attention_cls = DeepseekV41DSparkAttention + + +class DeepseekV41DSparkModel(DeepseekV4DSparkModel): + """Three serial draft blocks matching the checkpoint's ``mtp.*`` tree.""" + + def __init__(self, *, vllm_config, prefix="") -> None: + # Deliberately do not call the V4 dSPark constructor: V4.1 has delayed + # mHC state between blocks and no terminal hc_head parameters. + torch.nn.Module.__init__(self) + assert vllm_config.speculative_config is not None + self.vllm_config = vllm_config + validate_cache_runtime(vllm_config) + draft_model_config = vllm_config.speculative_config.draft_model_config + config = draft_model_config.hf_text_config + self.config = config + self.hc_mult = config.hc_mult + self.hidden_size = config.hidden_size + self.block_size = int(config.dspark_block_size) + self.target_layer_ids = list(config.dspark_target_layer_ids) + self.num_dspark_layers = _get_dspark_num_mtp_layers(config) + if self.num_dspark_layers != 3: + raise ValueError("Aurora's DSpark cache group requires exactly three draft layers") + self.mtp_start_layer_idx = config.num_hidden_layers + self.use_sequence_parallel = vllm_config.parallel_config.use_sequence_parallel_moe + + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + quant_config=vllm_config.quant_config, + prefix=maybe_prefix(prefix, "embed_tokens"), + ) + self.layers = torch.nn.ModuleDict( + { + str(self.mtp_start_layer_idx + idx): DeepseekV41DSparkDecoderLayer( + vllm_config, + prefix=f"mtp.{idx}", + config=config, + is_draft_layer=True, + ) + for idx in range(self.num_dspark_layers) + } + ) + + self.needs_moe_input_ids = any( + layer.mlp.gate.tid2eid is not None or layer.mlp.gate.bias_vl is not None for layer in self.layers.values() + ) + first_layer = self.layers[str(self.mtp_start_layer_idx)] + self.main_proj = ColumnParallelLinear( + config.hidden_size * len(self.target_layer_ids), + config.hidden_size, + bias=False, + return_bias=False, + quant_config=None, # Aurora stores this projection in BF16. + prefix=maybe_prefix(prefix, f"layers.{self.mtp_start_layer_idx}.main_proj"), + gather_output=True, + ) + self.main_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + first_layer.main_proj = self.main_proj + first_layer.main_norm = self.main_norm + + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + last_layer_idx = self.mtp_start_layer_idx + self.num_dspark_layers - 1 + self.markov_head = DSparkMarkovHead(config, maybe_prefix(prefix, f"layers.{last_layer_idx}.markov_head")) + self.confidence_head = DSparkConfidenceHead(config, maybe_prefix(prefix, "confidence_head")) + last_layer = self.layers[str(last_layer_idx)] + last_layer.norm = self.norm + last_layer.markov_head = self.markov_head + + def _store_standard_swa_kv(self, shared_kv, slot_mapping, attn=None): + if slot_mapping is None or slot_mapping.numel() == 0: + return + cache = attn.dsa_attn.swa_cache_layer + if slot_mapping.ndim == 1: + valid = slot_mapping >= 0 + physical = slot_mapping.clamp_min(0) + slot_mapping = torch.stack((physical // cache.block_size, physical % cache.block_size), dim=-1).to( + torch.int32 + ) + slot_mapping.masked_fill_(~valid.unsqueeze(-1), -1) + scatter_cache_sk(cache.kv_cache[0], slot_mapping, shared_kv.squeeze(1)) + + def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: + hidden_states = self.embed_tokens(input_ids).unsqueeze(-2).repeat(1, self.hc_mult, 1) + full_num_tokens = positions.shape[0] + if self.use_sequence_parallel: + if envs.VLLM_MOE_SKIP_PADDING and is_forward_context_available(): + forward_context = get_forward_context() + forward_context.is_padding = sp_padding_mask( + forward_context.is_padding, + hidden_states, + ) + hidden_states = sp_shard(hidden_states) + input_ids = sp_shard(input_ids) + pre_mix = hidden_states.new_zeros(hidden_states.shape[0], self.hc_mult, dtype=torch.float32) + pre_mix[:, 0] = 1.0 + last_layer = None + moe_input_ids = input_ids + if self.needs_moe_input_ids: + moe_input_ids = torch.where(input_ids == -1, 0, input_ids) + for layer in self.layers.values(): + last_layer = layer + hidden_states, pre_mix = layer( + positions, + hidden_states, + pre_mix, + llama_4_scaling=None, + input_ids=moe_input_ids, + ) + assert last_layer is not None + hidden_states = last_layer.hc_collapse(hidden_states, pre_mix) + if self.use_sequence_parallel: + hidden_states = sp_all_gather(hidden_states)[:full_num_tokens] + return hidden_states + + +@support_torch_compile +class DSparkDeepseekV41ForCausalLM(DSparkDeepseekV4ForCausalLM): + def __init__(self, *, vllm_config, prefix="") -> None: + torch.nn.Module.__init__(self) + assert vllm_config.speculative_config is not None + self.config = vllm_config.speculative_config.draft_model_config.hf_text_config + + from vllm_ascend.utils import get_rotation_path + + self.rotation_path = get_rotation_path(vllm_config) if vllm_config.quant_config is not None else None + self.model = DeepseekV41DSparkModel(vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")) + from vllm.model_executor.layers.logits_processor import LogitsProcessor + from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead + + self.lm_head = ParallelLMHead( + self.config.vocab_size, + self.config.hidden_size, + prefix=maybe_prefix(prefix, "lm_head"), + ) + self.logits_processor = LogitsProcessor(self.config.vocab_size) + self.set_moe_parameters() + + def _remap_dspark_name(self, name: str) -> str | None: + mapped = super()._remap_dspark_name(name) + if mapped is None: + return None + # Aurora names the low-rank Markov matrices after their operations, + # while the runtime uses explicit embedding/projection parameter names. + mapped = mapped.replace(".markov_head.embed.weight", ".markov_head.markov_w1.weight") + mapped = mapped.replace(".markov_head.head.weight", ".markov_head.markov_w2.weight") + mapped = mapped.replace("model.confidence_head.weight", "model.confidence_head.proj.weight") + return mapped diff --git a/vllm_ascend/models/deepseek_v41/engram_gate.py b/vllm_ascend/models/deepseek_v41/engram_gate.py new file mode 100644 index 000000000000..0fb2573139e8 --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/engram_gate.py @@ -0,0 +1,30 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import torch + + +def engram_gate( + hidden: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + channel_weight: torch.Tensor, + rotation_block: torch.Tensor, + token_mask: torch.Tensor, + eps: float, +) -> torch.Tensor: + """Apply original-basis gating to a rotated residual and rotated value. + + ``hidden`` and ``key`` have shape [tokens, hc_mult, hidden_size]. + The saved rotation consists of identical diagonal blocks. Restore hidden + in FP32; the value projection already includes the forward rotation. + """ + dim = hidden.shape[-1] + original = (hidden.float().unflatten(-1, (-1, rotation_block.shape[0])) @ rotation_block.float().T).flatten(-2) + key = key.float() + rstd = torch.rsqrt(original.square().mean(-1) + eps) + rstd *= torch.rsqrt(key.square().mean(-1) + eps) + dot = (original * channel_weight.float() * key).sum(-1) * rstd * dim**-0.5 + magnitude = dot.abs().clamp_min(1e-6).sqrt() + gate = torch.sigmoid(torch.where(torch.signbit(dot), -magnitude, magnitude)) + gate = gate.masked_fill(~token_mask.unsqueeze(-1), 0) + return (hidden.float() + gate.unsqueeze(-1) * value.float().unsqueeze(-2)).to(hidden.dtype) diff --git a/vllm_ascend/models/deepseek_v41/engram_hash.py b/vllm_ascend/models/deepseek_v41/engram_hash.py new file mode 100644 index 000000000000..2665ece1bfeb --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/engram_hash.py @@ -0,0 +1,277 @@ +# SPDX-License-Identifier: MIT +# Adapted from the DeepSeek V4.1 reference inference/engram.py. +from dataclasses import dataclass + +import numpy as np +import torch +from sympy import isprime # type: ignore[import-untyped] + +_HISTORY_SLAB_MIN_TOKENS = 16 +_PAGE_WRITE_NUMPY_MIN_TOKENS = 16 + + +def engram_history_metadata(metadata): + """Read full-request SWA pages, before attention's local CP slicing. + + V4.1 CP exposes local queries directly and keeps the replicated request + in global_metadata. Engram runs before token slicing, so all TP ranks + must use the full request lengths. + """ + request_metadata = getattr(metadata, "global_metadata", None) + if request_metadata is None: + request_metadata = metadata + boundaries = getattr(request_metadata, "query_start_loc_cpu", None) + block_table = getattr(request_metadata, "block_table_cpu", None) + if boundaries is None or block_table is None: + raise ValueError("Engram requires query_start_loc_cpu and block_table_cpu in request metadata") + if boundaries.device.type != "cpu" or block_table.device.type != "cpu": + raise ValueError("Engram request metadata mirrors must reside on CPU") + return boundaries.long(), block_table, request_metadata.storage_block_size + + +def valid_engram_token_mask( + input_ids: torch.Tensor, + image_token_id: int, + image_pad_token_id: int, +) -> torch.Tensor: + """Exclude the complete V4.1 image region from n-gram history.""" + return (input_ids != image_token_id) & (input_ids != image_pad_token_id) + + +def find_next_prime(start: int, seen_primes: set[int]) -> int: + """The smallest prime above `start` that has not been handed out yet.""" + candidate = start + 1 + while not isprime(candidate) or candidate in seen_primes: + candidate += 1 + return candidate + + +def build_compressed_token_map(tokenizer) -> tuple[list[int], int]: + """Map every token id onto a smaller id space where tokens that normalize alike + collapse together. + + N-grams are hashed over these compressed ids, so " The", "the" and "THE" all hash + the same way. + Returns the lookup plus the size of the compressed vocab -- and that size matters + beyond bounds + checking, because every hash multiplier is derived from it. + """ + from tokenizers import Regex, normalizers # type: ignore[import-untyped] + + # a private-use char, so a token that is exactly one space survives Strip() instead + # of + # collapsing to the empty string and merging with unrelated tokens + sentinel = "\ue000" + normalizer = normalizers.Sequence( + [ + normalizers.NFKC(), + normalizers.NFD(), + normalizers.StripAccents(), + normalizers.Lowercase(), + normalizers.Replace(Regex(r"[ \t\r\n]+"), " "), + normalizers.Replace(Regex(r"^ $"), sentinel), + normalizers.Strip(), + normalizers.Replace(sentinel, " "), + ] + ) + + # the raw Rust tokenizer, matching what training decodes with (no + # clean_up_tokenization_spaces) + backend = tokenizer.backend_tokenizer + key_to_new: dict[str, int] = {} + lookup = [0] * len(tokenizer) + for token_id in range(len(tokenizer)): + text = backend.decode([token_id], skip_special_tokens=False) + if "\ufffd" in text: + # a partial UTF-8 byte token: nothing to normalize, so key it by its raw + # form + key = backend.id_to_token(token_id) + else: + normalized = normalizer.normalize_str(text) + key = normalized if normalized else text + + new_id = key_to_new.get(key) + if new_id is None: + new_id = len(key_to_new) + key_to_new[key] = new_id + lookup[token_id] = new_id + + return lookup, len(key_to_new) + + +def compute_hash_multipliers( + layer_ids: tuple[int, ...], max_ngram_size: int, tokenizer_vocab_size: int +) -> torch.Tensor: + """Derive one multiplier per (layer, lookback) from a per-layer RNG. + + Kept odd, and bounded so that `token_id * multiplier` cannot overflow int64. + """ + max_long = np.iinfo(np.int64).max + multiplier_bound = max(1, (max_long // tokenizer_vocab_size) // 2) + rows = [] + for layer_id in layer_ids: + generator = np.random.default_rng(10007 * layer_id) + values = generator.integers( + low=0, + high=multiplier_bound, + size=(max_ngram_size,), + dtype=np.int64, + ) + rows.append(torch.tensor(values * 2 + 1)) + return torch.stack(rows) + + +@dataclass(frozen=True) +class EngramLayout: + """Bucket layout of the n-gram hash tables. + + A position uses `max_ngram_size - 1` n-grams, each split over `n_heads`. + Each (n-gram size, head) pair owns its own prime-sized bucket range in the + layer's table; the primes are drawn in order and never reused, which keeps the + ranges disjoint. + """ + + max_ngram_size: int + layer_ids: tuple[int, ...] + num_embeddings: tuple[int, ...] # table rows, per engram layer + primes: tuple[tuple[tuple[int, ...], ...], ...] # [layer][n-gram size][head] bucket modulus + n_heads: int + head_dim: int + + @classmethod + def from_args(cls, args) -> "EngramLayout | None": + layer_ids = tuple(args.engram_layer_ids) + if not layer_ids: + return None + max_ngram_size, n_heads = args.engram_max_ngram_size, args.engram_n_heads + primes = [] + seen: set[int] = set() + for _ in layer_ids: + per_ngram = [] + for _ in range(max_ngram_size - 1): + sizes, current = [], args.engram_vocab_size - 1 + for _ in range(n_heads): + current = find_next_prime(current, seen) + seen.add(current) + sizes.append(current) + per_ngram.append(tuple(sizes)) + primes.append(tuple(per_ngram)) + return cls( + max_ngram_size=max_ngram_size, + layer_ids=layer_ids, + num_embeddings=tuple(args.engram_num_embeddings), + primes=tuple(primes), + n_heads=n_heads, + head_dim=args.engram_head_dim, + ) + + +class PagedNgramHistory: + """Mirror token IDs in the scheduler's physical pages, including prefixes. + + A speculative suffix can be overwritten without rolling back a mutable + per-request tail. Hashes read only positions at or before the current query. + CPU residency supplies routing metadata without a device synchronization. + """ + + def __init__(self, config, tokenizer): + layout = EngramLayout.from_args(config) + if layout is None: + raise ValueError("PagedNgramHistory requires at least one Engram layer") + token_map, vocab_size = build_compressed_token_map(tokenizer) + if vocab_size != config.engram_compressed_vocab_size: + raise ValueError(f"Engram compressed vocabulary mismatch: {vocab_size}") + self.token_map = torch.tensor(token_map, dtype=torch.int64) + self.pad_id = token_map[config.engram_pad_id] + self.image_token_id = config.image_token_id + self.image_pad_token_id = getattr( + config, + "image_pad_token_id", + self.image_token_id + 1, + ) + self.primes = torch.tensor(layout.primes) + sizes = self.primes.flatten(1) + self.offsets = sizes.cumsum(-1) - sizes + self.multipliers = compute_hash_multipliers(layout.layer_ids, layout.max_ngram_size, vocab_size) + self.lookback = layout.max_ngram_size + self.pages: dict[int, torch.Tensor] = {} + + def update(self, input_ids, positions, request_ids, block_table, block_size): + """All arguments are CPU tensors; page numbers come from full SWA KV.""" + if input_ids.numel() == 0: + # Idle DP and empty prefill still follow the collective contract, + # but there is no page or hash state to update. + columns = (self.lookback - 1) * self.primes.shape[-1] + return ( + torch.empty((0, self.primes.shape[0], columns), dtype=torch.int64, device="cpu"), + torch.empty(0, dtype=torch.bool, device="cpu"), + ) + compressed = self.token_map[input_ids] + mask = valid_engram_token_mask( + input_ids, + self.image_token_id, + self.image_pad_token_id, + ) + compressed = compressed.masked_fill(~mask, -1) + # Materialize CPU lists once for sequential page writes and small-batch + # history reads. + compressed_list = compressed.tolist() + position_list = positions.tolist() + page_indices = block_table[request_ids, positions // block_size].tolist() + if len(input_ids) < _PAGE_WRITE_NUMPY_MIN_TOKENS: + for token, position, page in zip(compressed_list, position_list, page_indices): + if page not in self.pages: + self.pages[page] = torch.full((block_size,), -1, dtype=torch.int64, device="cpu") + self.pages[page][position % block_size] = token + else: + page_views: dict[int, np.ndarray] = {} + for token, position, page in zip(compressed_list, position_list, page_indices): + view = page_views.get(page) + if view is None: + if page not in self.pages: + self.pages[page] = torch.full((block_size,), -1, dtype=torch.int64, device="cpu") + view = self.pages[page].numpy() + page_views[page] = view + # Zero-copy CPU view avoids Torch dispatch per scalar write. Keep + # input order so repeated physical slots retain last-write wins. + view[position % block_size] = token + history = torch.full((len(input_ids), self.lookback), self.pad_id, dtype=torch.int64, device="cpu") + if len(input_ids) < _HISTORY_SLAB_MIN_TOKENS: + # Slab construction dominates decode and small batches; retain the + # page-row loop for this latency-sensitive path. + for row, (position, request) in enumerate(zip(position_list, request_ids.tolist())): + for shift in range(self.lookback): + previous = position - shift + if previous < 0: + break + page = block_table[request, previous // block_size].item() + token = self.pages[page][previous % block_size] + if token < 0: + break + history[row, shift] = token + else: + active = torch.ones(len(input_ids), dtype=torch.bool, device="cpu") + for shift in range(self.lookback): + previous = positions - shift + valid = active & (previous >= 0) + if not bool(valid.any()): + break + with torch.device("cpu"): + rows = torch.nonzero(valid, as_tuple=False).flatten() + page_ids = block_table[request_ids[rows], previous[rows] // block_size] + offsets = previous[rows] % block_size + with torch.device("cpu"): + unique_pages, slab_indices = torch.unique(page_ids, return_inverse=True) + # Reachable pages must exist, just as in the row path. Inactive + # rows never read past an image or unwritten-token barrier. + slab = torch.stack([self.pages[page] for page in unique_pages.tolist()]) + values = slab[slab_indices, offsets] + present = values >= 0 + history[rows[present], shift] = values[present] + active[rows] = present + products = history[:, None] * self.multipliers + rolling, hashes = products[..., 0], [] + for shift in range(1, self.lookback): + rolling = torch.bitwise_xor(rolling, products[..., shift]) + hashes.append(rolling[..., None] % self.primes[:, shift - 1]) + return torch.cat(hashes, -1) + self.offsets, mask diff --git a/vllm_ascend/models/deepseek_v41/engram_hbm.py b/vllm_ascend/models/deepseek_v41/engram_hbm.py new file mode 100644 index 000000000000..0f1ddf5457ed --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/engram_hbm.py @@ -0,0 +1,492 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Node-local BF16/INT8 Engram storage, independent of model TP.""" + +import json +import socket +from collections import OrderedDict +from pathlib import Path + +import torch +import torch.distributed as dist +from safetensors import safe_open +from torch import nn + +_OFFLOAD_BUFFER_CACHE_SIZE = 8 +_OFFLOAD_BUFFER_BYTES_LIMIT = 512 * 1024 * 1024 +_BF16_BYTES = 2 + + +def quantize_engram_rows(rows): + """Group32 symmetric INT8 with FP32 power-of-two scales and ties-to-even.""" + grouped = rows.float().unflatten(-1, (-1, 32)) + maximum = grouped.abs().amax(-1, keepdim=True) + scale = torch.where(maximum == 0, torch.ones_like(maximum), maximum / 127) + # NPU exp2 can return one ULP below an exact power of two, changing + # ties-to-even codes. ldexp constructs the binary scale exactly. + exponent = torch.ceil(torch.log2(scale)) + scale = torch.where(torch.isfinite(exponent), torch.ldexp(torch.ones_like(scale), exponent.int()), scale) + codes = torch.round(grouped / scale).clamp(-127, 127).to(torch.int8).flatten(-2) + return codes, scale.squeeze(-1) + + +def dequantize_engram_rows(codes, scale): + # Keep one FP32 work buffer: in-place scaling avoids the extra FP32 result + # allocation created by the broadcast multiply expression. + decoded = codes.float().unflatten(-1, (-1, 32)) + decoded.mul_(scale.unsqueeze(-1)) + return decoded.flatten(-2).bfloat16() + + +def pack_engram_int8_rows(codes, scale): + """Pack INT8 codes and FP32 group scales as one row-oriented wire buffer.""" + if codes.dtype != torch.int8 or scale.dtype != torch.float32: + raise TypeError("Engram INT8 wire packing expects int8 codes and FP32 scales") + return torch.cat((codes.view(torch.uint8), scale.view(torch.uint8)), dim=-1).contiguous() + + +def unpack_engram_int8_rows(payload, width): + """Decode the packed INT8 wire buffer without changing BF16 lookup semantics.""" + groups = width // 32 + expected = width + groups * 4 + if payload.dtype != torch.uint8 or payload.shape[-1] != expected: + raise ValueError(f"Invalid Engram INT8 wire payload: {tuple(payload.shape)}") + codes = payload[..., :width].contiguous().view(torch.int8) + scale = payload[..., width:].contiguous().view(torch.float32).reshape(*payload.shape[:-1], groups) + return dequantize_engram_rows(codes, scale) + + +class EngramQueryGroup: + """One node group shared by all Engram layers; TP leaders submit queries. + + All ranks (including idle DP replicas) must call lookup in the same order. + Counts and all-to-all split sizes are eager metadata, not graph inputs. + """ + + def __init__(self, group, cpu_group, tp_group, tp_source): + self.group = group + self.cpu_group = cpu_group + self.tp_group = tp_group + self.tp_source = tp_source + # HCCL metadata avoids the CPU/Gloo rendezvous; standalone Gloo/MPI + # probes keep CPU metadata. Callers may override this for rollback. + backend = str(dist.get_backend(group)).lower() + self.metadata_on_device = backend not in ("gloo", "mpi") + self.rank = dist.get_rank(group) + self.size = dist.get_world_size(group) + self.is_source = dist.get_rank() == tp_source + + @classmethod + def from_vllm(cls, parallel): + # Lazy imports keep the transport usable in standalone distributed probes. + from vllm.distributed import get_ep_group, get_tp_group + + if ( + parallel.pipeline_parallel_size != 1 + or parallel.prefill_context_parallel_size != 1 + or parallel.decode_context_parallel_size != 1 + or not parallel.enable_expert_parallel + ): + raise ValueError("Engram HBM sharing requires EP and PP=PCP=DCP=1") + ep, tp = get_ep_group(), get_tp_group() + hosts = [None] * ep.world_size + dist.all_gather_object(hosts, socket.gethostname(), group=ep.cpu_group) + node_groups = [[ep.ranks[i] for i, host in enumerate(hosts) if host == name] for name in dict.fromkeys(hosts)] + node_sizes = {len(ranks) for ranks in node_groups} + if len(node_sizes) != 1: + raise ValueError(f"Engram requires equal rank counts on every node: {node_groups}") + selected = None + for ranks in node_groups: + # Every world rank creates groups in the same order. + cpu = dist.new_group(ranks, backend="gloo") + device = dist.new_group(ranks, backend=dist.get_backend(ep.device_group)) + if dist.get_rank() in ranks: + if not set(tp.ranks).issubset(ranks): + raise ValueError("Engram requires each TP group to stay within one node") + selected = cls(device, cpu, tp.device_group, tp.ranks[0]) + return selected + + +class NodeShardedEngram(nn.Module): + """Contiguous row shards; only BF16 rows cross the node-local fabric.""" + + def __init__(self, rows, width, query_group, device=None, storage_format="bf16"): + super().__init__() + if storage_format not in ("bf16", "int8", "fp8", "mxfp8"): + raise ValueError("Engram storage_format must be bf16, int8, fp8, or mxfp8") + if storage_format in ("int8", "fp8", "mxfp8") and width % 32: + raise ValueError("INT8 Engram requires a width divisible by 32") + self.storage_format = storage_format + self.rows, self.width = rows, width + self.query_group = query_group + # Kept opt-in while the reduced-payload protocol is benchmarked. + self.compressed_int8_wire = False + # The fused gather/dequant kernel wins even for one row on A3; retain + # the threshold as a local rollback knob for future kernel changes. + self.use_triton_int8 = True + self.triton_int8_min_rows = 1 + self.offload_pinned = storage_format in ("fp8", "mxfp8") + self._offload_buffers = OrderedDict() + self._offload_buffer_bytes = 0 + self._offload_buffer_bytes_limit = _OFFLOAD_BUFFER_BYTES_LIMIT + self._offload_buffer_index = {} + self._offload_events = {} + # Ceil partition leaves at most size-1 unused rows, never a replica. + self.shard_rows = (rows + query_group.size - 1) // query_group.size + self.start = query_group.rank * self.shard_rows + self.end = min(self.start + self.shard_rows, rows) + if self.start >= rows: + raise ValueError("Engram table must have at least one row per rank") + self._empty_flat = torch.empty(0, dtype=torch.int64, device="cpu") + # Reuse fixed-size HCCL metadata buffers across requests. + self._metadata_device_buffers = {} + self._empty_metadata = torch.zeros(query_group.size + 1, dtype=torch.int64, device="cpu") + self.weight = nn.Parameter( + torch.empty( + self.end - self.start, + width, + dtype=( + torch.int8 + if storage_format == "int8" + else (torch.float8_e4m3fn if storage_format in ("fp8", "mxfp8") else torch.bfloat16) + ), + device=(torch.device("cpu") if storage_format in ("fp8", "mxfp8") else device), + pin_memory=False, + ), + requires_grad=False, + ) + if storage_format == "int8": + self.register_buffer( + "weight_scale", torch.empty(self.end - self.start, width // 32, dtype=torch.float32, device=device) + ) + elif storage_format in ("fp8", "mxfp8"): + self.register_buffer( + "weight_scale", + torch.empty( + self.end - self.start, + width // 32, + dtype=torch.float8_e8m0fnu, + device="cpu", + pin_memory=False, + ), + ) + + def set_rows(self, start, rows): + """Load BF16 rows into local storage without allocating a BF16 table copy.""" + end = start + rows.shape[0] + if self.storage_format == "int8": + codes, scales = quantize_engram_rows(rows.to(self.weight.device)) + if not bool((torch.isfinite(scales) & (scales > 0)).all()): + raise ValueError("INT8 Engram requires finite positive group scales") + self.weight.data[start:end].copy_(codes) + self.weight_scale[start:end].copy_(scales) + else: + self.weight.data[start:end].copy_(rows) + + def lookup_local(self, ids, *, pin_output=False): + # Idle DP replicas still enter routing collectives, but must not launch + # gather/dequant kernels for an empty owner request. + if ids.numel() == 0: + return torch.empty((*ids.shape, self.width), dtype=torch.bfloat16, device=self.weight.device) + original_shape = ids.shape + flat_ids = ids.reshape(-1) + if self.storage_format == "int8": + if ( + self.use_triton_int8 + and self.weight.device.type == "npu" + and self.width == 256 + and flat_ids.device.type == "npu" + and flat_ids.shape[0] >= self.triton_int8_min_rows + ): + # Importing the ops package initializes the active Triton + # backend, so keep it out of CPU-only routing and test workers. + from vllm_ascend.ops.triton.engram_int8 import gather_dequantize_engram_int8 + + rows = gather_dequantize_engram_int8(self.weight, self.weight_scale, flat_ids) + else: + codes = torch.index_select(self.weight, 0, flat_ids) + scales = torch.index_select(self.weight_scale, 0, flat_ids) + rows = dequantize_engram_rows(codes, scales) + elif self.storage_format in ("fp8", "mxfp8"): + # index_select avoids the extra advanced-indexing wrapper on the + # CPU-resident PLE table and keeps row selection explicit. + rows = torch.index_select(self.weight, 0, flat_ids) + decoded = rows.float().reshape(-1, self.width // 32, 32) + scales = torch.index_select(self.weight_scale, 0, flat_ids) + decoded.mul_(scales.float().unsqueeze(-1)) + decoded = decoded.reshape(-1, self.width) + if self.offload_pinned and pin_output: + key = decoded.shape[0] + slots = self._offload_buffers.get(key) + if slots is None: + slot_bytes = 2 * key * self.width * _BF16_BYTES + # A single shape may be larger than the configured cache + # budget. It cannot be split without changing the lookup + # contract, so keep that one shape as an explicit bound + # exception after evicting all prior shapes. + while self._offload_buffers and ( + len(self._offload_buffers) >= _OFFLOAD_BUFFER_CACHE_SIZE + or self._offload_buffer_bytes + slot_bytes > self._offload_buffer_bytes_limit + ): + evicted, evicted_slots = self._offload_buffers.popitem(last=False) + self._offload_buffer_bytes -= 2 * evicted * self.width * _BF16_BYTES + self._offload_buffer_index.pop(evicted, None) + for slot in evicted_slots: + event = self._offload_events.pop(slot.data_ptr(), None) + if event is not None: + event.synchronize() + slots = [torch.empty((key, self.width), dtype=torch.bfloat16, pin_memory=True) for _ in range(2)] + self._offload_buffers[key] = slots + self._offload_buffer_bytes += slot_bytes + self._offload_buffer_index[key] = 0 + else: + self._offload_buffers.move_to_end(key) + index = self._offload_buffer_index[key] + decoded_slot = slots[index] + self._offload_buffer_index[key] = 1 - index + event = self._offload_events.pop(decoded_slot.data_ptr(), None) + if event is not None: + event.synchronize() + # Convert directly into the reusable pinned destination. The + # previous expression materialized a separate BF16 tensor + # before this copy, doubling the temporary decoded allocation. + decoded_slot.copy_(decoded) + rows = decoded_slot + else: + rows = decoded.bfloat16() + else: + rows = torch.index_select(self.weight, 0, flat_ids) + return rows.view(*original_shape, self.width) + + def _record_offload_use(self, source_ptr, device): + """Keep a pinned staging slot alive until the submitted device work ends.""" + if not self.offload_pinned or device.type != "npu": + return + event = torch.npu.Event() + event.record(torch.npu.current_stream(device)) + self._offload_events[source_ptr] = event + + def load_checkpoint(self, model_path, key, chunk_rows=65536): + """Load BF16, INT8, FP8, or MXFP8 Engram tensors with bounded IO. + + FP8/MXFP8 remain CPU resident (PLE_OFFLOAD); only decoded BF16 rows + enter the node-local all-to-all response buffer. + """ + root = Path(model_path) + index = json.loads((root / "quant_model_weights.safetensors.index.json").read_text())["weight_map"] + scale_key = key.removesuffix(".weight") + ".scale" + with safe_open(root / index[key], framework="pt", device="cpu") as file: + tensor = file.get_slice(key) + if tensor.get_shape() != [self.rows, self.width]: + raise ValueError(f"{key}: expected BF16/FP8 [{self.rows}, {self.width}]") + source_dtype = tensor.get_dtype() + if self.storage_format == "int8" and source_dtype in ("I8", "INT8"): + if scale_key not in index: + raise ValueError(f"{key}: INT8 source requires .scale") + with safe_open(root / index[scale_key], framework="pt", device="cpu") as sf: + scale = sf.get_slice(scale_key) + if scale.get_shape() != [self.rows, self.width // 32] or scale.get_dtype() != "F32": + raise ValueError(f"{scale_key}: expected FP32 [{self.rows}, {self.width // 32}]") + for start in range(self.start, self.end, chunk_rows): + stop = min(start + chunk_rows, self.end) + self.weight.data[start - self.start : stop - self.start].copy_(tensor[start:stop]) + self.weight_scale[start - self.start : stop - self.start].copy_(scale[start:stop]) + return + if self.storage_format in ("fp8", "mxfp8"): + if source_dtype not in ("F8_E4M3", "F8_E4M3FN") or scale_key not in index: + raise ValueError(f"{key}: {self.storage_format} requires FP8 weight and .scale") + with safe_open(root / index[scale_key], framework="pt", device="cpu") as sf: + scale = sf.get_slice(scale_key) + if scale.get_shape() != [self.rows, self.width // 32]: + raise ValueError(f"{scale_key}: expected [{self.rows}, {self.width // 32}]") + for start in range(self.start, self.end, chunk_rows): + stop = min(start + chunk_rows, self.end) + self.weight.data[start - self.start : stop - self.start].copy_(tensor[start:stop]) + self.weight_scale[start - self.start : stop - self.start].copy_(scale[start:stop]) + return + if source_dtype != "BF16": + raise ValueError(f"{key}: expected BF16 source for {self.storage_format}") + for start in range(self.start, self.end, chunk_rows): + stop = min(start + chunk_rows, self.end) + self.set_rows(start - self.start, tensor[start:stop]) + + def _metadata(self, ids): + q = self.query_group + flat = ids.reshape(-1) if q.is_source else ids.new_empty(0) + if flat.numel() == 0: + return self._empty_flat, self._empty_flat, self._empty_metadata + invalid = bool(flat.min() < 0 or flat.max() >= self.rows) + owners = flat.clamp(0, self.rows - 1) // self.shard_rows if invalid else flat // self.shard_rows + # Only owner grouping is required; preserving equal-owner order adds + # avoidable CPU sort work because the same permutation restores rows. + order = owners.argsort(stable=False) + counts = torch.bincount(owners, minlength=q.size) + metadata = torch.empty(q.size + 1, dtype=torch.int64, device="cpu") + metadata[:-1].copy_(counts) + metadata[-1] = int(invalid) + return flat, order, metadata + + @torch.inference_mode() + def _forward_with_gathered(self, ids, gathered, routing=None, broadcast=True, output=None): + """Return ids.shape + [width], bit-preserving, even when a DP is idle. + + IDs reside on CPU; hashing/history already runs at the eager boundary. + Exactly one TP rank submits the DP's queries. All owners serve requests; + reverse all-to-all restores requester order before the TP broadcast. + """ + q = self.query_group + if ids.device.type != "cpu" or ids.dtype != torch.int64: + raise ValueError("Engram routing expects CPU int64 IDs") + if routing is None: + flat, order, metadata = self._metadata(ids) + else: + flat, order, metadata = routing + counts = [row.tolist() for row in gathered] + if any(row[-1] for row in counts): + raise IndexError("Engram hash ID outside table") + send = metadata[:-1].tolist() + recv = [row[q.rank] for row in counts] + total_recv = sum(recv) + total_requests = sum(send) + total_global = sum(sum(row[:-1]) for row in counts) + backend = str(dist.get_backend(q.group)).lower() + # HCCL collectives require NPU tensors even when the compressed table is CPU resident. + device = self.weight.device + if device.type == "cpu" and backend not in ("gloo", "mpi"): + device = torch.device("npu") + if device.type == "npu" and device.index is None: + device = torch.device("npu", torch.npu.current_device()) + # All-zero rounds are skipped identically on every rank (HCCL portability). + if total_global: + incoming = torch.empty(total_recv, dtype=torch.int64, device=device) + ordered_ids = torch.index_select(flat, 0, order).to(device) + dist.all_to_all_single(incoming, ordered_ids, recv, send, group=q.group) + local_ids = incoming - self.start + if self.storage_format == "int8" and self.compressed_int8_wire: + lookup_ids = local_ids.cpu() if self.weight.device.type == "cpu" else local_ids + codes = torch.index_select(self.weight, 0, lookup_ids) + scales = torch.index_select(self.weight_scale, 0, lookup_ids) + packed = pack_engram_int8_rows(codes.to(device), scales.to(device)) + wire_width = self.width + (self.width // 32) * 4 + returned = torch.empty(total_requests * wire_width, dtype=torch.uint8, device=device) + wire_send = [count * wire_width for count in send] + wire_recv = [count * wire_width for count in recv] + dist.all_to_all_single(returned, packed.flatten(), wire_send, wire_recv, group=q.group) + returned = unpack_engram_int8_rows(returned.reshape(total_requests, wire_width), self.width) + else: + local_values = self.lookup_local( + local_ids.cpu() if self.weight.device.type == "cpu" else local_ids, + pin_output=device.type == "npu", + ) + source_ptr = local_values.data_ptr() + values = local_values.to( + device=device, dtype=torch.bfloat16, non_blocking=self.offload_pinned + ).contiguous() + returned = torch.empty((total_requests, self.width), dtype=torch.bfloat16, device=device) + dist.all_to_all_single(returned, values, send, recv, group=q.group) + self._record_offload_use(source_ptr, device) + else: + returned = torch.empty((0, self.width), dtype=torch.bfloat16, device=device) + if output is None: + result = torch.empty((ids.numel(), self.width), dtype=torch.bfloat16, device=device) + else: + result = output + if result.shape != (ids.numel(), self.width) or result.device != device: + raise ValueError("Engram output buffer has an incompatible shape or device") + if q.is_source: + result[order.to(device)] = returned + if broadcast and result.numel(): + dist.broadcast(result, src=q.tp_source, group=q.tp_group) + return result.view(*ids.shape, self.width) + + def _gather_metadata_device(self, metadata, device, group): + """All-gather metadata through reusable device buffers.""" + key = (str(device), metadata.numel(), metadata.dtype) + buffers = self._metadata_device_buffers.get(key) + if buffers is None: + buffers = ( + torch.empty(metadata.numel(), dtype=metadata.dtype, device=device), + torch.empty(self.query_group.size * metadata.numel(), dtype=metadata.dtype, device=device), + ) + self._metadata_device_buffers[key] = buffers + metadata_device, gathered_device = buffers + metadata_device.copy_(metadata, non_blocking=False) + dist.all_gather_into_tensor(gathered_device, metadata_device, group=group) + # The reshape/unbind views are copied to CPU before returning. Keep + # consumption local to this route so the reusable device buffer remains + # safe for the next collective. + return list(gathered_device.reshape(self.query_group.size, *metadata.shape).cpu().unbind(0)) + + @torch.inference_mode() + def forward(self, ids): + q = self.query_group + if ids.device.type != "cpu" or ids.dtype != torch.int64: + raise ValueError("Engram routing expects CPU int64 IDs") + _, _, metadata = self._metadata(ids) + if q.metadata_on_device: + device = self.weight.device if self.weight.device.type == "npu" else torch.device("npu") + gathered = self._gather_metadata_device(metadata, device, q.group) + else: + gathered = [torch.empty_like(metadata) for _ in range(q.size)] + dist.all_gather(gathered, metadata, group=q.cpu_group) + return self._forward_with_gathered(ids, gathered) + + @torch.inference_mode() + def forward_many(self, ids_list): + """Route several Engram tables with one CPU metadata collective.""" + return self.route_many([self] * len(ids_list), ids_list) + + @torch.inference_mode() + def route_many(self, tables, ids_list): + """Route distinct tables while sharing their CPU metadata collective.""" + if not ids_list: + return [] + q = self.query_group + if len(tables) != len(ids_list): + raise ValueError("tables and ids_list must have the same length") + routing = [table._metadata(ids) for table, ids in zip(tables, ids_list)] + metadata = [item[2] for item in routing] + packed = torch.cat(metadata) + if q.metadata_on_device: + device = tables[0].weight.device if tables[0].weight.device.type == "npu" else torch.device("npu") + gathered_packed = self._gather_metadata_device(packed, device, q.group) + else: + gathered_packed = [torch.empty_like(packed) for _ in range(q.size)] + dist.all_gather(gathered_packed, packed, group=q.cpu_group) + width = q.size + 1 + gathered = [ + [row[offset : offset + width] for row in gathered_packed] + for offset in range(0, len(ids_list) * width, width) + ] + if len(tables) == 1: + result = tables[0]._forward_with_gathered(ids_list[0], gathered[0], routing[0]) + return [result] + total = sum(ids.numel() * table.width for table, ids in zip(tables, ids_list)) + device = tables[0].weight.device + if device.type == "cpu" and str(dist.get_backend(q.group)).lower() not in ("gloo", "mpi"): + device = torch.device("npu") + if device.type == "npu" and device.index is None: + device = torch.device("npu", torch.npu.current_device()) + combined = torch.empty(total, dtype=torch.bfloat16, device=device) + results = [] + offset = 0 + for table, ids, group, item in zip(tables, ids_list, gathered, routing): + size = ids.numel() * table.width + result = table._forward_with_gathered( + ids, + group, + item, + broadcast=False, + output=combined[offset : offset + size].view(ids.numel(), table.width), + ) + results.append(result) + offset += size + if total: + dist.broadcast(combined, src=q.tp_source, group=q.tp_group) + outputs = [] + offset = 0 + for result in results: + size = result.numel() + outputs.append(combined[offset : offset + size].view_as(result)) + offset += size + return outputs diff --git a/vllm_ascend/models/deepseek_v41/indexer.py b/vllm_ascend/models/deepseek_v41/indexer.py new file mode 100644 index 000000000000..9daad9572be1 --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/indexer.py @@ -0,0 +1,246 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""V4.1 index projections, quantized QLI and cross-layer candidate selection.""" + +import torch +import torch_npu +from torch import nn +from vllm.model_executor.layers.linear import ReplicatedLinear + +from vllm_ascend.attention.dsa_v41 import ( + DeepseekV41CacheLayer, + scatter_cache_sk, +) +from vllm_ascend.core.deepseek_v41 import DeepseekV41IndexerSpec +from vllm_ascend.ops.triton.prepare_indexer_indices import prepare_indexer_indices +from vllm_ascend.ops.triton.quantize_indexer_query import quantize_indexer_query +from vllm_ascend.worker.device_metadata import ( + DeviceMetadataStage, + wait_for_device_metadata, +) + +from .compressor import DeepseekV41RMSNorm, _read + + +class DeepseekV41Indexer(nn.Module): + """Small side attention that selects compressed KV positions. + + All index heads are replicated on each TP rank for the correctness path, + so every rank produces identical sparse indices without an all-reduce. + """ + + def __init__( + self, + config, + owns_k, + vllm_config, + prefix, + compress_ratio, + quant_config=None, + ): + super().__init__() + self.owns_k = owns_k + self.compress_ratio = compress_ratio + self.n_heads = int(_read(config, "index_n_heads")) + self.width = int(_read(config, "index_head_dim")) + self.rope_width = int(_read(config, "qk_rope_head_dim")) + self.index_topk = int(_read(config, "index_topk")) + self.softmax_scale = self.width**-0.5 + self.weights_scale = self.softmax_scale * self.n_heads**-0.5 + self.wq_b = ReplicatedLinear( + _read(config, "q_lora_rank"), + self.n_heads * self.width, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.wq_b", + return_bias=False, + ) + self.weights_proj = ReplicatedLinear( + _read(config, "hidden_size"), + self.n_heads, + bias=False, + quant_config=None, + prefix=f"{prefix}.weights_proj", + return_bias=False, + ) + if owns_k: + self.wk = nn.Linear( + _read(config, "head_dim"), + self.width, + bias=False, + dtype=torch.bfloat16, + ) + self.k_norm = DeepseekV41RMSNorm(self.width, _read(config, "rms_norm_eps")) + self.k_cache = DeepseekV41CacheLayer( + vllm_config, + f"{prefix}.k_cache", + DeepseekV41IndexerSpec( + block_size=vllm_config.cache_config.block_size, + num_kv_heads=1, + head_size=self.width, + dtype=torch.int8, + tokens_per_state=compress_ratio, + storage_block_size=(vllm_config.cache_config.block_size // compress_ratio), + scale_dim=1, + scale_dtype=torch.float16, + ), + ) + + @staticmethod + def _output(linear, value): + output = linear(value) + return output[0] if isinstance(output, tuple) else output + + def update_keys(self, latent, slots, cos, sin): + """Publish source-owned index K before latent is RoPE'd as long KV.""" + if not self.owns_k or latent.shape[0] == 0: + return + key = self.k_norm(self.wk(latent)).view(-1, 1, self.width) + torch.ops._C_ascend.inplace_partial_rotary_mul( + key.unsqueeze(1), + cos, + sin, + rotary_mode="interleave", + partial_slice=[self.width - self.rope_width, self.width], + ) + key = key.squeeze(1) + quantized, scale = torch_npu.npu_dynamic_quant(key, dst_type=torch.int8) + k_cache, scale_cache = self.k_cache.kv_cache[0] + scatter_cache_sk(k_cache, slots, quantized) + scatter_cache_sk( + scale_cache, + slots, + scale.unsqueeze(-1).to(torch.float16), + ) + + def select( + self, + hidden_states, + qr, + positions, + cos, + sin, + source_cache, + source_metadata, + *, + is_candidate_source, + uses_candidate_filter, + candidate_topk_blocks, + candidate_block_size, + candidates, + ): + """Score index K, optionally filter blocks, then return position TopK.""" + query = self._output(self.wq_b, qr).unflatten(-1, (self.n_heads, self.width)) + torch.ops._C_ascend.inplace_partial_rotary_mul( + query.unsqueeze(1), + cos, + sin, + rotary_mode="interleave", + partial_slice=[self.width - self.rope_width, self.width], + ) + weights = self._output(self.weights_proj, hidden_states) + weights = weights.float() * self.weights_scale + + return self.select_projected( + query, + weights, + positions, + source_cache, + source_metadata, + is_candidate_source=is_candidate_source, + uses_candidate_filter=uses_candidate_filter, + candidate_topk_blocks=candidate_topk_blocks, + candidate_block_size=candidate_block_size, + candidates=candidates, + ) + + def select_projected( + self, + query, + weights, + positions, + source_cache, + source_metadata, + *, + is_candidate_source, + uses_candidate_filter, + candidate_topk_blocks, + candidate_block_size, + candidates, + ): + """Run QLI V2 on paged INT8 K; candidates are block IDs, not positions. + + Source and consumer share [tokens, 1, candidate_topk_blocks] INT32 + block IDs only within this forward. Query quantization and position + ordering stay outside the native QLI/candidate operator. + """ + if is_candidate_source and uses_candidate_filter: + raise ValueError("A candidate source must use the unfiltered position TopK") + if uses_candidate_filter and candidates is None: + raise RuntimeError("V4.1 candidate-filtering indexer ran before its source") + if self.width != 128 or self.n_heads not in (32, 64): + raise ValueError("A3 QLI requires index_head_dim=128 and 32 or 64 index heads") + if not 1 <= self.index_topk <= 2048: + raise ValueError("A3 QLI requires index_topk in [1, 2048]") + if self.compress_ratio not in (1, 2): + raise ValueError("Aurora QLI supports compression ratios 1 and 2") + if is_candidate_source or uses_candidate_filter: + if not 0 < candidate_topk_blocks <= 2048 or candidate_topk_blocks % 64: + raise ValueError("candidate_topk_blocks must be a multiple of 64 in [64, 2048]") + if candidate_block_size != 8: + raise ValueError("The current A3 candidate kernel requires candidate_block_size=8") + candidate_shape = (query.shape[0], 1, candidate_topk_blocks) + if uses_candidate_filter and (candidates.shape != candidate_shape or candidates.dtype != torch.int32): + raise ValueError("Candidate consumer requires INT32 block IDs with matching query rows") + topk = self.index_topk + if query.shape[0] == 0: + selected = torch.full((0, topk), -1, dtype=torch.int32, device=query.device) + if is_candidate_source: + candidates = torch.full(candidate_shape, -1, dtype=torch.int32, device=query.device) + return selected, candidates + if source_metadata.max_cache_seq_len == 0: + selected = torch.full((query.shape[0], 0), -1, dtype=torch.int32, device=query.device) + if is_candidate_source: + candidates = torch.full(candidate_shape, -1, dtype=torch.int32, device=query.device) + return selected, candidates + + quantized_query, query_scale = quantize_indexer_query(query) + weights = weights.to(torch.float16) + key, key_scale = source_cache + key_scale = key_scale.squeeze(-1) # Preserve the Hybrid cache page stride. + cu_seqlens_q = source_metadata.query_start_loc + seqused_k = source_metadata.cache_seq_lens + residual = source_metadata.cmp_residual + common = dict( + cu_seqlens_q=cu_seqlens_q, + seqused_k=seqused_k, + cmp_residual_k=residual, + max_seqlen_q=source_metadata.max_query_len, + layout_q="TND", + layout_k="PA_BBND", + mask_mode=3, + cmp_ratio=self.compress_ratio, + ) + op_metadata = source_metadata.qli_metadata + if op_metadata is None: + raise RuntimeError("V4.1 QLI metadata was not built") + wait_for_device_metadata(DeviceMetadataStage.INDEXER, id(op_metadata)) + mode = 1 if is_candidate_source else 2 if uses_candidate_filter else 3 + selected, _, candidate_out = torch.ops._C_ascend.npu_quant_lightning_indexer_v3( + quantized_query, + key, + weights, + query_scale, + key_scale, + topk, + 2, + block_table=source_metadata.block_table, + metadata=op_metadata, + candidate_topk_index=candidates if uses_candidate_filter else None, + candidate_mode=mode, + candidate_topk_blocks=candidate_topk_blocks, + candidate_block_size=candidate_block_size, + **common, + ) + selected = prepare_indexer_indices(selected.squeeze(1), positions, self.compress_ratio) + return selected, candidate_out if is_candidate_source else candidates diff --git a/vllm_ascend/models/deepseek_v41/mm_preprocess.py b/vllm_ascend/models/deepseek_v41/mm_preprocess.py new file mode 100644 index 000000000000..594887997017 --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/mm_preprocess.py @@ -0,0 +1,440 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Ascend project +"""DeepSeek V4.1 multimodal preprocessing. + +V4.1 deliberately does not reuse V4's five-sentinel N-layout. The checkpoint +uses one token id for every image-span position and carries the +start/image/newline/end role in a separate tensor. vLLM adds at most one +reserved-token row before the span so ratio-2 compressor groups start at a +stable phase. +""" + +import copy +import math +import threading +from collections.abc import Mapping, Sequence +from typing import Any, cast + +import numpy as np +import torch +from PIL import Image, ImageOps +from transformers import BatchFeature +from vllm.config.multimodal import BaseDummyOptions, ImageDummyOptions +from vllm.inputs import MultiModalDataDict +from vllm.multimodal.inputs import MultiModalFieldConfig, MultiModalKwargsItems +from vllm.multimodal.parse import ImageSize, MultiModalDataItems +from vllm.multimodal.processing import ( + BaseDummyInputsBuilder, + BaseMultiModalProcessor, + BaseProcessingInfo, + PromptReplacement, + PromptUpdate, + PromptUpdateDetails, +) +from vllm.multimodal.processing.processor import ( + MultiModalPromptUpdates, + PlaceholderFeaturesInfo, +) + +from vllm_ascend.deepseek_v41_config import DeepseekV41Config + +IMAGE_START, IMAGE, IMAGE_NEW_LINE, IMAGE_END = range(4) +COMPRESS_PAD_TO = 2 + +IMAGE_PLACEHOLDER = "<|deepseek_image|>" +IMAGE_TOKEN_ID = 129264 +# A reserved in-vocabulary token used only for vLLM's leading compressor pad. +IMAGE_PAD_ID = 129265 +IMAGE_PAD_TOKEN_NAME = "<|place_holder_mm_span_0436|>" + +_TOKENIZER_THREAD_LOCAL = threading.local() + + +def _get_thread_local_tokenizer(tokenizer): + cached = getattr(_TOKENIZER_THREAD_LOCAL, "tokenizer", None) + source_id = getattr(_TOKENIZER_THREAD_LOCAL, "source_id", None) + if cached is None or source_id != id(tokenizer): + cached = copy.deepcopy(tokenizer) + _TOKENIZER_THREAD_LOCAL.tokenizer = cached + _TOKENIZER_THREAD_LOCAL.source_id = id(tokenizer) + return cached + + +def image_sentinel_mask(token_ids: torch.Tensor) -> torch.Tensor: + """Return image-span and compressor-pad positions.""" + return (token_ids == IMAGE_TOKEN_ID) | (token_ids == IMAGE_PAD_ID) + + +def validate_image_sentinel_ids(tokenizer) -> None: + image_id = tokenizer.convert_tokens_to_ids(IMAGE_PLACEHOLDER) + if image_id != IMAGE_TOKEN_ID: + raise ValueError(f"Image placeholder {IMAGE_PLACEHOLDER!r} has id {image_id}, expected {IMAGE_TOKEN_ID}.") + pad_id = tokenizer.convert_tokens_to_ids(IMAGE_PAD_TOKEN_NAME) + if pad_id != IMAGE_PAD_ID: + raise ValueError(f"Image pad token {IMAGE_PAD_TOKEN_NAME!r} has id {pad_id}, expected {IMAGE_PAD_ID}.") + + +def llm_grid(best_height, best_width, patch_size, downsample_ratio): + return ( + math.ceil((best_height // patch_size) / downsample_ratio), + math.ceil((best_width // patch_size) / downsample_ratio), + ) + + +def num_image_tokens(n_llm_h: int, n_llm_w: int) -> int: + return n_llm_h * (n_llm_w + 1) + 2 + + +def leading_compressor_pad(start_pos: int) -> int: + """Rows needed to align an image span to the V4.1 CR2 phase.""" + return COMPRESS_PAD_TO - 1 - start_pos % COMPRESS_PAD_TO + + +def solve_resize_ratio(height, width, patch_size, downsample_ratio, max_n_token): + r = height / width + max_w_float = math.sqrt((max_n_token - 2) / r + 0.25) - 0.5 + max_h_float = max_w_float * r + cell = patch_size * downsample_ratio + if max_w_float < 1.0: + return (max_n_token - 2) // 2 * cell, cell + if max_h_float < 1.0: + return cell, (max_n_token - 3) * cell + beta = min( + math.floor(max_w_float) * cell / width, + math.floor(max_h_float) * cell / height, + ) + return ( + math.floor(height * beta / patch_size) * patch_size, + math.floor(width * beta / patch_size) * patch_size, + ) + + +def safe_resize( + height, + width, + best_height, + best_width, + patch_size, + downsample_ratio, + max_n_token, +): + # Reserve the maximum one-token leading pad injected after tokenization. + max_n_token -= COMPRESS_PAD_TO - 1 + n_llm_h, n_llm_w = llm_grid(best_height, best_width, patch_size, downsample_ratio) + if num_image_tokens(n_llm_h, n_llm_w) > max_n_token: + best_height, best_width = solve_resize_ratio( + height, + width, + patch_size, + downsample_ratio, + max_n_token, + ) + n_llm_h, n_llm_w = llm_grid(best_height, best_width, patch_size, downsample_ratio) + assert num_image_tokens(n_llm_h, n_llm_w) <= max_n_token + return n_llm_h, n_llm_w, best_height, best_width + + +def load_image( + image: Image.Image, + *, + patch_size: int, + downsample_ratio: int, + max_n_token: int, + min_pixels: int, + max_wh_ratio: float | None, +): + p = patch_size + image = image.convert("RGB") + width, height = image.size + if max_wh_ratio is not None and width > height * max_wh_ratio: + width = height * max_wh_ratio + if 0 < width * height < min_pixels: + ratio = (min_pixels / (width * height)) ** 0.5 + width = int(width * ratio) + height = int(height * ratio) + best_width = math.ceil(width / p) * p + best_height = math.ceil(height / p) * p + n_llm_h, n_llm_w, best_height, best_width = safe_resize( + height, + width, + best_height, + best_width, + p, + downsample_ratio, + max_n_token, + ) + n_vit_h, n_vit_w = best_height // p, best_width // p + if max_wh_ratio is not None and image.width >= max_wh_ratio * image.height: + image = image.resize((best_width, best_height)) + else: + image = ImageOps.pad( + image, + (best_width, best_height), + color=(127, 127, 127), + ) + x = torch.from_numpy(np.asarray(image, dtype=np.float32)).permute(2, 0, 1) / 255 + x = ((x - 0.5) / 0.5).to(torch.bfloat16) + patches = x.reshape(3, n_vit_h, p, n_vit_w, p).permute(1, 3, 0, 2, 4).reshape(n_vit_h * n_vit_w, 3, p, p) + return patches, n_vit_h, n_vit_w, n_llm_h, n_llm_w + + +def image_token_types(n_llm_h: int, n_llm_w: int) -> torch.Tensor: + """Reference reading-order V4.1 image-span roles.""" + types = [IMAGE_START] + types += ([IMAGE] * n_llm_w + [IMAGE_NEW_LINE]) * n_llm_h + types.append(IMAGE_END) + return torch.tensor(types, dtype=torch.int64) + + +class DeepseekV41VLImageProcessor: + def __init__(self, config: DeepseekV41Config) -> None: + self.patch_size = config.vision_patch_size + self.downsample_ratio = config.vision_downsample_ratio + self.max_n_token = config.vision_max_n_token + self.min_pixels = config.vision_min_pixels + self.max_wh_ratio = config.vision_max_wh_ratio + + def __call__(self, image: Image.Image): + return load_image( + image, + patch_size=self.patch_size, + downsample_ratio=self.downsample_ratio, + max_n_token=self.max_n_token, + min_pixels=self.min_pixels, + max_wh_ratio=self.max_wh_ratio, + ) + + +class DeepseekV41VLProcessor: + def __init__(self, config: DeepseekV41Config) -> None: + self.config = config + self.image_processor = DeepseekV41VLImageProcessor(config) + + def __call__( + self, + text: str | None = None, + images: Sequence[Image.Image] | None = None, + return_tensors: str | None = None, + **kwargs: Any, + ) -> BatchFeature: + del text, return_tensors, kwargs + patches_list = [] + vit_grid = [] + llm_grid_list = [] + types_list = [] + for image in images or []: + patches, n_vit_h, n_vit_w, n_llm_h, n_llm_w = self.image_processor(image) + patches_list.append(patches) + vit_grid.append((n_vit_h, n_vit_w)) + llm_grid_list.append((n_llm_h, n_llm_w)) + types_list.append(image_token_types(n_llm_h, n_llm_w)) + + if not patches_list: + return BatchFeature({}) + + return BatchFeature( + { + "patches": torch.cat(patches_list), + "vit_grid": torch.tensor(vit_grid, dtype=torch.int64), + "llm_grid": torch.tensor(llm_grid_list, dtype=torch.int64), + "types": torch.cat(types_list), + } + ) + + +class DeepseekV41VLProcessingInfo(BaseProcessingInfo): + def get_hf_config(self) -> DeepseekV41Config: + return self.ctx.get_hf_config(DeepseekV41Config) + + def get_hf_processor(self, **kwargs: object) -> DeepseekV41VLProcessor: + if kwargs: + raise ValueError(f"Unexpected processor kwargs: {sorted(kwargs)}") + return DeepseekV41VLProcessor(self.get_hf_config()) + + def get_supported_mm_limits(self) -> Mapping[str, int | None]: + return {"image": None} + + def get_mm_max_tokens_per_item( + self, + seq_len: int, + mm_counts: Mapping[str, int], + ) -> Mapping[str, int]: + del seq_len, mm_counts + return {"image": self.get_hf_config().vision_max_n_token + COMPRESS_PAD_TO - 1} + + def get_image_placeholder_token_id(self) -> int: + token_id = self.get_tokenizer().convert_tokens_to_ids(IMAGE_PLACEHOLDER) + if token_id is None: + raise ValueError(f"Token not found in tokenizer: {IMAGE_PLACEHOLDER}") + return token_id + + def get_image_size_with_most_features(self) -> ImageSize: + config = self.get_hf_config() + budget = config.vision_max_n_token - (COMPRESS_PAD_TO - 1) + side = budget * config.vision_patch_size * config.vision_downsample_ratio + best_h, best_w = solve_resize_ratio( + side, + side, + config.vision_patch_size, + config.vision_downsample_ratio, + budget, + ) + return ImageSize(width=best_w, height=best_h) + + +class DeepseekV41VLDummyInputsBuilder(BaseDummyInputsBuilder[DeepseekV41VLProcessingInfo]): + def get_dummy_text(self, mm_counts: Mapping[str, int]) -> str: + return IMAGE_PLACEHOLDER * mm_counts.get("image", 0) + + def get_dummy_mm_data( + self, + seq_len: int, + mm_counts: Mapping[str, int], + mm_options: Mapping[str, BaseDummyOptions], + ) -> MultiModalDataDict: + del seq_len + size = self.info.get_image_size_with_most_features() + return { + "image": self._get_dummy_images( + width=size.width, + height=size.height, + num_images=mm_counts.get("image", 0), + overrides=cast( + ImageDummyOptions | None, + mm_options.get("image"), + ), + ), + } + + +class DeepseekV41VLMultiModalProcessor(BaseMultiModalProcessor[DeepseekV41VLProcessingInfo]): + def _call_hf_processor( + self, + prompt: str, + mm_data: Mapping[str, object], + mm_kwargs: Mapping[str, object], + tok_kwargs: Mapping[str, object] | None = None, + ) -> BatchFeature: + if tok_kwargs is None: + tok_kwargs = {} + processor = self.info.get_hf_processor(**mm_kwargs) + processed = processor( + text=prompt, + images=cast(Sequence[Image.Image] | None, mm_data.get("images")), + return_tensors="pt", + ) + tokenizer = _get_thread_local_tokenizer(self.info.get_tokenizer()) + processed["input_ids"] = tokenizer( + prompt, + return_tensors="pt", + **tok_kwargs, + )["input_ids"] + return processed + + def _hf_processor_applies_updates( + self, + prompt_text: str, + mm_items: MultiModalDataItems, + hf_processor_mm_kwargs: Mapping[str, object], + tokenization_kwargs: Mapping[str, object], + ) -> bool: + del prompt_text, mm_items, hf_processor_mm_kwargs, tokenization_kwargs + return False + + def _get_mm_fields_config( + self, + hf_inputs: BatchFeature, + hf_processor_mm_kwargs: Mapping[str, object], + ) -> Mapping[str, MultiModalFieldConfig]: + del hf_processor_mm_kwargs + vit_grid = hf_inputs.get("vit_grid") + llm_grid = hf_inputs.get("llm_grid") + if vit_grid is None or llm_grid is None: + empty = torch.empty(0, dtype=torch.long) + patch_sizes = types_sizes = empty + else: + patch_sizes = vit_grid.prod(-1) + n_llm_h, n_llm_w = llm_grid[:, 0], llm_grid[:, 1] + types_sizes = n_llm_h * (n_llm_w + 1) + 2 + return { + "patches": MultiModalFieldConfig.flat_from_sizes("image", patch_sizes), + "vit_grid": MultiModalFieldConfig.batched("image", keep_on_cpu=True), + "llm_grid": MultiModalFieldConfig.batched("image", keep_on_cpu=True), + "types": MultiModalFieldConfig.flat_from_sizes("image", types_sizes, keep_on_cpu=True), + } + + def _get_prompt_updates( + self, + mm_items: MultiModalDataItems, + hf_processor_mm_kwargs: Mapping[str, object], + out_mm_kwargs: MultiModalKwargsItems, + ) -> Sequence[PromptUpdate]: + del mm_items, hf_processor_mm_kwargs + image_token_id = self.info.get_image_placeholder_token_id() + validate_image_sentinel_ids(self.info.get_tokenizer()) + + def get_image_replacement(item_idx: int) -> PromptUpdateDetails: + types: torch.Tensor = out_mm_kwargs["image"][item_idx]["types"].data + full = [image_token_id] * types.numel() + return PromptUpdateDetails.select_token_id(full, image_token_id) + + return [ + PromptReplacement( + modality="image", + target=[image_token_id], + replacement=get_image_replacement, + ) + ] + + def _apply_prompt_updates( + self, + token_ids: list[int], + mm_prompt_updates: MultiModalPromptUpdates, + ) -> tuple[ + list[int], + Mapping[str, list[PlaceholderFeaturesInfo]], + ]: + """Apply replacements and prepend the position-dependent CR2 pad.""" + new_token_ids, base_placeholders = super()._apply_prompt_updates( + token_ids, + mm_prompt_updates, + ) + placeholders: dict[str, list[PlaceholderFeaturesInfo]] = {modality: [] for modality in base_placeholders} + ordered = sorted( + ( + placeholder.start_idx, + modality, + placeholder, + ) + for modality, items in base_placeholders.items() + for placeholder in items + ) + inserted = 0 + for _, modality, placeholder in ordered: + start_idx = placeholder.start_idx + inserted + tokens = list(placeholder.tokens) + is_embed = placeholder.is_embed + if modality == "image": + compress_pad = leading_compressor_pad(start_idx) + new_token_ids[start_idx:start_idx] = [IMAGE_PAD_ID] * compress_pad + tokens = [IMAGE_PAD_ID] * compress_pad + tokens + original_mask = ( + is_embed if is_embed is not None else torch.ones(len(placeholder.tokens), dtype=torch.bool) + ) + is_embed = torch.cat( + [ + torch.zeros(compress_pad, dtype=torch.bool), + original_mask, + ] + ) + inserted += compress_pad + placeholders[modality].append( + PlaceholderFeaturesInfo( + modality=modality, + item_idx=placeholder.item_idx, + start_idx=start_idx, + tokens=tokens, + is_embed=is_embed, + ) + ) + return new_token_ids, placeholders diff --git a/vllm_ascend/models/deepseek_v41/model.py b/vllm_ascend/models/deepseek_v41/model.py new file mode 100644 index 000000000000..f8fcd493778d --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/model.py @@ -0,0 +1,723 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""DeepSeek V4.1 text model and source-shared hybrid-cache graph.""" + +from collections.abc import Iterable, Iterator +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import torch +import vllm.envs as envs +from safetensors import safe_open +from transformers import AutoTokenizer +from vllm.distributed import get_pp_group +from vllm.forward_context import get_forward_context, is_forward_context_available + +from vllm_ascend.ascend_config import get_ascend_config +from vllm_ascend.attention.dsa_v41 import ( + DeepseekV41CacheBackend, + DeepseekV41CacheLayer, +) +from vllm_ascend.core.deepseek_v41 import ( + DeepseekV41FullSpec, + DeepseekV41SWASpec, + validate_cache_runtime, +) +from vllm_ascend.models.common.ops.sequence_parallel import ( + sp_all_gather, + sp_padding_mask, + sp_reduce_scatter, + sp_shard, +) +from vllm_ascend.models.deepseek_v4.model import ( + AscendDeepseekV4ForCausalLM, + AscendDeepseekV4SWACache, + DeepseekV2DecoderLayer, + DeepseekV4Attention, + DeepseekV4Model, +) + +from .compressor import DeepseekV41Compressor, _read, text_config_of +from .engram_gate import engram_gate +from .engram_hash import PagedNgramHistory, engram_history_metadata +from .engram_hbm import EngramQueryGroup, NodeShardedEngram +from .indexer import DeepseekV41Indexer + + +@dataclass(frozen=True) +class DeepseekV41LayerRole: + """The attention and future Engram responsibilities of one backbone layer.""" + + layer_idx: int + compress_ratio: int + kv_source_layer: int | None + index_source_layer: int | None + is_kv_source: bool + is_index_source: bool + is_candidate_source: bool + uses_candidate_filter: bool + engram_slot: int | None + + @property + def has_long_context(self) -> bool: + return self.compress_ratio > 0 + + +@dataclass(frozen=True) +class DeepseekV41Topology: + """Validated, immutable model-wide source/consumer topology.""" + + layers: tuple[DeepseekV41LayerRole, ...] + kv_source_layers: tuple[int, ...] + index_source_layers: tuple[int, ...] + candidate_source_layer: int + candidate_topk_blocks: int + candidate_block_size: int + index_topk: int + + def layer(self, layer_idx: int) -> DeepseekV41LayerRole: + return self.layers[layer_idx] + + def kv_consumers(self, source_layer: int) -> tuple[int, ...]: + return tuple(role.layer_idx for role in self.layers if role.kv_source_layer == source_layer) + + def index_consumers(self, source_layer: int) -> tuple[int, ...]: + return tuple(role.layer_idx for role in self.layers if role.index_source_layer == source_layer) + + +class DeepseekV41SharedAttentionState: + """Per-forward handoff between index sources and their consumer layers.""" + + def __init__(self, topk_indices, candidates): + self.topk_indices = topk_indices + self.candidates = candidates + + def reset(self): + # Source layers overwrite the active rows before any consumer reads + # them. Keeping the storage intact avoids replay depending on Python + # state mutation and preserves a fixed address for ACL Graph. + return None + + +def _as_int_tuple(config: Any, name: str) -> tuple[int, ...]: + value = _read(config, name) + if not isinstance(value, (list, tuple)) or any(not isinstance(item, int) for item in value): + raise ValueError(f"DeepSeek V4.1 {name} must be a list of integers") + return tuple(value) + + +def _latest_source(layer_idx: int, sources: tuple[int, ...]) -> int | None: + return next((source for source in reversed(sources) if source <= layer_idx), None) + + +def build_layer_plan(config: Any) -> DeepseekV41Topology: + """Build and validate the V4.1 layer-sharing graph from a text config. + + ``config`` may be a Transformers config object or the raw ``text_config`` + dictionary. Extra compression ratios for speculative layers are allowed, + but only the first ``num_hidden_layers`` entries describe the backbone. + """ + + config = text_config_of(config) + num_layers = int(_read(config, "num_hidden_layers")) + ratios = _as_int_tuple(config, "compress_ratios") + kv_sources = _as_int_tuple(config, "kv_source_layers") + index_sources = _as_int_tuple(config, "index_source_layers") + engram_layers = _as_int_tuple(config, "engram_layer_ids") + candidate_source = int(_read(config, "candidate_source_layer")) + candidate_topk_blocks = int(_read(config, "candidate_topk_blocks")) + candidate_block_size = int(_read(config, "candidate_block_size")) + index_topk = int(_read(config, "index_topk")) + + if num_layers <= 0: + raise ValueError("DeepSeek V4.1 num_hidden_layers must be positive") + if len(ratios) < num_layers: + raise ValueError( + "DeepSeek V4.1 compress_ratios must cover every backbone layer: " + f"got {len(ratios)} ratios for {num_layers} layers" + ) + ratios = ratios[:num_layers] + if any(ratio not in (0, 1, 2) for ratio in ratios): + raise ValueError(f"DeepSeek V4.1 backbone only supports compression ratios 0, 1 and 2; got {ratios}") + + for name, sources in (("kv_source_layers", kv_sources), ("index_source_layers", index_sources)): + if tuple(sorted(set(sources))) != sources: + raise ValueError(f"DeepSeek V4.1 {name} must be sorted and unique") + if any(source < 0 or source >= num_layers for source in sources): + raise ValueError(f"DeepSeek V4.1 {name} contains a layer outside the backbone") + if any(ratios[source] == 0 for source in sources): + raise ValueError(f"DeepSeek V4.1 {name} cannot point to a local-only layer") + + if not set(kv_sources).issubset(index_sources): + raise ValueError("Every DeepSeek V4.1 KV source must also be an index source") + if candidate_source not in kv_sources: + raise ValueError("DeepSeek V4.1 candidate_source_layer must be a KV source") + if candidate_topk_blocks <= 0 or candidate_block_size <= 0 or index_topk <= 0: + raise ValueError("DeepSeek V4.1 candidate and index TopK values must be positive") + if len(set(engram_layers)) != len(engram_layers): + raise ValueError("DeepSeek V4.1 engram_layer_ids must be unique") + if any(layer < 0 or layer >= num_layers for layer in engram_layers): + raise ValueError("DeepSeek V4.1 engram_layer_ids contains a layer outside the backbone") + + engram_slots = {layer_idx: slot for slot, layer_idx in enumerate(engram_layers)} + roles: list[DeepseekV41LayerRole] = [] + for layer_idx, ratio in enumerate(ratios): + kv_source = _latest_source(layer_idx, kv_sources) if ratio else None + index_source = _latest_source(layer_idx, index_sources) if ratio else None + if ratio and (kv_source is None or index_source is None): + raise ValueError(f"DeepSeek V4.1 layer {layer_idx} has long-context attention but no source layer") + if kv_source is not None and ratios[kv_source] != ratio: + raise ValueError( + f"DeepSeek V4.1 layer {layer_idx} has ratio {ratio}, but its KV source " + f"layer {kv_source} has ratio {ratios[kv_source]}" + ) + + roles.append( + DeepseekV41LayerRole( + layer_idx=layer_idx, + compress_ratio=ratio, + kv_source_layer=kv_source, + index_source_layer=index_source, + is_kv_source=layer_idx in kv_sources, + is_index_source=layer_idx in index_sources, + is_candidate_source=layer_idx == candidate_source, + # Consumer layers inherit the selection policy of their index + # source. For example, layer 26 reuses layer 24 TopK, and that + # TopK was computed inside layer 20's candidate blocks. + uses_candidate_filter=index_source is not None and index_source > candidate_source, + engram_slot=engram_slots.get(layer_idx), + ) + ) + + return DeepseekV41Topology( + layers=tuple(roles), + kv_source_layers=kv_sources, + index_source_layers=index_sources, + candidate_source_layer=candidate_source, + candidate_topk_blocks=candidate_topk_blocks, + candidate_block_size=candidate_block_size, + index_topk=index_topk, + ) + + +class AscendDeepseekV41SWACache(AscendDeepseekV4SWACache): + """V4 execution-compatible SWA plane participating in V4.1 grouping.""" + + def get_kv_cache_spec(self, vllm_config): + spec = super().get_kv_cache_spec(vllm_config) + return DeepseekV41SWASpec( + block_size=spec.block_size, + num_kv_heads=spec.num_kv_heads, + head_size=spec.head_size, + dtype=spec.dtype, + sliding_window=spec.sliding_window, + cache_dtype_str=spec.cache_dtype_str, + model_version="deepseek_v4", + alignment=spec.alignment, + ) + + def get_attn_backend(self): + return DeepseekV41CacheBackend + + +class DeepseekV41Attention(DeepseekV4Attention): + """V4.1 source-shared attention using V4 projections and CP adapters.""" + + swa_cache_cls = AscendDeepseekV41SWACache + + def __init__( + self, + vllm_config, + config, + max_position_embeddings=0, + cache_config=None, + quant_config=None, + prefix="", + topk_indices_buffer=None, + reduce_results=True, + need_gather_q_kv=False, + ): + config = text_config_of(config) + validate_cache_runtime(vllm_config) + layer_idx = int(prefix.split(".")[-2]) + topology = build_layer_plan(config) + role = topology.layer(layer_idx) + # Reuse V4's quant-aware projections and stable SWA eager backend. A + # zero ratio prevents V4 from creating its incompatible c4/c128 planes. + original_ratios = config.compress_ratios + config.compress_ratios = tuple(0 for _ in original_ratios) + try: + super().__init__( + vllm_config=vllm_config, + config=config, + max_position_embeddings=max_position_embeddings, + cache_config=cache_config, + quant_config=quant_config, + prefix=prefix, + topk_indices_buffer=topk_indices_buffer, + reduce_results=reduce_results, + need_gather_q_kv=need_gather_q_kv, + ) + finally: + config.compress_ratios = original_ratios + from vllm_ascend.ops.rope_dsv4 import ComplexExpRotaryEmbedding + + # V4.1 applies YaRN only to layers carrying long-context compressed KV. + # Pure SWA layers use the unscaled base RoPE even though the allocated + # lookup table still spans the configured maximum context length. + self.rotary_emb = ComplexExpRotaryEmbedding( + vllm_config=vllm_config, + layername=f"{prefix}.attn", + head_size=self.rope_head_dim, + rotary_dim=self.rope_head_dim, + max_position_embeddings=max_position_embeddings, + is_neox_style=False, + scaling_factor=config.rope_parameters["factor"], + base=(config.compress_rope_theta if role.has_long_context else config.rope_theta), + beta_fast=config.rope_parameters["beta_fast"], + beta_slow=config.rope_parameters["beta_slow"], + original_seq_len=(max_position_embeddings if role.has_long_context else 0), + rope_groups=["default"], + ) + block_size = vllm_config.cache_config.block_size + if block_size <= 0 or block_size % 2: + raise ValueError("V4.1 logical block_size must be a positive multiple of two") + owned: list[str] = [] + if role.is_kv_source: + owned.extend((f"{prefix}.long_kv_cache", f"{prefix}.indexer.k_cache")) + if role.compress_ratio == 2: + owned.append(f"{prefix}.compressor.state_cache") + duplicates = set(owned) & vllm_config.compilation_config.static_forward_context.keys() + if duplicates: + raise ValueError(f"Duplicate V4.1 cache prefixes: {sorted(duplicates)}") + self.role = role + self.topology = topology + self.shared_state = None + self.prefix = prefix + width = _read(config, "head_dim") + self.softmax_scale = width**-0.5 + if role.is_kv_source: + self.long_kv_cache = DeepseekV41CacheLayer( + vllm_config, + f"{prefix}.long_kv_cache", + DeepseekV41FullSpec( + block_size=block_size, + num_kv_heads=1, + head_size=width, + dtype=torch.bfloat16, + tokens_per_state=role.compress_ratio, + storage_block_size=block_size // role.compress_ratio, + ), + ) + self.compressor = ( + DeepseekV41Compressor(config, role.compress_ratio, vllm_config, f"{prefix}.compressor") + if role.is_kv_source + else None + ) + self.indexer = ( + DeepseekV41Indexer( + config, + role.is_kv_source, + vllm_config, + f"{prefix}.indexer", + role.compress_ratio, + quant_config=quant_config, + ) + if role.is_index_source + else None + ) + root = prefix.rsplit(".layers.", 1)[0] + source = f"{root}.layers.{role.kv_source_layer}.self_attn" + self.long_kv_source_prefix = f"{source}.long_kv_cache" if role.has_long_context else None + self.index_k_source_prefix = f"{source}.indexer.k_cache" if role.has_long_context else None + self.index_source_layer = role.index_source_layer + from vllm_ascend.attention.context_parallel.dsa_v41_cp import get_v41_cp_classes + + self.v41_impl = get_v41_cp_classes()[1]( + prefix=prefix, + role=role, + topology=topology, + long_kv_source_prefix=self.long_kv_source_prefix, + index_k_source_prefix=self.index_k_source_prefix, + ) + self.v41_layer_name = f"{prefix}.v41_attn" + context = vllm_config.compilation_config.static_forward_context + if self.v41_layer_name in context: + raise ValueError(f"Duplicate V4.1 attention layer: {self.v41_layer_name}") + context[self.v41_layer_name] = self + + def forward(self, positions, hidden_states, llama_4_scaling=None): + output = torch.empty_like(hidden_states) + torch.ops.vllm.dsa_v41_forward(hidden_states, output, self.v41_layer_name) + return output + + +class DeepseekV41DecoderLayer(DeepseekV2DecoderLayer): + """V4.1 block with the checkpoint's delayed mHC coefficient handoff.""" + + attention_cls = DeepseekV41Attention + + def __init__(self, vllm_config, prefix, **kwargs): + super().__init__(vllm_config, prefix, **kwargs) + self.use_sequence_parallel = vllm_config.parallel_config.use_sequence_parallel_moe + # Leave the TP partial sums for the reduce-scatter below. The mHC + # and MoE paths then stay sharded between attention calls. + if self.use_sequence_parallel: + self.self_attn.wo_b.reduce_results = False + config = vllm_config.model_config.hf_config + engram_enabled = get_ascend_config().enable_engram + if engram_enabled and self.layer_idx in config.engram_layer_ids: + self.engram = torch.nn.Module() + self.engram.wkv = torch.nn.Linear( + (config.engram_max_ngram_size - 1) * config.engram_n_heads * config.engram_head_dim, + (config.hc_mult + 1) * config.hidden_size, + bias=False, + dtype=torch.bfloat16, + ) + self.engram.q_weight = torch.nn.Parameter( + torch.empty(config.hc_mult, config.hidden_size, dtype=torch.bfloat16) + ) + self.engram.k_weight = torch.nn.Parameter( + torch.empty(config.hc_mult, config.hidden_size, dtype=torch.bfloat16) + ) + else: + self.engram = None + + @staticmethod + def hc_collapse(x, pre_mix): + return (pre_mix.unsqueeze(-1) * x.float()).sum(-2).to(x.dtype) + + def hc_pre(self, x, hc_fn, hc_scale, hc_base, pre_mix=None): + return torch.ops._C_ascend.npu_hc_pre_v3( + x, + hc_fn, + hc_scale, + hc_base, + pre_mix, + hc_mult=self.hc_mult, + hc_sinkhorn_iters=self.hc_sinkhorn_iters, + norm_eps=self.norm_eps, + hc_eps=self.hc_eps, + ) + + def hc_post(self, x, residual, post, comb): + return torch.ops._C_ascend.npu_hc_post( + x.unsqueeze(0), + residual.unsqueeze(0), + post.unsqueeze(0), + comb.unsqueeze(0), + ).squeeze(0) + + def forward( + self, + positions, + hidden_states, + pre_mix, + llama_4_scaling=None, + input_ids=None, + ): + use_sequence_parallel = getattr(self, "use_sequence_parallel", False) + residual = hidden_states + x, attn_post, attn_comb, attn_pre = self.hc_pre( + hidden_states, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + pre_mix, + ) + x = self.input_layernorm(x) + if use_sequence_parallel: + x = sp_all_gather(x)[: positions.shape[0]] + x = self.self_attn(positions, x, llama_4_scaling) + if use_sequence_parallel: + x = sp_reduce_scatter(x) + hidden_states = self.hc_post(x, residual, attn_post, attn_comb) + + residual = hidden_states + x, ffn_post, ffn_comb, ffn_pre = self.hc_pre( + hidden_states, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + attn_pre, + ) + x, x_fp32 = self.rms_norm_cast(x) + x = self.mlp( + x, + input_ids=input_ids, + hidden_states_fp32=x_fp32, + already_sequence_parallel=use_sequence_parallel, + ) + hidden_states = self.hc_post(x, residual, ffn_post, ffn_comb) + return hidden_states, ffn_pre + + +class DeepseekV41Model(DeepseekV4Model): + """Single V4.1 backbone entry, matching ``deepseek_v4/model.py``.""" + + decoder_layer_cls = DeepseekV41DecoderLayer + + def __init__(self, *, vllm_config, prefix=""): + if ( + get_ascend_config().enable_engram + and vllm_config.load_config.load_format != "dummy" + and vllm_config.load_config.safetensors_load_strategy != "lazy" + ): + raise ValueError("Engram HBM shards require --safetensors-load-strategy lazy") + super().__init__(vllm_config=vllm_config, prefix=prefix) + self.use_sequence_parallel = vllm_config.parallel_config.use_sequence_parallel_moe + # V4.1 collapses with the last block's ffn_pre; it has no hc_head + # projection in the checkpoint. + del self.hc_head_fn, self.hc_head_base, self.hc_head_scale, self.hc_norm + topology = build_layer_plan(self.config) + max_tokens = vllm_config.scheduler_config.max_num_batched_tokens + candidate_buffer = torch.full( + (max_tokens, 1, topology.candidate_topk_blocks), + -1, + dtype=torch.int32, + device=self.topk_indices_buffer.device, + ) + self.candidate_indices_buffer = candidate_buffer + self.shared_attention_state = DeepseekV41SharedAttentionState( + self.topk_indices_buffer, + candidate_buffer, + ) + for layer in self.layers: + if isinstance(layer, DeepseekV41DecoderLayer): + layer.self_attn.shared_state = self.shared_attention_state + self.engram_root = vllm_config.model_config.model + config = self.config + # Target storage is a loader/runtime choice. Checkpoint metadata is + # used only by load_checkpoint to validate the source representation. + # Read the storage choice after AscendConfig validation. + ascend_config = get_ascend_config() + storage_format = ascend_config.engram_storage + if ascend_config.enable_engram: + query_group = EngramQueryGroup.from_vllm(vllm_config.parallel_config) + for layer_id, rows in zip(config.engram_layer_ids, config.engram_num_embeddings): + self.layers[layer_id].engram.embed = NodeShardedEngram( + rows, + config.engram_head_dim, + query_group, + storage_format=storage_format, + ) + self.engram_history = None + self._engram_input_buffers = None + self._engram_max_tokens = max( + vllm_config.scheduler_config.max_num_batched_tokens, + vllm_config.compilation_config.max_cudagraph_capture_size or 0, + ) + self.register_buffer("engram_rotation", torch.eye(32), persistent=False) + if ascend_config.enable_engram and vllm_config.load_config.load_format != "dummy": + with torch.device("cpu"): + tokenizer = AutoTokenizer.from_pretrained(self.engram_root) + self.engram_history = PagedNgramHistory(config, tokenizer) + with safe_open(Path(self.engram_root) / "optional/quarot.safetensors", framework="pt") as file: + rotation = file.get_tensor("global_rotation") + block = rotation[:32, :32].contiguous() + if not torch.equal(rotation, torch.block_diag(*[block] * (config.hidden_size // 32))): + raise ValueError("Engram gate requires repeated block32 global rotation") + self.engram_rotation.copy_(block) + + def prepare_engram(self, input_ids, positions): + """Eager boundary: every DP participates, including metadata-free dummies.""" + config = self.config + if not get_ascend_config().enable_engram: + return {}, torch.empty(0, dtype=torch.bool, device=positions.device) + columns = (config.engram_max_ngram_size - 1) * config.engram_n_heads + hashes = torch.empty((0, len(config.engram_layer_ids), columns), dtype=torch.int64, device="cpu") + mask = torch.empty(0, dtype=torch.bool, device="cpu") + metadata = get_forward_context().attn_metadata + if metadata is not None and self.engram_history is not None: + first = self.layers[0].self_attn.dsa_attn.swa_cache_layer + meta = metadata[first.prefix] + boundaries, block_table, block_size = engram_history_metadata(meta) + n = int(boundaries[-1]) + requests = torch.repeat_interleave(torch.arange(len(boundaries) - 1, device="cpu"), boundaries.diff()) + hashes, mask = self.engram_history.update( + input_ids[:n].cpu().long(), + positions[:n].cpu().long(), + requests, + block_table, + block_size, + ) + lookups = {} + tables = [self.layers[layer_id].engram.embed for layer_id in config.engram_layer_ids] + ids_list = [hashes[:, slot] for slot in range(len(tables))] + if hasattr(tables[0], "route_many"): + routed = tables[0].route_many(tables, ids_list) + else: + routed = [table(ids) for table, ids in zip(tables, ids_list)] + for layer_id, values in zip(config.engram_layer_ids, routed): + lookups[layer_id] = values.flatten(1) + return lookups, mask.to(positions.device) + + def prepare_engram_inputs(self, input_ids, positions, padded_tokens=None): + """Refresh persistent inputs before main-model capture or replay.""" + lookups, mask = self.prepare_engram(input_ids, positions) + num_tokens = positions.shape[0] + # The compiled V4.1 backbone uses the scheduler's static token + # capacity for decode graphs (typically max_num_batched_tokens), even + # when the current request has one token. Keep lookup tensors at that + # capacity so every captured graph sees the same Engram shape. + output_tokens = max(self._engram_max_tokens, padded_tokens or 0) + if output_tokens < num_tokens: + raise ValueError("Engram padded token count is smaller than the input") + if self._engram_input_buffers is None: + capacity = self._engram_max_tokens + self._engram_input_buffers = ( + {layer: values.new_zeros((capacity, values.shape[1])) for layer, values in lookups.items()}, + mask.new_zeros(capacity), + ) + buffers, mask_buffer = self._engram_input_buffers + padded_mask = mask_buffer[:output_tokens] + padded_mask.zero_() + padded_mask[: mask.numel()].copy_(mask) + padded_lookups = {} + for layer, values in lookups.items(): + padded = buffers[layer][:output_tokens] + padded.zero_() + padded[: values.shape[0]].copy_(values) + padded_lookups[layer] = padded + return {"engram_lookups": padded_lookups, "engram_mask": padded_mask} + + def forward( + self, + input_ids, + positions, + intermediate_tensors, + inputs_embeds=None, + engram_lookups=None, + engram_mask=None, + ): + if not get_pp_group().is_first_rank or not get_pp_group().is_last_rank: + raise NotImplementedError("V4.1 eager milestone currently requires PP=1") + use_sequence_parallel = getattr(self, "use_sequence_parallel", False) + hidden_states = inputs_embeds if inputs_embeds is not None else self.embed_input_ids(input_ids) + if engram_lookups is None: + lookups, token_mask = self.prepare_engram(input_ids, positions) + else: + lookups, token_mask = engram_lookups, engram_mask + self.shared_attention_state.reset() + full_num_tokens = positions.shape[0] + if use_sequence_parallel: + if envs.VLLM_MOE_SKIP_PADDING and is_forward_context_available(): + forward_context = get_forward_context() + forward_context.is_padding = sp_padding_mask( + forward_context.is_padding, + hidden_states, + ) + hidden_states = sp_shard(hidden_states) + input_ids = sp_shard(input_ids) + token_mask = sp_shard(token_mask) + lookups = {layer_idx: sp_shard(lookup) for layer_idx, lookup in lookups.items()} + hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1) + pre_mix = hidden_states.new_zeros(hidden_states.shape[0], self.hc_mult, dtype=torch.float32) + pre_mix[:, 0] = 1.0 + last_layer = None + aux_hidden_states = [] + moe_input_ids = input_ids + if self.needs_moe_input_ids: + moe_input_ids = torch.where(input_ids == -1, 0, input_ids) + for layer in self.layers: + last_layer = layer + # DSpark consumes the residual stream entering its configured + # target layers. The runner expresses checkpoint IDs as one-based. + if layer.layer_idx + 1 in self.aux_hidden_state_layers: + aux_hidden_state = hidden_states.mean(dim=1) + if use_sequence_parallel: + aux_hidden_state = sp_all_gather(aux_hidden_state)[:full_num_tokens] + aux_hidden_states.append(aux_hidden_state) + if layer.engram is not None and token_mask.numel(): + n = hidden_states.shape[0] + # Graph captures keep lookup buffers at static capacity; the + # model's actual token dimension remains scheduler-dynamic. + lookup = lookups[layer.layer_idx][:n] + active_mask = token_mask[:n] + kv = layer.engram.wkv(lookup) + key, value = kv.split([self.hc_mult * self.config.hidden_size, self.config.hidden_size], -1) + hidden_states[:n] = engram_gate( + hidden_states[:n], + key.view(n, self.hc_mult, self.config.hidden_size), + value, + layer.engram.q_weight.float() * layer.engram.k_weight.float(), + self.engram_rotation, + active_mask, + self.config.rms_norm_eps, + ) + hidden_states, pre_mix = layer(positions, hidden_states, pre_mix, None, input_ids=moe_input_ids) + assert last_layer is not None + hidden_states = last_layer.hc_collapse(hidden_states, pre_mix) + if use_sequence_parallel: + hidden_states = sp_all_gather(hidden_states)[:full_num_tokens] + hidden_states = self.norm(hidden_states) + if aux_hidden_states: + return hidden_states, aux_hidden_states + return hidden_states + + +class AscendDeepseekV41ForCausalLM(AscendDeepseekV4ForCausalLM): + model_cls = DeepseekV41Model + requires_raw_input_tokens = True + _DEFERRED_WEIGHT_MARKERS: tuple[str, ...] = () + _DEFERRED_WEIGHT_PREFIXES = ("aligner.", "vision.", "image_", "mtp.") + + def prepare_engram_inputs(self, input_ids, positions, padded_tokens=None): + return self.model.prepare_engram_inputs(input_ids, positions, padded_tokens) + + def forward( + self, + input_ids, + positions, + intermediate_tensors=None, + inputs_embeds=None, + engram_lookups=None, + engram_mask=None, + ): + return self.model( + input_ids, + positions, + intermediate_tensors, + inputs_embeds, + engram_lookups=engram_lookups, + engram_mask=engram_mask, + ) + + @classmethod + def _is_milestone_weight(cls, name): + return not name.startswith(cls._DEFERRED_WEIGHT_PREFIXES) and not any( + marker in name for marker in cls._DEFERRED_WEIGHT_MARKERS + ) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + if not get_ascend_config().enable_engram: + return super().load_weights((name, tensor) for name, tensor in weights if ".engram." not in name) + engram_loaded: set[str] = set() + + def milestone_weights() -> Iterator[tuple[str, torch.Tensor]]: + for name, tensor in weights: + if ".engram." in name: + # Bypass V4's generic embed -> embed_tokens remapping and TP loader. + local_name = name.removeprefix("model.") + # FP8/MXFP8 Engram scales are consumed by the CPU loader. + if local_name.endswith(".engram.embed.scale"): + continue + parameter_name = "model." + local_name + if local_name.endswith(".engram.embed.weight"): + layer_id = int(local_name.split(".")[1]) + self.model.layers[layer_id].engram.embed.load_checkpoint(self.model.engram_root, local_name) + else: + param = self.get_parameter(parameter_name) + if tensor.dtype != torch.bfloat16 or tensor.shape != param.shape: + raise ValueError(f"Unexpected BF16 Engram parameter: {name}") + param.data.copy_(tensor) + engram_loaded.add(parameter_name) + elif self._is_milestone_weight(name): + yield name, tensor + + loaded = super().load_weights(milestone_weights()) + expected = {name for name, _ in self.named_parameters() if ".engram." in name} + if engram_loaded != expected: + raise ValueError(f"Missing Engram weights: {expected - engram_loaded}") + return loaded | engram_loaded diff --git a/vllm_ascend/models/deepseek_v41/vl_model.py b/vllm_ascend/models/deepseek_v41/vl_model.py new file mode 100644 index 000000000000..70c868f841db --- /dev/null +++ b/vllm_ascend/models/deepseek_v41/vl_model.py @@ -0,0 +1,231 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Ascend project +"""Ascend multimodal wrapper for DeepSeek V4.1.""" + +import torch +from torch import nn +from vllm.model_executor.models.interfaces import MultiModalEmbeddings +from vllm.model_executor.models.utils import maybe_prefix +from vllm.multimodal import MULTIMODAL_REGISTRY + +from vllm_ascend.models.deepseek_v4.vision import ( + DeepseekV4Aligner, + DeepseekV4ViT, +) +from vllm_ascend.models.deepseek_v4.vl_model import ( + AscendDeepseekV4ForConditionalGeneration, +) + +from .mm_preprocess import ( + IMAGE, + IMAGE_END, + IMAGE_NEW_LINE, + IMAGE_PAD_ID, + IMAGE_PLACEHOLDER, + IMAGE_START, + IMAGE_TOKEN_ID, + DeepseekV41VLDummyInputsBuilder, + DeepseekV41VLMultiModalProcessor, + DeepseekV41VLProcessingInfo, +) +from .model import AscendDeepseekV41ForCausalLM + + +@MULTIMODAL_REGISTRY.register_processor( + DeepseekV41VLMultiModalProcessor, + info=DeepseekV41VLProcessingInfo, + dummy_inputs=DeepseekV41VLDummyInputsBuilder, +) +class AscendDeepseekV41ForConditionalGeneration( + AscendDeepseekV4ForConditionalGeneration, +): + """V4.1 image-span semantics with the shared Ascend vision tower.""" + + language_model_cls = AscendDeepseekV41ForCausalLM + + @classmethod + def get_placeholder_str(cls, modality: str, i: int) -> str | None: + del i + if modality == "image": + return IMAGE_PLACEHOLDER + raise ValueError(f"Unsupported modality: {modality!r}") + + def __init__(self, *, vllm_config, prefix: str = "") -> None: + # Do not call the V4 wrapper constructor: V4 builds an image_pad + # parameter and uses five sentinel roles, neither of which exists in + # the V4.1 checkpoint. + nn.Module.__init__(self) + model_config = vllm_config.model_config + config = model_config.hf_config + if getattr(config, "vision_n_layers", 0) > 0: + config.is_mm_prefix_lm = True + config.mm_prefix_clamp_sliding_window = True + config.mm_prefix_span_leading_pad_modulus = 2 + self.config = config + self.multimodal_config = model_config.multimodal_config + assert self.multimodal_config is not None + + image_enabled = config.vision_n_layers > 0 and self.multimodal_config.get_limit_per_prompt("image") > 0 + with self._mark_tower_model(vllm_config, {"image"}): + self.vision: DeepseekV4ViT | None = None + self.aligner: DeepseekV4Aligner | None = None + self.image_start: nn.Parameter | None = None + self.image_end: nn.Parameter | None = None + self.image_newline: nn.Parameter | None = None + if image_enabled: + self.vision = DeepseekV4ViT(config) + self.aligner = DeepseekV4Aligner(config) + for name in ("image_start", "image_end", "image_newline"): + setattr( + self, + name, + nn.Parameter(torch.empty(config.hidden_size, dtype=torch.float32)), + ) + self.vision.to(dtype=model_config.dtype) + self.aligner.to(dtype=model_config.dtype) + + with self._mark_language_model(vllm_config): + self.language_model = self.language_model_cls( + vllm_config=vllm_config, + prefix=maybe_prefix(prefix, "language_model"), + ) + self.make_empty_intermediate_tensors = self.language_model.make_empty_intermediate_tensors + + def _parse_and_validate_image_input(self, **kwargs: object) -> dict | None: + patches = kwargs.pop("patches", None) + if patches is None: + return None + vit_grid = kwargs.pop("vit_grid", None) + llm_grid = kwargs.pop("llm_grid", None) + types = kwargs.pop("types", None) + if vit_grid is None or llm_grid is None or types is None: + raise ValueError("DeepSeek V4.1 vision input requires patches, vit_grid, llm_grid, and types.") + return { + "patches": patches, + "vit_grid": vit_grid, + "llm_grid": llm_grid, + "types": types, + } + + def _encode_image( + self, + patches: torch.Tensor, + n_vit_h: int, + n_vit_w: int, + ) -> torch.Tensor: + assert self.vision is not None and self.aligner is not None + return self.aligner( + self.vision(patches, n_vit_h, n_vit_w), + n_vit_h, + n_vit_w, + ) + + def _build_image_span( + self, + image_embeds: torch.Tensor, + types: torch.Tensor, + ) -> torch.Tensor: + types = types.to(image_embeds.device) + span = image_embeds.new_empty(types.numel(), image_embeds.shape[-1]) + dtype = image_embeds.dtype + assert self.image_start is not None + assert self.image_end is not None + assert self.image_newline is not None + span[types == IMAGE_START] = self.image_start.to(dtype) + span[types == IMAGE_END] = self.image_end.to(dtype) + span[types == IMAGE_NEW_LINE] = self.image_newline.to(dtype) + span[types == IMAGE] = image_embeds + return span + + def _process_image_input( + self, + patches: torch.Tensor, + vit_grid: torch.Tensor, + llm_grid: torch.Tensor, + types: torch.Tensor, + ) -> tuple[torch.Tensor, ...]: + assert self.aligner is not None + patches = patches.to(self.aligner.w1.weight.dtype) + embeds: list[torch.Tensor] = [] + vit_offset = 0 + span_offset = 0 + for (n_vit_h, n_vit_w), (n_llm_h, n_llm_w) in zip( + vit_grid.tolist(), + llm_grid.tolist(), + strict=True, + ): + n_vit = n_vit_h * n_vit_w + span_len = n_llm_h * (n_llm_w + 1) + 2 + image_embeds = self._encode_image( + patches[vit_offset : vit_offset + n_vit], + n_vit_h, + n_vit_w, + ) + embeds.append( + self._build_image_span( + image_embeds, + types[span_offset : span_offset + span_len], + ) + ) + vit_offset += n_vit + span_offset += span_len + return tuple(embeds) + + def embed_multimodal(self, **kwargs: object) -> MultiModalEmbeddings: + image_input = self._parse_and_validate_image_input(**kwargs) + if image_input is None or self.vision is None: + return [] + return self._process_image_input( + image_input["patches"], + image_input["vit_grid"], + image_input["llm_grid"], + image_input["types"], + ) + + def embed_input_ids( + self, + input_ids: torch.Tensor, + multimodal_embeddings: MultiModalEmbeddings | None = None, + *, + is_multimodal: torch.Tensor | None = None, + ) -> torch.Tensor: + from vllm.model_executor.models.utils import ( + _merge_multimodal_embeddings, + ) + + # The leading alignment row is not an image-feature position. It uses + # the checkpoint's ordinary image-token embedding instead. + embedding_ids = input_ids.masked_fill(input_ids == IMAGE_PAD_ID, IMAGE_TOKEN_ID) + inputs_embeds = self.language_model.embed_input_ids(embedding_ids) + if multimodal_embeddings is None or len(multimodal_embeddings) == 0: + return inputs_embeds + if is_multimodal is None: + raise ValueError("is_multimodal is required when merging image embeddings.") + return _merge_multimodal_embeddings( + inputs_embeds=inputs_embeds, + multimodal_embeddings=multimodal_embeddings, + is_multimodal=is_multimodal, + ) + + def prepare_engram_inputs(self, input_ids, positions, padded_tokens=None): + return self.language_model.prepare_engram_inputs( + input_ids, + positions, + padded_tokens, + ) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + intermediate_tensors=None, + inputs_embeds: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor: + return self.language_model( + input_ids, + positions, + intermediate_tensors, + inputs_embeds, + **kwargs, + ) diff --git a/vllm_ascend/models/minimax_m3/minimax_m3.py b/vllm_ascend/models/minimax_m3/minimax_m3.py index a9e35641d64a..b82b85366292 100644 --- a/vllm_ascend/models/minimax_m3/minimax_m3.py +++ b/vllm_ascend/models/minimax_m3/minimax_m3.py @@ -95,7 +95,6 @@ get_kv_quant_mode, ) -from vllm_ascend.device.device_op import DeviceOperator from vllm_ascend.models.minimax_m3.msa_m3 import ( AscendMiniMaxM3Indexer, AscendMiniMaxM3IndexerLinear, @@ -571,6 +570,11 @@ def forward(self, x: torch.Tensor) -> torch.Tensor | tuple[torch.Tensor, torch.T up = torch.clamp(x[..., d:], min=-self.limit, max=self.limit) activated = gate * torch.sigmoid(self.alpha * gate) * (up + self.beta) if self.use_mx_quant: + # Importing DeviceOperator initializes the Ascend ops package. + # Defer it until execution so model inspection cannot enter the + # ops package through a partially initialized device_op module. + from vllm_ascend.device.device_op import DeviceOperator + quantized_x, scale = DeviceOperator.npu_dynamic_quant( activated, act_quant_type=torch.float8_e4m3fn, diff --git a/vllm_ascend/ops/fused_moe/moe_comm_method.py b/vllm_ascend/ops/fused_moe/moe_comm_method.py index 44a44a4ae9ad..fa0b3364dca9 100644 --- a/vllm_ascend/ops/fused_moe/moe_comm_method.py +++ b/vllm_ascend/ops/fused_moe/moe_comm_method.py @@ -46,20 +46,87 @@ from vllm_ascend.quantization.quant_type import QuantType _MoECommMethods: dict[MoECommType | None, MoECommMethod] = {} +_MoECommMethodsByConfig: dict[tuple[MoECommType | None, tuple[int, ...]], MoECommMethod] = {} -def get_moe_comm_method(moe_comm_type: MoECommType | None) -> MoECommMethod | None: +def _moe_config_key(moe_config: FusedMoEConfig) -> tuple[int, ...]: + """Return the execution shape that owns mutable MoE comm state. + + Target and speculative draft models can have different expert counts and + top-k values. A single process-global dispatcher is therefore unsafe: the + model constructed last would overwrite the dispatcher's shape for every + other MoE layer. Configs with the same execution shape can still share the + existing stateful implementation because their forwards are sequential. + """ + fields = ( + "num_experts", + "num_local_experts", + "experts_per_token", + "hidden_dim", + "intermediate_size_per_partition", + "ep_size", + "tp_size", + "dp_size", + "pcp_size", + ) + return tuple(int(getattr(moe_config, field, 0) or 0) for field in fields) + + +def get_moe_comm_method( + moe_comm_type: MoECommType | None, + moe_config: FusedMoEConfig | None = None, +) -> MoECommMethod | None: + if moe_config is not None: + return _MoECommMethodsByConfig.get((moe_comm_type, _moe_config_key(moe_config))) return _MoECommMethods.get(moe_comm_type) def setup_moe_comm_method(moe_config): + implementations: dict[MoECommType, type[MoECommMethod]] if moe_config.ep_size > 1: - _MoECommMethods[MoECommType.ALLTOALL] = AlltoAllCommImpl(moe_config) - _MoECommMethods[MoECommType.ALLGATHER] = AllGatherCommImpl(moe_config) - _MoECommMethods[MoECommType.MC2] = MC2CommImpl(moe_config) - _MoECommMethods[MoECommType.FUSED_MC2] = FusedMC2CommImpl(moe_config) + implementations = { + MoECommType.ALLTOALL: AlltoAllCommImpl, + MoECommType.ALLGATHER: AllGatherCommImpl, + MoECommType.MC2: MC2CommImpl, + MoECommType.FUSED_MC2: FusedMC2CommImpl, + } else: - _MoECommMethods[MoECommType.ALLGATHER] = AllGatherCommImpl(moe_config) + implementations = {MoECommType.ALLGATHER: AllGatherCommImpl} + + config_key = _moe_config_key(moe_config) + for comm_type, implementation_cls in implementations.items(): + cache_key = (comm_type, config_key) + comm_method = _MoECommMethodsByConfig.get(cache_key) + if comm_method is None: + comm_method = implementation_cls(moe_config) + _MoECommMethodsByConfig[cache_key] = comm_method + # Preserve the legacy lookup for callers that do not own a layer + # config. Layer forwards use the shape-qualified cache below. + _MoECommMethods[comm_type] = comm_method + + +def activate_moe_comm_method(moe_comm_type: MoECommType | None, moe_config: FusedMoEConfig) -> MoECommMethod: + """Bind the communication implementation matching the active MoE layer.""" + matching_methods = [ + method for (comm_type, _), method in _MoECommMethodsByConfig.items() if comm_type == moe_comm_type + ] + if len(matching_methods) <= 1: + # Keep the upstream singleton path unchanged for ordinary models. In + # particular, mutating the forward context from inside a compiled MoE + # forward changes 310P ModelRunner V2 graph behavior. Per-layer + # rebinding is only needed when target and draft expert shapes coexist. + comm_method = get_moe_comm_method(moe_comm_type) + if comm_method is not None: + return comm_method + + comm_method = get_moe_comm_method(moe_comm_type, moe_config) + if comm_method is None: + setup_moe_comm_method(moe_config) + comm_method = get_moe_comm_method(moe_comm_type, moe_config) + if comm_method is None: + raise RuntimeError(f"No MoE communication method registered for {moe_comm_type}") + _EXTRA_CTX.moe_comm_method = comm_method + return comm_method @dataclass diff --git a/vllm_ascend/ops/fused_moe/moe_utils.py b/vllm_ascend/ops/fused_moe/moe_utils.py index ae21682a1f39..e62d7c011fd6 100644 --- a/vllm_ascend/ops/fused_moe/moe_utils.py +++ b/vllm_ascend/ops/fused_moe/moe_utils.py @@ -36,6 +36,22 @@ def async_all_to_all(input_, output_split_sizes, input_split_sizes, group, event=None): + # ``TokenDispatcherWithAll2AllV`` builds split sizes with torch/numpy. + # Normalize them before crossing the torch.distributed boundary: recent + # ProcessGroupHCCL versions do not reliably accept numpy scalar values and + # can report a misleading dimension mismatch instead of a type error. + if input_split_sizes is not None: + input_split_sizes = [int(size) for size in input_split_sizes] + if sum(input_split_sizes) != input_.size(0): + raise RuntimeError( + "MoE all-to-all input split mismatch: " + f"input_shape={tuple(input_.shape)}, " + f"input_splits={input_split_sizes}, " + f"input_split_sum={sum(input_split_sizes)}" + ) + if output_split_sizes is not None: + output_split_sizes = [int(size) for size in output_split_sizes] + if output_split_sizes is None: # Equal split (all2all) a2a_out = torch.empty_like(input_) diff --git a/vllm_ascend/ops/fused_moe/routed_experts.py b/vllm_ascend/ops/fused_moe/routed_experts.py index 30f0c4d71268..1bcc210a881e 100644 --- a/vllm_ascend/ops/fused_moe/routed_experts.py +++ b/vllm_ascend/ops/fused_moe/routed_experts.py @@ -37,7 +37,7 @@ from vllm_ascend.ops.fused_moe.dataclass.fused_experts import MoEWeights, build_fused_experts_input from vllm_ascend.ops.fused_moe.dataclass.moe_mlp import MoEMlpComputeInput from vllm_ascend.ops.fused_moe.force_eplb import get_force_eplb_topk -from vllm_ascend.ops.fused_moe.moe_comm_method import AllGatherCommImpl, FusedExpertsResult +from vllm_ascend.ops.fused_moe.moe_comm_method import AllGatherCommImpl, FusedExpertsResult, activate_moe_comm_method from vllm_ascend.ops.fused_moe.moe_utils import get_moe_num_logical_experts from vllm_ascend.ops.fused_moe.shared_experts import FusedMoEEvents from vllm_ascend.quantization.quant_type import QuantType @@ -622,6 +622,11 @@ def forward_impl( input_ids: torch.Tensor | None = None, ): forward_context = get_forward_context() + # Target and dSPark draft models may use different expert shapes in + # one process (Aurora uses 384/top-6 and 128/top-3 respectively). + # Select the dispatcher owned by this layer before prepare/apply use + # the forward-context communication method. + activate_moe_comm_method(_EXTRA_CTX.moe_comm_type, self.moe_config) # When static kernels are enabled, the forward pass runs twice # (compilation + capture), causing moe_layer_index to overflow. if self.enable_npugraph_ex_static_kernel and forward_context.all_moe_layers: diff --git a/vllm_ascend/ops/fused_moe/router/fused_topk_router.py b/vllm_ascend/ops/fused_moe/router/fused_topk_router.py index c09e002a5453..cc059af227b3 100644 --- a/vllm_ascend/ops/fused_moe/router/fused_topk_router.py +++ b/vllm_ascend/ops/fused_moe/router/fused_topk_router.py @@ -159,6 +159,8 @@ def _compute_routing( if self.tid2eid is not None or self.bias_vl is not None: if input_ids is None: raise ValueError("DeepSeek V4 vision/hash MoE routing requires input_ids.") + # The model sanitizes placeholder IDs once before layer-local + # communication, which only pads with zeros or shards IDs. input_ids = input_ids.to(torch.int64) tid2eid_ones = self.tid2eid.to(torch.int32) if self.tid2eid is not None else None if _EXTRA_CTX.moe_comm_type == MoECommType.ALLGATHER: @@ -172,31 +174,22 @@ def _compute_routing( # ids. Apply the identical TP chunk only when communication # has not already aligned ids with local router rows. input_ids = sequence_parallel_chunk(input_ids.reshape(-1, 1)).reshape(-1) - input_ids = torch.where(input_ids == -1, 0, input_ids) else: input_ids = None tid2eid_ones = None - if self.bias_vl is not None and input_ids is not None: - topk_weights, topk_ids = select_deepseek_v4_vision_experts( - router_logits=router_logits, - input_ids=input_ids, - tid2eid=tid2eid_ones, - bias_vl=self.bias_vl, - text_bias=self.e_score_correction_bias, - top_k=self.top_k, - renormalize=self.renormalize, - routed_scaling_factor=self.routed_scaling_factor, - image_sentinel_lo=self.image_sentinel_lo, - ) - return topk_weights.to(torch.float32), topk_ids.to( - torch.int32 if indices_type is None else indices_type - ) + bias_vl = self.bias_vl + if bias_vl is not None and bias_vl.dtype != router_logits.dtype: + bias_vl = bias_vl.to(router_logits.dtype) + text_bias = self.e_score_correction_bias + if text_bias is not None and text_bias.dtype != router_logits.dtype: + text_bias = text_bias.to(router_logits.dtype) topk_weights, topk_ids, _ = torch.ops._C_ascend.moe_gating_top_k_hash( x=router_logits, k=self.top_k, - bias=self.e_score_correction_bias, + bias=text_bias, input_ids=input_ids, tid2eid=tid2eid_ones, + bias_vl=bias_vl, k_group=topk_group, group_count=num_expert_group, routed_scaling_factor=self.routed_scaling_factor, @@ -207,8 +200,10 @@ def _compute_routing( renorm=0, norm_type=2, out_flag=False, + image_sentinel_lo=self.image_sentinel_lo, + image_sentinel_count=DEEPSEEK_V4_IMAGE_SENTINEL_COUNT, ) - return topk_weights, topk_ids + return topk_weights.to(torch.float32), topk_ids.to(torch.int32 if indices_type is None else indices_type) norm_type = 0 if self.scoring_func == "softmax" else 1 if self.e_score_correction_bias is not None and self.e_score_correction_bias.dtype != router_logits.dtype: self.e_score_correction_bias = self.e_score_correction_bias.to(router_logits.dtype) diff --git a/vllm_ascend/ops/rope_dsv4.py b/vllm_ascend/ops/rope_dsv4.py index f193d3a2e362..48d40524b1fa 100644 --- a/vllm_ascend/ops/rope_dsv4.py +++ b/vllm_ascend/ops/rope_dsv4.py @@ -174,6 +174,26 @@ def get_full_cos_and_sin_dsa(group_name: str) -> tuple[torch.Tensor, torch.Tenso return _ROPE_STATE.full_rope_cache[config_key] +def get_full_cos_and_sin_dsa_for_layer( + layer_name: str, +) -> tuple[torch.Tensor, torch.Tensor]: + """Return the full RoPE cache selected by one registered layer. + + A group name is not sufficient for V4.1 because pure-SWA and long-context + layers both use the ``default`` group while registering different RoPE + configurations. Resolving through ``layer_info`` keeps the compressor + metadata path tied to the exact table used by its source attention layer. + """ + info = _ROPE_STATE.layer_info.get(layer_name) + if info is None: + raise KeyError(f"Layer {layer_name} is not registered.") + config_key, _ = info + try: + return _ROPE_STATE.full_rope_cache[config_key] + except KeyError as exc: + raise KeyError(f"Rope cache for layer {layer_name} is not initialized.") from exc + + class ComplexExpRotaryEmbedding(nn.Module): def __init__( self, @@ -194,9 +214,11 @@ def __init__( self.rotary_dim = rotary_dim beta_fast = extra_kwargs.get("beta_fast", 32) beta_slow = extra_kwargs.get("beta_slow", 1) + original_seq_len = extra_kwargs.get("original_seq_len", max_position_embeddings) config_key = ( f"rotary_dim{rotary_dim}_max_position_embeddings{max_position_embeddings}_" - f"base{base}_scaling_factor{scaling_factor}_beta_fast{beta_fast}_beta_slow{beta_slow}" + f"original_seq_len{original_seq_len}_base{base}_scaling_factor{scaling_factor}_" + f"beta_fast{beta_fast}_beta_slow{beta_slow}" ) _ROPE_STATE.layer_info[layername] = (config_key, rope_groups) @@ -207,7 +229,7 @@ def __init__( if config_key not in _ROPE_STATE.full_rope_cache: inv_freq = self.precompute_freqs_cis( - rotary_dim, max_position_embeddings, max_position_embeddings, base, scaling_factor, beta_fast, beta_slow + rotary_dim, max_position_embeddings, original_seq_len, base, scaling_factor, beta_fast, beta_slow ) t = torch.arange( max_position_embeddings * scaling_factor, @@ -288,6 +310,8 @@ def yarn_linear_ramp_mask(low: float, high: float, dim: int, dtype: torch.dtype) pos_freqs = base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim) inv_freq_extrapolation = 1.0 / pos_freqs + if original_seq_len <= 0: + return inv_freq_extrapolation inv_freq_interpolation = 1.0 / (factor * pos_freqs) low, high = yarn_find_correction_range( diff --git a/vllm_ascend/ops/triton/compressor/__init__.py b/vllm_ascend/ops/triton/compressor/__init__.py new file mode 100644 index 000000000000..208f01a7cb5e --- /dev/null +++ b/vllm_ascend/ops/triton/compressor/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project diff --git a/vllm_ascend/ops/triton/compressor/compressor_triton.py b/vllm_ascend/ops/triton/compressor/compressor_triton.py new file mode 100644 index 000000000000..b1bb4b051299 --- /dev/null +++ b/vllm_ascend/ops/triton/compressor/compressor_triton.py @@ -0,0 +1,717 @@ +"""Compressor 算子 triton 实现。 + +kernel 结构: + K1 投影 GEMM ×2(bf16 权重,swizzle,block 由 autotune 按 M 择优) + K2 组池化(变长/残余拼接,两遍 softmax;段首残余从 cache 读取) + K3 cache 更新(只写每段最后 cache_size 个 token,每 slot 唯一写者) +K2 必须先于 K3 执行:段尾写入与段首残余读取会命中同一 slot,并发时读取 +结果会被本段写入覆盖。 +compressor_ref 为纯 torch 参考实现。 +""" + +import numpy as np +import torch +from vllm.triton_utils import tl, triton + +DEV = "npu" +MAX_CHUNK_ROWS = 8 # 池化组内分块行数上限(UB 容量:三遍循环中间量 <192KB) + + +# ============================================================================ +# K1: 投影 GEMM(bf16 权重,swizzle,fp32 输出) +# ============================================================================ + +# autotune 候选(key=['M'])。BK 影响 fp32 累加分组(ulp 级差异),在精度容差内。 +_proj_tune_configs = [ + triton.Config({"BM": 128, "BN": 256, "BK": 256, "GROUP": 8}), + triton.Config({"BM": 128, "BN": 256, "BK": 128, "GROUP": 8}), + triton.Config({"BM": 64, "BN": 256, "BK": 256, "GROUP": 8}), + triton.Config({"BM": 64, "BN": 256, "BK": 128, "GROUP": 8}), + triton.Config({"BM": 128, "BN": 128, "BK": 256, "GROUP": 8}), + triton.Config({"BM": 64, "BN": 128, "BK": 256, "GROUP": 8}), + triton.Config({"BM": 256, "BN": 128, "BK": 256, "GROUP": 8}), + triton.Config({"BM": 64, "BN": 256, "BK": 256, "GROUP": 4}), +] + + +@triton.autotune(configs=_proj_tune_configs, key=["M"]) +@triton.jit +def _proj_kernel( + out_ptr, + x_ptr, + w_ptr, + M, + IN_DIM: tl.constexpr, + OUT_DIM: tl.constexpr, + BM: tl.constexpr, + BN: tl.constexpr, + BK: tl.constexpr, + GROUP: tl.constexpr, +): + """x (M, IN_DIM) @ w^T,w 为 (OUT_DIM, IN_DIM) 行主序,输出 fp32。""" + pid = tl.program_id(0) + num_progs = tl.num_programs(0) + offs_m = tl.arange(0, BM) + offs_n = tl.arange(0, BN) + n_tile_count = OUT_DIM // BN + m_tile_count = tl.cdiv(M, BM) + total_tiles = m_tile_count * n_tile_count + tiles_per_group = GROUP * n_tile_count + for tile_id in range(pid, total_tiles, num_progs): + swizzle_group = tile_id // tiles_per_group + group_m_start = swizzle_group * GROUP + group_m_tiles = m_tile_count - group_m_start if (m_tile_count - group_m_start) < GROUP else GROUP + rank_in_group = tile_id - swizzle_group * tiles_per_group + m_tile = group_m_start + rank_in_group % group_m_tiles + n_tile = rank_in_group // group_m_tiles + rows = m_tile * BM + offs_m + cols = n_tile * BN + offs_n + row_valid = rows < M + acc = tl.zeros((BM, BN), dtype=tl.float32) + for k0 in range(0, IN_DIM, BK): + offs_k = k0 + tl.arange(0, BK) + x = tl.load(x_ptr + rows[:, None] * IN_DIM + offs_k[None, :], mask=row_valid[:, None], other=0.0) + w = tl.load(w_ptr + offs_k[:, None] + cols[None, :] * IN_DIM) + acc = tl.dot(x, w, acc) + tl.store(out_ptr + rows[:, None] * OUT_DIM + cols[None, :], acc, mask=row_valid[:, None]) + + +# ============================================================================ +# K2: 组池化(task_id -> (batch, local_group) O(1) 除法解码;控制量 kernel 内 load) +# ============================================================================ + + +@triton.jit +def _pooled_blocked( + group_idx, + start_pos, + seg_row_base, + cache_row, + kv_ptr, + score_ptr, + cache_ptr, + offs_h, + CACHE_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + RATIO: tl.constexpr, + RATIO_PAD: tl.constexpr, + CHUNK_ROWS: tl.constexpr, + HAS_RES, +): + """组内分块三遍循环(max / denom / 加权和)。HAS_RES=0 时不访问 cache。 + + 组覆盖全局 token [group_idx*RATIO, (group_idx+1)*RATIO):seg_off >= 0 读 + 投影大矩阵 (seg_row_base+seg_off),< 0 为段前残余,读 cache + (cache_row, token_pos % CACHE_SIZE)。""" + # Keep masked history lanes inside the projection allocation. Ascend may + # form a DMA address for a masked lane before applying the load mask. + # The original masks and ring selection still supply all history values. + score_max = tl.full((HEAD_DIM,), -float("inf"), dtype=tl.float32) + for chunk0 in range(0, RATIO_PAD, CHUNK_ROWS): + rows = chunk0 + tl.arange(0, CHUNK_ROWS) + row_valid = rows < RATIO + token_pos = group_idx * RATIO + rows + seg_off = token_pos - start_pos + in_seg = seg_off >= 0 + score_seg = tl.load( + score_ptr + (seg_row_base + tl.maximum(seg_off, 0))[:, None] * HEAD_DIM + offs_h[None, :], + mask=(row_valid & in_seg)[:, None], + other=0.0, + ) + if HAS_RES: + slot = token_pos % CACHE_SIZE + score_cache = tl.load( + cache_ptr + (cache_row * CACHE_SIZE + slot[:, None]) * 2 * HEAD_DIM + HEAD_DIM + offs_h[None, :], + mask=(row_valid & (seg_off < 0))[:, None], + other=0.0, + ) + score = tl.where(in_seg[:, None], score_seg, score_cache) + else: + score = score_seg + score = tl.where(row_valid[:, None], score, -float("inf")) + score_max = tl.maximum(score_max, tl.max(score, axis=0)) + exp_sum = tl.zeros((HEAD_DIM,), dtype=tl.float32) + for chunk0 in range(0, RATIO_PAD, CHUNK_ROWS): + rows = chunk0 + tl.arange(0, CHUNK_ROWS) + row_valid = rows < RATIO + token_pos = group_idx * RATIO + rows + seg_off = token_pos - start_pos + in_seg = seg_off >= 0 + score_seg = tl.load( + score_ptr + (seg_row_base + tl.maximum(seg_off, 0))[:, None] * HEAD_DIM + offs_h[None, :], + mask=(row_valid & in_seg)[:, None], + other=0.0, + ) + if HAS_RES: + slot = token_pos % CACHE_SIZE + score_cache = tl.load( + cache_ptr + (cache_row * CACHE_SIZE + slot[:, None]) * 2 * HEAD_DIM + HEAD_DIM + offs_h[None, :], + mask=(row_valid & (seg_off < 0))[:, None], + other=0.0, + ) + score = tl.where(in_seg[:, None], score_seg, score_cache) + else: + score = score_seg + score = tl.where(row_valid[:, None], score, -float("inf")) + exp_sum += tl.sum(tl.exp(score - score_max[None, :]), axis=0) + pooled = tl.zeros((HEAD_DIM,), dtype=tl.float32) + for chunk0 in range(0, RATIO_PAD, CHUNK_ROWS): + rows = chunk0 + tl.arange(0, CHUNK_ROWS) + row_valid = rows < RATIO + token_pos = group_idx * RATIO + rows + seg_off = token_pos - start_pos + in_seg = seg_off >= 0 + score_seg = tl.load( + score_ptr + (seg_row_base + tl.maximum(seg_off, 0))[:, None] * HEAD_DIM + offs_h[None, :], + mask=(row_valid & in_seg)[:, None], + other=0.0, + ) + if HAS_RES: + slot = token_pos % CACHE_SIZE + score_cache = tl.load( + cache_ptr + (cache_row * CACHE_SIZE + slot[:, None]) * 2 * HEAD_DIM + HEAD_DIM + offs_h[None, :], + mask=(row_valid & (seg_off < 0))[:, None], + other=0.0, + ) + score = tl.where(in_seg[:, None], score_seg, score_cache) + else: + score = score_seg + score = tl.where(row_valid[:, None], score, -float("inf")) + prob = tl.exp(score - score_max[None, :]) / exp_sum[None, :] + kv_seg = tl.load( + kv_ptr + (seg_row_base + tl.maximum(seg_off, 0))[:, None] * HEAD_DIM + offs_h[None, :], + mask=(row_valid & in_seg)[:, None], + other=0.0, + ) + if HAS_RES: + slot = token_pos % CACHE_SIZE + kv_cache = tl.load( + cache_ptr + (cache_row * CACHE_SIZE + slot[:, None]) * 2 * HEAD_DIM + offs_h[None, :], + mask=(row_valid & (seg_off < 0))[:, None], + other=0.0, + ) + kv_vals = tl.where(in_seg[:, None], kv_seg, kv_cache) + else: + kv_vals = kv_seg + pooled += tl.sum(prob * kv_vals, axis=0) + return pooled + + +@triton.jit +def _pool_kernel( + out_ptr, + kv_ptr, + score_ptr, + cache_ptr, + norm_w_ptr, + meta_ptr, # 拼包 [start_pos | used_len | out_row_offset | seg_row_base] + block_table_ptr, + total_group_slots, + GROUP_SLOTS_PER_BATCH, # 每 batch 组槽位上界 ceil(max_used_len / RATIO) + NUM_BATCH, + CACHE_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + RATIO: tl.constexpr, + RATIO_PAD: tl.constexpr, + CHUNK_ROWS: tl.constexpr, + eps: tl.constexpr, + SINGLE_BLOCK: tl.constexpr, + NO_PAD: tl.constexpr, + TOKEN_ALIGNED: tl.constexpr = False, + NUM_CACHE_BLOCKS: tl.constexpr = 0, +): + pid = tl.program_id(0) + offs_h = tl.arange(0, HEAD_DIM) + if not TOKEN_ALIGNED: + norm_w = tl.load(norm_w_ptr + offs_h).to(tl.float32) + for task_id in range(pid, total_group_slots, tl.num_programs(0)): + batch = task_id // GROUP_SLOTS_PER_BATCH + local_group = task_id - batch * GROUP_SLOTS_PER_BATCH + start_pos = tl.load(meta_ptr + batch) # [0, NUM_BATCH) + used_len = tl.load(meta_ptr + NUM_BATCH + batch) # [NUM_BATCH, 2*NUM_BATCH) + groups_in_batch = (start_pos + used_len) // RATIO - start_pos // RATIO + cache_row = tl.load(block_table_ptr + batch) + valid = local_group < groups_in_batch + if TOKEN_ALIGNED: + valid = valid & (cache_row > 0) & (cache_row < NUM_CACHE_BLOCKS) + if valid: + group_idx = start_pos // RATIO + local_group + seg_row_base = tl.load(meta_ptr + 3 * NUM_BATCH + batch) # [3*NUM_BATCH, 4*NUM_BATCH) + rows = tl.arange(0, CHUNK_ROWS) + token_pos = group_idx * RATIO + rows + seg_off = token_pos - start_pos + residual = start_pos - group_idx * RATIO # 段前残余数,仅 local_group==0 时可能 >0 + cache_row = tl.load(block_table_ptr + batch) + if residual > 0: + # 残余组(每 batch 至多 1 个):cache 拼接 + pooled = _pooled_blocked( + group_idx, + start_pos, + seg_row_base, + cache_row, + kv_ptr, + score_ptr, + cache_ptr, + offs_h, + CACHE_SIZE, + HEAD_DIM, + RATIO, + RATIO_PAD, + CHUNK_ROWS, + True, + ) + else: + # 非残余组:单块直读;RATIO 超出单块容量时走分块版 + if SINGLE_BLOCK: + if NO_PAD: + score = tl.load(score_ptr + (seg_row_base + seg_off)[:, None] * HEAD_DIM + offs_h[None, :]) + score_max = tl.max(score, axis=0) + e = tl.exp(score - score_max[None, :]) + exp_sum = tl.sum(e, axis=0) + prob = e / exp_sum[None, :] + kv_vals = tl.load(kv_ptr + (seg_row_base + seg_off)[:, None] * HEAD_DIM + offs_h[None, :]) + else: + row_valid = rows < RATIO + score = tl.load( + score_ptr + (seg_row_base + seg_off)[:, None] * HEAD_DIM + offs_h[None, :], + mask=row_valid[:, None], + other=-float("inf"), + ) + score_max = tl.max(score, axis=0) + e = tl.exp(score - score_max[None, :]) + e = tl.where(row_valid[:, None], e, 0.0) + exp_sum = tl.sum(e, axis=0) + prob = e / exp_sum[None, :] + kv_vals = tl.load( + kv_ptr + (seg_row_base + seg_off)[:, None] * HEAD_DIM + offs_h[None, :], + mask=row_valid[:, None], + other=0.0, + ) + pooled = tl.sum(prob * kv_vals, axis=0) + else: + pooled = _pooled_blocked( + group_idx, + start_pos, + seg_row_base, + cache_row, + kv_ptr, + score_ptr, + cache_ptr, + offs_h, + CACHE_SIZE, + HEAD_DIM, + RATIO, + RATIO_PAD, + CHUNK_ROWS, + False, + ) + if TOKEN_ALIGNED: + # Keep Aurora's existing RMSNorm outside this kernel. Place + # each completed pair at its original completion-token row. + out_row = seg_row_base + (group_idx + 1) * RATIO - 1 - start_pos + tl.store(out_ptr + out_row * HEAD_DIM + offs_h, pooled.to(tl.bfloat16)) + else: + pooled_bf16 = pooled.to(tl.bfloat16).to(tl.float32) + var = tl.sum(pooled_bf16 * pooled_bf16, axis=0) / HEAD_DIM + norm_out = pooled_bf16 * (1.0 / tl.sqrt(var + eps)) * norm_w + out_row = tl.load(meta_ptr + 2 * NUM_BATCH + batch) + local_group + tl.store(out_ptr + out_row * HEAD_DIM + offs_h, norm_out.to(tl.bfloat16)) + + +# ============================================================================ +# K3: cache 更新(只写每段最后 min(used_len, CACHE_SIZE) 个 token) +# ============================================================================ + + +@triton.jit +def _cache_update_kernel( + kv_ptr, + score_ptr, + cache_ptr, + meta_ptr, + block_table_ptr, + total_write_slots, + WRITE_SLOTS_PER_BATCH, # 每 batch 写槽位上界 min(max_used_len, CACHE_SIZE) + NUM_BATCH, + CACHE_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + PROTECT_NULL: tl.constexpr = False, + NUM_CACHE_BLOCKS: tl.constexpr = 0, +): + pid = tl.program_id(0) + offs_h = tl.arange(0, HEAD_DIM) + for task_id in range(pid, total_write_slots, tl.num_programs(0)): + batch = task_id // WRITE_SLOTS_PER_BATCH + write_idx = task_id - batch * WRITE_SLOTS_PER_BATCH + used_len = tl.load(meta_ptr + NUM_BATCH + batch) + tail_count = used_len if used_len < CACHE_SIZE else CACHE_SIZE + cache_row = tl.load(block_table_ptr + batch) + valid = write_idx < tail_count + if PROTECT_NULL: + valid = valid & (cache_row > 0) & (cache_row < NUM_CACHE_BLOCKS) + if valid: + token_in_seg = used_len - tail_count + write_idx # 最后 tail_count 个 token 中的第 write_idx 个 + proj_row = tl.load(meta_ptr + 3 * NUM_BATCH + batch) + token_in_seg + token_pos = tl.load(meta_ptr + batch) + token_in_seg + slot = token_pos % CACHE_SIZE + cache_row = tl.load(block_table_ptr + batch) + kv_vals = tl.load(kv_ptr + proj_row * HEAD_DIM + offs_h) + score_vals = tl.load(score_ptr + proj_row * HEAD_DIM + offs_h) + tl.store(cache_ptr + (cache_row * CACHE_SIZE + slot) * 2 * HEAD_DIM + offs_h, kv_vals) + tl.store(cache_ptr + (cache_row * CACHE_SIZE + slot) * 2 * HEAD_DIM + HEAD_DIM + offs_h, score_vals) + + +# ============================================================================ +# Host:布局归一、控制量拼包、launch +# ============================================================================ + + +def _ints(v): + """控制量解析:list/tuple/np.ndarray 直取;torch.Tensor 走一次 D2H(仅兼容)。""" + if v is None: + return None + if isinstance(v, torch.Tensor): + return v.detach().cpu().tolist() + if isinstance(v, np.ndarray): + return v.tolist() + return list(v) + + +# 控制量拼包的异步上传环:pinned + non_blocking 拷贝。 +# from_numpy().to(dev) 为同步 H2D,会阻塞至此前入队的全部 kernel 完成。 +_META_UPLOAD_POOLS = {} # (device, 拼包元素数) → _MetaUploadRing +_META_RING_SIZE = 32 # 槽位数,须大于 host 领先 device 的最大调用数 + + +class _MetaUploadRing: + """每槽位含锁页中转、device 副本与上传完成事件;复用槽位前等待其事件, + 避免写入中转缓冲时上一次 DMA 尚未完成。""" + + def __init__(self, n_int, dev): + self.pinned_staging = [torch.empty(n_int, dtype=torch.int32).pin_memory() for _ in range(_META_RING_SIZE)] + self.device_slots = [torch.empty(n_int, dtype=torch.int32, device=dev) for _ in range(_META_RING_SIZE)] + self.copy_done_events = [torch.npu.Event() for _ in range(_META_RING_SIZE)] + self.uploads_launched = 0 + + def upload(self, meta_pack): + slot = self.uploads_launched % _META_RING_SIZE + self.uploads_launched += 1 + if self.uploads_launched > _META_RING_SIZE: + self.copy_done_events[slot].synchronize() + self.pinned_staging[slot].numpy()[:] = meta_pack + self.device_slots[slot].copy_(self.pinned_staging[slot], non_blocking=True) + self.copy_done_events[slot].record() + return self.device_slots[slot] + + +def _meta_upload(meta_np, dev): + pool_key = (str(dev), meta_np.size) + if pool_key not in _META_UPLOAD_POOLS: + _META_UPLOAD_POOLS[pool_key] = _MetaUploadRing(meta_np.size, dev) + return _META_UPLOAD_POOLS[pool_key].upload(meta_np) + + +_CUBE_CORE_NUM = None # AI core 数,首次查询后缓存 + + +def _cube_core_num(): + """设备 AI core 数(grid 并行单位)。 + + multi_processor_count 为 vector core 数(num_aicore 的 2 倍),不可用作上限。""" + global _CUBE_CORE_NUM + if _CUBE_CORE_NUM is None: + driver = triton.runtime.driver + + _CUBE_CORE_NUM = driver.active.utils.get_device_properties(torch.npu.current_device())["num_aicore"] + return _CUBE_CORE_NUM + + +def compressor( + x, + wkv, + wgate, + state_cache, + cmp_ratio, + norm_w, + state_block_table=None, + cu_seqlens=None, + seqused=None, + start_pos=None, + num_cores=None, +): + """返回 cmp_kv(bf16)。 + + cu_seqlens / seqused / start_pos 为 host 控制量(推荐 list[int]/np.ndarray, + tensor 兼容但有 D2H 同步);state_cache / state_block_table 为 device tensor。 + num_cores 为并行核数,缺省取设备 AI core 数。 + """ + hidden_dim = x.shape[-1] + head_dim = wkv.shape[0] + assert hidden_dim == 5120 and head_dim == 512, f"规格约束 hidden=5120/D=512, got {hidden_dim}/{head_dim}" + ratio = int(cmp_ratio) + assert 2 <= ratio <= 128 + # batch 数与每段长度(布局归一) + cu_list = _ints(cu_seqlens) + if cu_list is not None: + num_batch = len(cu_list) - 1 + seg_len = [cu_list[i + 1] - cu_list[i] for i in range(num_batch)] + seg_row_base = [cu_list[i] for i in range(num_batch)] + max_seg_len = max(seg_len) + is_packed = True + else: + num_batch = x.shape[0] + seg_len = [x.shape[1]] * num_batch + seg_row_base = [i * x.shape[1] for i in range(num_batch)] + max_seg_len = x.shape[1] + is_packed = False + total_tokens = x.shape[0] if is_packed else num_batch * max_seg_len + used_lens = _ints(seqused) if seqused is not None else list(seg_len) + start_pos_list = _ints(start_pos) if start_pos is not None else [0] * num_batch + cache_size = state_cache.shape[1] + + # 槽位网格:task_id -> (batch, 槽位) O(1) 除法解码,越界槽位 kernel 内跳过; + # 控制量拼包 [start_pos | used_len | out_row_offset | seg_row_base | batch_idx] 一次 H2D + max_used_len = max(used_lens) + group_slots_per_batch = (max_used_len + ratio - 1) // ratio + write_slots_per_batch = min(max_used_len, cache_size) + total_group_slots = num_batch * group_slots_per_batch # ≥ Σ组数 + total_write_slots = num_batch * write_slots_per_batch + dev = x.device + # 每段有效组数与输出布局 + group_counts = [(spb + used_lens[b]) // ratio - spb // ratio for b, spb in enumerate(start_pos_list)] + out_rows_per_batch = (max_seg_len + ratio - 1) // ratio + if is_packed: + out_row_offsets, running = [], 0 + for b in range(num_batch): + out_row_offsets.append(running) + running += group_counts[b] + out = torch.zeros( + min(total_tokens, total_tokens // ratio + num_batch), head_dim, dtype=torch.bfloat16, device=dev + ) + else: + out_row_offsets = [b * out_rows_per_batch for b in range(num_batch)] + out = torch.zeros(num_batch, out_rows_per_batch, head_dim, dtype=torch.bfloat16, device=dev) + out = out.view(-1, head_dim) + meta_pack = np.concatenate( + [ + np.asarray(start_pos_list, dtype=np.int32), + np.asarray(used_lens, dtype=np.int32), + np.asarray(out_row_offsets, dtype=np.int32), + np.asarray(seg_row_base, dtype=np.int32), + np.asarray(range(num_batch), dtype=np.int32), + ] + ) + meta_dev = _meta_upload(meta_pack, dev) + block_table_ptr = ( + state_block_table.to(torch.int32).to(dev) + if state_block_table is not None + else meta_dev[4 * num_batch : 5 * num_batch] + ) + + # K1 投影 ×2 + kv = torch.empty(total_tokens, head_dim, dtype=torch.float32, device=dev) + score = torch.empty(total_tokens, head_dim, dtype=torch.float32, device=dev) + x_2d = x if is_packed else x.view(total_tokens, hidden_dim) + cores = num_cores or _cube_core_num() + gemm_grid = lambda META: (min(cores, triton.cdiv(total_tokens, META["BM"]) * (head_dim // META["BN"])),) + _proj_kernel[gemm_grid](kv, x_2d, wkv, total_tokens, IN_DIM=hidden_dim, OUT_DIM=head_dim) + _proj_kernel[gemm_grid](score, x_2d, wgate, total_tokens, IN_DIM=hidden_dim, OUT_DIM=head_dim) + + # K2 组池化 → K3 cache 更新(顺序执行,先读后写) + padded_ratio = max(2, triton.next_power_of_2(ratio)) + if total_group_slots > 0: + pool_grid = (min(cores, total_group_slots),) + _pool_kernel[pool_grid]( + out, + kv, + score, + state_cache, + norm_w, + meta_dev, + block_table_ptr, + total_group_slots, + group_slots_per_batch, + num_batch, + CACHE_SIZE=cache_size, + HEAD_DIM=head_dim, + RATIO=ratio, + RATIO_PAD=padded_ratio, + CHUNK_ROWS=min(padded_ratio, MAX_CHUNK_ROWS), + eps=1e-20, + SINGLE_BLOCK=(padded_ratio <= MAX_CHUNK_ROWS), + NO_PAD=(padded_ratio == ratio), + ) + if total_write_slots > 0: + cache_grid = (min(cores, total_write_slots),) + _cache_update_kernel[cache_grid]( + kv, + score, + state_cache, + meta_dev, + block_table_ptr, + total_write_slots, + write_slots_per_batch, + num_batch, + CACHE_SIZE=cache_size, + HEAD_DIM=head_dim, + ) + shape = (num_batch, out_rows_per_batch, head_dim) if not is_packed else (out.shape[0], head_dim) + return out.view(*shape) + + +# ============================================================================ +# 参考实现(纯 torch,逐 batch;cache 原地更新与实现一致) +# ============================================================================ + + +def compressor_ref( + x, wkv, wgate, state_cache, cmp_ratio, norm_w, state_block_table=None, cu_seqlens=None, seqused=None, start_pos=None +): + head_dim = wkv.shape[0] + ratio = int(cmp_ratio) + if cu_seqlens is not None: + cu_list = cu_seqlens.cpu().tolist() + num_batch = len(cu_list) - 1 + segs = [x[cu_list[i] : cu_list[i + 1]] for i in range(num_batch)] + else: + num_batch = x.shape[0] + segs = [x[i] for i in range(num_batch)] + used_lens = seqused.cpu().tolist() if seqused is not None else [s.shape[0] for s in segs] + start_pos_list = _ints(start_pos) if start_pos is not None else [0] * num_batch + block_table = state_block_table.cpu().tolist() if state_block_table is not None else list(range(num_batch)) + cache_size = state_cache.shape[1] + batch_outputs, group_counts_all = [], [] + for b in range(num_batch): + xs = segs[b][: used_lens[b]] + kv = torch.nn.functional.linear(xs.float(), wkv.float()) + scores = torch.nn.functional.linear(xs.float(), wgate.float()) + seg_start_pos = start_pos_list[b] + first_group = seg_start_pos // ratio + residual = seg_start_pos - first_group * ratio + residual_rows = [] + for j in range(residual): + pos = first_group * ratio + j + slot = pos % cache_size + residual_rows.append(state_cache[block_table[b], slot]) + seg_rows = ( + torch.stack([torch.cat([k, s]) for k, s in zip(kv, scores)]) + if len(kv) + else torch.empty(0, 2 * head_dim, dtype=torch.float32, device=x.device) + ) + tokens = torch.cat( + [ + torch.stack(residual_rows) + if residual_rows + else torch.empty(0, 2 * head_dim, dtype=torch.float32, device=x.device), + seg_rows, + ] + ) + token_count = tokens.shape[0] + group_count = token_count // ratio + out_batch = torch.empty(group_count, head_dim, dtype=torch.bfloat16, device=x.device) + for g in range(group_count): + group_tokens = tokens[g * ratio : (g + 1) * ratio] + kv_part = group_tokens[:, :head_dim] + score_part = group_tokens[:, head_dim:] + prob = score_part.softmax(dim=0) + pooled = (kv_part * prob).sum(dim=0) + pooled_bf16 = pooled.to(torch.bfloat16).float() + var = pooled_bf16.square().mean() + out_batch[g] = (pooled_bf16 * torch.rsqrt(var + 1e-20) * norm_w.float()).to(torch.bfloat16) + batch_outputs.append(out_batch) + group_counts_all.append(group_count) + for t in range(used_lens[b]): + pos = seg_start_pos + t + state_cache[block_table[b], pos % cache_size, :head_dim] = kv[t] + state_cache[block_table[b], pos % cache_size, head_dim:] = scores[t] + if cu_seqlens is None: + seg_len_max = x.shape[1] + rows_per_batch = (seg_len_max + ratio - 1) // ratio + out_full = torch.zeros(num_batch, rows_per_batch, head_dim, dtype=torch.bfloat16, device=x.device) + for b in range(num_batch): + out_full[b, : group_counts_all[b]] = batch_outputs[b] + return out_full + out_offset = 0 + out_full = torch.zeros( + min(x.shape[0], x.shape[0] // ratio + num_batch), head_dim, dtype=torch.bfloat16, device=x.device + ) + for b in range(num_batch): + out_full[out_offset : out_offset + group_counts_all[b]] = batch_outputs[b] + out_offset += group_counts_all[b] + return out_full + + +def compressor_from_projected(kv, scores, state_cache, metadata, out, *, max_query_len, num_cores): + """Pool C2 into token-aligned BF16 rows, then update a private FP32 ring. + + metadata is contiguous device INT32 [5, requests]: start positions, used + lengths, output bases (reserved), input bases, and actual global block IDs. + Controls must be bounded by the input rows and request table; padded requests + have used length and block ID zero. No tensor contents are read by the host. + """ + if kv.dtype != torch.float32 or scores.dtype != torch.float32 or state_cache.dtype != torch.float32: + raise ValueError("Aurora projections and ring state must be FP32") + if kv.ndim != 2 or scores.shape != kv.shape: + raise ValueError("Expected matching [tokens, width] projections") + tokens, width = kv.shape + if width < 1 or width & (width - 1): + raise ValueError("Compressor width must be a positive power of two") + if state_cache.ndim != 3 or state_cache.shape[1:] != (32, 2 * width): + raise ValueError("Expected [blocks, 32, 2*width] ring state") + if out.dtype != torch.bfloat16 or out.shape != kv.shape: + raise ValueError("Expected token-aligned BF16 pooled output") + if metadata.dtype != torch.int32 or metadata.ndim != 2 or metadata.shape[0] != 5: + raise ValueError("Expected INT32 [5, requests] device metadata") + tensors = (kv, scores, state_cache, metadata, out) + if not all(t.is_contiguous() for t in tensors): + raise ValueError("Compressor requires contiguous projections, ring pages and controls") + if not all(t.device == kv.device for t in tensors): + raise ValueError("Compressor tensors must share one device") + if num_cores <= 0 or not 0 <= max_query_len <= tokens: + raise ValueError("Invalid compressor launch bounds") + out.zero_() + batches = metadata.shape[1] + if tokens == 0 or batches == 0 or max_query_len == 0: + return out + if kv.device.type != "npu": + raise ValueError("Projected Triton compressor requires an NPU") + groups = (max_query_len + 1) // 2 + writes = min(max_query_len, 32) + # Same-stream launch order is required: residual reads precede ring writes. + _pool_kernel[(min(num_cores, batches * groups),)]( + out, + kv, + scores, + state_cache, + out, + metadata, + metadata[4], + batches * groups, + groups, + batches, + CACHE_SIZE=32, + HEAD_DIM=width, + RATIO=2, + RATIO_PAD=2, + CHUNK_ROWS=2, + eps=0.0, + SINGLE_BLOCK=True, + NO_PAD=True, + TOKEN_ALIGNED=True, + NUM_CACHE_BLOCKS=state_cache.shape[0], + ) + _cache_update_kernel[(min(num_cores, batches * writes),)]( + kv, + scores, + state_cache, + metadata, + metadata[4], + batches * writes, + writes, + batches, + CACHE_SIZE=32, + HEAD_DIM=width, + PROTECT_NULL=True, + NUM_CACHE_BLOCKS=state_cache.shape[0], + ) + return out diff --git a/vllm_ascend/ops/triton/engram_int8.py b/vllm_ascend/ops/triton/engram_int8.py new file mode 100644 index 000000000000..bf501aeb762e --- /dev/null +++ b/vllm_ascend/ops/triton/engram_int8.py @@ -0,0 +1,61 @@ +import torch +from vllm.triton_utils import tl, triton + +from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton + + +@triton.jit +def _engram_int8_dequant_kernel(codes_ptr, scales_ptr, output_ptr, rows, WIDTH: tl.constexpr): + row = tl.program_id(0) + if row >= rows: + return + offsets = tl.arange(0, WIDTH) + codes = tl.load(codes_ptr + row * WIDTH + offsets).to(tl.float32) + scales = tl.load(scales_ptr + row * (WIDTH // 32) + offsets // 32) + tl.store(output_ptr + row * WIDTH + offsets, (codes * scales).to(tl.bfloat16)) + + +@triton.jit +def _engram_int8_gather_dequant_kernel(weight_ptr, scale_ptr, ids_ptr, output_ptr, rows, WIDTH: tl.constexpr): + row = tl.program_id(0) + if row >= rows: + return + offsets = tl.arange(0, WIDTH) + source_row = tl.load(ids_ptr + row) + codes = tl.load(weight_ptr + source_row * WIDTH + offsets).to(tl.float32) + scales = tl.load(scale_ptr + source_row * (WIDTH // 32) + offsets // 32) + tl.store(output_ptr + row * WIDTH + offsets, (codes * scales).to(tl.bfloat16)) + + +def dequantize_engram_int8(codes: torch.Tensor, scales: torch.Tensor) -> torch.Tensor: + if codes.device.type != "npu" or codes.dtype != torch.int8 or scales.dtype != torch.float32: + raise ValueError("Engram Triton INT8 dequant expects NPU int8 codes and FP32 scales") + if codes.ndim != 2 or codes.shape[1] != 256 or scales.shape != (codes.shape[0], 8): + raise ValueError("Engram Triton INT8 dequant expects [rows, 256] and [rows, 8]") + init_device_properties_triton() + output = torch.empty(codes.shape, dtype=torch.bfloat16, device=codes.device) + _engram_int8_dequant_kernel[(codes.shape[0],)](codes, scales, output, codes.shape[0], WIDTH=256, num_warps=4) + return output + + +def gather_dequantize_engram_int8(weight: torch.Tensor, scales: torch.Tensor, ids: torch.Tensor) -> torch.Tensor: + if ( + weight.device.type != "npu" + or weight.dtype != torch.int8 + or scales.dtype != torch.float32 + or ids.device.type != "npu" + or ids.dtype != torch.int64 + ): + raise ValueError("Engram fused INT8 gather expects NPU int8/FP32/int64 tensors") + if weight.ndim != 2 or weight.shape[1] != 256 or scales.shape != (weight.shape[0], 8): + raise ValueError("Engram fused INT8 gather expects table [rows, 256] and [rows, 8] scales") + if ids.ndim != 1: + raise ValueError("Engram fused INT8 gather expects flat IDs") + if ids.numel() == 0: + return torch.empty((0, 256), dtype=torch.bfloat16, device=weight.device) + init_device_properties_triton() + output = torch.empty((ids.shape[0], 256), dtype=torch.bfloat16, device=weight.device) + _engram_int8_gather_dequant_kernel[(ids.shape[0],)]( + weight, scales, ids, output, ids.shape[0], WIDTH=256, num_warps=4 + ) + return output diff --git a/vllm_ascend/ops/triton/prepare_indexer_indices.py b/vllm_ascend/ops/triton/prepare_indexer_indices.py new file mode 100644 index 000000000000..e479e8bc273e --- /dev/null +++ b/vllm_ascend/ops/triton/prepare_indexer_indices.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Convert score-ordered QLI indices into visible, chronological positions.""" + +import math + +import torch +from vllm.triton_utils import tl, triton + +from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num + + +@triton.jit(do_not_specialize=["num_rows", "blocks_per_core"]) +def _prepare_indexer_indices_kernel( + selected_ptr, + positions_ptr, + output_ptr, + num_rows, + TOPK: tl.constexpr, + COMPRESS_RATIO: tl.constexpr, + blocks_per_core, + BLOCK_ROWS: tl.constexpr, + BLOCK_COLS: tl.constexpr, + SENTINEL: tl.constexpr, + SORT_KEY_SHIFT: tl.constexpr, + NEGATIVE_KEY_BASE: tl.constexpr, +): + first_block = tl.program_id(0) * blocks_per_core + last_block = tl.minimum(first_block + blocks_per_core, tl.cdiv(num_rows, BLOCK_ROWS)) + columns = tl.arange(0, BLOCK_COLS) + for block in range(first_block, last_block): + rows = block * BLOCK_ROWS + tl.arange(0, BLOCK_ROWS) + offsets = rows[:, None] * TOPK + columns[None, :] + mask = (rows[:, None] < num_rows) & (columns[None, :] < TOPK) + selected = tl.load(selected_ptr + offsets, mask, other=SENTINEL) + positions = tl.load(positions_ptr + rows, rows < num_rows, other=-1) + visible = (positions + 1) // COMPRESS_RATIO + valid = (selected >= 0) & (selected < visible[:, None]) + selected = tl.where(valid, selected, SENTINEL) + # Use native FP32 sort even when the INT32 implementation is absent. + # Encode ordered, normal FP32 bit patterns without numeric conversion: + # [0, 2**24) uses negative keys; larger indices use positive keys. + key_bits = tl.where(selected < 2 * SORT_KEY_SHIFT, NEGATIVE_KEY_BASE - selected, selected - SORT_KEY_SHIFT) + sort_keys = key_bits.to(tl.float32, bitcast=True) + sort_keys = tl.extra.cann.extension.sort(sort_keys, dim=1, descending=False) + key_bits = sort_keys.to(tl.int32, bitcast=True) + selected = tl.where(key_bits < 0, NEGATIVE_KEY_BASE - key_bits, key_bits + SORT_KEY_SHIFT) + selected = tl.where(selected == SENTINEL, -1, selected) + tl.store(output_ptr + offsets, selected, mask) + + +def prepare_indexer_indices(selected: torch.Tensor, positions: torch.Tensor, compress_ratio: int) -> torch.Tensor: + """Filter and sort [tokens, topk] INT32 indices, with invalid slots last.""" + assert selected.ndim == 2 and selected.dtype == torch.int32 + assert positions.ndim == 1 and positions.shape[0] == selected.shape[0] + assert positions.dtype in (torch.int32, torch.int64) + assert compress_ratio in (1, 2) + num_rows, topk = selected.shape + assert 1 <= topk <= 2048 + selected = selected.contiguous() + positions = positions.contiguous() + output = torch.empty_like(selected) + if num_rows == 0: + return output + + num_cores = get_vectorcore_num() + block_cols = triton.next_power_of_2(topk) + # Budget 128 KiB for INT32 sort data and scratch (eight buffers). + max_block_rows = 128 * 1024 // (block_cols * 4 * 8) + # Core boundaries must also align to 32 bytes for arbitrary TopK widths. + aligned_rows = 8 // math.gcd(topk, 8) + block_rows = min(max_block_rows, triton.next_power_of_2(triton.cdiv(num_rows, num_cores))) + num_blocks = triton.cdiv(num_rows, block_rows) + aligned_blocks = triton.cdiv(aligned_rows, block_rows) + grid = min(triton.cdiv(num_blocks, aligned_blocks), num_cores) + blocks_per_core = triton.cdiv(triton.cdiv(num_blocks, grid), aligned_blocks) * aligned_blocks + _prepare_indexer_indices_kernel[(grid,)]( + selected, + positions, + output, + num_rows, + TOPK=topk, + COMPRESS_RATIO=compress_ratio, + blocks_per_core=blocks_per_core, + BLOCK_ROWS=block_rows, + BLOCK_COLS=block_cols, + SENTINEL=torch.iinfo(torch.int32).max, + SORT_KEY_SHIFT=1 << 23, + NEGATIVE_KEY_BASE=0x81800000 - (1 << 32), + multibuffer=False, + unit_flag=False, + ) + return output diff --git a/vllm_ascend/ops/triton/quantize_indexer_query.py b/vllm_ascend/ops/triton/quantize_indexer_query.py new file mode 100644 index 000000000000..b6fd6cd99f61 --- /dev/null +++ b/vllm_ascend/ops/triton/quantize_indexer_query.py @@ -0,0 +1,70 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Per-head INT8 query quantization with the FP16 scales consumed by QLI.""" + +import torch +from vllm.triton_utils import tl, triton + +from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num + + +@triton.jit(do_not_specialize=["num_rows", "blocks_per_core"]) +def _quantize_indexer_query_kernel( + query_ptr, + quantized_ptr, + scale_ptr, + num_rows, + blocks_per_core, + BLOCK_ROWS: tl.constexpr, + HEAD_DIM: tl.constexpr, + QUANT_MAX: tl.constexpr, + MIN_SCALE: tl.constexpr, +): + first_block = tl.program_id(0) * blocks_per_core + columns = tl.arange(0, HEAD_DIM) + for block in range(blocks_per_core): + rows = (first_block + block) * BLOCK_ROWS + tl.arange(0, BLOCK_ROWS) + offsets = rows[:, None] * HEAD_DIM + columns[None, :] + query = tl.load(query_ptr + offsets, rows[:, None] < num_rows, other=0).to(tl.float32) + abs_max = tl.max(tl.abs(query), axis=1) + # Quantize with the rounded FP16 scale, not the original FP32 value. + scale = tl.div_rn(abs_max, QUANT_MAX).to(tl.float16).to(tl.float32) + scale = tl.maximum(scale, MIN_SCALE) + normalized = tl.div_rn(query, scale[:, None]) + quantized = tl.extra.cann.libdevice.nearbyint(normalized) + quantized = tl.minimum(tl.maximum(quantized, -QUANT_MAX), QUANT_MAX).to(tl.int8) + tl.store(quantized_ptr + offsets, quantized, rows[:, None] < num_rows) + tl.store(scale_ptr + rows, scale, rows < num_rows) + + +def quantize_indexer_query(query: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Quantize [tokens, heads, 128] queries to INT8 and FP16 per-head scales. + + Match the indexer's round-to-even and symmetric [-127, 127] range. Zero + heads use the smallest positive FP16 subnormal so division stays defined. + """ + assert query.ndim == 3 and query.shape[-1] == 128 + query = query.contiguous() + quantized = torch.empty_like(query, dtype=torch.int8) + scale = torch.empty(query.shape[:-1], dtype=torch.float16, device=query.device) + num_rows = query.numel() // query.shape[-1] + if num_rows == 0: + return quantized, scale + + # Each core writes whole 32-byte groups of FP16 scales. Contiguous blocks + # keep neighboring cores from racing on a partial output cache line. + block_rows = 16 + num_blocks = triton.cdiv(num_rows, block_rows) + grid = min(num_blocks, get_vectorcore_num()) + _quantize_indexer_query_kernel[(grid,)]( + query, + quantized, + scale, + num_rows, + blocks_per_core=triton.cdiv(num_blocks, grid), + BLOCK_ROWS=block_rows, + HEAD_DIM=query.shape[-1], + QUANT_MAX=127.0, + MIN_SCALE=2.0**-24, + ) + return quantized, scale diff --git a/vllm_ascend/patch/__init__.py b/vllm_ascend/patch/__init__.py index f7a2895129d3..eb10afda605f 100644 --- a/vllm_ascend/patch/__init__.py +++ b/vllm_ascend/patch/__init__.py @@ -1279,3 +1279,45 @@ # Remove this patch once upstream `load_dspark_model` inherits the target # quant config for same-checkpoint drafts. # +# ** 34. File: platform/patch_vision.py** +# ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +# 1. `vllm.model_executor.models.vision.FusedInputNorm.forward` +# Why: +# Upstream vLLM uses PyTorch 2.13.0, which requires eps > 0 for training +# but allows eps >= 0 for inference. vllm-ascend bundles PyTorch 2.10.0, +# which does not distinguish scenarios and requires eps > 0 in all cases. +# So when upstream FusedInputNorm passes eps=0.0 to F.batch_norm it works +# fine upstream, but fails on vllm-ascend with "batch_norm eps must be +# positive". +# How: +# Monkey-patch FusedInputNorm.forward to use eps=1e-5 instead of 0.0. +# The patch is guarded with contextlib.suppress(ImportError) so it does +# not crash on release wheels (v0.26.0) where FusedInputNorm does not exist. +# Upstream PR #51734 (dc5101fb1b, Aug 10) rewrote FusedInputNorm.forward to +# use a broadcast multiply-add (x * weight + bias) instead of F.batch_norm, +# removing running_mean/running_var. That commit is included in the target +# 16cfe728, so the patch is gated to v0.27.1 only via vllm_version_is; +# on newer versions FusedInputNorm.forward is used as-is (multiply-add). +# Related PR (if no, explain why): +# https://github.com/vllm-project/vllm/pull/50411 +# https://github.com/vllm-project/vllm/pull/51734 +# Future Plan: +# Remove this patch once vllm-ascend's bundled PyTorch >= 2.13.0 +# (which, like upstream, allows eps >= 0 for inference). +# + +# ** File: platform/patch_deepseek_v41_frontend/ ** +# ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +# Target: vLLM tokenizer, renderer, reasoning and tool parser registries. +# Why: +# Day0 DeepSeek V4.1 encoding differs from the upstream V4 protocol. +# How: +# Register lazy deepseek_v41 implementations when global patches load; +# reuse vLLM rendering and parsing infrastructure with checkpoint encoding. +# Test: +# tests/ut/patch/platform/deepseek_v41 and CPU HTTP render/derender. +# Related PR: +# No upstream PR yet; this is temporary day0 protocol support. +# Future Plan: +# Upstream V4.1 frontend support to vLLM and remove this patch once the +# pinned vLLM version implements the same checkpoint protocol. diff --git a/vllm_ascend/patch/platform/__init__.py b/vllm_ascend/patch/platform/__init__.py index c6dc818a4700..a87ad351888f 100644 --- a/vllm_ascend/patch/platform/__init__.py +++ b/vllm_ascend/patch/platform/__init__.py @@ -16,7 +16,10 @@ import os +import vllm_ascend.patch.platform.patch_circular_buffer # noqa import vllm_ascend.patch.platform.patch_deepseek_v4_vision # noqa +import vllm_ascend.patch.platform.patch_deepseek_v41_config # noqa +import vllm_ascend.patch.platform.patch_deepseek_v41_frontend # noqa import vllm_ascend.patch.platform.patch_distributed # noqa import vllm_ascend.patch.platform.patch_kv_cache_utils # noqa import vllm_ascend.patch.platform.patch_mamba_block_aligned_split # noqa diff --git a/vllm_ascend/patch/platform/patch_circular_buffer.py b/vllm_ascend/patch/platform/patch_circular_buffer.py new file mode 100644 index 000000000000..fd451d50aa70 --- /dev/null +++ b/vllm_ascend/patch/platform/patch_circular_buffer.py @@ -0,0 +1,30 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Prefix lookup results never contain circular scratch pages.""" + +from vllm.v1.core.kv_cache_manager import KVCacheManager +from vllm.v1.core.kv_cache_utils import KVCacheBlock + +from vllm_ascend.core.circular_buffer import prefix_cacheable + +_original_truncate = KVCacheManager.truncate_computed_blocks + + +def _truncate_computed_blocks(self, blocks, num_computed_tokens): + groups = self.kv_cache_config.kv_cache_groups + if all(prefix_cacheable(g.kv_cache_spec) for g in groups): + return _original_truncate(self, blocks, num_computed_tokens) + truncated: list[list[KVCacheBlock]] = [] + for group_blocks, manager, group in zip(blocks.blocks, self.coordinator.single_type_managers, groups, strict=True): + if not prefix_cacheable(group.kv_cache_spec): + assert not group_blocks, "Scratch pages cannot be prefix-cache hits" + truncated.append([]) + continue + assert num_computed_tokens % manager.block_size == 0 + count = num_computed_tokens // manager.block_size + assert count <= len(group_blocks) + truncated.append(list(group_blocks[:count])) + return self.create_kv_cache_blocks(tuple(truncated)) + + +KVCacheManager.truncate_computed_blocks = _truncate_computed_blocks diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_config.py b/vllm_ascend/patch/platform/patch_deepseek_v41_config.py new file mode 100644 index 000000000000..ad14c82ea3d4 --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_config.py @@ -0,0 +1,32 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Register downstream DeepSeek V4.1 config classes before model parsing.""" + +from vllm.config.model import ModelConfig +from vllm.transformers_utils.config import _CONFIG_REGISTRY + +from vllm_ascend.deepseek_v41_config import ( + DeepseekV41Config, + DeepseekV41TextConfig, + DeepseekV41VisionConfig, +) + +_CONFIG_REGISTRY["deepseek_v4.1"] = DeepseekV41Config +_CONFIG_REGISTRY["deepseek_v4.1_text"] = DeepseekV41TextConfig +_CONFIG_REGISTRY["deepseek_v4.1_vision"] = DeepseekV41VisionConfig +_CONFIG_REGISTRY["deepseek_v41"] = DeepseekV41Config +_CONFIG_REGISTRY["deepseek_v41_text"] = DeepseekV41TextConfig +_CONFIG_REGISTRY["deepseek_v41_vision"] = DeepseekV41VisionConfig + +_original_is_deepseek_mla = ModelConfig.is_deepseek_mla.fget # type: ignore[attr-defined] + + +def _is_deepseek_mla(self: ModelConfig) -> bool: + if _original_is_deepseek_mla(self): + return True + return getattr(self.hf_text_config, "model_type", None) in ( + "deepseek_v4.1_text", + "deepseek_v41_text", + ) + + +ModelConfig.is_deepseek_mla = property(_is_deepseek_mla) # type: ignore[method-assign,assignment] diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/LICENSE b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/LICENSE new file mode 100644 index 000000000000..d84f527e101b --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2023 DeepSeek + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/README.md b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/README.md new file mode 100644 index 000000000000..557e8526772d --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/README.md @@ -0,0 +1,75 @@ +# DeepSeek V4.1 前端 monkey patch + +这是对 vLLM 前端的 day0 临时适配,放在 `patch/platform/`,随现有 +`adapt_patch(is_global_patch=True)` 加载,注册独立的 `deepseek_v41` +tokenizer、renderer、reasoning parser 和 tool parser。原有 `deepseek_v4` +注册和编码保持不变。待 vLLM 上游支持相同协议后移除该 patch。 + +## 格式来源 + +`encoding.py`、编码测试和 fixtures 来自正式 V4.1 权重仓库的 `encoding/` +及 `inference/examples/`,保留 DeepSeek MIT 许可。编码器仅调整 Python +格式、类型注解、异常捕获及使用仓库要求的 `regex`,格式语义由原始 golden +outputs 验证;运行时不读取权重目录中的 Python 文件。 + +- thinking 默认开启;`low/high/xhigh/max` 对应 `25/50/75/100`,默认 50。 +- `reasoning_effort="none"` 或显式关闭 thinking 切换至 chat 模式。 +- 数字预算 1–100 通过 `chat_template_kwargs.reasoning_effort` 传入。 + `minimal/medium` 不是参考格式的有效预算,返回参数错误。 +- 显式编码 `<|System|>`,保留中途 system、工具结果顺序和历史 thinking + 的参考语义;接受 `reasoning_content`,兼容 vLLM 的 `reasoning` 字段。 +- 工具使用带空格的 `calls`、`invoke`、`parameter` 标签;解析保留 + `string="true|false"` 的类型语义及 reasoning/content 空白。 +- 图片使用 `<|deepseek_image|>` 占位符。renderer 保留原始内容块的 + 双换行和图片顺序,使用 vLLM 的媒体加载、安全限制和 UUID 通道。 +- `response_format=json_schema` 的 schema 同时写入参考 system 提示。 + +## 使用 + +在 day0 的正常模型启动命令后添加: + +```bash +--tokenizer-mode deepseek_v41 \ +--reasoning-parser deepseek_v41 \ +--tool-call-parser deepseek_v41 \ +--enable-auto-tool-choice +``` + +数字预算请求示例: + +```json +{ + "model": "v41", + "messages": [{"role": "user", "content": "计算 17×23"}], + "chat_template_kwargs": {"thinking": true, "reasoning_effort": 42} +} +``` + +与 vLLM 一致,`tool_choice=none` 是否从提示中移除工具由 +`--exclude-tools-when-tool-choice-none` 控制。 + +`required`、指定函数及 strict auto tools 使用 V4.1 structural tag。 +支持声明顺序的必选/可选参数、字符串、数值、布尔、null、数组、JSON 对象、 +字符串 enum/pattern 以及类型 union。顶层跨参数条件和开放属性 schema +显式报错;不会退回旧版 DSML grammar。保留 vLLM 的 +`VLLM_ENFORCE_STRICT_TOOL_CALLING` 开关语义。 + +## 验证与边界 + +```bash +python -m pytest -q tests/ut/patch/platform/deepseek_v41 \ + tests/ut/patch/platform/test_deepseek_v4_thinking.py +``` + +测试涵盖参考编码、参数映射、原始请求不被修改、同步/异步图片加载顺序、 +完整/分块工具解析、类型语义及 grammar 接受/拒绝。另使用权重目录真实 +tokenizer 验证 prompt token IDs 和逐 token 解析,并通过 CPU-only +`vllm launch render` 检查 HTTP render/derender。 + +这是前端协议适配。day0 当前 V4.1 模型类仍是文本执行路径,视觉 processor、 +视觉权重、Engram、整模精度及性能验收属于模型适配;图片编码测试不代表 +图片端到端推理通过。 + +本地 HTTP 验证使用 vLLM `6e448d0ea9`。day0 的固定 vLLM revision +`ba07e4a48f` 的全局 patch 加载和完整模型检查会因基线引用已迁移的 +`vllm.model_executor.layers.attention.pcp` 而失败;需由 day0 基线另行对齐。 diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/__init__.py b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/__init__.py new file mode 100644 index 000000000000..b1461df7b07d --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/__init__.py @@ -0,0 +1,18 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Day0 patch for vLLM DeepSeek V4.1 frontend registries.""" + + +def register_frontend(): + from vllm.reasoning import ReasoningParserManager + from vllm.renderers.registry import RENDERER_REGISTRY + from vllm.tokenizers.registry import TokenizerRegistry + from vllm.tool_parsers import ToolParserManager + + package = "vllm_ascend.patch.platform.patch_deepseek_v41_frontend" + TokenizerRegistry.register("deepseek_v41", f"{package}.tokenizer", "DeepseekV41Tokenizer") + RENDERER_REGISTRY.register("deepseek_v41", f"{package}.renderer", "DeepseekV41Renderer") + ReasoningParserManager.register_lazy_module("deepseek_v41", f"{package}.parser", "DeepseekV41ReasoningParser") + ToolParserManager.register_lazy_module("deepseek_v41", f"{package}.parser", "DeepseekV41ToolParser") + + +register_frontend() diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/encoding.py b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/encoding.py new file mode 100644 index 000000000000..95fc1b7be956 --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/encoding.py @@ -0,0 +1,972 @@ +# SPDX-License-Identifier: MIT +# ruff: noqa: E501 +# Copyright (c) 2023 DeepSeek +# Vendored from the V4.1 checkpoint encoding/encoding.py; see LICENSE. +""" +DeepSeek-V4.1 Text and Vision Encoding + +A fully self-contained implementation for encoding/decoding DeepSeek-V4.1 chat +messages with tool calling, thinking mode, quick instruction tasks, and image +content blocks. No dependency on encoding_dsv4. + +V4.1 changes relative to V4: + +1. DSML tag names: tool calls are wrapped in "<|DSML| calls>" blocks with + "<|DSML| invoke>" / "<|DSML| parameter>" tags (leading-space tag names). +2. Numeric reasoning effort: "Reasoning Effort: {budget} (range 1-100, ...)". + Accepts an int in [1, 100] or one of "low"/"high"/"xhigh"/"max" + (mapped to 25/50/75/100). Defaults to "high". Only rendered in thinking mode. +3. Mid-conversation system messages are supported via the "<|System|>" token. + A mid-conversation system message behaves like a user message for the purpose + of appending the assistant generation header. +""" + +import copy +import json +from typing import Any + +import regex as re + +# ============================================================ +# Special Tokens +# ============================================================ + +bos_token: str = "<|begin▁of▁sentence|>" +eos_token: str = "<|end▁of▁sentence|>" +thinking_start_token: str = "" +thinking_end_token: str = "" +dsml_token: str = "|DSML|" + +USER_SP_TOKEN = "<|User|>" +ASSISTANT_SP_TOKEN = "<|Assistant|>" +LATEST_REMINDER_SP_TOKEN = "<|latest_reminder|>" + +IMAGE_PLACEHOLDER = "<|deepseek_image|>" +IMAGE_TAG_PATTERN = re.compile(r"(.*?)", re.DOTALL) + +# Task special tokens for internal classification tasks +DS_TASK_SP_TOKENS = { + "action": "<|action|>", + "query": "<|query|>", + "authority": "<|authority|>", + "domain": "<|domain|>", + "title": "<|title|>", + "read_url": "<|read_url|>", +} +VALID_TASKS = set(DS_TASK_SP_TOKENS.keys()) + +# ============================================================ +# Templates +# ============================================================ + +system_msg_template: str = "{content}" +user_msg_template: str = "{content}" +latest_reminder_msg_template: str = "{content}" +assistant_msg_template: str = "{reasoning}{content}{tool_calls}" + eos_token +assistant_msg_wo_eos_template: str = "{reasoning}{content}{tool_calls}" +thinking_template: str = "{reasoning_content}" + +response_format_template: str = ( + "## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n{schema}" +) + +tool_output_template: str = "{content}" + +# ============================================================ +# Utility Functions +# ============================================================ + + +def to_json(value: Any) -> str: + """Serialize a value to JSON string.""" + try: + return json.dumps(value, ensure_ascii=False) + except Exception: + return json.dumps(value, ensure_ascii=True) + + +def tools_from_openai_format(tools): + """Extract function definitions from OpenAI-format tool list.""" + return [tool["function"] for tool in tools] + + +def tool_calls_from_openai_format(tool_calls): + """Convert OpenAI-format tool calls to internal format.""" + return [ + { + "name": tool_call["function"]["name"], + "arguments": tool_call["function"]["arguments"], + } + for tool_call in tool_calls + ] + + +def tool_calls_to_openai_format(tool_calls): + """Convert internal tool calls to OpenAI format.""" + return [ + { + "type": "function", + "function": { + "name": tool_call["name"], + "arguments": tool_call["arguments"], + }, + } + for tool_call in tool_calls + ] + + +def decode_dsml_to_arguments(tool_name: str, tool_args: dict[str, tuple[str, str]]) -> dict[str, str]: + """ + Decode DSML parameters back to a tool call dict. + + Args: + tool_name: Name of the tool. + tool_args: Dict mapping param_name -> (value, is_string_flag). + + Returns: + Dict with "name" and "arguments" (JSON string) keys. + """ + + def _decode_value(key: str, value: str, string: str): + if string == "true": + value = to_json(value) + return f"{to_json(key)}: {value}" + + tool_args_json = "{" + ", ".join([_decode_value(k, v, string=is_str) for k, (v, is_str) in tool_args.items()]) + "}" + return dict(name=tool_name, arguments=tool_args_json) + + +# ============================================================ +# Preprocessing +# ============================================================ + + +def merge_tool_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: + """ + Merge tool messages into the preceding user message using content_blocks format. + + DeepSeek-V4.1 does not have a standalone "tool" role; instead, tool results + are encoded as blocks within user messages. + """ + merged: list[dict[str, Any]] = [] + + for msg in messages: + msg = copy.deepcopy(msg) + role = msg.get("role") + + if role == "tool": + # Convert tool message to a user message with tool_result block + tool_block = { + "type": "tool_result", + "tool_use_id": msg.get("tool_call_id", ""), + "content": msg.get("content", ""), + } + # Merge into previous message if it's already a user (merged tool) + if merged and merged[-1].get("role") == "user" and "content_blocks" in merged[-1]: + merged[-1]["content_blocks"].append(tool_block) + else: + merged.append( + { + "role": "user", + "content_blocks": [tool_block], + } + ) + elif role == "user": + content_blocks = msg.get("content_blocks") + if content_blocks is None: + content_blocks = [{"type": "text", "text": msg.get("content", "")}] + if ( + merged + and merged[-1].get("role") == "user" + and "content_blocks" in merged[-1] + and merged[-1].get("task") is None + ): + merged[-1]["content_blocks"].extend(content_blocks) + else: + # Preserve structured content and all message-level metadata. + new_msg = msg + new_msg["content_blocks"] = content_blocks + merged.append(new_msg) + else: + merged.append(msg) + + return merged + + +def sort_tool_results_by_call_order(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: + """ + Sort tool_result blocks within user messages by the order of tool_calls + in the preceding assistant message. + """ + last_tool_call_order: dict[str, int] = {} + + for msg in messages: + role = msg.get("role") + if role == "assistant" and msg.get("tool_calls"): + last_tool_call_order = {} + for idx, tc in enumerate(msg["tool_calls"]): + tc_id = tc.get("id") or tc.get("function", {}).get("id", "") + if tc_id: + last_tool_call_order[tc_id] = idx + + elif role == "user" and msg.get("content_blocks"): + tool_blocks = [b for b in msg["content_blocks"] if b.get("type") == "tool_result"] + if len(tool_blocks) > 1 and last_tool_call_order: + sorted_blocks = sorted(tool_blocks, key=lambda b: last_tool_call_order.get(b.get("tool_use_id", ""), 0)) + sorted_idx = 0 + new_blocks = [] + for block in msg["content_blocks"]: + if block.get("type") == "tool_result": + new_blocks.append(sorted_blocks[sorted_idx]) + sorted_idx += 1 + else: + new_blocks.append(block) + msg["content_blocks"] = new_blocks + + return messages + + +# ============================================================ +# Vision Message Preprocessing +# ============================================================ + + +def parse_tagged_text(text: str) -> str | list[dict[str, Any]]: + """Convert ``path`` text into standard content blocks.""" + matches = list(IMAGE_TAG_PATTERN.finditer(text)) + remaining = IMAGE_TAG_PATTERN.sub("", text) + if "" in remaining or "" in remaining: + raise ValueError("Malformed path tag") + if not matches: + return text + + blocks: list[dict[str, Any]] = [] + cursor = 0 + for match in matches: + if match.start() > cursor: + blocks.append({"type": "text", "text": text[cursor : match.start()]}) + path = match.group(1) + if not path: + raise ValueError("Image path must not be empty") + blocks.append( + { + "type": "image_url", + "image_url": {"url": path}, + } + ) + cursor = match.end() + if cursor < len(text): + blocks.append({"type": "text", "text": text[cursor:]}) + return blocks + + +def _is_image_block(block: dict[str, Any]) -> bool: + """Return whether a content block is an OpenAI/Anthropic/internal image.""" + return isinstance(block, dict) and block.get("type") in ("image", "image_url") + + +def _extract_image(block: dict[str, Any]) -> dict[str, Any]: + """Normalize a supported image block into an internal image record.""" + record: dict[str, Any] = {"type": "image"} + if block.get("type") == "image_url": + image_url = block.get("image_url") + if isinstance(image_url, str): + record["url"] = image_url + else: + record["url"] = (image_url or {}).get("url", "") + else: + for key in ("source", "url", "data"): + if key in block: + record[key] = block[key] + if not any(record.get(key) for key in ("source", "url", "data")): + raise ValueError("Image block does not contain a valid source") + return record + + +def _process_image_blocks( + blocks: list[Any], image_placeholder: str = IMAGE_PLACEHOLDER +) -> tuple[list[Any], list[dict[str, Any]]]: + """Replace image blocks and collect their records in one ordered traversal.""" + new_blocks: list[Any] = [] + images: list[dict[str, Any]] = [] + for block in blocks: + if not isinstance(block, dict): + new_blocks.append(block) + continue + if _is_image_block(block): + new_blocks.append({"type": "text", "text": image_placeholder}) + images.append(_extract_image(block)) + elif block.get("type") == "tool_result" and isinstance(block.get("content"), list): + block = copy.copy(block) + block["content"], nested_images = _process_image_blocks(block["content"], image_placeholder) + new_blocks.append(block) + images.extend(nested_images) + elif block.get("type") == "text": + text = block.get("text") or "" + if IMAGE_PLACEHOLDER in text: + raise ValueError( + f"Text block contains image placeholder '{IMAGE_PLACEHOLDER}': " + f"'{text[:100]}'. Images should be separate content blocks." + ) + new_blocks.append(block) + else: + new_blocks.append(block) + return new_blocks, images + + +def _validate_no_image_sp_tokens(msg: dict[str, Any]) -> None: + """Reject user-supplied image placeholder tokens in textual fields.""" + content = msg.get("content") + if isinstance(content, str) and IMAGE_PLACEHOLDER in content: + raise ValueError( + f"Message content contains image special token '{IMAGE_PLACEHOLDER}'. " + "Images should be provided as image content blocks." + ) + reasoning_content = msg.get("reasoning_content") + if isinstance(reasoning_content, str) and IMAGE_PLACEHOLDER in reasoning_content: + raise ValueError(f"reasoning_content contains image special token '{IMAGE_PLACEHOLDER}'") + + +def process_image_messages( + messages: list[dict[str, Any]], +) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + """Normalize image blocks and return their records in prompt order.""" + processed: list[dict[str, Any]] = [] + images: list[dict[str, Any]] = [] + for msg in messages: + msg = copy.deepcopy(msg) + _validate_no_image_sp_tokens(msg) + + if isinstance(msg.get("content"), list) and "content_blocks" not in msg: + msg["content_blocks"] = msg.pop("content") + + if msg.get("content_blocks"): + msg["content_blocks"], message_images = _process_image_blocks(msg["content_blocks"]) + images.extend(message_images) + if not isinstance(msg.get("content"), str): + texts = [ + block.get("text", "") + for block in msg["content_blocks"] + if isinstance(block, dict) and block.get("type") == "text" + ] + msg["content"] = "\n\n".join(texts) + + processed.append(msg) + return processed, images + + +def _read_until_stop(index: int, text: str, stop: list[str]) -> tuple[int, str, str | None]: + """ + Read text from index until one of the stop strings is found. + + Returns: + Tuple of (new_index, content_before_stop, matched_stop_string_or_None). + """ + min_pos = len(text) + matched_stop = None + + for s in stop: + pos = text.find(s, index) + if pos != -1 and pos < min_pos: + min_pos = pos + matched_stop = s + + if matched_stop: + content = text[index:min_pos] + return min_pos + len(matched_stop), content, matched_stop + else: + content = text[index:] + return len(text), content, None + + +# ============================================================ +# V4.1 Special Tokens and DSML Tag Names +# ============================================================ + +SYSTEM_SP_TOKEN = "<|System|>" + +tool_calls_block_name: str = " calls" +tool_call_tag_name: str = " invoke" +tool_parameter_tag_name: str = " parameter" + +tool_call_template: str = ( + '<{dsml_token}{tool_call_tag_name} name="{name}">\n{arguments}\n' +) +tool_calls_template = "<{dsml_token}{tc_block_name}>\n{tool_calls}\n" + +# ============================================================ +# Reasoning Effort (numeric budget) +# ============================================================ + +REASONING_EFFORT_TEMPLATE = ( + "Reasoning Effort: {budget} (range 1-100, the higher the value, the more thorough the reasoning)\n\n" +) + +REASONING_EFFORT_MAPPINGS: dict[str, int] = { + "low": 25, + "high": 50, + "xhigh": 75, + "max": 100, +} +DEFAULT_REASONING_EFFORT = "high" + + +def render_reasoning_effort( + index: int, + thinking_mode: str, + effort: str | int | None, +) -> str: + """Render the V4.1 numeric reasoning effort prefix (thinking mode, index 0 only).""" + if effort is None: + effort = DEFAULT_REASONING_EFFORT + assert (type(effort) is int and 1 <= effort <= 100) or effort in REASONING_EFFORT_MAPPINGS, ( + "Invalid reasoning effort for deepseek_v41: " + f"{effort}, should be int within [1,100] or {list(REASONING_EFFORT_MAPPINGS)}" + ) + if type(effort) is str: + effort = REASONING_EFFORT_MAPPINGS[effort] + if index == 0 and thinking_mode == "thinking": + return REASONING_EFFORT_TEMPLATE.format(budget=effort) + return "" + + +# ============================================================ +# Tools rendering +# ============================================================ + +TOOLS_TEMPLATE = """## Tools + +You have access to a set of tools to help answer the user's question. You can invoke tools by writing a "<{dsml_token}{tc_block_name}>" block like the following: + +<{dsml_token}{tc_block_name}> +<{dsml_token}{tool_call_tag_name} name="$TOOL_NAME"> +<{dsml_token}{tool_parameter_tag_name} name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE +... + +<{dsml_token}{tool_call_tag_name} name="$TOOL_NAME2"> +... + + + +String parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`. + +If thinking_mode is enabled (triggered by {thinking_start_token}), you MUST output your complete reasoning inside {thinking_start_token}...{thinking_end_token} BEFORE any tool calls or final response. + +Otherwise, output directly after {thinking_end_token} with tool calls or final response. + +### Available Tool Schemas + +{tool_schemas} + +You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls. +""" + + +def render_tools(tools: list[dict[str, str | dict[str, Any]]]) -> str: + """Render tool schemas into the V4.1 system prompt format.""" + tools_json = [to_json(t) for t in tools] + + return TOOLS_TEMPLATE.format( + tool_schemas="\n".join(tools_json), + dsml_token=dsml_token, + tc_block_name=tool_calls_block_name, + tool_call_tag_name=tool_call_tag_name, + tool_parameter_tag_name=tool_parameter_tag_name, + thinking_start_token=thinking_start_token, + thinking_end_token=thinking_end_token, + ) + + +def encode_arguments_to_dsml(tool_call: dict[str, Any]) -> str: + """Encode tool call arguments into V4.1 DSML parameter format.""" + p_dsml_template = ( + '<{dsml_token}{tool_parameter_tag_name} name="{key}" string="{is_str}">' + "{value}" + ) + P_dsml_strs = [] + + arguments = tool_call["arguments"] + if not isinstance(arguments, dict): + # Tolerate JSON strings, including double-encoded ones. + for _ in range(2): + if isinstance(arguments, str): + try: + arguments = json.loads(arguments) + except Exception: + break + else: + break + if not isinstance(arguments, dict): + arguments = {"arguments": tool_call["arguments"]} + + for k, v in arguments.items(): + P_dsml_strs.append( + p_dsml_template.format( + dsml_token=dsml_token, + tool_parameter_tag_name=tool_parameter_tag_name, + key=k, + is_str="true" if isinstance(v, str) else "false", + value=v if isinstance(v, str) else to_json(v), + ) + ) + + return "\n".join(P_dsml_strs) + + +# ============================================================ +# Message Rendering +# ============================================================ + + +def find_last_user_index(messages: list[dict[str, Any]]) -> int: + """ + Find the index of the last user/developer message. + + V4.1 supports mid-conversation system messages, which count as user + messages for the purposes of the assistant generation header. + """ + last_user_index = -1 + for idx in range(len(messages) - 1, -1, -1): + role = messages[idx].get("role") + if role in ["user", "developer"] or (role == "system" and idx > 0): + last_user_index = idx + break + return last_user_index + + +def render_message( + index: int, + messages: list[dict[str, Any]], + thinking_mode: str, + drop_thinking: bool = True, + reasoning_effort: str | int | None = None, +) -> str: + """ + Render a single message at the given index into its V4.1 encoded string form. + """ + assert 0 <= index < len(messages) + assert thinking_mode in ["chat", "thinking"], f"Invalid thinking_mode `{thinking_mode}`" + + msg = messages[index] + last_user_idx = find_last_user_index(messages) + + role = msg.get("role") + content = msg.get("content") + tools = msg.get("tools") + response_format = msg.get("response_format") + tool_calls = msg.get("tool_calls") + reasoning_content = msg.get("reasoning_content") + wo_eos = msg.get("wo_eos", False) + + if tools: + tools = tools_from_openai_format(tools) + if tool_calls: + tool_calls = tool_calls_from_openai_format(tool_calls) + + # Reasoning effort prefix (thinking mode, index 0 only) + reasoning_effort_prompt = render_reasoning_effort(index, thinking_mode, reasoning_effort) + # System token leads the conversation when there is a reasoning effort prompt + # or the first message is a system message. + prompt = SYSTEM_SP_TOKEN if index == 0 and (reasoning_effort_prompt or role == "system") else "" + prompt += reasoning_effort_prompt + + if role == "system": + if index > 0: + # Mid-conversation system message + prompt += SYSTEM_SP_TOKEN + prompt += system_msg_template.format(content=content or "") + if tools: + prompt += "\n\n" + render_tools(tools) + if response_format: + prompt += "\n\n" + response_format_template.format(schema=to_json(response_format)) + + elif role == "developer": + assert content, f"Invalid message for role `{role}`: {msg}" + + content_developer = USER_SP_TOKEN + content_developer += content + + if tools: + content_developer += "\n\n" + render_tools(tools) + if response_format: + content_developer += "\n\n" + response_format_template.format(schema=to_json(response_format)) + + prompt += user_msg_template.format(content=content_developer) + + elif role == "user": + prompt += USER_SP_TOKEN + + # Handle content blocks (tool results mixed with text) + content_blocks = msg.get("content_blocks") + if content_blocks: + parts = [] + for block in content_blocks: + block_type = block.get("type") + if block_type == "text": + parts.append(block.get("text", "")) + elif block_type == "tool_result": + tool_content = block.get("content", "") + if isinstance(tool_content, list): + text_parts = [] + for b in tool_content: + if b.get("type") == "text": + text_parts.append(b.get("text", "")) + else: + text_parts.append(f"[Unsupported {b.get('type')}]") + tool_content = "\n\n".join(text_parts) + parts.append(tool_output_template.format(content=tool_content)) + else: + parts.append(f"[Unsupported {block_type}]") + prompt += "\n\n".join(parts) + else: + prompt += content or "" + + elif role == "latest_reminder": + prompt += LATEST_REMINDER_SP_TOKEN + latest_reminder_msg_template.format(content=content) + + elif role == "tool": + raise NotImplementedError( + "deepseek_v41 merges tool messages into user; please preprocess with merge_tool_messages()" + ) + + elif role == "assistant": + thinking_part = "" + tc_content = "" + + if tool_calls: + tc_list = [ + tool_call_template.format( + dsml_token=dsml_token, + tool_call_tag_name=tool_call_tag_name, + name=tc.get("name"), + arguments=encode_arguments_to_dsml(tc), + ) + for tc in tool_calls + ] + tc_content += "\n\n" + tool_calls_template.format( + dsml_token=dsml_token, + tool_calls="\n".join(tc_list), + tc_block_name=tool_calls_block_name, + ) + + summary_content = content or "" + rc = reasoning_content or "" + + # Check if previous message has a task - if so, this is a task output (no thinking) + prev_has_task = index - 1 >= 0 and messages[index - 1].get("task") is not None + + if thinking_mode == "thinking" and not prev_has_task: + if not drop_thinking or index > last_user_idx: + thinking_part = thinking_template.format(reasoning_content=rc) + thinking_end_token + else: + thinking_part = "" + + if wo_eos: + prompt += assistant_msg_wo_eos_template.format( + reasoning=thinking_part, + content=summary_content, + tool_calls=tc_content, + ) + else: + prompt += assistant_msg_template.format( + reasoning=thinking_part, + content=summary_content, + tool_calls=tc_content, + ) + else: + raise NotImplementedError(f"Unknown role: {role}") + + # Append transition tokens based on what follows + if index + 1 < len(messages) and messages[index + 1].get("role") not in ["assistant", "latest_reminder"]: + return prompt + + task = messages[index].get("task") + if task is not None: + # Task special token for internal classification tasks + assert task in VALID_TASKS, f"Invalid task: '{task}'. Valid tasks are: {list(VALID_TASKS)}" + task_sp_token = DS_TASK_SP_TOKENS[task] + + if task != "action": + # Non-action tasks: append task sp token directly after the message + prompt += task_sp_token + else: + # Action task: append Assistant + thinking token + action sp token + prompt += ASSISTANT_SP_TOKEN + prompt += thinking_end_token if thinking_mode != "thinking" else thinking_start_token + prompt += task_sp_token + + elif messages[index].get("role") in ["user", "developer"] or ( + messages[index].get("role") == "system" and index > 0 + ): + # Normal generation: append Assistant + thinking token + # (mid-conversation system messages also trigger the assistant header) + prompt += ASSISTANT_SP_TOKEN + if ( + not drop_thinking + and thinking_mode == "thinking" + or drop_thinking + and thinking_mode == "thinking" + and index >= last_user_idx + ): + prompt += thinking_start_token + else: + prompt += thinking_end_token + + return prompt + + +# ============================================================ +# Main Encoding Function +# ============================================================ + + +def _drop_thinking_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: + """ + Drop reasoning_content and non-essential messages before the last user message. + Same as V4, but uses the V4.1 last-user definition (mid systems count). + """ + last_user_idx = find_last_user_index(messages) + result = [] + keep_roles = {"user", "system", "tool", "latest_reminder", "direct_search_results"} + + for idx, msg in enumerate(messages): + role = msg.get("role") + if role in keep_roles or idx >= last_user_idx: + result.append(msg) + elif role == "assistant": + msg = copy.copy(msg) + msg.pop("reasoning_content", None) + result.append(msg) + # developer and other roles before last_user_idx are dropped + + return result + + +def _encode_messages_text( + messages: list[dict[str, Any]], + thinking_mode: str, + context: list[dict[str, Any]] | None = None, + drop_thinking: bool = True, + add_default_bos_token: bool = True, + reasoning_effort: str | int | None = None, +) -> str: + """Encode preprocessed (text-only) messages into the V4.1 prompt format.""" + context = context if context else [] + + # Preprocess: merge tool messages and sort tool results + messages = merge_tool_messages(messages) + messages = sort_tool_results_by_call_order(context + messages)[len(context) :] + if context: + context = merge_tool_messages(context) + context = sort_tool_results_by_call_order(context) + + full_messages = context + messages + + prompt = bos_token if add_default_bos_token and len(context) == 0 else "" + + # Resolve drop_thinking: if any message has tools defined, don't drop thinking + effective_drop_thinking = drop_thinking + if any(m.get("tools") for m in full_messages): + effective_drop_thinking = False + + if thinking_mode == "thinking" and effective_drop_thinking: + full_messages = _drop_thinking_messages(full_messages) + num_to_render = len(full_messages) - len(_drop_thinking_messages(context)) + context_len = len(full_messages) - num_to_render + else: + num_to_render = len(messages) + context_len = len(context) + + for idx in range(num_to_render): + prompt += render_message( + idx + context_len, + full_messages, + thinking_mode=thinking_mode, + drop_thinking=effective_drop_thinking, + reasoning_effort=reasoning_effort, + ) + + return prompt + + +def encode_messages( + messages: list[dict[str, Any]], + thinking_mode: str, + context: list[dict[str, Any]] | None = None, + drop_thinking: bool = True, + add_default_bos_token: bool = True, + reasoning_effort: str | int | None = None, + return_multi_modal_data: bool = False, +) -> Any: + """Encode text or multimodal messages into the DeepSeek-V4.1 prompt format. + + Text-only calls return the prompt string. When return_multi_modal_data is + true, the result is ``(prompt, media_data)``. + """ + context = context or [] + processed_context, _ = process_image_messages(context) if context else ([], []) + processed_messages, images = process_image_messages(messages) + prompt = _encode_messages_text( + processed_messages, + thinking_mode=thinking_mode, + context=processed_context if processed_context else None, + drop_thinking=drop_thinking, + add_default_bos_token=add_default_bos_token, + reasoning_effort=reasoning_effort, + ) + if return_multi_modal_data: + return prompt, {"images": images} + return prompt + + +def load_cases(input_file: str) -> list[dict[str, Any]]: + """Load one or more OpenAI-format conversation cases from JSON.""" + with open(input_file) as file: + data = json.load(file) + if isinstance(data, dict): + data = [data] + elif data and isinstance(data[0], dict) and "role" in data[0]: + data = [{"messages": data}] + + cases = [] + for case in data: + messages = copy.deepcopy(case["messages"]) + if "tools" in case: + if not messages: + raise ValueError("A case with tools must contain at least one message") + messages[0]["tools"] = case["tools"] + cases.append( + { + "messages": messages, + "context": case.get("context"), + "thinking_mode": case.get("thinking_mode"), + "reasoning_effort": case.get("reasoning_effort"), + } + ) + return cases + + +def encode_case(case: dict[str, Any], thinking_mode: str) -> tuple[str, list[dict[str, Any]]]: + """Encode one JSON case and return its current-turn image records.""" + prompt, media_data = encode_messages( + case["messages"], + thinking_mode=case.get("thinking_mode") or thinking_mode, + context=case.get("context"), + reasoning_effort=case.get("reasoning_effort"), + return_multi_modal_data=True, + ) + return prompt, media_data["images"] + + +# ============================================================ +# Parsing (Decoding model output) +# ============================================================ + + +def parse_tool_calls(index: int, text: str) -> tuple[int, str | None, list[dict[str, str]]]: + """ + Parse V4.1 DSML tool calls from text starting at the given index. + + Returns: + Tuple of (new_index, last_stop_token, list_of_tool_call_dicts). + """ + tool_calls: list[dict[str, Any]] = [] + stop_token = None + tool_calls_end_token = f"" + tool_call_start_token = f"<{dsml_token}{tool_call_tag_name}" + tool_call_end_token = f"\n": + raise ValueError(f"Tool call format error: expected '>\\n' but got '{_}'") + + if stop_token == tool_calls_end_token: + break + + if stop_token is None: + raise ValueError("Missing special token in tool calls") + + index, tool_name_content, stop_token = _read_until_stop( + index, text, [tool_parameter_start_token, tool_call_end_token] + ) + + p_tool_name = re.findall(r'^\s*name="(.*?)">\n$', tool_name_content, flags=re.DOTALL) + if len(p_tool_name) != 1: + raise ValueError(f"Tool name format error: '{tool_name_content}'") + tool_name = p_tool_name[0] + + tool_args: dict[str, tuple[str, str]] = {} + while stop_token == tool_parameter_start_token: + index, param_content, stop_token = _read_until_stop(index, text, [tool_parameter_end_token]) + + param_kv = re.findall(r'^ name="(.*?)" string="(true|false)">(.*?)<$', param_content, flags=re.DOTALL) + if len(param_kv) != 1: + raise ValueError(f"Parameter format error: '{param_content}'") + param_name, string, param_value = param_kv[0] + + if param_name in tool_args: + raise ValueError(f"Duplicate parameter name: '{param_name}'") + tool_args[param_name] = (param_value, string) + + index, content, stop_token = _read_until_stop( + index, text, [tool_parameter_start_token, tool_call_end_token] + ) + if content != ">\n": + raise ValueError(f"Parameter format error: expected '>\\n' but got '{content}'") + + tool_call = decode_dsml_to_arguments(tool_name=tool_name, tool_args=tool_args) + tool_calls.append(tool_call) + + return index, stop_token, tool_calls + + +def parse_message_from_completion_text(text: str, thinking_mode: str) -> dict[str, Any]: + """ + Parse a model completion text into a structured assistant message (V4.1 format). + + Returns: + Dict with keys: "role", "content", "reasoning_content", "tool_calls". + tool_calls are in OpenAI format. + """ + summary_content, reasoning_content = "", "" + tool_calls: list[dict[str, Any]] = [] + index, stop_token = 0, None + tool_calls_start_token = f"\n\n<{dsml_token}{tool_calls_block_name}" + + is_thinking = thinking_mode == "thinking" + is_tool_calling = False + + if is_thinking: + index, content_delta, stop_token = _read_until_stop(index, text, [thinking_end_token, tool_calls_start_token]) + reasoning_content = content_delta + assert stop_token == thinking_end_token, "Invalid thinking format: missing " + + index, content_delta, stop_token = _read_until_stop(index, text, [eos_token, tool_calls_start_token]) + summary_content = content_delta + if stop_token == tool_calls_start_token: + is_tool_calling = True + else: + assert stop_token == eos_token, "Invalid format: missing EOS token" + + if is_tool_calling: + index, stop_token, tool_calls = parse_tool_calls(index, text) + + index, tool_ends_text, stop_token = _read_until_stop(index, text, [eos_token]) + assert not tool_ends_text, "Unexpected content after tool calls" + + assert len(text) == index and stop_token in [eos_token, None], "Unexpected content at end" + + for sp_token in [bos_token, eos_token, thinking_start_token, thinking_end_token, dsml_token]: + assert sp_token not in summary_content and sp_token not in reasoning_content, ( + f"Unexpected special token '{sp_token}' in content" + ) + + return { + "role": "assistant", + "content": summary_content, + "reasoning_content": reasoning_content, + "tool_calls": tool_calls_to_openai_format(tool_calls), + } diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/parser.py b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/parser.py new file mode 100644 index 000000000000..36e196664986 --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/parser.py @@ -0,0 +1,97 @@ +# SPDX-License-Identifier: Apache-2.0 +"""V4.1 DSML terminals on the vLLM streaming parser engine.""" + +import contextlib +import json +from dataclasses import replace +from functools import cache + +import regex as re +from vllm.parser.deepseek_v4 import deepseek_v4_config +from vllm.parser.engine.adapters import ParserEngineReasoningAdapter, ParserEngineToolAdapter +from vllm.parser.engine.parser_engine import ParserEngine + +from .tokenizer import thinking_enabled + +PARAMETER_PATTERN = re.compile( + r'<|DSML| parameter name="([^"]+)" string="(true|false)">(.*?)', re.DOTALL +) +PARTIAL_PARAMETER_PATTERN = re.compile(r'<|DSML| parameter name="([^"]+)" string="(true|false)">(.*)$', re.DOTALL) + + +def convert_arguments(raw_args, partial): + params = {} + end = 0 + for match in PARAMETER_PATTERN.finditer(raw_args): + name, string, value = match.groups() + if name in params: + raise ValueError(f"Duplicate V4.1 tool parameter: {name}") + params[name] = value if string == "true" else json.loads(value) + end = match.end() + if partial and (partial_match := PARTIAL_PARAMETER_PATTERN.search(raw_args, end)): + name, string, value = partial_match.groups() + if string == "true": + params[name] = value + else: + with contextlib.suppress(json.JSONDecodeError): + params[name] = json.loads(value) + return json.dumps(params, ensure_ascii=False) + + +def v41_terminal(text): + return ( + text.replace("|DSML|tool_calls", "|DSML| calls") + .replace("|DSML|invoke", "|DSML| invoke") + .replace("|DSML|parameter", "|DSML| parameter") + ) + + +@cache +def deepseek_v41_config(thinking): + config = deepseek_v4_config(thinking=thinking) + terminals = {name: v41_terminal(text) for name, text in config.terminals.items()} + # The reference decoder treats these two newlines as the tool delimiter, + # not as part of the assistant's summary content. + terminals["TOOL_START"] = "\n\n<|DSML| calls>" + return replace( + config, + name="deepseek_v41", + terminals=terminals, + token_id_terminals={ + name: v41_terminal(text) for name, text in config.token_id_terminals.items() if name != "TOOL_START" + }, + arg_converter=convert_arguments, + strip_trailing_reasoning_whitespace=False, + drop_whitespace_only_content_before_tools=False, + strip_content_whitespace_with_tools=False, + # The DSML marker may be a special token inside a multi-token tag. + preserve_tokens=config.preserve_tokens | {"|DSML|"}, + ) + + +class DeepseekV41Parser(ParserEngine): + def __init__(self, tokenizer, tools=None, **kwargs): + chat_kwargs = kwargs.pop("chat_template_kwargs", None) or {} + super().__init__( + tokenizer, tools, parser_engine_config=deepseek_v41_config(thinking_enabled(chat_kwargs)), **kwargs + ) + + def _fix_arg_types(self, args_json, func_name): + # The wire format's string flag is authoritative, including values + # such as "true" and "42" that happen to resemble another JSON type. + return args_json + + +class DeepseekV41ReasoningParser(ParserEngineReasoningAdapter): + _parser_engine_cls = DeepseekV41Parser + + +class DeepseekV41ToolParser(ParserEngineToolAdapter): + _parser_engine_cls = DeepseekV41Parser + structural_tag_model = "deepseek_v41" + + def get_structural_tag(self, request, *, reasoning=False): + # Imported lazily: grammar compilation is unnecessary for auto tools. + from .structural_tag import get_structural_tag + + return get_structural_tag(request, reasoning=reasoning) diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/renderer.py b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/renderer.py new file mode 100644 index 000000000000..06b6d7bd6ef9 --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/renderer.py @@ -0,0 +1,77 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Reuse vLLM media loading while encoding the original V4.1 messages.""" + +import copy + +from vllm.entrypoints.chat_utils import parse_chat_messages, parse_chat_messages_async +from vllm.renderers.deepseek_v4 import DeepseekV4Renderer +from vllm.renderers.inputs.preprocess import parse_dec_only_prompt + + +def _media_blocks(blocks): + """Expose images nested in reference tool_result/content_blocks to vLLM.""" + result = [] + for block in blocks: + if block.get("type") == "tool_result": + content = block.get("content", "") + result.extend(_media_blocks(content) if isinstance(content, list) else [{"type": "text", "text": content}]) + elif block.get("type") == "image": + source = block.get("source") or {} + url = block.get("url") or source.get("url") + if source.get("type") == "base64": + url = f"data:{source['media_type']};base64,{source['data']}" + url = url or block.get("data") + result.append({"type": "image_url", "image_url": {"url": url}}) + elif block.get("type") == "image_url" and isinstance(block.get("image_url"), str): + result.append({**block, "image_url": {"url": block["image_url"]}}) + else: + result.append(block) + return result + + +def _media_messages(messages): + messages = copy.deepcopy(messages) + for message in messages: + content = message.pop("content_blocks", message.get("content")) + if isinstance(content, list): + message["content"] = _media_blocks(content) + if "reasoning_content" in message and "reasoning" not in message: + message["reasoning"] = message["reasoning_content"] + return messages + + +class DeepseekV41Renderer(DeepseekV4Renderer): + def _template_kwargs(self, params): + return {**params.get_apply_chat_template_kwargs(), "response_format": params.response_format} + + @staticmethod + def _prompt(raw, media, uuids): + prompt = parse_dec_only_prompt(raw) + if media is not None: + prompt["multi_modal_data"] = media + if uuids is not None: + prompt["multi_modal_uuids"] = uuids + return prompt + + def render_messages(self, messages, params): + conversation, media, uuids = parse_chat_messages( + _media_messages(messages), + self.model_config, + content_format="string", + media_io_kwargs=params.media_io_kwargs, + mm_processor_kwargs=params.mm_processor_kwargs, + ) + raw = self._apply_chat_template(messages=messages, **self._template_kwargs(params)) + return conversation, self._prompt(raw, media, uuids) + + async def render_messages_async(self, messages, params): + conversation, media, uuids = await parse_chat_messages_async( + _media_messages(messages), + self.model_config, + content_format="string", + media_io_kwargs=params.media_io_kwargs, + mm_processor_kwargs=params.mm_processor_kwargs, + ) + raw = await self._apply_chat_template_async(messages=messages, **self._template_kwargs(params)) + return conversation, self._prompt(raw, media, uuids) diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/structural_tag.py b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/structural_tag.py new file mode 100644 index 000000000000..1a5edb1c4315 --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/structural_tag.py @@ -0,0 +1,120 @@ +# SPDX-License-Identifier: Apache-2.0 +"""V4.1 tool constraints: raw string parameters and JSON for other values.""" + +import copy +import json + +from vllm import envs +from vllm.tool_parsers.structural_tag_registry import ( + _any_tool_strict, + _dump_tool_choice_for_xgrammar, + _dump_tool_for_xgrammar, +) +from xgrammar import Grammar, StructuralTag, normalize_tool_choice +from xgrammar.structural_tag import ( + AnyTextFormat, + ConstStringFormat, + GrammarFormat, + JSONSchemaFormat, + OptionalFormat, + OrFormat, + SequenceFormat, + TagFormat, + TagsWithSeparatorFormat, + TriggeredTagsFormat, +) + +PARAMETER_END = "" + + +def _parameter_formats(name, schema, definitions): + if "$ref" in schema: + ref = schema["$ref"] + if not ref.startswith("#/$defs/") or ref[8:] not in definitions: + raise ValueError(f"Unsupported V4.1 tool parameter reference: {ref}") + schema = {**definitions[ref[8:]], **{key: value for key, value in schema.items() if key != "$ref"}} + if "anyOf" in schema: + return [item for branch in schema["anyOf"] for item in _parameter_formats(name, branch, definitions)] + types = schema.get("type", ["string", "number", "boolean", "null", "array", "object"]) + if isinstance(types, str): + types = [types] + formats = [] + for value_type in types: + value_schema = {**schema, "type": value_type, "$defs": definitions} + if value_type == "string": + # Use xgrammar's raw XML string constraints (enum/pattern/length), + # changing only its terminator, never user-provided schema literals. + grammar = str( + Grammar.from_structural_tag( + StructuralTag(format=JSONSchemaFormat(json_schema=value_schema, style="deepseek_xml")) + ) + ) + old_end = json.dumps("") + grammar = grammar.replace(f"excludes=({old_end})", f"excludes=({json.dumps(PARAMETER_END)})") + content = GrammarFormat(grammar=grammar) + else: + content = JSONSchemaFormat(json_schema=value_schema) + formats.append( + TagFormat( + begin=f'<|DSML| parameter name="{name}" string="{str(value_type == "string").lower()}">', + content=content, + end=PARAMETER_END + "\n", + ) + ) + return formats + + +def _tool_tag(tool): + function = tool.function + schema = copy.deepcopy(function.parameters or {"type": "object", "properties": {}}) + # Cross-property constraints cannot be enforced by independent parameter + # grammars. Reject them rather than silently weakening strict tool schemas. + supported = {"type", "properties", "required", "additionalProperties", "$defs", "title", "description", "$schema"} + if set(schema) - supported or schema.get("type", "object") != "object": + raise ValueError("V4.1 strict tools require an object schema with properties and required fields") + if schema.get("additionalProperties") not in (None, False): + raise ValueError("V4.1 strict tools require named properties (additionalProperties=false)") + parameters = [] + properties = schema.get("properties", {}) + required = schema.get("required", []) + if not set(required).issubset(properties): + raise ValueError("V4.1 required parameters must be declared in properties") + for name, value_schema in properties.items(): + if '"' in name: + raise ValueError("V4.1 DSML parameter names cannot contain quotes") + alternatives = _parameter_formats(name, value_schema, schema.get("$defs", {})) + parameter = alternatives[0] if len(alternatives) == 1 else OrFormat(elements=alternatives) + if name not in required: + parameter = OptionalFormat(content=parameter) + parameters.append(parameter) + return TagFormat( + begin=f'<|DSML| invoke name="{function.name}">\n', + content=SequenceFormat(elements=parameters) if parameters else ConstStringFormat(value="\n"), + end="\n", + ) + + +def get_structural_tag(request, *, reasoning): + if not envs.VLLM_ENFORCE_STRICT_TOOL_CALLING or not request.tools or request.tool_choice == "none": + return None + if request.tool_choice == "auto" and not _any_tool_strict(request.tools): + return None + tools, builtin_tools, choice = normalize_tool_choice( + [_dump_tool_for_xgrammar(tool) for tool in request.tools], + _dump_tool_choice_for_xgrammar(request.tool_choice), + ) + if builtin_tools: + raise ValueError("V4.1 supports function tools only") + tags = [_tool_tag(tool) for tool in tools] + calls = TagFormat( + begin="<|DSML| calls>\n", + content=tags[0] if choice == "forced" else TagsWithSeparatorFormat(tags=tags, separator="", at_least_one=True), + end="", + ) + if choice == "auto": + suffix = TriggeredTagsFormat(triggers=["<|DSML| calls>"], tags=[calls], excludes=["", ""]) + else: + suffix = SequenceFormat(elements=[ConstStringFormat(value="\n\n"), calls]) + if reasoning: + suffix = SequenceFormat(elements=[TagFormat(begin="", content=AnyTextFormat(), end=""), suffix]) + return StructuralTag(format=suffix) diff --git a/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/tokenizer.py b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/tokenizer.py new file mode 100644 index 000000000000..e9414173983d --- /dev/null +++ b/vllm_ascend/patch/platform/patch_deepseek_v41_frontend/tokenizer.py @@ -0,0 +1,88 @@ +# SPDX-License-Identifier: Apache-2.0 +"""OpenAI request adapter for the checkpoint's standalone encoder.""" + +import copy + +from transformers import TokenizersBackend +from vllm.tokenizers.hf import get_cached_tokenizer +from vllm.tokenizers.protocol import TokenizerLike + +from .encoding import encode_messages, render_reasoning_effort + + +def thinking_enabled(kwargs): + if kwargs.get("reasoning_effort") == "none": + return False + if "thinking" not in kwargs and "enable_thinking" not in kwargs: + return True + return bool(kwargs.get("thinking") or kwargs.get("enable_thinking")) + + +def get_deepseek_v41_tokenizer(tokenizer): + wrapped = copy.copy(tokenizer) + + class _DeepseekV41Tokenizer(tokenizer.__class__): # type: ignore[name-defined] + def apply_chat_template(self, messages, tools=None, **kwargs): + # Keep original content blocks: vLLM's flattened conversation loses + # reference separators, image positions and reasoning_content. + messages = copy.deepcopy(messages) + for message in messages: + if "reasoning_content" not in message and "reasoning" in message: + message["reasoning_content"] = message["reasoning"] + + response_format = kwargs.get("response_format") + if response_format is not None and hasattr(response_format, "model_dump"): + response_format = response_format.model_dump(by_alias=True) + schema = None + if response_format and response_format.get("type") == "json_schema": + schema = response_format["json_schema"]["schema"] + if tools or schema is not None: + if not messages or messages[0]["role"] != "system": + messages.insert(0, {"role": "system", "content": ""}) + if tools: + # OpenAI request validation inserts absent optional fields + # as None; those are not part of the reference tool schema. + messages[0]["tools"] = [ + { + **tool, + "function": {key: value for key, value in tool["function"].items() if value is not None}, + } + for tool in tools + ] + if schema is not None: + messages[0]["response_format"] = schema + + effort = kwargs.get("reasoning_effort") + effort = None if effort == "none" else effort + try: + render_reasoning_effort(0, "chat", effort) + except (AssertionError, TypeError) as error: + raise ValueError(str(error)) from error + prompt = encode_messages( + messages, + thinking_mode="thinking" if thinking_enabled(kwargs) else "chat", + reasoning_effort=effort, + drop_thinking=kwargs.get("drop_thinking", True), + context=kwargs.get("context"), + add_default_bos_token=kwargs.get("add_default_bos_token", True), + ) + if not kwargs.get("tokenize", True): + return prompt + return self.encode( + prompt, + add_special_tokens=False, + **{key: kwargs[key] for key in ("truncation", "max_length") if key in kwargs}, + ) + + def __reduce__(self): + return get_deepseek_v41_tokenizer, (tokenizer,) + + wrapped.__class__ = _DeepseekV41Tokenizer + return wrapped + + +class DeepseekV41Tokenizer(TokenizerLike): + @classmethod + def from_pretrained(cls, *args, **kwargs): + tokenizer = TokenizersBackend.from_pretrained(*args, **kwargs) + return get_cached_tokenizer(get_deepseek_v41_tokenizer(tokenizer)) diff --git a/vllm_ascend/patch/platform/patch_deepseek_v4_vision.py b/vllm_ascend/patch/platform/patch_deepseek_v4_vision.py index a6d6bfae8512..cfee0a2f519c 100644 --- a/vllm_ascend/patch/platform/patch_deepseek_v4_vision.py +++ b/vllm_ascend/patch/platform/patch_deepseek_v4_vision.py @@ -29,6 +29,9 @@ def register_deepseek_v4_vision_config_convertor() -> None: class AscendDeepseekV4ModelArchConfigConvertor(ModelArchConfigConvertorBase): """Route vision checkpoints to the Ascend multimodal wrapper.""" + architecture = "DeepseekV4ForConditionalGeneration" + mm_prefix_span_leading_pad_modulus = 4 + def __init__( self, hf_config: "PretrainedConfig", @@ -36,7 +39,7 @@ def __init__( revision: str | None = None, ) -> None: if getattr(hf_config, "vision_n_layers", 0) > 0: - hf_config.architectures = ["DeepseekV4ForConditionalGeneration"] + hf_config.architectures = [self.architecture] hf_config.mm_prefix_clamp_sliding_window = True hf_config.mm_prefix_span_leading_pad_modulus = 4 super().__init__(hf_config, hf_text_config, revision) @@ -44,5 +47,14 @@ def __init__( def is_mm_prefix_lm(self, supports_multimodal: bool = True) -> bool: return supports_multimodal and (getattr(self.hf_config, "vision_n_layers", 0) > 0) + class AscendDeepseekV41ModelArchConfigConvertor( + AscendDeepseekV4ModelArchConfigConvertor, + ): + """Route V4.1 vision checkpoints to their multimodal wrapper.""" + + architecture = "DeepseekV41ForConditionalGeneration" + mm_prefix_span_leading_pad_modulus = 2 + MODEL_ARCH_CONFIG_CONVERTORS["deepseek_v4"] = AscendDeepseekV4ModelArchConfigConvertor + MODEL_ARCH_CONFIG_CONVERTORS["deepseek_v4.1"] = AscendDeepseekV41ModelArchConfigConvertor _REGISTERED = True diff --git a/vllm_ascend/patch/platform/patch_fused_moe.py b/vllm_ascend/patch/platform/patch_fused_moe.py index f34ed781a542..a881e810b32b 100644 --- a/vllm_ascend/patch/platform/patch_fused_moe.py +++ b/vllm_ascend/patch/platform/patch_fused_moe.py @@ -42,10 +42,17 @@ from vllm_ascend.ascend_config import get_ascend_config from vllm_ascend.distributed.eplb.state import AscendEplbLayerState -from vllm_ascend.ops.fused_moe.router.router_factory import create_ascend_fused_moe_router _EPLB_ROUTER_ADAPTED = "_vllm_ascend_eplb_router_adapted" + +def create_ascend_fused_moe_router(*args, **kwargs): + """Load Ascend router ops only when a model constructs its MoE.""" + from vllm_ascend.ops.fused_moe.router.router_factory import create_ascend_fused_moe_router as factory + + return factory(*args, **kwargs) + + # Capture the real original before fused_moe.py's module-level code runs. _original_FusedMoE = _fused_moe_layer.FusedMoEFactory diff --git a/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py b/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py index 14cf16de94a5..28eb27ada66c 100644 --- a/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py +++ b/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py @@ -32,6 +32,7 @@ MambaSpec, ) +from vllm_ascend.core.circular_buffer import prefix_cacheable from vllm_ascend.utils import vllm_version_is USE_MULTI_GROUPS_KV_CACHE = True @@ -172,7 +173,7 @@ def __init__( # type: ignore[misc] assert all( self._get_effective_block_size(g.kv_cache_spec) % hash_block_size == 0 for g in kv_cache_config.kv_cache_groups - if getattr(g.kv_cache_spec, "participates_in_prefix_caching", True) + if prefix_cacheable(g.kv_cache_spec) ), "block_size must be divisible by hash_block_size" self.enable_partial_hash_hits = dcp_world_size == 1 and any( isinstance(g.kv_cache_spec, MambaSpec) @@ -197,7 +198,9 @@ def __init__( # type: ignore[misc] def _cache_hit_alignment_tokens(self) -> int: if self.enable_partial_hash_hits: return self.hash_block_size - return self.scheduler_block_size or self.lcm_block_size + alignment = self.scheduler_block_size or self.lcm_block_size + assert alignment is not None + return alignment def _get_effective_block_size(self, kv_cache_spec: KVCacheSpec) -> int: block_size = kv_cache_spec.block_size @@ -214,6 +217,8 @@ def verify_and_split_kv_cache_groups(self) -> None: """ self.attention_groups: list[SpecGroup] = [] for i, g in enumerate(self.kv_cache_config.kv_cache_groups): + if not prefix_cacheable(g.kv_cache_spec): + continue manager_cls = self.single_type_managers[i].__class__ spec = g.kv_cache_spec use_eagle = i in self.eagle_group_ids @@ -229,7 +234,11 @@ def verify_and_split_kv_cache_groups(self) -> None: else: self.attention_groups.append(SpecGroup(spec, [i], manager_cls, use_eagle)) - assert len(self.attention_groups) > 1, "HybridKVCacheCoordinator requires at least two attention groups." + self.full_attention_group_id: int | None + if not self.attention_groups: + self.full_attention_group_id = None + self.lcm_block_size = self.scheduler_block_size + return # Put full attention first: its efficient left-to-right scan provides # a tighter initial bound, reducing work for subsequent groups. @@ -240,9 +249,7 @@ def verify_and_split_kv_cache_groups(self) -> None: # so any group reporting a longer per-group hit implies the union of # per-group hits is not consistent at a single boundary (#46453). first = self.attention_groups[0] - self.full_attention_group_id: int | None = ( - first.group_ids[0] if isinstance(first.spec, FullAttentionSpec) else None - ) + self.full_attention_group_id = first.group_ids[0] if isinstance(first.spec, FullAttentionSpec) else None # Propagate the eagle bit to every manager in an eagle-containing # attention group, mirroring upstream @@ -303,6 +310,8 @@ def _get_block_hashes(kv_cache_spec: KVCacheSpec) -> BlockHashList: return block_hashes num_groups = len(self.kv_cache_config.kv_cache_groups) + if not self.attention_groups: + return tuple([] for _ in range(num_groups)), 0 hit_length = max_cache_hit_length longest_hit_length = 0 hit_blocks_by_group: list[list[KVCacheBlock] | None] = [None] * num_groups diff --git a/vllm_ascend/patch/platform/patch_kv_cache_utils.py b/vllm_ascend/patch/platform/patch_kv_cache_utils.py index 94a13eef6a7a..29e2664e0a09 100644 --- a/vllm_ascend/patch/platform/patch_kv_cache_utils.py +++ b/vllm_ascend/patch/platform/patch_kv_cache_utils.py @@ -2,6 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM-Ascend project import math from collections import defaultdict +from dataclasses import replace import vllm.v1.core.kv_cache_utils from vllm.config import VllmConfig @@ -22,6 +23,26 @@ get_kv_cache_spec_kind, ) +from vllm_ascend.core.circular_buffer import prefix_cacheable +from vllm_ascend.core.deepseek_v41 import ( + allocate_cache_config as allocate_v41_cache_config, +) +from vllm_ascend.core.deepseek_v41 import ( + group_cache_specs as group_v41_cache_specs, +) +from vllm_ascend.core.deepseek_v41 import ( + has_v41_groups, + is_v41_spec, +) +from vllm_ascend.core.deepseek_v41 import ( + make_cache_groups as make_v41_cache_groups, +) +from vllm_ascend.core.deepseek_v41 import ( + pool_bytes_per_block as v41_pool_bytes_per_block, +) +from vllm_ascend.core.deepseek_v41 import ( + request_blocks as v41_request_blocks, +) from vllm_ascend.models.glm5next.cache_config import ( _get_glm5_next_cache_layout, get_glm5_next_kv_cache_config, @@ -44,6 +65,15 @@ _orig_get_kv_cache_config_from_groups = vllm.v1.core.kv_cache_utils.get_kv_cache_config_from_groups _orig_max_memory_usage_bytes_from_groups = vllm.v1.core.kv_cache_utils._max_memory_usage_bytes_from_groups _orig_pool_bytes_per_block = vllm.v1.core.kv_cache_utils._pool_bytes_per_block +_orig_max_concurrency = vllm.v1.core.kv_cache_utils.get_max_concurrency_for_kv_cache_config + + +def _ascend_max_concurrency(vllm_config, kv_cache_config): + groups = kv_cache_config.kv_cache_groups + + if has_v41_groups(groups): + return max(0, kv_cache_config.num_blocks - 1) / v41_request_blocks(vllm_config, groups) + return _orig_max_concurrency(vllm_config, kv_cache_config) if UniformTypeKVCacheSpecs.max_num_blocks_per_req is KVCacheSpec.max_num_blocks_per_req: @@ -94,6 +124,18 @@ def _ascend_resolve_kv_cache_block_sizes( dcp = vllm_config.parallel_config.decode_context_parallel_size groups = kv_cache_config.kv_cache_groups + cacheable_groups = [g for g in groups if prefix_cacheable(g.kv_cache_spec)] + if len(cacheable_groups) != len(groups): + scheduler_block_size = math.lcm(*(g.kv_cache_spec.block_size for g in groups)) * dcp + if not cache_config.enable_prefix_caching or not cacheable_groups: + return scheduler_block_size, scheduler_block_size + filtered = replace(kv_cache_config, kv_cache_groups=cacheable_groups) + if dcp == 1: + _, hash_block_size = _orig_resolve_kv_cache_block_sizes(filtered, vllm_config) + else: + hash_block_size = math.gcd(*(g.kv_cache_spec.block_size for g in cacheable_groups)) + return scheduler_block_size, hash_block_size + if len(groups) <= 1: bs = cache_config.block_size * dcp return bs, bs @@ -226,6 +268,8 @@ def group_and_unify_kv_cache_specs( Group the KV cache specs and unify each group into one UniformTypeKVCacheSpecs. Currently, this is only used for DeepseekV4. """ + if (v41_groups := group_v41_cache_specs(kv_cache_spec)) is not None: + return v41_groups if not any(isinstance(spec, SlidingWindowMLASpec) for spec in kv_cache_spec.values()): return None @@ -260,6 +304,8 @@ def _get_kv_cache_groups_uniform_groups( Generate the KV cache groups from the grouped specs. """ assert len(grouped_specs) > 0 and all(isinstance(spec, UniformTypeKVCacheSpecs) for spec in grouped_specs) + if any(is_v41_spec(s) for g in grouped_specs for s in g.kv_cache_specs.values()): + return make_v41_cache_groups(grouped_specs) # For now, we restrict the first grouped_spec to be UniformTypeKVCacheSpecs # containing only MLAAttentionSpec. full_mla_spec = grouped_specs[0] @@ -425,6 +471,8 @@ def _get_kv_cache_config_deepseek_v4( available_memory: int, ) -> tuple[int, list[KVCacheTensor]]: """Plan v0.28.0 DSV4 tensors using the shared_by contract.""" + if has_v41_groups(kv_cache_groups): + return allocate_v41_cache_config(vllm_config, kv_cache_groups, available_memory) page_sizes, bucketed, mtp_layer_names, mtp_page_size, num_layer_tuples = _get_deepseek_v4_cache_layout( kv_cache_groups ) @@ -458,6 +506,8 @@ def _get_kv_cache_config_deepseek_v4_main( kv_cache_groups: list[KVCacheGroupSpec], available_memory: int, ) -> tuple[int, list[KVCacheTensor]]: + if has_v41_groups(kv_cache_groups): + return allocate_v41_cache_config(vllm_config, kv_cache_groups, available_memory) ( page_sizes, bucketed, @@ -541,6 +591,8 @@ def _ascend_pool_bytes_per_block(kv_cache_groups: list[KVCacheGroupSpec]) -> int layout, so using the upstream value changes ``num_blocks`` during the re-plan and leaves ranks inconsistent. """ + if has_v41_groups(kv_cache_groups): + return v41_pool_bytes_per_block(kv_cache_groups) if _get_glm5_next_cache_layout(kv_cache_groups) is not None: return get_glm5_next_pool_bytes_per_block(kv_cache_groups) if not _is_deepseek_v4_groups(kv_cache_groups): @@ -555,6 +607,8 @@ def _ascend_max_memory_usage_bytes_from_groups( kv_cache_groups: list[KVCacheGroupSpec], ) -> int: """Keep the pre-#51718 DSV4 admission formula for its shared tuples.""" + if has_v41_groups(kv_cache_groups): + return (v41_request_blocks(vllm_config, kv_cache_groups) + 1) * v41_pool_bytes_per_block(kv_cache_groups) if _get_glm5_next_cache_layout(kv_cache_groups) is not None: return get_glm5_next_max_memory_usage(vllm_config, kv_cache_groups) if vllm_version_is("0.28.0") or not _is_deepseek_v4_groups(kv_cache_groups): @@ -591,6 +645,14 @@ def _ascend_get_kv_cache_config_from_groups( available_memory: int, ) -> KVCacheConfig: """Restore Ascend's DSV4 shared-tuple planner removed by vLLM #51718.""" + if has_v41_groups(kv_cache_groups): + num_blocks, kv_cache_tensors = allocate_v41_cache_config(vllm_config, kv_cache_groups, available_memory) + return KVCacheConfig( + num_blocks=num_blocks, + kv_cache_tensors=kv_cache_tensors, + kv_cache_groups=kv_cache_groups, + prefix_cache_retention_interval=(vllm_config.cache_config.prefix_cache_retention_interval), + ) if _get_glm5_next_cache_layout(kv_cache_groups) is not None: return get_glm5_next_kv_cache_config(vllm_config, kv_cache_groups, available_memory) if vllm_version_is("0.28.0") or not _is_deepseek_v4_groups(kv_cache_groups): @@ -616,6 +678,7 @@ def _ascend_get_kv_cache_config_from_groups( else: assert _orig_get_packed_kv_cache_groups is not None vllm.v1.core.kv_cache_utils._get_packed_kv_cache_groups = _ascend_get_packed_kv_cache_groups +vllm.v1.core.kv_cache_utils.get_max_concurrency_for_kv_cache_config = _ascend_max_concurrency vllm.v1.core.kv_cache_utils._get_kv_cache_groups_uniform_page_size = _get_kv_cache_groups_uniform_page_size # vLLM v0.24.0 renamed _get_kv_cache_config_deepseek_v4 to # _get_kv_cache_config_packed. The v0.28.0 planner still consumes shared_by; diff --git a/vllm_ascend/patch/platform/patch_speculative_config.py b/vllm_ascend/patch/platform/patch_speculative_config.py index d8f3d6b19398..813195a0ecd0 100644 --- a/vllm_ascend/patch/platform/patch_speculative_config.py +++ b/vllm_ascend/patch/platform/patch_speculative_config.py @@ -63,22 +63,62 @@ def _normalize_deepseek_v4_dspark_draft(draft_model_config) -> None: multimodal architecture conversion. """ hf_config = getattr(draft_model_config, "hf_config", None) + text_config = getattr(hf_config, "text_config", None) + draft_hf_config = text_config if text_config is not None else hf_config + root_model_type = getattr(hf_config, "model_type", None) + text_model_type = getattr(draft_hf_config, "model_type", None) + is_v41 = root_model_type in ("deepseek_v4.1", "deepseek_v41") or text_model_type in ( + "deepseek_v4.1_text", + "deepseek_v41_text", + ) if ( hf_config is None - or getattr(hf_config, "model_type", None) != "deepseek_v4" - or getattr(hf_config, "dspark_target_layer_ids", None) is None + or (root_model_type != "deepseek_v4" and not is_v41) + or getattr(draft_hf_config, "dspark_target_layer_ids", None) is None ): return - hf_config.update({"architectures": ["DSparkDraftModel"]}) + architecture = "DeepseekV41DSparkDraftModel" if is_v41 else "DSparkDraftModel" + if is_v41: + # The Aurora target and draft experts intentionally have different + # widths. SpeculativeConfig owns a private config copy, so adapting + # these fields cannot alter the target model. + draft_hf_config.update( + { + "n_routed_experts": draft_hf_config.dspark_n_routed_experts, + "num_experts_per_tok": getattr(draft_hf_config, "dspark_num_experts_per_tok", None) + or draft_hf_config.dspark_n_activated_experts, + "n_mtp_layers": getattr(draft_hf_config, "num_nextn_predict_layers", 3), + } + ) + if is_v41: + uses_released_name = root_model_type == "deepseek_v41" or text_model_type == "deepseek_v41_text" + normalized_model_type = "deepseek_v41" if uses_released_name else "deepseek_v4.1" + else: + normalized_model_type = str(root_model_type) + hf_config.update( + { + "architectures": [architecture], + "model_type": normalized_model_type, + } + ) + arch_updates = dict( + architectures=[architecture], + model_type=normalized_model_type, + is_mm_prefix_lm=False, + ) + if is_v41: + arch_updates.update( + num_experts=draft_hf_config.n_routed_experts, + num_experts_per_token=draft_hf_config.num_experts_per_tok, + ) draft_model_config.model_arch_config = replace( draft_model_config.model_arch_config, - architectures=["DSparkDraftModel"], - model_type="deepseek_v4", - is_mm_prefix_lm=False, + **arch_updates, ) + architectures = draft_model_config.model_arch_config.architectures model_info, architecture = draft_model_config.registry.inspect_model_cls( - draft_model_config.architectures, + architectures, draft_model_config, ) draft_model_config._model_info = model_info diff --git a/vllm_ascend/quantization/configs/modelslim_config.py b/vllm_ascend/quantization/configs/modelslim_config.py index abbc4866c661..3d8df6ebb01a 100644 --- a/vllm_ascend/quantization/configs/modelslim_config.py +++ b/vllm_ascend/quantization/configs/modelslim_config.py @@ -92,6 +92,14 @@ def modelslim_moe_weight_loader( # Note: Currently, only models that do not have the `packed_modules_mapping` attribute # in the vLLM upstream need to be added here. UPDATED_PACKED_MODULES_MAPPING: dict[str, dict[str, list[str]]] = { + "deepseek_v4.1": { + "gate_up_proj": ["w1", "w3"], + "experts": ["experts.0.w1", "experts.0.w2", "experts.0.w3"], + }, + "deepseek_v41": { + "gate_up_proj": ["w1", "w3"], + "experts": ["experts.0.w1", "experts.0.w2", "experts.0.w3"], + }, # GLM-5.3-Flash (glm5_next): KDA layers ship a fused q/k/v/b/f_a/g_a # projection; sparse-MLA layers keep the DeepSeek-style q_a/kv_a pair. # Native HF FP8 checkpoints leave the KDA projections in bf16 via @@ -162,6 +170,21 @@ def modelslim_moe_weight_loader( "embed.": "model.embed_tokens.", "head.": "lm_head.", }, + "deepseek_v4.1": { + # V4.1 ModelSlim descriptions keep the original checkpoint names, + # while the runtime reuses the V4 module tree. Map runtime prefixes + # back to the checkpoint namespace for quant-scheme lookup. + "language_model.model.layers.": "layers.", + "language_model.model.embed_tokens.": "embed.", + "language_model.model.embed_tokens": "embed", + "language_model.lm_head.": "head.", + "language_model.lm_head": "head", + "model.layers.": "layers.", + "model.embed_tokens.": "embed.", + "model.embed_tokens": "embed", + "lm_head.": "head.", + "lm_head": "head", + }, } @@ -175,6 +198,20 @@ def modelslim_moe_weight_loader( ".ffn_norm.": ".post_attention_layernorm.", ".attn_norm.": ".input_layernorm.", }, + "deepseek_v4.1": { + ".self_attn.": ".attn.", + ".gate_proj.": ".w1.", + ".gate_proj": ".w1", + ".down_proj.": ".w2.", + ".down_proj": ".w2", + ".up_proj.": ".w3.", + ".up_proj": ".w3", + ".mlp.": ".ffn.", + ".post_attention_layernorm.": ".ffn_norm.", + ".post_attention_layernorm": ".ffn_norm", + ".input_layernorm.": ".attn_norm.", + ".input_layernorm": ".attn_norm", + }, # The step3.5 MTP draft nests its decoder block under ".mtp_block.", but the # checkpoint's quant_model_description.json keys it without that infix # (e.g. "model.layers.45.self_attn.q_proj.weight"). Strip it so the quant @@ -199,6 +236,11 @@ def modelslim_moe_weight_loader( }, } +# The released config renamed the V4.1 model type without changing its +# ModelSlim module namespace. Keep pre-release checkpoints compatible. +QUANT_MODEL_PREFIX_MAPPINGS["deepseek_v41"] = QUANT_MODEL_PREFIX_MAPPINGS["deepseek_v4.1"] +QUANT_MODEL_SUBSTR_MAPPINGS["deepseek_v41"] = QUANT_MODEL_SUBSTR_MAPPINGS["deepseek_v4.1"] + def _is_missing_v_shard(shard_key: str, quant_description: dict[str, Any]) -> bool: """Return whether the missing shard is Gemma4's replicated v_proj. @@ -345,6 +387,15 @@ def get_config_filenames(cls) -> list[str]: @classmethod def from_config(cls, config: dict[str, Any]) -> "AscendModelSlimConfig": + # Some ModelSlim checkpoints keep only format metadata in config.json + # and store the per-parameter description in + # quant_model_description.json. Treat that metadata-only form as a + # deferred file load; otherwise maybe_update_config() sees a non-empty + # dict and never reads the actual layer descriptions. + if config.get("quant_method") == ASCEND_QUANTIZATION_METHOD and not any( + isinstance(name, str) and name.endswith(".weight") for name in config + ): + return cls() return cls(config) @classmethod diff --git a/vllm_ascend/spec_decode/dspark_proposer.py b/vllm_ascend/spec_decode/dspark_proposer.py index 9b5bca6d806d..106c570522bb 100644 --- a/vllm_ascend/spec_decode/dspark_proposer.py +++ b/vllm_ascend/spec_decode/dspark_proposer.py @@ -36,7 +36,9 @@ def __init__( ): super().__init__(vllm_config, device, runner=runner) assert vllm_config.speculative_config is not None - self.sample_from_anchor = getattr(self.draft_model_config.hf_config, "sample_from_anchor", True) + hf_config = self.draft_model_config.hf_config + hf_config = getattr(hf_config, "text_config", hf_config) + self.sample_from_anchor = getattr(hf_config, "sample_from_anchor", True) if self.sample_from_anchor: self.num_query_per_req = self.num_speculative_tokens else: @@ -187,6 +189,11 @@ def initialize_attn_backend( builder = attn_group.get_metadata_builder() if isinstance(builder, AscendDSAMetadataBuilder): builder.enable_dspark_device_metadata(self.max_query_tokens) + else: + from vllm_ascend.attention.dsa_v41 import DeepseekV41MetadataBuilder + + if isinstance(builder, DeepseekV41MetadataBuilder): + builder.enable_device_metadata() self.kv_cache_gid = self.draft_attn_groups[0].kv_cache_group_id self.kernel_block_size = self._per_group_kernel_block_sizes[self.kv_cache_gid] diff --git a/vllm_ascend/spec_decode/llm_base_proposer.py b/vllm_ascend/spec_decode/llm_base_proposer.py index 99f0fc2d393a..a51aa9670af7 100644 --- a/vllm_ascend/spec_decode/llm_base_proposer.py +++ b/vllm_ascend/spec_decode/llm_base_proposer.py @@ -726,7 +726,7 @@ def dummy_run( # This is used to hold a position. slot_mapping=self.runner.input_batch.block_table[self.kv_cache_gid].slot_mapping.gpu, positions=self.runner.positions, - positions_cpu=self.runner._dsa_positions_cpu_buf if self.use_compress else None, + positions_cpu=None, attn_state=self.runner.attn_state, decode_token_per_req=self.runner.decode_token_per_req, is_prefilling=torch.zeros(num_reqs, dtype=torch.bool), diff --git a/vllm_ascend/worker/block_table.py b/vllm_ascend/worker/block_table.py index aca379a28f86..fffd92da9aba 100644 --- a/vllm_ascend/worker/block_table.py +++ b/vllm_ascend/worker/block_table.py @@ -10,6 +10,7 @@ ) from vllm.v1.utils import CpuGpuBuffer +from vllm_ascend.core.circular_buffer import is_circular_spec from vllm_ascend.distributed.utils import get_decode_context_model_parallel_world_size from vllm_ascend.ops.triton.compute_slot_mapping import ( _compute_slot_mapping_kernel, @@ -48,6 +49,7 @@ def __init__( self.device = device self.physical_block_size = block_size self.is_mamba_group = is_mamba_group + self.is_circular_group = kv_cache_group is not None and is_circular_spec(kv_cache_group.kv_cache_spec) # If kernel_sizes is None or [0], use physical block size (no splitting) if kernel_sizes is None or kernel_sizes == [0]: @@ -144,6 +146,9 @@ def compute_slot_mapping( query_start_loc: torch.Tensor, positions: torch.Tensor, ) -> None: + if self.is_circular_group: + self.slot_mapping.gpu.fill_(PAD_SLOT_ID) + return num_tokens = positions.shape[0] total_cp_world_size = self.dcp_world_size total_cp_rank = self.dcp_rank @@ -191,6 +196,9 @@ def compute_slot_mapping_draft( # here because M (max_model_len) is not necessarily divisible by # block_size. + if self.is_circular_group: + self.slot_mapping.gpu.fill_(PAD_SLOT_ID) + return if self.dcp_world_size > 1: if not isinstance(req_indices, torch.Tensor): req_indices = torch.from_numpy(req_indices) diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index 55c4f53c25a8..433ef429b520 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -137,6 +137,10 @@ from vllm_ascend.attention.context_parallel.dsa_cp import AscendDSACPMetadataBuilder from vllm_ascend.attention.context_parallel.sfa_cp import AscendSFADCPMetadataBuilder from vllm_ascend.attention.dsa_v1 import AscendDSAMetadataBuilder +from vllm_ascend.attention.dsa_v41 import ( + DeepseekV41CacheLayer, + DeepseekV41MetadataBuilder, +) from vllm_ascend.attention.mla_v1 import AscendMLABackend from vllm_ascend.attention.utils import ( AscendCommonAttentionMetadata, @@ -153,6 +157,12 @@ update_full_graph_params, ) from vllm_ascend.compilation.breakable_aclgraph import BreakableACLGraphWrapper +from vllm_ascend.core.circular_buffer import is_circular_spec +from vllm_ascend.core.deepseek_v41 import ( + is_v41_spec, + plan_cache_slots, + reshape_cache, +) from vllm_ascend.device.hardware_profile import HardwareCapability, get_current_hardware_profile from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.layerwise_cache_layout import ( apply_layerwise_kv_cache_plan, @@ -740,6 +750,7 @@ def _get_eagle3_aux_layers_from_config(self) -> tuple[int, ...] | None: return layer_ids if self.speculative_config.use_dspark(): hf_config = self.speculative_config.draft_model_config.hf_config + hf_config = getattr(hf_config, "text_config", hf_config) # deepseek v4 dspark dspark_layer_ids = getattr(hf_config, "dspark_target_layer_ids", None) if dspark_layer_ids: @@ -1409,7 +1420,9 @@ def _prepare_inputs( # _update_states_after_model_execute for hybrid models). if self.num_accepted_tokens_event is not None: self.num_accepted_tokens_event.synchronize() - self._sync_num_accepted_tokens(num_reqs, has_prev_mapping=bool(prev_req_id_to_index)) + self._sync_num_accepted_tokens( + num_reqs, has_prev_mapping=bool(prev_req_id_to_index) + ) self.num_accepted_tokens.np[num_reqs:].fill(1) self.num_accepted_tokens.copy_to_gpu() else: @@ -2362,13 +2375,18 @@ def execute_model( self.requests, self.mamba_state_idx, ) + if self.use_compress: if deferred_state_corrections_fn: deferred_state_corrections_fn() deferred_state_corrections_fn = None num_reqs = self.input_batch.num_reqs - req_indices = np.repeat(self.arange_np[:num_reqs], num_scheduled_tokens_np) - dsa_positions_np = self._dsa_positions_np_buf[:total_num_scheduled_tokens] + req_indices = np.repeat( + self.arange_np[:num_reqs], num_scheduled_tokens_np + ) + dsa_positions_np = self._dsa_positions_np_buf[ + :total_num_scheduled_tokens + ] np.add( self.input_batch.num_computed_tokens_cpu[req_indices], self.query_pos.np[:total_num_scheduled_tokens], @@ -2447,6 +2465,22 @@ def execute_model( self.model_config.is_encoder_decoder or self.model_config.requires_raw_input_tokens ) + # V4.1's Python reference compressor/indexer path is correctness-safe + # in eager mode, while only uniform decode is prepared for a full ACL + # graph. FULL_DECODE_ONLY dispatches prefills and unsupported decode + # shapes as runtime NONE; bypass the compiled model for those calls so + # the mode is genuinely "eager prefill + full-graph decode". + hf_model_type = getattr(self.model_config.hf_config, "model_type", None) + hf_text_model_type = getattr( + self.model_config.hf_text_config, "model_type", None + ) + is_deepseek_v41 = ( + hf_model_type in ("deepseek_v4.1", "deepseek_v41") + or hf_text_model_type in ("deepseek_v4.1_text", "deepseek_v41_text") + ) + v41_eager_fallback = ( + is_deepseek_v41 and cudagraph_mode == CUDAGraphMode.NONE + ) # Run forward pass defer_kv_connector_finalize = self.speculative_config is not None and ( @@ -2465,7 +2499,7 @@ def execute_model( num_actual_tokens=scheduler_output.total_num_scheduled_tokens, model_instance=self.model, device_metadata_executor=active_device_metadata_executor, - skip_compiled=has_encoder_input, + skip_compiled=has_encoder_input or v41_eager_fallback, has_sinks=self._has_sinks, eplb_heat_collection_status=self.eplb_heat_collection_status if self.dynamic_eplb else False, ), @@ -3052,6 +3086,11 @@ def _model_forward( "inputs_embeds": inputs_embeds, **model_kwargs, } + # Variable Engram routing must run on every DP before ACLGraph capture + # or replay; only its persistent BF16 inputs enter the model graph. + prepare_engram = getattr(self.model, "prepare_engram_inputs", None) + if prepare_engram is not None: + model_inputs.update(prepare_engram(input_ids, positions, num_tokens_padded)) run_model = partial(self.model, **model_inputs) if self.enable_enpu: @@ -3463,6 +3502,8 @@ def _build_attn_group_metadata( attn_gid: int, common_attn_metadata: CommonAttentionMetadata, common_ratio_to_sas_metadata: dict, + common_v41_metadata: dict, + common_v41_batch_metadata: dict, ubid: int | None = None, ) -> None: attn_group = self.attn_groups[kv_cache_gid][attn_gid] @@ -3526,11 +3567,20 @@ def _build_attn_group_metadata( common_ratio_to_sas_metadata=common_ratio_to_sas_metadata, full_graph_mode=cudagraph_runtime_mode == CUDAGraphMode.FULL, ) + elif isinstance(builder, DeepseekV41MetadataBuilder): + extra_attn_metadata_args = dict( + num_actual_reqs=num_reqs, + skip_ring_state_update=skip_gdn_state_update, + common_v41_metadata=common_v41_metadata, + common_v41_batch_metadata=common_v41_batch_metadata, + full_graph_mode=cudagraph_runtime_mode == CUDAGraphMode.FULL, + ) if (for_cudagraph_capture and not isinstance(builder, ( AscendDSAMetadataBuilder, AscendDSACPMetadataBuilder, AscendSFADCPMetadataBuilder, + DeepseekV41MetadataBuilder, ))): attn_metadata_i = builder.build_for_cudagraph_capture(common_attn_metadata) else: @@ -3564,8 +3614,13 @@ def _build_attn_group_metadata( # Prepare the attention metadata for each KV cache group and make layers # in the same group share the same metadata. common_ratio_to_sas_metadata: dict[Any, Any] = {} + common_v41_batch_metadata: dict[str, Any] = {} spec_decode_common_attn_metadata = None for kv_cache_gid, kv_cache_group in enumerate(self.kv_cache_config.kv_cache_groups): + # V4.1 cache coordinates are shared only inside one framework KV + # cache group. This lets a source's LongKV and Indexer reuse the + # same [T, 2] mapping without aliasing any SWA group's mapping. + common_v41_metadata: dict[str, Any] = {} cm = copy(cm_base) # shallow copy # Basically only the encoder seq_lens, block_table and slot_mapping change # for each kv_cache_group. @@ -3590,6 +3645,15 @@ def _build_attn_group_metadata( cm.block_table_tensor, cm.slot_mapping = _get_block_table_and_slot_mapping( kv_cache_gid ) + if isinstance(kv_cache_group.kv_cache_spec, EncoderOnlyAttentionSpec): + cm.block_table_cpu = torch.zeros((num_reqs_padded, 1), dtype=torch.int32, device="cpu") + else: + cm.block_table_cpu = self.input_batch.block_table[kv_cache_gid].get_cpu_tensor()[:num_reqs_padded] + if num_reqs < num_reqs_padded: + # Match the device padding without modifying an H2D source + # that may still be in flight. + cm.block_table_cpu = cm.block_table_cpu.clone() + cm.block_table_cpu[num_reqs:num_reqs_padded].zero_() if self.speculative_config and isinstance(self.drafter, (AscendStep3p5MTPProposer, AscendDSparkProposer)): # step3p5 MTP draft layers span multiple KV cache groups; capture # each group's block table / slot mapping so the proposer can @@ -3617,6 +3681,8 @@ def _build_attn_group_metadata( attn_gid, cm, common_ratio_to_sas_metadata, + common_v41_metadata, + common_v41_batch_metadata, ) if req_doc_ranges is not None: if isinstance(attn_metadata, list): @@ -3752,8 +3818,14 @@ def _dummy_run( num_reqs_padded = batch_desc.num_reqs if batch_desc.num_reqs is not None else num_reqs if num_tokens_across_dp is not None and num_tokens_padded != num_tokens: # pad is needed if the pad of `num_tokens` is triggered inside CudagraphDispatcher - num_tokens_across_dp[:] = num_tokens_padded - num_scheduled_tokens = num_scheduled_tokens.repeat(num_reqs_padded) + if _cudagraph_mode == CUDAGraphMode.NONE: + # Preserve the eager DP dummy-run contract used by legacy DSA + # CP. Graph-only padding below must not turn its replicated + # request metadata into an empty request on idle DP ranks. + num_scheduled_tokens = num_scheduled_tokens.repeat(num_reqs_padded) + else: + # Requests added by CudagraphDispatcher have zero query length. + num_scheduled_tokens = np.pad(num_scheduled_tokens, (0, num_reqs_padded - num_reqs)) if self.dynamic_eplb: self.update_eplb_heat_collection_status(num_tokens_padded) @@ -3840,6 +3912,19 @@ def _dummy_run( # Dummy graph runs do not go through _prepare_inputs(), but GDN/Mamba # metadata reads block_table[:num_reqs_padded] below. Sync padded # rows as well so device-side metadata does not see stale block ids. + # Dummy requests bypass scheduler allocation. Give each active + # request a distinct non-null state ID before metadata/capture. + for gid, group in enumerate(self.kv_cache_config.kv_cache_groups): + if skip_gdn_state_update or not is_circular_spec(group.kv_cache_spec): + continue + if num_reqs >= self.kv_cache_config.num_blocks: + raise ValueError("Insufficient ring pages for dummy graph requests") + table = self.input_batch.block_table[gid] + table.block_table.np[:num_reqs_padded].fill(0) + table.block_table.np[:num_reqs, 0] = np.arange(1, num_reqs + 1) + context = self.compilation_config.static_forward_context + for name in group.layer_names: + context[name].kv_cache[0][1:num_reqs + 1].zero_() self.input_batch.block_table.commit_block_table(num_reqs_padded) # Invalidate real-request slots before attention backends derive @@ -3848,6 +3933,18 @@ def _dummy_run( for kv_cache_gid in range(len(self.kv_cache_config.kv_cache_groups)): blk_table = self.input_batch.block_table[kv_cache_gid] blk_table.slot_mapping.gpu.fill_(-1) + else: + for kv_cache_gid, group in enumerate(self.kv_cache_config.kv_cache_groups): + group_spec = group.kv_cache_spec + if not isinstance(group_spec, UniformTypeKVCacheSpecs) or not any( + is_v41_spec(spec) for spec in group_spec.kv_cache_specs.values() + ): + continue + # V4.1 derives backend-specific 2D slot mappings from + # this buffer. Dummy capture has no scheduler-owned + # slots, so stale active entries must not write caches. + blk_table = self.input_batch.block_table[kv_cache_gid] + blk_table.slot_mapping.gpu[:num_tokens_padded].fill_(-1) pad_attn = cudagraph_runtime_mode == CUDAGraphMode.FULL # check how to build dummy @@ -4248,6 +4345,14 @@ def initialize_kv_cache(self, kv_cache_config: KVCacheConfig) -> None: self.sparse_kv_offload_config, ) kv_caches = self.initialize_kv_cache_tensors(kv_cache_config) + if any(is_circular_spec(g.kv_cache_spec) for g in kv_cache_config.kv_cache_groups): + # Lazy import avoids the model/cache registration cycle. + from vllm_ascend.models.deepseek_v41.compressor import DeepseekV41Compressor + + for module in self.model.modules(): + if isinstance(module, DeepseekV41Compressor) and module.ratio == 2: + module.prepare_ring_compressor(self.max_num_tokens, self.device) + # TODO: refactor the logic of attention if ( self.speculative_config @@ -4326,7 +4431,16 @@ def initialize_kv_cache_tensors(self, kv_cache_config: KVCacheConfig) -> dict[st logger.debug("%s reuses KV cache of %s", layer_name, target_layer_name) kv_caches[layer_name] = kv_caches[target_layer_name] - if self.model_config.hf_text_config.model_type == "deepseek_v4": + if any( + isinstance(self.compilation_config.static_forward_context.get(name), DeepseekV41CacheLayer) + for name in kv_caches + ): + if self.kv_caches: + raise ValueError("V4.1 cache tensors were already bound") + for name in sorted(kv_caches): + self.compilation_config.static_forward_context[name].kv_cache = [kv_caches[name]] + self.kv_caches.append(kv_caches[name]) + elif self.model_config.hf_text_config.model_type == "deepseek_v4": from vllm_ascend.utils import extract_dsv4_layer_index assert len(self.kv_caches) == 0 @@ -4511,6 +4625,29 @@ def _allocate_kv_cache_tensors(self, kv_cache_config: KVCacheConfig) -> dict[str # prefill disaggregation need the addr of cache tensor be aligned with 2M alignment = 2 * 1024 * 1024 layer_kv_cache_spec = self._get_layer_kv_cache_specs(kv_cache_config) + if any(is_v41_spec(spec) for spec in layer_kv_cache_spec.values()): + if not all(is_v41_spec(spec) for spec in layer_kv_cache_spec.values()): + raise ValueError("Mixed V4.1 cache allocation is not supported") + slots = plan_cache_slots(layer_kv_cache_spec) + if len(kv_cache_config.kv_cache_tensors) != len(slots): + raise ValueError("V4.1 requires one allocation per layer slot") + for allocation, slot in zip(kv_cache_config.kv_cache_tensors, slots): + allocation_layers = get_kv_cache_tensor_layers(allocation) + if ( + allocation.offset + or allocation.block_stride != slot.page_size_bytes + or allocation.size != kv_cache_config.num_blocks * slot.page_size_bytes + or allocation_layers != [p.name for p in slot.placements] + ): + raise ValueError("V4.1 allocation disagrees with its layer slot") + backing = self._allocate_int8_cache_tensor(allocation.size, alignment) + for name in allocation_layers: + kv_cache_raw_tensors[name] = backing + expected = set(layer_kv_cache_spec) + if set(kv_cache_raw_tensors) != expected: + raise ValueError("V4.1 cache descriptors do not cover every resource") + return kv_cache_raw_tensors + # v0.28.0 keeps the legacy ``shared_by`` contract: one allocation per # descriptor, shared by every listed layer. Main uses #51718's # standardized descriptors, whose ``size`` is the size of one common @@ -4998,6 +5135,14 @@ def _reshape_kv_cache_tensors( """ kv_caches: dict[str, torch.Tensor] = {} layer_kv_cache_spec = self._get_layer_kv_cache_specs(kv_cache_config) + layer_placements = {} + if any(is_v41_spec(spec) for spec in layer_kv_cache_spec.values()): + layer_placements = { + p.name: (p.offset, slot.page_size_bytes) + for slot in plan_cache_slots(layer_kv_cache_spec) + for p in slot.placements + } + for group in self._kv_cache_spec_attn_group_iterator(): attn_backend = group.backend current_kv_cache_spec = group.kv_cache_spec @@ -5007,6 +5152,17 @@ def _reshape_kv_cache_tensors( current_kv_cache_spec = layer_kv_cache_spec[layer_name] + if is_v41_spec(current_kv_cache_spec): + offset, block_stride = layer_placements[layer_name] + kv_caches[layer_name] = reshape_cache( + kv_cache_raw_tensors[layer_name], + current_kv_cache_spec, + num_blocks=kv_cache_config.num_blocks, + offset=offset, + block_stride=block_stride, + ) + continue + # TODO: remove this after the OOM issue is located and fixed, otherwise, some model may # encounter OOM issue if self._uses_page_strided_kv_layout(current_kv_cache_spec): @@ -5427,6 +5583,8 @@ def may_reinitialize_input_batch(self, kv_cache_config: KVCacheConfig) -> None: kv_cache_spec = next(iter(kv_cache_spec.kv_cache_specs.values())) if isinstance(kv_cache_spec, EncoderOnlyAttentionSpec): continue + elif is_circular_spec(kv_cache_spec): + self.kernel_block_sizes.append([kv_cache_spec.block_size]) elif isinstance(kv_cache_spec, AttentionSpec): # This is an attention backend that supports virtual # block splitting. Get the supported block sizes from @@ -5646,6 +5804,8 @@ def get_kv_cache_spec(self) -> dict[str, KVCacheSpec]: # or enable more requests to be processed simultaneously. self.shared_kv_cache_layers[layer_name] = kv_tgt_layer continue + elif isinstance(attn_module, DeepseekV41CacheLayer): + kv_cache_spec[layer_name] = attn_module.get_kv_cache_spec(self.vllm_config) elif self.use_compress: # Skip modules that don't need KV cache (eg encoder-only attention) if spec := attn_module.get_kv_cache_spec(self.vllm_config): @@ -5804,25 +5964,15 @@ def _check_and_update_cudagraph_mode( tensor_parallel_size = self.parallel_config.tensor_parallel_size resolver_tensor_parallel_size = tensor_parallel_size if ( - self.compilation_config.pass_config.enable_sp - and self.uniform_decode_query_len > 1 - and tensor_parallel_size > 1 + self.compilation_config.cudagraph_mode.decode_mode() == CUDAGraphMode.FULL + and (enable_dsa_cp() or enable_sp(self.vllm_config) or self.compilation_config.pass_config.enable_sp) ): - graph_alignment = math.lcm( - self.uniform_decode_query_len, - tensor_parallel_size, - ) - capture_sizes = self.compilation_config.cudagraph_capture_sizes - # vLLM 0.27 has no path for explicit capture sizes that are - # already aligned to both speculative steps and SP. Skip its - # redundant TP adjustment only for that exact case. - if ( - graph_alignment - > max(self.uniform_decode_query_len, tensor_parallel_size) - and capture_sizes - and all(size % graph_alignment == 0 for size in capture_sizes) - ): - resolver_tensor_parallel_size = 1 + # CP and SP pad tokens to TP. Align capture keys to both TP + # and the speculative query length before the v0.27 resolver, + # whose max(query_len, TP) rejects non-divisible pairs (6, 8). + graph_alignment = math.lcm(self.uniform_decode_query_len, tensor_parallel_size) + self.compilation_config.adjust_cudagraph_sizes_for_spec_decode(graph_alignment, 1) + resolver_tensor_parallel_size = 1 cudagraph_mode = self.compilation_config.resolve_cudagraph_mode_and_sizes( min_cg_support=min_cg_support, min_cg_attn_backend=min_cg_attn_backend, diff --git a/vllm_ascend/worker/worker.py b/vllm_ascend/worker/worker.py index dfff01429ce8..e499c9b57fce 100644 --- a/vllm_ascend/worker/worker.py +++ b/vllm_ascend/worker/worker.py @@ -1218,7 +1218,7 @@ def reset_encoder_cache(self) -> None: def execute_dummy_batch(self) -> None: self.log_memory_stats() num_tokens = getattr(self.model_runner, "uniform_decode_query_len", 1) - self.model_runner._dummy_run(num_tokens, uniform_decode=True) + self.model_runner._dummy_run(num_tokens, uniform_decode=True, skip_gdn_state_update=True) def _init_worker_distributed_environment(self) -> None: """Initialize the distributed environment."""