From ff7c84b0804500005c747eee6a81279cbbdcc453 Mon Sep 17 00:00:00 2001 From: Tuomas Karna Date: Thu, 27 Aug 2026 14:41:09 +0300 Subject: [PATCH 1/4] reduction schedule: use upstream folding patterns --- lighthouse/schedule/xegpu/reduction_schedule.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/lighthouse/schedule/xegpu/reduction_schedule.py b/lighthouse/schedule/xegpu/reduction_schedule.py index dc64a110..14347147 100644 --- a/lighthouse/schedule/xegpu/reduction_schedule.py +++ b/lighthouse/schedule/xegpu/reduction_schedule.py @@ -2,7 +2,7 @@ from mlir import ir from mlir.dialects import transform -from mlir.dialects.transform import structured, xegpu +from mlir.dialects.transform import structured, xegpu, tensor import lighthouse.transform as lh_transform from .lowering_common import ( get_payload_func, @@ -114,8 +114,14 @@ def bundle_xegpu_reduction_schedule( # Normalize possible singleton dimensions so tile+fuse logic works. with ir.InsertionPoint(transform.apply_patterns(func).patterns): + # fold unit dims in linalg.generic op inputs structured.apply_patterns_linalg_fold_unit_extent_dims_via_slices() - transform_ext.fold_singleton_extract_slice(func) + # fold tensor.extract_slice(tensor.expand_shape(x)) into x + tensor.apply_patterns_tensor_reassociative_reshape_folding() + # swap tensor.extract_slice(linalg.fill(...)) ops + structured.apply_patterns_linalg_swap_extract_slice_with_fill() + # fold tensor.extract_slice(tensor.empty(...)) into tensor.tensor_empty(...) + tensor.apply_patterns_tensor_fold_tensor_empty(fold_single_use_only=True) lh_transform.cleanup(func) # Fuse elementwise ops, also removes unused linalg op results (if any). From 5f3b916ff9598ad89f896a35b8022c8478ed3d4e Mon Sep 17 00:00:00 2001 From: Tuomas Karna Date: Thu, 27 Aug 2026 14:43:07 +0300 Subject: [PATCH 2/4] transform_ext: remove fold_singleton_extract_slice op --- .../transform/transform_ext/__init__.py | 2 - .../ops/fold_singleton_extract_slice.py | 200 ------------------ 2 files changed, 202 deletions(-) delete mode 100644 lighthouse/dialects/transform/transform_ext/ops/fold_singleton_extract_slice.py diff --git a/lighthouse/dialects/transform/transform_ext/__init__.py b/lighthouse/dialects/transform/transform_ext/__init__.py index 363bd6b4..45aa49d0 100644 --- a/lighthouse/dialects/transform/transform_ext/__init__.py +++ b/lighthouse/dialects/transform/transform_ext/__init__.py @@ -21,7 +21,6 @@ from .ops.filter_reduction_ops import filter_reduction_ops from .ops.get_leading_unit_tile_sizes import get_leading_unit_tile_sizes from .ops.move_offsets_to_subview import move_offsets_to_subview -from .ops.fold_singleton_extract_slice import fold_singleton_extract_slice from .ops.clear_tile_and_fuse_annotations import clear_tile_and_fuse_annotations from .ops.get_fusion_roots import get_fusion_roots from .ops.propagate_tile_sizes import propagate_tile_sizes @@ -36,7 +35,6 @@ "filter_elementwise", "filter_num_loops", "filter_reduction_ops", - "fold_singleton_extract_slice", "get_fusion_roots", "get_leading_unit_tile_sizes", "get_named_attribute", diff --git a/lighthouse/dialects/transform/transform_ext/ops/fold_singleton_extract_slice.py b/lighthouse/dialects/transform/transform_ext/ops/fold_singleton_extract_slice.py deleted file mode 100644 index e211512c..00000000 --- a/lighthouse/dialects/transform/transform_ext/ops/fold_singleton_extract_slice.py +++ /dev/null @@ -1,200 +0,0 @@ -from mlir import ir -from mlir.dialects import ext, transform, tensor, linalg -from mlir.dialects.transform import DiagnosedSilenceableFailure - -from lighthouse.dialects.transform.transform_ext import TransformExtensionDialect - - -def _is_dynamic_dim(dim: int) -> bool: - return dim < 0 - - -def _dim_compatible(a: int, b: int) -> bool: - return a == b or _is_dynamic_dim(a) or _is_dynamic_dim(b) - - -def _same_ranked_tensor_type(lhs: ir.Type, rhs: ir.Type) -> bool: - if not isinstance(lhs, ir.RankedTensorType) or not isinstance( - rhs, ir.RankedTensorType - ): - return False - if lhs.rank != rhs.rank: - return False - if lhs.element_type != rhs.element_type: - return False - return all(_dim_compatible(dl, dr) for dl, dr in zip(lhs.shape, rhs.shape)) - - -def _is_rank_reduction_by_unit_dims( - expanded: ir.RankedTensorType, reduced: ir.RankedTensorType -) -> bool: - if expanded.rank <= reduced.rank: - return False - - reduced_idx = 0 - skipped = 0 - for dim in expanded.shape: - if reduced_idx < reduced.rank and _dim_compatible( - dim, reduced.shape[reduced_idx] - ): - reduced_idx += 1 - continue - if dim == 1: - skipped += 1 - continue - return False - - return reduced_idx == reduced.rank and skipped == (expanded.rank - reduced.rank) - - -def _as_single_result_value(op_or_val): - if isinstance(op_or_val, ir.Value): - return op_or_val - if hasattr(op_or_val, "result"): - return op_or_val.result - if hasattr(op_or_val, "results") and len(op_or_val.results) == 1: - return op_or_val.results[0] - raise ValueError("Expected a value or single-result operation") - - -def _rewrite_extract_of_expand( - extract_op: ir.Operation, rewriter: transform.TransformRewriter -) -> bool: - source = extract_op.operands[0] - producer = source.owner - if producer is None or producer.name != "tensor.expand_shape": - return False - - expanded_source = producer.operands[0] - extract_result_ty = extract_op.results[0].type - expanded_source_ty = expanded_source.type - expanded_ty = source.type - - if not ( - isinstance(extract_result_ty, ir.RankedTensorType) - and isinstance(expanded_source_ty, ir.RankedTensorType) - and isinstance(expanded_ty, ir.RankedTensorType) - ): - return False - - if not _same_ranked_tensor_type(extract_result_ty, expanded_source_ty): - return False - - if not _is_rank_reduction_by_unit_dims(expanded_ty, extract_result_ty): - return False - - rewriter.replace_op(extract_op, [expanded_source]) - return True - - -def _rewrite_extract_of_fill( - extract_op: ir.Operation, rewriter: transform.TransformRewriter -) -> bool: - source = extract_op.operands[0] - producer = source.owner - if producer is None or producer.name != "linalg.fill": - return False - - fill_result_ty = source.type - extract_result_ty = extract_op.results[0].type - if not ( - isinstance(fill_result_ty, ir.RankedTensorType) - and isinstance(extract_result_ty, ir.RankedTensorType) - ): - return False - - if not _is_rank_reduction_by_unit_dims(fill_result_ty, extract_result_ty): - return False - - if any(_is_dynamic_dim(dim) for dim in extract_result_ty.shape): - # Keep conservative behavior for dynamic shapes. - return False - - fill_value = producer.operands[0] - with ir.InsertionPoint(extract_op), extract_op.location: - empty = tensor.EmptyOp( - tuple(extract_result_ty.shape), extract_result_ty.element_type - ) - filled = linalg.fill(fill_value, outs=[empty.result]) - rewriter.replace_op(extract_op, [_as_single_result_value(filled)]) - return True - - -def _collect_extract_slice_ops(root: ir.Operation) -> list[ir.Operation]: - extract_slice_ops = [] - - def collect(op: ir.Operation) -> ir.WalkResult: - if op.name == "tensor.extract_slice": - extract_slice_ops.append(op) - return ir.WalkResult.ADVANCE - - root.walk(collect, ir.WalkOrder.PRE_ORDER) - return extract_slice_ops - - -class FoldSingletonExtractSliceOp( - TransformExtensionDialect.Operation, name="fold_singleton_extract_slice" -): - """ - Rewrites redundant singleton-dimension tensor slice patterns. - - Rewrites: - 1) tensor.extract_slice(tensor.expand_shape(x)) -> x - 2) tensor.extract_slice(linalg.fill(... : tensor<...x1x...>)) - -> linalg.fill on rank-reduced tensor output. - - Args: - target: Handle to root ops to rewrite within (e.g. func.func). - Returns: - Handle to (possibly) rewritten extract_slice ops. - """ - - target: ext.Operand[transform.AnyOpType] - rewritten_ops: ext.Result[transform.AnyOpType[()]] = ext.infer_result() - - @classmethod - def attach_interface_impls(cls, context=None): - cls.TransformOpInterfaceModel.attach(cls.OPERATION_NAME, context=context) - cls.MemoryEffectsOpInterfaceModel.attach(cls.OPERATION_NAME, context=context) - - class TransformOpInterfaceModel(transform.TransformOpInterface): - @staticmethod - def apply( - op: "FoldSingletonExtractSliceOp", - rewriter: transform.TransformRewriter, - results: transform.TransformResults, - state: transform.TransformState, - ) -> DiagnosedSilenceableFailure: - targets = state.get_payload_ops(op.target) - - for target in targets: - extract_ops = _collect_extract_slice_ops(target) - for extract_op in extract_ops: - did_rewrite = _rewrite_extract_of_expand(extract_op, rewriter) - if not did_rewrite: - _rewrite_extract_of_fill(extract_op, rewriter) - - # Return stable handles to the transformed target roots. - results.set_ops(op.rewritten_ops, targets) - return DiagnosedSilenceableFailure.Success - - @staticmethod - def allow_repeated_handle_operands(_op: "FoldSingletonExtractSliceOp") -> bool: - return False - - class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): - @staticmethod - def get_effects(op: ir.Operation): - return ( - transform.only_reads_handle(op.op_operands) - + transform.produces_handle(op.results) - + transform.modifies_payload() - ) - - -def fold_singleton_extract_slice( - target: ir.Value[transform.AnyOpType], -) -> ir.Value[transform.AnyOpType]: - """snake_case wrapper to create FoldSingletonExtractSliceOp.""" - op = FoldSingletonExtractSliceOp(target=target) - return op.rewritten_ops From e670c029b9d288335f63cf4511620915c9ba9f7d Mon Sep 17 00:00:00 2001 From: Tuomas Karna Date: Thu, 27 Aug 2026 14:49:46 +0300 Subject: [PATCH 3/4] update llvm version --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index bf3ce5f5..a8cad4af 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,7 +3,7 @@ name = "lighthouse" dynamic = ["version"] requires-python = ">=3.10,<3.13" # Bounds are due to torch-mlir's packaging dependencies = [ - "mlir-python-bindings==20260820+002905df0", + "mlir-python-bindings==20260826+6c81b990d", "pyyaml>=6.0", ] From dd38f50a222933661a5c1b2108bcf50341a93ea6 Mon Sep 17 00:00:00 2001 From: Tuomas Karna Date: Thu, 27 Aug 2026 21:03:53 +0300 Subject: [PATCH 4/4] fused attention schedule: use upstream folding patterns --- lighthouse/schedule/xegpu/fused_attention_schedule.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/lighthouse/schedule/xegpu/fused_attention_schedule.py b/lighthouse/schedule/xegpu/fused_attention_schedule.py index eec3b083..43ede10d 100644 --- a/lighthouse/schedule/xegpu/fused_attention_schedule.py +++ b/lighthouse/schedule/xegpu/fused_attention_schedule.py @@ -2,7 +2,7 @@ from mlir import ir from mlir.dialects import transform -from mlir.dialects.transform import structured, xegpu +from mlir.dialects.transform import structured, xegpu, tensor import lighthouse.transform as lh_transform from lighthouse.pipeline.helper import ( apply_registered_pass, @@ -152,8 +152,14 @@ def bundle_xegpu_fused_attention_schedule( # Normalize possible singleton dimensions so tile+fuse logic works. with ir.InsertionPoint(transform.apply_patterns(func).patterns): + # fold unit dims in linalg.generic op inputs structured.apply_patterns_linalg_fold_unit_extent_dims_via_slices() - transform_ext.fold_singleton_extract_slice(func) + # fold tensor.extract_slice(tensor.expand_shape(x)) into x + tensor.apply_patterns_tensor_reassociative_reshape_folding() + # swap tensor.extract_slice(linalg.fill(...)) ops + structured.apply_patterns_linalg_swap_extract_slice_with_fill() + # fold tensor.extract_slice(tensor.empty(...)) into tensor.tensor_empty(...) + tensor.apply_patterns_tensor_fold_tensor_empty(fold_single_use_only=True) lh_transform.cleanup(func) # Fuse elementwise ops, also removes unused linalg op results (if any).