Skip to content

[VMI] bf16 quantize 路径应将 values/inverse 推导为 d2 layout,避免 vintlv 扩宽 inverse #1568

Description

@Zhendong404

背景

在 per_token_cast 的 bf16 量化路径中,一个 256-element tile 的逻辑计算是:

inverse = V.vload(
    sf_inv_ub[row, tile * 8],
    size=256,
    dist_mode="brc",
    group=8,
    stride=1,
)
values = V.vload(x_ub[row, col], size=256)
quantized = V.vmul(values, inverse, mask256)

这个 VMI 表达本身没有问题。问题是当前 layout 推导把 values、inverse、quantized 都推成了普通 contiguous,而 ASC 对应路径实际使用的是 d2 layout。

ASC 对应实现

ASC 的 quantize_bf16 是:

x0, x1 = S.vld2(x_ub[row, col], dist="DINTLV_B16")
inverse = S.vld(sf_inv_ub[row, tile * 8], dist="E2B_B16")
quantized0 = S.vmul(x0, inverse)
quantized1 = S.vmul(x1, inverse)

关键点:

  1. vld2 DINTLV_B16 一次 load 出两个 deinterleaved 的 128-lane bf16 向量 x0/x1。
  2. x0/x1 是 d2 物理布局。
  3. inverse 是一个 128-lane E2B_B16 broadcast packet。
  4. x0/x1 共用同一个 inverse packet。
  5. 后续乘法直接在 d2 的两个 physical part 上完成,不需要把 inverse 通过 vintlv 扩成两个 part。

当前 VMI layout 推导结果

当前 VMILayoutAssignment 后,计算链是:

%80 = pto.vmi.group_broadcast_load ...
  : !pto.ptr<bf16, ub> -> !pto.vmi.vreg<256xbf16, #pto.vmi.layout<contiguous>>

%84 = pto.vmi.load ...
  : !pto.ptr<bf16, ub> -> !pto.vmi.vreg<256xbf16, #pto.vmi.layout<contiguous>>

%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
  : !pto.vmi.vreg<256xf32, #pto.vmi.layout<contiguous>>
  -> !pto.vmi.vreg<256xf8E4M3FN,
      #pto.vmi.layout<contiguous, lane_stride = 4>>

也就是 value 和 inverse 都被推成了普通 contiguous。

当前 lowering 的额外开销

为了把 contiguous inverse 变成两个 128-lane physical part,当前 lowering 需要:

1 x E2B_B16
+ 1 x vintlv(inverse, inverse)

而不是 ASC 的:

1 x E2B_B16
+ 两个 physical part 共用同一份 inverse

进入 BF16 -> F32 -> FP8 后,还会继续产生:

vintlv(zero_bf16, quantized)
+ vbitcast
+ 4 x vcvt(part=P0)
+ 4 x vsts PK4_B32

而不是 ASC 的 d2 路径:

x0/x1 d2
+ inverse E2B
+ 2 x vmul
+ part-wise vcvt
+ vor 合并
+ 1 x NORM_B8 store

核心问题

当前布局选择和 ASC 的物理数据流不一致:

ASC:
  values       -> d2
  inverse      -> d2 / E2B shared
  quantized    -> d2
  downstream   -> d2-part compatible

VMI current:
  values       -> contiguous
  inverse      -> contiguous + vintlv expansion
  quantized    -> contiguous
  vector layout conversion

所以问题不是 VMI 语法该怎么写,而是 layout assignment 应该识别:

256xbf16 group_broadcast_load
+ 256xbf16 contiguous load
+ vmul

这个模式,把 values、inverse、quantized 推导成 d2 layout,而不是 contiguous。

期望的推导结果

对于目标形状:

element type = bf16
lane count   = 256
num_groups   = 8
group size   = 32
lanes/part   = 128
source stride = 1

期望 layout:

values     = deinterleaved = 2
inverse    = deinterleaved = 2
quantized  = deinterleaved = 2

对应 lowering 期望:

; one E2B packet, shared by both d2 parts
%inv = pto.vlds %sf[...] {dist = "E2B_B16"}
  : !pto.ptr<bf16, ub> -> !pto.vreg<128xbf16>

; values load in d2
%x0, %x1 = pto.vldsx2 %x[...] {dist = "DINTLV_B16"}
  : !pto.ptr<bf16, ub>, index
  -> !pto.vreg<128xbf16>, !pto.vreg<128xbf16>

%q0 = pto.vmul %x0, %inv
%q1 = pto.vmul %x1, %inv

而不是:

values contiguous
+ inverse contiguous
+ vintlv inverse

影响

当前这个问题会导致:

  • inverse 多一次 vintlv;
  • extf 前多 vintlv + vbitcast 的 part 重组;
  • E4M3 输出最终走 4 个 vcvt(part=P0) 加 4 个 PK4_B32 store。

典型慢 case:

bf16 -> e4m3, h=16384, tokens=512, packed SF, 非 TMA
  ASC/VMI ≈ 0.966

bf16 -> e4m3, h=3072, tokens=8001, packed SF, 非 TMA
  ASC/VMI ≈ 0.972

e2m1 -> e4m3, h=65536, tokens=512, packed SF, TMA
  ASC/VMI ≈ 0.912

其中 e2m1 -> e4m3 还叠加了 dequantize 的转换成本,但 bf16 -> e4m3 没有 input SF dequant,仍然回退,说明 d2 layout 传播是主要问题。

建议

  1. 在 VMILayoutAssignment 中增加该 256xbf16 group-broadcast + vmul 模式的 d2 candidate。
  2. cost model 同时评估:
    • contiguous + inverse vintlv + downstream part regrouping;
    • d2 + E2B shared inverse。
  3. 如果 d2 candidate 被选中,让 group_broadcast_load、vload、vmul 在同一 d2 layout 上传播,避免把 inverse 扩成两个独立 part。
  4. 后续 E4M3 输出路径再基于 d2 做 part-wise conversion,尽量避免先转回 contiguous 再重组。

当前 PTOAS 中已有一个 256xbf16 group broadcast 的临时优化:用 E2B + VINTLV 得到一个 logical 256 结果。这个优化解决了正确性和部分性能问题,但本质上仍然是“contiguous + inverse 扩宽”。如果 layout assignment 能直接选 d2,这个路径应该可以进一步简化为一个共享的 E2B packet。

Activity

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

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions