Skip to content

[xegpu] Add modular, parameter-free pipeline - #293

Open
tkarna wants to merge 31 commits into
llvm:mainfrom
tkarna:xegpu-modular-pipeline
Open

tkarna wants to merge 31 commits into
llvm:mainfrom
tkarna:xegpu-modular-pipeline

Conversation

@tkarna

@tkarna tkarna commented Oct 1, 2026 •

Copy link
Copy Markdown
Contributor

Introduces a new XeGPU pipeline that consists of parameter-free sub-schedules.

The full pipeline is defined as

schedules = [
    xegpu.cleanup_schedule(),
    xegpu.wg_tiling_schedule(),
    xegpu.vectorize_schedule(),
    xegpu.bufferize_schedule(),
    xegpu.outline_gpu_func_schedule(),
    xegpu.vector_to_xegpu_schedule(),
    xegpu.annotate_layouts_schedule(),
    xegpu.xegpu_to_binary_schedule(),
]
  • Each lowering stage has been defined as a standalone sub-schedule.
  • The schedules do not take any parameters by default. Instead, the schedules analyze the IR at lowering time and infer required tile sizes etc.
  • Consequently, the tile selection logic is now defined using the transform dialect; In the schedule, parameters, like tile sizes, are no longer concrete Python int values, but transform dialect parameters. They do not get concrete values until the transform schedule is invoked, i.e. at lowering time.

The long-term goal is that we could replace all the existing xegpu schedules, e.g., mlp_schedule, reduction_schedule, fused_attetnion_schedule, with a single pipeline that can consume any kind of input IR.

This PR is just a first step toward that direction, and still in progress. Currently the pipeline supports basic matmul, reduction and attention payloads (from the KernelBench suite). The matmul support is most mature - it implements the existing GEMM tile size selection mechanism completely. At the moment, the new pipeline is only used in kernel_bench.py script.

Some schedules, like xepgu.wg_tiling(), implement different lowering implementations for gemm, reduction, and attention -like payloads. These implementations are handled by the transform.alternatives op that tries to apply each defined "alternative" block until one succeeds.

As the parameters are not concrete Python values, all calculations that involve parameters must be done inside a dedicated transform op (see e.g., compute_sg_layout.py), or via the transform dialect SMT extension (see annotate_layouts schedule for example).

@tkarna tkarna changed the title [xegpu] Add pipeline [xegpu] Add modular, parameter-free pipeline Oct 1, 2026
@tkarna
tkarna force-pushed the xegpu-modular-pipeline branch from 8c28598 to 9e6a3dc Compare October 2, 2026 12:46
@tkarna

tkarna commented Oct 2, 2026 •

Copy link
Copy Markdown
Contributor Author

Currently supported KernelBench kernels. The new pipeline produces identical xegpu-wg level IR as the existing mlp/reduction/fused_attention schedules.

Type Level Benchmark Status Parameter selection
GEMM 1 1_Square_matrix_multiplication OK Automatic
GEMM 1 2_Standard_matrix_multiplication OK Automatic
GEMM 1 6_Matmul_with_large_K_dimension OK Automatic
GEMM 1 7_Matmul_with_small_K_dimension OK Automatic
GEMM 1 9_Tall_skinny_matrix_multiplication OK Automatic
GEMM 1 13_Matmul_for_symmetric_matrices OK Automatic
GEMM 1 16_Matmul_with_transposed_A OK Automatic
GEMM 1 17_Matmul_with_transposed_B OK Automatic
GEMM 1 18_Matmul_with_transposed_both OK Automatic
GEMM 2 9_Matmul_Subtract_Multiply_ReLU OK Automatic
GEMM 2 12_Gemm_Multiply_LeakyReLU OK Automatic
GEMM 2 29_Matmul_Mish_Mish OK Automatic
GEMM 2 40_Matmul_Scaling_ResidualAdd OK Automatic
GEMM 2 59_Matmul_Swish_Scaling OK Automatic
GEMM 2 63_Gemm_ReLU_Divide OK Automatic
GEMM 2 70_Gemm_Sigmoid_Scaling_ResidualAdd OK Automatic
GEMM 2 81_Gemm_Swish_Divide_Clamp_Tanh_Clamp OK Automatic
MLP 3 1_MLP OK Automatic
MLP 3 2_ShallowWideMLP OK Automatic
MLP 3 3_DeepNarrowMLP OK Automatic
Reduction 1 23_Softmax OK Hard-coded
Reduction 1 24_LogSoftmax OK Hard-coded
Attention 1 97_ScaledDotProductAttention OK Hard-coded

@tkarna
tkarna force-pushed the xegpu-modular-pipeline branch from 9e6a3dc to 2431d89 Compare October 5, 2026 16:47
@tkarna
tkarna marked this pull request as ready for review October 5, 2026 16:51
@tkarna
tkarna requested a review from adam-smnk October 5, 2026 16:51
@tkarna
tkarna force-pushed the xegpu-modular-pipeline branch from 2431d89 to f634cea Compare October 5, 2026 20:13

@adam-smnk adam-smnk 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.

Could you also add more granular tests per new transform op?
Probably not all really need it but at least the infer ones.

Comment thread lighthouse/dialects/transform/transform_ext/ops/compute_num_threads.py Outdated
Comment thread lighthouse/schedule/xegpu/annotate_layouts_schedule.py
Comment thread lighthouse/dialects/transform/transform_ext/utils/matmul_analysis.py Outdated
op_name = op.operation.name
wg_tile = None
k_tile = None
if op_name == "linalg.matmul":

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.

I assume it only supports plain variant without broadcasts, transposes etc.
I'd be good to at least add assert to document these.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

I can add some asserts. But in general I think we should adopt test-driven development philosophy. It's not possible to safe-guard such analysis against all possible input IR variations. The test suite shows which cases are supported and for the rest we assume the method is broken. We'll update the test suite as we improve coverage.

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.

Test coverage is often limited, I wouldn't take them as ground truth.

Also, the question is how quickly you realize that it's broken. Will it blow up instantly, 10 passes later, or give incorrect results?
Specialized transform are fine and bugs are unavoidable but at least sth traceable directly in code will make the made assumptions explicit.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Hmm, actually I do not know what the expected behavior is if the input IR has a transpose in the matmul op's indexing maps, or if an operand is produced by a broadcast op. The correct xegpu metadata depends on what the vector/xegpu dialect IR looks like. I can figure that out by cooking an example, which brings me back to TTD - address it when there's a concrete use case.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

added assertion and tests

Comment thread lighthouse/schedule/xegpu/bufferize.py Outdated
Comment thread lighthouse/schedule/xegpu/wg_tiling.py Outdated
Comment thread lighthouse/schedule/xegpu/wg_tiling.py
@tkarna
tkarna force-pushed the xegpu-modular-pipeline branch from 812c343 to c0327ac Compare October 6, 2026 17:35

Copilot AI left a comment

Copy link
Copy Markdown

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 parameter-inference and lowering failures prevent reliable execution of the new pipeline.

Review effort: Balanced
Findings: 3 High severity · 6 Medium severity

Open (9)
What changed in this PR

Introduces modular XeGPU lowering schedules that select parameters during transformation, moving toward a shared pipeline for matmul, reduction, and attention payloads.

Changes:

  • Adds standalone lowering stages and fallback alternatives.
  • Introduces transform-based parameter inference and layout calculations.
  • Integrates the pipeline into KernelBench and adds matmul-analysis tests.
File Description
test/​transform/​test_matmul_analysis.py Tests shape and tile inference.
lighthouse/​utils/​mlir.py Stops producer traversal at matmuls.
lighthouse/​transform/​alternatives.py Wraps fallback transform regions.
lighthouse/​transform/​__init__.py Exports alternatives wrapper.
lighthouse/​schedule/​xegpu/​xegpu_parameter_selector.py Supports constrained parameter selection.
lighthouse/​schedule/​xegpu/​wg_tiling_schedule.py Adds tiling alternatives.
lighthouse/​schedule/​xegpu/​vectorize_schedule.py Adds standalone vectorization.
lighthouse/​schedule/​xegpu/​vector_to_xegpu_schedule.py Adds vector-to-XeGPU conversion stage.
lighthouse/​schedule/​xegpu/​outline_gpu_func_schedule.py Infers launch threads and outlines kernels.
lighthouse/​schedule/​xegpu/​matmul_costmodel.py Supports fixed tile constraints.
lighthouse/​schedule/​xegpu/​lowering_common.py Generalizes shared lowering helpers.
lighthouse/​schedule/​xegpu/​cleanup_schedule.py Normalizes input operations.
lighthouse/​schedule/​xegpu/​bufferize_schedule.py Adds standalone GPU bufferization.
lighthouse/​schedule/​xegpu/​annotate_layouts_schedule.py Adds layout and prefetch annotation.
lighthouse/​schedule/​xegpu/​__init__.py Exports modular schedules.
lighthouse/​dialects/​transform/​transform_ext/​utils/​xegpu_param_selection.py Connects analysis to parameter selection.
lighthouse/​dialects/​transform/​transform_ext/​utils/​matmul_analysis.py Infers matmul shapes and tiles.
lighthouse/​dialects/​transform/​transform_ext/​ops/​replace_with_fused_attention.py Accepts transform-valued tile sizes.
lighthouse/​dialects/​transform/​transform_ext/​ops/​infer_xegpu_reduction_params.py Adds placeholder reduction parameters.
lighthouse/​dialects/​transform/​transform_ext/​ops/​infer_xegpu_gemm_params.py Infers GEMM parameter dictionaries.
lighthouse/​dialects/​transform/​transform_ext/​ops/​infer_xegpu_attention_params.py Adds placeholder attention parameters.
lighthouse/​dialects/​transform/​transform_ext/​ops/​get_param_dict_entry.py Extracts dictionary parameters.
lighthouse/​dialects/​transform/​transform_ext/​ops/​extract_handle.py Supports recoverable extraction failures.
lighthouse/​dialects/​transform/​transform_ext/​ops/​emit_definite_failure.py Reports exhausted alternatives.
lighthouse/​dialects/​transform/​transform_ext/​ops/​compute_sg_layout.py Computes subgroup layouts.
lighthouse/​dialects/​transform/​transform_ext/​ops/​compute_num_threads.py Computes launch thread counts.
lighthouse/​dialects/​transform/​transform_ext/​__init__.py Exports new transform operations.
lighthouse/​dialects/​transform/​smt_ext/​ops/​constrain_params.py Attaches interfaces per context.
examples/​xegpu/​llama3_schedule.py Updates helper call to keyword syntax.
examples/​xegpu/​kernel_bench.py Adds optional modular pipeline execution.

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

stop_at_stage=stop_at_stage,
)
if use_gpu_pipeline:
pipeline = get_xegpu_pipeline(stop_at_stage)
stop_at_stage=stop_at_stage,
)
if use_gpu_pipeline:
pipeline = get_xegpu_pipeline(stop_at_stage)
wg_loop = match_and_split(func, ops={"scf.forall"}, nhandles=1)[0]
generic_ops = match(wg_loop, ops={"linalg.generic"})
elemwise_ops = transform_ext.filter_elementwise(generic_ops)
leaf_elemwise = transform_ext.extract_handle(elemwise_ops, -1, silenceable=True)
Comment on lines +149 to +151
slice_op = _first_producer_named(value, "tensor.extract_slice")
source = slice_op.operands[0] if slice_op is not None else value
return list(ir.ShapedType(source.type).shape)
# TODO use op name as xegpu dialect python bindings are missing
op_name = op.operation.name
if op_name == "linalg.matmul":
return _linalg_matmul_shape_and_transpose(op)
Comment on lines +225 to +227
if op.parent.name != "scf.forall":
# target is not within a scf.forall loop, so it's not workgroup tiled
return None, None
"""Normalize singleton dimensions and fuse elementwise ops in the payload."""

with schedule_boilerplate() as (schedule, named_seq):
op_names = ["linalg.generic", "linalg.matmul"]
op_name="builtin.module",
deduplicate=True,
)
lowering_common.vectorize(payload_mod, payload_func=func)
Comment on lines +94 to +96
if (fixed_wg_tile is not None and wg_tile != fixed_wg_tile) or (
fixed_k_tile is not None and params["k_tile"] != fixed_k_tile
):

@charithaintc charithaintc 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.

thanks!
I have some clarification comments.

nit: It would be nice if the new transfrom ext ops show a simple usage of the op in docstring.

"""
Compute the workgroup thread count from wg and sg tile params.

Returns `base * prod_i(wg_i // sg_i)` over the tile dimensions, treating an

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.

good to define what base, wg_i and sg_i specifically refer to here. I believe wg_i and sg_i refer to tile sizes?

TransformExtensionDialect.Operation, name="infer_xegpu_attention_params"
):
"""
Infer XeGPU attention tiling parameters for an attention anchor op.

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.

what is an attention anchor op? Do you mean to say that we match the attention pattern and then use this op to derive tile shapes?

n_ctx = ir.RankedTensorType(k.type).shape[-2]
tile_size = ir.IntegerAttr(op.tile_size).value
tile_size_params = state.get_params(op.tile_size)
if len(tile_size_params) != 1 or not isinstance(

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.

how would this look like when we transition to dependent reduction fusion approach? I believe you will still pass the derived attention tiling params to the sequence of transform ops?

lh_transform.cleanup(func)

# Fuse elementwise ops, also removes unused linalg op results (if any).
func = apply_registered_pass(func, "linalg-fuse-elementwise-ops")

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.

I wonder if we can avoid this pass and use your upstreamed change for "unused linalg op results (if any)"

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.

4 participants