maskType
MaskType枚举
位置:“include/atb/infer_op_params.h”
enum MaskType : int {
MASK_TYPE_UNDEFINED = 0, // 全 0 mask
MASK_TYPE_NORM, // 倒三角 causal
MASK_TYPE_ALIBI, // ALIBI 相对位置偏置
MASK_TYPE_NORM_COMPRESS, // 倒三角压缩
MASK_TYPE_ALIBI_COMPRESS, // ALIBI 压缩
MASK_TYPE_ALIBI_COMPRESS_SQRT, // ALIBI 压缩 + 开方
MASK_TYPE_ALIBI_COMPRESS_LEFT_ALIGN, // ALIBI 压缩左对齐(仅 910B)
MASK_TYPE_SLIDING_WINDOW_NORM, // 滑动窗口 + 倒三角
MASK_TYPE_SLIDING_WINDOW_COMPRESS, // 滑动窗口压缩
MASK_TYPE_CAUSAL_MASK, // 算子内部生成 causal
};
值 |
名称 |
语义摘要 |
|---|---|---|
0 |
UNDEFINED |
无 mask,等价全 0。 |
1 |
NORM |
mask[i,j]=0 if i≥j else -inf;int8/bf16使用1做填充值;float16使用-inf。 |
2 |
ALIBI |
线性位置偏置。 |
3 |
NORM_COMPRESS |
块级倒三角;shape 分平台([128, 128], [1, 16, 128, 16] 或 [2048, 2048])。 |
4-6 |
ALIBI_COMPRESS |
压缩 ALIBI + slopes [headNum]。 |
7 |
SWA_NORM |
滑动窗口 + causal;DECODER 可无 mask 输入,内部生成。 |
8 |
SWA_COMPRESS |
滑动窗口压缩;PA 专用。 |
9 |
CAUSAL_MASK |
内部生成;仅 PREFIX。 |
术语与mask输入分类
输入名 |
何时需要 |
典型 Shape |
说明 |
|---|---|---|---|
无 mask 输入 |
UNDEFINED; CAUSAL_MASK; DECODER+SWA_NORM |
- |
CAUSAL 由算子内部生成。 |
attentionMask |
NORM NORM_COMPRESS SWA_* 910B PA ALIBI 等 |
见 Mask Shape 约束 |
主要使用的mask。 |
slopes |
ALIBI_COMPRESS / SQRT / LEFT_ALIGN |
[headNum] |
alibi叠加压缩mask。 |
关联param字段
字段 |
作用 |
|---|---|
maskType |
mask 类型枚举 |
isTriuMask |
倒三角优化; UNDEFINED 时不得为 1; 310P NORM_COMPRESS 构造时强制置 1 |
windowSize |
>0 时须配合 SLIDING_WINDOW_NORM/COMPRESS |
硬件 × MaskType 支持矩阵
MaskType |
PA |
ENCODER/DECODER |
备注 |
|---|---|---|---|
UNDEFINED |
√ |
√ |
|
NORM |
√ |
√ |
|
ALIBI |
√ |
√ |
attentionMask,非 slopes |
NORM_COMPRESS |
√ |
√ |
[128,128] |
ALIBI_COMPRESS* |
√ |
- |
+ slopes |
ALIBI_LEFT_ALIGN |
√ |
- |
仅 |
SWA_NORM/COMPRESS |
√ |
DECODER(NORM)/PA |
|
CAUSAL_MASK |
× |
- |
仅 PREFIX,内部生成mask |
说明 |
备注 |
|---|---|
UNDEFINED、NORM、NORM_COMPRESS(NZ) |
仅 PA_ENCODER |
MaskType |
PA |
备注 |
|---|---|---|
UNDEFINED/NORM/NORM_COMPRESS |
√ |
|
SWA_* |
√ |
NZ 4D |
Atlas 350 加速卡
MaskType |
PA |
备注 |
|---|---|---|
UNDEFINED |
√ |
仅 PA、BSND |
NORM |
√ |
prefill 需 mask;decode 行为不同 |
ALIBI |
√ |
pseShift,非 slopes |
NORM_COMPRESS |
√ |
ND [2048,2048] |