if (!value.has_value()) {
TORCH_CHECK(
IsSupportSparseFlashAttentionNoneValue(),
"npu_sparse_flash_attention: value=None is only supported with CANN >= 9.2.0.",
OPS_ERROR(ErrCode::NOT_SUPPORT));
}
value 继续直接透传给 aclnnSparseFlashAttention / aclnnSparseFlashAttentionV2(经 EXEC_NPU_NO_FORMAT_CHECK_CMD 宏统一参数转换),可空语义由算子层处理;sinks 形态(V2 接口)的调用路径保持不变。
提交提案之前,请先检索仓库内是否已有相同的提案,如已有请在同一提案中进行讨论。
💻 需求背景、当前现状、期望实现的功能内容、具体的设计方案、以及测试方案
1 背景与目标
1.1 问题背景
npu_sparse_flash_attention是面向稀疏注意力(QSFA 等)场景的融合算子,CANN 9.2.0 起,aclnnSparseFlashAttention算子层支持可空value输入,本提案在框架层放开value=None。1.2 当前现状
value为必填Tensor(位置参数,位于sparse_indices之前);value做非空校验并透传 aclnn;value做维度校验并构造d_value;value列入必填参数校验,value=None在 fake/meta 路径下会直接报错;2 总体设计
2.1 对外接口设计
前向 schema(仅
value参数发生变化,其余不变):- npu_sparse_flash_attention(Tensor query, Tensor key, Tensor value, Tensor sparse_indices, float scale_value, *, ...) + npu_sparse_flash_attention(Tensor query, Tensor key, Tensor? value, Tensor sparse_indices, float scale_value, *, ...)反向 schema 同步变更:
- npu_sparse_flash_attention_grad(Tensor query, Tensor key, Tensor value, Tensor sparse_indices, Tensor d_out, ...) + npu_sparse_flash_attention_grad(Tensor query, Tensor key, Tensor? value, Tensor sparse_indices, Tensor d_out, ...)2.2 版本门控设计
新增 C++ 门控函数(
op_plugin/utils/op_api_common.h):// Check if npu_sparse_flash_attention supports value=None: CANN >= 9.2.0. inline bool IsSupportSparseFlashAttentionNoneValue() { static const bool is_support = []() -> bool { return op_plugin::utils::is_gte_cann_version_920() }(); return is_support; }2.3 前向 kernel(SparseFlashAttentionKernelNpuOpApi.cpp)
std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_sparse_flash_attention( const at::Tensor& query, const at::Tensor& key, const c10::optional<at::Tensor>& value, // 原 const at::Tensor& ...query/key/sparse_indices非空校验之后新增门控:if (!value.has_value()) { TORCH_CHECK( IsSupportSparseFlashAttentionNoneValue(), "npu_sparse_flash_attention: value=None is only supported with CANN >= 9.2.0.", OPS_ERROR(ErrCode::NOT_SUPPORT)); }value继续直接透传给aclnnSparseFlashAttention/aclnnSparseFlashAttentionV2(经EXEC_NPU_NO_FORMAT_CHECK_CMD宏统一参数转换),可空语义由算子层处理;sinks形态(V2 接口)的调用路径保持不变。2.4 反向 kernel(SparseFlashAttentionGradKernelNpuOpApi.cpp)
const at::Tensor& value→const c10::optional<at::Tensor>& value;const at::Tensor& value_const = value.value_or(at::Tensor());value_const.defined():保留原 3/4 维校验;value=None):走与 2.3 节相同的门控校验,不满足则抛NOT_SUPPORT;value=None时d_value = at::empty({0}, query.options()),否则沿用apply_tensor_without_format(value_const);value_const。3 测试设计
新增用例
test_sfa_value_none(test/test_custom_ops/test_npu_sparse_flash_attention.py):替代方案
补充说明
欢迎加入社区,感谢您对社区的贡献 🎉!