Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
248 changes: 45 additions & 203 deletions csrc/attention/msa_index_score/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,9 @@

| Product | Supported |
| --------------------------------------------------------------------- | :-------: |
| <term>Atlas A2 Products</term> | √ |
| <term>Atlas A3 Products</term> | √ |
| <term>950PR&950DT Products</term> | √ |
| <term>Atlas A2 Products</term> | √ |
| <term>Atlas A3 Products</term> | √ |
| <term>950PR&950DT Products</term> | √ |

## Function Description

Expand Down Expand Up @@ -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

Expand All @@ -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

Expand All @@ -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

Expand All @@ -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)
Loading
Loading