diff --git a/csrc/attention/msa_index_score/README.md b/csrc/attention/msa_index_score/README.md index d17707ca5f26..ea943d4e6b71 100644 --- a/csrc/attention/msa_index_score/README.md +++ b/csrc/attention/msa_index_score/README.md @@ -4,9 +4,9 @@ | Product | Supported | | --------------------------------------------------------------------- | :-------: | -| Atlas A2 Products | √ | -| Atlas A3 Products | √ | -| 950PR&950DT Products | √ | +| Atlas A2 Products | √ | +| Atlas A3 Products | √ | +| 950PR&950DT Products | √ | ## Function Description @@ -37,8 +37,6 @@ $block\_size$. `start_loc`, `init_blocks`, and `local_blocks` generate $local\_mask`, which assigns high scores to leading blocks and blocks around the current query so that TopK always retains them. Set both block-count attributes to 0 to disable this behavior and match the Triton raw-score kernel. -When the two windows overlap, the local-window score (`1e29`) overrides -the leading-block score (`1e30`). ## Parameters @@ -53,29 +51,21 @@ Notation: is the token count per page, and `maxBlockNumPerSeq` is the width of `block_table`. -| Parameter | ACLNN Name | Kind | Description | Data Type | Format | -| --------- | ---------- | ---- | ----------- | --------- | ------ | -| `query` | `query` | Input | Query tensor in TND layout, shape `[T1, N1, D]`. | BFLOAT16, FLOAT16, HIFLOAT8, FLOAT8_E5M2, FLOAT8_E4M3FN | ND | -| `key` | `key` | Input | Key tensor in TND `[T2, N2, D]`, BNBD `[block_num, N2, block_size, D]`, or BBND `[block_num, block_size, N2, D]` layout. | BFLOAT16, FLOAT16, INT8, HIFLOAT8, FLOAT8_E5M2, FLOAT8_E4M3FN | ND | -| `block_table` | `blockTableOptional` | Optional input | PageAttention logical-block-to-physical-page mapping, shape `[B, maxBlockNumPerSeq]`. Required for BBND and BNBD. | INT32 | ND | -| `scale` | `scaleOptional` | Optional input | INT8 dequantization scale. PageAttention shape: `[block_num, N2, block_size]` or `[block_num, block_size, N2]`; TND shape: `[T2, N2]`, or `[T2]` when N2 is 1. | FLOAT | ND | -| `atten_mask` | `attenMaskOptional` | Optional input | Base causal mask used by `sparse_mode=3`, shape `[2048, 2048]`. A value of 1 excludes a position and 0 includes it. | INT8 | ND | -| `actual_seq_qlen` | `actualSeqQlenOptional` | Input | Non-decreasing query prefix sums, shape `[B+1]`. | INT32 | ND | -| `actual_seq_klen` | `actualSeqKlenOptional` | Input | TND key prefix sums `[B+1]`, or visible key lengths `[B]` for PageAttention. | INT32 | ND | -| `start_loc` | `startLoc` | Input | Logical-block index containing the current query, shape `[B]`. | INT32 | ND | -| `layout_key` | `layoutKeyOptional` | Attribute | Key layout: `"TND"`, `"BBND"`, or `"BNBD"`. The aclnn parameter is `layoutKeyOptional` and defaults to `"BBND"`. | STRING | - | -| `sparse_mode` | `sparseMode` | Attribute | 0: `defaultMask`; 3: `rightDownCausal`. | INT64 | - | -| `init_blocks` | `initBlocks` | Attribute | Number of leading blocks assigned `1e30`. Default: 0. | INT64 | - | -| `local_blocks` | `localBlocks` | Attribute | Size of the local window `[max(0, start_loc+1-local_blocks), start_loc]`, assigned `1e29`. Default: 1. | INT64 | - | -| `score` | `score` | Output | Block scores, shape `[N1, T1, RoundUp(maxBlockNumPerSeq, 16)]`. | FLOAT | ND | - -The defaults above describe the operator schema. The ACLNN C++ calls take -explicit attribute arguments. The vLLM binding in -[msa_index_score_torch_adpt.h](./msa_index_score_torch_adpt.h) supports the -BBND non-INT8 path and uses `init_blocks=0, local_blocks=0` by default, as -registered in [torch_binding.cpp](../../torch_binding.cpp). Keep both at -zero for TP-sharded scoring, where the subsequent TopK stage applies global -block forcing. +| Parameter | Kind | Description | Data Type | Format | +| --------- | ---- | ----------- | --------- | ------ | +| `query` | Input | Query tensor in TND layout, shape `[T1, N1, D]`. | BFLOAT16, FLOAT16, HIFLOAT8, FLOAT8_E5M2, FLOAT8_E4M3FN | ND | +| `key` | Input | Key tensor in TND `[T2, N2, D]`, BNBD `[block_num, N2, block_size, D]`, or BBND `[block_num, block_size, N2, D]` layout. | BFLOAT16, FLOAT16, INT8, HIFLOAT8, FLOAT8_E5M2, FLOAT8_E4M3FN | ND | +| `block_table` | Optional input | PageAttention logical-block-to-physical-page mapping, shape `[B, maxBlockNumPerSeq]`. Required for BBND and BNBD. | INT32 | ND | +| `scale` | Optional input | INT8 dequantization scale. PageAttention shape: `[block_num, N2, block_size]` or `[block_num, block_size, N2]`; TND shape: `[T2, N2]`. | FLOAT | ND | +| `atten_mask` | Optional input | Base causal mask used by `sparse_mode=3`, shape `[2048, 2048]`. A value of 1 excludes a position and 0 includes it. | INT8 | ND | +| `actual_seq_qlen` | Input | Non-decreasing query prefix sums, shape `[B+1]`. | INT32 | ND | +| `actual_seq_klen` | Input | TND key prefix sums `[B+1]`, or visible key lengths `[B]` for PageAttention. | INT32 | ND | +| `start_loc` | Input | Logical-block index containing the current query, shape `[B]`. | INT32 | ND | +| `layout_key` | Attribute | Key layout: `"TND"`, `"BBND"`, or `"BNBD"`. The aclnn parameter is `layoutKeyOptional` and defaults to `"BBND"`. | STRING | - | +| `sparse_mode` | Attribute | 0: `defaultMask`; 3: `rightDownCausal`. | INT64 | - | +| `init_blocks` | Attribute | Number of leading blocks assigned `1e30`. Default: 0. | INT64 | - | +| `local_blocks` | Attribute | Size of the local window `[max(0, start_loc+1-local_blocks), start_loc]`, assigned `1e29`. Default: 1. | INT64 | - | +| `score` | Output | Block scores, shape `[N1, T1, RoundUp(maxBlockNumPerSeq, 16)]`. | FLOAT | ND | ## Constraints @@ -102,129 +92,43 @@ block forcing. - PageAttention BBND/BNBD keys may be non-contiguous on the physical-page axis on A2/A3 and 950PR&950DT Products. All inner axes must remain contiguous. TND keys must be contiguous, and `scale` remains tightly packed by logical page. -- A2/A3 and 950PR&950DT Products size the MIX launch from the estimated M-task count. - When short-M decode cannot fill the 950PR&950DT Products AICs and spans multiple visible - KV S-tiles, `kvChunks` partitions the visible KV range across additional MIX - tasks. Short-KV and wide-table inputs keep a single KV chunk. - The operator returns block scores only and does not perform TopK. -## ACLNN Interface - -### Function Prototypes - -This operator uses a two-stage interface. Call -`aclnnMsaIndexScoreGetWorkspaceSize` to validate the inputs and obtain an -executor and the required workspace size. Then call `aclnnMsaIndexScore` on the -same stream context to execute the computation. +## Build and Run -```cpp -aclnnStatus aclnnMsaIndexScoreGetWorkspaceSize( - const aclTensor *query, - const aclTensor *key, - const aclTensor *blockTableOptional, - const aclTensor *scaleOptional, - const aclTensor *attenMaskOptional, - const aclTensor *actualSeqQlenOptional, - const aclTensor *actualSeqKlenOptional, - const aclTensor *startLoc, - char *layoutKeyOptional, - int64_t sparseMode, - int64_t initBlocks, - int64_t localBlocks, - const aclTensor *score, - uint64_t *workspaceSize, - aclOpExecutor **executor); +For Atlas A2/A3: -aclnnStatus aclnnMsaIndexScore( - void *workspace, - uint64_t workspaceSize, - aclOpExecutor *executor, - aclrtStream stream); +```bash +bash build.sh --pkg --soc=ascend910b --ops=msa_index_score -j32 +bash ./build_out/cann-ops-transformer-custom_linux-x86_64.run \ + --quiet --install-path=/tmp/msa_opp +export ASCEND_CUSTOM_OPP_PATH=/tmp/msa_opp/vendors/custom_transformer +bash build.sh --run_example msa_index_score eager cust \ + --vendor_name=custom --soc=ascend910b ``` -### Workspace Query Outputs - -The input tensors and attributes of `aclnnMsaIndexScoreGetWorkspaceSize` -are described in [Parameters](#parameters). The caller supplies `score` -with the documented output shape and receives: - -| Parameter | Type | Description | -| --------- | ---- | ----------- | -| `workspaceSize` | `uint64_t*` | Required device workspace size in bytes. | -| `executor` | `aclOpExecutor**` | Executor passed to the second-stage call. | - -### Workspace Query Return Values - -| Return Code | Error Code | Description | -| ----------- | ---------- | ----------- | -| `ACLNN_SUCCESS` | 0 | Validation succeeded. | -| `ACLNN_ERR_PARAM_NULLPTR` | 161001 | A required input or output is null. | -| `ACLNN_ERR_PARAM_INVALID` | 161002 | A dtype, format, dimension, stride, or value violates a constraint. | - -### Execution: aclnnMsaIndexScore - -#### Parameters +For 950PR&950DT Products: -| Parameter | Kind | Description | -| --------- | ---- | ----------- | -| `workspace` | Input | Device workspace address. | -| `workspaceSize` | Input | Workspace size returned by the first-stage interface. | -| `executor` | Input | Executor returned by the first-stage interface. | -| `stream` | Input | ACL stream used to execute the operator. | - -#### Return Values - -| Return Code | Error Code | Description | -| ----------- | ---------- | ----------- | -| `ACLNN_SUCCESS` | 0 | Execution succeeded. | -| `ACLNN_ERR_PARAM_INVALID` | 161002 | A parameter is invalid. | - -### Invocation Example - -The following excerpt shows the two-stage invocation for a non-quantized BBND -PageAttention input. It assumes the input and output tensors have already -been allocated with the shapes described above. - -```cpp -static char layoutKey[] = "BBND"; -int64_t sparseMode = 3; -int64_t initBlocks = 0; -int64_t localBlocks = 1; -void *workspace = nullptr; -uint64_t workspaceSize = 0; -aclOpExecutor *executor = nullptr; - -aclnnStatus ret = aclnnMsaIndexScoreGetWorkspaceSize( - queryTensor, - keyTensor, - blockTableTensor, - nullptr, // scaleOptional: null for non-quantized input - attenMaskTensor, - actualSeqQlenTensor, - actualSeqKlenTensor, - startLocTensor, - layoutKey, - sparseMode, - initBlocks, - localBlocks, - scoreTensor, - &workspaceSize, - &executor); - -if (ret == ACLNN_SUCCESS && workspaceSize > 0) { - ret = aclrtMalloc(&workspace, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); -} -if (ret == ACLNN_SUCCESS) { - ret = aclnnMsaIndexScore(workspace, workspaceSize, executor, stream); -} +```bash +bash build.sh --pkg --soc=ascend950 --ops=msa_index_score -j32 +bash ./build_out/cann-ops-transformer-custom_linux-x86_64.run \ + --quiet --install-path=/tmp/msa_opp +source /tmp/msa_opp/vendors/custom_transformer/bin/set_env.bash +export ASCEND_CUSTOM_OPP_PATH=/tmp/msa_opp/vendors/custom_transformer +bash build.sh --run_example msa_index_score eager cust \ + --vendor_name=custom --soc=ascend950 ``` -For BNBD, set `layoutKey="BNBD"` and use a -`[block_num, N2, block_size, D]` key. For TND, set `layoutKey="TND"`, omit -`blockTableOptional`, and provide `[B+1]` key-length prefix sums. +The expected result is 40/40 cases on 950PR&950DT Products. A2/A3 skip the four FP8 +cases and run 36 cases. FLOAT16/BFLOAT16/INT8 use a tolerance of `1e-3`; FP8 +uses `2e-2`. -Keep attribute storage alive through execution and graph capture. Synchronize -the stream before reading the output or releasing the workspace and tensors. +## References + +- [aclnn interface documentation](./docs/aclnnMsaIndexScore.md) +- [End-to-end example](./examples/test_aclnn_msa_index_score.cpp) +- [Test guide](./tests/README.md) +- [torch extension documentation](../../torch_extension/cann_ops_transformer/docs/zh/msa_index_score.md) ## Implementation Notes @@ -239,65 +143,3 @@ the stream before reading the output or releasing the workspace and tensors. `op_kernel/catlass`, derived from v1.3.1-notla. A2/A3 continue to use the repository Catlass submodule. The `msa_` prefix isolates only the A5-specific snapshot because its interfaces and implementation differ. -- For 950PR&950DT Products short-M/long-KV decode, host tiling derives `kvChunks` from - visible KV S-tiles. Both arch22 and arch35 schedulers split the S range and - only the final chunk writes the aligned tail fill. - -## Testing - -### CPU Reference - -The NumPy CPU reference is in -[msa_index_score_golden.py](./tests/golden/msa_index_score_golden.py): - -```python -inputs = MsaIndexScoreGoldenInputs( - query=query, - key=key, - block_table=block_table, - actual_seq_qlen=actual_seq_qlen, - actual_seq_klen=actual_seq_klen, - start_loc=start_loc, - sparse_mode=3, - scale=None, -) -score = msa_index_score_golden(inputs) -``` - -### vLLM Single-Operator Precision Tests - -From the vllm-ascend repository root, with the custom operator installed, run: - -```bash -pytest -sv tests/e2e/nightly/single_node/ops/singlecard_ops/test_msa_index_score.py -``` - -The [nightly test](../../../tests/e2e/nightly/single_node/ops/singlecard_ops/test_msa_index_score.py) -uses a self-contained CPU FP32 reference. Its eight scenarios cover prefill, -decode, non-contiguous page axes, long KV with a wide block table, wide-table -padding, empty query/KV requests, forced blocks, and dense TP chunks. Each is -parameterized over FLOAT16, BFLOAT16, and FLOAT8_E4M3FN, giving 24 cases. FP8 -cases skip when `HardwareCapability.FP8_ATTENTION` is unavailable. - -The existing [PR operator test](../../../tests/e2e/pull_request/one_card/test_msa_index_score.py) -also compares operator outputs against the NumPy reference. The -[MiniMax unit tests](../../../tests/ut/models/minimax_m3/test_msa_m3.py) -cover integration and dispatch behavior. These checks do not need model -weights and do not measure GPQA answer accuracy. - -### Acceptance Criteria - -- Masked and padded score positions must match the reference fill positions. -- Blocks forced by `local_mask` must be at least `1e28` on both sides. -- The NumPy reference's optional `compare` helper allows an error ratio no - greater than `1e-3`, with an error threshold of - `atol + rtol * max(abs(golden), 1)` for ordinary score values. -- The nightly test checks fill/forced-block masks exactly and uses - `torch.testing.assert_close` for every remaining score, with `atol=rtol=1e-3` - for FLOAT16/BFLOAT16 and `2e-2` for FP8. It uses inputs exactly representable - in the tested dtypes to isolate computation and scheduling differences. - -## References - -- [Upstream ops-transformer implementation](https://gitcode.com/cann/ops-transformer/tree/master/attention/msa_index_score) -- [Upstream Python interface](https://gitcode.com/cann/ops-transformer/blob/master/attention/msa_index_score/docs/torchapi_msa_index_score.md) diff --git a/csrc/attention/msa_index_score/docs/aclnnMsaIndexScore.md b/csrc/attention/msa_index_score/docs/aclnnMsaIndexScore.md new file mode 100644 index 000000000000..1ba1f0b893b8 --- /dev/null +++ b/csrc/attention/msa_index_score/docs/aclnnMsaIndexScore.md @@ -0,0 +1,205 @@ +# aclnnMsaIndexScore + +[View upstream source](https://gitcode.com/cann/ops-transformer/tree/master/attention/msa_index_score) + +## Product Support + +| Product | Supported | +| --------------------------------------------------------------------- | :-------: | +| Atlas A2 Products | √ | +| Atlas A3 Products | √ | +| 950PR&950DT Products | √ | + +## Function Description + +`aclnnMsaIndexScore` computes per-block importance scores for the Index Branch +of MiniMax Sparse Attention (MSA). For every query token and sparse KV block, +the operator performs matrix multiplication and max pooling over the causally +visible tokens in that block. It optionally dequantizes an INT8 key. The result +is consumed by a subsequent TopK operation; TopK itself is not part of this +operator. + +The complete formula is: + +$$ +score = Maxpool[(scale \cdot)Q_{idx}@K_{idx}^{T} + atten\_mask] + local\_mask +$$ + +`local_mask` is generated from `startLoc`, `initBlocks`, and `localBlocks`. +Logical blocks in `[0, initBlocks)` receive `1e30`. Blocks in +`[max(0, startLoc+1-localBlocks), startLoc]` receive `1e29`, overriding an +init-block score at the same position. Setting both block-count attributes to +zero disables `local_mask`. + +Notation used below: + +- B is the batch size. +- S1 and S2 are the query and key sequence lengths. +- T1 and T2 are the sums of query and key lengths across the batch. +- N1 and N2 are the query-head and key-head counts. +- D is the head dimension. +- `block_num` is the number of physical PageAttention pages. +- `maxBlockNumPerSeq` is the width of `blockTableOptional`. + +## Function Prototypes + +This operator uses a two-stage interface. Call +`aclnnMsaIndexScoreGetWorkspaceSize` to validate the inputs and obtain an +executor and the required workspace size. Then call `aclnnMsaIndexScore` on the +same stream context to execute the computation. + +```cpp +aclnnStatus aclnnMsaIndexScoreGetWorkspaceSize( + const aclTensor *query, + const aclTensor *key, + const aclTensor *blockTableOptional, + const aclTensor *scaleOptional, + const aclTensor *attenMaskOptional, + const aclTensor *actualSeqQlenOptional, + const aclTensor *actualSeqKlenOptional, + const aclTensor *startLoc, + char *layoutKeyOptional, + int64_t sparseMode, + int64_t initBlocks, + int64_t localBlocks, + const aclTensor *score, + uint64_t *workspaceSize, + aclOpExecutor **executor); + +aclnnStatus aclnnMsaIndexScore( + void *workspace, + uint64_t workspaceSize, + aclOpExecutor *executor, + aclrtStream stream); +``` + +## aclnnMsaIndexScoreGetWorkspaceSize + +### Parameters + +| Parameter | Kind | Description | Data Type | Format | Shape | +| --------- | ---- | ----------- | --------- | ------ | ----- | +| `query` | Input | Query in TND layout. | BFLOAT16, FLOAT16, HIFLOAT8, FLOAT8_E5M2, FLOAT8_E4M3FN | ND | `[T1, N1, D]` | +| `key` | Input | Key in TND, BNBD, or BBND layout. | BFLOAT16, FLOAT16, INT8, HIFLOAT8, FLOAT8_E5M2, FLOAT8_E4M3FN | ND | `[T2, N2, D]`, `[block_num, N2, block_size, D]`, or `[block_num, block_size, N2, D]` | +| `blockTableOptional` | Optional input | Logical-block-to-physical-page mapping. Required for PageAttention and omitted for TND. | INT32 | ND | `[B, maxBlockNumPerSeq]` | +| `scaleOptional` | Optional input | INT8 dequantization scale. Pass `nullptr` for non-quantized and native FP8 inputs. | FLOAT | ND | PageAttention: `[block_num, N2, block_size]` or `[block_num, block_size, N2]`; TND: `[T2, N2]` or `[T2]` when N2 is 1 | +| `attenMaskOptional` | Optional input | Base causal mask for `sparseMode=3`. A value of 1 excludes a position and 0 includes it. | INT8 | ND | `[2048, 2048]` | +| `actualSeqQlenOptional` | Input | Non-decreasing query-length prefix sums. | INT32 | ND | `[B+1]` | +| `actualSeqKlenOptional` | Input | TND key-length prefix sums, or visible key lengths for PageAttention. | INT32 | ND | TND: `[B+1]`; PageAttention: `[B]` | +| `startLoc` | Input | Logical-block index containing the current query. | INT32 | ND | `[B]` | +| `layoutKeyOptional` | Attribute | `"TND"`, `"BBND"`, or `"BNBD"`. Defaults to `"BBND"` when omitted or empty. | CHAR* | - | - | +| `sparseMode` | Attribute | 0 selects `defaultMask`; 3 selects `rightDownCausal`. | INT64 | - | - | +| `initBlocks` | Attribute | Number of leading blocks assigned `1e30`. Default: 0. | INT64 | - | - | +| `localBlocks` | Attribute | Local-window length assigned `1e29`. Default: 1. | INT64 | - | - | +| `score` | Output | Per-block importance scores. | FLOAT | ND | `[N1, T1, RoundUp(maxBlockNumPerSeq, 16)]` | +| `workspaceSize` | Output | Required workspace size in bytes. | uint64_t | - | - | +| `executor` | Output | Operator executor returned by the first-stage interface. | aclOpExecutor** | - | - | + +### Return Values + +| Return Code | Error Code | Description | +| ----------- | ---------- | ----------- | +| `ACLNN_SUCCESS` | 0 | Validation succeeded. | +| `ACLNN_ERR_PARAM_NULLPTR` | 161001 | A required input or output is null. | +| `ACLNN_ERR_PARAM_INVALID` | 161002 | A dtype, format, dimension, stride, or value violates a constraint. | + +## aclnnMsaIndexScore + +### Parameters + +| Parameter | Kind | Description | +| --------- | ---- | ----------- | +| `workspace` | Input | Device workspace address. | +| `workspaceSize` | Input | Workspace size returned by the first-stage interface. | +| `executor` | Input | Executor returned by the first-stage interface. | +| `stream` | Input | ACL stream used to execute the operator. | + +### Return Values + +| Return Code | Error Code | Description | +| ----------- | ---------- | ----------- | +| `ACLNN_SUCCESS` | 0 | Execution succeeded. | +| `ACLNN_ERR_PARAM_INVALID` | 161002 | A parameter is invalid. | + +## Constraints + +- Only `block_size=128` is supported. +- `layoutKeyOptional` must match the key shape. BBND is + `[block_num, block_size, N2, D]`, BNBD is + `[block_num, N2, block_size, D]`, and TND is `[T2, N2, D]`. +- PageAttention requires `blockTableOptional`. TND requires a null + `blockTableOptional` and `[B+1]` prefix sums in `actualSeqKlenOptional`. +- For non-quantized input, query and key must use the same dtype and + `scaleOptional` must be null. A2/A3 support FLOAT16 and BFLOAT16. 950PR&950DT Products + additionally supports HIFLOAT8, FLOAT8_E5M2, and FLOAT8_E4M3FN. +- The quantized path supports a FLOAT16 query, an INT8 key, and a required + FLOAT `scaleOptional`. Native FP8 is a non-quantized 950PR&950DT Products path: query + and key must use the same FP8 dtype and `scaleOptional` must be null. +- `sparseMode=0` requires a null `attenMaskOptional`. `sparseMode=3` requires + an INT8 `[2048, 2048]` mask. +- `initBlocks` and `localBlocks` must be non-negative and no greater than the + logical block width. Setting both to zero disables `local_mask`. +- `q_len` and `kv_len` may be zero, including for the entire batch. The kernel + skips empty-query computation and fills scores for empty KV requests. An + all-empty query batch launches with one block. +- A PageAttention block table may be wider than the actual logical KV length. + Score width is `RoundUp(blockTableOptional.shape[1], 16)`. On 950PR&950DT Products, + widths above 256 are flushed in 256-column windows. +- PageAttention BBND/BNBD keys may be non-contiguous only on the physical-page + axis. The operator reads the first-axis element stride from tensor metadata; + all inner axes must be contiguous. TND keys must be contiguous, and + `scaleOptional` remains tightly packed by logical page. + +## Invocation Example + +The following excerpt shows the two-stage invocation for a non-quantized BBND +PageAttention input. See +[test_aclnn_msa_index_score.cpp](../examples/test_aclnn_msa_index_score.cpp) +for complete BBND, BNBD, TND, INT8, FP8, empty-sequence, strided-page, and wide +block-table accuracy cases. + +```cpp +char layoutKey[] = "BBND"; +int64_t sparseMode = 3; +int64_t initBlocks = 0; +int64_t localBlocks = 1; +uint64_t workspaceSize = 0; +aclOpExecutor *executor = nullptr; + +aclnnStatus ret = aclnnMsaIndexScoreGetWorkspaceSize( + queryTensor, + keyTensor, + blockTableTensor, + nullptr, // scaleOptional: null for non-quantized input + attenMaskTensor, + actualSeqQlenTensor, + actualSeqKlenTensor, + startLocTensor, + layoutKey, + sparseMode, + initBlocks, + localBlocks, + scoreTensor, + &workspaceSize, + &executor); + +if (ret == ACLNN_SUCCESS && workspaceSize > 0) { + ret = aclrtMalloc(&workspace, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); +} +if (ret == ACLNN_SUCCESS) { + ret = aclnnMsaIndexScore(workspace, workspaceSize, executor, stream); +} +``` + +For BNBD, set `layoutKey="BNBD"` and use a +`[block_num, N2, block_size, D]` key. For TND, set `layoutKey="TND"`, omit +`blockTableOptional`, and provide `[B+1]` key-length prefix sums. + +## Validation Matrix + +The standalone example runs 40 cases on 950PR&950DT Products: 36 +FLOAT16/BFLOAT16/INT8 cases plus four FP8 cases. A2/A3 skip FP8 and run 36 +cases. The matrix includes all supported layouts, empty sequences, +non-contiguous PageAttention page axes, and a block-table width of 257. + +FLOAT16/BFLOAT16/INT8 use `atol=rtol=1e-3`; FP8 uses `atol=rtol=2e-2`. diff --git a/csrc/attention/msa_index_score/examples/test_aclnn_msa_index_score.cpp b/csrc/attention/msa_index_score/examples/test_aclnn_msa_index_score.cpp new file mode 100644 index 000000000000..6ac9486f64c6 --- /dev/null +++ b/csrc/attention/msa_index_score/examples/test_aclnn_msa_index_score.cpp @@ -0,0 +1,939 @@ +/** + * 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. + */ + +/*! + * \file test_aclnn_msa_index_score.cpp + * \brief aclnnMsaIndexScore 调用示例,内置 CPU golden 做端到端精度自验证。 + * + * 用例矩阵覆盖:Prefill 多 M-tile、prefix 非 128 对齐的边界 block、varlen 多 batch、 + * Decode(q_len=1)、投机解码(q_len>1)、长序列多 S-tile 轮转、block_table 乱序、 + * 无效尾填充、q_len/kv_len=0 的 mixed-batch pad、bf16 / fp16 双 dtype、 + * int8 key 前融合反量化、PA BNBD、TND packed key、A2/A3 与 950 PA key dim0 非连续。 + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "acl/acl.h" +#include "aclnnop/aclnn_msa_index_score.h" +#include "securec.h" + +namespace { + +constexpr int64_t BLOCK_SIZE = 128; +constexpr int64_t NUM_KV_HEADS = 1; // MSA index cache 为单头共享,P0 仅支持 1 +constexpr int64_t SCORE_STRIDE_ALIGN = 16; +constexpr float kNegInf = -3.4028234663852886e+38F; +constexpr float kAtol = 1e-3F; +constexpr float kRtol = 1e-3F; +constexpr float kAtolFp8 = 2e-2F; +constexpr float kRtolFp8 = 2e-2F; +constexpr int64_t kSparseModeRightDown = 3; +constexpr uint32_t kBf16Shift = 16; +constexpr uint32_t kPrngMulA = 1664525U; +constexpr uint32_t kPrngAddA = 1013904223U; +constexpr uint32_t kPrngXorShiftA = 16; +constexpr uint32_t kPrngMulB = 2246822519U; +constexpr uint32_t kPrngXorShiftB = 13; +constexpr uint32_t kPrngMod = 20000U; +constexpr float kPrngDiv = 10000.0F; +constexpr float kQueryAmp = 0.5F; +constexpr float kFillCmpScale = 0.5F; +constexpr float kBoostThr = 1.0e28F; +constexpr float kInt8Scale = 64.0F; +constexpr float kDeqScaleBase = 0.01F; +constexpr float kDeqScaleSpan = 0.02F; +constexpr float kKeyStridePoisonFp = 99.0F; +constexpr int8_t kKeyStridePoisonI8 = 90; +constexpr uint32_t kKeySeedOff = 999983U; +constexpr uint32_t kDeqSeedOff = 424242U; +constexpr int64_t kBlockTableMulB = 7; +constexpr int64_t kBlockTableMulK = 3; +constexpr int64_t kTraceExtraTok = 2; +constexpr int64_t kPrevDimOffset = 2; // 从末维向前推 stride 时的偏移 +constexpr size_t kLayoutKeyBufSize = 8; +constexpr int kFp8E4M3ExpBits = 4; +constexpr int kFp8E4M3MantissaBits = 3; +constexpr int kFp8E4M3ExpBias = (1 << (kFp8E4M3ExpBits - 1)) - 1; +constexpr float kFp8E4M3MantissaScale = static_cast(1 << kFp8E4M3MantissaBits); +constexpr int kFp8E5M2ExpBits = 5; +constexpr int kFp8E5M2MantissaBits = 2; +constexpr int kFp8E5M2ExpBias = (1 << (kFp8E5M2ExpBits - 1)) - 1; +constexpr float kFp8E5M2MantissaScale = static_cast(1 << kFp8E5M2MantissaBits); +constexpr int kFp8MinNormalExp = 1; // IEEE 次正规数 ldexp 指数为 1-bias +constexpr int kFp8CodeCount = static_cast(std::numeric_limits::max()) + 1; + +enum class KeyLayout { + BBND = 0, + BNBD = 1, + TND = 2, +}; + +std::string LayoutKeyName(KeyLayout layout) +{ + switch (layout) { + case KeyLayout::TND: + return "TND"; + case KeyLayout::BNBD: + return "BNBD"; + case KeyLayout::BBND: + default: + return "BBND"; + } +} + +struct TestCase { + std::string name; + int64_t numQHeads; + int64_t headDim; + int64_t numPages; + std::vector qLen; // 每个请求的 query 长度 + std::vector kvLen; // 每个请求可见的 kv 长度 + std::vector startLoc; // 当前 query 所在逻辑 block 索引(local_mask) + bool useBf16; + bool useInt8Key; // true: key=int8,scale=[NP,N_kv,P] 或 TND [T2,N2] + int64_t sparseMode = kSparseModeRightDown; + KeyLayout keyLayout = KeyLayout::BBND; + // 0=无;1=float8_e4m3fn;2=float8_e5m2。仅 950,query/key 同型。 + int fp8Kind = 0; + // PA key dim0 间隔:1=紧凑;>1 时 storage 为 key|gap|key|...(A2/A3 与 950)。 + int64_t keyDim0Gap = 1; + // PA block_table 第二维;0 表示按 kv_len 取 maxBlocks。>256 覆盖 950 C2UB 滑窗 flush。 + int64_t tableWidth = 0; +}; + +constexpr float kLocalScoreInit = 1.0e30F; +constexpr float kLocalScoreLocal = 1.0e29F; +constexpr int64_t kInitBlocks = 0; +constexpr int64_t kLocalBlocks = 1; +constexpr int64_t kAttenMaskSize = 2048; + +int64_t CeilDivI64(int64_t a, int64_t b) +{ + return (a + b - 1) / b; +} + +int64_t RoundUpI64(int64_t a, int64_t b) +{ + return CeilDivI64(a, b) * b; +} + +// 简单可复现的伪随机源,取值落在 [-1, 1)。 +float PseudoRandom(uint32_t seed) +{ + seed = seed * kPrngMulA + kPrngAddA; + seed ^= seed >> kPrngXorShiftA; + seed = seed * kPrngMulB; + seed ^= seed >> kPrngXorShiftB; + return static_cast(seed % kPrngMod) / kPrngDiv - 1.0F; +} + +// bf16 <-> fp32:截断低 16 位尾数(round-to-nearest-even 对本用例不必要)。 +uint16_t FloatToBf16(float v) +{ + uint32_t bits = 0; + const errno_t ret = memcpy_s(&bits, sizeof(bits), &v, sizeof(v)); + if (ret != 0) { + return 0; + } + return static_cast(bits >> kBf16Shift); +} + +float Bf16ToFloat(uint16_t v) +{ + const uint32_t bits = static_cast(v) << kBf16Shift; + float out = 0.0F; + const errno_t ret = memcpy_s(&out, sizeof(out), &bits, sizeof(bits)); + if (ret != 0) { + return 0.0F; + } + return out; +} + +// OCP E4M3FN:exp=15 且 mant=7 为 NaN,其余有限;bias=7。 +float Fp8E4M3fnToFloat(uint8_t x) +{ + const uint32_t s = (static_cast(x) >> 7) & 1U; + const uint32_t e = (static_cast(x) >> 3) & 0xFU; + const uint32_t m = static_cast(x) & 7U; + if (e == 0xFU && m == 7U) { + return std::numeric_limits::quiet_NaN(); + } + const float sign = (s != 0U) ? -1.0F : 1.0F; + if (e == 0U) { + return sign * std::ldexp(static_cast(m) / kFp8E4M3MantissaScale, kFp8MinNormalExp - kFp8E4M3ExpBias); + } + return sign * + std::ldexp(1.0F + static_cast(m) / kFp8E4M3MantissaScale, static_cast(e) - kFp8E4M3ExpBias); +} + +// IEEE-like E5M2:bias=15,exp=31 为 Inf/NaN。 +float Fp8E5M2ToFloat(uint8_t x) +{ + const uint32_t s = (static_cast(x) >> 7) & 1U; + const uint32_t e = (static_cast(x) >> 2) & 0x1FU; + const uint32_t m = static_cast(x) & 3U; + const float sign = (s != 0U) ? -1.0F : 1.0F; + if (e == 0x1FU) { + return (m == 0U) ? (sign * std::numeric_limits::infinity()) : std::numeric_limits::quiet_NaN(); + } + if (e == 0U) { + return sign * std::ldexp(static_cast(m) / kFp8E5M2MantissaScale, kFp8MinNormalExp - kFp8E5M2ExpBias); + } + return sign * + std::ldexp(1.0F + static_cast(m) / kFp8E5M2MantissaScale, static_cast(e) - kFp8E5M2ExpBias); +} + +uint8_t FloatToFp8Nearest(float v, bool e5m2) +{ + uint8_t best = 0; + float bestDiff = std::numeric_limits::infinity(); + for (int i = 0; i < kFp8CodeCount; ++i) { + const uint8_t b = static_cast(i); + const float d = e5m2 ? Fp8E5M2ToFloat(b) : Fp8E4M3fnToFloat(b); + if (!std::isfinite(d)) { + continue; + } + const float diff = std::fabs(d - v); + if (diff < bestDiff) { + bestDiff = diff; + best = b; + } + } + return best; +} + +template +std::vector ScatterPaKeyDim0(const std::vector &packed, int64_t numPages, int64_t pageElems, int64_t gap, + T poison) +{ + std::vector wide(static_cast(numPages * gap * pageElems), poison); + for (int64_t p = 0; p < numPages; ++p) { + const auto src = packed.begin() + static_cast(p * pageElems); + const auto dst = wide.begin() + static_cast(p * gap * pageElems); + std::copy(src, src + static_cast(pageElems), dst); + } + return wide; +} + +std::vector ContiguousStrides(const std::vector &shape) +{ + std::vector strides(shape.size(), 1); + for (int64_t i = static_cast(shape.size()) - kPrevDimOffset; i >= 0; i--) { + strides[static_cast(i)] = shape[static_cast(i + 1)] * strides[static_cast(i + 1)]; + } + return strides; +} + +class DeviceBuffer { +public: + ~DeviceBuffer() + { + for (auto *t : tensors_) { + if (t != nullptr) { + (void)aclDestroyTensor(t); + } + } + for (auto *p : addrs_) { + if (p != nullptr) { + (void)aclrtFree(p); + } + } + } + + template + aclTensor *Create(const std::vector &host, const std::vector &shape, aclDataType dtype, + void **addrOut = nullptr, const std::vector *storageShape = nullptr, + const std::vector *stridesIn = nullptr) + { + void *devAddr = nullptr; + const size_t bytes = host.size() * sizeof(T); + // 空张量(T1=0 / T2=0)仍需要合法 device 指针;aclrtMalloc(0) 会失败。 + const size_t allocBytes = (bytes == 0) ? static_cast(32) : bytes; + if (aclrtMalloc(&devAddr, allocBytes, ACL_MEM_MALLOC_HUGE_FIRST) != ACL_SUCCESS) { + return nullptr; + } + addrs_.push_back(devAddr); + if (bytes > 0 && aclrtMemcpy(devAddr, bytes, host.data(), bytes, ACL_MEMCPY_HOST_TO_DEVICE) != ACL_SUCCESS) { + return nullptr; + } + const std::vector &stShape = (storageShape != nullptr) ? *storageShape : shape; + std::vector strides; + if (stridesIn != nullptr) { + strides = *stridesIn; + } else { + strides.assign(shape.size(), 1); + for (int64_t i = static_cast(shape.size()) - kPrevDimOffset; i >= 0; i--) { + strides[static_cast(i)] = + shape[static_cast(i + 1)] * strides[static_cast(i + 1)]; + } + } + aclTensor *t = aclCreateTensor(shape.data(), shape.size(), dtype, strides.data(), 0, ACL_FORMAT_ND, + stShape.data(), stShape.size(), devAddr); + tensors_.push_back(t); + if (addrOut != nullptr) { + *addrOut = devAddr; + } + return t; + } + +private: + std::vector addrs_; + std::vector tensors_; +}; + +/// CPU 参考:score = Maxpool[(scale·)Q@Kᵀ + atten_mask] + local_mask +void ComputeGolden(const TestCase &tc, const std::vector &actualSeqQlen, + const std::vector &actualSeqKlen, const std::vector &blockTable, int64_t maxBlocks, + int64_t scoreStride, int64_t totalQ, const std::vector &queryF, + const std::vector &keyF, const std::vector &deqScale, std::vector &golden) +{ + const int64_t batch = static_cast(tc.qLen.size()); + golden.assign(static_cast(tc.numQHeads * totalQ * scoreStride), kNegInf); + + auto keyAt = [&](int64_t pageOrTok, int64_t n, int64_t d) -> float { + if (tc.keyLayout == KeyLayout::TND) { + return keyF[static_cast(((pageOrTok + n) * NUM_KV_HEADS) * tc.headDim + d)]; + } + // BBND [NP,P,1,D] 与 BNBD [NP,1,P,D] 在 N2=1 时同址 + return keyF[static_cast(((pageOrTok * BLOCK_SIZE) + n) * tc.headDim + d)]; + }; + auto scaleAt = [&](int64_t pageOrTok, int64_t n) -> float { + if (!tc.useInt8Key) { + return 1.0F; + } + if (tc.keyLayout == KeyLayout::TND) { + return deqScale[static_cast(pageOrTok + n)]; + } + return deqScale[static_cast(pageOrTok * BLOCK_SIZE + n)]; + }; + + for (int64_t b = 0; b < batch; ++b) { + const int32_t qBegin = actualSeqQlen[b]; + const int32_t qEnd = actualSeqQlen[b + 1]; + const int32_t qLen = qEnd - qBegin; + const int32_t kvLen = tc.kvLen[b]; + const int64_t numBlocks = CeilDivI64(kvLen, BLOCK_SIZE); + const int32_t qBlock = tc.startLoc[b]; + const int32_t localStart = + (qBlock + 1 > static_cast(kLocalBlocks)) ? (qBlock + 1 - static_cast(kLocalBlocks)) : 0; + const int32_t cuK = (tc.keyLayout == KeyLayout::TND) ? actualSeqKlen[b] : 0; + + for (int32_t t = qBegin; t < qEnd; ++t) { + const int32_t tOff = t - qBegin; + int32_t visibleKeyEnd = kvLen; + if (tc.sparseMode == kSparseModeRightDown) { + visibleKeyEnd = kvLen - qLen + tOff + 1; + if (visibleKeyEnd < 0) { + visibleKeyEnd = 0; + } + if (visibleKeyEnd > kvLen) { + visibleKeyEnd = kvLen; + } + } + for (int64_t h = 0; h < tc.numQHeads; ++h) { + for (int64_t blk = 0; blk < numBlocks; ++blk) { + const int32_t pageOrTok = (tc.keyLayout == KeyLayout::TND) ? + (cuK + static_cast(blk * BLOCK_SIZE)) : + blockTable[b * maxBlocks + blk]; + float best = 0.0F; + bool any = false; + for (int64_t n = 0; n < BLOCK_SIZE; ++n) { + if (blk * BLOCK_SIZE + n >= visibleKeyEnd) { + break; + } + float acc = 0.0F; + const float s = scaleAt(pageOrTok, n); + for (int64_t d = 0; d < tc.headDim; ++d) { + acc += queryF[((t * tc.numQHeads) + h) * tc.headDim + d] * (keyAt(pageOrTok, n, d) * s); + } + if (!any || acc > best) { + best = acc; + any = true; + } + } + if (!any) { + continue; + } + golden[(h * totalQ + t) * scoreStride + blk] = best; + } + for (int64_t blk = 0; blk < numBlocks; ++blk) { + float boost = 0.0F; + if (blk < kInitBlocks) { + boost = kLocalScoreInit; + } + if (blk >= localStart && blk <= qBlock) { + boost = kLocalScoreLocal; + } + if (boost != 0.0F) { + golden[(h * totalQ + t) * scoreStride + blk] = boost; + } + } + } + } + } +} + +void PrintVec(const std::string &tag, const float *data, int64_t n) +{ + (void)printf(" %s[", tag.c_str()); + for (int64_t i = 0; i < n; ++i) { + (void)printf("%s%.6g", i == 0 ? "" : ", ", static_cast(data[i])); + } + (void)printf("]\n"); +} + +/// 小尺寸用例的 host 侧逐步拆解:Q/K -> DOT/DEQUANT -> MASK -> MAX。 +void PrintTracePipeline(const TestCase &tc, const std::vector &actualSeqQlen, + const std::vector &blockTable, int64_t maxBlocks, int64_t scoreStride, int64_t totalQ, + const std::vector &queryF, const std::vector &keyF, + const std::vector &deqScale, const std::vector &actual) +{ + (void)printf("\n======== HOST TRACE: %s (int8=%d) ========\n", tc.name.c_str(), tc.useInt8Key ? 1 : 0); + (void)printf("shape: Hq=%ld D=%ld qLen=%d kvLen=%d startLoc=%d blockSize=%ld maxBlocks=%ld scoreStride=%ld\n", + tc.numQHeads, tc.headDim, tc.qLen[0], tc.kvLen[0], tc.startLoc[0], BLOCK_SIZE, maxBlocks, scoreStride); + const int32_t page = blockTable[0]; + for (int64_t t = 0; t < tc.qLen[0]; ++t) { + for (int64_t h = 0; h < tc.numQHeads; ++h) { + const int64_t flat = t * tc.numQHeads + h; + const float *q = &queryF[static_cast(flat * tc.headDim)]; + (void)printf("\n-- row flat=%ld (token=%ld head=%ld) --\n", flat, t, h); + PrintVec("Q", q, tc.headDim); + const int32_t tOff = static_cast(t); + int32_t visibleKeyEnd = tc.kvLen[0]; + if (tc.sparseMode == kSparseModeRightDown) { + visibleKeyEnd = tc.kvLen[0] - tc.qLen[0] + tOff + 1; + if (visibleKeyEnd < 0) { + visibleKeyEnd = 0; + } + if (visibleKeyEnd > tc.kvLen[0]) { + visibleKeyEnd = tc.kvLen[0]; + } + } + float best = kNegInf; + bool any = false; + for (int64_t n = 0; n < visibleKeyEnd + kTraceExtraTok && n < BLOCK_SIZE; ++n) { + const float *k = &keyF[static_cast((page * BLOCK_SIZE + n) * tc.headDim)]; + const float ds = tc.useInt8Key ? deqScale[static_cast(page) * BLOCK_SIZE + n] : 1.0F; + float acc = 0.0F; + for (int64_t d = 0; d < tc.headDim; ++d) { + acc += q[d] * (k[d] * ds); + } + const bool visible = (n < visibleKeyEnd); + (void)printf(" S[%ld]=%.6g deqScale=%.6g %s\n", n, static_cast(acc), static_cast(ds), + visible ? "KEEP" : "MASK"); + if (visible && (!any || acc > best)) { + best = acc; + any = true; + } + } + const float deviceOut = actual[static_cast((h * totalQ + (actualSeqQlen[0] + t)) * scoreStride)]; + (void)printf(" [MAX]=%.6g [OUT]=%.6g\n", static_cast(any ? best : kNegInf), + static_cast(deviceOut)); + } + } + (void)printf("======== END HOST TRACE ========\n\n"); + (void)maxBlocks; +} + +bool Compare(const std::string &name, const std::vector &actual, const std::vector &golden, float atol, + float rtol) +{ + size_t badCount = 0; + size_t infBad = 0; + size_t total = 0; + float maxAbsDiff = 0.0F; + size_t firstBad = 0; + bool hasBad = false; + + for (size_t i = 0; i < golden.size(); ++i) { + const float g = golden[i]; + const float a = actual[i]; + if (g <= kNegInf * kFillCmpScale) { // 填充位:要求实测同样是极小值 + if (!(std::isfinite(a) && a <= kNegInf * kFillCmpScale) && !(std::isinf(a) && a < 0)) { + if (!hasBad) { + firstBad = i; + hasBad = true; + } + ++infBad; + } + continue; + } + if (g >= kBoostThr) { // local_mask 强制高分 + if (!(a >= kBoostThr)) { + if (!hasBad) { + firstBad = i; + hasBad = true; + } + ++badCount; + } + continue; + } + ++total; + if (!std::isfinite(a) || !std::isfinite(g)) { + if (!hasBad) { + firstBad = i; + hasBad = true; + } + ++badCount; + continue; + } + const float diff = std::fabs(a - g); + maxAbsDiff = diff > maxAbsDiff ? diff : maxAbsDiff; + if (diff > atol + rtol * std::fabs(g)) { + if (!hasBad) { + firstBad = i; + hasBad = true; + } + ++badCount; + } + } + + const bool pass = (badCount == 0) && (infBad == 0); + (void)printf(" [%s] valid=%zu mismatch=%zu fill_mismatch=%zu max_abs_diff=%.6g -> %s\n", name.c_str(), total, + badCount, infBad, static_cast(maxAbsDiff), pass ? "PASS" : "FAIL"); + if (!pass) { + (void)printf(" first mismatch at %zu: actual=%g golden=%g\n", firstBad, + static_cast(actual[firstBad]), static_cast(golden[firstBad])); + } + return pass; +} + +bool RunCase(const TestCase &tc, aclrtStream stream) +{ + const int64_t batch = static_cast(tc.qLen.size()); + std::vector actualSeqQlen(batch + 1, 0); + for (int64_t b = 0; b < batch; ++b) { + actualSeqQlen[b + 1] = actualSeqQlen[b] + tc.qLen[b]; + } + const int64_t totalQ = actualSeqQlen[batch]; + + int64_t maxBlocks = 1; + for (int64_t b = 0; b < batch; ++b) { + maxBlocks = std::max(maxBlocks, CeilDivI64(tc.kvLen[b], BLOCK_SIZE)); + } + if (tc.tableWidth > maxBlocks) { + maxBlocks = tc.tableWidth; + } + const int64_t scoreStride = RoundUpI64(maxBlocks, SCORE_STRIDE_ALIGN); + + // block_table 故意打乱,验证 paged 间接寻址。 + std::vector blockTable(static_cast(batch * maxBlocks), 0); + if (tc.numPages > 0) { + for (int64_t b = 0; b < batch; ++b) { + for (int64_t k = 0; k < maxBlocks; ++k) { + blockTable[b * maxBlocks + k] = + static_cast((b * kBlockTableMulB + k * kBlockTableMulK + 1) % tc.numPages); + } + } + } + + // fp32 参考数据;再按 dtype 转成低精度输入,golden 用转换后的值以对齐数值路径。 + const bool isTnd = (tc.keyLayout == KeyLayout::TND); + int64_t totalK = 0; + std::vector actualSeqKlenPrefix(static_cast(batch + 1), 0); + if (isTnd) { + for (int64_t b = 0; b < batch; ++b) { + actualSeqKlenPrefix[b + 1] = actualSeqKlenPrefix[b] + tc.kvLen[b]; + } + totalK = actualSeqKlenPrefix[batch]; + } + const int64_t keyTokens = isTnd ? totalK : (tc.numPages * BLOCK_SIZE); + std::vector queryF(static_cast(totalQ * tc.numQHeads * tc.headDim)); + std::vector keyF(static_cast(keyTokens * tc.headDim)); + std::vector queryBf(queryF.size()); + std::vector keyBf(keyF.size()); + std::vector queryHf(queryF.size()); + std::vector keyHf(keyF.size()); + std::vector queryFp8(queryF.size()); + std::vector keyFp8(keyF.size()); + const bool isFp8 = (tc.fp8Kind != 0); + const bool isE5M2 = (tc.fp8Kind == 2); + + for (size_t i = 0; i < queryF.size(); ++i) { + const float v = PseudoRandom(static_cast(i) + 1U) * kQueryAmp; + if (isFp8) { + queryFp8[i] = FloatToFp8Nearest(v, isE5M2); + queryF[i] = isE5M2 ? Fp8E5M2ToFloat(queryFp8[i]) : Fp8E4M3fnToFloat(queryFp8[i]); + } else if (tc.useBf16) { + queryBf[i] = FloatToBf16(v); + queryF[i] = Bf16ToFloat(queryBf[i]); + } else { + queryHf[i] = aclFloatToFloat16(v); + queryF[i] = aclFloat16ToFloat(queryHf[i]); + } + } + for (size_t i = 0; i < keyF.size(); ++i) { + if (tc.useInt8Key) { + // int8 量化值:落在 [-64, 63],golden 用同一整数值的 fp32。 + const int8_t qv = static_cast( + static_cast(PseudoRandom(static_cast(i) + kKeySeedOff) * kInt8Scale)); + keyF[i] = static_cast(qv); + } else { + const float v = PseudoRandom(static_cast(i) + kKeySeedOff) * kQueryAmp; + if (isFp8) { + keyFp8[i] = FloatToFp8Nearest(v, isE5M2); + keyF[i] = isE5M2 ? Fp8E5M2ToFloat(keyFp8[i]) : Fp8E4M3fnToFloat(keyFp8[i]); + } else if (tc.useBf16) { + keyBf[i] = FloatToBf16(v); + keyF[i] = Bf16ToFloat(keyBf[i]); + } else { + keyHf[i] = aclFloatToFloat16(v); + keyF[i] = aclFloat16ToFloat(keyHf[i]); + } + } + } + // 反量化 scale:PA [NP, N_kv=1, P];TND [T2, N2=1]。 + std::vector deqScale(static_cast(keyTokens), 1.0F); + std::vector keyI8; + if (tc.useInt8Key) { + keyI8.resize(keyF.size()); + for (size_t i = 0; i < keyF.size(); ++i) { + keyI8[i] = static_cast(keyF[i]); + } + for (size_t i = 0; i < deqScale.size(); ++i) { + deqScale[i] = kDeqScaleBase + kDeqScaleSpan * (PseudoRandom(static_cast(i) + kDeqSeedOff) + 1.0F); + } + } + + DeviceBuffer buf; + const std::vector queryShape = {totalQ, tc.numQHeads, tc.headDim}; + std::vector keyShape; + if (isTnd) { + keyShape = {totalK, NUM_KV_HEADS, tc.headDim}; + } else if (tc.keyLayout == KeyLayout::BNBD) { + keyShape = {tc.numPages, NUM_KV_HEADS, BLOCK_SIZE, tc.headDim}; + } else { + keyShape = {tc.numPages, BLOCK_SIZE, NUM_KV_HEADS, tc.headDim}; + } + const std::vector scoreShape = {tc.numQHeads, totalQ, scoreStride}; + + aclTensor *queryT = nullptr; + if (isFp8) { + const aclDataType fp8Dt = isE5M2 ? ACL_FLOAT8_E5M2 : ACL_FLOAT8_E4M3FN; + queryT = buf.Create(queryFp8, queryShape, fp8Dt); + } else { + queryT = tc.useBf16 ? buf.Create(queryBf, queryShape, ACL_BF16) : buf.Create(queryHf, queryShape, ACL_FLOAT16); + } + + const int64_t keyDim0Gap = (isTnd || tc.keyDim0Gap < 1) ? 1 : tc.keyDim0Gap; + std::vector keyStorageShape = keyShape; + std::vector keyStrides = ContiguousStrides(keyShape); + const std::vector *keyStoragePtr = nullptr; + const std::vector *keyStridePtr = nullptr; + int64_t pageElems = 1; + if (keyDim0Gap > 1) { + for (size_t d = 1; d < keyShape.size(); ++d) { + pageElems *= keyShape[d]; + } + keyStorageShape[0] = tc.numPages * keyDim0Gap; + keyStrides[0] = keyDim0Gap * keyStrides[0]; + keyStoragePtr = &keyStorageShape; + keyStridePtr = &keyStrides; + } + + aclTensor *keyT = nullptr; + if (tc.useInt8Key) { + if (keyDim0Gap > 1) { + const auto wide = ScatterPaKeyDim0(keyI8, tc.numPages, pageElems, keyDim0Gap, kKeyStridePoisonI8); + keyT = buf.Create(wide, keyShape, ACL_INT8, nullptr, keyStoragePtr, keyStridePtr); + } else { + keyT = buf.Create(keyI8, keyShape, ACL_INT8); + } + } else if (isFp8) { + const aclDataType fp8Dt = isE5M2 ? ACL_FLOAT8_E5M2 : ACL_FLOAT8_E4M3FN; + if (keyDim0Gap > 1) { + const uint8_t poison = FloatToFp8Nearest(kKeyStridePoisonFp, isE5M2); + const auto wide = ScatterPaKeyDim0(keyFp8, tc.numPages, pageElems, keyDim0Gap, poison); + keyT = buf.Create(wide, keyShape, fp8Dt, nullptr, keyStoragePtr, keyStridePtr); + } else { + keyT = buf.Create(keyFp8, keyShape, fp8Dt); + } + } else if (tc.useBf16) { + if (keyDim0Gap > 1) { + const auto wide = + ScatterPaKeyDim0(keyBf, tc.numPages, pageElems, keyDim0Gap, FloatToBf16(kKeyStridePoisonFp)); + keyT = buf.Create(wide, keyShape, ACL_BF16, nullptr, keyStoragePtr, keyStridePtr); + } else { + keyT = buf.Create(keyBf, keyShape, ACL_BF16); + } + } else { + if (keyDim0Gap > 1) { + const auto wide = + ScatterPaKeyDim0(keyHf, tc.numPages, pageElems, keyDim0Gap, aclFloatToFloat16(kKeyStridePoisonFp)); + keyT = buf.Create(wide, keyShape, ACL_FLOAT16, nullptr, keyStoragePtr, keyStridePtr); + } else { + keyT = buf.Create(keyHf, keyShape, ACL_FLOAT16); + } + } + aclTensor *blockTableT = nullptr; + if (!isTnd) { + blockTableT = buf.Create(blockTable, {batch, maxBlocks}, ACL_INT32); + } + aclTensor *scaleT = nullptr; + if (tc.useInt8Key) { + if (isTnd) { + scaleT = buf.Create(deqScale, {totalK, NUM_KV_HEADS}, ACL_FLOAT); + } else { + scaleT = buf.Create(deqScale, {tc.numPages, NUM_KV_HEADS, BLOCK_SIZE}, ACL_FLOAT); + } + } + aclTensor *actualSeqQlenT = buf.Create(actualSeqQlen, {batch + 1}, ACL_INT32); + aclTensor *actualSeqKlenT = + isTnd ? buf.Create(actualSeqKlenPrefix, {batch + 1}, ACL_INT32) : buf.Create(tc.kvLen, {batch}, ACL_INT32); + aclTensor *startLocT = buf.Create(tc.startLoc, {batch}, ACL_INT32); + std::vector attenMaskHost(static_cast(kAttenMaskSize * kAttenMaskSize), 0); + aclTensor *attenMaskT = nullptr; + if (tc.sparseMode == kSparseModeRightDown) { + attenMaskT = buf.Create(attenMaskHost, {kAttenMaskSize, kAttenMaskSize}, ACL_INT8); + } + + std::vector scoreInit(static_cast(tc.numQHeads * totalQ * scoreStride), 0.0F); + void *scoreDev = nullptr; + aclTensor *scoreT = buf.Create(scoreInit, scoreShape, ACL_FLOAT, &scoreDev); + + if (queryT == nullptr || keyT == nullptr || actualSeqQlenT == nullptr || actualSeqKlenT == nullptr || + startLocT == nullptr || scoreT == nullptr || (!isTnd && blockTableT == nullptr) || + (tc.sparseMode == kSparseModeRightDown && attenMaskT == nullptr)) { + (void)printf(" [%s] create tensor failed -> FAIL\n", tc.name.c_str()); + return false; + } + + const std::string layoutKey = LayoutKeyName(tc.keyLayout); + char layoutKeyBuf[kLayoutKeyBufSize] = {}; + const errno_t cpyRet = memcpy_s(layoutKeyBuf, sizeof(layoutKeyBuf), layoutKey.c_str(), layoutKey.size() + 1); + if (cpyRet != 0) { + (void)printf(" [%s] memcpy_s layout_key failed -> FAIL\n", tc.name.c_str()); + return false; + } + + uint64_t workspaceSize = 0; + aclOpExecutor *executor = nullptr; + void *workspaceAddr = nullptr; + int ret = aclnnMsaIndexScoreGetWorkspaceSize(queryT, keyT, blockTableT, scaleT, attenMaskT, actualSeqQlenT, + actualSeqKlenT, startLocT, layoutKeyBuf, tc.sparseMode, kInitBlocks, + kLocalBlocks, scoreT, &workspaceSize, &executor); + if (ret != ACL_SUCCESS) { + (void)printf(" [%s] GetWorkspaceSize failed, ERROR %d -> FAIL\n", tc.name.c_str(), ret); + return false; + } + if (workspaceSize > 0ULL && aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST) != ACL_SUCCESS) { + (void)printf(" [%s] malloc workspace failed -> FAIL\n", tc.name.c_str()); + return false; + } + ret = aclnnMsaIndexScore(workspaceAddr, workspaceSize, executor, stream); + if (ret != ACL_SUCCESS) { + (void)printf(" [%s] aclnnMsaIndexScore failed, ERROR %d -> FAIL\n", tc.name.c_str(), ret); + (void)aclrtFree(workspaceAddr); + return false; + } + (void)printf(" [%s] kernel launched, synchronizing...\n", tc.name.c_str()); + (void)aclrtSynchronizeStream(stream); + (void)printf(" [%s] synchronized\n", tc.name.c_str()); + + std::vector actual(scoreInit.size(), 0.0F); + (void)aclrtMemcpy(actual.data(), actual.size() * sizeof(float), scoreDev, actual.size() * sizeof(float), + ACL_MEMCPY_DEVICE_TO_HOST); + if (workspaceAddr != nullptr) { + (void)aclrtFree(workspaceAddr); + } + + std::vector golden; + ComputeGolden(tc, actualSeqQlen, isTnd ? actualSeqKlenPrefix : tc.kvLen, blockTable, maxBlocks, scoreStride, totalQ, + queryF, keyF, deqScale, golden); + if (tc.name.find("debug-trace") != std::string::npos) { + PrintTracePipeline(tc, actualSeqQlen, blockTable, maxBlocks, scoreStride, totalQ, queryF, keyF, deqScale, + actual); + } + return Compare(tc.name, actual, golden, isFp8 ? kAtolFp8 : kAtol, isFp8 ? kRtolFp8 : kRtol); +} + +} // namespace + +int main() +{ + setvbuf(stdout, nullptr, _IONBF, 0); +#if defined(__linux__) + (void)fedisableexcept(FE_ALL_EXCEPT); + fenv_t fenvHold; + (void)feholdexcept(&fenvHold); +#endif + const int32_t deviceId = 0; + aclrtStream stream = nullptr; + if (aclInit(nullptr) != ACL_SUCCESS || aclrtSetDevice(deviceId) != ACL_SUCCESS || + aclrtCreateStream(&stream) != ACL_SUCCESS) { + (void)printf("[FAIL] init acl failed\n"); + return -1; + } + + // startLoc 为逻辑 block 索引;因果由 sparseMode=3(rightDownCausal)承担。 + const std::vector cases = { + // name Hq D pages qLen kvLen startLoc(block) bf16 int8 + {"L0-debug-trace", 2, 16, 2, {2}, {5}, {16}, false, false}, + {"L0-int8-dequant-trace", 2, 16, 2, {2}, {5}, {16}, false, true}, + {"L0-prefill-aligned", 8, 128, 8, {32}, {256}, {1}, false, false}, + {"L1-prefill-unaligned", 8, 128, 8, {32, 17}, {300, 130}, {2, 0}, false, false}, + {"L1-prefill-multi-mtile", 8, 128, 16, {64, 48}, {700, 520}, {4, 3}, false, false}, + {"L1-decode-lq1", + 8, + 128, + 16, + {1, 1, 1, 1, 1, 1}, + {900, 512, 128, 129, 1, 4096}, + {7, 3, 0, 1, 0, 31}, + false, + false}, + {"L1-decode-speculative", 8, 128, 16, {4, 2}, {1024, 260}, {7, 2}, false, false}, + {"L1-long-seq-multi-stile", 8, 128, 40, {8}, {4096}, {31}, false, false}, + {"L1-bf16", 8, 128, 8, {32, 17}, {300, 130}, {2, 0}, true, false}, + {"L1-head-dim-64", 8, 64, 8, {32, 17}, {300, 130}, {2, 0}, false, false}, + {"L1-heads-16", 16, 128, 8, {13}, {300}, {2}, false, false}, + {"L1-int8-dequant", 8, 128, 8, {32, 17}, {300, 130}, {2, 0}, false, true}, + {"L2-tiny-kv", 8, 128, 4, {1, 1}, {1, 3}, {0, 0}, false, false}, + {"L1-bnbd", 8, 128, 8, {32, 17}, {300, 130}, {2, 0}, false, false, kSparseModeRightDown, KeyLayout::BNBD}, + {"L1-bnbd-int8", 8, 128, 8, {32, 17}, {300, 130}, {2, 0}, false, true, kSparseModeRightDown, KeyLayout::BNBD}, + {"L1-tnd-unaligned", + 8, + 128, + 0, + {32, 17}, + {300, 130}, + {2, 0}, + false, + false, + kSparseModeRightDown, + KeyLayout::TND}, + {"L1-tnd-int8", 8, 128, 0, {32, 17}, {300, 130}, {2, 0}, false, true, kSparseModeRightDown, KeyLayout::TND}, + {"L0-tnd-tiny", 2, 16, 0, {2}, {5}, {16}, false, false, kSparseModeRightDown, KeyLayout::TND}, + // FP8:C0=32,headDim 用 128 对齐;D=16 半个 C0 数值不可靠。 + {"L0-fp8-e4m3fn", 2, 128, 2, {2}, {5}, {16}, false, false, kSparseModeRightDown, KeyLayout::BBND, 1}, + {"L0-fp8-e5m2", 2, 128, 2, {2}, {5}, {16}, false, false, kSparseModeRightDown, KeyLayout::BBND, 2}, + {"L1-fp8-e4m3fn-prefill", 8, 128, 8, {32}, {256}, {1}, false, false, kSparseModeRightDown, KeyLayout::BBND, 1}, + // mixed-batch 零长度:部分请求 pad。整 batch 全 0 见文末 L0-all-* / L0-tnd-all-*。 + {"L1-pad-q0", 8, 128, 8, {0, 32}, {256, 256}, {1, 1}, false, false}, + {"L1-pad-kv0", 8, 128, 8, {32, 32}, {0, 256}, {0, 1}, false, false}, + {"L1-pad-q0-kv0", 8, 128, 8, {0, 32}, {0, 256}, {0, 1}, false, false}, + {"L1-pad-mid-q0", 8, 128, 8, {16, 0, 16}, {200, 0, 200}, {1, 0, 1}, false, false}, + {"L1-tnd-pad-q0-kv0", 8, 128, 0, {0, 32}, {0, 256}, {0, 1}, false, false, kSparseModeRightDown, KeyLayout::TND}, + // 整 batch 空序列:host 跳过计算(A2/A3 与 950)。 + {"L0-all-q0", 8, 128, 8, {0}, {256}, {1}, false, false}, + {"L0-all-kv0", 8, 128, 8, {32}, {0}, {0}, false, false}, + {"L0-all-q0-kv0", 8, 128, 8, {0}, {0}, {0}, false, false}, + {"L1-all-q0", 8, 128, 8, {0, 0}, {256, 256}, {1, 1}, false, false}, + {"L0-tnd-all-q0", 8, 128, 0, {0}, {256}, {1}, false, false, kSparseModeRightDown, KeyLayout::TND}, + {"L0-tnd-all-kv0", 8, 128, 0, {32}, {0}, {0}, false, false, kSparseModeRightDown, KeyLayout::TND}, + {"L0-tnd-all-q0-kv0", 8, 128, 0, {0}, {0}, {0}, false, false, kSparseModeRightDown, KeyLayout::TND}, + // PA key dim0 非连续(A2/A3 与 950):间隔槽下毒,错 stride 会读到 poison。 + {"L0-stride-bbnd", 2, 16, 2, {2}, {5}, {16}, false, false, kSparseModeRightDown, KeyLayout::BBND, 0, 2}, + {"L1-stride-bbnd", + 8, + 128, + 8, + {32, 17}, + {300, 130}, + {2, 0}, + false, + false, + kSparseModeRightDown, + KeyLayout::BBND, + 0, + 2}, + {"L1-stride-bnbd", + 8, + 128, + 8, + {32, 17}, + {300, 130}, + {2, 0}, + false, + false, + kSparseModeRightDown, + KeyLayout::BNBD, + 0, + 2}, + {"L1-stride-int8", + 8, + 128, + 8, + {32, 17}, + {300, 130}, + {2, 0}, + false, + true, + kSparseModeRightDown, + KeyLayout::BBND, + 0, + 2}, + // block_table 宽 257:score 末维 RoundUp=272 > UB 单窗 256,950 C2UB 滑窗 flush。 + {"L0-wide-table-257", 2, 16, 2, {2}, {5}, {16}, false, false, kSparseModeRightDown, KeyLayout::BBND, 0, 1, 257}, + {"L0-fp8-wide-table-257", + 2, + 128, + 2, + {2}, + {5}, + {16}, + false, + false, + kSparseModeRightDown, + KeyLayout::BBND, + 1, + 1, + 257}, + {"L1-wide-table-257-bf16", + 8, + 128, + 8, + {32}, + {256}, + {1}, + true, + false, + kSparseModeRightDown, + KeyLayout::BBND, + 0, + 1, + 257}, + }; + + size_t passed = 0; + size_t ran = 0; + size_t skipped = 0; + const char *socName = aclrtGetSocName(); + const std::string soc = (socName == nullptr) ? "" : socName; + const bool isAscend950 = (soc.find("950") != std::string::npos) || (soc.find("910_95") != std::string::npos); + (void)printf("running %zu MsaIndexScore cases (soc=%s)\n", cases.size(), soc.empty() ? "?" : soc.c_str()); + for (const auto &tc : cases) { + if (tc.fp8Kind != 0 && !isAscend950) { + (void)printf(" [%s] SKIP: FP8 is Ascend 950 only\n", tc.name.c_str()); + ++skipped; + continue; + } + ++ran; + if (RunCase(tc, stream)) { + ++passed; + } + } + const std::string extra = (skipped == 0) ? "" : (" (skipped " + std::to_string(skipped) + ")"); + (void)printf("%s: %zu/%zu cases passed%s\n", passed == ran ? "[PASS]" : "[FAIL]", passed, ran, extra.c_str()); + + (void)aclrtDestroyStream(stream); + (void)aclrtResetDevice(deviceId); + (void)aclFinalize(); + return passed == ran ? 0 : -1; +} diff --git a/csrc/attention/msa_index_score/op_host/msa_index_score_def.cpp b/csrc/attention/msa_index_score/op_host/msa_index_score_def.cpp index 5d3467feaf14..431f7eacf1c6 100644 --- a/csrc/attention/msa_index_score/op_host/msa_index_score_def.cpp +++ b/csrc/attention/msa_index_score/op_host/msa_index_score_def.cpp @@ -10,7 +10,7 @@ /*! * \file msa_index_score_def.cpp - * \brief MsaIndexScore 算子原型注册(对齐 ../README.md)。 + * \brief MsaIndexScore 算子原型注册(对齐 docs/aclnnMsaIndexScore.md)。 * * dtype 组合(按列表下标对齐): * A2/A3(3 组,全 ND):0 bf16/bf16、1 fp16/fp16、2 fp16/int8 diff --git a/csrc/attention/msa_index_score/op_host/msa_index_score_tiling.cpp b/csrc/attention/msa_index_score/op_host/msa_index_score_tiling.cpp index a313317212e6..5d91cb5c3c3b 100644 --- a/csrc/attention/msa_index_score/op_host/msa_index_score_tiling.cpp +++ b/csrc/attention/msa_index_score/op_host/msa_index_score_tiling.cpp @@ -58,118 +58,6 @@ inline uint32_t RoundUpU32(uint32_t value, uint32_t align) return (value + align - 1) / align * align; } -// MIX:按估计 M-task 数启动 AIC,避免短 decode 打满空核(scheduler.Init GetValue / MODE4 VC 空转)。 -// Host 读不到 per-request q_len。kernel TotalTasks = Σ CeilDiv(qLen_b * Hq, MSA_ROW_TILE_M)。 -// Σ ceil(x_i) <= ceil(Σ x_i) + B;B=1 时 packed == actual。B>1 用 packed+B 上界,避免少估把 -// 多 tile 请求串到已启动核上。多 M-tile / 大 batch 仍截到全量 AIC。 -inline uint32_t EstPackedMTasks(const MsaIndexScoreInfo &info, uint32_t aicNum) -{ - if (info.totalQ == 0U || aicNum == 0U) { - return 1U; - } - const uint64_t packedRows = static_cast(info.totalQ) * static_cast(info.numQHeads); - const uint32_t packedM = - static_cast((packedRows + static_cast(MSA_ROW_TILE_M) - 1U) / MSA_ROW_TILE_M); - uint32_t est = packedM; - if (info.batch > 1U) { - const uint32_t add = packedM + info.batch; - est = (add < packedM) ? aicNum : add; - } - if (est < 1U) { - est = 1U; - } - if (est > aicNum) { - est = aicNum; - } - return est; -} - -// 可见 KV page 上界:优先 host 可读的 actual_seq_klen;否则 min(表宽, PA numPages)。 -// GetData() 对 device tensor 也返回非空指针,解引用会段错误;仅 kOnHost/kFollowing 可读。 -inline bool IsHostVisibleTensor(const gert::Tensor *tensor) -{ - if (tensor == nullptr) { - return false; - } - const auto placement = tensor->GetPlacement(); - return (placement == gert::kOnHost) || (placement == gert::kFollowing); -} - -inline uint32_t EstMaxVisibleBlocks(gert::TilingContext *context, const MsaIndexScoreInfo &info) -{ - uint32_t vis = info.maxBlocksPerBatch; - if (info.keyLayout != MSA_KEY_LAYOUT_TND && info.numPages > 0U && info.numPages < vis) { - vis = info.numPages; - } - auto klenTensor = context->GetOptionalInputTensor(MSA_IDX_ACTUAL_SEQ_KLEN); - if (!IsHostVisibleTensor(klenTensor)) { - return vis; - } - const int32_t *data = klenTensor->GetData(); - if (data == nullptr) { - return vis; - } - uint32_t maxTok = 0U; - if (info.keyLayout == MSA_KEY_LAYOUT_TND) { - for (uint32_t b = 0; b < info.batch; ++b) { - const int32_t delta = data[b + 1U] - data[b]; - if (delta > 0 && static_cast(delta) > maxTok) { - maxTok = static_cast(delta); - } - } - } else { - for (uint32_t b = 0; b < info.batch; ++b) { - if (data[b] > 0 && static_cast(data[b]) > maxTok) { - maxTok = static_cast(data[b]); - } - } - } - const uint32_t fromLen = (maxTok + MSA_BLOCK_SIZE - 1U) / MSA_BLOCK_SIZE; - if (fromLen < vis) { - vis = fromLen; - } - return vis; -} - -// 950:M-task 填不满 AIC 且可见 KV stile 足够时,沿 S 切开打满核。A1 类 packedM≈aicNum 时 kvChunks=1。 -inline uint32_t EstKvChunks(const MsaIndexScoreInfo &info, uint32_t aicNum, uint32_t packedM, uint32_t visBlocks) -{ - if (!info.isAscend950 || packedM == 0U || aicNum == 0U || packedM >= aicNum) { - return 1U; - } - const uint32_t nStiles = (visBlocks + MSA_BLOCKS_PER_STILE - 1U) / MSA_BLOCKS_PER_STILE; - if (nStiles <= 1U) { - return 1U; - } - uint32_t chunks = aicNum / packedM; - if (chunks > nStiles) { - chunks = nStiles; - } - if (chunks < 1U) { - chunks = 1U; - } - return chunks; -} - -inline uint32_t EstLaunchAic(uint32_t aicNum, uint32_t packedM, uint32_t kvChunks) -{ - if (aicNum == 0U) { - return 1U; - } - uint32_t est = packedM; - if (kvChunks > 1U) { - const uint32_t prod = packedM * kvChunks; - est = (prod / kvChunks != packedM) ? aicNum : prod; - } - if (est < 1U) { - est = 1U; - } - if (est > aicNum) { - est = aicNum; - } - return est; -} - inline uint64_t GetDefaultStride0(const gert::Shape &shape) { uint64_t stride = 1; @@ -202,31 +90,30 @@ inline const gert::Stride *GetKeyStrideDesc(gert::TilingContext *context) return context->GetInputStride(MSA_IDX_KEY); } -// 对齐 MlaProlog GetCacheStride0:dim1…末维须等于紧凑 stride,仅 dim0 可非连续。 -// size-1 轴不参与寻址:PyTorch / torch_npu 对其 stride 不做连续约束(BBND↔BNBD 且 N2=1 -// 时 permute().contiguous() 仍可能是 dim1 stride=P 而非 N2*P*D),跳过该轴比较。 -inline ge::graphStatus ValidateKeyInnerAxesContiguous(gert::TilingContext *context, const gert::Stride *stride, - const gert::Shape &shape) +inline bool ValidateKeyInnerAxesContiguous(gert::TilingContext *context, const gert::Stride *stride, + const gert::Shape &shape) { if (stride == nullptr || stride->GetDimNum() != shape.GetDimNum()) { - return ge::GRAPH_SUCCESS; + return true; } - uint64_t expectedStride = 1U; - for (int64_t dim = static_cast(shape.GetDimNum()) - 1; dim >= 1; --dim) { - const int64_t dimSize = shape.GetDim(static_cast(dim)); - OP_CHECK_IF(dimSize < 0, OP_LOGE(context, "key dim%ld size is negative.", dim), return ge::GRAPH_FAILED); - if (dimSize > 1) { - const uint64_t actualStride = static_cast(stride->GetStride(static_cast(dim))); - OP_CHECK_IF(actualStride != expectedStride, - OP_LOGE(context, - "key dim%ld must be contiguous, actual stride is %lu, expected stride is %lu. " - "Only dim0 may be non-contiguous.", - dim, actualStride, expectedStride), - return ge::GRAPH_FAILED); + uint64_t expectedStride = 1; + for (int64_t i = static_cast(shape.GetDimNum()) - 1; i >= 1; --i) { + const int64_t dimSize = shape.GetDim(static_cast(i)); + if (dimSize <= 1) { + // size-1 轴不参与寻址;PyTorch 对其 stride 不做连续约束(BNBD 且 N2=1 时常见)。 + continue; + } + const uint64_t actualStride = static_cast(stride->GetStride(static_cast(i))); + if (actualStride != expectedStride) { + OP_LOGE(context, + "key dim%ld must be contiguous, actual stride=%lu, expected=%lu. " + "Only dim0 (PA page axis) may be non-contiguous.", + i, actualStride, expectedStride); + return false; } expectedStride *= static_cast(dimSize); } - return ge::GRAPH_SUCCESS; + return true; } // PA:strideKvBlock 取 key dim0 元素 stride。连续输入写入值与 shape 推算相同。 @@ -249,9 +136,8 @@ inline ge::graphStatus ResolveStrideKvBlock(gert::TilingContext *context, const strideKvBlock); return ge::GRAPH_SUCCESS; } - if (ValidateKeyInnerAxesContiguous(context, stride, kc) != ge::GRAPH_SUCCESS) { - return ge::GRAPH_FAILED; - } + OP_CHECK_IF(!ValidateKeyInnerAxesContiguous(context, stride, kc), + OP_LOGE(context, "key inner axes must stay contiguous."), return ge::GRAPH_FAILED); const int64_t actualStride0 = stride->GetStride(MSA_DIM_0); if (info.keyLayout == MSA_KEY_LAYOUT_TND) { @@ -546,17 +432,6 @@ ge::graphStatus DoTiling(gert::TilingContext *context, const MsaIndexScoreInfo & const uint32_t kScratchBytes = useKScratch ? (aicNum * kScratchElems * scratchElemBytes) : 0U; const uint32_t kScratchOffsetElems = sWsBytes / sizeof(float); - uint32_t launchAic = aicNum; - uint32_t kvChunks = 1U; - if (info.totalQ == 0U) { - launchAic = 1U; - } else { - const uint32_t packedM = EstPackedMTasks(info, aicNum); - const uint32_t visBlocks = EstMaxVisibleBlocks(context, info); - kvChunks = EstKvChunks(info, aicNum, packedM, visBlocks); - launchAic = EstLaunchAic(aicNum, packedM, kvChunks); - } - MsaIndexScoreTilingData tilingData; tilingData.set_batch(info.batch); tilingData.set_totalQ(info.totalQ); @@ -567,7 +442,7 @@ ge::graphStatus DoTiling(gert::TilingContext *context, const MsaIndexScoreInfo & tilingData.set_blockSize(info.blockSize); tilingData.set_maxBlocksPerBatch(info.maxBlocksPerBatch); tilingData.set_scoreBlockStride(scoreBlockStride); - tilingData.set_usedCoreNum(launchAic); + tilingData.set_usedCoreNum(aicNum); tilingData.set_isQuant(info.isQuant ? 1U : 0U); tilingData.set_sparseMode(info.sparseMode); tilingData.set_initBlocks(info.initBlocks); @@ -575,7 +450,6 @@ ge::graphStatus DoTiling(gert::TilingContext *context, const MsaIndexScoreInfo & tilingData.set_numPages(info.numPages); tilingData.set_keyLayout(info.keyLayout); tilingData.set_totalK(info.totalK); - tilingData.set_kvChunks(kvChunks); tilingData.set_strideQt(info.numQHeads * info.headDim); tilingData.set_strideQn(info.headDim); @@ -610,17 +484,12 @@ ge::graphStatus DoTiling(gert::TilingContext *context, const MsaIndexScoreInfo & // MIX 1AIC:2AIV:CalcTschBlockDim 的 sliceNum 按 AIV 计数,内部再 / (aiv/aic)。 // 传入 aicNum 会再除一次得到 blockDim=aic/2,只能打一半 Cube。 + // sliceNum = aivNum。 // 整 batch q_len=0:totalQ==0 → BlockDim=1,避免 totalTaskNum=0。 - // sliceNum = launchAic * 2:短 decode 只起实际 M-task(及 950 的 S-chunk)对应的 MIX; - // 多 M-tile / 大 batch 仍打满 AIC。 if (info.totalQ == 0U) { context->SetBlockDim(1); } else { - uint32_t sliceAiv = launchAic * MSA_AIV_PER_AIC; - if (sliceAiv > aivNum) { - sliceAiv = aivNum; - } - context->SetBlockDim(ascendcPlatform.CalcTschBlockDim(sliceAiv, aicNum, aivNum)); + context->SetBlockDim(ascendcPlatform.CalcTschBlockDim(aivNum, aicNum, aivNum)); } size_t *workspaces = context->GetWorkspaceSizes(1); @@ -645,9 +514,8 @@ ge::graphStatus DoTiling(gert::TilingContext *context, const MsaIndexScoreInfo & tilingKey = (info.queryDtype == ge::DT_BF16) ? MSA_TILING_KEY_BF16 : MSA_TILING_KEY_FP16; } context->SetTilingKey(tilingKey); - OP_LOGI(context->GetNodeName(), - "MsaIndexScore tilingKey=%lu headDim=%u layout=%u strideKvBlock=%u launchAic=%u kvChunks=%u aicNum=%u", - tilingKey, info.headDim, info.keyLayout, strideKvBlock, launchAic, kvChunks, aicNum); + OP_LOGI(context->GetNodeName(), "MsaIndexScore tilingKey=%lu headDim=%u layout=%u strideKvBlock=%u", tilingKey, + info.headDim, info.keyLayout, strideKvBlock); return ge::GRAPH_SUCCESS; } } // namespace diff --git a/csrc/attention/msa_index_score/op_host/msa_index_score_tiling.h b/csrc/attention/msa_index_score/op_host/msa_index_score_tiling.h index d939353d589f..14923e73e82f 100644 --- a/csrc/attention/msa_index_score/op_host/msa_index_score_tiling.h +++ b/csrc/attention/msa_index_score/op_host/msa_index_score_tiling.h @@ -66,8 +66,6 @@ TILING_DATA_FIELD_DEF(uint32_t, strideOutToken) TILING_DATA_FIELD_DEF(uint32_t, kScratchOffsetElems) TILING_DATA_FIELD_DEF(uint32_t, keyLayout) // MSA_KEY_LAYOUT_BBND / BNBD / TND TILING_DATA_FIELD_DEF(uint32_t, totalK) // TND:key 第 0 维 T2; PA:0 -// 950 短 M 长 S:每个 M-tile 沿 KV page 切的 chunk 数(stile 对齐)。A2/A3 为 1。 -TILING_DATA_FIELD_DEF(uint32_t, kvChunks) END_TILING_DATA_DEF REGISTER_TILING_DATA_CLASS(MsaIndexScore, MsaIndexScoreTilingData) diff --git a/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_epilogue.h b/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_epilogue.h index cbb8666481a9..80fc91502e19 100644 --- a/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_epilogue.h +++ b/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_epilogue.h @@ -128,30 +128,23 @@ class MsaSegRowMaxEpilogue { const uint32_t subM = MsaCeilDiv(task.mActual, subBlockNum); mOff_ = subIdx * subM; mSub_ = (mOff_ >= task.mActual) ? 0U : MsaMinU32(subM, task.mActual - mOff_); - // 暂存区按半个 M-tile 静态分配;行数超出则退回逐 pass 直写。 - // score 末维超过一窗(如 275 page → stride 288)时也走直写:滑窗 flush 会和 - // S16 乒乓共用硬件事件,第二窗首个 stile(GM 下标 256)会丢掉变成 -inf。 - stageOn_ = (mSub_ > 0U) && (mSub_ <= MSA_STAGE_ROWS) && (strideOutToken_ <= MSA_STAGE_BLOCKS); + // 暂存区按半个 M-tile 静态分配;若 subcore 数变化导致行数超出则退回逐 pass 直写。 + stageOn_ = (mSub_ > 0U) && (mSub_ <= MSA_STAGE_ROWS); stageBlkBase_ = 0; stageBlkEnd_ = 0; rowInReqBase_ = task.mStart + mOff_; tokenBase0_ = task.cuQStart; - writeLo_ = task.sBlkBegin; - writeHi_ = task.writeTailFill ? strideOutToken_ : task.sBlkEnd; if constexpr (!IS_QUANT) { // 本任务两级 S16 的 V_MTE2 余额;EndTask 对称 Wait,避免跨任务/跨启动泄漏。 AscendC::SetFlag(EVENT_ID0); AscendC::SetFlag(EVENT_ID1); } - if (stageOn_) { - OpenStageWindow(0); - } } /// 任务结束:把暂存窗口里剩余的 score 写回 GM。 __aicore__ inline void EndTask() { - FlushStageToStrideEnd(); + FlushStage(); if constexpr (!IS_QUANT) { AscendC::WaitFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID1); @@ -173,7 +166,11 @@ class MsaSegRowMaxEpilogue { if (mSub_ == 0U) { return; } - AdvanceStageWindow(blkBase); + if (stageOn_ && (blkBase >= stageBlkBase_ + MSA_STAGE_BLOCKS)) { + FlushStage(); + stageBlkBase_ = blkBase; + stageBlkEnd_ = blkBase; + } } __aicore__ inline void FinishSTile(uint32_t blkBase) @@ -327,11 +324,11 @@ class MsaSegRowMaxEpilogue { if (stageOn_) { StageScore(blkBase, rowOff - mOff_, rows); } else { - AscendC::SetFlag(EVENT_ID2); - AscendC::WaitFlag(EVENT_ID2); + AscendC::SetFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0); StoreScore(task, blkBase, rowOff, rows); - AscendC::SetFlag(EVENT_ID2); - AscendC::WaitFlag(EVENT_ID2); + AscendC::SetFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0); } } @@ -347,62 +344,28 @@ class MsaSegRowMaxEpilogue { AscendC::PipeBarrier(); } - /// 打开以 base 为起点的 score 暂存窗口,整窗先填 -inf。 - __aicore__ inline void OpenStageWindow(uint32_t base) - { - stageBlkBase_ = base; - stageBlkEnd_ = MsaMinU32(base + MSA_STAGE_BLOCKS, strideOutToken_); - AscendC::Duplicate(ubStage_, MSA_FILL_VALUE, mSub_ * MSA_STAGE_BLOCKS); - AscendC::PipeBarrier(); - } - - /// blk 越过当前窗时先 flush,再开下一窗。MTE3 用 EVENT_ID2,避开 S16 乒乓的 EVENT_ID0/1。 - __aicore__ inline void AdvanceStageWindow(uint32_t blk) - { - if (!stageOn_) { - return; - } - while (blk >= stageBlkBase_ + MSA_STAGE_BLOCKS) { - FlushStage(); - OpenStageWindow(stageBlkBase_ + MSA_STAGE_BLOCKS); - } - } - - /// 当前窗中本 chunk 负责的列;writeTailFill 时 writeHi_ 延到 score 末维。 - __aicore__ inline void FlushStageToStrideEnd() - { - FlushStage(); - } - - /// 把 score 暂存窗口按行写回 GM:每行一次 ≥32B 对齐的连续 DataCopy。 - /// MTE3 用 EVENT_ID2,避开 S16 乒乓的 EVENT_ID0/1。 + /// 把 score 暂存窗口按行写回 GM:每行一次 ≥1KB 的连续 DataCopy。 __aicore__ inline void FlushStage() { if (!stageOn_ || stageBlkEnd_ <= stageBlkBase_) { return; } - const uint32_t lo = (writeLo_ > stageBlkBase_) ? writeLo_ : stageBlkBase_; - const uint32_t hi = (writeHi_ < stageBlkEnd_) ? writeHi_ : stageBlkEnd_; - if (hi <= lo) { - return; - } - const uint32_t count = hi - lo; - const uint32_t ubOff = lo - stageBlkBase_; - AscendC::SetFlag(EVENT_ID2); - AscendC::WaitFlag(EVENT_ID2); + const uint32_t count = stageBlkEnd_ - stageBlkBase_; + AscendC::SetFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0); uint32_t tOff = rowInReqBase_ / numQHeads_; uint32_t h = rowInReqBase_ - tOff * numQHeads_; - uint64_t tokenBase = static_cast(tokenBase0_ + tOff) * strideOutToken_ + lo; + uint64_t tokenBase = static_cast(tokenBase0_ + tOff) * strideOutToken_ + stageBlkBase_; for (uint32_t r = 0; r < mSub_; ++r) { AscendC::DataCopy(gScore_[tokenBase + static_cast(h) * strideOutHead_], - ubStage_[r * MSA_STAGE_BLOCKS + ubOff], count); + ubStage_[r * MSA_STAGE_BLOCKS], count); if (++h == numQHeads_) { h = 0; tokenBase += strideOutToken_; } } - AscendC::SetFlag(EVENT_ID2); - AscendC::WaitFlag(EVENT_ID2); + AscendC::SetFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0); } /// sparse_mode=3:rightDownCausal;否则仅 kv_len 截断。 @@ -724,8 +687,6 @@ class MsaSegRowMaxEpilogue { uint32_t tokenBase0_ = 0; uint32_t stageBlkBase_ = 0; uint32_t stageBlkEnd_ = 0; - uint32_t writeLo_ = 0; - uint32_t writeHi_ = 0; bool stageOn_ = false; uint32_t numQHeads_ = 1; diff --git a/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_kernel.h b/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_kernel.h index e4b9d9d813c0..7a179f3f6de5 100644 --- a/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_kernel.h +++ b/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_kernel.h @@ -104,7 +104,7 @@ class MsaIndexScoreKernel { scheduler_.Init(tiling_->batch, tiling_->numQHeads, tiling_->maxBlocksPerBatch, tiling_->scoreBlockStride, tiling_->sparseMode, tiling_->initBlocks, tiling_->localBlocks, tiling_->keyLayout, - tiling_->kvChunks, gActualSeqQlen_, gActualSeqKlen_, gStartLoc_); + gActualSeqQlen_, gActualSeqKlen_, gStartLoc_); } __aicore__ inline void Process() @@ -141,7 +141,7 @@ class MsaIndexScoreKernel { scheduler_.Decode(taskIdx, task); // 只对可见 S-tile 做 QKᵀ + 握手;因果不可见尾由 AIV 直接写 -inf,避免空转同步。 bool needLoadQ = true; - for (uint32_t st = task.sStileBegin; st < task.sStileEnd; ++st) { + for (uint32_t st = 0; st < task.numComputeSTiles; ++st) { const bool needKScratch = StileNeedsKScratch(task, st * MSA_BLOCKS_PER_STILE); if (needKScratch) { // MIX 1AIC:2AIV:AIV→AIC 的 0x2 flag 需两个 AIV 都 Set 后才放行。 @@ -503,7 +503,7 @@ class MsaIndexScoreKernel { ++tileSeq; } } else { - for (uint32_t st = task.sStileBegin; st < task.sStileEnd; ++st) { + for (uint32_t st = 0; st < task.numComputeSTiles; ++st) { if (StileNeedsKScratch(task, st * MSA_BLOCKS_PER_STILE)) { IssueKGatherAndNotify(resource, task, st, coreKScratch, subIdx, subBlockNum, flagK0, flagK1, flagK2, flagK3); @@ -516,10 +516,8 @@ class MsaIndexScoreKernel { } } // 因果不可见尾:不握手,直接写 -inf(与 AIC 跳过这些 tile 对齐)。 - if (task.writeTailFill) { - for (uint32_t st = task.numComputeSTiles; st < task.numSTiles; ++st) { - epilogue.ProcessSTile(gWorkspace_[0], task, st * MSA_BLOCKS_PER_STILE); - } + for (uint32_t st = task.numComputeSTiles; st < task.numSTiles; ++st) { + epilogue.ProcessSTile(gWorkspace_[0], task, st * MSA_BLOCKS_PER_STILE); } epilogue.EndTask(); } diff --git a/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_task.h b/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_task.h index c8f51be65a29..e2d213a36636 100644 --- a/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_task.h +++ b/csrc/attention/msa_index_score/op_kernel/arch22/msa_index_score_task.h @@ -60,7 +60,7 @@ __aicore__ inline int32_t MsaClampI32(int32_t v, int32_t lo, int32_t hi) return v; } -/// 一个任务 = 一个请求内的一个 M-tile,可选再沿 KV stile 切开(950 短 decode)。 +/// 一个任务 = 一个请求内的一个 M-tile。 /// M 维 = 请求内 (token, head) 扁平行索引,head 为低位。 struct MsaTask { uint32_t batchIdx; @@ -77,13 +77,8 @@ struct MsaTask { uint32_t localBlocks; uint32_t fullEndBlk; uint32_t visibleEndBlk; - uint32_t sBlkBegin; // 本 chunk 计算的 page 半开区间 - uint32_t sBlkEnd; - uint32_t sStileBegin; // GM 路径 stile 循环 - uint32_t sStileEnd; uint32_t numSTiles; // 含不可见尾,必须写 -inf 的 S-tile 数 uint32_t numComputeSTiles; // AIC 真正做 QKᵀ 的 S-tile 数(<= numSTiles) - bool writeTailFill; // 本 M-tile 的最后 chunk 负责 [visible, stride) 填 -inf }; class MsaTaskScheduler { @@ -92,7 +87,7 @@ class MsaTaskScheduler { __aicore__ inline void Init(uint32_t batch, uint32_t numQHeads, uint32_t maxBlocksPerBatch, uint32_t scoreBlockStride, uint32_t sparseMode, uint32_t initBlocks, - uint32_t localBlocks, uint32_t keyLayout, uint32_t kvChunks, + uint32_t localBlocks, uint32_t keyLayout, const AscendC::GlobalTensor &actualSeqQlen, const AscendC::GlobalTensor &actualSeqKlen, const AscendC::GlobalTensor &startLoc) @@ -105,15 +100,14 @@ class MsaTaskScheduler { initBlocks_ = initBlocks; localBlocks_ = localBlocks; keyLayout_ = keyLayout; - kvChunks_ = (kvChunks < 1U) ? 1U : kvChunks; gActualSeqQlen_ = actualSeqQlen; gActualSeqKlen_ = actualSeqKlen; gStartLoc_ = startLoc; totalTasks_ = 0; for (uint32_t b = 0; b < batch_; ++b) { - LoadBatch(b); - totalTasks_ += curTasks_; + const uint32_t qLen = static_cast(gActualSeqQlen_.GetValue(b + 1) - gActualSeqQlen_.GetValue(b)); + totalTasks_ += MsaCeilDiv(qLen * numQHeads_, MSA_ROW_TILE_M); } Reset(); } @@ -149,99 +143,6 @@ class MsaTaskScheduler { LoadBatch(curBatch_); } - const uint32_t local = taskIdx - taskBase_; - uint32_t mStart = 0; - uint32_t mActual = 0; - uint32_t sChunkIdx = 0; - uint32_t sChunks = 1; - if (kvChunks_ <= 1U) { - GetMTile(local, mStart, mActual); - } else { - uint32_t acc = 0; - const uint32_t mTiles = MsaCeilDiv(curRows_, MSA_ROW_TILE_M); - for (uint32_t t = 0; t < mTiles; ++t) { - GetMTile(t, mStart, mActual); - const uint32_t vis = VisibleEndBlkOf(mStart, mActual); - sChunks = KvChunksOf(vis); - if (local < acc + sChunks) { - sChunkIdx = local - acc; - break; - } - acc += sChunks; - } - } - FillTask(task, mStart, mActual, sChunkIdx, sChunks); - } - -private: - __aicore__ inline void GetMTile(uint32_t tileIdx, uint32_t &mStart, uint32_t &mActual) const - { - mStart = tileIdx * MSA_ROW_TILE_M; - mActual = MsaMinU32(MSA_ROW_TILE_M, (curRows_ > mStart) ? (curRows_ - mStart) : 0U); - if (mActual > 0U && mActual < MSA_M_ALIGN && mStart >= MSA_M_ALIGN) { - mStart -= (MSA_M_ALIGN - mActual); - mActual = MSA_M_ALIGN; - } - } - - __aicore__ inline uint32_t VisibleEndBlkOf(uint32_t mStart, uint32_t mActual) const - { - if (mActual == 0U) { - return 0U; - } - const uint32_t tokenHigh = (mStart + mActual - 1U) / numQHeads_; - const int32_t visibleKeyEndHi = VisibleKeyEndOf(static_cast(tokenHigh)); - uint32_t visibleEndBlk = MsaCeilDiv(static_cast(visibleKeyEndHi), MSA_BLOCK_SIZE); - return MsaMinU32(visibleEndBlk, maxBlocksPerBatch_); - } - - __aicore__ inline uint32_t KvChunksOf(uint32_t visibleEndBlk) const - { - if (kvChunks_ <= 1U) { - return 1U; - } - const uint32_t nStiles = MsaCeilDiv(visibleEndBlk, MSA_BLOCKS_PER_STILE); - if (nStiles <= 1U) { - return 1U; - } - return MsaMinU32(kvChunks_, nStiles); - } - - __aicore__ inline void AssignSRange(uint32_t visibleEndBlk, uint32_t sChunkIdx, uint32_t sChunks, - MsaTask &task) const - { - const uint32_t nStiles = MsaCeilDiv(visibleEndBlk, MSA_BLOCKS_PER_STILE); - if (sChunks <= 1U || nStiles == 0U) { - task.sBlkBegin = 0; - task.sBlkEnd = visibleEndBlk; - task.writeTailFill = true; - } else { - uint32_t sc = sChunks; - if (sc > nStiles) { - sc = nStiles; - } - const uint32_t base = nStiles / sc; - const uint32_t extra = nStiles % sc; - uint32_t st0 = 0; - uint32_t nst = 0; - if (sChunkIdx < extra) { - st0 = sChunkIdx * (base + 1U); - nst = base + 1U; - } else { - st0 = extra * (base + 1U) + (sChunkIdx - extra) * base; - nst = base; - } - task.sBlkBegin = st0 * MSA_BLOCKS_PER_STILE; - task.sBlkEnd = MsaMinU32((st0 + nst) * MSA_BLOCKS_PER_STILE, visibleEndBlk); - task.writeTailFill = (sChunkIdx + 1U == sc); - } - task.sStileBegin = task.sBlkBegin / MSA_BLOCKS_PER_STILE; - task.sStileEnd = MsaCeilDiv(task.sBlkEnd, MSA_BLOCKS_PER_STILE); - } - - __aicore__ inline void FillTask(MsaTask &task, uint32_t mStart, uint32_t mActual, uint32_t sChunkIdx, - uint32_t sChunks) const - { task.batchIdx = curBatch_; task.cuQStart = cuQStart_; task.startLoc = startLoc_; @@ -251,13 +152,23 @@ class MsaTaskScheduler { task.sparseMode = sparseMode_; task.initBlocks = initBlocks_; task.localBlocks = localBlocks_; - task.mStart = mStart; - task.mActual = mActual; + task.mStart = (taskIdx - taskBase_) * MSA_ROW_TILE_M; + task.mActual = MsaMinU32(MSA_ROW_TILE_M, curRows_ - task.mStart); + // Cube L0A fractal 为 MSA_M_ALIGN 行。整请求不足对齐宽度时保持原样(L0 已覆盖); + // 否则把末尾短 tile 向前重叠到 MSA_M_ALIGN 行,避免 mActual ∈ (0,16) 的 int8 路径出错。 + if (task.mActual > 0U && task.mActual < MSA_M_ALIGN && task.mStart >= MSA_M_ALIGN) { + task.mStart -= (MSA_M_ALIGN - task.mActual); + task.mActual = MSA_M_ALIGN; + } task.globalRowBase = cuQStart_ * numQHeads_ + task.mStart; - const uint32_t tLo = (mActual == 0U) ? 0U : (task.mStart / numQHeads_); - const int32_t visibleKeyEndLo = VisibleKeyEndOf(static_cast(tLo)); - uint32_t visibleEndBlk = VisibleEndBlkOf(mStart, mActual); + const uint32_t tokenLo = task.mStart / numQHeads_; + const uint32_t tokenHi = (task.mStart + task.mActual - 1) / numQHeads_; + const int32_t visibleKeyEndHi = VisibleKeyEndOf(static_cast(tokenHi)); + const int32_t visibleKeyEndLo = VisibleKeyEndOf(static_cast(tokenLo)); + + uint32_t visibleEndBlk = MsaCeilDiv(static_cast(visibleKeyEndHi), MSA_BLOCK_SIZE); + visibleEndBlk = MsaMinU32(visibleEndBlk, maxBlocksPerBatch_); const uint32_t causalFull = static_cast(visibleKeyEndLo) / MSA_BLOCK_SIZE; const uint32_t seqFull = (kvLen_ < 0 ? 0U : static_cast(kvLen_)) / MSA_BLOCK_SIZE; @@ -265,29 +176,16 @@ class MsaTaskScheduler { task.visibleEndBlk = visibleEndBlk; task.fullEndBlk = MsaMinU32(fullEndBlk, visibleEndBlk); + + // 不可见尾一律写 -inf,末维对齐到 scoreBlockStride task.numSTiles = MsaCeilDiv(scoreBlockStride_, MSA_BLOCKS_PER_STILE); task.numComputeSTiles = MsaCeilDiv(task.visibleEndBlk, MSA_BLOCKS_PER_STILE); if (task.numComputeSTiles > task.numSTiles) { task.numComputeSTiles = task.numSTiles; } - AssignSRange(visibleEndBlk, sChunkIdx, sChunks, task); } - __aicore__ inline uint32_t CountBatchTasks() const - { - const uint32_t mTiles = MsaCeilDiv(curRows_, MSA_ROW_TILE_M); - if (kvChunks_ <= 1U) { - return mTiles; - } - uint32_t n = 0; - uint32_t mStart = 0; - uint32_t mActual = 0; - for (uint32_t t = 0; t < mTiles; ++t) { - GetMTile(t, mStart, mActual); - n += KvChunksOf(VisibleEndBlkOf(mStart, mActual)); - } - return n; - } +private: __aicore__ inline void LoadBatch(uint32_t b) { if (b >= batch_) { @@ -306,6 +204,7 @@ class MsaTaskScheduler { qLen_ = 0; } curRows_ = static_cast(qLen_) * numQHeads_; + curTasks_ = MsaCeilDiv(curRows_, MSA_ROW_TILE_M); if (keyLayout_ == MSA_KEY_LAYOUT_TND) { cuKStart_ = static_cast(gActualSeqKlen_.GetValue(b)); kvLen_ = static_cast(gActualSeqKlen_.GetValue(b + 1)) - static_cast(cuKStart_); @@ -317,7 +216,6 @@ class MsaTaskScheduler { kvLen_ = gActualSeqKlen_.GetValue(b); } startLoc_ = gStartLoc_.GetValue(b); - curTasks_ = CountBatchTasks(); } AscendC::GlobalTensor gActualSeqQlen_; @@ -332,7 +230,6 @@ class MsaTaskScheduler { uint32_t initBlocks_ = MSA_DEFAULT_INIT_BLOCKS; uint32_t localBlocks_ = MSA_DEFAULT_LOCAL_BLOCKS; uint32_t keyLayout_ = MSA_KEY_LAYOUT_BBND; - uint32_t kvChunks_ = 1; uint32_t totalTasks_ = 0; uint32_t curBatch_ = 0; diff --git a/csrc/attention/msa_index_score/op_kernel/arch35/msa_index_score_kernel.h b/csrc/attention/msa_index_score/op_kernel/arch35/msa_index_score_kernel.h index e46053470e7f..332a63e2afc8 100644 --- a/csrc/attention/msa_index_score/op_kernel/arch35/msa_index_score_kernel.h +++ b/csrc/attention/msa_index_score/op_kernel/arch35/msa_index_score_kernel.h @@ -114,7 +114,7 @@ class MsaIndexScoreKernel { scheduler_.Init(tiling_->batch, tiling_->numQHeads, tiling_->maxBlocksPerBatch, tiling_->scoreBlockStride, tiling_->sparseMode, tiling_->initBlocks, tiling_->localBlocks, tiling_->keyLayout, - tiling_->kvChunks, gActualSeqQlen_, gActualSeqKlen_, gStartLoc_); + gActualSeqQlen_, gActualSeqKlen_, gStartLoc_); } __aicore__ inline void Process() @@ -154,7 +154,7 @@ class MsaIndexScoreKernel { for (uint32_t taskIdx = coreIdx; taskIdx < totalTasks; taskIdx += coreNum) { scheduler_.Decode(taskIdx, task); bool needLoadQ = true; - for (uint32_t st = task.sStileBegin; st < task.sStileEnd; ++st) { + for (uint32_t st = 0; st < task.numComputeSTiles; ++st) { const bool needKScratch = StileNeedsKScratch(task, st * MSA_BLOCKS_PER_STILE); if (needKScratch) { Catlass::Arch::CrossCoreWaitFlag(flagKReady); @@ -178,12 +178,9 @@ class MsaIndexScoreKernel { MsaTask task; for (uint32_t taskIdx = coreIdx; taskIdx < totalTasks; taskIdx += coreNum) { scheduler_.Decode(taskIdx, task); - if (task.sBlkBegin >= task.sBlkEnd) { - continue; - } bool needLoadQ = true; uint32_t pageSeq = 0; - for (uint32_t blk = task.sBlkBegin; blk < task.sBlkEnd; ++blk) { + for (uint32_t blk = 0; blk < task.visibleEndBlk; ++blk) { const uint32_t ping = pageSeq % MSA_A5_S_STAGES; if (KeyBlockNeedsScratch(task, blk)) { AscendC::CrossCoreWaitFlag(MSA_A5_FLAG_K); @@ -458,7 +455,7 @@ class MsaIndexScoreKernel { for (uint32_t taskIdx = coreIdx; taskIdx < totalTasks; taskIdx += coreNum) { scheduler_.Decode(taskIdx, task); epilogue.BeginTask(task, subIdx, subBlockNum); - for (uint32_t st = task.sStileBegin; st < task.sStileEnd; ++st) { + for (uint32_t st = 0; st < task.numComputeSTiles; ++st) { if (StileNeedsKScratch(task, st * MSA_BLOCKS_PER_STILE)) { if (subIdx == 0) { GatherKeySTileToScratch(resource, task, st * MSA_BLOCKS_PER_STILE, coreKScratch); @@ -472,10 +469,8 @@ class MsaIndexScoreKernel { epilogue.ProcessSTile(gWorkspace_[sBase], task, st * MSA_BLOCKS_PER_STILE); ++tileSeq; } - if (task.writeTailFill) { - for (uint32_t st = task.numComputeSTiles; st < task.numSTiles; ++st) { - epilogue.ProcessSTile(gWorkspace_[0], task, st * MSA_BLOCKS_PER_STILE); - } + for (uint32_t st = task.numComputeSTiles; st < task.numSTiles; ++st) { + epilogue.ProcessSTile(gWorkspace_[0], task, st * MSA_BLOCKS_PER_STILE); } epilogue.EndTask(); } @@ -492,12 +487,9 @@ class MsaIndexScoreKernel { MsaTask task; for (uint32_t taskIdx = coreIdx; taskIdx < totalTasks; taskIdx += coreNum) { scheduler_.Decode(taskIdx, task); - if (task.sBlkBegin >= task.sBlkEnd && !task.writeTailFill) { - continue; - } epilogue.BeginTask(task, subIdx, subBlockNum); uint32_t pageSeq = 0; - for (uint32_t blk = task.sBlkBegin; blk < task.sBlkEnd; ++blk) { + for (uint32_t blk = 0; blk < task.visibleEndBlk; ++blk) { const uint32_t ping = pageSeq % MSA_A5_S_STAGES; if (KeyBlockNeedsScratch(task, blk)) { if (subIdx == 0) { diff --git a/csrc/attention/msa_index_score/op_kernel/arch35/msa_seg_row_max_epilogue.h b/csrc/attention/msa_index_score/op_kernel/arch35/msa_seg_row_max_epilogue.h index 5f12074a76fc..c612d1e1c018 100644 --- a/csrc/attention/msa_index_score/op_kernel/arch35/msa_seg_row_max_epilogue.h +++ b/csrc/attention/msa_index_score/op_kernel/arch35/msa_seg_row_max_epilogue.h @@ -159,8 +159,6 @@ class MsaSegRowMaxEpilogue { stageBlkEnd_ = 0; rowInReqBase_ = task.mStart + mOff_; tokenBase0_ = task.cuQStart; - writeLo_ = task.sBlkBegin; - writeHi_ = task.writeTailFill ? strideOutToken_ : task.sBlkEnd; if constexpr (MSA_A5_USE_C2UB) { // 禁止 4B DataCopyPad:MTE3 会按 32B 写出,把后面的 fill 槽写成 0。 // UB 一次最多攒 MSA_STAGE_BLOCKS 列;block_table 宽 >256 时 score 末维 @@ -288,7 +286,7 @@ class MsaSegRowMaxEpilogue { AscendC::PipeBarrier(); } } - if (col == (MSA_BLOCKS_PER_STILE - 1U) || (blk + 1U) == task.sBlkEnd || (blk + 1U) == task.visibleEndBlk) { + if (col == (MSA_BLOCKS_PER_STILE - 1U) || (blk + 1U) == task.visibleEndBlk) { ReduceStileToStage(task, blk - col); } } @@ -423,12 +421,12 @@ class MsaSegRowMaxEpilogue { } } - /// 当前窗中本 chunk 负责的列;writeTailFill 时 writeHi_ 延到 score 末维。 + /// 当前窗 + 一直到 score 末维的后续纯 -inf 窗。 __aicore__ inline void FlushStageToStrideEnd() { FlushStage(); if constexpr (MSA_A5_USE_C2UB) { - while (stageOn_ && stageBlkEnd_ < writeHi_) { + while (stageOn_ && stageBlkEnd_ < strideOutToken_) { OpenStageWindow(stageBlkEnd_); FlushStage(); } @@ -441,22 +439,16 @@ class MsaSegRowMaxEpilogue { if (!stageOn_ || stageBlkEnd_ <= stageBlkBase_) { return; } - const uint32_t lo = (writeLo_ > stageBlkBase_) ? writeLo_ : stageBlkBase_; - const uint32_t hi = (writeHi_ < stageBlkEnd_) ? writeHi_ : stageBlkEnd_; - if (hi <= lo) { - return; - } - const uint32_t count = hi - lo; - const uint32_t ubOff = lo - stageBlkBase_; + const uint32_t count = MsaMinU32(stageBlkEnd_ - stageBlkBase_, MSA_STAGE_BLOCKS); AscendC::PipeBarrier(); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); uint32_t tOff = rowInReqBase_ / numQHeads_; uint32_t h = rowInReqBase_ - tOff * numQHeads_; - uint64_t tokenBase = static_cast(tokenBase0_ + tOff) * strideOutToken_ + lo; + uint64_t tokenBase = static_cast(tokenBase0_ + tOff) * strideOutToken_ + stageBlkBase_; for (uint32_t r = 0; r < mSub_; ++r) { AscendC::DataCopy(gScore_[tokenBase + static_cast(h) * strideOutHead_], - ubStage_[r * MSA_STAGE_BLOCKS + ubOff], count); + ubStage_[r * MSA_STAGE_BLOCKS], count); if (++h == numQHeads_) { h = 0; tokenBase += strideOutToken_; @@ -975,8 +967,6 @@ class MsaSegRowMaxEpilogue { uint32_t tokenBase0_ = 0; uint32_t stageBlkBase_ = 0; uint32_t stageBlkEnd_ = 0; - uint32_t writeLo_ = 0; - uint32_t writeHi_ = 0; bool stageOn_ = false; uint32_t numQHeads_ = 1; diff --git a/csrc/attention/msa_index_score/op_kernel/msa_index_score.cpp b/csrc/attention/msa_index_score/op_kernel/msa_index_score.cpp index 583ce9d7a495..fdcd9a18044e 100644 --- a/csrc/attention/msa_index_score/op_kernel/msa_index_score.cpp +++ b/csrc/attention/msa_index_score/op_kernel/msa_index_score.cpp @@ -17,7 +17,7 @@ * atten_mask 仅做 host 校验;device 侧按 sparse_mode 解析因果,不消费该 GM。 */ -#include "kernel_operator.h" // force-rebuild-arch22-v67-wide-score-direct-store +#include "kernel_operator.h" // force-rebuild-arch22-v65-int8-4slot-dual-aiv #include "lib/matmul_intf.h" #include "msa_index_score_common.h" #if (__CCE_AICORE__ == 310) diff --git a/csrc/attention/msa_index_score/tests/README.md b/csrc/attention/msa_index_score/tests/README.md new file mode 100644 index 000000000000..b4beb7850d70 --- /dev/null +++ b/csrc/attention/msa_index_score/tests/README.md @@ -0,0 +1,97 @@ +# MsaIndexScore Test Guide + +## 1. End-to-End Accuracy Self-Check + +`examples/test_aclnn_msa_index_score.cpp` contains the aclnn invocation and a +CPU golden implementation. + +For Atlas A2/A3: + +```bash +bash build.sh --pkg --soc=ascend910b --ops=msa_index_score -j32 +bash ./build_out/cann-ops-transformer-custom_linux-x86_64.run \ + --quiet --install-path=/tmp/msa_opp +export ASCEND_CUSTOM_OPP_PATH=/tmp/msa_opp/vendors/custom_transformer +bash build.sh --run_example msa_index_score eager cust \ + --vendor_name=custom --soc=ascend910b +``` + +For Ascend 950, pass `--soc=ascend950` explicitly and source the installed +environment before running the example: + +```bash +bash build.sh --pkg --soc=ascend950 --ops=msa_index_score -j32 +bash ./build_out/cann-ops-transformer-custom_linux-x86_64.run \ + --quiet --install-path=/tmp/msa_opp +source /tmp/msa_opp/vendors/custom_transformer/bin/set_env.bash +export ASCEND_CUSTOM_OPP_PATH=/tmp/msa_opp/vendors/custom_transformer +bash build.sh --run_example msa_index_score eager cust \ + --vendor_name=custom --soc=ascend950 +``` + +The expected summary is 40/40 passing cases on Ascend 950. A2/A3 skip four +FP8 cases and run 36 cases. + +## 2. Test Matrix + +`start_loc` is a logical-block index, and `sparse_mode=3` applies +right-down-causal masking. + +| Test Case | Scenario | Coverage | +| --------- | -------- | -------- | +| `L0-debug-trace` | Minimal dimensions | Main path and trace | +| `L0-int8-dequant-trace` | INT8 key with scale | Fused dequantization | +| `L0-prefill-aligned` | Aligned chunked prefill | Causal and local masks | +| `L1-prefill-unaligned` | Variable-length batch | Boundary-block mask | +| `L1-prefill-multi-mtile` | Row count greater than M-tile | M-tile partitioning | +| `L1-decode-lq1` | Decode with `q_len=1` | Multiple sequence lengths | +| `L1-decode-speculative` | Decode with `q_len>1` | Speculative decoding | +| `L1-long-seq-multi-stile` | `kv_len=4096` | Multiple S-tiles | +| `L1-bf16` / `L1-int8-dequant` | Data type | Non-quantized and quantized paths | +| `L2-tiny-kv` | Minimal KV length | Tail padding | +| `L1-bnbd` / `L1-bnbd-int8` | PageAttention BNBD | `[NP, N2, P, D]` layout | +| `L1-tnd-unaligned` / `L1-tnd-int8` / `L0-tnd-tiny` | Packed TND | No block table and key-length prefix sums | +| `L0-fp8-e4m3fn` / `L0-fp8-e5m2` / `L1-fp8-e4m3fn-prefill` | Ascend 950 FP8 | Native E4M3FN/E5M2 Cube paths; HIFLOAT8 is kernel-only | +| `L1-pad-q0` / `L1-pad-kv0` | Empty request in a mixed batch | Skip empty query or fill empty KV scores | +| `L1-pad-q0-kv0` / `L1-tnd-pad-q0-kv0` / `L1-pad-mid-q0` | Empty request at an edge or in the middle | PageAttention and TND padding | +| `L0-all-q0` / `L1-all-q0` | Entire batch has `q_len=0` | Host acceptance and skipped computation | +| `L0-all-kv0` / `L0-all-q0-kv0` | Entire batch has empty KV | Fill scores and fully empty input | +| `L0-tnd-all-q0` / `L0-tnd-all-kv0` / `L0-tnd-all-q0-kv0` | Empty packed TND batch | Empty query and key tensors | +| `L0-stride-bbnd` / `L1-stride-bbnd` / `L1-stride-bnbd` | Page axis has a gap of two | Non-contiguous physical-page addressing | +| `L1-stride-int8` | INT8 page axis has a gap of two | Quantized page copy with a stride | +| `L0-wide-table-257` / `L1-wide-table-257-bf16` | Block-table width 257 | Ascend 950 C2UB windowed flush | +| `L0-fp8-wide-table-257` | Width 257 with FP8 | Aligned score width and fill positions | + +The full matrix runs by default. The key layout is selected by `layout_key` +(`layoutKeyOptional` in aclnn) and is not inferred from tensor rank. + +## 3. Python Reference and Unit Tests + +The CPU reference is in `tests/golden/msa_index_score_golden.py`: + +```python +inputs = MsaIndexScoreGoldenInputs( + query=query, + key=key, + block_table=block_table, + actual_seq_qlen=actual_seq_qlen, + actual_seq_klen=actual_seq_klen, + start_loc=start_loc, + sparse_mode=3, + scale=None, +) +score = msa_index_score_golden(inputs) +``` + +The repository unit test +`tests/ut/ops/test_msa_index_score_golden.py` executes this imported reference +for non-quantized PageAttention, INT8 dequantization, and a non-contiguous +physical-page axis. + +## 4. Acceptance Criteria + +- Fill positions use the negative fill value on both sides. +- Blocks forced by `local_mask` are at least `1e28` on both sides. +- FLOAT16, BFLOAT16, and INT8 use `atol=rtol=1e-3` and an error ratio no + greater than `1e-3`. +- Ascend 950 FP8 uses `atol=rtol=2e-2`. diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_msa_index_score.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_msa_index_score.py deleted file mode 100644 index 9cc6b86e276f..000000000000 --- a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_msa_index_score.py +++ /dev/null @@ -1,120 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project - -import math - -import pytest -import torch -import torch_npu # noqa: F401 -import vllm_ascend.vllm_ascend_C # type: ignore[import-untyped] # noqa: F401 - -from vllm_ascend.device.hardware_profile import HardwareCapability, get_current_hardware_profile - -BLOCK_SIZE = 128 -HEAD_DIM = 128 -NUM_HEADS = 8 -SCORE_ALIGNMENT = 16 -MASK_SIZE = 2048 -FILL_THRESHOLD = -1.0e30 -BOOST_THRESHOLD = 1.0e28 - - -def _reference_scores(query, key, block_table, q_lens, kv_lens, sparse_mode, init_blocks, local_blocks): - """CPU FP32 QK, right-aligned causal masking, then max over each KV block.""" - width = math.ceil(block_table.shape[1] / SCORE_ALIGNMENT) * SCORE_ALIGNMENT - scores = torch.full((NUM_HEADS, sum(q_lens), width), -torch.inf) - q_begin = 0 - for batch, (q_len, kv_len) in enumerate(zip(q_lens, kv_lens)): - block_count = math.ceil(kv_len / BLOCK_SIZE) - if q_len and block_count: - pages = key[block_table[batch, :block_count].long(), :, 0, :] - logits = torch.einsum("qhd,bkd->hqbk", query[q_begin : q_begin + q_len], pages) - positions = torch.arange(block_count * BLOCK_SIZE).view(block_count, BLOCK_SIZE) - visible = torch.full((q_len,), kv_len) - if sparse_mode == 3: - visible = (kv_len - q_len + torch.arange(q_len) + 1).clamp(0, kv_len) - logits.masked_fill_(positions[None, None] >= visible[None, :, None, None], -torch.inf) - block_scores = logits.amax(dim=-1) - block_scores[..., :init_blocks] = 1.0e30 - if local_blocks: - block_scores[..., max(0, block_count - local_blocks) :] = 1.0e29 - scores[:, q_begin : q_begin + q_len, :block_count] = block_scores - q_begin += q_len - return scores - - -def _assert_scores_close(actual, expected, dtype): - assert actual.dtype == torch.float32 - assert actual.shape == expected.shape - actual = actual.cpu() - # Kernels may represent masked/forced scores with finite sentinels or inf. - fill = expected <= FILL_THRESHOLD - boost = expected >= BOOST_THRESHOLD - torch.testing.assert_close(actual <= FILL_THRESHOLD, fill, rtol=0, atol=0) - torch.testing.assert_close(actual >= BOOST_THRESHOLD, boost, rtol=0, atol=0) - valid = ~(fill | boost) - tolerance = 2.0e-2 if dtype == torch.float8_e4m3fn else 1.0e-3 - torch.testing.assert_close(actual[valid], expected[valid], rtol=tolerance, atol=tolerance) - - -@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float8_e4m3fn]) -@pytest.mark.parametrize( - ("q_lens", "kv_lens", "table_width", "page_axis_gap", "sparse_mode", "init_blocks", "local_blocks"), - [ - pytest.param((32, 17), (300, 130), 8, 1, 3, 0, 0, id="prefill"), - pytest.param((32, 17), (300, 130), 8, 2, 3, 0, 0, id="strided-page-axis"), - pytest.param((1, 2), (900, 257), 8, 1, 3, 0, 0, id="decode"), - pytest.param((4,), (32769,), 275, 1, 3, 0, 0, id="long-kv-wide-table"), - pytest.param((2,), (5,), 257, 1, 3, 0, 0, id="wide-table-padding"), - pytest.param((0, 1), (129, 0), 2, 1, 3, 0, 0, id="empty-query-and-kv"), - pytest.param((2,), (513,), 8, 1, 3, 1, 2, id="forced-blocks"), - pytest.param((2,), (129,), 8, 1, 0, 0, 0, id="dense-tp-chunk"), - ], -) -@torch.inference_mode() -def test_msa_index_score_precision( - dtype, q_lens, kv_lens, table_width, page_axis_gap, sparse_mode, init_blocks, local_blocks -): - """Compare the custom operator with an independent reference, without model weights.""" - if dtype == torch.float8_e4m3fn and not get_current_hardware_profile().supports(HardwareCapability.FP8_ATTENTION): - pytest.skip("MsaIndexScore FP8 requires FP8 attention support") - - generator = torch.Generator().manual_seed(2026) - num_pages = max(math.ceil(max(kv_lens) / BLOCK_SIZE) + 3, 4) - # Binary fractions are exactly representable in all tested input dtypes. - query_cpu = torch.randint(-8, 9, (sum(q_lens), NUM_HEADS, HEAD_DIM), generator=generator).float() / 8 - key_cpu = torch.randint(-8, 9, (num_pages, BLOCK_SIZE, 1, HEAD_DIM), generator=generator).float() / 8 - block_table = torch.zeros((len(q_lens), table_width), dtype=torch.int32) - for batch, kv_len in enumerate(kv_lens): - blocks = math.ceil(kv_len / BLOCK_SIZE) - block_table[batch, :blocks] = torch.randperm(num_pages, generator=generator)[:blocks].int() - - query = query_cpu.npu().to(dtype) - storage = torch.full((num_pages * page_axis_gap, BLOCK_SIZE, 1, HEAD_DIM), 7.0, device="npu") - storage[::page_axis_gap].copy_(key_cpu) - key = storage.to(dtype)[::page_axis_gap] - assert key.stride(0) == page_axis_gap * BLOCK_SIZE * HEAD_DIM - - cu_seqlens = torch.tensor([0, *q_lens], dtype=torch.int32).cumsum(0, dtype=torch.int32).npu() - seq_lens = torch.tensor(kv_lens, dtype=torch.int32, device="npu") - start_loc = torch.tensor( - [max(0, math.ceil(length / BLOCK_SIZE) - 1) for length in kv_lens], dtype=torch.int32, device="npu" - ) - mask = torch.zeros((MASK_SIZE, MASK_SIZE), dtype=torch.int8, device="npu") if sparse_mode == 3 else None - expected = _reference_scores( - query_cpu, key_cpu, block_table, q_lens, kv_lens, sparse_mode, init_blocks, local_blocks - ) - actual = torch.ops._C_ascend.npu_msa_index_score( - query, - key, - block_table.npu(), - start_loc, - atten_mask=mask, - actual_seq_qlen=cu_seqlens, - actual_seq_klen=seq_lens, - layout_key="BBND", - sparse_mode=sparse_mode, - init_blocks=init_blocks, - local_blocks=local_blocks, - ) - _assert_scores_close(actual, expected, dtype) diff --git a/tests/ut/models/minimax_m3/test_msa_m3.py b/tests/ut/models/minimax_m3/test_msa_m3.py index cf378442583f..936262186a0b 100644 --- a/tests/ut/models/minimax_m3/test_msa_m3.py +++ b/tests/ut/models/minimax_m3/test_msa_m3.py @@ -544,31 +544,32 @@ def test_sparse_prepare_bypasses_fused_qkv_norm_rope_on_a5() -> None: assert "1.0 + self.q_norm.weight" in source -def test_index_score_uses_ascendc_prefill_and_decode() -> None: +def test_a5_index_score_uses_ascendc_prefill_and_triton_decode() -> None: module_source = inspect.getsource(msa_m3_module) - fp8_branch_start = module_source.index( - "if get_current_hardware_profile().supports(HardwareCapability.FP8_ATTENTION):" - ) - fp8_branch_end = module_source.index("\n\ndef _should_use_tp_sharded_index_decode", fp8_branch_start) - import_branches = module_source[fp8_branch_start:fp8_branch_end] + a5_branch_start = module_source.index("if get_ascend_device_type() == AscendDeviceType.A5:") + a5_branch_end = module_source.index("\n\ndef _should_use_tp_sharded_index_decode", a5_branch_start) + import_branches = module_source[a5_branch_start:a5_branch_end] assert msa_m3_module._USE_ASCENDC_INDEX_SCORE_PREFILL is True - assert msa_m3_module._USE_ASCENDC_INDEX_SCORE_DECODE is True + assert msa_m3_module._USE_ASCENDC_INDEX_SCORE_DECODE is ( + msa_m3_module.get_ascend_device_type() != AscendDeviceType.A5 + ) assert import_branches.count("minimax_m3_index_decode") == 1 assert "msa_m3_triton_a5" in import_branches assert "msa_m3_triton" not in import_branches.replace("msa_m3_triton_a5", "") assert "_USE_ASCENDC_INDEX_SCORE_PREFILL = True" in module_source - assert "_USE_ASCENDC_INDEX_SCORE_DECODE = True" in module_source - with patch.object( - msa_m3_module, "get_current_hardware_profile", return_value=get_hardware_profile(AscendDeviceType.A5) + assert "_USE_ASCENDC_INDEX_SCORE_DECODE = get_ascend_device_type() != AscendDeviceType.A5" in module_source + with patch( + "vllm_ascend.models.minimax_m3.msa_m3.get_ascend_device_type", + return_value=AscendDeviceType.A5, ): assert not _should_use_tp_sharded_index_decode(tp_size=4, num_prefills=0) def test_non_a5_decode_keeps_tp_block_sharding() -> None: with patch( - "vllm_ascend.models.minimax_m3.msa_m3.get_current_hardware_profile", - return_value=get_hardware_profile(AscendDeviceType.A3), + "vllm_ascend.models.minimax_m3.msa_m3.get_ascend_device_type", + return_value=AscendDeviceType.A3, ): assert _should_use_tp_sharded_index_decode(tp_size=4, num_prefills=0) assert not _should_use_tp_sharded_index_decode(tp_size=1, num_prefills=0) @@ -1003,9 +1004,11 @@ def test_bundled_ascendc_index_score_flushes_wide_a5_block_tables() -> None: repo_root = Path(msa_m3_module.__file__).parents[3] op_root = repo_root / "csrc" / "attention" / "msa_index_score" epilogue = (op_root / "op_kernel" / "arch35" / "msa_seg_row_max_epilogue.h").read_text(encoding="utf-8") + example = (op_root / "examples" / "test_aclnn_msa_index_score.cpp").read_text(encoding="utf-8") assert "AdvanceStageWindow" in epilogue assert "FlushStageToStrideEnd" in epilogue + assert "L0-fp8-wide-table-257" in example def test_ascendc_index_score_uses_dense_mode_without_mask() -> None: diff --git a/vllm_ascend/models/minimax_m3/msa_m3.py b/vllm_ascend/models/minimax_m3/msa_m3.py index 2404e798081f..73aa6e68bacf 100644 --- a/vllm_ascend/models/minimax_m3/msa_m3.py +++ b/vllm_ascend/models/minimax_m3/msa_m3.py @@ -39,10 +39,8 @@ from vllm_ascend.attention.attention_mask import AttentionMaskBuilder from vllm_ascend.core.kv_cache_interface import AscendSFAIndexerCacheSpec -from vllm_ascend.device.hardware_profile import HardwareCapability, get_current_hardware_profile from vllm_ascend.models.minimax_m3.ops.msa_m3_npu import ( MiniMaxM3TPDecodeScoreMetadata, - minimax_m3_index_decode_replicated, minimax_m3_index_tp_block_parallel_decode, minimax_m3_sparse_attn, minimax_m3_sparse_attn_decode, @@ -55,13 +53,14 @@ ) from vllm_ascend.ops.linear import AscendColumnParallelLinear from vllm_ascend.ops.linear_op import get_parallel_op +from vllm_ascend.utils import AscendDeviceType, get_ascend_device_type -# The bundled MsaIndexScore supports FP8. Use AscendC for both prefill and -# decode; FP8-capable hardware scores the full local table without TP collectives. +# The bundled MsaIndexScore includes the Ascend 950 arch35 FP8 kernel. Keep it +# enabled for A5 prefill, while A5 decode uses its lower-latency Triton path. _USE_ASCENDC_INDEX_SCORE_PREFILL = True -_USE_ASCENDC_INDEX_SCORE_DECODE = True +_USE_ASCENDC_INDEX_SCORE_DECODE = get_ascend_device_type() != AscendDeviceType.A5 -if get_current_hardware_profile().supports(HardwareCapability.FP8_ATTENTION): +if get_ascend_device_type() == AscendDeviceType.A5: from vllm_ascend.models.minimax_m3.ops.msa_m3_triton_a5 import ( minimax_m3_index_decode, minimax_m3_index_score, @@ -70,13 +69,10 @@ def _should_use_tp_sharded_index_decode(tp_size: int, num_prefills: int) -> bool: - # The FP8-capable decode backend scores the complete, replicated index-K - # cache on every TP rank. Other profiles retain TP block sharding. - return ( - not get_current_hardware_profile().supports(HardwareCapability.FP8_ATTENTION) - and tp_size > 1 - and num_prefills == 0 - ) + # The A5 Triton decode kernel operates on the complete, replicated index-K + # cache on every TP rank. Keep the mainline block-sharded optimization for + # the other device families only. + return get_ascend_device_type() != AscendDeviceType.A5 and tp_size > 1 and num_prefills == 0 def _active_decode_num_reqs( @@ -368,12 +364,7 @@ def build( cu_seqlens_q=decode_cu_seqlens_q, context_lens=decode_context_lens, ) - if ( - _USE_ASCENDC_INDEX_SCORE_DECODE - and not get_current_hardware_profile().supports(HardwareCapability.FP8_ATTENTION) - and self.tp_size > 1 - and active_prefills == 0 - ): + if _USE_ASCENDC_INDEX_SCORE_DECODE and self.tp_size > 1 and active_prefills == 0: decode_metadata.tp_score = self._build_tp_score_metadata( decode_metadata.block_table, decode_cu_seqlens_q, @@ -516,26 +507,7 @@ def forward( tp_group = get_tp_group() decode_iq = iq[:num_decode_tokens] if _USE_ASCENDC_INDEX_SCORE_DECODE: - if get_current_hardware_profile().supports(HardwareCapability.FP8_ATTENTION): - decode_start_loc = torch.div( - d.context_lens, - self.block_size, - rounding_mode="floor", - ).to(dtype=torch.int32) - decode_topk, decode_select_num_idx = minimax_m3_index_decode_replicated( - decode_iq, - kv, - d.block_table, - d.cu_seqlens_q, - d.seq_lens, - decode_start_loc, - index_md.causal_mask, - topk=self.topk_blocks, - init_blocks=self.init_blocks, - local_blocks=self.local_blocks, - decode_query_len=d.decode_query_len, - ) - elif tp_group.world_size > 1 and index_md.num_prefills == 0: + if tp_group.world_size > 1 and index_md.num_prefills == 0: decode_topk = minimax_m3_index_tp_block_parallel_decode( decode_iq, kv, diff --git a/vllm_ascend/models/minimax_m3/ops/msa_m3_npu.py b/vllm_ascend/models/minimax_m3/ops/msa_m3_npu.py index 9a1bb9e5280d..2be23796f054 100644 --- a/vllm_ascend/models/minimax_m3/ops/msa_m3_npu.py +++ b/vllm_ascend/models/minimax_m3/ops/msa_m3_npu.py @@ -26,9 +26,6 @@ _FP8_E4M3_MAX = 448.0 if get_current_hardware_profile().supports(HardwareCapability.FP8_ATTENTION): - from vllm_ascend.models.minimax_m3.ops.msa_m3_triton_a5 import ( - _index_topk_postprocess_kernel, - ) from vllm_ascend.models.minimax_m3.ops.msa_m3_triton_a5 import ( minimax_m3_index_topk as _minimax_m3_index_prefill_topk, ) @@ -174,7 +171,6 @@ def _minimax_m3_index_score( *, init_blocks: int = 0, local_blocks: int = 0, - force_blocks_in_kernel: bool = False, ) -> torch.Tensor: """Compute MSA index scores with the bundled AscendC operator. @@ -184,15 +180,12 @@ def _minimax_m3_index_score( A causal mask selects sparse mode 3. Passing no mask selects dense mode 0 for a TP chunk that is entirely before the current query positions. - ``init_blocks`` and ``local_blocks`` can be applied in the AscendC - epilogue for a complete block table. TP-sharded scoring keeps them at - zero because the operator attributes do not carry a global block offset; - that path applies global forcing in its candidate TopK stage instead. + ``init_blocks`` and ``local_blocks`` are kept for parity with the index + scoring interface; candidate forcing is applied by the TopK stage. """ index_kv_cache = _as_ascendc_index_kv_cache(index_kv_cache) if index_kv_cache.dtype == torch.float8_e4m3fn and idx_q.dtype != index_kv_cache.dtype: idx_q = _to_fp8_e4m3(idx_q) - block_force_kwargs = {"init_blocks": init_blocks, "local_blocks": local_blocks} if force_blocks_in_kernel else {} return torch.ops._C_ascend.npu_msa_index_score( idx_q, index_kv_cache, @@ -203,72 +196,7 @@ def _minimax_m3_index_score( actual_seq_klen=seq_lens, layout_key="BBND", sparse_mode=3 if causal_mask is not None else 0, - **block_force_kwargs, - ) - - -@torch.no_grad() -def minimax_m3_index_decode_replicated( - idx_q: torch.Tensor, - index_kv_cache: torch.Tensor | tuple[torch.Tensor], - block_table: torch.Tensor, - cu_seqlens_q: torch.Tensor, - seq_lens: torch.Tensor, - start_loc: torch.Tensor, - causal_mask: torch.Tensor, - *, - topk: int, - init_blocks: int, - local_blocks: int, - decode_query_len: int, -) -> tuple[torch.Tensor, torch.Tensor]: - """Score the full replicated table with fused block forcing and TopK cleanup.""" - assert get_current_hardware_profile().supports(HardwareCapability.FP8_ATTENTION) - score = _minimax_m3_index_score( - idx_q, - index_kv_cache, - block_table, - cu_seqlens_q, - seq_lens, - start_loc, - causal_mask, - init_blocks=init_blocks, - local_blocks=local_blocks, - force_blocks_in_kernel=True, - ) - if score.shape[-1] < topk: - raise ValueError(f"MsaIndexScore width {score.shape[-1]} is smaller than topk {topk}") - - _, raw_topk = torch.topk(score, k=topk, dim=-1) - num_index_heads, total_q, _ = raw_topk.shape - topk_indices = torch.empty( - raw_topk.shape, - dtype=torch.int32, - device=raw_topk.device, - ) - select_num_idx = torch.empty( - (num_index_heads, total_q), - dtype=torch.int32, - device=raw_topk.device, - ) - _index_topk_postprocess_kernel[(total_q, num_index_heads)]( - raw_topk, - topk_indices, - select_num_idx, - seq_lens, - _MSA_INDEX_BLOCK_SIZE, - topk, - decode_query_len, - raw_topk.stride(0), - raw_topk.stride(1), - raw_topk.stride(2), - topk_indices.stride(0), - topk_indices.stride(1), - topk_indices.stride(2), - select_num_idx.stride(0), - select_num_idx.stride(1), ) - return topk_indices, select_num_idx @torch.no_grad()