Skip to content

[VMI] group_reduce_maxf group=4 needs consumer-aware lowering: packed vector vs direct point store #1559

Description

@Zhendong404

Summary

pto.vmi.group_reduce_maxf with num_groups = 4 currently has a single generic lowering for contiguous f32 input. That lowering materializes the four reduced scalars into a packed contiguous vector using:

4 x pto.vcmax
4 x pto.vdup {position = "LOWEST"}
3 x pto.vsel
+ several mask operations

This is a good lowering when the reduce result is reused as a dense/broadcast vector, but it is a poor lowering when the result is immediately written as four scale values. The same VMI operation has opposite optimal lowerings depending on its consumer.

History

The contiguous f32 group_reduce_maxf group=4 path was introduced by:

  • 27d93e322 fix(vmi): keep group4 reduce and fp8 stores contiguous

and later shortened by:

  • 836fd277c perf(vmi): shorten contiguous group4 max packing

Both changes were motivated by per-block / per-channel kernels, where the reduce is hoisted outside the row loop and the packed result is reused.

Scenario A: per-block / per-channel

Typical shape:

for each 128-lane column chunk:
    amax128 = max over many rows
    amax4 = vcmax(amax128, group=4)
    inverse_vec = vbrc(amax4, size=128, group=4)
    use inverse_vec for many row-wise quant operations

Here one amax4 packing serves many rows. The vdup + vsel packing is amortized and is useful because a dense vector is immediately needed by vbrc / subsequent vector operations.

Scenario B: per-token

Typical shape:

for row:
    for pair in K / 128:
        amax4 = vcmax(vabs(vload(row, pair)), group=4)
        vstore(amax4, amax_ub[row, pair * 4], mask4)

Minimal MLIR pattern:

%x = pto.vmi.vload %src[%off]
    : !pto.ptr<f32, ub> -> !pto.vmi.vreg<128xf32>
%abs = pto.vmi.vabs %x
    : !pto.vmi.vreg<128xf32> -> !pto.vmi.vreg<128xf32>
%mask = pto.vmi.create_mask %c128 : index -> !pto.vmi.mask<128xb32>
%amax = pto.vmi.vcmax %abs, %mask {group = 4}
    : !pto.vmi.vreg<128xf32>, !pto.vmi.mask<128xb32>
   -> !pto.vmi.vreg<4xf32>
pto.vmi.vstore %amax, %scale_out[%off], %c1 {group = 4}
    : !pto.vmi.vreg<4xf32>, !pto.ptr<f32, ub>

Here the reduce is in the innermost row x pair loop. The packed vector is only used by a 4-lane store. The 4 x vdup + 3 x vsel + mask work has no reuse and becomes pure overhead.

Observed performance impact

On the per-token cast kernel, non-packed bf16 -> e4m3, round_sf=True path:

tokens=8001, hidden=3072:
    ASC: 38.55 us
    VMI: 82.53 us

tokens=8001, hidden=16384:
    ASC: 187.71 us
    VMI: 423.12 us

tokens=8001, hidden=65536:
    ASC: 749.83 us
    VMI: 1689.91 us

ASC implements this as four contiguous 64-lane masked vcmax operations followed by four scalar point stores. It does not build a packed 4-lane f32 vector.

Why the current lowering is chosen

The current lowering is correct and efficient for Scenario A, but VMIToVPTO selects it unconditionally for contiguous f32 group=4 reduction. It does not inspect whether the result has a dense/broadcast consumer or only a small uniform store consumer.

Proposed fix

Make the lowering consumer-aware, or expose the choice explicitly.

Option 1: fuse reduce + store in VMIToVPTO.

If the group_reduce_maxf group=4 result has a single pto.vmi.group_store user whose result is a 4-lane contiguous store (or a direct point-store pattern), lower directly to:

4 x row_reduce / vcmax
4 x point store / vsts_1st

and skip the packed-vector vdup + vsel phase.

Option 2: allow the result to remain in a group-slot layout.

For the direct-store case, preserve the four group slots and let group_store emit slot-wise point stores, instead of forcing a contiguous 4-lane vector.

Option 3: add an explicit attribute / dedicated VMI form.

For example, distinguish:

%amax = pto.vmi.vcmax ... {group = 4}

from a form that requests direct slot-store lowering. This avoids relying only on consumer pattern matching.

Expected outcome

  • Keep the existing packed lowering for per-block / per-channel, where it is beneficial.
  • Add a direct reduce-store lowering for per-token, where packing is hot-loop overhead.
  • The same VMI expression should then be able to lower efficiently in both contexts.

Additional context

I also tried two alternative VMI expressions for Scenario B:

  • Splitting the 128-lane input into two 64-lane halves triggers VMI-RESIDUAL-OP in the current pipeline.
  • A 256-lane group=8 expression is much faster in an experiment but is not value-correct with the current layout lowering.

Neither should be required if the 128-lane group=4 reduce can select the direct store lowering when its consumer is a small store.

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