Files
enginex-ascend-910-vllm/csrc/attention/compressor/op_host/compressor_def.cpp
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

192 lines
7.8 KiB
C++
Raw Blame History

This file contains invisible Unicode characters

This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/**
 * 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 "register/op_def_registry.h"
namespace ops {
class Compressor : public OpDef {
public:
static constexpr uint32_t ROPE_HEAD_DIM_VALUE = 64;
static constexpr uint32_t CMP_RATIO_VALUE = 4;
static constexpr uint32_t COFF_VALUE = 1;
static constexpr uint32_t ROTARY_MODE_VALUE = 1;
static constexpr uint32_t CACHE_MODE_VALUE = 1;
static constexpr uint32_t STATE_CACHE_STRIDE_DIM0 = 0;
explicit Compressor(const char *name) : OpDef(name)
{
this->Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("wkv")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("wgate")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("state_cache")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.IgnoreContiguous();
this->Input("ape")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("norm_weight")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("rope_sin")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("rope_cos")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("state_block_table")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("cu_seqlens")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("seqused")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Input("start_pos")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
this->Output("cmp_kv")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND});
this->Output("state_cache")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND});
this->Attr("rope_head_dim").AttrType(REQUIRED).Int(ROPE_HEAD_DIM_VALUE);
this->Attr("cmp_ratio").AttrType(REQUIRED).Int(CMP_RATIO_VALUE);
this->Attr("coff").AttrType(OPTIONAL).Int(COFF_VALUE);
this->Attr("norm_eps").AttrType(OPTIONAL).Float(1e-6f);
this->Attr("rotary_mode").AttrType(OPTIONAL).Int(ROTARY_MODE_VALUE);
this->Attr("cache_mode").AttrType(OPTIONAL).Int(CACHE_MODE_VALUE);
this->Attr("state_cache_stride_dim0").AttrType(OPTIONAL).Int(STATE_CACHE_STRIDE_DIM0);
OpAICoreConfig aicore_config;
aicore_config.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.PrecisionReduceFlag(true)
.ExtendCfgInfo("aclnnSupport.value", "support_aclnn"); // set value of aclnn support
OpAICoreConfig config910;
config910.Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("wkv")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("wgate")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("state_cache")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.IgnoreContiguous();
config910.Input("ape")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("norm_weight")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("rope_sin")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("rope_cos")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("state_block_table")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("cu_seqlens")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("seqused")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Input("start_pos")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32})
.FormatList({ge::FORMAT_ND})
.AutoContiguous();
config910.Output("cmp_kv")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
.FormatList({ge::FORMAT_ND});
config910.Output("state_cache")
.ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT})
.FormatList({ge::FORMAT_ND});
config910.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.PrecisionReduceFlag(true)
.ExtendCfgInfo("aclnnSupport.value", "support_aclnn");
this->AICore().AddConfig("ascend910b", config910);
this->AICore().AddConfig("ascend910_93", config910);
this->AICore().AddConfig("ascend950", aicore_config);
}
};
OP_ADD(Compressor, optiling::CompressorCompileInfo);
} // namespace ops