diff --git a/src/xtc/backends/jir/JIRScheduler.py b/src/xtc/backends/jir/JIRScheduler.py index dde8312c5..e12c358ca 100644 --- a/src/xtc/backends/jir/JIRScheduler.py +++ b/src/xtc/backends/jir/JIRScheduler.py @@ -318,6 +318,11 @@ def fuse_producer_at( # TODO: not implemented for now pass + @override + def fuse_consumer_at(self, axis: str, root: str = DEFAULT_ROOT) -> None: + # TODO: not implemented for now + pass + @override def define_memory_mesh(self, axes: dict[str, int]) -> None: # TODO: not implemented for now diff --git a/src/xtc/backends/mlir/MlirScheduler.py b/src/xtc/backends/mlir/MlirScheduler.py index c47a49535..25d163748 100644 --- a/src/xtc/backends/mlir/MlirScheduler.py +++ b/src/xtc/backends/mlir/MlirScheduler.py @@ -174,6 +174,11 @@ def fuse_producer_at( ) -> None: self._current_scheduler.fuse_producer_at(axis, input_idx, root=root) + @override + def fuse_consumer_at(self, axis: str, root: str = DEFAULT_ROOT) -> None: + # TODO: not implemented for now + pass + @override def define_memory_mesh(self, axes: dict[str, int]) -> None: self._require_extension("sdist") diff --git a/src/xtc/backends/tvm/TVMBackend.py b/src/xtc/backends/tvm/TVMBackend.py index 325edcc05..69e612be3 100644 --- a/src/xtc/backends/tvm/TVMBackend.py +++ b/src/xtc/backends/tvm/TVMBackend.py @@ -29,6 +29,7 @@ def __init__( reduction_dims: list[str] | None = None, **kwargs: Any, ) -> None: + self._tir_schedule = kwargs.get("tir_schedule", False) self._graph: Graph | None = None self._tvm_base: TVMBaseExpr if isinstance(source_op, XTCGraph): diff --git a/src/xtc/backends/tvm/TVMCompiler.py b/src/xtc/backends/tvm/TVMCompiler.py index d4eb012f1..60195d63a 100644 --- a/src/xtc/backends/tvm/TVMCompiler.py +++ b/src/xtc/backends/tvm/TVMCompiler.py @@ -4,7 +4,6 @@ # from typing import Any, cast from typing_extensions import override -from collections.abc import Sequence import tempfile from pathlib import Path import shutil @@ -12,6 +11,7 @@ import shlex import sys from functools import partial +from packaging.version import Version from xtc.targets.host import HostModule @@ -22,7 +22,15 @@ from xtc.utils.host_tools import disassemble, target_triple -from .TVMOps import TVMBaseExpr, TESchedule, TETensor +from .TVMOpsCompiler import ( + TVMExprCompiler, + TVMScheduledExpr, + TVMScheduledExprTE, + TVMScheduledExprTIR, +) +from .TVMOps import ( + TVMBaseExpr, +) import tvm @@ -31,6 +39,8 @@ "TVMCompiler", ] +TVM_VERSION = Version(tvm.__version__.split("+", 1)[0]) + class TVMCompiler(itf.comp.Compiler): def __init__( @@ -110,10 +120,10 @@ def compile(self, schedule: itf.schd.Schedule) -> itf.comp.Module: else: packed_lib_path = lib_path emit_c_packed_base = emit_c_base - operation = op.generate() + expr_compiler = TVMExprCompiler(op, tir_schedule=self._backend._tir_schedule) + schedulable = expr_compiler.generate() if self.print_source_ir or self.save_temps: - sch = op.schedule(operation) - lowered = str(op.lower(operation, sch)) + lowered = schedulable.schedule().dumps() if self.print_source_ir: self._print(lowered) save_temp(f"{dump_base}.initial.txt", lowered) @@ -121,18 +131,20 @@ def compile(self, schedule: itf.schd.Schedule) -> itf.comp.Module: save_temp(f"{dump_base}.sched.txt", str(schedule)) if self.print_transformed_ir: self._print(schedule) - sch = op.schedule(operation, schedule) + sch = schedulable.schedule(schedule) if self.print_transformed_ir or self.save_temps: - lowered = str(op.lower(operation, sch)) + lowered = sch.dumps() if self.print_transformed_ir: self._print(lowered) save_temp(f"{dump_base}.scheduled.txt", lowered) if self.emit_c: - self._build_te_c( - op, operation, sch, func_name=packed_func_name, fname=emit_c_packed_base + self._build_c( + sch, + func_name=packed_func_name, + fname=emit_c_packed_base, ) if type in ["shlib", "arlib"]: - built = self._build_te(op, operation, sch, func_name=packed_func_name) + built = self._build(sch, func_name=packed_func_name) if self.save_temps: for idx, mod in enumerate(built._collect_dso_modules()): llvm_ir = str(mod.get_source("ll")) @@ -213,35 +225,31 @@ def compile(self, schedule: itf.schd.Schedule) -> itf.comp.Module: def _print(self, *content: Any) -> None: print(*content, flush=True, file=self.print_file) - def _build_te( + def _build( self, - op: TVMBaseExpr, - tensors: Sequence[TETensor], - sch: TESchedule, + sch: TVMScheduledExpr, func_name: str | None = None, ) -> Any: + op = sch.schedulable.expr if func_name is None: func_name = op.name return self._tvm_build_crt( sch, - tensors, cname=func_name, target=self.tvm_tgt, ) - def _build_te_c( + def _build_c( self, - op: TVMBaseExpr, - tensors: Sequence[TETensor], - sch: TESchedule, + sch: TVMScheduledExpr, func_name: str | None = None, fname: str | None = None, ) -> None: if func_name is None: - func_name = op.name + func_name = sch.schedulable.expr.name if fname is None: fname = func_name - self._tvm_emit_c(sch, tensors, self.tvm_tgt, func_name, fname) + self._tvm_emit_c(sch, self.tvm_tgt, func_name, fname) @classmethod def _get_tvm_target_options(cls, target: str, arch: str) -> str: @@ -320,16 +328,16 @@ def _tvm_build_crt_args(cls, target: str) -> dict[str, Any]: } except: runtime_kwargs = {} - target = f"{target} --system-lib --runtime=c" + if TVM_VERSION < Version("0.21"): + target = f"{target} --system-lib --runtime=c" + return { "target": target, **runtime_kwargs, } @classmethod - def _tvm_build_crt( - cls, sch: Any, tensors: Sequence[TETensor], target: str, cname: str - ) -> Any: + def _tvm_build_crt(cls, sch: TVMScheduledExpr, target: str, cname: str) -> Any: build_kwargs = cls._tvm_build_crt_args(target) config = {} if target.startswith("c "): @@ -339,20 +347,28 @@ def _tvm_build_crt( } ) with tvm.transform.PassContext(opt_level=3, config=config): - return tvm.build(sch, list(tensors), name=cname, **build_kwargs) + if isinstance(sch, TVMScheduledExprTE): + tensors = sch.schedulable._params + built = tvm.build(sch._schedule, tensors, name=cname, **build_kwargs) # type: ignore + else: + assert isinstance(sch, TVMScheduledExprTIR) + func = sch._schedule.mod[sch.schedulable.expr.name] + func = func.with_attr("global_symbol", cname) + mod = tvm.IRModule({cname: func}) + built = tvm.build(mod, **build_kwargs) + return built @classmethod def _tvm_emit_c( cls, - sch: TESchedule, - tensors: Sequence[TETensor], + sch: TVMScheduledExpr, target: str, cname: str, fname: str, ) -> Any: # Ignore initial target as of now and generate target agnostic C target = "c -keys=arch -march=generic -mcpu=generic" - built = cls._tvm_build_crt(sch, tensors, target, cname) + built = cls._tvm_build_crt(sch, target, cname) out_dir = Path(fname).parent out_base = Path(fname).stem with tempfile.TemporaryDirectory() as tmp_dir: @@ -363,7 +379,10 @@ def _tvm_emit_c( members = [info for info in tf.getmembers() if info.name.endswith(".c")] tf.extractall(tmp_dir, members=members, filter="data") out_dir.mkdir(parents=True, exist_ok=True) - shutil.copy(tmp_dir_path / "lib1.c", out_dir / f"{out_base}.c") + cfile = tmp_dir_path / "lib1.c" + if not cfile.exists(): + cfile = tmp_dir_path / "lib0.c" + shutil.copy(cfile, out_dir / f"{out_base}.c") def _cc_prefix(self) -> str: map = { diff --git a/src/xtc/backends/tvm/TVMOps.py b/src/xtc/backends/tvm/TVMOps.py index 3ef40e975..43b0d7a15 100644 --- a/src/xtc/backends/tvm/TVMOps.py +++ b/src/xtc/backends/tvm/TVMOps.py @@ -5,7 +5,7 @@ from abc import ABC, abstractmethod from collections.abc import Sequence from typing_extensions import override -from typing import Any, Type, TypeAlias +from typing import Any, Type, TypeAlias, cast from xtc.utils.math import mulall from xtc.utils.text import to_cname @@ -25,29 +25,19 @@ import tvm.te as te import tvm.topi as topi -TETensor: TypeAlias = Any # Use instead of te.Tensor to avoids type errors +TETensor: TypeAlias = te.Tensor TEIndexVar: TypeAlias = Any -TESchedule: TypeAlias = te.Schedule class TVMBaseExpr(ABC): def __init__(self, name: str) -> None: self.name = name - @abstractmethod - def generate(self) -> tuple[TETensor, ...]: ... - @abstractmethod - def schedule( - self, tensors: Sequence[TETensor], schedule: Any = None - ) -> te.Schedule: ... @abstractmethod def np_outputs_spec(self) -> list[dict[str, Any]]: ... @abstractmethod def np_inputs_spec(self) -> list[dict[str, Any]]: ... - def lower(self, tensors: Sequence[TETensor], sch: te.Schedule) -> str: - return tvm.lower(sch, list(tensors), simple_mode=True) - @classmethod def from_operation(cls, xtc_op: Operation, name: str | None) -> "TVMOperation": for typ in xtc_op.inputs_types: @@ -88,24 +78,6 @@ def __init__( name=self.operator.name if name is None else name, ) - @override - def generate(self) -> tuple[Any, ...]: - return self.operator.generate_op() - - @override - def schedule( - self, tensors: Sequence[TETensor], schedule: Any = None - ) -> te.Schedule: - sch = te.create_schedule(tensors[-1].op) - if schedule is None: - return sch - schedule_map = schedule.schedule_impl - tensors_map = {t.name: t for t in tensors} - for sched in schedule_map.values(): - if sched: - exec(sched, {"sch": sch, "obj": tensors_map}, {}) - return sch - @override def np_inputs_spec(self) -> list[dict[str, Any]]: inputs_spec = [ @@ -197,25 +169,6 @@ def _te_expr_from_graph(self) -> tuple[dict[str, TETensor], dict[str, TETensor]] } return variables, params - @override - def generate(self) -> tuple[TETensor, ...]: - self._variables, self._params = self._te_expr_from_graph() - return tuple(self._params.values()) - - @override - def schedule( - self, tensors: Sequence[TETensor], schedule: Any = None - ) -> te.Schedule: - sch = te.create_schedule(tensors[-1].op) - if schedule is None: - return sch - schedule_map = schedule.schedule_impl - tensors_map = {t.name: t for t in self._variables.values()} - for sched in schedule_map.values(): - if sched: - exec(sched, {"sch": sch, "obj": tensors_map}, {}) - return sch - @override def np_inputs_spec(self) -> list[dict[str, Any]]: inputs_spec = [ @@ -303,7 +256,7 @@ def generate_op( A = te.placeholder((Ki, Kk), name="A", dtype=dtype) B = te.placeholder((Kk, Kj), name="B", dtype=dtype) else: - A, B = inputs + A, B = cast(Sequence[Any], inputs) Ashape = tuple(A.shape) Bshape = tuple(B.shape) Anewshape = (Ashape[0], mulall(list(Ashape[1:]))) @@ -323,7 +276,7 @@ def generate_op( ), name=self.name, ) - return A, B, O + return cast(tuple[TETensor], (A, B, O)) @override def inputs_dims(self) -> tuple[tuple[int, ...], ...]: @@ -375,7 +328,7 @@ def generate_op( if inputs is None: A = te.placeholder((Ki,), name="A", dtype=dtype) else: - (A,) = inputs + (A,) = cast(Sequence[Any], inputs) shape = tuple(A.shape) size = mulall(A.shape) newshape = (size,) @@ -389,7 +342,7 @@ def generate_op( ) if shape != newshape: O = topi.reshape(O, newshape=shape) - return A, O + return cast(tuple[TETensor], (A, O)) @override def inputs_dims(self) -> tuple[tuple[int, ...], ...]: @@ -444,7 +397,7 @@ def generate_op( A = te.placeholder(inps_dims[0], name="A", dtype=dtype) W = te.placeholder(inps_dims[1], name="W", dtype=dtype) else: - A, W = inputs + A, W = cast(Sequence[Any], inputs) r = te.reduce_axis((0, Kr), "r") s = te.reduce_axis((0, Ks), "s") c = te.reduce_axis((0, Kc), "c") @@ -460,7 +413,7 @@ def generate_op( ), name=self.name, ) - return A, W, O + return cast(tuple[TETensor], (A, W, O)) @override def inputs_dims(self) -> tuple[tuple[int, ...], ...]: @@ -522,7 +475,7 @@ def generate_op( if inputs is None: A = te.placeholder(tuple(dims_values_all), name="A", dtype=dtype) else: - (A,) = inputs + (A,) = cast(Sequence[Any], inputs) def get_indexes(*args: int) -> tuple[int, ...]: indexes = list(args) @@ -563,7 +516,7 @@ def get_args_bounds(*args: int) -> list[tvm.tir.PrimExpr]: ), name=self.name, ) - return A, O + return cast(tuple[TETensor], (A, O)) @override def inputs_dims(self) -> tuple[tuple[int, ...], ...]: @@ -641,7 +594,7 @@ def generate_op( ] A = te.placeholder(tuple(dims_values_before_unpad), name="A", dtype=dtype) else: - (A,) = inputs + (A,) = cast(Sequence[Any], inputs) def get_indexes(*args: int) -> tuple[int, ...]: indexes = list(args) @@ -657,7 +610,7 @@ def get_indexes(*args: int) -> tuple[int, ...]: lambda *args: A[get_indexes(*args)], name=self.name, ) - return A, O + return cast(tuple[TETensor], (A, O)) @override def inputs_dims(self) -> tuple[tuple[int, ...], ...]: @@ -713,7 +666,7 @@ def generate_op( if inputs is None: A = te.placeholder(self.attrs["inp_shape"], name="A", dtype=dtype) else: - (A,) = inputs + (A,) = cast(Sequence[Any], inputs) shape = A.shape axes = self.attrs["axes"] if axes == (): @@ -736,7 +689,7 @@ def transpose_i(*ivs: TEIndexVar): O = te.compute((i,), transpose_i, name=self.name) if out_shape != (i,): O = topi.reshape(O, newshape=out_shape) - return A, O + return cast(tuple[TETensor], (A, O)) @override def inputs_dims(self) -> tuple[tuple[int, ...], ...]: diff --git a/src/xtc/backends/tvm/TVMOpsCompiler.py b/src/xtc/backends/tvm/TVMOpsCompiler.py new file mode 100644 index 000000000..b597990c3 --- /dev/null +++ b/src/xtc/backends/tvm/TVMOpsCompiler.py @@ -0,0 +1,159 @@ +# +# SPDX-License-Identifier: BSD-3-Clause +# Copyright (c) 2024-2026 The XTC Project Authors +# +from abc import ABC, abstractmethod +from collections.abc import Sequence +from typing_extensions import override +from typing import Any, TypeAlias + +import tvm +import tvm.te as te + +from .TVMOps import ( + TVMBaseExpr, + TVMOperation, + TVMGraph, +) + +__all__ = [ + "TVMExprCompiler", + "TVMSchedulableExpr", + "TVMSchedulableExpr", + "TVMSchedulableExprTE", + "TVMSchedulableExprTIR", + "TVMScheduledExpr", + "TVMScheduledExprTE", + "TVMScheduledExprTIR", +] + + +TETensor: TypeAlias = te.Tensor +TIRSchedule: TypeAlias = tvm.tir.Schedule +TIRFunc: TypeAlias = tvm.tir.PrimFunc +TESchedule: TypeAlias = Any # te.Schedule not available on tvm > 0.19 + + +class TVMExprCompiler: + def __init__(self, expr: TVMBaseExpr, tir_schedule: bool = True): + self._expr = expr + self._tir_schedule = tir_schedule + + def generate(self) -> "TVMSchedulableExpr": + if isinstance(self._expr, TVMGraph): + vars, params = [ + list(vars.values()) for vars in self._expr._te_expr_from_graph() + ] + else: + assert isinstance(self._expr, TVMOperation) + params = list(self._expr.operator.generate_op()) + vars = params + if self._tir_schedule: + prim_func = te.create_prim_func(params) + return TVMSchedulableExprTIR(self._expr, prim_func) + return TVMSchedulableExprTE(self._expr, params, vars) + + +class TVMSchedulableExpr(ABC): + @abstractmethod + def schedule(self, schedule: Any = None) -> "TVMScheduledExpr": ... + + @property + @abstractmethod + def expr(self) -> TVMBaseExpr: ... + + +class TVMSchedulableExprTE(TVMSchedulableExpr): + def __init__( + self, + expr: TVMBaseExpr, + params: Sequence[TETensor], + tensors: Sequence[TETensor] | None = None, + ): + self._expr = expr + self._params = list(params) + self._tensors = list(params) if tensors is None else list(tensors) + + @property + @override + def expr(self) -> TVMBaseExpr: + return self._expr + + @override + def schedule(self, schedule: Any = None) -> "TVMScheduledExprTE": + sch = te.create_schedule(self._params[-1].op) # type: ignore + if schedule is not None: + schedule_map = schedule.schedule_impl + tensors_map = {t.name: t for t in self._tensors} + for sched in schedule_map.values(): + if sched: + exec(sched, {"sch": sch, "obj": tensors_map}, {}) + return TVMScheduledExprTE(self, sch) + + +class TVMSchedulableExprTIR(TVMSchedulableExpr): + def __init__(self, expr: TVMBaseExpr, func: TIRFunc): + self._expr = expr + self._func = func + + @property + @override + def expr(self) -> TVMBaseExpr: + return self._expr + + @override + def schedule(self, schedule: Any = None) -> "TVMScheduledExprTIR": + func_name = self._expr.name + func = self._func.with_attr("global_symbol", self._expr.name) + mod = tvm.IRModule({func_name: func}) + sch = tvm.tir.Schedule(mod) + if schedule is None: + return TVMScheduledExprTIR(self, sch) + # TODO: schedule TIR + schedule_map = schedule.schedule_impl + sch.work_on(func_name) + for sched in schedule_map.values(): + if sched: + exec(sched, {"sch": sch}, {}) + return TVMScheduledExprTIR(self, sch) + + +class TVMScheduledExpr(ABC): + @property + @abstractmethod + def schedulable(self) -> TVMSchedulableExpr: ... + + @abstractmethod + def dumps(self) -> str: ... + + +class TVMScheduledExprTE(TVMScheduledExpr): + def __init__(self, schedulable: TVMSchedulableExprTE, schedule: TESchedule): + self._schedulable = schedulable + self._schedule = schedule + + @property + @override + def schedulable(self) -> TVMSchedulableExprTE: + return self._schedulable + + @override + def dumps(self) -> str: + return str( + tvm.lower(self._schedule, self._schedulable._params, simple_mode=True) # type: ignore + ) + + +class TVMScheduledExprTIR(TVMScheduledExpr): + def __init__(self, schedulable: TVMSchedulableExprTIR, schedule: TIRSchedule): + self._schedulable = schedulable + self._schedule = schedule + + @property + @override + def schedulable(self) -> TVMSchedulableExprTIR: + return self._schedulable + + @override + def dumps(self) -> str: + return str(self._schedule.mod) diff --git a/src/xtc/backends/tvm/TVMScheduler.py b/src/xtc/backends/tvm/TVMScheduler.py index 165525708..911360cea 100644 --- a/src/xtc/backends/tvm/TVMScheduler.py +++ b/src/xtc/backends/tvm/TVMScheduler.py @@ -3,16 +3,18 @@ # Copyright (c) 2024-2026 The XTC Project Authors # import sys +from abc import ABC, abstractmethod from typing_extensions import override from typing import TextIO, TypeAlias from io import StringIO import numpy as np from dataclasses import dataclass from copy import deepcopy +import functools from xtc.utils.math import pow2divisor 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.tvm as backend import xtc.itf as itf @@ -36,7 +38,12 @@ class TVMPlainSchedule: fused: list[tuple[str, int]] -class TVMScheduleEmitter: +class TVMScheduleEmitter(ABC): + @abstractmethod + def emit(self, scheduler: "TVMScheduler"): ... + + +class TVMScheduleEmitterTE(TVMScheduleEmitter): def __init__( self, op: TVMOperation, @@ -303,12 +310,196 @@ def _dump_schedule(self, sched: TVMPlainSchedule): file=outf, ) - def emit(self, sched: TVMPlainSchedule): + @override + def emit(self, scheduler: "TVMScheduler"): + sched = scheduler._get_plain_schedule() # First adjust schedule to fix code gen limitations before emit sched = self._update_schedule_for_codegen(sched) self._dump_schedule(sched) +class TVMScheduleEmitterTIR(TVMScheduleEmitter): + def __init__( + self, + op: TVMOperation, + obj_var: str = "obj", + sch_var: str = "sch", + outf: TextIO = sys.stdout, + ): + self._op = op + self._obj_var = obj_var + self._sch_var = sch_var + self._outf = outf + + def _cache_read_factor_offset( + self, input_idx: int, pad: bool + ) -> tuple[int, int, int]: + if not pad: + return 0, 0, 0 + input_spec = self._op.np_inputs_spec()[input_idx] + if len(input_spec["shape"]) < 2: + return 0, 0, 0 + # Assume for CPU common number of sets and line size for L1 + # Except to minimize conflicts by setting the inner axis + # size to a factor of num_sets and adding +1 + num_sets, line_size = 64, 64 + elt_size = np.dtype(input_spec["dtype"]).itemsize + elts_per_line = line_size // elt_size + return -2, elts_per_line * num_sets, elts_per_line + + def _dump_schedule(self, sched: LoopNest): + root = sched.root_node + if root is None: + return + self._dump_schedule_node(sched, root) + + def _dump_schedule_node(self, sched: LoopNest, node: LoopNestNode): + assert node is not None, "unexpected undefined node" + assert not node.splits, "node split not implemented for this backend" + sch = self._sch_var + outf = self._outf + dims = sched.abstract_dims + block = "O" + print(f'{block} = {sch}.get_block("{self._op.name}")', file=outf) + print(f"{', '.join(dims)}, = {sch}.get_loops({block})", file=outf) + if node.fuse_consumer_at: + print(f"O_F0 = {sch}.get_consumers({block})[0]", file=outf) + if node.pack_at: + inputs = list({inp[0]: None for inp in node.pack_at.values()}) + for inp_idx in inputs: + print( + f'I_R{inp_idx} = {sch}.cache_read({block}, {inp_idx}, "global")', + file=outf, + ) + if node.buffer_at: + print(f'O_W0 = {sch}.cache_write({block}, 0, "global")', file=outf) + if node.fuse_producer_at: + producers = list({idx: None for idx in node.fuse_producer_at.values()}) + for prod_idx in producers: + print( + f"I_F{prod_idx} = {sch}.get_producers({block})[{prod_idx}]", + file=outf, + ) + for t_axis, t_tiles in [(k, v) for k, v in node.tiles.items() if v]: + t_names = [t_axis] + list(t_tiles) + factors = functools.reduce( + lambda acc, x: acc + [x // acc[-1]], reversed(t_tiles.values()), [1] + ) + t_factors = ["None"] + [str(f) for f in factors[:0:-1]] + print( + f"{', '.join(t_names)}, = {sch}.split({t_axis}, factors=[{', '.join(t_factors)}])", + file=outf, + ) + print(f"{sch}.reorder({', '.join(node.interchange)})", file=outf) + if node.buffer_at: + for axis in node.buffer_at: + print(f"{sch}.reverse_compute_at(O_W0, {axis})", file=outf) + if node.pack_at: + for axis, (inp_idx, mtype, pad) in node.pack_at.items(): + print(f"{sch}.compute_at(I_R{inp_idx}, {axis})", file=outf) + dim, factor, offset = self._cache_read_factor_offset(inp_idx, pad) + if factor != 0: + print( + f"{sch}.storage_align(I_R{inp_idx}, 0, ", + f"axis={dim}, factor={factor}, offset={offset})", + file=outf, + ) + if node.fuse_producer_at: + for axis, prod_idx in node.fuse_producer_at.items(): + print(f"{sch}.compute_at(I_F{prod_idx}, {axis})", file=outf) + if node.fuse_consumer_at: + for axis in node.fuse_consumer_at: + print(f"{sch}.reverse_compute_at(O_F0, {axis})", file=outf) + for u_axis, u_factor in node.unroll.items(): + print(f"{sch}.unroll({u_axis})", file=outf) + for v_axis in node.vectorize: + print(f"{sch}.vectorize({v_axis})", file=outf) + if node.parallelize: + if len(node.parallelize) > 1: + print( + f"{node.parallelize[-1]} = {sch}.fuse({', '.join(node.parallelize)})", + file=outf, + ) + print( + f"{sch}.parallel({node.parallelize[-1]})", + file=outf, + ) + + @classmethod + def _update_loopnest_for_codegen(cls, sched: LoopNest): + def _update_loopnode(node: LoopNestNode) -> LoopNestNode: + adjusted_tiles = {} + adjusted_unrolling = { + k: v for k, v in node.unroll.items() if k not in node.vectorize + } + adjusted_unrolls = list(adjusted_unrolling) + adjusted_vectorization = node.vectorize[:] + adjusted_permutation = node.interchange[:] + for dim, dim_tiles in node.tiles.items(): + adjusted_dim_tiles = {} + for axis, size in dim_tiles.items(): + adjusted_dim_tiles.update({axis: size}) + if axis in adjusted_unrolling: + assert axis not in adjusted_vectorization + unroll = adjusted_unrolling[axis] + if unroll < size: + axis_idx = adjusted_unrolls.index(axis) + new_axis = f"__u_{axis}" + adjusted_dim_tiles.update({new_axis: unroll}) + adjusted_unrolls[axis_idx] = new_axis + del adjusted_unrolling[axis] + adjusted_unrolling.update({new_axis: unroll}) + adjusted_permutation.insert( + adjusted_permutation.index(axis) + 1, + new_axis, + ) + elif axis in adjusted_vectorization: + assert axis not in adjusted_unrolling + pow2 = pow2divisor(size) + unroll = size // pow2 + if unroll > 1: + axis_idx = adjusted_vectorization.index(axis) + new_axis = f"__v_{axis}" + adjusted_dim_tiles.update({new_axis: pow2}) + adjusted_vectorization[axis_idx] = new_axis + adjusted_unrolls.append(axis) + adjusted_unrolling.update({axis: unroll}) + adjusted_permutation.insert( + adjusted_permutation.index(axis) + 1, + new_axis, + ) + adjusted_tiles[dim] = adjusted_dim_tiles + adjusted_unrolling = {u: adjusted_unrolling[u] for u in adjusted_unrolls} + return LoopNestNode( + root=node.root, + tiles=adjusted_tiles, + splits=deepcopy(node.splits), + interchange=adjusted_permutation, + vectorize=adjusted_vectorization, + parallelize=deepcopy(node.parallelize), + unroll=adjusted_unrolling, + buffer_at=deepcopy(node.buffer_at), + pack_at=deepcopy(node.pack_at), + fuse_producer_at=deepcopy(node.fuse_producer_at), + fuse_consumer_at=deepcopy(node.fuse_consumer_at), + ) + + root = sched.root_node + if root is not None: + root = _update_loopnode(root) + return LoopNest( + abstract_dims=sched.abstract_dims, + root_node=root, + ) + + @override + def emit(self, scheduler: "TVMScheduler"): + sched = scheduler.get_loop_nest() + sched = self._update_loopnest_for_codegen(sched) + sched.check() + self._dump_schedule(sched) + + class TVMScheduler(itf.schd.Scheduler): def __init__( self, @@ -343,6 +534,7 @@ def __init__( self.write_caches: list[str] = [] self.read_buffers: list[tuple[str, int, bool]] = [] self.fused: list[tuple[str, int]] = [] + self.fused_consumers: list[str] = [] self._update_loops() @property @@ -365,9 +557,11 @@ def backend(self) -> itf.back.Backend: @override def schedule(self) -> itf.schd.Schedule: io = StringIO() - emitter = TVMScheduleEmitter(op=self._op, outf=io) - schedule = self._get_plain_schedule() - emitter.emit(schedule) + if self._backend._tir_schedule: + emitter: TVMScheduleEmitter = TVMScheduleEmitterTIR(op=self._op, outf=io) + else: + emitter = TVMScheduleEmitterTE(op=self._op, outf=io) + emitter.emit(self) sched = io.getvalue() assert self._op.name is not None schedule_impl = {self._op.name: sched} @@ -467,6 +661,10 @@ def fuse_producer_at( assert input_idx >= 0 and input_idx < len(self._op.np_inputs_spec()) self.fused.append((axis, input_idx)) + @override + def fuse_consumer_at(self, axis: str, root: str = DEFAULT_ROOT) -> None: + self.fused_consumers.append(axis) + @override def define_memory_mesh(self, axes: dict[str, int]) -> None: # TODO: not implemented for now @@ -498,11 +696,11 @@ def distributed_buffer_at( def _get_plain_schedule(self) -> TVMPlainSchedule: return TVMPlainSchedule( dims=deepcopy(self.dims), - tiles=self.tiles, - permutation=self.permutation, + tiles=deepcopy(self.tiles), + permutation=deepcopy(self.permutation), parallelization=deepcopy(self.parallelization), - unrolling=self.unrolling, - vectorization=self.vectorization, + unrolling=deepcopy(self.unrolling), + vectorization=deepcopy(self.vectorization), write_caches=deepcopy(self.write_caches), read_buffers=deepcopy(self.read_buffers), fused=deepcopy(self.fused), @@ -543,6 +741,12 @@ def get_loop_nest(self) -> LoopNest: axis: (input_idx, None, pad) for axis, input_idx, pad in self.read_buffers } + # Build fuse_producer_at mapping + root_node.fuse_producer_at = dict(self.fused) + + # Build fuse_consumer_at list + root_node.fuse_consumer_at = list(self.fused_consumers) + return loop_nest diff --git a/src/xtc/cli/explore.py b/src/xtc/cli/explore.py index 6c9ecbe2f..68cf34683 100644 --- a/src/xtc/cli/explore.py +++ b/src/xtc/cli/explore.py @@ -277,6 +277,12 @@ def main(): default=defaults.use_tensors, help="use tensors instead of memref for the mlir backend", ) + parser.add_argument( + "--tir-schedule", + action=argparse.BooleanOptionalAction, + default=defaults.tir_schedule, + help="use TIR schedule instead of TE schedule for tvm backend", + ) parser.add_argument( "--batch", type=int, default=defaults.batch, help="batch size for optimizer" ) diff --git a/src/xtc/itf/schd/scheduler.py b/src/xtc/itf/schd/scheduler.py index 0f6c51d98..97e17a4f2 100644 --- a/src/xtc/itf/schd/scheduler.py +++ b/src/xtc/itf/schd/scheduler.py @@ -232,6 +232,19 @@ def fuse_producer_at( """ ... + @abstractmethod + def fuse_consumer_at(self, axis: str, root: str = DEFAULT_ROOT) -> None: + """Fuse the consumer computation at the given producer location. + + The consumer of output zero is fused at the given scheduled producer + axis. Other outputs are not currently supported. + + Args: + axis: localisation of the fusion in the producer + root: the parent split (or the operator's absolute root) + """ + ... + @abstractmethod def define_memory_mesh(self, axes: dict[str, int]) -> None: """Define a memory mesh. diff --git a/src/xtc/schedules/loop_nest.py b/src/xtc/schedules/loop_nest.py index 306c0c910..701baf844 100644 --- a/src/xtc/schedules/loop_nest.py +++ b/src/xtc/schedules/loop_nest.py @@ -97,6 +97,9 @@ class LoopNestNode(Node["LoopNestNode"]): pack_at: Pack configuration per axis. Maps axis names to tuples of (input_idx, mtype, pad). input_idx is the input buffer index, mtype is the memory type (None for default), pad enables padding. + fuse_producer_at: Producer fusion configuration per axis. Maps axis + names to producer indices. + fuse_consumer_at: List of axes where the output consumer is fused. """ root: str @@ -108,6 +111,8 @@ class LoopNestNode(Node["LoopNestNode"]): unroll: dict[str, int] = field(default_factory=dict) buffer_at: dict[str, str | None] = field(default_factory=dict) pack_at: dict[str, tuple[int, str | None, bool]] = field(default_factory=dict) + fuse_producer_at: dict[str, int] = field(default_factory=dict) + fuse_consumer_at: list[str] = field(default_factory=list) def pretty_print(self, indent: int = 0) -> str: """Return a human-readable representation of the loop nest. @@ -208,7 +213,7 @@ def pretty_print(self, indent: int = 0) -> str: return "\n".join(lines) def _add_annotations(self, line: str, loop_name: str) -> str: - """Add annotations (parallelized, vectorized, unroll, buffer, pack) to a loop line.""" + """Add loop annotations to a loop line.""" annotations: list[str] = [] if loop_name in self.parallelize: annotations.append("parallelized") @@ -230,6 +235,11 @@ def _add_annotations(self, line: str, loop_name: str) -> str: if pad: parts.append("pad") annotations.append(f"pack({', '.join(parts)})") + if loop_name in self.fuse_producer_at: + prod_idx = self.fuse_producer_at[loop_name] + annotations.append(f"fuse_producer({prod_idx})") + if loop_name in self.fuse_consumer_at: + annotations.append("fuse_consumer") if annotations: line += " // " + ", ".join(annotations) return line diff --git a/src/xtc/search/explore.py b/src/xtc/search/explore.py index 1cbfbc5ef..e47bcc7a8 100644 --- a/src/xtc/search/explore.py +++ b/src/xtc/search/explore.py @@ -111,6 +111,7 @@ class ExplorationConfig: results: list[Sequence] = field(default_factory=list) descript: str | None = None use_tensors: bool = False + tir_schedule: bool = False progress_cls: str = "tqdm" def __post_init__(self): @@ -349,10 +350,15 @@ def compile_one( args = self.config assert isinstance(in_x, list), f"X not a list: {in_x} ({type(in_x)})" logger.debug("Compile: %s: %s: %s...", ident, backend, in_x) + kwargs = {} + if backend == "tvm": + kwargs.update({"tir_schedule": args.tir_schedule}) + if backend == "mlir": + kwargs.update({"use_tensor_dialect": args.use_tensors}) impl, backend_name = self.graph_implementer( graph, backend, - use_tensor_dialect=args.use_tensors, + **kwargs, ) assert backend_name == backend scheduler = impl.get_scheduler() diff --git a/tests/filecheck/backends/test_matmul_relu_fused_tvm_tir.py b/tests/filecheck/backends/test_matmul_relu_fused_tvm_tir.py new file mode 100644 index 000000000..86ce8cb10 --- /dev/null +++ b/tests/filecheck/backends/test_matmul_relu_fused_tvm_tir.py @@ -0,0 +1,45 @@ +# RUN: python %s 2>&1 | filecheck %s +# REQUIRES: module_tvm + +import xtc.graphs.xtc.op as O +from xtc.backends.tvm import Backend + +I, J, K, dtype = 4, 32, 512, "float32" +a = O.tensor((I, K), dtype, name="A") +b = O.tensor((K, J), dtype, name="B") + +with O.graph(name="matmul_relu") as gb: + m = O.matmul(a, b, name="matmul") + O.relu(m, name="relu") + +graph = gb.graph +print(graph) + +impl = Backend(graph, tir_schedule=True) + +sch = impl.get_scheduler(default_node="matmul") +sch.tile("i", {"i1": 2}) +sch.tile("j", {"j1": 16}) +sch.interchange(["i", "j", "i1", "j1", "k"]) +sch.fuse_consumer_at("j1") +sched = sch.schedule() + +comp = impl.get_compiler( + shared_lib=True, + dump_file="matmul_relu_fused_tvm_tir", + print_source_ir=True, + print_transformed_ir=True, +) +module = comp.compile(sched) +executor = module.get_executor(validate=True) +res = executor.execute() +print(f"CODE: {res}") + +# CHECK: O = sch.get_block("matmul") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_F0 = sch.get_consumers(O)[0] +# CHECK-NEXT: i, i1, = sch.split(i, factors=[None, 2]) +# CHECK-NEXT: j, j1, = sch.split(j, factors=[None, 16]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k) +# CHECK-NEXT: sch.reverse_compute_at(O_F0, j1) +# CHECK: CODE: 0 diff --git a/tests/filecheck/backends/test_matmul_tvm_tir.py b/tests/filecheck/backends/test_matmul_tvm_tir.py new file mode 100644 index 000000000..401806b66 --- /dev/null +++ b/tests/filecheck/backends/test_matmul_tvm_tir.py @@ -0,0 +1,128 @@ +# RUN: python %s 2>&1 | filecheck %s +# REQUIRES: module_tvm + +import xtc.graphs.xtc.op as O +from xtc.backends.tvm import Backend + +I, J, K, dtype = 64, 192, 256, "float32" +a = O.tensor((I, K), dtype, name="A") +b = O.tensor((K, J), dtype, name="B") + +with O.graph(name="matmul") as gb: + O.matmul(a, b, name="C") + +graph = gb.graph +print(graph) + +impl = Backend(graph, tir_schedule=True) + +sch = impl.get_scheduler() +sch.tile("i", {"i1": 8, "i2": 4}) +sch.tile("j", {"j1": 96, "j2": 48}) +sch.tile("k", {"k1": 16}) +sch.interchange(["i", "j", "k", "i1", "j1", "k1", "i2", "j2"]) +sch.buffer_at("j") +sch.pack_at("k", 1, pad=True) +sch.vectorize(["j2"]) +sch.unroll({"i2": 2}) +sch.parallelize(["i", "j"]) +sched = sch.schedule() + +comp = impl.get_compiler( + shared_lib=True, + dump_file="matmul_tvm_tir", + print_source_ir=True, + print_transformed_ir=True, +) +module = comp.compile(sched) +executor = module.get_executor(validate=True) +res = executor.execute() +print(f"CODE: {res}") + +# CHECK: graph: +# CHECK-NEXT: name: matmul +# CHECK-NEXT: inputs: +# CHECK-NEXT: - %0 : 64x256xfloat32 +# CHECK-NEXT: - %1 : 256x192xfloat32 +# CHECK-NEXT: outputs: +# CHECK-NEXT: - %2 : 64x192xfloat32 +# CHECK-NEXT: nodes: +# CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [64x256xfloat32, 256x192xfloat32] -> [64x192xfloat32] +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: def matmul(_0: T.Buffer((64, 256), "float32"), _1: T.Buffer((256, 192), "float32"), C: T.Buffer((64, 192), "float32")): +# CHECK-NEXT: T.func_attr({"tir.noalias": T.bool(True)}) +# CHECK-NEXT: # with T.block("root"): +# CHECK-NEXT: for i, j, k in T.grid(64, 192, 256): +# CHECK-NEXT: with T.block("C"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(_0[v_i, v_k], _1[v_k, v_j]) +# CHECK-NEXT: T.writes(C[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: C[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: C[v_i, v_j] = C[v_i, v_j] + _0[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: O = sch.get_block("C") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: I_R1 = sch.cache_read(O, 1, "global") +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, __u_i2, = sch.split(i, factors=[None, 4, 2, 2]) +# CHECK-NEXT: j, j1, j2, __v_j2, = sch.split(j, factors=[None, 32, 3, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 16]) +# CHECK-NEXT: sch.reorder(i, j, k, i1, j1, k1, i2, __u_i2, j2, __v_j2) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j) +# CHECK-NEXT: sch.compute_at(I_R1, k) +# CHECK-NEXT: sch.storage_align(I_R1, 0, axis=-2, factor=1024, offset=16) +# CHECK-NEXT: sch.unroll(__u_i2) +# CHECK-NEXT: sch.unroll(j2) +# CHECK-NEXT: sch.vectorize(__v_j2) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: def matmul(_0: T.Buffer((64, 256), "float32"), _1: T.Buffer((256, 192), "float32"), C: T.Buffer((64, 192), "float32")): +# CHECK-NEXT: T.func_attr({"tir.noalias": T.bool(True)}) +# CHECK-NEXT: # with T.block("root"): +# CHECK-NEXT: _1_global = T.alloc_buffer((256, 192)) +# CHECK-NEXT: C_global = T.alloc_buffer((64, 192)) +# CHECK-NEXT: for i_0_j_0_fused in T.parallel(4): +# CHECK-NEXT: for k_0 in range(16): +# CHECK-NEXT: for ax0, ax1 in T.grid(16, 192): +# CHECK-NEXT: with T.block("_1_global"): +# CHECK-NEXT: v0 = T.axis.spatial(256, k_0 * 16 + ax0) +# CHECK-NEXT: v1 = T.axis.spatial(192, ax1) +# CHECK-NEXT: T.reads(_1[v0, v1]) +# CHECK-NEXT: T.writes(_1_global[v0, v1]) +# CHECK-NEXT: T.block_attr({"buffer_dim_align": [[0, 0, 1024, 16]]}) +# CHECK-NEXT: _1_global[v0, v1] = _1[v0, v1] +# CHECK-NEXT: for i_1, j_1, k_1, i_2 in T.grid(4, 32, 16, 2): +# CHECK-NEXT: for i_3 in T.unroll(2): +# CHECK-NEXT: for j_2 in T.unroll(3): +# CHECK-NEXT: for j_3 in T.vectorized(16): +# CHECK-NEXT: with T.block("C"): +# CHECK-NEXT: v_i = T.axis.spatial(64, i_0_j_0_fused * 16 + i_1 * 4 + i_2 * 2 + i_3) +# CHECK-NEXT: v_j = T.axis.spatial(192, j_1 * 48 + j_2 * 16 + j_3) +# CHECK-NEXT: v_k = T.axis.reduce(256, k_0 * 16 + k_1) +# CHECK-NEXT: T.where(((T.Mul(0, 32) + j_1) * 3 + j_2) * 16 + j_3 < 192) +# CHECK-NEXT: T.reads(_0[v_i, v_k], _1_global[v_k, v_j]) +# CHECK-NEXT: T.writes(C_global[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: C_global[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: C_global[v_i, v_j] = C_global[v_i, v_j] + _0[v_i, v_k] * _1_global[v_k, v_j] +# CHECK-NEXT: for ax0, ax1 in T.grid(16, 192): +# CHECK-NEXT: with T.block("C_global"): +# CHECK-NEXT: v0 = T.axis.spatial(64, i_0_j_0_fused * 16 + ax0) +# CHECK-NEXT: v1 = T.axis.spatial(192, ax1) +# CHECK-NEXT: T.reads(C_global[v0, v1]) +# CHECK-NEXT: T.writes(C[v0, v1]) +# CHECK-NEXT: C[v0, v1] = C_global[v0, v1] +# CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/test_pad_conv2d_relu_fused_tvm_tir.py b/tests/filecheck/backends/test_pad_conv2d_relu_fused_tvm_tir.py new file mode 100644 index 000000000..7c587877c --- /dev/null +++ b/tests/filecheck/backends/test_pad_conv2d_relu_fused_tvm_tir.py @@ -0,0 +1,156 @@ +# RUN: python %s 2>&1 | filecheck %s +# REQUIRES: module_tvm + +import xtc.graphs.xtc.op as O +from xtc.backends.tvm import TVMBackend as Backend + +# Small conv2d +N, H, W, F, R, S, C, SH, SW, dtype = 1, 8, 8, 16, 5, 5, 3, 2, 2, "float32" +a = O.tensor((N, H, W, C), dtype, name="I") +b = O.tensor((R, S, C, F), dtype, name="W") + +with O.graph(name="pad_conv2d_nhwc_mini") as gb: + p = O.pad2d(a, padding=2, name="pad") + c = O.conv2d(p, b, stride=(SH, SW), name="conv") + O.relu(c, name="relu") + +graph = gb.graph +print(graph) + +impl = Backend(graph, tir_schedule=True) + +sch = impl.get_scheduler(default_node="conv") +sch.interchange(["b", "h", "w", "r", "s", "c", "f"]) +sch.fuse_producer_at("r", 0) +sch.fuse_consumer_at("w") +sch.vectorize(["f"]) +sched = sch.schedule() +comp = impl.get_compiler( + shared_lib=True, + dump_file="pad_conv2d_relu_fused_tvm_tir", + print_source_ir=True, + print_transformed_ir=True, +) +module = comp.compile(sched) +executor = module.get_executor(validate=True) +res = executor.execute() +print(f"CODE: {res}") + +# CHECK: graph: +# CHECK-NEXT: name: pad_conv2d_nhwc_mini +# CHECK-NEXT: inputs: +# CHECK-NEXT: - %0 : 1x8x8x3xfloat32 +# CHECK-NEXT: - %1 : 5x5x3x16xfloat32 +# CHECK-NEXT: outputs: +# CHECK-NEXT: - %4 : 1x4x4x16xfloat32 +# CHECK-NEXT: nodes: +# CHECK-NEXT: - %2: pad2d(%0, padding={-3: (2, 2), -2: (2, 2)}, constant_value=0) {name = 'pad'} : [1x8x8x3xfloat32] -> [1x12x12x3xfloat32] +# CHECK-NEXT: - %3: conv2d(%2, %1, stride=(2, 2)) {name = 'conv'} : [1x12x12x3xfloat32, 5x5x3x16xfloat32] -> [1x4x4x16xfloat32] +# CHECK-NEXT: - %4: relu(%3) {name = 'relu'} : [1x4x4x16xfloat32] -> [1x4x4x16xfloat32] +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), T_reshape: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tir.noalias": T.bool(True)}) +# CHECK-NEXT: # with T.block("root"): +# CHECK-NEXT: pad = T.alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: conv = T.alloc_buffer((1, 4, 4, 16)) +# CHECK-NEXT: T_reshape_1 = T.alloc_buffer((256,)) +# CHECK-NEXT: relu = T.alloc_buffer((256,)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.block("pad"): +# CHECK-NEXT: v_i0, v_i1, v_i2, v_i3 = T.axis.remap("SSSS", [i0, i1, i2, i3]) +# CHECK-NEXT: T.reads(_0[v_i0, v_i1 - 2, v_i2 - 2, v_i3]) +# CHECK-NEXT: T.writes(pad[v_i0, v_i1, v_i2, v_i3]) +# CHECK-NEXT: pad[v_i0, v_i1, v_i2, v_i3] = T.if_then_else(2 <= v_i1 and v_i1 < 10 and 2 <= v_i2 and v_i2 < 10, _0[v_i0, v_i1 - 2, v_i2 - 2, v_i3], T.float32(0.0)) +# CHECK-NEXT: for b, h, w, f, r, s, c in T.grid(1, 4, 4, 16, 5, 5, 3): +# CHECK-NEXT: with T.block("conv"): +# CHECK-NEXT: v_b, v_h, v_w, v_f, v_r, v_s, v_c = T.axis.remap("SSSSRRR", [b, h, w, f, r, s, c]) +# CHECK-NEXT: T.reads(pad[v_b, v_h * 2 + v_r, v_w * 2 + v_s, v_c], _1[v_r, v_s, v_c, v_f]) +# CHECK-NEXT: T.writes(conv[v_b, v_h, v_w, v_f]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: conv[v_b, v_h, v_w, v_f] = T.float32(0.0) +# CHECK-NEXT: conv[v_b, v_h, v_w, v_f] = conv[v_b, v_h, v_w, v_f] + pad[v_b, v_h * 2 + v_r, v_w * 2 + v_s, v_c] * _1[v_r, v_s, v_c, v_f] +# CHECK-NEXT: for ax0 in range(256): +# CHECK-NEXT: with T.block("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(256, ax0) +# CHECK-NEXT: T.reads(conv[0, v_ax0 % 256 // 64, v_ax0 % 64 // 16, v_ax0 % 16]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0]) +# CHECK-NEXT: T_reshape_1[v_ax0] = conv[0, v_ax0 % 256 // 64, v_ax0 % 64 // 16, v_ax0 % 16] +# CHECK-NEXT: for i in range(256): +# CHECK-NEXT: with T.block("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(256, i) +# CHECK-NEXT: T.reads(T_reshape_1[v_i]) +# CHECK-NEXT: T.writes(relu[v_i]) +# CHECK-NEXT: relu[v_i] = T.max(T.float32(0.0), T_reshape_1[v_i]) +# CHECK-NEXT: for ax0, ax1, ax2, ax3 in T.grid(1, 4, 4, 16): +# CHECK-NEXT: with T.block("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1, v_ax2, v_ax3 = T.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) +# CHECK-NEXT: T.reads(relu[(v_ax1 * 64 + v_ax2 * 16 + v_ax3) % 256]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0, v_ax1, v_ax2, v_ax3]) +# CHECK-NEXT: T_reshape[v_ax0, v_ax1, v_ax2, v_ax3] = relu[(v_ax1 * 64 + v_ax2 * 16 + v_ax3) % 256] +# CHECK-NEXT: O = sch.get_block("conv") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_F0 = sch.get_consumers(O)[0] +# CHECK-NEXT: I_F0 = sch.get_producers(O)[0] +# CHECK-NEXT: sch.reorder(b, h, w, r, s, c, f) +# CHECK-NEXT: sch.compute_at(I_F0, r) +# CHECK-NEXT: sch.reverse_compute_at(O_F0, w) +# CHECK-NEXT: sch.vectorize(f) +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), T_reshape: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tir.noalias": T.bool(True)}) +# CHECK-NEXT: # with T.block("root"): +# CHECK-NEXT: pad = T.alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: conv = T.alloc_buffer((1, 4, 4, 16)) +# CHECK-NEXT: T_reshape_1 = T.alloc_buffer((256,)) +# CHECK-NEXT: relu = T.alloc_buffer((256,)) +# CHECK-NEXT: for b, h, w in T.grid(1, 4, 4): +# CHECK-NEXT: for r in range(5): +# CHECK-NEXT: for ax0, ax1 in T.grid(5, 3): +# CHECK-NEXT: with T.block("pad"): +# CHECK-NEXT: v_i0 = T.axis.spatial(1, 0) +# CHECK-NEXT: v_i1 = T.axis.spatial(12, h * 2 + r) +# CHECK-NEXT: v_i2 = T.axis.spatial(12, w * 2 + ax0) +# CHECK-NEXT: v_i3 = T.axis.spatial(3, ax1) +# CHECK-NEXT: T.reads(_0[v_i0, v_i1 - 2, v_i2 - 2, v_i3]) +# CHECK-NEXT: T.writes(pad[v_i0, v_i1, v_i2, v_i3]) +# CHECK-NEXT: pad[v_i0, v_i1, v_i2, v_i3] = T.if_then_else(2 <= v_i1 and v_i1 < 10 and 2 <= v_i2 and v_i2 < 10, _0[v_i0, v_i1 - 2, v_i2 - 2, v_i3], T.float32(0.0)) +# CHECK-NEXT: for s, c in T.grid(5, 3): +# CHECK-NEXT: for f in T.vectorized(16): +# CHECK-NEXT: with T.block("conv"): +# CHECK-NEXT: v_b, v_h, v_w, v_f, v_r, v_s, v_c = T.axis.remap("SSSSRRR", [b, h, w, f, r, s, c]) +# CHECK-NEXT: T.reads(pad[v_b, v_h * 2 + v_r, v_w * 2 + v_s, v_c], _1[v_r, v_s, v_c, v_f]) +# CHECK-NEXT: T.writes(conv[v_b, v_h, v_w, v_f]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: conv[v_b, v_h, v_w, v_f] = T.float32(0.0) +# CHECK-NEXT: conv[v_b, v_h, v_w, v_f] = conv[v_b, v_h, v_w, v_f] + pad[v_b, v_h * 2 + v_r, v_w * 2 + v_s, v_c] * _1[v_r, v_s, v_c, v_f] +# CHECK-NEXT: for ax0 in range(16): +# CHECK-NEXT: with T.block("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(256, h * 64 + w * 16 + ax0) +# CHECK-NEXT: T.reads(conv[0, v_ax0 % 256 // 64, v_ax0 % 64 // 16, v_ax0 % 16]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0]) +# CHECK-NEXT: T_reshape_1[v_ax0] = conv[0, v_ax0 % 256 // 64, v_ax0 % 64 // 16, v_ax0 % 16] +# CHECK-NEXT: for i in range(256): +# CHECK-NEXT: with T.block("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(256, i) +# CHECK-NEXT: T.reads(T_reshape_1[v_i]) +# CHECK-NEXT: T.writes(relu[v_i]) +# CHECK-NEXT: relu[v_i] = T.max(T.float32(0.0), T_reshape_1[v_i]) +# CHECK-NEXT: for ax0, ax1, ax2, ax3 in T.grid(1, 4, 4, 16): +# CHECK-NEXT: with T.block("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1, v_ax2, v_ax3 = T.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) +# CHECK-NEXT: T.reads(relu[(v_ax1 * 64 + v_ax2 * 16 + v_ax3) % 256]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0, v_ax1, v_ax2, v_ax3]) +# CHECK-NEXT: T_reshape[v_ax0, v_ax1, v_ax2, v_ax3] = relu[(v_ax1 * 64 + v_ax2 * 16 + v_ax3) % 256] +# CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/test_relu_matmul_fused_tvm.py b/tests/filecheck/backends/test_relu_matmul_fused_tvm.py new file mode 100644 index 000000000..6b072a4af --- /dev/null +++ b/tests/filecheck/backends/test_relu_matmul_fused_tvm.py @@ -0,0 +1,158 @@ +# RUN: python %s 2>&1 | filecheck %s +# REQUIRES: module_tvm + +import xtc.graphs.xtc.op as O +from xtc.backends.tvm import Backend + +I, J, K, dtype = 64, 64, 64, "float32" +a = O.tensor((I, K), dtype, name="A") +b = O.tensor((K, J), dtype, name="B") + +with O.graph(name="matmul") as gb: + p = O.relu(a, name="relu") + O.matmul(p, b, name="C") + +graph = gb.graph +print(graph) + +impl = Backend(graph) + +sch = impl.get_scheduler() +sch.tile("i", {"i1": 8, "i2": 4}) +sch.tile("j", {"j1": 32, "j2": 16}) +sch.tile("k", {"k1": 16}) +sch.interchange(["i", "j", "k", "i1", "j1", "k1", "i2", "j2"]) +sch.buffer_at("j") +sch.pack_at("k", 1, pad=True) +sch.fuse_producer_at("k", 0) +sch.vectorize(["j2"]) +sch.unroll({"i2": 4}) +sch.parallelize(["i", "j"]) +sched = sch.schedule() + +comp = impl.get_compiler( + shared_lib=True, + dump_file="relu_matmul_tvm_fused", + print_source_ir=True, + print_transformed_ir=True, +) +module = comp.compile(sched) +executor = module.get_executor(validate=True) +res = executor.execute() +print(f"CODE: {res}") + + +# CHECK: graph: +# CHECK-NEXT: name: matmul +# CHECK-NEXT: inputs: +# CHECK-NEXT: - %0 : 64x64xfloat32 +# CHECK-NEXT: - %1 : 64x64xfloat32 +# CHECK-NEXT: outputs: +# CHECK-NEXT: - %3 : 64x64xfloat32 +# CHECK-NEXT: nodes: +# CHECK-NEXT: - %2: relu(%0) {name = 'relu'} : [64x64xfloat32] -> [64x64xfloat32] +# CHECK-NEXT: - %3: matmul(%2, %1) {name = 'C'} : [64x64xfloat32, 64x64xfloat32] -> [64x64xfloat32] +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: def main(_0: T.Buffer((64, 64), "float32"), _1: T.Buffer((64, 64), "float32"), C: T.Buffer((64, 64), "float32")): +# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) +# CHECK-NEXT: T_reshape = T.allocate([4096], "float32", "global") +# CHECK-NEXT: T_reshape_1 = T.Buffer((4096,), data=T_reshape) +# CHECK-NEXT: for ax0 in range(4096): +# CHECK-NEXT: _0_1 = T.Buffer((4096,), data=_0.data) +# CHECK-NEXT: T_reshape_1[ax0] = _0_1[ax0] +# CHECK-NEXT: for i in range(4096): +# CHECK-NEXT: T_reshape_2 = T.Buffer((4096,), data=T_reshape) +# CHECK-NEXT: T_reshape_2[i] = T.max(T.float32(0.0), T_reshape_1[i]) +# CHECK-NEXT: for i, j in T.grid(64, 64): +# CHECK-NEXT: C_1 = T.Buffer((4096,), data=C.data) +# CHECK-NEXT: C_1[i * 64 + j] = T.float32(0.0) +# CHECK-NEXT: for k in range(64): +# CHECK-NEXT: cse_var_2: T.int32 = i * 64 +# CHECK-NEXT: cse_var_1: T.int32 = cse_var_2 + j +# CHECK-NEXT: T_reshape_2 = T.Buffer((4096,), data=T_reshape) +# CHECK-NEXT: _1_1 = T.Buffer((4096,), data=_1.data) +# CHECK-NEXT: C_1[cse_var_1] = C_1[cse_var_1] + T_reshape_2[cse_var_2 + k] * _1_1[k * 64 + j] +# CHECK-NEXT: INPS = list(obj.values())[:-1] +# CHECK-NEXT: O = obj['C'] +# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") +# CHECK-NEXT: I_R1 = sch.cache_read(INPS[1], "global", [O_W0]) +# CHECK-NEXT: I_F0 = O_W0.op.input_tensors[0] +# CHECK-NEXT: i, j, = O.op.axis +# CHECK-NEXT: k, = O.op.reduce_axis +# CHECK-NEXT: i, i_ = sch[O].split(i, factor=8) +# CHECK-NEXT: j, j_ = sch[O].split(j, factor=32) +# CHECK-NEXT: sch[O].reorder(i, j, i_, j_) +# CHECK-NEXT: j = sch[O].fuse(i, j) +# CHECK-NEXT: sch[O].parallel(j) +# CHECK-NEXT: sch[O_W0].compute_at(sch[O], j) +# CHECK-NEXT: i, j, = O_W0.op.axis +# CHECK-NEXT: k, = O_W0.op.reduce_axis +# CHECK-NEXT: i1 = i +# CHECK-NEXT: j1 = j +# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=16) +# CHECK-NEXT: i1, i2 = sch[O_W0].split(i1, factor=4) +# CHECK-NEXT: j1, j2 = sch[O_W0].split(j1, factor=16) +# CHECK-NEXT: sch[O_W0].reorder(k, i1, j1, k1, i2, j2) +# CHECK-NEXT: sch[I_R1].compute_at(sch[O_W0], k) +# CHECK-NEXT: sch[I_R1].storage_align(I_R1.op.axis[-2], factor=1024, offset=16) +# CHECK-NEXT: sch[I_F0].compute_at(sch[O_W0], k) +# CHECK-NEXT: sch[O_W0].unroll(i2) +# CHECK-NEXT: sch[O_W0].vectorize(j2) +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: def main(_0: T.Buffer((64, 64), "float32"), _1: T.Buffer((64, 64), "float32"), C: T.Buffer((64, 64), "float32")): +# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) +# CHECK-NEXT: T_reshape = T.allocate([4096], "float32", "global") +# CHECK-NEXT: T_reshape_1 = T.Buffer((4096,), data=T_reshape) +# CHECK-NEXT: for ax0 in range(4096): +# CHECK-NEXT: _0_1 = T.Buffer((4096,), data=_0.data) +# CHECK-NEXT: T_reshape_1[ax0] = _0_1[ax0] +# CHECK-NEXT: T_reshape_2 = T.Buffer((4096,), data=T_reshape) +# CHECK-NEXT: for i in range(4096): +# CHECK-NEXT: T_reshape_2[i] = T.max(T.float32(0.0), T_reshape_1[i]) +# CHECK-NEXT: for i_outer_j_outer_fused in T.parallel(16): +# CHECK-NEXT: C_global = T.allocate([256], "float32", "global") +# CHECK-NEXT: T_reshape_3 = T.allocate([128], "float32", "global") +# CHECK-NEXT: _1_global = T.allocate([16640], "float32", "global") +# CHECK-NEXT: C_global_1 = T.Buffer((256,), data=C_global) +# CHECK-NEXT: for i_c_outer_init, j_c_outer_init in T.grid(2, 2): +# CHECK-NEXT: cse_var_1: T.int32 = i_c_outer_init * 128 + j_c_outer_init * 16 +# CHECK-NEXT: C_global_1[cse_var_1:cse_var_1 + 16] = T.Broadcast(T.float32(0.0), 16) +# CHECK-NEXT: C_global_1[cse_var_1 + 32:cse_var_1 + 32 + 16] = T.Broadcast(T.float32(0.0), 16) +# CHECK-NEXT: C_global_1[cse_var_1 + 64:cse_var_1 + 64 + 16] = T.Broadcast(T.float32(0.0), 16) +# CHECK-NEXT: C_global_1[cse_var_1 + 96:cse_var_1 + 96 + 16] = T.Broadcast(T.float32(0.0), 16) +# CHECK-NEXT: for k_outer in range(4): +# CHECK-NEXT: T_reshape_4 = T.Buffer((128,), data=T_reshape_3) +# CHECK-NEXT: for ax0, ax1 in T.grid(8, 16): +# CHECK-NEXT: T_reshape_4[ax0 * 16 + ax1] = T_reshape_2[i_outer_j_outer_fused // 2 * 512 + ax0 * 64 + k_outer * 16 + ax1] +# CHECK-NEXT: _1_global_1 = T.Buffer((16640,), data=_1_global) +# CHECK-NEXT: for ax0, ax1 in T.grid(16, 32): +# CHECK-NEXT: _1_1 = T.Buffer((4096,), data=_1.data) +# CHECK-NEXT: _1_global_1[ax0 * 1040 + ax1] = _1_1[k_outer * 1024 + ax0 * 64 + i_outer_j_outer_fused % 2 * 32 + ax1] +# CHECK-NEXT: for i_c_outer, j_c_outer, k_inner in T.grid(2, 2, 16): +# CHECK-NEXT: cse_var_8: T.int32 = j_c_outer * 16 +# CHECK-NEXT: cse_var_7: T.int32 = i_c_outer * 64 + k_inner +# CHECK-NEXT: cse_var_6: T.int32 = k_inner * 1040 + cse_var_8 +# CHECK-NEXT: cse_var_5: T.int32 = i_c_outer * 128 + cse_var_8 +# CHECK-NEXT: cse_var_4: T.int32 = cse_var_5 + 96 +# CHECK-NEXT: cse_var_3: T.int32 = cse_var_5 + 64 +# CHECK-NEXT: cse_var_2: T.int32 = cse_var_5 + 32 +# CHECK-NEXT: C_global_1[cse_var_5:cse_var_5 + 16] = C_global_1[cse_var_5:cse_var_5 + 16] + T.Broadcast(T_reshape_4[cse_var_7], 16) * _1_global_1[cse_var_6:cse_var_6 + 16] +# CHECK-NEXT: C_global_1[cse_var_2:cse_var_2 + 16] = C_global_1[cse_var_2:cse_var_2 + 16] + T.Broadcast(T_reshape_4[cse_var_7 + 16], 16) * _1_global_1[cse_var_6:cse_var_6 + 16] +# CHECK-NEXT: C_global_1[cse_var_3:cse_var_3 + 16] = C_global_1[cse_var_3:cse_var_3 + 16] + T.Broadcast(T_reshape_4[cse_var_7 + 32], 16) * _1_global_1[cse_var_6:cse_var_6 + 16] +# CHECK-NEXT: C_global_1[cse_var_4:cse_var_4 + 16] = C_global_1[cse_var_4:cse_var_4 + 16] + T.Broadcast(T_reshape_4[cse_var_7 + 48], 16) * _1_global_1[cse_var_6:cse_var_6 + 16] +# CHECK-NEXT: for i_inner, j_inner in T.grid(8, 32): +# CHECK-NEXT: C_1 = T.Buffer((4096,), data=C.data) +# CHECK-NEXT: C_1[i_outer_j_outer_fused // 2 * 512 + i_inner * 64 + i_outer_j_outer_fused % 2 * 32 + j_inner] = C_global_1[i_inner * 32 + j_inner] +# CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/test_relu_matmul_fused_tvm_tir.py b/tests/filecheck/backends/test_relu_matmul_fused_tvm_tir.py new file mode 100644 index 000000000..5afeea78c --- /dev/null +++ b/tests/filecheck/backends/test_relu_matmul_fused_tvm_tir.py @@ -0,0 +1,174 @@ +# RUN: python %s 2>&1 | filecheck %s +# REQUIRES: module_tvm + +import xtc.graphs.xtc.op as O +from xtc.backends.tvm import Backend + +I, J, K, dtype = 64, 64, 64, "float32" +a = O.tensor((I, K), dtype, name="A") +b = O.tensor((K, J), dtype, name="B") + +with O.graph(name="matmul") as gb: + p = O.relu(a, name="relu") + O.matmul(p, b, name="C") + +graph = gb.graph +print(graph) + +impl = Backend(graph, tir_schedule=True) + +sch = impl.get_scheduler() +sch.tile("i", {"i1": 8, "i2": 4}) +sch.tile("j", {"j1": 32, "j2": 16}) +sch.tile("k", {"k1": 16}) +sch.interchange(["i", "j", "k", "i1", "j1", "k1", "i2", "j2"]) +sch.buffer_at("j") +sch.pack_at("k", 1, pad=True) +sch.fuse_producer_at("k", 0) +sch.vectorize(["j2"]) +sch.unroll({"i2": 4}) +sch.parallelize(["i", "j"]) +sched = sch.schedule() + +comp = impl.get_compiler( + shared_lib=True, + dump_file="relu_matmul_fused_tvm_tir", + print_source_ir=True, + print_transformed_ir=True, +) +module = comp.compile(sched) +executor = module.get_executor(validate=True) +res = executor.execute() +print(f"CODE: {res}") + + +# CHECK: graph: +# CHECK-NEXT: name: matmul +# CHECK-NEXT: inputs: +# CHECK-NEXT: - %0 : 64x64xfloat32 +# CHECK-NEXT: - %1 : 64x64xfloat32 +# CHECK-NEXT: outputs: +# CHECK-NEXT: - %3 : 64x64xfloat32 +# CHECK-NEXT: nodes: +# CHECK-NEXT: - %2: relu(%0) {name = 'relu'} : [64x64xfloat32] -> [64x64xfloat32] +# CHECK-NEXT: - %3: matmul(%2, %1) {name = 'C'} : [64x64xfloat32, 64x64xfloat32] -> [64x64xfloat32] +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: def matmul(_0: T.Buffer((64, 64), "float32"), _1: T.Buffer((64, 64), "float32"), C: T.Buffer((64, 64), "float32")): +# CHECK-NEXT: T.func_attr({"tir.noalias": T.bool(True)}) +# CHECK-NEXT: # with T.block("root"): +# CHECK-NEXT: T_reshape = T.alloc_buffer((4096,)) +# CHECK-NEXT: relu = T.alloc_buffer((4096,)) +# CHECK-NEXT: T_reshape_1 = T.alloc_buffer((64, 64)) +# CHECK-NEXT: for ax0 in range(4096): +# CHECK-NEXT: with T.block("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(4096, ax0) +# CHECK-NEXT: T.reads(_0[v_ax0 % 4096 // 64, v_ax0 % 64]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0]) +# CHECK-NEXT: T_reshape[v_ax0] = _0[v_ax0 % 4096 // 64, v_ax0 % 64] +# CHECK-NEXT: for i in range(4096): +# CHECK-NEXT: with T.block("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(4096, i) +# CHECK-NEXT: T.reads(T_reshape[v_i]) +# CHECK-NEXT: T.writes(relu[v_i]) +# CHECK-NEXT: relu[v_i] = T.max(T.float32(0.0), T_reshape[v_i]) +# CHECK-NEXT: for ax0, ax1 in T.grid(64, 64): +# CHECK-NEXT: with T.block("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) +# CHECK-NEXT: T.reads(relu[(v_ax0 * 64 + v_ax1) % 4096]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0, v_ax1]) +# CHECK-NEXT: T_reshape_1[v_ax0, v_ax1] = relu[(v_ax0 * 64 + v_ax1) % 4096] +# CHECK-NEXT: for i, j, k in T.grid(64, 64, 64): +# CHECK-NEXT: with T.block("C"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(T_reshape_1[v_i, v_k], _1[v_k, v_j]) +# CHECK-NEXT: T.writes(C[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: C[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: C[v_i, v_j] = C[v_i, v_j] + T_reshape_1[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: O = sch.get_block("C") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: I_R1 = sch.cache_read(O, 1, "global") +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: I_F0 = sch.get_producers(O)[0] +# CHECK-NEXT: i, i1, i2, = sch.split(i, factors=[None, 2, 4]) +# CHECK-NEXT: j, j1, j2, = sch.split(j, factors=[None, 2, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 16]) +# CHECK-NEXT: sch.reorder(i, j, k, i1, j1, k1, i2, j2) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j) +# CHECK-NEXT: sch.compute_at(I_R1, k) +# CHECK-NEXT: sch.storage_align(I_R1, 0, axis=-2, factor=1024, offset=16) +# CHECK-NEXT: sch.compute_at(I_F0, k) +# CHECK-NEXT: sch.unroll(i2) +# CHECK-NEXT: sch.vectorize(j2) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: def matmul(_0: T.Buffer((64, 64), "float32"), _1: T.Buffer((64, 64), "float32"), C: T.Buffer((64, 64), "float32")): +# CHECK-NEXT: T.func_attr({"tir.noalias": T.bool(True)}) +# CHECK-NEXT: # with T.block("root"): +# CHECK-NEXT: T_reshape = T.alloc_buffer((4096,)) +# CHECK-NEXT: relu = T.alloc_buffer((4096,)) +# CHECK-NEXT: T_reshape_1 = T.alloc_buffer((64, 64)) +# CHECK-NEXT: _1_global = T.alloc_buffer((64, 64)) +# CHECK-NEXT: C_global = T.alloc_buffer((64, 64)) +# CHECK-NEXT: for ax0 in range(4096): +# CHECK-NEXT: with T.block("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(4096, ax0) +# CHECK-NEXT: T.reads(_0[v_ax0 % 4096 // 64, v_ax0 % 64]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0]) +# CHECK-NEXT: T_reshape[v_ax0] = _0[v_ax0 % 4096 // 64, v_ax0 % 64] +# CHECK-NEXT: for i in range(4096): +# CHECK-NEXT: with T.block("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(4096, i) +# CHECK-NEXT: T.reads(T_reshape[v_i]) +# CHECK-NEXT: T.writes(relu[v_i]) +# CHECK-NEXT: relu[v_i] = T.max(T.float32(0.0), T_reshape[v_i]) +# CHECK-NEXT: for i_0_j_0_fused in T.parallel(16): +# CHECK-NEXT: for k_0 in range(4): +# CHECK-NEXT: for ax0, ax1 in T.grid(16, 32): +# CHECK-NEXT: with T.block("_1_global"): +# CHECK-NEXT: v0 = T.axis.spatial(64, k_0 * 16 + ax0) +# CHECK-NEXT: v1 = T.axis.spatial(64, i_0_j_0_fused % 2 * 32 + ax1) +# CHECK-NEXT: T.reads(_1[v0, v1]) +# CHECK-NEXT: T.writes(_1_global[v0, v1]) +# CHECK-NEXT: T.block_attr({"buffer_dim_align": [[0, 0, 1024, 16]]}) +# CHECK-NEXT: _1_global[v0, v1] = _1[v0, v1] +# CHECK-NEXT: for ax0, ax1 in T.grid(8, 16): +# CHECK-NEXT: with T.block("T_reshape_1"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(64, i_0_j_0_fused // 2 * 8 + ax0) +# CHECK-NEXT: v_ax1 = T.axis.spatial(64, k_0 * 16 + ax1) +# CHECK-NEXT: T.reads(relu[(v_ax0 * 64 + v_ax1) % 4096]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0, v_ax1]) +# CHECK-NEXT: T_reshape_1[v_ax0, v_ax1] = relu[(v_ax0 * 64 + v_ax1) % 4096] +# CHECK-NEXT: for i_1, j_1, k_1 in T.grid(2, 2, 16): +# CHECK-NEXT: for i_2 in T.unroll(4): +# CHECK-NEXT: for j_2 in T.vectorized(16): +# CHECK-NEXT: with T.block("C"): +# CHECK-NEXT: v_i = T.axis.spatial(64, i_0_j_0_fused // 2 * 8 + i_1 * 4 + i_2) +# CHECK-NEXT: v_j = T.axis.spatial(64, i_0_j_0_fused % 2 * 32 + j_1 * 16 + j_2) +# CHECK-NEXT: v_k = T.axis.reduce(64, k_0 * 16 + k_1) +# CHECK-NEXT: T.reads(T_reshape_1[v_i, v_k], _1_global[v_k, v_j]) +# CHECK-NEXT: T.writes(C_global[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: C_global[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: C_global[v_i, v_j] = C_global[v_i, v_j] + T_reshape_1[v_i, v_k] * _1_global[v_k, v_j] +# CHECK-NEXT: for ax0, ax1 in T.grid(8, 32): +# CHECK-NEXT: with T.block("C_global"): +# CHECK-NEXT: v0 = T.axis.spatial(64, i_0_j_0_fused // 2 * 8 + ax0) +# CHECK-NEXT: v1 = T.axis.spatial(64, i_0_j_0_fused % 2 * 32 + ax1) +# CHECK-NEXT: T.reads(C_global[v0, v1]) +# CHECK-NEXT: T.writes(C[v0, v1]) +# CHECK-NEXT: C[v0, v1] = C_global[v0, v1] +# CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/schedules/test_get_descript.py b/tests/filecheck/schedules/test_get_descript.py index 47888bb41..67385ad57 100644 --- a/tests/filecheck/schedules/test_get_descript.py +++ b/tests/filecheck/schedules/test_get_descript.py @@ -5,12 +5,14 @@ import sys import xtc.graphs.xtc.op as O -I, J, K, dtype = 4, 32, 512, "float32" +I, J, K, dtype = 128, 256, 512, "float32" a = O.tensor((I, K), dtype, name="A") b = O.tensor((K, J), dtype, name="B") with O.graph(name="matmul") as gb: - O.matmul(a, b, name="C") + p = O.relu(a, name="relu") + m = O.matmul(p, b, name="matmul") + O.relu(m, name="relu") graph = gb.graph @@ -22,42 +24,52 @@ else: assert False - + impl = Backend(graph) -sch = impl.get_scheduler() +sch = impl.get_scheduler(default_node="matmul") sch.set_dims(["I", "J", "K"]) -sch.tile("I", {"I0": 2}) -sch.tile("J", {"J0": 16}) -sch.interchange(["K", "I", "J", "I0", "J0"]) -sch.unroll({"I0": 2}) +sch.tile("I", {"I1": 64, "I0": 2}) +sch.tile("J", {"J1": 128, "J0": 16}) +sch.tile("K", {"K0": 32}) +sch.interchange(["J", "K", "I", "J1", "I1", "K0", "I0", "J0"]) +sch.unroll({"K0": 8, "I0": 2}) sch.vectorize(["J0"]) +sch.parallelize(["J"]) if "--tvm" in sys.argv: sch.buffer_at("J") - sch.pack_at("I", 0, pad=True) + sch.pack_at("K", 1, pad=True) + sch.fuse_producer_at("I", 0) + sch.fuse_consumer_at("J") loop_nest = sch.get_loop_nest() print(loop_nest.root_node.pretty_print()) -# CHECK-MLIR: loop K -# CHECK-MLIR-NEXT: loop I -# CHECK-MLIR-NEXT: loop J -# CHECK-MLIR-NEXT: tile(I, 2) // unroll(2) -# CHECK-MLIR-NEXT: tile(J, 16) // vectorized -# CHECK-MLIR-NEXT: ... +# CHECK-MLIR: loop J // parallelized +# CHECK-MLIR-NEXT: loop K +# CHECK-MLIR-NEXT: loop I +# CHECK-MLIR-NEXT: tile(J, 128) +# CHECK-MLIR-NEXT: tile(I, 64) +# CHECK-MLIR-NEXT: tile(K, 32) // unroll(8) +# CHECK-MLIR-NEXT: tile(I, 2) // unroll(2) +# CHECK-MLIR-NEXT: tile(J, 16) // vectorized +# CHECK-MLIR-NEXT: ... -# CHECK-TVM: loop K -# CHECK-TVM-NEXT: loop I // pack(0, pad) -# CHECK-TVM-NEXT: loop J // buffer -# CHECK-TVM-NEXT: tile(I, 2) // unroll(2) -# CHECK-TVM-NEXT: tile(J, 16) // vectorized -# CHECK-TVM-NEXT: ... +# CHECK-TVM: loop J // parallelized, buffer, fuse_consumer +# CHECK-TVM-NEXT: loop K // pack(1, pad) +# CHECK-TVM-NEXT: loop I // fuse_producer(0) +# CHECK-TVM-NEXT: tile(J, 128) +# CHECK-TVM-NEXT: tile(I, 64) +# CHECK-TVM-NEXT: tile(K, 32) // unroll(8) +# CHECK-TVM-NEXT: tile(I, 2) // unroll(2) +# CHECK-TVM-NEXT: tile(J, 16) // vectorized +# CHECK-TVM-NEXT: ... # Test with split (MLIR only - TVM does not support split) if "--mlir" in sys.argv: print("---") impl2 = Backend(graph) - sch2 = impl2.get_scheduler() + sch2 = impl2.get_scheduler(default_node="matmul") sch2.set_dims(["I", "J", "K"]) sch2.split("I", {"I_lo": 0, "I_hi": 2}) sch2.tile("J", {"J0": 16}, root="./I_lo")