Repository navigation
[VMI] fp8 layout cost model: contiguous vs contiguous+lane_stride=4 for f32->fp8 masked_store #1567
Description
Activity
补充定位:问题本质是 VMI 的 layout 推导是“局部最优、全局非最优”
前面的描述里把备选方案写成
f8 = contiguous, lane_stride=1,进一步验证后需要修正:当前路径不是简单的f32 c -> f8 c,而实际推导结果是:bf16 c -> f32 c -> f8 c, lane_stride = 4 -> masked_store c, lane_stride = 4layout assignment 后的核心类型是:
%85 = pto.vmi.mulf %84, %80 : ..., -> !pto.vmi.vreg<256xbf16, #pto.vmi.layout<contiguous>> %86 = pto.vmi.extf %85 : !pto.vmi.vreg<256xbf16, #pto.vmi.layout<contiguous>> -> !pto.vmi.vreg<256xf32, #pto.vmi.layout<contiguous>> %87 = pto.vmi.truncf %86 : !pto.vmi.vreg<256xf32, #pto.vmi.layout<contiguous>> -> !pto.vmi.vreg<256xf8E4M3FN, #pto.vmi.layout<contiguous, lane_stride = 4>>
因此 VMIToVPTO 走的是:
4 x vcvt P0 + 4 x PK4_B32 store + 前置 vintlv / vbitcast而 ASC 的
P0/P1/P2/P3 -> vor -> NORM_B8对应的是:f32 <deinterleaved = 4> -> f8 <contiguous>不是
f32 <contiguous> -> f8 <contiguous>。这两者的 physical part 语义不同。验证 1:直接强制
f32 c -> f8 c不正确在 legal table 中加入
{f32, f8, c, c}并强制该 relation 后,IR 可以变成:4 x vcvt P0/P1/P2/P3 + 3 x vor + NORM_B8但结果错误:
out_equal=False out_mismatch=29671057原因是通用
VMIToVPTOlowering 在resultLaneStride=1时直接按:source part 0 -> P0 source part 1 -> P1 source part 2 -> P2 source part 3 -> P3做
vor,默认 source 已经是d4排列。f32 c的四个 part 是连续 chunk,不是d4interleave,所以会稳定产生 byte permutation。也就是:只改 layout preference 不够,part 排列也必须是语义正确的。验证 2:专用
f32 c -> f8 clowering 正确但更慢补充实现了专用 lowering:
f32 c -> 寄存器内 c -> d4 重排 -> 4 x vcvt P0/P1/P2/P3 -> 3 x vor -> NORM_B8结果:
out_equal=True sf_equal=True out_mismatch=0 sf_mismatch=0但性能更差:
ASC baseline: 25.9289 us current VMI: ~28.4 us dedicated f32 c -> f8 c: 40.18 us也就是说,把
ls4 + PK4_B32换成c -> d4 -> P-part + vor + NORM_B8时,省掉的是 4 次 store,但新增的c -> d4寄存器重排(多级vintlv)代价远高于省掉的 store。验证 3:把当前 VMI 路径完整移植回 ASC
把当前 VMI 的完整物理路径移植到 ASC:
E2B_B16 inverse + vintlv(inverse, inverse) + 2 x vlds + 2 x vmul + 2 x vintlv(zero, q) + 4 x vcvt P0 + 4 x PK4_B32该 ASC 版本正确:
out_equal=True sf_equal=True性能从:
ASC -> aligned VMI-style ASC 25.9289 us -> 28.6094 us稳定复现约 9.4% 回退。
重新表述问题本质
这不是简单的“静态偏好选错了 lane_stride”,而是:
layout assignment 对当前 producer/consumer 链只做了局部最优选择,保持了
bf16 c -> f32 c,并在f32 -> f8时选择ls=4以直接复用PK4_B32store;但这个选择把本可以在d2/d4物理 part 上完成的 ASC 式part + vor + NORM_B8路径挤掉了。代价项至少包括:
当前 c/ls4 路径: vintlv/b vbitcast part 重组 + 4 x vcvt P0 + 4 x PK4_B32 潜在 d4/c 路径: c -> d4 寄存器重排 + 4 x vcvt P0/P1/P2/P3 + 3 x vor + 1 x NORM_B8需要判断的不是“store 少就一定快”,而是整条 layout/part transition 的总代价。当前 cost model 没有把
c -> d4重排、d2 -> d4传播和vor + NORM_B8合并作为一个整体来选择。修正后的建议
-
不能简单地移除
f32 -> f8的 lane-stride preference,或只强制f8 c。
这会得到错误 byte order,或者引入昂贵的c -> d4重排。 -
需要的是 producer/consumer-aware layout 推导:
- 对
bf16 -> f32 -> f8 -> masked_store整段一起评估; - 同时考虑 source 已存在的
d2/d4语义,而不是把f32强制压成c。
- 对
-
或者在
VMIToVPTO增加针对当前f32 cpart 序列的专用重组:- 直接识别
extf产生的四个连续 part; - 用正确的 lane permutation 生成
P0/P1/P2/P3 + vor + NORM_B8; - 避免完整的
c -> d4通用重排。
- 直接识别
结论:当前问题的本质确实是 VMI layout 推导非最优,但不是“选
ls4而不是ls1”这么简单;关键是 layout solver 没有把d2/d4物理 part 信息和最终 store family 联合起来做全局选择。-
背景
在
per_token_cast的 bf16 量化路径中,有一个典型的:场景。当前 VMI layout assignment 为 fp8 结果选择了:
而不是普通的:
这会导致 VMIToVPTO 走 4 个 part-specific
vcvt加 4 次PK4_B32store,而 ASC 的等价路径是 part-wisevcvt、vor合并、1 次NORM_B8store。这里存在一个 layout/cost model 问题。复现配置
代表性 case:
VMILayoutAssignment 结果
layout assignment 后,核心 IR 是:
Layout 推导来源
该结果来自 lane-stride narrowing 的 preferred cast layout:
在
vmi-prefer-lane-stride-narrowing=true时,getPreferredCastLayoutFact()会优先命中:之后
VMIMaskedStore看到 value 已经是lane_stride=4,为了不 rematerialize value,会保留它并推导 matching mask layout:当前 VMIToVPTO lowering
OneToNVMITruncFOpPattern命中:于是选择:
part = "P0";后续生成 4 个 part-specific 转换:
masked store 看到
lane_stride=4后,走 packed store:此外,前置的 f32 part 获取还产生了:
ASC 的等价路径
ASC 对应代码是:
也就是:
问题本质
当前选择:
避免把 value 重新 materialize 成连续 fp8,但代价是:
备选选择:
则可能 lower 成:
这是一个 cost model 问题:
lane_stride=4:避免 rematerialize,store 次数多。lane_stride=1:需要vor/ merge,store 次数少。当前 preference 是静态的:
没有根据 consumer 是 dense masked store、输出 lane 数、元素类型和 store family 做动态选择。
建议
在 layout assignment 中加入 consumer-aware cost model:
f32 -> f8 -> dense masked_store这类模式同时评估:contiguous, lane_stride=4contiguous, lane_stride=1vcvt、vintlv/vbitcast、vor和vsts数量选择更便宜的 layout。或者增加 VMIToVPTO peephole:
vcvt(part=P0)加 4 个PK4_B32store 的组合;vcvt、vor和 1 次NORM_B8store。影响
当前此路径在代表性 bf16→E4M3 case 上比 ASC 慢约 2%~3%;部分
e2m1/e4m3 -> e4m3case 在 5%~9% 范围。FP4 输出路径没有同类问题。