Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
38 changes: 38 additions & 0 deletions test/transform/test_get_fusion_roots.py
Original file line number Diff line number Diff line change
Expand Up @@ -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<i64: 0, 0>}
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<i64: 32, 32>}
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")
44 changes: 44 additions & 0 deletions test/transform/test_tile_and_fuse.py
Original file line number Diff line number Diff line change
Expand Up @@ -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<i64: 0>}
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<i64: 32, 32>}
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)
Loading