已关闭
[Feature]: npu_sparse_flash_attention 支持 value=None 特性 #423
nomiz创建于  8月18日关闭于  8月25日
nomiz成员
8月18日 创建

提交提案之前,请先检索仓库内是否已有相同的提案,如已有请在同一提案中进行讨论。

💻 需求背景、当前现状、期望实现的功能内容、具体的设计方案、以及测试方案

1 背景与目标

1.1 问题背景

npu_sparse_flash_attention 是面向稀疏注意力(QSFA 等)场景的融合算子,CANN 9.2.0 起,aclnnSparseFlashAttention 算子层支持可空 value 输入,本提案在框架层放开 value=None 。

1.2 当前现状

  • schema 中 value 为必填 Tensor(位置参数,位于 sparse_indices 之前);
  • 前向 kernel 直接对 value 做非空校验并透传 aclnn;
  • 反向 kernel 对 value 做维度校验并构造 d_value;
  • Python meta 注册将 value 列入必填参数校验,value=None 在 fake/meta 路径下会直接报错;
  • 无版本/SoC 门控,是否支持完全取决于底层 aclnn 接口。

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);
  • aclnn 调用处统一传入归一化后的 value_const。

3 测试设计

新增用例 test_sfa_value_none(test/test_custom_ops/test_npu_sparse_flash_attention.py):

替代方案

补充说明

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
Nnomiz成员
8月18日 添加了label:feature
Nnomiz成员
8月21日 修改了issue 的描述
ascend-robotascend-robot成员
8月25日 关闭了 issue
ascend-robotascend-robot成员
8月25日 添加了label:resolved
Nnomiz成员
18 天前 修改了issue 的描述
ascend-robotascend-robot成员
18 天前 issue状态由 TODO 改变为 DONE