Skip to content

Commit c45acc0

Browse files
authored
Merge pull request #8 from coredac/feature/type-system-modification
Simplify TileArray replacement and model typed buffers
2 parents 3c74929 + 097a833 commit c45acc0

15 files changed

Lines changed: 200 additions & 264 deletions

File tree

‎python/synapse/compiler/compiler.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,15 +6,15 @@
66
from tempfile import TemporaryDirectory
77

88
from synapse.frontend.lowering import lower
9-
from synapse.language.types import TensorType
9+
from synapse.language.types import BufferType
1010
from synapse.patterns import TileArrayProgramPattern
1111

1212

1313
def compile(
1414
program: Callable | str,
1515
*,
1616
target: str,
17-
argument_types: tuple[TensorType, ...] = (),
17+
argument_types: tuple[BufferType, ...] = (),
1818
patterns: Sequence[type[TileArrayProgramPattern]] | None = None,
1919
) -> str:
2020
"""Compiles a TileArray function or bufferized task IR for the backend."""

‎python/synapse/compiler/pattern_replacement.py‎

Lines changed: 12 additions & 88 deletions
Original file line numberDiff line numberDiff line change
@@ -11,12 +11,11 @@
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
2019
from 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)

‎python/synapse/frontend/lowering.py‎

Lines changed: 15 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010
from typing import TYPE_CHECKING, cast
1111

1212
from synapse.language.spatial import Tile
13-
from synapse.language.tensor import Tensor, TensorAccess
13+
from synapse.language.values import Buffer, BufferSlice
1414
from synapse.language.tile_array_program import (
1515
AddOp,
1616
ConstantOp,
@@ -21,27 +21,27 @@
2121
TileArrayOp,
2222
TileArrayProgram,
2323
)
24-
from synapse.language.types import DType, TensorType
24+
from synapse.language.types import BufferType, DType
2525

2626
if TYPE_CHECKING:
2727
from taskflow_mlir.ir import AffineMapAttr, DenseI64ArrayAttr, DictAttr, Value
2828

2929

30-
def lower(program_fn: Callable, *, argument_types: tuple[TensorType, ...] = ()) -> str:
30+
def lower(program_fn: Callable, *, argument_types: tuple[BufferType, ...] = ()) -> str:
3131
"""Lowers a standalone TileArray program to pre-mapping Taskflow and Neura IR."""
3232
program = build_tile_array_program(program_fn, argument_types=argument_types)
3333
return TileArrayProgramLowering(program).lower_to_single_task(program_fn.__name__)
3434

3535

3636
def build_tile_array_program(
37-
program_fn: Callable, *, argument_types: tuple[TensorType, ...] = ()
37+
program_fn: Callable, *, argument_types: tuple[BufferType, ...] = ()
3838
) -> TileArrayProgram:
39-
"""Runs a TileArray function and records its operations and tensor accesses.
39+
"""Runs a TileArray function and records its operations and buffer slices.
4040
4141
Direct compilation and pattern replacement share this construction step.
4242
It creates a TileArrayProgram without importing or constructing MLIR.
4343
"""
44-
# Program arguments currently carry tensor types. Scalar argument capture
44+
# Program arguments currently carry buffer types. Scalar argument capture
4545
# will use Taskflow value dependencies when that frontend path is added.
4646

4747
function_signature = signature(program_fn)
@@ -54,12 +54,12 @@ def build_tile_array_program(
5454
)
5555

5656
if any(
57-
not isinstance(argument_type, TensorType) for argument_type in argument_types
57+
not isinstance(argument_type, BufferType) for argument_type in argument_types
5858
):
59-
raise TypeError("program argument types must be TensorType values")
59+
raise TypeError("program argument types must be BufferType values")
6060

6161
arguments = [
62-
Tensor(name=name, type=argument_type)
62+
Buffer(name=name, type=argument_type)
6363
for name, argument_type in zip(parameter_names, argument_types)
6464
]
6565

@@ -199,7 +199,7 @@ def lower_to_single_task(self, program_name: str) -> str:
199199

200200
return str(module)
201201

202-
def lower_to_kernel(self, argument_values: dict[Tensor, Value]) -> None:
202+
def lower_to_kernel(self, argument_values: dict[Buffer, Value]) -> None:
203203
"""Creates one kernel at the caller's insertion point in an existing task.
204204
205205
The caller owns the current context, location, task, and terminator.
@@ -217,7 +217,7 @@ def lower_to_kernel(self, argument_values: dict[Tensor, Value]) -> None:
217217

218218
# Configuration validation precedes IR insertion.
219219
metadata = self.get_kernel_metadata()
220-
memory_configs: dict[int, tuple[TensorAccess, DenseI64ArrayAttr]] = {}
220+
memory_configs: dict[int, tuple[BufferSlice, DenseI64ArrayAttr]] = {}
221221
for index, operation in enumerate(self.program.operations):
222222
if isinstance(operation, LoadOp):
223223
access = operation.source
@@ -277,16 +277,16 @@ def get_mlir_type(self, dtype: DType):
277277

278278
raise NotImplementedError(f"unsupported tile-array data type: {dtype.value}")
279279

280-
def get_memref_type(self, tensor_type: TensorType):
281-
"""Translates a TensorType into a Taskflow MemRef type."""
280+
def get_memref_type(self, buffer_type: BufferType):
281+
"""Translates a BufferType into a Taskflow MemRef type."""
282282

283283
from taskflow_mlir.ir import MemRefType
284284

285285
return MemRefType.get(
286-
list(tensor_type.shape), self.get_mlir_type(tensor_type.dtype)
286+
list(buffer_type.shape), self.get_mlir_type(buffer_type.dtype)
287287
)
288288

289-
def get_memory_offsets(self, access: TensorAccess) -> DenseI64ArrayAttr:
289+
def get_memory_offsets(self, access: BufferSlice) -> DenseI64ArrayAttr:
290290
"""Converts a static access to Neura's constant element-offset array."""
291291
from taskflow_mlir.ir import DenseI64ArrayAttr
292292

‎python/synapse/language/__init__.py‎

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,20 +1,27 @@
11
"""Public Synapse language API."""
22

33
from .spatial import TileArray
4-
from .tensor import Tensor
54
from .tile_array_program import (
65
add,
76
constant,
87
load,
98
mac,
109
store,
1110
)
12-
from .types import DType, TensorType, f32, i32
11+
from .types import (
12+
BufferType,
13+
DType,
14+
f32,
15+
i32,
16+
)
17+
from .values import Buffer, BufferSlice, SynapseValue
1318

1419
__all__ = [
20+
"Buffer",
21+
"BufferSlice",
22+
"BufferType",
1523
"DType",
16-
"Tensor",
17-
"TensorType",
24+
"SynapseValue",
1825
"TileArray",
1926
"add",
2027
"constant",

‎python/synapse/language/tensor.py‎

Lines changed: 0 additions & 70 deletions
This file was deleted.

0 commit comments

Comments
 (0)