Skip to content
Open
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
100 changes: 85 additions & 15 deletions src/xtc/backends/mlir/MlirCompilerPasses.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,8 @@
)
from mlir.passmanager import PassManager
from mlir.ir import Module
import mlir.xtc_transform

from mlir.xtc_transform import FuseConsumerOp

# Import SDist if available
try:
Expand Down Expand Up @@ -235,12 +236,17 @@ def _generate_scheduling(self) -> OpResult:
schedule=schedule,
root=list(schedule.permutation)[0],
handle=handle,
fuse_axes=fused_producers.get(schedule.node_ident),
producer_fuse_axes=fused_producers.get(schedule.node_ident),
)
if schedule.vectorization or self._always_vectorize:
self._post_vectorize(scheduling_state, schedule)
handle = scheduling_state.handle

if schedule.fused_consumers:
self._fuse_consumers_into_loops(
schedule, scheduling_state, unscheduled_handles
)

assert handle, "At least 1 operation should have been processed"
return handle

Expand Down Expand Up @@ -312,7 +318,7 @@ def _generate_node_scheduling(
schedule: MlirNodeSchedule,
root: str,
handle: OpResult,
fuse_axes: dict[str, list[str]] | None,
producer_fuse_axes: dict[str, list[str]] | None,
) -> SchedulingState:
sched_state = SchedulingState({}, handle, None)
split_state = SplitState(schedule.splits, root)
Expand Down Expand Up @@ -365,9 +371,9 @@ def _generate_node_scheduling(
if loop_name in schedule.distribution:
self._distribute_loop(loop_name, schedule, sched_state)
# Fuse the producers
if fuse_axes and loop_name in fuse_axes:
if producer_fuse_axes and loop_name in producer_fuse_axes:
self._fuse_producers_into_loop(
loop_name, fuse_axes, schedule, sched_state
loop_name, producer_fuse_axes, schedule, sched_state
)

# For now on, the focus is on the outermost loop
Expand All @@ -380,6 +386,45 @@ def _generate_node_scheduling(

return sched_state

def _fuse_consumers_into_loops(
self,
schedule: MlirNodeSchedule,
sched_state: SchedulingState,
unscheduled_handles: set[str | None],
):
assert self._named_sequence is not None
assert len(schedule.fused_consumers) == 1

fuse_root = parent_name(schedule.fused_consumers[0])
fuse_axis = schedule.fused_consumers[0]
# derive handle of consumer
consumer_handles = find_consumer_handles(
self._mlir_program.mlir_module, schedule.node_ident
)
if not consumer_handles:
return
consumer_id = consumer_handles[0]
unscheduled_handles.add(consumer_id)
# fuse consumer into all loops until the fuse axis
fuse_loops = []
for loop_dim in schedule.permutation[fuse_root]:
transform_result = sched_state.all_loops[loop_dim]
fuse_loops.append(transform_result)
if loop_dim == fuse_axis:
break
consumer_handle = structured_match(
results_=transform.AnyOpType.get(),
target=self._named_sequence.bodyTarget,
op_attrs={consumer_id: UnitAttr.get()},
)
op = FuseConsumerOp(consumer_handle, fuse_loops)
# re-annotate the loops that were touched by the fusion
for i, loop_dim in enumerate(schedule.permutation[fuse_root]):
sched_state.all_loops[loop_dim] = op.new_loops[i]
transform.AnnotateOp(op.new_loops[i], loop_dim)
if loop_dim == fuse_axis:
break

def _fuse_producers_into_loop(
self,
loop_name: str,
Expand Down Expand Up @@ -484,7 +529,10 @@ def _recursive_scheduling(
self, schedule: MlirNodeSchedule, root: str, sched_state: SchedulingState
):
inner_sched_state = self._generate_node_scheduling(
schedule=schedule, root=root, handle=sched_state.handle, fuse_axes=None
schedule=schedule,
root=root,
handle=sched_state.handle,
producer_fuse_axes=None,
)
sched_state.all_loops.update(inner_sched_state.all_loops)
sched_state.handle = inner_sched_state.handle
Expand Down Expand Up @@ -673,35 +721,58 @@ def _pack_buffer(
)

def _collect_fused_producers(self, unscheduled_handles: set[str | None]):
# maps each fused consumer op to the producer handles that must be
# maps each fused containing op to the producer handles that must be
# fused through each loop dimension to reach their target fusion depth.
fused_producers = {}
fused_producer_handles = {}

for schedule in self._nodes_schedules:
if schedule.fused:
if schedule.fused_producers:
prods = find_producer_handles(
self._mlir_program.mlir_module, schedule.node_ident
)
fuse_root = parent_name(schedule.fused[0][0])
fuse_root = parent_name(schedule.fused_producers[0][0])
unscheduled_handles.update(set(prods))
op_axes = {idx: ax for ax, idx in schedule.fused}
op_axes = {idx: ax for ax, idx in schedule.fused_producers}

fuse_destinations = {}
for idx, prod_handle in enumerate(prods):
if not prod_handle:
continue
if idx in op_axes:
fuse_destinations[prod_handle] = op_axes[idx]
# get outer dims to fuse, assumes fuse no splitting avove loop dim
# get outer dims to fuse, assumes fuse no splitting above loop dim
dim_fuse_handles: dict[str, list[str]] = {}
for fuse_handle, fuse_dest in fuse_destinations.items():
for dim in schedule.permutation[fuse_root]:
dim_fuse_handles.setdefault(dim, []).append(fuse_handle)
if dim == fuse_dest:
break
fused_producers[schedule.node_ident] = dim_fuse_handles
fused_producer_handles[schedule.node_ident] = dim_fuse_handles

return fused_producer_handles


def find_consumer_handles(module: Module, root_handle: str) -> list[str | None]:
# returns the handles for each consumer op of the operation specified by root_handle
consumer_handles: list[str | None] = []
root_op = None
for func_op in module.body.operations:
for op in func_op.regions[0].blocks[0].operations:
if root_handle in op.attributes:
root_op = op
break
if root_op:
break

if not root_op:
return consumer_handles

return fused_producers
for use in root_op.results[0].uses:
consumer_op = use.owner
for attr in consumer_op.attributes:
if attr.startswith("__xtc_id_"):
consumer_handles.append(attr)
return consumer_handles


def find_producer_handles(module: Module, root_handle: str) -> list[str | None]:
Expand Down Expand Up @@ -779,7 +850,6 @@ def run(self, pass_names: list[str]) -> None:


def apply_bufferization_passes(mlir_program: RawMlirProgram, mlir_install_dir: str):
assert mlir.xtc_transform
bufferize_options = [
"bufferize-function-boundaries",
"function-boundary-type-conversion=identity-layout-map",
Expand Down
3 changes: 1 addition & 2 deletions src/xtc/backends/mlir/MlirScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -185,8 +185,7 @@ def fuse_producer_at(

@override
def fuse_consumer_at(self, axis: str, root: str = DEFAULT_ROOT) -> None:
# TODO: not implemented for now
pass
self._current_scheduler.fuse_consumer_at(axis, root=root)

@override
def define_memory_mesh(self, axes: dict[str, int]) -> None:
Expand Down
2 changes: 1 addition & 1 deletion src/xtc/schedules/loop_nest_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,7 +79,7 @@ def localize_axis_tuples(
if is_node_axis(axis)
}
# TODO: loop nest supports only one fuse per axis
fuse_producer_at = dict(localize_axis_tuples(node_sched.fused))
fuse_producer_at = dict(localize_axis_tuples(node_sched.fused_producers))
# TODO: loop nest supports only one fuse consumer per axis
fuse_consumer_at = localize_axis_list(node_sched.fused_consumers)

Expand Down
8 changes: 4 additions & 4 deletions src/xtc/schedules/plain_schedule.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@ class PlainNodeSchedule:
processor_mesh: dict[str, int]
distribution: dict[str, str]
distributed_buffers: dict[str, dict]
fused: list[tuple[str, int]]
fused_producers: list[tuple[str, int]]
fused_consumers: list[str]
# Optional caller-provided vector sizes, keyed by vectorized axis name.
# When an axis has a size, its dimension is vectorized with masking for
Expand Down Expand Up @@ -114,7 +114,7 @@ def __init__(
self.processor_mesh: dict[str, int] = {}
self.distribution: dict[str, str] = {}
self.distributed_buffers: dict[str, dict] = {}
self.fused: list[tuple[str, int]] = []
self.fused_producers: list[tuple[str, int]] = []
self.fused_consumers: list[str] = []

def get_plain_schedule(self) -> PlainNodeSchedule:
Expand All @@ -135,7 +135,7 @@ def get_plain_schedule(self) -> PlainNodeSchedule:
processor_mesh=deepcopy(self.processor_mesh),
distribution=deepcopy(self.distribution),
distributed_buffers=deepcopy(self.distributed_buffers),
fused=deepcopy(self.fused),
fused_producers=deepcopy(self.fused_producers),
fused_consumers=deepcopy(self.fused_consumers),
vectorization_sizes=deepcopy(self.vectorization_sizes),
)
Expand Down Expand Up @@ -269,7 +269,7 @@ def fuse_producer_at(
self, axis: str, input_idx: int, root: str = DEFAULT_ROOT
) -> None:
fuse_axis = make_loop_name(root, axis)
self.fused.append((fuse_axis, input_idx))
self.fused_producers.append((fuse_axis, input_idx))

def fuse_consumer_at(self, axis: str, root: str = DEFAULT_ROOT) -> None:
fuse_axis = make_loop_name(root, axis)
Expand Down
Loading
Loading