From 8466cf5ba8eac3e7952c7dea3b6a65449f1b7621 Mon Sep 17 00:00:00 2001 From: Adam Siemieniuk Date: Thu, 1 Oct 2026 20:08:14 +0200 Subject: [PATCH] [transform] Skip fusion roots with all zero tiles Skips potential fusion root ops when there is no tiling to apply. This prevents hard failure due to assertion in upstream tile and fuse rewrite during attempt to tile with all zero tile sizes. Such tiling configuration can occur when tiling strategy disables unprofitable tiling dimensions which may zero out all tile sizes. Assisted-by: Claude --- .../transform_ext/ops/get_fusion_roots.py | 4 +- test/transform/test_get_fusion_roots.py | 38 ++++++++++++++++ test/transform/test_tile_and_fuse.py | 44 +++++++++++++++++++ 3 files changed, 85 insertions(+), 1 deletion(-) diff --git a/lighthouse/dialects/transform/transform_ext/ops/get_fusion_roots.py b/lighthouse/dialects/transform/transform_ext/ops/get_fusion_roots.py index 6958f95a..2913d928 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/get_fusion_roots.py +++ b/lighthouse/dialects/transform/transform_ext/ops/get_fusion_roots.py @@ -132,7 +132,9 @@ def apply( roots = [] for target_op in target_ops: - if tsa.get_tile_sizes_attr(target_op) is None: + # Skip ops with all-zero tile sizes which raises assertion + # in upstream tile-and-fuse rewrite. + if not any(tsa.get_tile_sizes_attr(target_op) or []): continue if GetFusionRootsOp._is_fusion_root(target_op): roots.append(target_op) diff --git a/test/transform/test_get_fusion_roots.py b/test/transform/test_get_fusion_roots.py index 7b46f7f1..c51ece8a 100644 --- a/test/transform/test_get_fusion_roots.py +++ b/test/transform/test_get_fusion_roots.py @@ -241,3 +241,41 @@ def all_linalg_roots(named_seq): # CHECK-NOT: math.exp # CHECK: math.sqrt apply_schedule(COMPATIBLE_CHAIN, all_linalg_roots, "COMPATIBLE_CHAIN") + + +# An op whose tile sizes are all zero (nothing to tile) is never a root, even +# when it terminates its own group; upstream fusion asserts on such a tiling. +ZERO_TILES = """ +#id = affine_map<(d0, d1) -> (d0, d1)> +module { + func.func @main(%a: tensor<8x8xf32>, %b: tensor<64x64xf32>) + -> (tensor<8x8xf32>, tensor<64x64xf32>) { + %e0 = tensor.empty() : tensor<8x8xf32> + %small = linalg.generic {indexing_maps = [#id, #id], + iterator_types = ["parallel", "parallel"], + transform_ext.tile_sizes = array} + ins(%a : tensor<8x8xf32>) + outs(%e0 : tensor<8x8xf32>) { + ^bb0(%i: f32, %o: f32): + linalg.yield %i : f32 + } -> tensor<8x8xf32> + %e1 = tensor.empty() : tensor<64x64xf32> + %big = linalg.generic {indexing_maps = [#id, #id], + iterator_types = ["parallel", "parallel"], + transform_ext.tile_sizes = array} + ins(%b : tensor<64x64xf32>) + outs(%e1 : tensor<64x64xf32>) { + ^bb0(%i: f32, %o: f32): + linalg.yield %i : f32 + } -> tensor<64x64xf32> + return %small, %big : tensor<8x8xf32>, tensor<64x64xf32> + } +} +""" + + +# Only the 64x64 op is a root; the all-zero 8x8 op (first in program order) is not. +# CHECK: IR printer: ZERO_TILES +# CHECK-NOT: tensor<8x8xf32> +# CHECK: tensor<64x64xf32> +apply_schedule(ZERO_TILES, all_linalg_roots, "ZERO_TILES") diff --git a/test/transform/test_tile_and_fuse.py b/test/transform/test_tile_and_fuse.py index b0c3d4bc..e233d499 100644 --- a/test/transform/test_tile_and_fuse.py +++ b/test/transform/test_tile_and_fuse.py @@ -922,3 +922,47 @@ def tile_and_fuse_for(): # CHECK-NOT: transform_ext # CHECK: scf.yield run("elementwise_scf_for_nested_clears", ELTWISE, assign_elementwise, tile_and_fuse_for) + + +# An op with all-zero tile sizes (every dim below the cache tile) next to a tiled +# one: upstream fusion asserts on an all-zero tiling, so the op is left untiled +# while the other is still tiled. +ZERO_TILES = """ +#id1 = affine_map<(d0) -> (d0)> +#id2 = affine_map<(d0, d1) -> (d0, d1)> +module { + func.func @main(%a: tensor<16xf32>, %b: tensor<64x64xf32>) + -> (tensor<16xf32>, tensor<64x64xf32>) { + %e0 = tensor.empty() : tensor<16xf32> + %small = linalg.generic {indexing_maps = [#id1, #id1], + iterator_types = ["parallel"], + transform_ext.tile_sizes = array} + ins(%a : tensor<16xf32>) + outs(%e0 : tensor<16xf32>) { + ^bb0(%i: f32, %o: f32): + %e = math.exp %i : f32 + linalg.yield %e : f32 + } -> tensor<16xf32> + %e1 = tensor.empty() : tensor<64x64xf32> + %big = linalg.generic {indexing_maps = [#id2, #id2], + iterator_types = ["parallel", "parallel"], + transform_ext.tile_sizes = array} + ins(%b : tensor<64x64xf32>) + outs(%e1 : tensor<64x64xf32>) { + ^bb0(%i: f32, %o: f32): + %s = math.sqrt %i : f32 + linalg.yield %s : f32 + } -> tensor<64x64xf32> + return %small, %big : tensor<16xf32>, tensor<64x64xf32> + } +} +""" + + +# CHECK-LABEL: Test: all_zero_tile_sizes_untiled +# CHECK-NOT: scf.forall +# CHECK: math.exp +# CHECK: -> tensor<16xf32> +# CHECK: scf.forall ({{.*}}) = (0, 0) to (64, 64) step (32, 32) +# CHECK: math.sqrt +run("all_zero_tile_sizes_untiled", ZERO_TILES, tile_and_fuse)