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)