Skip to content
Closed
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
6 changes: 3 additions & 3 deletions src/xtc/backends/mlir/MlirCompilerPasses.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,9 @@
# SPDX-License-Identifier: BSD-3-Clause
# Copyright (c) 2024-2026 The XTC Project Authors
#
from .MlirLoopNames import parent_name
from dataclasses import dataclass
import subprocess

from mlir.dialects import transform
from mlir.dialects.transform import (
NamedSequenceOp,
Expand Down Expand Up @@ -33,7 +34,6 @@
)
from mlir.passmanager import PassManager
from mlir.ir import Module
import subprocess

# Import SDist if available
try:
Expand All @@ -42,7 +42,7 @@
sdist_transform = None
pass

from .MlirLoopNames import make_loop_name
from xtc.schedules.loop_names import make_loop_name, parent_name
from xtc.utils.ext_tools import transform_opts

from .MlirProgram import RawMlirProgram
Expand Down
194 changes: 42 additions & 152 deletions src/xtc/backends/mlir/MlirNodeScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,9 @@
#
from typing_extensions import override
from dataclasses import dataclass, asdict
from pprint import pformat

from xtc.itf.schd.scheduler import DEFAULT_ROOT
from .MlirLoopNames import make_loop_name, basename
from xtc.schedules.plain_schedule import PlainNodeSchedule, PlainNodeScheduler

__all__ = [
"MlirNodeScheduler",
Expand All @@ -15,58 +15,8 @@


@dataclass(frozen=True)
class MlirNodeSchedule:
node_name: str
node_ident: str
dims: list[str]
loop_stamps: list[str]
splits: dict[str, dict[str, int]]
tiles: dict[str, dict[str, int]]
permutation: dict[str, list[str]]
vectorization: list[str]
parallelization: list[str]
unrolling: dict[str, int]
packed_buffers: dict[str, list[int]]
memory_mesh: dict[str, int]
processor_mesh: dict[str, int]
distribution: dict[str, str]
distributed_buffers: dict[str, dict]
fused: list[tuple[str, int]]

def index_of_dim(self, dim: str) -> int:
return list(self.dims).index(dim)

def is_tile(self, loop_name: str) -> bool:
for tiles in self.tiles.values():
for tile in tiles:
if loop_name == tile:
return True
return False

def is_base(self, loop_name: str) -> bool:
return basename(loop_name) in self.dims

def dim_of_tile(self, loop_name: str) -> str:
# Base dimension
bn = basename(loop_name)
if bn in self.dims:
return bn
# Tiled dimension
for dim, tiles in self.tiles.items():
for tile in tiles:
if bn == dim or loop_name == tile:
return dim
assert False

def size_of_tile(self, tile_name: str) -> int | None:
for tiles in self.tiles.values():
if tile_name in tiles:
return tiles[tile_name]
return None

@override
def __str__(self):
return pformat(asdict(self))
class MlirNodeSchedule(PlainNodeSchedule):
pass


class MlirNodeScheduler:
Expand All @@ -93,81 +43,39 @@ def __init__(
self.distribution: dict[str, str] = {}
self.distributed_buffers: dict[str, dict] = {}
self.fused: list[tuple[str, int]] = []

def mlir_node_schedule(self) -> MlirNodeSchedule:
if not self.permutation:
self.permutation[DEFAULT_ROOT] = self.get_default_interchange(DEFAULT_ROOT)

for fuse_axis in self.fused:
assert fuse_axis[0] in self.permutation[next(iter(self.permutation))], (
"Fusion must be to an axis in the base root not the result of a split."
)

return MlirNodeSchedule(
node_name=self.node_name,
node_ident=self.node_ident,
dims=self.dims,
loop_stamps=self.loop_stamps,
tiles=self.tiles,
splits=self.splits,
permutation=self.permutation,
vectorization=self.vectorization,
parallelization=self.parallelization,
unrolling=self.unrolling,
memory_mesh=self.memory_mesh,
packed_buffers=self.packed_buffers,
processor_mesh=self.processor_mesh,
distribution=self.distribution,
distributed_buffers=self.distributed_buffers,
fused=self.fused,
self._plain_sch = PlainNodeScheduler(
node_name,
node_ident,
dims,
)

@override
def __str__(self) -> str:
return str(self.mlir_node_schedule())

def get_default_interchange(self, root: str) -> list[str]:
ret = [make_loop_name(root, d) for d in self.dims]
for tile_level in range(len(max(self.tiles.values(), key=len))):
for _, v in self.tiles.items():
if tile_level >= len(v):
continue
dim_name = list(v.keys())[tile_level]
ret.append(dim_name)
return ret

def set_dims(self, dims: list[str]) -> None:
assert len(dims) == len(self.dims)
self.dims = dims[:]
self.tiles = {k: {} for k in self.dims}
self._plain_sch.set_dims(dims)

def split(
self, dim: str, segments: dict[str, int], root: str = DEFAULT_ROOT
) -> None:
segments_renamed = {
make_loop_name(root, key): val for key, val in segments.items()
}
self.splits[dim] = segments_renamed
for s in segments_renamed:
self.tiles[s] = {}
self._plain_sch.split(dim, segments, root)

def tile(self, dim: str, tiles: dict[str, int], root: str = DEFAULT_ROOT):
for d, s in tiles.items():
tile_name = make_loop_name(root, d)
self.tiles[dim][tile_name] = s
def tile(self, dim: str, tiles: dict[str, int], root: str = DEFAULT_ROOT) -> None:
self._plain_sch.tile(dim, tiles, root)

def interchange(self, permutation: list[str], root: str = DEFAULT_ROOT):
self.permutation[root] = [make_loop_name(root, a) for a in permutation]
def interchange(self, permutation: list[str], root: str = DEFAULT_ROOT) -> None:
self._plain_sch.interchange(permutation, root)

def vectorize(self, axes: list[str], root: str = DEFAULT_ROOT):
self.vectorization += [make_loop_name(root, a) for a in axes]
def vectorize(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
self._plain_sch.vectorize(axes, root)

def parallelize(self, axes: list[str], root: str = DEFAULT_ROOT):
self.parallelization = [make_loop_name(root, a) for a in axes]
def parallelize(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
self._plain_sch.parallelize(axes, root)

def unroll(self, unrolls: dict[str, int], root: str = DEFAULT_ROOT):
for dim, ufactor in unrolls.items():
self.unrolling[make_loop_name(root, dim)] = ufactor
def unroll(self, unrolls: dict[str, int], root: str = DEFAULT_ROOT) -> None:
self._plain_sch.unroll(unrolls, root)

def buffer_at(
self, axis: str, mtype: str | None = None, root: str = DEFAULT_ROOT
) -> None:
self._plain_sch.buffer_at(axis, mtype, root)

def pack_at(
self,
Expand All @@ -177,36 +85,16 @@ def pack_at(
pad: bool = False,
root: str = DEFAULT_ROOT,
):
axis_key = make_loop_name(root, axis)
if axis_key not in self.packed_buffers.keys():
self.packed_buffers[axis_key] = [input_idx]
else:
self.packed_buffers[axis_key].append(input_idx)
self._plain_sch.pack_at(axis, input_idx, mtype, pad, root)

def define_memory_mesh(self, axes: dict[str, int]):
assert len(self.memory_mesh) == 0, "Memory mesh has already been defined"
self.memory_mesh = axes
self._plain_sch.define_memory_mesh(axes)

def define_processor_mesh(self, axes: dict[str, int]):
assert len(self.processor_mesh) == 0, "Processor mesh has already been defined"
assert self.memory_mesh, "Memory mesh has not been defined"
assert len(self.memory_mesh) <= len(axes), (
"Memory mesh must be a subset of the processor mesh"
)
for i, memory_size in enumerate(self.memory_mesh.values()):
assert list(axes.values())[i] == memory_size, (
"Memory mesh must be a subset of the processor mesh"
)
self.processor_mesh = axes
self._plain_sch.define_processor_mesh(axes)

def distribute(self, axis: str, processor_axis: str, root: str = DEFAULT_ROOT):
assert self.processor_mesh, "Processor mesh has not been defined"
assert processor_axis in self.processor_mesh or processor_axis == "*", (
"Processor axis not found in processor mesh"
)
axis_key = make_loop_name(root, axis)
self.parallelization.append(axis_key)
self.distribution[axis_key] = processor_axis
self._plain_sch.distribute(axis, processor_axis, root)

def distributed_buffer_at(
self,
Expand All @@ -215,18 +103,20 @@ def distributed_buffer_at(
memory_axes: list[str],
root: str = DEFAULT_ROOT,
):
assert self.memory_mesh, "Memory mesh has not been defined"
for ma in memory_axes:
assert ma in self.memory_mesh or ma == "*", (
"Memory axis not found in memory mesh"
)
axis_key = make_loop_name(root, axis)
self.distributed_buffers[axis_key] = {
"input_idx": input_idx,
"memory_axes": memory_axes,
}
self._plain_sch.distributed_buffer_at(axis, input_idx, memory_axes, root)

def fuse_producer_at(
self, axis: str, input_idx: int, root: str = DEFAULT_ROOT
) -> None:
self.fused.append((make_loop_name(root, axis), input_idx))
self._plain_sch.fuse_producer_at(axis, input_idx, root)

def fuse_consumer_at(self, axis: str, root: str = DEFAULT_ROOT) -> None:
self._plain_sch.fuse_consumer_at(axis, root)

def get_node_schedule(self) -> MlirNodeSchedule:
plain_schedule = self._plain_sch.get_plain_schedule()
return MlirNodeSchedule(**asdict(plain_schedule))

@override
def __str__(self) -> str:
return str(self.get_node_schedule())
62 changes: 7 additions & 55 deletions src/xtc/backends/mlir/MlirScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,12 +5,13 @@
from typing_extensions import override

from xtc.itf.schd.scheduler import DEFAULT_ROOT
from xtc.schedules.loop_nest import LoopNest, LoopNestNode, LoopInfo, SplitOrigin
from xtc.schedules.loop_nest import LoopNest
from xtc.schedules.loop_nest_builder import LoopNestBuilder

import xtc.itf as itf
import xtc.backends.mlir as backend

from .MlirNodeScheduler import MlirNodeScheduler, MlirNodeSchedule
from .MlirLoopNames import basename

__all__ = [
"MlirScheduler",
Expand Down Expand Up @@ -98,12 +99,12 @@ def schedule(self) -> itf.schd.Schedule:

if isinstance(self._backend, MlirGraphBackend):
nodes_schedules = [
scheduler._current_scheduler.mlir_node_schedule()
scheduler._current_scheduler.get_node_schedule()
for scheduler in self._nodes_schedulers
]
else:
assert isinstance(self._backend, MlirNodeBackend)
nodes_schedules = [self._current_scheduler.mlir_node_schedule()]
nodes_schedules = [self._current_scheduler.get_node_schedule()]
return MlirSchedule(
scheduler=self,
nodes_schedules=nodes_schedules,
Expand Down Expand Up @@ -211,57 +212,8 @@ def distributed_buffer_at(

@override
def get_loop_nest(self) -> LoopNest:
node_sched = self._current_scheduler
dims = node_sched.dims[:]

loop_nest = LoopNest(abstract_dims=dims)
root_node = loop_nest.build_root_node(node_sched.node_name)

# Assign splits to root_node first, stripping the root prefix from names
for axis, axis_splits in node_sched.splits.items():
root_node.splits[axis] = {basename(k): v for k, v in axis_splits.items()}

# Build mapper to get splits_info
mapper = LoopInfo.build_from_node(root_node)

def populate_node(node: LoopNestNode, perm: list[str]) -> None:
"""Populate node with data for loops in its permutation."""
perm_set = set(perm)
node.interchange = [basename(n) for n in perm]
for axis, axis_tiles in node_sched.tiles.items():
for tile_name, size in axis_tiles.items():
if tile_name in perm_set:
if axis not in node.tiles:
node.tiles[axis] = {}
node.tiles[axis][basename(tile_name)] = size
node.vectorize = [
basename(v) for v in node_sched.vectorization if v in perm_set
]
node.parallelize = [
basename(p) for p in node_sched.parallelization if p in perm_set
]
node.unroll = {
basename(k): v for k, v in node_sched.unrolling.items() if k in perm_set
}

# Process each root in permutation
for root, perm in node_sched.permutation.items():
root_name = basename(root)
if root_name in mapper.splits_info:
# This root is a split - create child node
axis, start, end = mapper.splits_info[root_name]
child = LoopNestNode(
root=root_name,
tiles={d: {} for d in dims},
split_origin=SplitOrigin(axis=axis, start=start, end=end),
)
populate_node(child, perm)
root_node.add_child(child)
else:
# This is the main root
populate_node(root_node, perm)

return loop_nest
node_schedule = self._current_scheduler.get_node_schedule()
return LoopNestBuilder.from_plain_node_schedule(node_schedule)


class MlirSchedule(itf.schd.Schedule):
Expand Down
Loading
Loading