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
36 changes: 22 additions & 14 deletions src/xtc/backends/jir/JIRScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@

import xtc.itf as itf
from xtc.itf.schd.scheduler import DEFAULT_ROOT
from xtc.schedules.loop_nest import LoopNest
from xtc.schedules.loop_nest import LoopNest, LoopNestNode
import xtc.backends.jir as backend

__all__ = [
Expand Down Expand Up @@ -359,26 +359,34 @@ def get_loop_nest(self) -> LoopNest:
transformer = self._transformer
dims = list(transformer.dims.keys())

loop_nest = LoopNest(abstract_dims=dims)
root_node = loop_nest.build_root_node(self._backend.payload_name)

# Build tiles mapping
for axis, axis_tiles in transformer.tiles.items():
for tile_name, size in axis_tiles.items():
root_node.tiles[axis][tile_name] = size
tiles = {
axis: {tile_name: size for tile_name, size in axis_tiles.items()}
for axis, axis_tiles in transformer.tiles.items()
}

# Build interchange
root_node.interchange = (
list(transformer.order) if transformer.order else dims[:]
)
interchange = list(transformer.order) if transformer.order else dims[:]

# Build vectorization list
root_node.vectorize = list(transformer.vectorized)
vectorize = list(transformer.vectorized)

# Build parallelization list
root_node.parallelize = list(transformer.parallelized)
parallelize = list(transformer.parallelized)

# Build unroll mapping
root_node.unroll = dict(transformer.unrolled)

unroll = dict(transformer.unrolled)

root_node = LoopNestNode(
root=self._backend.payload_name,
tiles=tiles,
interchange=interchange,
vectorize=vectorize,
parallelize=parallelize,
unroll=unroll,
)
loop_nest = LoopNest(
abstract_dims=dims,
root_node=root_node,
)
return loop_nest
4 changes: 2 additions & 2 deletions src/xtc/backends/mlir/MlirCompilerPasses.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@
sdist_transform = None
pass

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

from .MlirProgram import RawMlirProgram
Expand Down Expand Up @@ -465,7 +465,7 @@ def _split_section(
dim_to_split = split_state.loop_dim_by_split[loop_name]
split_handle = structured.SplitOp(
target=sched_state.handle,
dimension=schedule.dims.index(dim_to_split),
dimension=schedule.dims.index(basename(dim_to_split)),
chunk_sizes=split_state.chunk_size(loop_name),
)
split_command = SplitHandleOp(
Expand Down
4 changes: 4 additions & 0 deletions src/xtc/backends/mlir/MlirScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,10 @@ def __init__(
loop_stamps=backend.loop_stamps,
)

@property
def _default_node_name(self) -> str:
return self._current_scheduler.node_name

@property
def _current_scheduler(self) -> MlirNodeScheduler:
assert self._node_scheduler is not None
Expand Down
13 changes: 11 additions & 2 deletions src/xtc/backends/tvm/TVMScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,7 +81,7 @@ def _full_tilings(self, sched: LoopNestNode) -> dict[str, tuple[str, str, int]]:
for idx, (axis, size) in enumerate(dim_tiles.items()):
tilings[axis] = (dim, prev_axis, size)
prev_axis = axis
tilings = {axis: tilings[axis] for axis in order}
tilings = {axis: tilings.get(axis, (axis, "", 0)) for axis in order}
return tilings

def _write_buffer_tiling(
Expand Down Expand Up @@ -417,7 +417,7 @@ def _update_loopnode(node: LoopNestNode) -> LoopNestNode:
)
adjusted_tiles[dim] = adjusted_dim_tiles
adjusted_unrolling = {u: adjusted_unrolling[u] for u in adjusted_unrolls}
return LoopNestNode(
updated_node = LoopNestNode(
root=node.root,
tiles=adjusted_tiles,
splits=deepcopy(node.splits),
Expand All @@ -429,7 +429,12 @@ def _update_loopnode(node: LoopNestNode) -> LoopNestNode:
pack_at=deepcopy(node.pack_at),
fuse_producer_at=deepcopy(node.fuse_producer_at),
fuse_consumer_at=deepcopy(node.fuse_consumer_at),
split_origin=deepcopy(node.split_origin),
)
for child in node.children:
updated_child = _update_loopnode(child)
updated_node.add_child(updated_child)
return updated_node

root = sched.root_node
if root is not None:
Expand Down Expand Up @@ -467,6 +472,10 @@ def __init__(
list(self._op.operator.dims()),
)

@property
def _default_node_name(self) -> str:
return self._op.name

@property
@override
def backend(self) -> itf.back.Backend:
Expand Down
17 changes: 16 additions & 1 deletion src/xtc/schedules/loop_names.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,10 @@
# SPDX-License-Identifier: BSD-3-Clause
# Copyright (c) 2024-2026 The XTC Project Authors
#
__all__ = ["make_loop_name", "basename", "parent_name"]
__all__ = ["make_loop_name", "basename", "parent_name", "path_names", "rooted_name"]

_LOOP_SEP = "/"
_LOOP_ROOT = "."


def make_loop_name(root: str, name: str) -> str:
Expand All @@ -17,3 +18,17 @@ def basename(loop_name: str) -> str:

def parent_name(loop_name: str) -> str:
return loop_name.rsplit(_LOOP_SEP, 1)[0]


def path_names(loop_name: str) -> list[str]:
return loop_name.split(_LOOP_SEP)


def rooted_name(loop_name: str, root: str) -> str:
assert root != _LOOP_ROOT
if loop_name == _LOOP_ROOT:
return root
prefix = f"{_LOOP_ROOT}{_LOOP_SEP}"
if loop_name.startswith(prefix):
return make_loop_name(root, loop_name[len(prefix) :])
return loop_name
27 changes: 13 additions & 14 deletions src/xtc/schedules/loop_nest.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from .exceptions import ScheduleValidationError


@dataclass
@dataclass(frozen=True)
class SplitOrigin:
"""Describes how a node was created via a split from its parent.

Expand All @@ -28,7 +28,7 @@ class SplitOrigin:
NodeT = TypeVar("NodeT", bound="Node")


@dataclass(kw_only=True)
@dataclass(frozen=True, kw_only=True)
class Node(Generic[NodeT]):
"""Base class for tree nodes with parent/child relationships.

Expand All @@ -52,8 +52,13 @@ def is_root(self) -> bool:
return self.parent is None

def add_child(self, child: NodeT) -> None:
"""Add a child node and set its parent to this node."""
child.parent = self # type: ignore[assignment]
"""Add a child node and set its parent to this node.

If the parent was already set, assert it's same.
"""
if child.parent is not None:
assert child.parent is self
object.__setattr__(child, "parent", self)
self.children.append(child)

def ancestors(self) -> list[NodeT]:
Expand All @@ -74,7 +79,7 @@ def descendants_dfs(self) -> list[NodeT]:
return result


@dataclass
@dataclass(frozen=True)
class LoopNestNode(Node["LoopNestNode"]):
"""Represents a node in the loop nest tree with its transformations.

Expand Down Expand Up @@ -103,7 +108,7 @@ class LoopNestNode(Node["LoopNestNode"]):
"""

root: str
tiles: dict[str, dict[str, int]]
tiles: dict[str, dict[str, int]] = field(default_factory=dict)
splits: dict[str, dict[str, int]] = field(default_factory=dict)
interchange: list[str] = field(default_factory=list)
vectorize: list[str] = field(default_factory=list)
Expand Down Expand Up @@ -245,7 +250,7 @@ def _add_annotations(self, line: str, loop_name: str) -> str:
return line


@dataclass
@dataclass(frozen=True)
class LoopInfo:
"""Maps loop names to their corresponding axis names and metadata.

Expand Down Expand Up @@ -327,7 +332,7 @@ def collect(n: LoopNestNode) -> None:
return LoopInfo(list(dims), tiles_info, splits_info)


@dataclass
@dataclass(frozen=True)
class LoopNest:
"""Represents a complete loop nest structure for scheduling.

Expand All @@ -353,12 +358,6 @@ def nodes(self) -> list[LoopNestNode]:
return []
return [self.root_node] + self.root_node.descendants_dfs()

def build_root_node(self, root: str) -> LoopNestNode:
"""Build and set the root node of the loop nest tree."""
node = LoopNestNode(root=root, tiles={a: {} for a in self.abstract_dims})
self.root_node = node
return node

def check(self):
assert self.root_node is not None
info = LoopInfo.build_from_node(self.root_node)
Expand Down
Loading
Loading