---
title: MLA
description: "支持通过C接口直调接入PyTorch，在整网中进行亲和算子替换。"
url: https://www.hiascend.com/document/detail/zh/canncommercial/latest/API/ascendtb/ascendtb_01_0314.html
sourcePath: /source/zh/canncommercial/900/API/ascendtb/ascendtb_01_0314.html
indexId: 39b5bda6c42d7811009d01af319635d4a8c531bceec2471f9c8e1d4e8d0891b564
---
# MLA

支持通过C接口直调接入PyTorch，在整网中进行亲和算子替换。


#### 定义

```
atb::Status AtbMLAGetWorkspaceSize(const aclTensor *qNope, const aclTensor *qRope, const aclTensor *ctKV,
const aclTensor *kRope, const aclTensor *blockTables, const aclTensor *contextLens,
const aclTensor *mask, const aclTensor *qSeqLen, const aclTensor *qkDescale,
const aclTensor *pvDescale, int32_t headNum, float qkScale, int32_t kvHeadNum,
int maskType, int calcType, uint8_t cacheMode, aclTensor *attenOut, aclTensor *ise,
uint64_t *workspaceSize, atb::Operation **op, atb::Context *context);
atb::Status AtbMLA(void* workspace, uint64_t workspaceSize, atb::Operation *op, atb::Context *context);
```


#### AtbMLAGetWorkspaceSize成员

| 参数 | 标量/张量 | 维度 | 数据类型 | 格式 | 默认值 | 是否必选 | 描述 |
| --- | --- | --- | --- | --- | --- | --- | --- |
| qNope | 张量 | [num\_tokens, num\_heads, 512] | float16/bf16/int8 | ND | \- | 是 | 无位置编码query。 |
| qRope | 张量 | [num\_tokens, num\_heads, 64] | float16/bf16 | ND | \- | 是 | 旋转位置编码query。 |
| ctKV | 张量 | [num\_blocks , block\_size, kv\_heads, 512] cacheMode为2： [blockNum, kv\_heads\*512/32,block\_size, 32] cacheMode为3： [blockNum, kv\_heads\*512/16,block\_size, 16] | float16/bf16/int8 | ND/NZ | \- | 是 | 无位置编码ctkv。 cacheMode为2：int8，NZ cacheMode为3：float16/bf16，NZ |
| kRope | 张量 | [num\_blocks , block\_size, kv\_heads, 64] cacheMode为2或3： [blockNum, kv\_heads\*64 / 16 ,block\_size, 16] | float16/bf16 | ND/NZ | \- | 是 | 旋转位置编码key。 cacheMode为2或3：NZ |
| blockTables | 张量 | [batch, max\_num\_blocks\_per\_query] | int32 | ND | \- | 是 | 每个query的kvcache的block映射表。 |
| contextLens | 张量 | [batch] | int32 | ND | \- | 是 | 每个query对应的上下文长度，kseqlen。host侧tensor。 |
| mask | 张量 | MASK\_TYPE\_SPEC： [num\_tokens(合轴), max\_seq\_len] MASK\_TYPE\_MASK\_FREE： [125 + 2 \* qseqlen, 128] | float16/bf16 | ND | \- | 否 | 注意力掩码，maskType为0时传入nullptr。 |
| qseqlen | 张量 | [batch] | int32 | ND | \- | 否 | calcType为1时传入，每个batch对应的seqLen。 |
| qkDescale | 张量 | [num\_heads] | float | ND | \- | 否 | cacheMode为2时传入。 |
| pvDescale | 张量 | [num\_heads] | float | ND | \- | 否 | cacheMode为2时传入。 |
| headNum | 标量 | \- | int32\_t | \- | 0 | 是 | query头数量。 |
| qkScale | 标量 | \- | float | \- | 1.0 | 是 | Q\*K^T后乘以的缩放系数。 |
| kvHeadNum | 标量 | \- | int32\_t | \- | 0 | 是 | kv头数量。 |
| maskType | 标量 | \- | int | \- | 0 | 否 | mask类型。 UNDEFINED：默认值，全0的mask。 MASK\_TYPE\_SPEC： qseqlen > 1时的mask。 MASK\_TYPE\_MASK\_FREE：maskfree功能。 |
| calcType | 标量 | \- | int | \- | 0 | 否 | 计算类型。 CALC\_TYPE\_UNDEFINED = 0：默认值。 CALC\_TYPE\_SPEC：支持传入大于1的qseqlen。 CALC\_TYPE\_RING： ringAttention。 CALC\_TYPE\_PREFILL：prefill全量。 |
| cacheMode | 标量 | \- | int | \- | 0 | 是 | KVCACHE = 0：拼接cache。 KROPE\_CTKV：分离cache。 INT8\_NZCACHE：高性能分离cache。 NZCACHE：非量化NZcache。 |
| attenOut | 张量 | [num\_tokens, num\_heads, 512] | float16/bf16 | ND | \- | 是 | attention输出。 |
| lse | 张量 | [num\_tokens, num\_heads, 1] | float16/bf16 | ND | \- | 否 | LSE输出，calcType不为CALC\_TYPE\_RING时可以传入nullptr。 |


#### 原始接口

请参见MLA输入输出。
