Repository navigation
Fuse expert gathers, interleaved swiglu, and the expert sum into the fused_reduce kernels - #5348
Merged
Merged
Conversation
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
added this pull request to stack #5323
October 2, 2026 16:12
Contributor
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Unresolved memory-safety, numerical-correctness, and compilation failures block approval.
Review effort: Balanced
Findings: 7
Open (7)
Empty outer axes produce an invalid reduction · New Dynamic reduction chains crash the rewrite matcher · New Zero-element outputs cause division by zero · New Distributed scaling causes fp16 overflow and NaN · New Conflicting shared-data gathers cannot be fused · New Gathered block_batch tiles use an invalid axis · New Unsafe wide loads for strided and offset views · New
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.
| return nullopt; | ||
| result.push_back(*it); | ||
| } | ||
| return result; |
Comment on lines
+330
to
+331
| auto [ops, outer] = *chain; | ||
| auto inner_axes = reduce_axes(inner); |
Collaborator
Author
There was a problem hiding this comment.
It doesnt accept dynamic shapes.
Comment on lines
+332
to
+333
| auto outer_elements = | ||
| outer->inputs().front()->get_shape().elements() / outer->get_shape().elements(); |
Collaborator
Author
There was a problem hiding this comment.
We shouldn't have zero element tensors here.
| 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; |
bdevorem
approved these changes
Oct 9, 2026
bdevorem
left a comment
Member
There was a problem hiding this comment.
took a while to get through this, but it looks good to me
Codecov Report❌ Patch coverage is
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
🚀 New features to boost your workflow:
|
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.

Motivation
Compiling a gpt-oss style int4 MoE decode model with
enable_skinny_dotfailed withfused_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_stridedloads the S-aligned block around the (possibly misaligned) element pointer and selects lanesphase + 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 indexx[i]directly and keep the plain path.reducer_base::make_inner_sliceconverts viaload_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 agather_view(gathered lens, data strides) via thegather_arg<Axis, Data, Indices>()transform; the index is resolved once per output slice, or per element when the gather axis is reduced (make_slice_viewoverload). 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 throughreduce_dims/normalize_permutationwith 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 lastfuse_reducerun (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, generalizedinsert_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 asreduce_sum[K∪E](x * bcast(w)) + reduce_sum[E](reshape(w*b))with scale/shift derived symbolically (notf(1)-f(0), which loses precision in fp16). Both reductions are emitted in the inner layout sofind_reduce_reducemerges them (merge_reduce_axesnow allows a reduce over a subset of axes when its inputs are unit along the rest).prepare_reduce::fuse_reductionscompares operator and input lens so a vec32 and a scalar reduction are not paired, andsplit_reduceskips modules whose reductions differ.Tests — kernel tests
test/gpu/kernels/strided_vec.cpp,gather_view.cpp; pass tests intest/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 teststest_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.mdentry for any option other thanNot ApplicableFollow the LLVM AI Tool Use Policy for contributions using AI.