Skip to content
4 changes: 2 additions & 2 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@ check-lit-c:
env XTC_MLIR_TARGET=c lit -v tests/filecheck/backends tests/filecheck/mlir_loop

check-lit-nvgpu:
[ `uname -s` = Darwin ] || env XTC_MLIR_TARGET=nvgpu lit -v tests/filecheck/backends tests/filecheck/mlir_loop tests/filecheck/evaluation
[ `uname -s` = Darwin ] || env XTC_MLIR_TARGET=nvgpu lit -v tests/filecheck/backends tests/filecheck/mlir_loop tests/filecheck/evaluation tests/filecheck/schedules

check-lit-mppa:
[ `uname -s` = Darwin ] || env XTC_MLIR_TARGET=mppa lit -v -j 1 tests/filecheck/backends/target_mppa tests/filecheck/evaluation/test_matmul_pmu_counters_mppa.py
Expand Down Expand Up @@ -105,5 +105,5 @@ claude:
run-tutorial:
marimo run docs/tutorials/xtc_101.py

.PHONY: help test check check-lit-all check-lit check-lit-c check-lit-nvpgu check-pytest check-type check-pyright check-mypy check-format check-format-ruff check-license check-banwords format format-ruff format-license agents claude check-tutorials run-tutorial check-dependencies dependencies wheel pages
.PHONY: help test check check-lit-all check-lit check-lit-c check-lit-nvgpu check-pytest check-type check-pyright check-mypy check-format check-format-ruff check-license check-banwords format format-ruff format-license agents claude check-tutorials run-tutorial check-dependencies dependencies wheel pages
.SUFFIXES:
20 changes: 20 additions & 0 deletions src/xtc/backends/jir/JIRScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -361,6 +361,26 @@ def distributed_buffer_at(
# TODO: not implemented for now
pass

@override
def gpu_lane(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
# TODO: not implemented for now
pass

@override
def gpu_warp(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
# TODO: not implemented for now
pass

@override
def gpu_thread(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
# TODO: not implemented for now
pass

@override
def gpu_block(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
# TODO: not implemented for now
pass

def get_schedule_str(self) -> str:
return str(JIRSchedule(scheduler=self))

Expand Down
225 changes: 213 additions & 12 deletions src/xtc/backends/mlir/MlirCompilerPasses.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,13 +23,19 @@
MatchInterfaceEnum,
FuseIntoContainingOp,
)
from mlir.dialects.transform.gpu import (
MapForallToBlocks,
MapNestedForallToThreads,
)
from mlir.dialects.transform.loop import loop_unroll
from mlir.dialects.transform import SplitHandleOp
from mlir.ir import (
Location,
InsertionPoint,
UnitAttr,
OpResult,
Attribute,
ArrayAttr,
)
from mlir.passmanager import PassManager
from mlir.ir import Module
Expand All @@ -53,6 +59,7 @@
_VECTO_SEQ_NAME = "_vecto"
_SUPER_VECTORIZE_SEQ_NAME = "_super_vectorize"
_POST_BUFFERIZE_SEQ_NAME = "_post_bufferize"
_GPU_DIM = ["x", "y", "z"]


@dataclass
Expand Down Expand Up @@ -135,6 +142,7 @@ def __init__(
self._super_vectorize_sequence: NamedSequenceOp | None = None
self._post_bufferize_sequence: NamedSequenceOp | None = None
self._named_sequence: NamedSequenceOp | None = None
self._gpu_block_order: ArrayAttr | None = None
self._nodes_schedules = (
self._mlir_schedule.schedule_impl if self._mlir_schedule is not None else []
)
Expand Down Expand Up @@ -240,6 +248,13 @@ def _generate_scheduling(self) -> OpResult:
)
if schedule.vectorization or self._always_vectorize:
self._post_vectorize(scheduling_state, schedule)

# GPU mapping
if schedule.gpu_blocks:
self._gpu_mapping(
schedule,
scheduling_state,
)
handle = scheduling_state.handle

if schedule.fused_consumers:
Expand Down Expand Up @@ -326,7 +341,9 @@ def _generate_node_scheduling(
permutation = schedule.permutation[root]
if not permutation:
return sched_state

gpu_material = True
gpu_mat_thread = True
gpu_warp_thread = True
# Materialize the loops
for loop_name in permutation:
# Manage the splits
Expand Down Expand Up @@ -362,12 +379,58 @@ def _generate_node_scheduling(
self._vectorize(sched_state, self._vector_sizes_for(schedule))
break
elif loop_name in tiles_sizes_by_loops:
self._strip_mine(
loop_name=loop_name,
tiling_vector=tiles_sizes_by_loops[loop_name],
schedule=schedule,
sched_state=sched_state,
)
if loop_name in schedule.gpu_blocks:
if gpu_material:
self._gpu_strip_mine(
loop_name=loop_name,
schedule=schedule,
sched_state=sched_state,
gpu_list=schedule.gpu_blocks,
permutation=permutation,
tiles_sizes_by_loops=tiles_sizes_by_loops,
)
gpu_material = False
elif loop_name in schedule.gpu_warps:
if gpu_warp_thread:
self._gpu_strip_mine(
loop_name=loop_name,
schedule=schedule,
sched_state=sched_state,
gpu_list=schedule.gpu_warps,
permutation=permutation,
tiles_sizes_by_loops=tiles_sizes_by_loops,
)
gpu_warp_thread = False
elif loop_name in schedule.gpu_threads:
if gpu_mat_thread:
self._gpu_strip_mine(
loop_name=loop_name,
schedule=schedule,
sched_state=sched_state,
gpu_list=schedule.gpu_threads,
permutation=permutation,
tiles_sizes_by_loops=tiles_sizes_by_loops,
)
gpu_mat_thread = False
elif loop_name in schedule.gpu_lanes:
if gpu_mat_thread:
self._gpu_strip_mine(
loop_name=loop_name,
schedule=schedule,
sched_state=sched_state,
gpu_list=schedule.gpu_lanes,
permutation=permutation,
tiles_sizes_by_loops=tiles_sizes_by_loops,
)
gpu_mat_thread = False
else:
self._strip_mine(
loop_name=loop_name,
tiling_vector=tiles_sizes_by_loops[loop_name],
mapping_order=[],
schedule=schedule,
sched_state=sched_state,
)
if loop_name in schedule.distribution:
self._distribute_loop(loop_name, schedule, sched_state)
# Fuse the producers
Expand Down Expand Up @@ -546,20 +609,44 @@ def _strip_mine(
self,
loop_name: str,
tiling_vector: list[int],
mapping_order: list[int],
schedule: MlirNodeSchedule,
sched_state: SchedulingState,
) -> OpResult:
if loop_name in schedule.parallelization:
tiling_command = TileUsingForallOp(
sched_state.handle, tile_sizes=tiling_vector
attr_array = {}
attr_array["tile_sizes"] = tiling_vector
if loop_name in schedule.gpu_blocks:
attr_array["mapping"] = ArrayAttr.get(
[self._get_block_id(index) for index in mapping_order]
)
self._gpu_block_order = attr_array["mapping"]
tiling_command = TileUsingForallOp(sched_state.handle, **attr_array)
elif loop_name in schedule.gpu_threads:
attr_array["mapping"] = ArrayAttr.get(
[self._get_thread_id(index) for index in mapping_order]
)
tiling_command = TileUsingForallOp(sched_state.handle, **attr_array)
elif loop_name in schedule.gpu_warps:
attr_array["mapping"] = ArrayAttr.get(
[self._get_warp_id(index) for index in mapping_order]
)
tiling_command = TileUsingForallOp(sched_state.handle, **attr_array)
elif loop_name in schedule.gpu_lanes:
attr_array["mapping"] = ArrayAttr.get(
[self._get_lane_id(index) for index in mapping_order]
)
tiling_command = TileUsingForallOp(sched_state.handle, **attr_array)
elif loop_name in schedule.parallelization:
tiling_command = TileUsingForallOp(sched_state.handle, **attr_array)
else:
tiling_command = TileUsingForOp(sched_state.handle, sizes=tiling_vector)
# Extract the results
sched_state.handle = tiling_command.results[0]
assert len(tiling_command.results) == 2
new_loop = tiling_command.results[-1]
sched_state.all_loops[loop_name] = new_loop
if loop_name in schedule.gpu_blocks:
loop_name = schedule.gpu_blocks[0]
# Annotate the resulting loop if successfully generated
transform.AnnotateOp(new_loop, loop_name)

Expand Down Expand Up @@ -634,11 +721,12 @@ def _post_vectorize(self, sched_state: SchedulingState, schedule: MlirNodeSchedu
vector.ApplyTransferPermutationPatternsOp()

# the remaining patterns must be applied post-bufferization to work properly
if not self._post_bufferize_sequence:
if not self._post_bufferize_sequence and not schedule.gpu_blocks:
with InsertionPoint(transform.ApplyPatternsOp(parent_op).patterns):
vector.ApplyLowerOuterProductPatternsOp()
vector.ApplyLowerContractionPatternsOp()
else:
# Do not lower vector contract as it can be useful for gpu optimisation
elif self._post_bufferize_sequence and not schedule.gpu_blocks:
func_name = self._mlir_program.mlir_module.body.operations[0].attributes[
"sym_name"
]
Expand Down Expand Up @@ -772,6 +860,119 @@ def _collect_fused_producers(self, unscheduled_handles: set[str | None]):

return fused_producer_handles

def _get_lane_id(self, index: int) -> Attribute:
ctx = self._mlir_program.mlir_context
return Attribute.parse(f"#gpu.lane<linear_dim_{index}>", context=ctx)

def _get_warp_id(self, index: int) -> Attribute:
ctx = self._mlir_program.mlir_context
return Attribute.parse(f"#gpu.warp<{_GPU_DIM[index]}>", context=ctx)

def _get_thread_id(self, index: int) -> Attribute:
ctx = self._mlir_program.mlir_context
return Attribute.parse(f"#gpu.thread<{_GPU_DIM[index]}>", context=ctx)

def _get_block_id(self, index: int) -> Attribute:
ctx = self._mlir_program.mlir_context
return Attribute.parse(f"#gpu.block<{_GPU_DIM[index]}>", context=ctx)

def _gpu_mapping(
self,
schedule: MlirNodeSchedule,
sched_state: SchedulingState,
):
if schedule.gpu_blocks and not self._using_tensors:
new_loop = next(
(
sched_state.all_loops[loop_name]
for loop_name in schedule.gpu_blocks
if loop_name in sched_state.all_loops
),
None,
)
self._gpu_mapping_helper(schedule, new_loop)
elif (
schedule.gpu_blocks
and self._using_tensors
and self._post_bufferize_sequence
and self._gpu_block_order is not None
):
with (
InsertionPoint.at_block_begin(self._post_bufferize_sequence.body),
self._mlir_program.mlir_context,
self._loc,
):
gpu_block_handle = structured_match(
results_=transform.AnyOpType.get(),
target=self._post_bufferize_sequence.bodyTarget,
op_attrs={
schedule.gpu_blocks[0]: UnitAttr.get(),
"mapping": self._gpu_block_order,
},
)
self._gpu_mapping_helper(schedule, gpu_block_handle)

def _gpu_mapping_helper(self, schedule: MlirNodeSchedule, handle: OpResult):
tiles_sizes_by_loops = self._generate_tiling_insns(schedule)
new_loop = MapForallToBlocks(
handle,
generate_gpu_launch=True,
).result
block_dims: list[int] = []
for curType, gpu_list in enumerate(
[schedule.gpu_threads, schedule.gpu_lanes, schedule.gpu_warps]
):
if not gpu_list:
continue
# If there is a something in gpu warp multiply it by 32
thread_size = 1
if curType == 2:
thread_size = 32
block_dims = []
for loop_name in gpu_list:
tile_size = schedule.size_of_tile(loop_name)

if tile_size is None:
block_dims.append(1)
else:
block_dims.append(
thread_size
* (tile_size // max(tiles_sizes_by_loops[loop_name]))
)
if block_dims:
block_dims = block_dims + [1] * (3 - len(block_dims))
MapNestedForallToThreads(
new_loop,
block_dims=block_dims,
)

def _gpu_strip_mine(
self,
loop_name: str,
schedule: MlirNodeSchedule,
sched_state: SchedulingState,
gpu_list: list[str],
permutation: list[str],
tiles_sizes_by_loops: dict[str, list[int]],
):
tile_vect = [
sum(values)
for values in zip(*[tiles_sizes_by_loops[loop] for loop in gpu_list])
]
tile_vect = tile_vect + [0] * (3 - len(tile_vect))
# TODO: Make it work with splitting
position_index = [permutation.index(loop) for loop in gpu_list]
mapping_order = sorted(
range(len(position_index)), key=lambda i: position_index[i]
)
self._strip_mine(
loop_name=loop_name,
tiling_vector=tile_vect,
mapping_order=mapping_order,
schedule=schedule,
sched_state=sched_state,
)


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
Expand Down
12 changes: 12 additions & 0 deletions src/xtc/backends/mlir/MlirNodeScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,18 @@ def fuse_producer_at(
def fuse_consumer_at(self, axis: str, root: str = DEFAULT_ROOT) -> None:
self._plain_sch.fuse_consumer_at(axis, root)

def map_gpu_threads(self, axes: list[str], root: str = DEFAULT_ROOT):
self._plain_sch.gpu_thread(axes, root)

def map_gpu_blocks(self, axes: list[str], root: str = DEFAULT_ROOT):
self._plain_sch.gpu_block(axes, root)

def map_gpu_lanes(self, axes: list[str], root: str = DEFAULT_ROOT):
self._plain_sch.gpu_lane(axes, root)

def map_gpu_warps(self, axes: list[str], root: str = DEFAULT_ROOT):
self._plain_sch.gpu_warp(axes, root)

def get_node_schedule(self) -> MlirNodeSchedule:
plain_schedule = self._plain_sch.get_plain_schedule()
return MlirNodeSchedule(**asdict(plain_schedule))
Expand Down
16 changes: 16 additions & 0 deletions src/xtc/backends/mlir/MlirScheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -217,6 +217,22 @@ def distributed_buffer_at(
axis, input_idx, memory_axes, root=root
)

@override
def gpu_lane(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
self._current_scheduler.map_gpu_lanes(axes, root=root)

@override
def gpu_warp(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
self._current_scheduler.map_gpu_warps(axes, root=root)

@override
def gpu_thread(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
self._current_scheduler.map_gpu_threads(axes, root=root)

@override
def gpu_block(self, axes: list[str], root: str = DEFAULT_ROOT) -> None:
self._current_scheduler.map_gpu_blocks(axes, root=root)

@override
def get_loop_nest(self) -> LoopNest:
node_schedule = self._current_scheduler.get_node_schedule()
Expand Down
Loading
Loading