MaterializeTensorStrides Pass¶
将程序中所有 TensorType / DistributedTensorType 上的 view.has_value() && view.stride.empty() 槽位按对应 layout 的 packed canonical 公式填入显式 stride(参考 RFC #1300 §2.4)。Pass 运行后即满足 codegen 入口契约:每个存在的 TensorView 都带显式 stride 与其 layout / shape 一致,严格模式 TensorViewCanonical verifier 也会通过。
状态:本 Pass 已注册(
passes.materialize_tensor_strides())、有单测覆盖,并自 RFC #1300 P6 起接入默认 tile/PTO pipeline,位置在CanonicalizeIOOrder与InitMemRef之间。
概述¶
PyPTO IR 上 TensorType.tensor_view_ 当前可以处于两种等价形态:
- 隐式 ——
view.has_value() && view.stride.empty():layout 标签已设(如DN),但每维 stride 为空。下游消费方需把空 stride 当作「该 layout 的 packed canonical stride」。 - 显式 —— 每个维度的 stride
ExprPtr都已写出。
为了让 codegen 看到单一可机械读取的契约,MaterializeTensorStrides 遍历整个程序,把所有隐式 TensorView 用 tensor_view_semantics.h 中的 BuildLogicalStridesFromLayout 改写为显式 packed canonical 形式。裸 TensorType(!view.has_value())不被改写:TensorViewCanonical verifier 在弱/严格模式下都把它当作 ND-packed 接受,本身就无歧义。输入类型是 DistributedTensorType 时,重建后的类型仍保持 distributed wrapper,并保留 memref、TensorView.pad 等非 stride 元数据以及 window_buffer 反向引用。
Requirements:
SSAForm、SplitIncoreOrch、IncoreTileOps、TileOps2D、TileMemoryInferred、NormalizedStmtStructure
Produces:
TensorViewCanonical——PassPipeline在 Pass 之后自动用 registry 中的严格模式 verifier 校验(拒绝view.has_value() && stride.empty()—— 正是本 Pass 负责消除的状态)
默认 pipeline 中的位置(自 RFC #1300 P6 起激活):CanonicalizeIOOrder 与 InitMemRef 之间。这是 codegen-prep 边界 —— 所有 layout-mutating pass(ResolveBackendOpLayouts / ExpandMixedKernel / SplitVectorKernel)已结束,InitMemRef 是第一个依赖显式 stride 的消费者。
API¶
| C++ | Python | 级别 |
|---|---|---|
pass::MaterializeTensorStrides() |
passes.materialize_tensor_strides() |
Program-level |
from pypto.pypto_core import passes
mat_pass = passes.materialize_tensor_strides()
program_canon = mat_pass(program)
算法¶
Pass 使用带 Var 替换缓存的 IRMutator,结构与 InferTileMemorySpace 一致。它遍历程序中可达的每个 TypePtr:
- 逐函数重建 —— 重新构造形参 / 返回类型 / 函数体:
- 遍历形参类型;若某形参
TensorType经MaterializeType后变为不同的类型,构造新Var(保留name_hint与span)并登记替换。 - 同理处理返回类型。
-
通过
IRMutator::VisitStmt遍历函数体:VisitExpr_(VarPtr):若该 Var 的类型经MaterializeType改变,构造新 Var(查var_cache_,确保对同一个原 Var 的所有引用都解析到同一新 Var)。VisitExpr_(IterArgPtr):与 Var 同理,附加递归处理init_value_。VisitExpr_(CallPtr):注册 op 走OpRegistry重建;GlobalVar调用 / 未注册 op 走直接Call构造路径。VisitStmt_(AssignStmtPtr):先重建 RHS;若 RHS Call 的返回类型比 LHS Var 当前类型更显式(已物化),同步 LHS Var。
-
类型重写 ——
MaterializeType(type): TensorType/DistributedTensorType满足view.has_value() && view.stride.empty() && layout != NZ:用BuildLogicalStridesFromLayout(shape, layout)重建,并保留 distributed wrapper 与可选元数据(memref、TensorView.pad、window_buffer)。其他 tensor 形态原样返回(保持指针身份)。TensorType且layout == NZ:原样返回(NZ 在TensorType上属于非法 IR;交由 verifier 报错而非在这里BuildLogicalStridesFromLayoutCHECK-fail)。TupleType:递归处理元素类型;任一子类型变化时重建。- 其它:原样返回。
Pass 幂等 —— 在已物化的 IR 上重跑等于无操作(类型比较走指针身份就短路;无变化时 MutableCopy 也被跳过)。
| 行为 | 触发条件 |
|---|---|
| 用 packed canonical 填入 stride | view.has_value() && view.stride.empty() 且 layout in {ND, DN} |
| 原样直通 | !view.has_value()(裸 tensor) |
| 原样直通 | view.has_value() && !view.stride.empty()(已显式) |
| 原样直通 | view.layout == NZ(由 verifier 单独拒绝) |
示例¶
Before —— InCore 形参带有空 stride 的 DN view(用户写的 pl.Tensor[..., pl.DN] 未给显式 stride 提示):
@pl.function(type=pl.FunctionType.InCore)
def kernel(b: pl.Tensor[[2, 4, 8], pl.FP32, pl.TensorView(stride=[], layout=pl.TensorLayout.DN)],
out: pl.Out[pl.Tensor[[2, 4, 8], pl.FP32]]) -> pl.Tensor[[2, 4, 8], pl.FP32]:
...
After:
@pl.function(type=pl.FunctionType.InCore)
def kernel(b: pl.Tensor[[2, 4, 8], pl.FP32, pl.TensorView(stride=[32, 1, 4], layout=pl.TensorLayout.DN)],
out: pl.Out[pl.Tensor[[2, 4, 8], pl.FP32]]) -> pl.Tensor[[2, 4, 8], pl.FP32]:
...
shape [2, 4, 8] 的 DN packed canonical stride:
stride[1] = 1(DN 内层对中较小的那一维)stride[2] = shape[1] = 4stride[0] = shape[1] * shape[2] = 32
ND 情况下公式退化为标准行主序 packed stride。
Stride 公式¶
详见 tensor_view_semantics.h 中的 BuildLogicalStridesFromLayout。
| Layout | 公式 |
|---|---|
ND |
stride[n-1] = 1; stride[k] = stride[k+1] * shape[k+1],k = n-2 .. 0 |
DN(n ≥ 2) |
stride[n-2] = 1;stride[n-1] = shape[n-2];stride[n-3] = shape[n-2] * shape[n-1];stride[k] = stride[k+1] * shape[k+1],k = n-4 .. 0 |
NZ |
无法用 flat stride 表达(分形,仅 tile 使用)—— BuildLogicalStridesFromLayout CHECK-fail |
MakeIndexMul 对 ConstInt * ConstInt 做常量折叠(带 __builtin_mul_overflow 守卫,溢出时回退到符号 Mul 而不是静默 wrap),并消除 × 1 单位元;这样符号维保留为 Mul 表达式,静态常量链折叠为单个 ConstInt。
与 verifier 的协同¶
由于 Pass 声明 produced = {... ∪ TensorViewCanonical},PassPipeline 在 Pass 完成后自动调用 registry 中的 TensorViewCanonical verifier。registry 默认是严格模式 verifier(RFC #1300 §2.4 codegen 入口契约):它拒绝 view.has_value() && stride.empty() —— 因为本 Pass 就是负责物化这些 stride 的。裸 TensorType(!view.has_value())仍然接受 —— 隐式 ND-packed 自然 canonical。同一 verifier 也可通过 passes.verify_tensor_view_canonical(program, require_materialized=True) 显式调用;传 require_materialized=False 切换到弱模式(用于物化之前的解析期 / 前期 pass 窗口)。
相关¶
CanonicalizeIOOrder—— 紧邻其前;产生本 Pass 消费的程序状态InitMemRef—— 第一个依赖显式 stride 的下游消费者tensor_view_semantics.h—— 工具函数(BuildLogicalStridesFromLayout/CheckCanonicalView/CanonicalizeView)- RFC #1300 —— IR Tensor Layout 自洽表示方案