已合并
SFA资料修改/补充sparse_size为0的用例 #3753
zzzyh22创建于 2025年12月10日
SFA资料修改/补充sparse_size为0的用例 #3753
已合并
zzzyh22创建于 2025年12月10日
共 3 个文件变更+75-19
@@ -25,10 +25,9 @@ torch_npu.npu_kv_quant_sparse_flash_attention(query, key, value, sparse_indices,
25 25 
26## 参数说明26## 参数说明
27 27 
28->**说明:**<br> 28+> [!NOTE]
29->29+> - query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。
30->- query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。30+> - Q\_S和S1表示query shape中的S,KV\_S和S2表示key shape中的S,Q\_N表示num\_query\_heads,KV\_N表示num\_key\_value\_heads。
31->- Q\_S和S1表示query shape中的S,KV\_S和S2表示key shape中的S,Q\_N表示num\_query\_heads,KV\_N表示num\_key\_value\_heads。
32 31 
33- **query**(`Tensor`):必选参数,表示attention结构的Q输入,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`和`float16`,query相同dtype的q_nope和q_rope按D维度拼接得到,query的N支持1、2、4、8、16、32、64、128。32- **query**(`Tensor`):必选参数,表示attention结构的Q输入,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`和`float16`,query相同dtype的q_nope和q_rope按D维度拼接得到,query的N支持1、2、4、8、16、32、64、128。
34 33 
@@ -52,11 +51,9 @@ torch_npu.npu_kv_quant_sparse_flash_attention(query, key, value, sparse_indices,
52 51 
53- **block\_table**(`Tensor`):可选参数,表示PageAttention中kvCache存储使用的block映射表。数据格式支持$ND$,数据类型支持`int32`,shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的s2对应的block数量,即s2\_max / block\_size向上取整。52- **block\_table**(`Tensor`):可选参数,表示PageAttention中kvCache存储使用的block映射表。数据格式支持$ND$,数据类型支持`int32`,shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的s2对应的block数量,即s2\_max / block\_size向上取整。
54 53 
55-- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。54+- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该参数中每个Batch的有效token数不超过`query`中的维度S大小。支持长度为B的一维tensor。<br>当`query`的input\_layout为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。不能出现负值。
56- >该入参中每个Batch的有效token数不超过`query`中的维度S大小。支持长度为B的一维tensor。当`query`的input\_layout为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。不能出现负值。
57 55 
58-- **actual\_seq\_lengths\_kv**(`Tensor`):可选参数,表示不同Batch中`key`和`value`的有效token数,数据类型支持`int32`。如果不指定None,表示和key的shape的S长度相同。56+- **actual\_seq\_lengths\_kv**(`Tensor`):可选参数,表示不同Batch中`key`和`value`的有效token数,数据类型支持`int32`。如果不指定None,表示和key的shape的S长度相同。该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。
59- >该入参中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。
60 57 
61- **sparse\_block\_size**(`int`):可选参数,代表sparse阶段的block大小,在计算importance score时使用,支持范围为[1, 16],数据类型支持`int64`。58- **sparse\_block\_size**(`int`):可选参数,代表sparse阶段的block大小,在计算importance score时使用,支持范围为[1, 16],数据类型支持`int64`。
62 59 
@@ -26,8 +26,9 @@ torch_npu.npu_sparse_flash_attention(query, key, value, sparse_indices, scale_va
26 26 
27## 参数说明27## 参数说明
28 28 
29-> query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。\29+> [!NOTE]
30-> Q\_S和S1表示query shape中的S,KV\_S和S2表示key shape中的S,Q\_N和N1表示num\_query\_heads,KV\_N和N2表示num\_key\_value\_heads。30+>- query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。
31+>- Q\_S和S1表示query shape中的S,KV\_S和S2表示key shape中的S,Q\_N和N1表示num\_query\_heads,KV\_N和N2表示num\_key\_value\_heads。
31- **query**(`Tensor`):必选参数,对应公式中的$Q$,不支持非连续,维度N支持1/2/4/8/16/32/64/128,数据格式支持ND,数据类型支持`bfloat16`和`float16`。 32- **query**(`Tensor`):必选参数,对应公式中的$Q$,不支持非连续,维度N支持1/2/4/8/16/32/64/128,数据格式支持ND,数据类型支持`bfloat16`和`float16`。
32- **key**(`Tensor`):必选参数,对应公式中的$\tilde{K}$,不支持非连续,维度N只支持1,数据格式支持ND,数据类型支持`bfloat16`和`float16`,layout\_kv为PA\_BSND时shape为[block\_num, block\_size, KV\_N, D],其中block\_num为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的倍数,最大支持1024。33- **key**(`Tensor`):必选参数,对应公式中的$\tilde{K}$,不支持非连续,维度N只支持1,数据格式支持ND,数据类型支持`bfloat16`和`float16`,layout\_kv为PA\_BSND时shape为[block\_num, block\_size, KV\_N, D],其中block\_num为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的倍数,最大支持1024。
33 34 
@@ -37,21 +38,21 @@ torch_npu.npu_sparse_flash_attention(query, key, value, sparse_indices, scale_va
37 38 
38- **scale\_value**(`double`):必选参数,代表缩放系数,作为query和key矩阵乘后Muls的scalar值,数据类型支持`float`。39- **scale\_value**(`double`):必选参数,代表缩放系数,作为query和key矩阵乘后Muls的scalar值,数据类型支持`float`。
39 40 
40-- <strong>*</strong>:代表其之前的参数是位置相关的,必须按照顺序输入,属于必选参数;其之后的参数是键值对赋值,与位置无关,属于可选参数(不传入会使用默认值)。41+- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
41 42 
42- **block\_table**(`Tensor`):可选参数,表示PageAttention中kvCache存储使用的block映射表。数据格式支持ND,数据类型支持`int32`,shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的S2对应的block数量,即S2\_max / block\_size向上取整。43- **block\_table**(`Tensor`):可选参数,表示PageAttention中kvCache存储使用的block映射表。数据格式支持ND,数据类型支持`int32`,shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的S2对应的block数量,即S2\_max / block\_size向上取整。
43 44 
44-- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。45+- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该入参中每个Batch的有效token数不超过`query`中的维度S大小。支持长度为B的一维tensor。<br>当`query`的input\_layout为TND时,该入参必须传入,且以该入参元素的数量作为B值,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须>=前一个元素的值。不能出现负值。
45- >该入参中每个Batch的有效token数不超过`query`中的维度S大小。支持长度为B的一维tensor。当`query`的input\_layout为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须>=前一个元素的值。不能出现负值。
46 46 
47-- **actual\_seq\_lengths\_kv**(`Tensor`):可选参数,表示不同Batch中`key`和`value`的有效token数,数据类型支持`int32`。如果不指定None,表示和key的shape的S长度相同。47+- **actual\_seq\_lengths\_kv**(`Tensor`):可选参数,表示不同Batch中`key`和`value`的有效token数,数据类型支持`int32`。如果不指定None,表示和key的shape的S长度相同。该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。
48- >该入参中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。
49 48 
50- **query\_rope**(`Tensor`):可选参数,表示MLA结构中的query的rope信息,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。49- **query\_rope**(`Tensor`):可选参数,表示MLA结构中的query的rope信息,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。
51 50
52- **key\_rope**(`Tensor`):可选参数,表示MLA结构中的key的rope信息,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。51- **key\_rope**(`Tensor`):可选参数,表示MLA结构中的key的rope信息,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。
53 52 
54-- **sparse\_block\_size**(`int`):可选参数,代表sparse阶段的block大小,在计算importance score时使用,数据类型支持`int64`。53+- **sparse\_block\_size**(`int`):可选参数,代表sparse阶段的block大小,在计算importance score时使用,数据类型支持`int64`,取值范围为[1,128]。
54+ - sparse_block_size为1时,为Token-wise稀疏化场景,将每个token视为独立单元,在计算重要性分数时,评估每个查询token与每个键值token之间的独立关联程度。
55+ - sparse_block_size为大于1小于等于128时,为Block-wise稀疏化场景,将token序列划分为固定大小的连续块,以块为单位进行重要性评估,块内token共享相同的稀疏化决策。
55 56 
56- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,用户不特意指定时可传入默认值"BSND",支持传入BSND和TND。57- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,用户不特意指定时可传入默认值"BSND",支持传入BSND和TND。
57 58 
@@ -65,7 +66,7 @@ torch_npu.npu_sparse_flash_attention(query, key, value, sparse_indices, scale_va
65 66 
66- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。67- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。
67 68 
68-- **attention\_mode**(`int`):可选参数,表示attention的模式,数据类型支持`int64`,仅支持传入2,表示MLA-absorb模式,即query和key的D包含rope和nope两部分,且key和value是同一份。69+- **attention\_mode**(`int`):可选参数,表示attention的模式,数据类型支持`int64`,仅支持传入2,表示MLA-absorb模式,即计算过程中会将query和key的nope部分分别和query_rope和key_rope的rope部分沿头维度(D)拼接,合并形成最终的query和key用于后续计算,且key和value共享同一份底层张量数据。
69 70 
70- **return\_softamx\_lse**(`bool`):可选参数,用于表示是否返回softmax_max和softmax_sum,默认值`False`。71- **return\_softamx\_lse**(`bool`):可选参数,用于表示是否返回softmax_max和softmax_sum,默认值`False`。
71 72 
@@ -223,7 +223,6 @@ class TestSparseFlashAttention(TestCase):
223 act_seq_q = torch.tensor(act_seq_q).to(torch.int32).npu()223 act_seq_q = torch.tensor(act_seq_q).to(torch.int32).npu()
224 act_seq_kv = torch.tensor(act_seq_kv).to(torch.int32).npu()224 act_seq_kv = torch.tensor(act_seq_kv).to(torch.int32).npu()
225 225 
226- print(f'======================== PTA eager BEGIN ========================')
227 # start run custom ops226 # start run custom ops
228 npu_out, npu_softmax_max, npu_softmax_sum = torch_npu.npu_sparse_flash_attention(227 npu_out, npu_softmax_max, npu_softmax_sum = torch_npu.npu_sparse_flash_attention(
229 query, key, value, sparse_indices, scale_value, block_table=None, 228 query, key, value, sparse_indices, scale_value, block_table=None,
@@ -247,7 +246,66 @@ class TestSparseFlashAttention(TestCase):
247 print("cpu output:\n", cpu_out, cpu_out.shape)246 print("cpu output:\n", cpu_out, cpu_out.shape)
248 print("correct ratio of cpu vs npu is:", true_ratio * 100, "%")247 print("correct ratio of cpu vs npu is:", true_ratio * 100, "%")
249 self.assertTrue(true_ratio > 0.99, "precision compare fail")248 self.assertTrue(true_ratio > 0.99, "precision compare fail")
250- print(f'======================== PTA eager FINISH ========================')249+ 
250+ 
251+ @unittest.skip("Skipping due to outdated CANN version; please update CANN to the latest version and remove this skip")
252+ def test_sfa_saprse_size_zero(self, device = "npu"):
253+ scale_value = 0.041666666666666664
254+ sparse_block_size = 1
255+ query_type = torch.float16
256+ scale_value = 0.041666666666666664
257+ sparse_block_size = 1
258+ sparse_size = 0
259+ t = 10
260+ b = 4
261+ s1 = 1
262+ s2 = 8192
263+ n1 = 128
264+ n2 = 1
265+ dn = 512
266+ dr = 64
267+ tile_size = 128
268+ block_size = 256
269+ s2_act = 4096
270+ attention_mode = 2
271+ return_softmax_lse = False
272+ 
273+ query = torch.tensor(np.random.uniform(-10, 10, (b, s1, n1, dn))).to(query_type)
274+ key = torch.tensor(np.random.uniform(-5, 10, (b, s2, n2, dn))).to(query_type)
275+ value = key.clone()
276+ idxs = random.sample(range(s2_act - s1 + 1), sparse_size)
277+ sparse_indices = torch.tensor([idxs for _ in range(b * s1 * n2)]).reshape(b, s1, n2, sparse_size). \
278+ to(torch.int32)
279+ query_rope = torch.tensor(np.random.uniform(-10, 10, (b, s1, n1, dr))).to(query_type)
280+ key_rope = torch.tensor(np.random.uniform(-10, 10, (b, s2, n2, dr))).to(query_type)
281+ act_seq_q = [s1] * b
282+ act_seq_kv = [s2_act] * b
283+ 
284+ query = query.npu()
285+ key = key.npu()
286+ value = value.npu()
287+ sparse_indices = sparse_indices.npu()
288+ query_rope = query_rope.npu()
289+ key_rope = key_rope.npu()
290+ act_seq_q = torch.tensor(act_seq_q).to(torch.int32).npu()
291+ act_seq_kv = torch.tensor(act_seq_kv).to(torch.int32).npu()
292+ 
293+ # start run custom ops
294+ npu_out, npu_softmax_max, npu_softmax_sum = torch_npu.npu_sparse_flash_attention(
295+ query, key, value, sparse_indices, scale_value, block_table=None,
296+ actual_seq_lengths_query=act_seq_q, actual_seq_lengths_kv=act_seq_kv,
297+ query_rope=query_rope, key_rope=key_rope, sparse_block_size=sparse_block_size,
298+ layout_query='BSND', layout_kv='BSND', sparse_mode=3, pre_tokens=(1<<63)-1, next_tokens=(1<<63)-1,
299+ attention_mode = attention_mode, return_softmax_lse = return_softmax_lse)
300+ 
301+ # compare result
302+ cpu_out = self.cpu_sparse_flash_attention(
303+ query, key, value, sparse_indices, scale_value, sparse_block_size,
304+ actual_seq_lengths_query=act_seq_q, actual_seq_lengths_kv=act_seq_kv,
305+ query_rope=query_rope, key_rope=key_rope,
306+ layout_query='BSND', layout_kv='BSND', sparse_mode=3, block_table=None)
307+ npu_out = npu_out.cpu().to(torch.float32).numpy()
308+ self.assertRtolEqual(cpu_out, npu_out, prec=0.01)
251 309 
252if __name__ == "__main__":310if __name__ == "__main__":
253 run_tests()311 run_tests()