Repository navigation
Conversation
8c28598 to
9e6a3dc
Compare
|
Currently supported KernelBench kernels. The new pipeline produces identical xegpu-wg level IR as the existing mlp/reduction/fused_attention schedules.
|
9e6a3dc to
2431d89
Compare
Interfaces must be registered for every new context, the python class bool _interfaces_attached persists across contexts.
2431d89 to
f634cea
Compare
adam-smnk
left a comment
There was a problem hiding this comment.
Could you also add more granular tests per new transform op?
Probably not all really need it but at least the infer ones.
| op_name = op.operation.name | ||
| wg_tile = None | ||
| k_tile = None | ||
| if op_name == "linalg.matmul": |
There was a problem hiding this comment.
I assume it only supports plain variant without broadcasts, transposes etc.
I'd be good to at least add assert to document these.
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
added assertion and tests
812c343 to
c0327ac
Compare
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Unresolved parameter-inference and lowering failures prevent reliable execution of the new pipeline.
Review effort: Balanced
Findings: 3
Open (9)
GPU pipeline drops the requested payload function name · New Empty GPU pipeline crashes before returning the module · New Reduction-only payloads fail epilogue extraction · New Transpose slice shape recovery returns incorrect dimensions · New Untiled matmul shape analysis ignores indexing maps · New Parameter inference rejects nested workgroup and K loops · New Payload lookup misses operations normalized by cleanup · New Unsupported multi-reduction-to-contract patterns remain enabled · New Fast path ignores forced subgroup tile sizes · New
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) |
| 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) |
| 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) |
| 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 | ||
| ): |
| """ | ||
| 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 |
There was a problem hiding this comment.
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. |
There was a problem hiding this comment.
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( |
There was a problem hiding this comment.
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") |
There was a problem hiding this comment.
I wonder if we can avoid this pass and use your upstreamed change for "unused linalg op results (if any)"


Introduces a new XeGPU pipeline that consists of parameter-free sub-schedules.
The full pipeline is defined as
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.pyscript.Some schedules, like
xepgu.wg_tiling(), implement different lowering implementations for gemm, reduction, and attention -like payloads. These implementations are handled by thetransform.alternativesop 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).