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()