compressor
产品支持情况
- Ascend 950PR/Ascend 950DT:支持
- Atlas A3 训练系列产品/Atlas A3 推理系列产品:支持
- Atlas A2 训练系列产品/Atlas A2 推理系列产品:支持
- Atlas 200I/500 A2 推理产品:不支持
- Atlas 推理系列产品:不支持
- Atlas 训练系列产品:不支持
功能说明
-
接口功能:Compressor是推理场景下SMLA和QLI的前处理算子,用于将每4或128个token的KV cache压缩成一个,然后每个token与这些压缩的KV cache进行DSA计算。在长序列的情况下,Compressor可以有效地减少计算开销。主要计算过程为:
- 将输入XX与WKVW^{KV}做Matmul运算得到kv_statekv\_state,将输入XX与WGateW^{Gate}做Matmul运算后再与ApeApe做Add运算得到score_statescore\_state,kv_statekv\_state与score_statescore\_state根据输入的start_pos及cu_seqlens完成更新。
- 在coff为2的情况下对kv_statekv\_state和score_statescore\_state进行数据重排。
- 对score_statescore\_state进行softmax运算将softmax结果与kv_statekv\_state做Mul计算,后进行ReduceSum运算。
-
计算公式:
-
计算矩阵乘法:
C4A:[kv_statea,score_statea]=X@[WaKV,WaGate],[kv_stateb,score_stateb]=X@[WbKV,WbGate];C4A:\left[kv\_state^a, score\_state^a\right] = X @ \left[W^{aKV}, W^{aGate}\right], \left[kv\_state^b, score\_state^b\right] = X @ \left[W^{bKV}, W^{bGate}\right];
C128A:[kv_state,score_state]=X@[WKV,WGate]C128A:\left[kv\_state, score\_state\right] = X @ \left[W^{KV}, W^{Gate}\right]
-
计算分组加法:
C4A:score_statei′=[score_state[4(i−1)+1:4i,:]a;score_state[4i+1:4(i+1),:]b]+Ape, i=1,2,⋯ ,s4;C4A:score\_state_i^\prime = \left[score\_state_{\left[4(i-1)+1:4i,:\right]}^a; score\_state_{\left[4i+1:4(i+1),:\right]}^b\right] + Ape,~i=1,2,\cdots, \frac{s}{4};
C128A:score_statei′=score_state[128(i−1)+1:128i,:]+Ape, i=1,2,⋯ ,s128;C128A:score\_state_i^\prime = score\_state_{\left[128(i-1)+1:128i,:\right]} + Ape,~i=1,2,\cdots, \frac{s}{128};
-
计算分组Softmax:
C4A:Si′=softmax(score_statei′), i=1,2,⋯ ,s4;C4A:S_i^\prime = softmax(score\_state_i^\prime),~i=1,2,\cdots, \frac{s}{4};
C128A:Si′=softmax(score_statei′), i=1,2,⋯ ,s128;C128A:S_i^\prime = softmax(score\_state_i^\prime),~i=1,2,\cdots, \frac{s}{128};
-
计算Hadamard乘积:
C4A:(SH)i=Si′⊙[kv_state[4(i−1)+1:4i,:]a;kv_state[4i+1:4(i+1),:]b], i=1,2,⋯ ,s4;C4A:(S_H)_i = S_i^\prime \odot \left[kv\_state^a_{\left[4(i-1)+1:4i,:\right]} ;kv\_state^b_{\left[4i+1:4(i+1),:\right]}\right],~i=1,2,\cdots, \frac{s}{4};
C128A:SH=Si′⊙kv_state;C128A:S_H = S_i^\prime \odot kv\_state;
-
沿着压缩轴分组求和:
C4A:CiComp=[1]1×8@(SH)i, i=1,2,⋯ ,s4;C4A:C_{i}^{\text{Comp}} = \left[1\right]_{1\times8} @ (S_H)_i, ~i=1,2,\cdots, \frac{s}{4};
C128A:CiComp=[1]1×128@(SH)i, i=1,2,⋯ ,s128;C128A:C_{i}^{\text{Comp}} = \left[1\right]_{1\times128} @ (S_H)_i, ~i=1,2,\cdots, \frac{s}{128};
-
函数原型
cann_ops_transformer.compressor(
x,
wkv,
wgate,
state_cache,
ape,
cmp_ratio,
*,
state_block_table=None,
cu_seqlens=None,
seqused=None,
start_pos=None,
coff=1,
cache_mode=1) -> Tensor
参数说明
| 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 维度(shape) |
|---|---|---|---|---|---|
| x | Tensor | 必选 | 原始不经压缩的数据,对应公式中的 XX。不支持非连续,数据格式支持ND。 | bfloat16、float16 | [B,S,H]、[T,H] |
| wkv | Tensor | 必选 | kv压缩权重,对应公式中的 WKVW^{KV}。不支持非连续,数据格式支持ND。 | bfloat16、float16 | [coff*D,H] |
| wgate | Tensor | 必选 | gate压缩权重,对应公式中的 WGateW^{Gate}。不支持非连续,数据格式支持ND。 | bfloat16、float16 | [coff*D,H] |
| state_cache | Tensor | 必选 | kv_state和score_state的历史数据,对应公式中的 [kv_state,score_state]\left[kv\_state, score\_state\right]。支持0轴非连续,数据格式支持ND。计算后 kv_state 和 score_state 会原位更新到此 Tensor | float32 | [block_num, block_size, 2*coff*D],要求block_num>0 |
| ape | Tensor | 必选 | positional biases,对应公式中的 ApeApe。不支持非连续,数据格式支持ND。 | float32 | [cmp_ratio,coff*D] |
| cmp_ratio | int | 必选 | 数据压缩率。取值范围为[2, 128]内的整数。 | - | - |
| state_block_table | Tensor | 可选 | state_cache存储使用的block映射表。不支持非连续,数据格式支持ND。 | int32 | cache_mode=1时,shape为[B,ceil(Smax/block_size)],Smax为每个Batch中最大的Sequence Length,当x的shape为[B,S,H]时,Smax=max(start_pos)+S。当x的shape为[T,H]时,Smax=max(start_pos)+max(cu_seqlens[n+1] - cu_seqlens[n])。cache_mode=2时,shape为[B] |
| cu_seqlens | Tensor | 可选 | 不同Batch上的有效token数。不支持非连续,数据格式支持ND。 当x的shape为[B,S,H]时,参数必须为空。 当x的shape为[T,H]时,输入shape必须为[B+1,],该参数为前缀和数组,后一个元素≥前一个元素,第一位必须为0,最后一位必须为T。 |
int32 | [B+1,] |
| seqused | Tensor | 可选 | 不同Batch中实际参与压缩的token数。不支持非连续,数据格式支持ND。 指定为None时,数值等于每个Batch上的Sequence Length。 [B,S,H]场景:0 ≤ seqused[n] ≤ S [T,H]场景:0 ≤ seqused[n] ≤ cu_seqlens[n+1] - cu_seqlens[n]。 |
int32 | [B,] |
| start_pos | Tensor | 可选 | 计算起始位置。不支持非连续,数据格式支持ND,输入为None时从0开始计算 | int32 | [B,] |
| coff | int | 可选 | 表示是否进行overlap数据重排,默认值为1。 仅支持1/2: coff=1:无需进行overlap数据重排。 coff=2:需要进行overlap数据重排。 |
int | - |
| cache_mode | int | 可选 | state_cache的存储模式,默认值为1。 1:连续buffer。 2:循环buffer。 |
int | - |
- Atlas A3 训练系列产品/Atlas A3 推理系列产品:cmp_ratio仅支持2/4/8/16/32/64/128;gradEnabled不支持为true。
返回值说明
| 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 维度(shape) |
|---|---|---|---|---|---|
| cmp_kv | Tensor | 必选 | 压缩后的数据。不支持非连续,数据格式支持ND; 当x的shape为[B,S,H]时,输出拼接:(<batch0>compressed_tokens+pad0) + (<batch1>compressed_tokens+pad1) + ... + (<batchN>compressed_tokens+padN); 当x的shape为[T,H]时,输出拼接:<batch0>compressed_tokens + <batch1>compressed_tokens + ... + <batchN>compressed_tokens + pad。 |
bfloat16、float16 | x=[B,S,H]:[B,ceil(S/cmp_ratio),D] x=[T,H]:[min(T,T//cmp_ratio+B),D] |
约束说明
- 该接口支持推理场景下使用。
- 该接口支持单算子模式和TorchAir图模式(aclgraph)调用。
- x参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、D(Head Dim)表示hidden层的最小单元大小、T表示所有Batch输入样本序列长度的累加和。
- 该接口支持B、S泛化,且存在如下场景限制:
- 部分长序列场景下,如果计算量过大可能会导致出现超过NPU内存的报错,注:这里计算量会受x输入shape的影响,值越大计算量越大。典型的长序列(即B、S的乘积或T较大)场景包括但不限于:
B S H 100 65525 4096 25 261120 4096 100 131072 4096 100 261120 4096
- 部分长序列场景下,如果计算量过大可能会导致出现超过NPU内存的报错,注:这里计算量会受x输入shape的影响,值越大计算量越大。典型的长序列(即B、S的乘积或T较大)场景包括但不限于:
- 该接口支持B、S、T取0,即shape与B、S、T值相关的入参允许传入空tensor,其余入参不支持传入空tensor。该场景下state_cache不做更新,输出cmp_kv为空tensor。
- state_block_table元素取值范围为[0, block_num),block_num为state_cache第0维大小。元素值直接用作state_cache的block索引,越界会导致内存非法访问。元素值为0时:cache_mode=1(连续buffer)下写state_cache操作跳过该位置;cache_mode=2(循环buffer)下读写操作均不跳过。算子不做重复校验,需由调用方保证元素值唯一性:cache_mode=1下元素值0为"未分配"哨兵值可重复出现,非0值须全局唯一;cache_mode=2下0为有效物理块号,所有元素值须全局唯一。重复会导致多个逻辑块/batch写同一物理块区域,造成state_cache数据踩踏覆盖。
- 支持D为128/512。
- 支持H为1K~10K,512对齐。
- 支持block_size为1~1024。
- 支持如下三种典型组合场景:
- C4A: D=512, coff=2, cmp_ratio=4;
- C4Li: D=128, coff=2, cmp_ratio=4;
- C128A: D=512, coff=1, cmp_ratio=128。
确定性计算
- 默认支持确定性计算。
- Ascend 950PR/Ascend 950DT:batch一致性:通过torch_npu.npu.set_deterministic_level()设置确定性级别为3开启batch一致性,开启后可以满足计算结果和所在批次大小和所在批次位置无关。
调用示例
说明:
- 以下示例以C128A场景为例(B=1、S=128、H=4096、D=512、coff=1、cmp_ratio=128),更多参数组合请参考约束说明。
-
单算子模式调用:
import torch import torch_npu import numpy as np import cann_ops_transformer # 参数设置 B = 1 S = 128 H = 4096 D = 512 coff = 1 # 1: no overlap 2: overlap cmp_ratio = 128 cache_mode = 1 block_size = 128 # block_table构造:cache_mode=1时shape为[B, ceil(Smax/block_size)] block_num = (S + block_size - 1) // block_size block_table = torch.zeros(size=(B, block_num), dtype=torch.int32) next_block_id = 1 for i in range(B): for j in range(block_num): block_table[i][j] = next_block_id next_block_id = next_block_id + 1 # 构造输入 x = torch.randn((B, S, H), dtype=torch.bfloat16).npu() wkv = torch.randn((coff * D, H), dtype=torch.bfloat16).npu() wgate = torch.randn((coff * D, H), dtype=torch.bfloat16).npu() ape = torch.randn((cmp_ratio, coff * D), dtype=torch.float32).npu() state_cache = torch.zeros((torch.max(block_table).item() + 1, block_size, 2 * coff * D), dtype=torch.float32).npu() start_pos = torch.zeros((B,), dtype=torch.int32).npu() block_table = block_table.npu() # 调用compressor执行压缩计算 cmp_kv = cann_ops_transformer.compressor( x, wkv, wgate, state_cache, ape, cmp_ratio=cmp_ratio, state_block_table=block_table, cu_seqlens=None, seqused=None, start_pos=start_pos, coff=coff, cache_mode=cache_mode ) print(f"cmp_kv shape: {cmp_kv.shape}") -
TorchAir图模式调用:
import torch import torch_npu import numpy as np import torchair import cann_ops_transformer from torchair.configs.compiler_config import CompilerConfig # 参数设置 B = 1 S = 128 H = 4096 D = 512 coff = 1 cmp_ratio = 128 cache_mode = 1 block_size = 128 block_num = (S + block_size - 1) // block_size block_table = torch.zeros(size=(B, block_num), dtype=torch.int32) next_block_id = 1 for i in range(B): for j in range(block_num): block_table[i][j] = next_block_id next_block_id = next_block_id + 1 x = torch.randn((B, S, H), dtype=torch.bfloat16).npu() wkv = torch.randn((coff * D, H), dtype=torch.bfloat16).npu() wgate = torch.randn((coff * D, H), dtype=torch.bfloat16).npu() ape = torch.randn((cmp_ratio, coff * D), dtype=torch.float32).npu() state_cache = torch.zeros((torch.max(block_table).item() + 1, block_size, 2 * coff * D), dtype=torch.float32).npu() start_pos = torch.zeros((B,), dtype=torch.int32).npu() block_table = block_table.npu() class CompressorNetwork(torch.nn.Module): def __init__(self): super().__init__() def forward(self, x, wkv, wgate, state_cache, ape, block_table, start_pos): return torch.ops.cann_ops_transformer.compressor( x, wkv, wgate, state_cache, ape, cmp_ratio=cmp_ratio, state_block_table=block_table, cu_seqlens=None, seqused=None, start_pos=start_pos, coff=coff, cache_mode=cache_mode ) config = CompilerConfig() config.mode = "reduce-overhead" npu_backend = torchair.get_npu_backend(compiler_config=config) torch._dynamo.reset() npu_mode = torch.compile(CompressorNetwork(), fullgraph=True, backend=npu_backend, dynamic=False) cmp_kv = npu_mode(x, wkv, wgate, state_cache, ape, block_table, start_pos) print(f"cmp_kv shape: {cmp_kv.shape}")