1111 IntegerType ,
1212 MemRefType ,
1313 Module ,
14- OpResult ,
1514 OpView ,
1615 Value ,
1716)
1817
19- from synapse .language .types import DType , TensorType
18+ from synapse .language .types import BufferType , DType
2019from synapse .patterns import TileArrayProgramPattern
2120
2221
@@ -74,10 +73,9 @@ def replace_with_tile_array_program(
7473 ) -> bool :
7574 """Replaces a compatible buffer computation with a TileArray kernel.
7675
77- The pattern establishes computation semantics. Shared compiler checks
78- establish type, task, and memory compatibility. An unsuitable candidate
79- returns False without mutation; invalid implementations raise errors.
80- Staging validates the generated kernel before the source IR is changed.
76+ The pattern establishes applicability and semantic preconditions. The
77+ shared path adapts argument types and validates the generated kernel
78+ before changing the source IR.
8179 """
8280 from taskflow_mlir .dialects import func
8381
@@ -88,14 +86,12 @@ def replace_with_tile_array_program(
8886
8987 if operation .operation != self ._root .operation :
9088 raise ValueError ("TileArray replacement requires the pattern root" )
91- inferred_types = _task_argument_types (operation , arguments )
92- if inferred_types is None :
89+ buffer_types = _task_buffer_types (operation , arguments )
90+ if buffer_types is None :
9391 return False
9492
95- tile_program = build_tile_array_program (program , argument_types = inferred_types )
93+ tile_program = build_tile_array_program (program , argument_types = buffer_types )
9694 lowering = TileArrayProgramLowering (tile_program )
97- if not _replacement_memory_is_legal (operation , arguments , lowering ):
98- return False
9995
10096 staged = Module .create ()
10197 with InsertionPoint (staged .body ):
@@ -157,59 +153,13 @@ def _apply_patterns(
157153 return replace_count
158154
159155
160- def _base_buffer (value ):
161- """Traces task captures and view-like operations to their memory origin."""
162- while True :
163- if BlockArgument .isinstance (value ):
164- argument = BlockArgument (value )
165- parent = argument .owner .owner .operation
166- if parent .name == "taskflow.task" :
167- value = parent .operands [argument .arg_number ]
168- continue
169- return value
170- if not OpResult .isinstance (value ):
171- return value
172- producer = OpResult (value ).owner
173- if producer .name in (
174- "memref.cast" ,
175- "memref.subview" ,
176- "memref.reinterpret_cast" ,
177- ):
178- value = producer .operands [0 ]
179- continue
180- return value
181-
182-
183- def _disjoint_buffers (lhs , rhs ):
184- """Proves disjointness for fresh allocations and incoming function buffers."""
185- lhs , rhs = _base_buffer (lhs ), _base_buffer (rhs )
186- if lhs == rhs :
187- return False
188-
189- def is_allocation (value ):
190- return OpResult .isinstance (value ) and OpResult (value ).owner .name in (
191- "memref.alloc" ,
192- "memref.alloca" ,
193- )
194-
195- def is_function_argument (value ):
196- return (
197- BlockArgument .isinstance (value )
198- and BlockArgument (value ).owner .owner .operation .name == "func.func"
199- )
200-
201- return (
202- is_allocation (lhs ) and (is_allocation (rhs ) or is_function_argument (rhs ))
203- ) or (is_allocation (rhs ) and is_function_argument (lhs ))
204-
205-
206- def _task_argument_types (operation , arguments ):
207- """Returns supported capture types, or None when the task boundary is unsuitable."""
156+ def _task_buffer_types (operation , arguments ):
157+ """Returns supported buffer types, or None when the task boundary is unsuitable."""
208158 parent = operation .operation .parent
209159 if parent is None or parent .name != "taskflow.task" or len (operation .results ):
210160 return None
211161 block = parent .regions [0 ].blocks [0 ]
212- types = []
162+ buffer_types = []
213163 for value in arguments :
214164 if (
215165 not BlockArgument .isinstance (value )
@@ -228,31 +178,5 @@ def _task_argument_types(operation, arguments):
228178 return None
229179 if memref != MemRefType .get (list (memref .shape ), memref .element_type ):
230180 return None
231- types .append (TensorType (tuple (memref .shape ), dtype ))
232- return tuple (types )
233-
234-
235- def _replacement_memory_is_legal (operation , arguments , lowering ):
236- """Checks inferred implementation effects against task declarations and aliasing."""
237- if lowering .has_dynamic_memory :
238- return False
239- task = operation .operation .parent .opview
240- block_arguments = tuple (task .body .blocks [0 ].arguments )
241- read_count = len (task .will_reads )
242- write_count = len (task .will_writes )
243- declared_reads = block_arguments [:read_count ]
244- declared_writes = block_arguments [read_count : read_count + write_count ]
245- values = dict (zip (lowering .program .arguments , arguments ))
246- reads = lowering .read_arguments
247- writes = lowering .write_arguments
248- if any (values [argument ] not in declared_reads for argument in reads ):
249- return False
250- if any (values [argument ] not in declared_writes for argument in writes ):
251- return False
252- for output in writes :
253- for other in reads + writes :
254- if output is other :
255- continue
256- if not _disjoint_buffers (values [output ], values [other ]):
257- return False
258- return True
181+ buffer_types .append (BufferType (tuple (memref .shape ), dtype ))
182+ return tuple (buffer_types )
0 commit comments