TilingKey#
本节介绍TilingKey的声明和使用方法。TilingKey使用有限的编译期配置,为同一份 PyPTO Pro Kernel生成多个专用实例,并在启动时选择目标实例。TilingKey适用于会改变 代码路径、Tile模板或数据布局的离散模式。
所有示例使用以下导入:
from pypto_pro.runtime.tilingkey import TilingKeyField
import pypto_pro.language as pl
何时使用TilingKey#
TilingKey字段在编译阶段会被折叠为常量。因此,每个具体Key都会生成独立Kernel, 不会在单个Kernel中保留对应的运行时分支。TilingKey适合以下场景:
可选功能会改变较大代码路径,例如是否应用attention mask;
Tile尺寸、Layout或Dtype模板只有有限候选值;
需要在启动前拒绝不支持的字段组合;
二进制交付时需要枚举可编译的TilingKey。
TilingKey不适用于任意运行时shape。候选值的笛卡尔积决定可枚举的组合数量,因此 TilingKey应仅描述有限且会影响代码生成的模式。
声明Schema#
TilingKey是一个普通Python类。类属性使用TilingKeyField(bits=..., values=...)声明:
class AttentionKey:
# 二进制开关,取值只能为 0 或 1。
HasAtten = TilingKeyField(bits=1, values=[0, 1])
# 两种固定 Tile 模板。
BlockM = TilingKeyField(bits=8, values=[64, 128])
def is_valid(self, key):
has_atten, block_m = key
# 仅举例:mask 模式只支持 128 行 Tile。
return has_atten == 0 or block_m == 128
字段按类定义顺序收集。该顺序同时决定:
is_valid()中key元组的元素顺序;64-bit编码中各字段的bit offset;
二进制交付头文件中的字段和selector顺序。
is_valid()是可选校验函数,参数key使用按照字段定义顺序排列的元组。
该函数既用于在JIT启动时校验具体Key,也用于在二进制交付时过滤枚举组合。
字段约束#
框架在应用@pl.jit装饰器时检查schema:
约束 |
说明 |
|---|---|
|
必须是class,且至少包含一个 |
|
必须大于0。 |
|
必须非空、元素必须是互不重复的 |
编码容量 |
候选数量不得超过 |
总位宽 |
所有字段位宽之和不得超过64。 |
|
若定义,必须可调用。 |
字段位宽用于为候选下标分配编码空间,不限制候选值本身的数值范围。实际值可以是稀疏的
模板编号,例如bits=3的候选集合可以是[16, 64, 128];它们分别编码为0、1、2。
编码与AscendC对齐#
values的顺序具有语义:编码后的TilingKey字段保存该值在values中的下标,解码后
得到Kernel中使用的实际值。该行为与Ascend C模板TilingKey的Selector一致。
class MaskKey:
NeedAttnMask = TilingKeyField(bits=1, values=[1, 0])
此字段的映射为:
实际值 |
|
字段Bit / TilingKey |
|---|---|---|
|
|
|
|
|
|
因此,TilingKey为1时,NeedAttnMask的实际值为0。启动时仍传入实际值,
例如{"NeedAttnMask": 0},不需要手工传入下标。
将Schema绑定到Kernel#
通过@pl.jit(tiling_key=...)关联schema。TilingKey字段在Kernel中作为编译期变量直接引用,
无需在Kernel形参列表中声明:
@pl.jit(auto_mutex=True, tiling_key=AttentionKey)
def attention_kernel(
q: pl.Ptr[pl.DT_FP16],
k: pl.Ptr[pl.DT_FP16],
out: pl.Ptr[pl.DT_FP16],
):
for qi in pl.range(0, 32):
kv_end = 32
# HasAtten 在此处是当前 TilingKey 实例对应的编译期常量。
if HasAtten == 1:
kv_end = qi + 1
for ki in pl.range(0, kv_end):
# ...
pass
例如,HasAtten=1时,parser保留kv_end = qi + 1分支;HasAtten=0时,
编译阶段移除该分支并保留初始值kv_end = 32。二者共享源码,但最终Kernel中不会保留对
HasAtten的运行时判断。
TilingKey字段名不得与Kernel形参或模块级普通变量冲突。字段名应描述其编译期语义,
例如HasAtten、S1TemplateType,并避免使用n、shape等可能与运行时变量冲突的名称。
选择实例并启动#
带TilingKey的Kernel必须在方括号启动参数中提供完整的Key字典,位置在stream和
block_dim之后:
key = {"HasAtten": 1, "BlockM": 128}
attention_kernel[None, num_cores, key](q, k, out)
Key字典的字段必须与Schema 完全一致:
每个字段都必须出现,且不能有额外字段;
值必须属于该字段的
values;整个组合必须通过
is_valid();不能直接调用
attention_kernel(...),也不能使用list或tuple代替Key字典。
若同时使用datatype特化,TilingKey仍是第三个参数,datatype dict紧随其后:
attention_kernel[None, num_cores, key, datatype](q, k, out)
框架按照字段定义顺序,将Key字典中的实际值转换为对应的values下标,再打包为唯一的
64-bit Key,并缓存对应的专用编译结果。
FlashAttention特化示例#
以下示例使用FaTilingKey为flash_attention_score生成causal attention和
full attention两种专用实例。其他Kernel实参不参与TilingKey Schema,也不影响具体Key的选择。
FaTilingKey的字段#
FaTilingKey声明14个编译期字段:
字段 |
bits |
候选值 |
此用例的固定值 |
|---|---|---|---|
|
2 |
0, 1 |
0 |
|
2 |
0, 1, 2 |
0 |
|
4 |
0, 1, 2, 3, 4 |
1 |
|
10 |
0, 16, 64, 128, 256 |
128 |
|
10 |
0, 16, 32, 64, 128, 256, 512 |
128 |
|
12 |
0, 16, 32, 48, 64, 80, 96, 128, 160, 192, 256, 768 |
128 |
|
12 |
同 |
128 |
|
4 |
0, 1, 2, 3, 4, 9 |
9 |
|
1 |
0, 1 |
0或1 |
|
1 |
0, 1 |
0 |
|
1 |
0, 1 |
0 |
|
2 |
0, 1, 2 |
0 |
|
1 |
0, 1 |
1 |
|
1 |
0, 1 |
0 |
这些字段总计63 bits。is_valid()将除HasAtten之外的字段限制为表中的固定值,
因此候选值的笛卡尔积最终只保留两个合法Key:HasAtten=0和HasAtten=1。
Kernel在Cube和Vector两个循环中均使用HasAtten选择causal和full专用路径:
causal_skv = skv_tiles
if HasAtten == 1:
causal_skv = qi + 1
HasAtten=0的专用Kernel仅保留full attention的skv_tiles路径;HasAtten=1的
专用Kernel仅保留causal attention的qi + 1路径。源码中可以保留清晰的if/else结构,
而每个具体Key的最终代码只包含可达分支。
启动两个专用实例#
基础Key如下:
base_key = {
"KernelTypeKey": 0, "ImplMode": 0, "Layout": 1,
"S1TemplateType": 128, "S2TemplateType": 128,
"DTemplateType": 128, "DvTemplateType": 128,
"PseMode": 9, "HasAtten": 0, "HasDrop": 0, "HasRope": 0,
"OutDtype": 0, "Regbase": 1, "OptionalDn": 0,
}
causal_key = {**base_key, "HasAtten": 1}
flash_attention_score[None, actual_num_cores, causal_key, datatype](
query, key, value, ...
)
full_key = {**base_key, "HasAtten": 0}
flash_attention_score[None, actual_num_cores, full_key, datatype](
query, key, value, ...
)
两次启动使用相同的大部分Key字段,仅改变HasAtten,分别选择causal attention和
full attention专用实例。该示例支持FP16和BF16。
二进制交付#
对带TilingKey的Kernel调用generate_binary_headers()可生成TilingKey头文件:
from pypto_pro.runtime.opc.pypto_compile import generate_binary_headers
binary_dir = generate_binary_headers(flash_attention_score)
生成的FaTilingKey_tilingkey.h包含字段声明及通过is_valid()
的Key Selector。该文件使用ASCENDC_TPL_ARGS_DECL描述各字段和允许值,并以
ASCENDC_TPL_SEL仅列出合法组合。字段bit选择values中对应下标的实际值;因此应尽量
收紧values,并在存在字段关联约束时实现is_valid(),避免生成无用的二进制实例。
常见错误#
现象 |
原因与处理 |
|---|---|
应用Kernel装饰器时失败 |
检查 |
启动时字段不匹配 |
Key字典必须包含所有字段且不能包含未知字段;字段名大小写必须与类属性一致。 |
启动值不在候选集中 |
将实际值加入声明的 |
启动被 |
按字段定义顺序检查组合约束。 |
Kernel中找不到字段名 |
在 |
为每个shape新增Key |
仅将真正影响代码生成的有限模式放入TilingKey。 |
最小示例#
import pypto_pro.language as pl
from pypto_pro.runtime.tilingkey import TilingKeyField
class MyKey:
UseFastPath = TilingKeyField(bits=1, values=[0, 1])
def is_valid(self, key):
(use_fast_path,) = key
return use_fast_path in (0, 1)
@pl.jit(auto_mutex=True, tiling_key=MyKey)
def kernel(x: pl.Ptr[pl.DT_FP16]):
if UseFastPath == 1: # 编译期常量
pass
for i in pl.range(0, 8):
pass
kernel[None, 1, {"UseFastPath": 1}](x)
用TilingKey选择有限的专用实现,从而消除关键模式分支的运行时开销。