maskShape约束
MaskType |
dtype |
Shape |
Extra |
备注 |
|---|---|---|---|---|
UNDEFINED |
- |
无输入 |
- |
|
NORM |
同query的dtype |
[maxSeqLen,maxSeqLen] / [batch,maxSeqLen,maxSeqLen] |
- |
3D 时 batch 与 seqLen 一致 |
ALIBI |
同query的dtype |
1. [headNum, maxSeqLen, maxSeqLen] 2. [batch, headNum, maxSeqLen, maxSeqLen] |
- |
分别对应: 1. 每个batch相同的mask; 2. 每个batch不同的mask |
NORM_COMPRESS |
同query的dtype |
[128,128] |
- |
不支持 quant |
ALIBI_COMPRESS/SQRT |
同query的dtype |
[256,256] 或 [headNum,seqlen,128] |
slopes [headNum] |
fp16 mask 须 HIGH_PRECISION |
ALIBI_LEFT_ALIGN |
同query的dtype |
仅 [256,256] |
slopes [headNum] |
- |
SWA_NORM |
同query的dtype |
[S,S] |
- |
windowSize>0 |
SWA_COMPRESS |
同query的dtype |
[512,512] |
- |
windowSize>0 |
MaskType |
输入数 |
mask shape |
slopes |
|---|---|---|---|
ALIBI_COMPRESS/SQRT |
8 |
同 PA ALIBI compress |
√ |
NORM_COMPRESS |
7 |
同 PA [128,128] |
× |
CAUSAL_MASK |
6 |
无(内部生成) |
× |
Atlas 350 加速卡 PA_ENCODER(ND)
MaskType |
dtype |
Shape |
|---|---|---|
NORM |
int8 |
[maxSeqLen,maxSeqLen] / [batch,maxSeqLen,maxSeqLen] |
NORM_COMPRESS |
int8 |
[2048,2048] |
ALIBI |
int8 |
[batch,headNum,maxSeqLen,maxSeqLen] |
MaskType |
Format |
Shape |
备注 |
|---|---|---|---|
NORM_COMPRESS |
同query的dtype |
[1,8,128,16] |
legacy |
NORM_COMPRESS long |
同query的dtype |
[1,128,2048,16] |
仅 |
SWA_COMPRESS |
同query的dtype |
[1,32,512,16] |