Skip to content

Fuse expert gathers, interleaved swiglu, and the expert sum into the fused_reduce kernels - #5348

Merged
pfultz2 merged 60 commits into
developfrom
fuse-gather
Oct 11, 2026
Merged

pfultz2 merged 60 commits into
developfrom
fuse-gather

Conversation

@pfultz2

@pfultz2 pfultz2 commented Oct 2, 2026

Copy link
Copy Markdown
Collaborator

Motivation

Compiling a gpt-oss style int4 MoE decode model with enable_skinny_dot failed with fused_reduce: packed inputs require vectorization. The swiglu over the interleaved gate/up columns (reshape {N/2, 2} + slice on the last axis) fused as a prologue of the int4 down_proj reduce, leaving stride-2 inputs along the vector axis that the packed-input vectorizer refused. Gating the fusion would have kept the model compiling but left the decode path slow: each MoE layer ran 13 kernels, 8 of which were plain gather copies of expert weights (~50MB per layer).

This PR keeps the fusion and makes the kernel read those layouts, then fuses the rest of the MoE combine so a decode layer runs 4 kernels (router, convert, gate_up, down_proj incl. bias and expert sums). Driver time per MoE layer goes from 0.565ms to 0.114–0.127ms on gfx1201, and the GPU result verifies against the reference target.

Technical Details

Strided vector loads (strided_vec<T, N, S>) — kernels/vec.hpp, nontemporal.hpp, vectorize.hpp. An input with stride S along the vector axis is viewed as blocks of N*S elements. load_strided loads the S-aligned block around the (possibly misaligned) element pointer and selects lanes phase + i*S, so loads stay wide and aligned and never leave the buffer. Only the fused_reduce plan enables it (vectorize<N, Axis, true>, gen::vectorize::elements(..., strided=true)), gated to strides 2–4 with a power-of-two byte block; pointwise kernels index x[i] directly and keep the plain path. reducer_base::make_inner_slice converts via load_type.

Gather fusion (find_gather_reduce, gather_view) — a gather along a non-fastest axis reached through reshape/unsqueeze/broadcast/slice views moves into the fused_reduce submodule. The data is replayed through the view chain with the gather axis at the data length, and the 1-D indices become an input. The kernel wraps the data in a gather_view (gathered lens, data strides) via the gather_arg<Axis, Data, Indices>() transform; the index is resolved once per output slice, or per element when the gather axis is reduced (make_slice_view overload). Indices are clamped so compile-time benchmark runs on arbitrary buffers stay in bounds. The JIT plans on the gathered logical shape and tracks the gather axis through reduce_dims/normalize_permutation with a marker shape; tiling or vectorizing along the gather axis is refused. Because a reduce with fused gathers can no longer be remapped by the reshape/broadcast rewrites, gathers fuse only in the last fuse_reduce run (fuse_reduce::enable_gather), after the fixed-point fusion loop.

Interleaved slice split (find_reduce_slice) — slices that cut one part of a reduce output axis a reshape split in two (reduce {.., 5760} -> reshape {.., 2880, 2} -> slice axis 3) now split that axis on every input and in the submodule (split_reduce_axis, split_axis_op, generalized insert_split_axis(..., inner)), then slice in the same apply so both halves are visible before the pointwise is visited. The swiglu then fuses as the epilogue of the gate_up reduce and down_proj reads a contiguous activation.

Expert-sum fold (find_reduce_affine_reduce, rewrite_reduce) — reduce_sum[E](mul(add(reduce_sum[K](x), b), w)) is rewritten as reduce_sum[K∪E](x * bcast(w)) + reduce_sum[E](reshape(w*b)) with scale/shift derived symbolically (not f(1)-f(0), which loses precision in fp16). Both reductions are emitted in the inner layout so find_reduce_reduce merges them (merge_reduce_axes now allows a reduce over a subset of axes when its inputs are unit along the rest). prepare_reduce::fuse_reductions compares operator and input lens so a vec32 and a scalar reduction are not paired, and split_reduce skips modules whose reductions differ.

Tests — kernel tests test/gpu/kernels/strided_vec.cpp, gather_view.cpp; pass tests in test/fuse_reduce.cpp (gather_*, reduce_slice_interleaved, reduce_reduce_subset_axes, reduce_reduce_different_axes), test/rewrite_reduce.cpp (reduce_affine_reduce[_unfused]), test/gpu/prepare_reduce.cpp, test/gpu/compile_gen.cpp (vectorize_strided); verify tests test_gather_reduce_moe, test_gather_unpack_int4_dequant_reduce, test_gather_reduce_expert_sum, test_unpack_int4_dequant_reduce_interleaved_swiglu, test_unpack_int4_dequant_reduce_interleaved_input.

Changelog Category

Add a CHANGELOG.md entry for any option other than Not Applicable

    • Added: New functionality.
    • Changed: Changes to existing functionality.
    • Removed: Functionality or support that has been removed. (Compared to a previous release)
    • Optimized: Component performance that has been optimized or improved.
    • Resolved Issues: Known issues from a previous version that have been resolved.
    • Not Applicable: This PR is not to be included in the changelog.

Follow the LLVM AI Tool Use Policy for contributions using AI.

pfultz2 and others added 30 commits September 2, 2026 10:51
A packed input holds two elements per byte, so a full 16-byte load needs
a vector of 32 logical elements. Offer sizes up to 32 for packed inputs
instead of the occupancy-based sizes used for unpacked inputs.

Moved from opt-int4-2 (Optimize int4 further), where it was written on
top of develop's fixed vector sizes; adapted to this branch's
occupancy-based selection for unpacked inputs.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_0128FE5m6fKHLoApAk32eyy1
Update unary condition to check for one input and one output.

Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>
Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>
@pfultz2
pfultz2 added this pull request to stack #5323 October 2, 2026 16:12
@pfultz2
pfultz2 requested a review from causten as a code owner October 2, 2026 16:12
Copilot AI balanced review requested due to automatic review settings October 5, 2026 16:37

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot review overview

🟡 Changes recommended

Unresolved memory-safety, numerical-correctness, and compilation failures block approval.

Review effort: Balanced
Findings: 7 High severity

Open (7)
What changed in this PR

Extends MIGraphX reduction fusion to reduce kernel launches and intermediate copies in int4 mixture-of-experts decode workloads.

Changes:

  • Adds strided vector loads and gathered tensor views.
  • Fuses interleaved SwiGLU slices and weighted expert sums into reductions.
  • Adds kernel, optimization-pass, and numerical verification tests.
File Description
test/​verify/​test_unpack_int4_dequant_reduce_interleaved_swiglu.cpp Verifies interleaved SwiGLU between matvecs.
test/​verify/​test_unpack_int4_dequant_reduce_interleaved_input.cpp Verifies stride-two activation inputs.
test/​verify/​test_gather_unpack_int4_dequant_reduce.cpp Verifies gathered int4 weights and scales.
test/​verify/​test_gather_reduce_moe.cpp Verifies expert weight and bias gathers.
test/​verify/​test_gather_reduce_expert_sum.cpp Verifies weighted expert combination.
test/​rewrite_reduce.cpp Tests affine reduction folding.
test/​gpu/​prepare_reduce.cpp Tests incompatible reduction lengths.
test/​gpu/​kernels/​strided_vec.cpp Tests strided loads and vector traits.
test/​gpu/​kernels/​gather_view.cpp Tests gathered indexing and vectorization.
test/​gpu/​compile_gen.cpp Tests strided vectorization selection.
test/​fuse_reduce.cpp Tests gather, slice, and subset-axis fusion.
src/​targets/​gpu/​prepare_reduce.cpp Tightens parallel-reduction compatibility.
src/​targets/​gpu/​kernels/​include/​migraphx/​kernels/​vectorize.hpp Adds strided tensor vectorization.
src/​targets/​gpu/​kernels/​include/​migraphx/​kernels/​vec.hpp Defines strided vectors and load types.
src/​targets/​gpu/​kernels/​include/​migraphx/​kernels/​reduce.hpp Loads strided reduction inputs.
src/​targets/​gpu/​kernels/​include/​migraphx/​kernels/​nontemporal.hpp Implements strided block loads.
src/​targets/​gpu/​kernels/​include/​migraphx/​kernels/​gather_view.hpp Adds gathered tensor views.
src/​targets/​gpu/​jit/​reduce.cpp Plans gathered inputs and kernel transforms.
src/​targets/​gpu/​include/​migraphx/​gpu/​compile_gen.hpp Exposes strided vectorization options.
src/​targets/​gpu/​compile_gen.cpp Updates vectorization and reduction generation.
src/​split_reduce.cpp Excludes unsupported reduction splitting.
src/​rewrite_reduce.cpp Adds affine expert-sum folding.
src/​include/​migraphx/​fuse_reduce.hpp Adds gather-fusion control.
src/​fuse_reduce.cpp Implements gather and interleaved-slice fusion.
src/​fuse_pointwise_reduce.cpp Defers gather fusion until the final pass.

💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread src/rewrite_reduce.cpp
return nullopt;
result.push_back(*it);
}
return result;
Comment thread src/rewrite_reduce.cpp
Comment on lines +330 to +331
auto [ops, outer] = *chain;
auto inner_axes = reduce_axes(inner);

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It doesnt accept dynamic shapes.

Comment thread src/rewrite_reduce.cpp
Comment on lines +332 to +333
auto outer_elements =
outer->inputs().front()->get_shape().elements() / outer->get_shape().elements();

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We shouldn't have zero element tensors here.

Comment thread src/rewrite_reduce.cpp
outer, make_op("reshape", {{"dims", inner->get_shape().lens()}}), *a.scale);
scale = m.insert_instruction(
outer, make_op("multibroadcast", {{"out_lens", x->get_shape().lens()}}), scale);
x = m.insert_instruction(outer, make_op("mul"), x, scale);
Comment on lines +655 to +656
map_ins[data] = result.add_parameter(name, gathered);
map_ins[ins] = map_ins[data];
Comment on lines +1095 to +1099
// A tile along the gather axis would read adjacent data rows instead
// of the gathered ones, so a cached solution tiling it falls back
if(contains({"block_tile", "block_batch"}, algo) and
plan.is_gather_axis(v.at("tile_axis").to<std::size_t>()))
algo = "block";
Comment on lines +100 to +108
index_int phase = block_phase<S>(elements);
const auto* block = as_vec<N * S>(elements - phase);
MIGRAPHX_ASSERT(bit_cast<uintptr_t>(block) % alignof(block_type) == 0);
block_type v;

if constexpr(Stream)
v = nontemporal_load(block);
else
v = *block;
Base automatically changed from fuse-topk3 to develop October 7, 2026 17:38

@bdevorem bdevorem left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

took a while to get through this, but it looks good to me

@codecov

codecov Bot commented Oct 11, 2026

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 87.45098% with 64 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/fuse_reduce.cpp 89.39% 40 Missing ⚠️
src/rewrite_reduce.cpp 83.20% 21 Missing ⚠️
src/fuse_pointwise_reduce.cpp 0.00% 2 Missing ⚠️
src/split_reduce.cpp 83.33% 1 Missing ⚠️
Additional details and impacted files
@@             Coverage Diff             @@
##           develop    #5348      +/-   ##
===========================================
- Coverage    92.70%   92.64%   -0.07%     
===========================================
  Files          635      635              
  Lines        36204    36613     +409     
===========================================
+ Hits         33562    33917     +355     
- Misses        2642     2696      +54     
Files with missing lines Coverage Δ
src/include/migraphx/fuse_reduce.hpp 100.00% <ø> (ø)
src/include/migraphx/rewrite_reduce.hpp 100.00% <ø> (ø)
src/split_reduce.cpp 93.81% <83.33%> (-0.69%) ⬇️
src/fuse_pointwise_reduce.cpp 0.00% <0.00%> (ø)
src/rewrite_reduce.cpp 93.05% <83.20%> (-4.21%) ⬇️
src/fuse_reduce.cpp 91.70% <89.39%> (-1.38%) ⬇️
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.
  • 📦 JS Bundle Analysis: Save yourself from yourself by tracking and limiting bundle sizes in JS merges.

@pfultz2
pfultz2 merged commit 8510d02 into develop Oct 11, 2026
33 of 35 checks passed
@pfultz2
pfultz2 deleted the fuse-gather branch October 11, 2026 14:55
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants