已合并
SFA资料修改/补充sparse_size为0的用例 #3753
zzzyh22创建于 2025年12月10日
SFA资料修改/补充sparse_size为0的用例 #3753
已合并
共 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 ops | 226 | # 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 | + | ||
| 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 | ||
| 252 | if __name__ == "__main__": | 310 | if __name__ == "__main__": |
| 253 | run_tests() | 311 | run_tests() |