Skip to content

[VMI] 优化 256xbf16 grouped broadcast load:E2B_B16 + VINTLV #1565

Description

@Zhendong404

背景

在 per_token_cast 的 bf16 量化路径中,需要把 8 个 inverse scale 展开成 256-lane 逻辑向量:

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)

语义上期望得到:

inverse = [s0 x32, s1 x32, s2 x32, ..., s7 x32]

当前 lowering 对 group_broadcast_load 的 256xbf16, num_groups=8, group_size=32 形状没有高效路径。

当前问题

1. fallback 使用未对齐的 vsldb

当前 fallback 会把逻辑 256-lane 结果拆成两个 128-lane physical part,并使用:

%result = pto.vsldb %ptr, %zero_i16, %zero_i16, %mask
  : !pto.ptr<bf16, ub>, i16, i16, !pto.vmask -> !pto.vreg<128xbf16>
%lo = pto.vselr %result, %idx_lo
%hi = pto.vselr %result, %idx_hi

vsldb 是 32B block load,要求 base 32B 对齐;但这里 source offset 是:

row * 64 + tile * 8   (bf16 elements)

tile * 8 是 16B 步进,奇数 tile 的地址不满足 32B 对齐,导致读取到错误的 scale 数据。

2. 走 vgather2 可以避错,但性能明显回退

为了避开 vsldb 对齐问题,当前另一条路径会拆成两条 vgather2:

%v0 = pto.vgather2 %src, %idx0, %mask
%v1 = pto.vgather2 %src, %idx1, %mask

该路径精度正确,但比 ASC 的 E2B broadcast 路径慢。代表性案例:

bf16 -> E4M3, h=3072, tokens=512, packed SF, TMA
ASC: 4.429 us
VMI: 4.698 us
ASC/VMI ≈ 0.943

期望优化

对于以下形状:

result element type = bf16
result lane count   = 256
num_groups          = 8
group_size          = 32
lanes_per_part      = 128
layout              = contiguous, lane_stride = 1
source_group_stride = 1

期望 lowering 为:

%packet = pto.vlds %src[%offset] {dist = "E2B_B16"}
  : !pto.ptr<bf16, ub> -> !pto.vreg<128xbf16>

%lo, %hi = pto.vintlv %packet, %packet
  : !pto.vreg<128xbf16>, !pto.vreg<128xbf16>
  -> !pto.vreg<128xbf16>, !pto.vreg<128xbf16>

即:

1 x E2B_B16 (单物理 part)
+ 1 x VINTLV

这里:

  • E2B_B16 将一个 128-lane packet 中的 8 个 group slot 广播到 16 lane/group。
  • VINTLV 将每个 group 扩展到 32 lane/group,并拆成两个 128-lane half:
    • lo = [s0 x32, s1 x32, s2 x32, s3 x32]
    • hi = [s4 x32, s5 x32, s6 x32, s7 x32]

结果正好等价于逻辑上的:

[s0 x32, s1 x32, ..., s7 x32]

已做的尝试

在 OneToNVMIGroupBroadcastLoadOpPattern 中增加固定形状 match:

  • 条件:256xbf16, num_groups=8, group_size=32, lanes_per_part=128, contiguous, stride=1
  • Lower:
    • 创建一次 VldsOp(dist="E2B_B16")
    • 创建一次 VintlvOp(packet, packet)
    • 将 low/high 作为两个 128-lane physical result parts

验证结果:

bf16 -> E4M3 TMA: out_equal=True, sf_equal=True
bf16 -> FP4  TMA: out_equal=True, sf_equal=True

bf16 -> E4M3 TMA:
  ASC: 4.429 us
  VMI: 4.532 us
  ASC/VMI ≈ 0.9774

bf16 -> FP4 TMA:
  ASC: 4.008 us
  VMI: 3.943 us
  ASC/VMI ≈ 1.0164

希望 PTOAS 处理的内容

建议把该固定形状的优化正式收录到 PTOAS 的 group_broadcast_load lowering 中,最好是通用化:

  • 256-lane 逻辑向量 broadcast -> 一个 128-lane E2B packet + vintlv 展开;
  • 避免 fallback 使用未对齐的 vsldb;
  • 避免拆成两条 vgather2。

如果该 pattern 只适用于 bf16/16-bit,也可以先做精确匹配,后续再推广到其他 16-bit 类型。

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