Skip to content

[VMI] fp8 layout cost model: contiguous vs contiguous+lane_stride=4 for f32->fp8 masked_store #1567

Description

@Zhendong404

背景

在 per_token_cast 的 bf16 量化路径中,有一个典型的:

bf16x256 -> f32x256 -> f8E4M3FNx256 -> masked_store

场景。当前 VMI layout assignment 为 fp8 结果选择了:

#pto.vmi.layout<contiguous, lane_stride = 4>

而不是普通的:

#pto.vmi.layout<contiguous>

这会导致 VMIToVPTO 走 4 个 part-specific vcvt 加 4 次 PK4_B32 store,而 ASC 的等价路径是 part-wise vcvt、vor 合并、1 次 NORM_B8 store。这里存在一个 layout/cost model 问题。

复现配置

代表性 case:

per_token_cast
fmt=e4m3
in_dtype=bf16
hidden=16384
num_tokens=512
round_sf=true
use_packed_ue8m0=true
use_tma_aligned_col_major_sf=false

VMILayoutAssignment 结果

layout assignment 后,核心 IR 是:

%85 = pto.vmi.mulf %84, %80
  : !pto.vmi.vreg<256xbf16, #pto.vmi.layout<contiguous>>,
    !pto.vmi.vreg<256xbf16, #pto.vmi.layout<contiguous>>
  -> !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 {rounding = "R", saturate = "SAT"}
  : !pto.vmi.vreg<256xf32, #pto.vmi.layout<contiguous>>
  -> !pto.vmi.vreg<256xf8E4M3FN,
      #pto.vmi.layout<contiguous, lane_stride = 4>>

%mask = pto.vmi.create_mask %c256
  : index -> !pto.vmi.mask<256xb8,
      #pto.vmi.layout<contiguous, lane_stride = 4>>

pto.vmi.masked_store %87, %out[%offset], %mask
  : !pto.vmi.vreg<256xf8E4M3FN,
      #pto.vmi.layout<contiguous, lane_stride = 4>>,
    !pto.ptr<f8E4M3FN, ub>,
    !pto.vmi.mask<256xb8,
      #pto.vmi.layout<contiguous, lane_stride = 4>>

Layout 推导来源

该结果来自 lane-stride narrowing 的 preferred cast layout:

kPreferredLaneStrideNarrowCastLayoutPatterns[] = {
    {bits<16>(), bits<8>(), 0, c(), ls(2)},
    {bits<32>(), bits<16>(), 0, c(), ls(2)},
    {bits<32>(), bits<8>(), 0, c(), ls(4)},
};

在 vmi-prefer-lane-stride-narrowing=true 时,getPreferredCastLayoutFact() 会优先命中:

f32 -> f8:
  source = contiguous
  result = contiguous, lane_stride = 4

之后 VMIMaskedStore 看到 value 已经是 lane_stride=4,为了不 rematerialize value,会保留它并推导 matching mask layout:

mask = contiguous, lane_stride = 4

当前 VMIToVPTO lowering

OneToNVMITruncFOpPattern 命中:

sourceBits == 32
resultBits == 8
resultLayout.isContiguous()
resultLayout.getLaneStride() == 4

于是选择:

part = "P0";

后续生成 4 个 part-specific 转换:

%114 = pto.vcvt %110 {part = "P0", rnd = "R", sat = "SAT"}
%115 = pto.vcvt %111 {part = "P0", rnd = "R", sat = "SAT"}
%116 = pto.vcvt %112 {part = "P0", rnd = "R", sat = "SAT"}
%117 = pto.vcvt %113 {part = "P0", rnd = "R", sat = "SAT"}

masked store 看到 lane_stride=4 后,走 packed store:

pto.vsts %114 ... {dist = "PK4_B32"}
pto.vsts %115 ... {dist = "PK4_B32"}
pto.vsts %116 ... {dist = "PK4_B32"}
pto.vsts %117 ... {dist = "PK4_B32"}

此外,前置的 f32 part 获取还产生了:

2 x vintlv
4 x vbitcast

ASC 的等价路径

ASC 对应代码是:

q0 = vcvt(vcvt(x0, f32, part=0), f8, part=0)
q2 = vcvt(vcvt(x0, f32, part=1), f8, part=2)
q1 = vcvt(vcvt(x1, f32, part=0), f8, part=1)
q3 = vcvt(vcvt(x1, f32, part=1), f8, part=3)

merged = vor(vor(q0, q2), vor(q1, q3))
vsts(out, reinterpret(merged, f8x256), dist="NORM_B8")

也就是:

part-wise vcvt
+ vor 合并成完整 fp8x256
+ 1 次 NORM_B8 store

问题本质

当前选择:

fp8 = contiguous + lane_stride=4

避免把 value 重新 materialize 成连续 fp8,但代价是:

4 x vcvt part=P0
+ 4 x PK4_B32 store
+ 前置 vintlv / vbitcast

备选选择:

fp8 = contiguous + lane_stride=1

则可能 lower 成:

4 x part-wise vcvt
+ vor 合并
+ 1 x NORM_B8 store

这是一个 cost model 问题:

  • lane_stride=4:避免 rematerialize,store 次数多。
  • lane_stride=1:需要 vor / merge,store 次数少。

当前 preference 是静态的:

vmi-prefer-lane-stride-narrowing = true

没有根据 consumer 是 dense masked store、输出 lane 数、元素类型和 store family 做动态选择。

建议

  1. 在 layout assignment 中加入 consumer-aware cost model:

    • 对 f32 -> f8 -> dense masked_store 这类模式同时评估:
      • contiguous, lane_stride=4
      • contiguous, lane_stride=1
    • 根据预计的 vcvt、vintlv/vbitcast、vor 和 vsts 数量选择更便宜的 layout。
  2. 或者增加 VMIToVPTO peephole:

    • 识别 4 个 vcvt(part=P0) 加 4 个 PK4_B32 store 的组合;
    • 合并成 part-wise vcvt、vor 和 1 次 NORM_B8 store。

影响

当前此路径在代表性 bf16→E4M3 case 上比 ASC 慢约 2%~3%;部分 e2m1/e4m3 -> e4m3 case 在 5%~9% 范围。FP4 输出路径没有同类问题。

Activity

  1. Zhendong404 commented on Sep 21, 2026

    @Zhendong404
    CollaboratorAuthor

    补充定位:问题本质是 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 = 4
    

    layout 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
    

    原因是通用 VMIToVPTO lowering 在 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,不是 d4 interleave,所以会稳定产生 byte permutation。也就是:只改 layout preference 不够,part 排列也必须是语义正确的。

    验证 2:专用 f32 c -> f8 c lowering 正确但更慢

    补充实现了专用 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_B32 store;但这个选择把本可以在 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 合并作为一个整体来选择。

    修正后的建议

    1. 不能简单地移除 f32 -> f8 的 lane-stride preference,或只强制 f8 c。
      这会得到错误 byte order,或者引入昂贵的 c -> d4 重排。

    2. 需要的是 producer/consumer-aware layout 推导:

      • 对 bf16 -> f32 -> f8 -> masked_store 整段一起评估;
      • 同时考虑 source 已存在的 d2/d4 语义,而不是把 f32 强制压成 c。
    3. 或者在 VMIToVPTO 增加针对当前 f32 c part 序列的专用重组:

      • 直接识别 extf 产生的四个连续 part;
      • 用正确的 lane permutation 生成 P0/P1/P2/P3 + vor + NORM_B8;
      • 避免完整的 c -> d4 通用重排。

    结论:当前问题的本质确实是 VMI layout 推导非最优,但不是“选 ls4 而不是 ls1”这么简单;关键是 layout solver 没有把 d2/d4 物理 part 信息和最终 store family 联合起来做全局选择。

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions