diff --git a/.gitignore b/.gitignore index f1d2c33d2..d275360f3 100644 --- a/.gitignore +++ b/.gitignore @@ -3,6 +3,7 @@ __marimo__ *.so *.csv *.dylib +*.a venv .venv build diff --git a/dependencies.toml b/dependencies.toml index 54939d2e0..7db05a90c 100644 --- a/dependencies.toml +++ b/dependencies.toml @@ -51,7 +51,7 @@ dependencies = [ ] [groups.tvm] -include = ["llvm"] dependencies = [ - 'xtc-tvm-python-bindings==0.19.0.12', + 'apache-tvm==0.26.0', + 'apache-tvm-ffi==0.1.13.post3', ] diff --git a/docs/tutorials/xtc_101.py b/docs/tutorials/xtc_101.py index 170681598..4cdb6901f 100644 --- a/docs/tutorials/xtc_101.py +++ b/docs/tutorials/xtc_101.py @@ -740,7 +740,6 @@ def _(mo, run_exploration): value= '''import xtc.graphs.xtc.op as O from xtc.graphs.xtc.graph import XTCGraph -from xtc.backends.tvm import Backend as TVM_Backend from xtc.backends.mlir import Backend as MLIR_Backend from xtc.schedules.descript import descript_scheduler from xtc.runtimes.host import HostRuntime diff --git a/pyproject.toml b/pyproject.toml index d377562cd..0aa14da1d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,11 +21,11 @@ dependencies = [ ] [project.optional-dependencies] -default = ['xtc-llvm-tools==22.1.8.2', 'xtc-mlir-tools==22.1.8.3', 'xtc-mlir-python-bindings==22.1.8.3', 'xtc-mlir-extra-tools==22.1.8.7', 'xtc-tvm-python-bindings==0.19.0.12'] -dev = ['xtc-llvm-tools==22.1.8.2', 'xtc-mlir-tools==22.1.8.3', 'xtc-mlir-python-bindings==22.1.8.3', 'xtc-mlir-extra-tools==22.1.8.7', 'xtc-tvm-python-bindings==0.19.0.12', "coverage>=7.8.0", "filecheck==1.0.3", "lit", "mkdocs", "mkdocs-material", "mkdocstrings-python", "mypy==1.15.0", "pytest", "pytest-xdist", "pyright==1.1.407", "types-PyYAML", "marimo==0.19.6", "ruff==0.14.10"] +default = ['xtc-llvm-tools==22.1.8.2', 'xtc-mlir-tools==22.1.8.3', 'xtc-mlir-python-bindings==22.1.8.3', 'xtc-mlir-extra-tools==22.1.8.7', 'apache-tvm==0.26.0', 'apache-tvm-ffi==0.1.13.post3'] +dev = ['xtc-llvm-tools==22.1.8.2', 'xtc-mlir-tools==22.1.8.3', 'xtc-mlir-python-bindings==22.1.8.3', 'xtc-mlir-extra-tools==22.1.8.7', 'apache-tvm==0.26.0', 'apache-tvm-ffi==0.1.13.post3', "coverage>=7.8.0", "filecheck==1.0.3", "lit", "mkdocs", "mkdocs-material", "mkdocstrings-python", "mypy==1.15.0", "pytest", "pytest-xdist", "pyright==1.1.407", "types-PyYAML", "marimo==0.19.6", "ruff==0.14.10"] test = ["coverage>=7.8.0", "filecheck==1.0.3", "lit", "mkdocs", "mkdocs-material", "mkdocstrings-python", "mypy==1.15.0", "pytest", "pytest-xdist", "pyright==1.1.407", "types-PyYAML", "marimo==0.19.6", "ruff==0.14.10"] mlir = ['xtc-llvm-tools==22.1.8.2', 'xtc-mlir-tools==22.1.8.3', 'xtc-mlir-python-bindings==22.1.8.3', 'xtc-mlir-extra-tools==22.1.8.7'] -tvm = ['xtc-llvm-tools==22.1.8.2', 'xtc-tvm-python-bindings==0.19.0.12'] +tvm = ['apache-tvm==0.26.0', 'apache-tvm-ffi==0.1.13.post3'] [tool.setuptools] platforms = ["Linux"] diff --git a/src/xtc/backends/tvm/TVMBackend.py b/src/xtc/backends/tvm/TVMBackend.py index 9c35062ba..6995c3ff7 100644 --- a/src/xtc/backends/tvm/TVMBackend.py +++ b/src/xtc/backends/tvm/TVMBackend.py @@ -27,7 +27,6 @@ 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 04a1661d0..59b2c40a5 100644 --- a/src/xtc/backends/tvm/TVMCompiler.py +++ b/src/xtc/backends/tvm/TVMCompiler.py @@ -19,13 +19,18 @@ import xtc.itf as itf from xtc.utils.text import jinja_generate_file from xtc.utils.tarfile import TarFile +from xtc.utils.files import relative_to -from xtc.utils.host_tools import disassemble, target_triple +from xtc.utils.host_tools import ( + disassemble, + target_triple, + cc_command, + binutils_command, +) from .TVMOpsCompiler import ( TVMExprCompiler, TVMScheduledExpr, - TVMScheduledExprTE, TVMScheduledExprTIR, ) from .TVMOps import ( @@ -33,13 +38,16 @@ ) import tvm +import tvm_ffi __all__ = [ "TVMCompiler", ] + TVM_VERSION = Version(tvm.__version__.split("+", 1)[0]) +assert TVM_VERSION >= Version("0.26") class TVMCompiler(itf.comp.Compiler): @@ -66,9 +74,9 @@ def __init__( self.emit_c = kwargs.get("emit_c", False) self.target = kwargs.get("target", "native") self.arch = kwargs.get("arch", "native") - self.tvm_target_options = self._get_tvm_target_options(self.target, self.arch) self.tvm_target = "llvm" - self.tvm_tgt = f"{self.tvm_target} {self.tvm_target_options}" + self.tvm_target_options = self._get_tvm_target_options(self.target, self.arch) + self.tvm_tgt = self._get_tvm_target(self.tvm_target, self.tvm_target_options) assert not self.executable, f"executable generation not supported yet for TVM" assert self.shared_lib or self.emit_c or self.ar_lib, ( f"shared_lib/ar_lib or C generation is mandatory for TVM" @@ -86,7 +94,7 @@ def backend(self) -> itf.back.Backend: def get_source_ir(self, schedule: itf.schd.Schedule) -> str: # The initial lowered Tensor IR, before the schedule is applied. op = self._backend._tvm_base - expr_compiler = TVMExprCompiler(op, tir_schedule=self._backend._tir_schedule) + expr_compiler = TVMExprCompiler(op) return expr_compiler.generate().schedule().dumps() def _save_temp(self, fname: str, content: str) -> None: @@ -103,7 +111,8 @@ def compile(self, schedule: itf.schd.Schedule) -> itf.comp.Module: save_temp = self._save_temp op = self._backend._tvm_base func_name = self.payload_name - packed_func_name = f"packed_{func_name}" if self.bare_ptr else func_name + tvm_ffi_func_name = f"__tvm_ffi_{func_name}" + compute_func_name = f"{func_name}_compute_" if self.shared_lib: type = "shlib" @@ -118,16 +127,12 @@ def compile(self, schedule: itf.schd.Schedule) -> itf.comp.Module: dump_base = Path(self.dump_file).stem lib_path = self.dump_file if type in ["arlib", "shlib"]: - emit_c_base = f"{lib_path}.export_c" + emit_c_base = f"{lib_path}_export_c" else: emit_c_base = lib_path - if self.bare_ptr: - packed_lib_path = f"{lib_path}_packed" - emit_c_packed_base = f"{emit_c_base}_packed" - else: - packed_lib_path = lib_path - emit_c_packed_base = emit_c_base - expr_compiler = TVMExprCompiler(op, tir_schedule=self._backend._tir_schedule) + packed_lib_path = f"{lib_path}_tvm_ffi" + emit_c_packed_base = f"{emit_c_base}_tvm_ffi" + expr_compiler = TVMExprCompiler(op) schedulable = expr_compiler.generate() if self.print_source_ir or self.save_temps: lowered = schedulable.schedule().dumps() @@ -147,85 +152,46 @@ def compile(self, schedule: itf.schd.Schedule) -> itf.comp.Module: if self.emit_c: self._build_c( sch, - func_name=packed_func_name, + func_name=func_name, fname=emit_c_packed_base, ) if type in ["shlib", "arlib"]: - built = self._build(sch, func_name=packed_func_name) + built = self._build(sch, func_name=func_name) if self.save_temps: for idx, mod in enumerate(built._collect_dso_modules()): - llvm_ir = str(mod.get_source("ll")) + llvm_ir = str(mod.inspect_source("ll")) save_temp(f"{dump_base}.lib{idx}.ll", llvm_ir) # This will generate a .tar with the .o files # built.export_library(f"{save_temps_dir}/{packed_lib_path}.tar") - if self.print_assembly: - with tempfile.TemporaryDirectory() as tdir: - tmpname = f"{tdir}/built" - fname = f"{packed_func_name}_compute_" - self._export_library(built, tmpname, type="shlib") - ext = ".dylib" if sys.platform == "darwin" else ".so" - disassembly = disassemble( - f"{tmpname}{ext}", - function=fname, - section=".text", - color=self.color, - ) - print(disassembly, flush=True) - self._export_library(built, packed_lib_path, type=type) - - csrcs, shlibs, arlibs, headers = [], [], [], [] - tvm_prefix = Path(tvm.__path__[0]) - if self.bare_ptr: - wrapper = PackedOperatorWrapper( - op, - func_name, - packed_func_name, - cc_prefix=self._cc_prefix(), - ) - if type == "shlib": - ext = ".dylib" if sys.platform == "darwin" else ".so" - wrapper.build(lib_path, packed_lib_path, type=type) - shlibs = [f"{packed_lib_path}{ext}"] - elif type == "arlib": - wrapper.build(lib_path, packed_lib_path, type=type) - arlibs = [f"{packed_lib_path}.a"] - if self.emit_c: - wrapper.build(emit_c_base, emit_c_packed_base, type="csrc") - csrcs = [f"{emit_c_packed_base}.c"] - if Path(f"{lib_path}.h").exists(): - headers = [f"{lib_path}.h"] - if type in ["shlib", "arlib"]: - ext = ".dylib" if sys.platform == "darwin" else ".so" - shlibs += [f"{tvm_prefix}/libtvm_runtime{ext}"] - # As of now shared_lib/ar_lib takes priority over emit_c - return HostModule( - dump_base, - func_name, - f"{lib_path}{ext}" if type == "shlib" else f"{lib_path}.a", - type, - shlibs=shlibs, - arlibs=arlibs, - headers=headers, - bare_ptr=self.bare_ptr, - graph=self._backend._graph, + self._export_archive(built, f"{packed_lib_path}.a") + + wrapper = PackedOperatorWrapper( + op, + func_name, + tvm_ffi_func_name, + arch=self.target, + bare_ptr=self.bare_ptr, + ) + if type in ["shlib", "arlib"] and self.emit_c: + wrapper.build(emit_c_base, emit_c_packed_base, type="csrc") + module_file, module_args = wrapper.build(lib_path, packed_lib_path, type=type) + assert Path(module_file).with_suffix("") == Path(lib_path) + if type == "shlib" and self.print_assembly: + disassembly = disassemble( + module_file, + function=compute_func_name, + section="text", + color=self.color, + arch=self.target, ) - assert type == "csrc" - headers_path = [ - str(path) - for path in [ - tvm_prefix / "include", - tvm_prefix / "3rdparty" / "dlpack" / "include", - ] - ] + print(disassembly, flush=True) return HostModule( dump_base, func_name, - f"{lib_path}.c", - "csrc", + module_file, + type, + **module_args, bare_ptr=self.bare_ptr, - csrcs=csrcs, - headers=headers, - headers_path=headers_path, graph=self._backend._graph, ) @@ -259,7 +225,14 @@ def _build_c( self._tvm_emit_c(sch, self.tvm_tgt, func_name, fname) @classmethod - def _get_tvm_target_options(cls, target: str, arch: str) -> str: + def _get_tvm_target(cls, kind: str, options: dict) -> dict: + return { + **dict(kind=kind), + **options, + } + + @classmethod + def _get_tvm_target_options(cls, target: str, arch: str) -> dict: """ Returm the tvm target options given the target and arch """ @@ -269,7 +242,7 @@ def _get_tvm_target_options(cls, target: str, arch: str) -> str: else: assert arch != "native", f"can't pass native arch for non native target" tvm_cpu = "" - tvm_attrs = "" + tvm_attrs = [] tvm_triple = target_triple(target) if target in ["x86_64"]: if arch == "avx512": @@ -279,18 +252,15 @@ def _get_tvm_target_options(cls, target: str, arch: str) -> str: elif target in ["aarch64"]: if arch == "neon": tvm_cpu = "cortex-a72" - tvm_attrs = "+neon" - target_options = [] - if tvm_triple: - target_options.append(f"-mtriple={tvm_triple}") - if tvm_cpu: - target_options.append(f"-mcpu={tvm_cpu}") - if tvm_attrs: - target_options.append(f"-mattr={tvm_attrs}") - return " ".join(target_options) + tvm_attrs += ["+neon"] + return { + **(dict(mtriple=tvm_triple) if tvm_triple else {}), + **(dict(mcpu=tvm_cpu) if tvm_cpu else {}), + **(dict(mattr=tvm_attrs) if tvm_attrs else {}), + } @classmethod - def _get_tvm_native_target_options(cls) -> str: + def _get_tvm_native_target_options(cls) -> dict: """ Returm the tvm target options to pass to llvm. """ @@ -299,82 +269,64 @@ def _get_tvm_native_target_options(cls) -> str: info = get_cpu_info() arch = info["arch_string_raw"] flags = info.get("flags", []) - triple = target_triple(arch) - cpu, attrs = "", "" + tvm_triple = target_triple(arch) + tvm_cpu, tvm_attrs = "", [] if arch == "x86_64": if "avx512f" in flags: - cpu = "skylake-avx512" + tvm_cpu = "skylake-avx512" elif "avx2" in flags: - cpu = "core-avx2" + tvm_cpu = "core-avx2" elif arch == "aarch64": if "asimd" in flags: - cpu = "cortex-a72" - attrs = "+neon" - target_options = [] - if triple: - target_options.append(f"-mtriple={triple}") - if cpu: - target_options.append(f"-mcpu={cpu}") - if attrs: - target_options.append(f"-mattr={attrs}") - return " ".join(target_options) + tvm_cpu = "cortex-a72" + tvm_attrs += ["+neon"] + return { + **(dict(mtriple=tvm_triple) if tvm_triple else {}), + **(dict(mcpu=tvm_cpu) if tvm_cpu else {}), + **(dict(mattr=tvm_attrs) if tvm_attrs else {}), + } @classmethod - def _tvm_build_crt_args(cls, target: str) -> dict[str, Any]: + def _tvm_build_crt_args(cls, target: dict) -> dict[str, Any]: # We use system-lib with crt runtime such that DSO loading works + # As of TVM >= 0.26 this is not needed anymore # The generated .so can then be used: # - for static compilation as soon as the tvm runtime is provided # - for dynamic loading from python - # Recent version of tvm (i.e. 0.19) have a Runtime object - # Older version (i.e. 0.16) support passing runtime options in target - try: - from tvm.relay.backend import Runtime - - runtime_kwargs = { - "runtime": Runtime("crt", {"system-lib": True}), - } - except: - runtime_kwargs = {} - if TVM_VERSION < Version("0.21"): - target = f"{target} --system-lib --runtime=c" - + runtime_kwargs: dict[str, Any] = {} return { "target": target, **runtime_kwargs, } @classmethod - def _tvm_build_crt(cls, sch: TVMScheduledExpr, target: str, cname: str) -> Any: + def _tvm_build_crt(cls, sch: TVMScheduledExpr, target: dict, cname: str) -> Any: build_kwargs = cls._tvm_build_crt_args(target) config = {} - if target.startswith("c "): + if target["kind"] == "c": config.update( { - "tir.disable_vectorize": True, + "tirx.disable_vectorize": True, } ) with tvm.transform.PassContext(opt_level=3, config=config): - 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) + 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.tirx.build(mod, **build_kwargs) return built @classmethod def _tvm_emit_c( cls, sch: TVMScheduledExpr, - target: str, + target: dict, cname: str, fname: str, ) -> Any: # Ignore initial target as of now and generate target agnostic C - target = "c -keys=arch -march=generic -mcpu=generic" + target = dict(kind="c", keys=["arch"], march="generic", mcpu="generic") built = cls._tvm_build_crt(sch, target, cname) out_dir = Path(fname).parent out_base = Path(fname).stem @@ -391,32 +343,15 @@ def _tvm_emit_c( cfile = tmp_dir_path / "lib0.c" shutil.copy(cfile, out_dir / f"{out_base}.c") - def _cc_prefix(self) -> str: - map = { - "x86_64": "x86_64-linux-gnu-", - "aarch64": "aarch64-linux-gnu-", - "native": "", - } - assert self.target in map, ( - f"unsupported target for cross compilation: {self.target}" - ) - return map[self.target] - - def _export_library(self, mod: Any, basename: str, type: str): - from tvm.contrib import cc + def _export_archive(self, mod: Any, archive_path: str): + from tvm.support import cc - prefix = self._cc_prefix() - if type == "arlib": - fcompile = partial(cc.create_staticlib, ar=f"{prefix}ar") - mod.export_library(f"{basename}.a", fcompile=fcompile) - assert Path(f"{basename}.a").exists() - else: - fcompile = None - if prefix != "": - fcompile = cc.cross_compiler(f"{prefix}g++") - assert type == "shlib" - ext = ".dylib" if sys.platform == "darwin" else ".so" - mod.export_library(f"{basename}{ext}", fcompile=fcompile) + fcompile = partial( + cc.create_staticlib, + ar=binutils_command("ar", self.target), + ) + mod.export_library(archive_path, fcompile=fcompile) + assert Path(archive_path).exists() class PackedOperatorWrapper: @@ -427,68 +362,113 @@ def __init__( operation: TVMBaseExpr, func_name: str, packed_func_name: str, - cc_prefix: str = "", + arch: str = "", + bare_ptr: bool = False, ) -> None: self.operation = operation self.func_name = func_name self.packed_func_name = packed_func_name - self._cc_prefix = cc_prefix + self._arch = arch + self._bare_ptr = bare_ptr - def generate_c(self, output_base: str) -> None: + def generate_c(self, output_base: str, header: bool = False) -> None: config = { "inputs": self.operation.np_inputs_spec(), "outputs": self.operation.np_outputs_spec(), "func_name": self.func_name, - "packed_func_name": self.packed_func_name, + "tvm_ffi_func_name": self.packed_func_name, } - jinja_generate_file( - f"{output_base}.c", - str(self.TEMPLATES_DIR / "packed_op_wrapper.c.jinja"), - **config, - ) - jinja_generate_file( - f"{output_base}.h", - str(self.TEMPLATES_DIR / "unpacked_op.h.jinja"), - **config, - ) + if self._bare_ptr: + jinja_generate_file( + f"{output_base}.c", + str(self.TEMPLATES_DIR / "packed_op_wrapper.c.jinja"), + **config, + ) + if header: + jinja_generate_file( + f"{output_base}.h", + str(self.TEMPLATES_DIR / "unpacked_op.h.jinja"), + **config, + ) + else: + jinja_generate_file( + f"{output_base}.c", + str(self.TEMPLATES_DIR / "tvm_ffi_op_wrapper.c.jinja"), + **config, + ) - def build(self, lib_fname: str, packed_lib_fname: str, type: str) -> None: + def build( + self, lib_fname: str, packed_lib_fname: str, type: str + ) -> tuple[str, dict[str, Any]]: + ext = ".dylib" if sys.platform == "darwin" else ".so" unpacked_lib_dir = Path(lib_fname).parent unpacked_lib_base = Path(lib_fname).stem packed_lib_dir = Path(packed_lib_fname).parent packed_lib_name = Path(packed_lib_fname).stem + packed_ar_name = f"{packed_lib_fname}.a" assert packed_lib_dir == unpacked_lib_dir, ( f"must generate wrapper at the same location as packed lib" ) - with tempfile.TemporaryDirectory() as tdir: + if type in ["shlib", "arlib"]: + assert Path(packed_ar_name).exists() + tdir = tempfile.mkdtemp(dir=unpacked_lib_dir) + try: + tvm_ffi_prefix = Path(tvm_ffi.__path__[0]) + tvm_ffi_libdir = tvm_ffi_prefix / "lib" + tvm_prefix = Path(tvm.__path__[0]) + tvm_libdir = tvm_prefix / "lib" output_base = str(Path(tdir) / Path(lib_fname).stem) - self.generate_c(output_base) - shutil.copy(f"{output_base}.h", unpacked_lib_dir / f"{unpacked_lib_base}.h") + host_runtime_dir = Path(__file__).parents[2] / "csrcs" / "runtimes" / "host" + tvm_runtime_init_c = str(host_runtime_dir / "tvm_runtime_init.c") + self.generate_c(output_base, header=(type == "csrc")) + headers, headers_path, csrcs, shlibs, arlibs = [], [], [], [], [] + if self._bare_ptr and type == "csrc": + shutil.copy( + f"{output_base}.h", unpacked_lib_dir / f"{unpacked_lib_base}.h" + ) + headers += [f"{lib_fname}.h"] if type == "csrc": shutil.copy( f"{output_base}.c", unpacked_lib_dir / f"{unpacked_lib_base}.c" ) + module_file = f"{lib_fname}.c" + csrcs += [ + tvm_runtime_init_c, + f"{packed_lib_fname}.c", + ] + headers_path += [ + str(path) + for path in [ + tvm_prefix / "include", + tvm_ffi_prefix / "include", + ] + ] + shlibs += [ + f"{tvm_libdir}/libtvm_runtime{ext}", + f"{tvm_ffi_libdir}/libtvm_ffi{ext}", + ] elif type == "shlib": - ext = ".dylib" if sys.platform == "darwin" else ".so" - cmd = ( - f"{self._cc_prefix}gcc --shared -fPIC -O2 {output_base}.c " - f"-o {unpacked_lib_base}{ext} " - f"{packed_lib_name}{ext} -Wl,--rpath,$ORIGIN" - ) - p = subprocess.run( - shlex.split(cmd), - text=True, - capture_output=True, - cwd=unpacked_lib_dir, - ) - if p.returncode != 0: - raise RuntimeError( - f"Failed command {cmd}:\n{p.stdout}\n{p.stderr}\n" + output_dir = unpacked_lib_dir + object_fnames = [ + str(relative_to(fname, output_dir)) + for fname in self._build_objects( + [f"{output_base}.c"] + self._runtime_sources(), + tdir, ) - elif type == "arlib": + ] + opts = "-O2" + sh_opts = "--shared -fPIC" + ext = ".so" + if sys.platform == "darwin": + sh_opts += " -undefined dynamic_lookup" + ext = ".dylib" + shlib_fname = f"{unpacked_lib_base}{ext}" + shlib_dest = str(relative_to(shlib_fname, output_dir)) cmd = ( - f"{self._cc_prefix}gcc -c -O2 {output_base}.c " - f"-o {unpacked_lib_base}.o" + f"{cc_command(self._arch)} {sh_opts} {opts} " + f"{' '.join(object_fnames)} " + f"{packed_lib_fname}.a " + f"-o {unpacked_lib_base}{ext}" ) p = subprocess.run( shlex.split(cmd), @@ -498,19 +478,103 @@ def build(self, lib_fname: str, packed_lib_fname: str, type: str) -> None: ) if p.returncode != 0: raise RuntimeError( - f"Failed command {cmd}:\n{p.stdout}\n{p.stderr}\n" + f"Failed command {cmd} (cwd: {unpacked_lib_dir}:\n" + f"{p.stdout}\n" + f"{p.stderr}\n" ) - cmd = ( - f"{self._cc_prefix}ar -crs {unpacked_lib_base}.a " - f"{unpacked_lib_base}.o" + module_file = f"{lib_fname}{ext}" + shlibs += [ + f"{tvm_libdir}/libtvm_runtime{ext}", + f"{tvm_ffi_libdir}/libtvm_ffi{ext}", + ] + else: + assert type == "arlib" + archive_fname = self._build_archive( + [f"{output_base}.c"] + self._runtime_sources(), + f"{lib_fname}.a", ) - p = subprocess.run( - shlex.split(cmd), - text=True, - capture_output=True, - cwd=unpacked_lib_dir, + module_file = archive_fname + arlibs += [f"{packed_lib_fname}.a"] + shlibs += [ + f"{tvm_libdir}/libtvm_runtime{ext}", + f"{tvm_ffi_libdir}/libtvm_ffi{ext}", + ] + except Exception: + raise + else: + shutil.rmtree(tdir) + module_args = { + **(dict(headers=headers) if headers else {}), + **(dict(headers_path=headers_path) if headers_path else {}), + **(dict(csrcs=csrcs) if csrcs else {}), + **(dict(shlibs=shlibs) if shlibs else {}), + **(dict(arlibs=arlibs) if arlibs else {}), + } + return module_file, module_args + + def _runtime_sources(self) -> list[str]: + host_runtime_dir = Path(__file__).parents[2] / "csrcs" / "runtimes" / "host" + tvm_runtime_init_c = str(host_runtime_dir / "tvm_runtime_init.c") + return [tvm_runtime_init_c] + + def _build_object(self, source_fname: str, object_fname: str) -> str: + assert object_fname.endswith(".o") + opts = "-O2" + pic_opts = "-fPIC" + output_dir = Path(object_fname).parent + object_dest = str(relative_to(object_fname, output_dir)) + source_inp = str(relative_to(source_fname, output_dir)) + cmd = ( + f"{cc_command(self._arch)} -c {opts} {pic_opts} " + f"{source_inp} " + f"-o {object_dest}" + ) + p = subprocess.run( + shlex.split(cmd), + text=True, + capture_output=True, + cwd=output_dir, + ) + if p.returncode != 0: + raise RuntimeError( + f"Failed command {cmd} (cwd: {output_dir} :\n{p.stdout}\n{p.stderr}\n" + ) + return object_fname + + def _build_objects(self, source_fnames: list[str], output_dir: str) -> list[str]: + return [ + self._build_object(fname, str(Path(output_dir) / f"{Path(fname).stem}.o")) + for fname in source_fnames + ] + + def _build_archive(self, source_fnames: list[str], archive_fname: str) -> str: + assert archive_fname.endswith(".a") + output_dir = Path(archive_fname).parent + archive_dest = str(relative_to(archive_fname, output_dir)) + tdir = tempfile.mkdtemp(dir=output_dir) + try: + object_fnames = [ + str(relative_to(fname, output_dir)) + for fname in self._build_objects(source_fnames, tdir) + ] + cmd = ( + f"{binutils_command('ar', self._arch)} -crs {archive_dest} " + f"{' '.join(object_fnames)} " + ) + p = subprocess.run( + shlex.split(cmd), + text=True, + capture_output=True, + cwd=output_dir, + ) + if p.returncode != 0: + raise RuntimeError( + f"Failed command {cmd} (cwd: {output_dir}:\n" + f"{p.stdout}\n" + f"{p.stderr}\n" ) - if p.returncode != 0: - raise RuntimeError( - f"Failed command {cmd}:\n{p.stdout}\n{p.stderr}\n" - ) + except Exception: + raise + else: + shutil.rmtree(tdir) + return archive_fname diff --git a/src/xtc/backends/tvm/TVMOps.py b/src/xtc/backends/tvm/TVMOps.py index 43b0d7a15..242684859 100644 --- a/src/xtc/backends/tvm/TVMOps.py +++ b/src/xtc/backends/tvm/TVMOps.py @@ -337,7 +337,7 @@ def generate_op( O = topi.reshape(A, newshape=(size,)) O = te.compute( (Ki,), - lambda i,: tvm.tir.max(self.attrs["threshold"], O[i]), + lambda i,: tvm.tirx.max(self.attrs["threshold"], O[i]), name=self.name, ) if shape != newshape: @@ -486,11 +486,11 @@ def get_indexes(*args: int) -> tuple[int, ...]: indexes = [i - padding[0] for i in indexes] return tuple(indexes) - def get_args_bounds(*args: int) -> list[tvm.tir.PrimExpr]: + def get_args_bounds(*args: int) -> list[Any]: indexes = list(args) if isinstance(padding, dict): return [ - tvm.tir.all( + tvm.tirx.all( indexes[i] - pad_b >= 0, indexes[i] - pad_b < dims_values_all[i], ) @@ -498,7 +498,7 @@ def get_args_bounds(*args: int) -> list[tvm.tir.PrimExpr]: ] else: return [ - tvm.tir.all( + tvm.tirx.all( index - padding[0] >= 0, index - padding[0] < dims_values_all[i], ) @@ -508,8 +508,8 @@ def get_args_bounds(*args: int) -> list[tvm.tir.PrimExpr]: O = te.compute( tuple(dims_values), lambda *args: ( - tvm.tir.if_then_else( - tvm.tir.all(*get_args_bounds(*args)), + tvm.tirx.if_then_else( + tvm.tirx.all(*get_args_bounds(*args)), A[get_indexes(*args)], constant_value, ) diff --git a/src/xtc/backends/tvm/TVMOpsCompiler.py b/src/xtc/backends/tvm/TVMOpsCompiler.py index b597990c3..e95ded651 100644 --- a/src/xtc/backends/tvm/TVMOpsCompiler.py +++ b/src/xtc/backends/tvm/TVMOpsCompiler.py @@ -3,12 +3,12 @@ # 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 +from typing import Any, TypeAlias, cast import tvm import tvm.te as te +import tvm.s_tir from .TVMOps import ( TVMBaseExpr, @@ -20,24 +20,22 @@ "TVMExprCompiler", "TVMSchedulableExpr", "TVMSchedulableExpr", - "TVMSchedulableExprTE", "TVMSchedulableExprTIR", "TVMScheduledExpr", - "TVMScheduledExprTE", "TVMScheduledExprTIR", ] TETensor: TypeAlias = te.Tensor -TIRSchedule: TypeAlias = tvm.tir.Schedule -TIRFunc: TypeAlias = tvm.tir.PrimFunc +TIRSchedule: TypeAlias = tvm.s_tir.Schedule +TIRFunc: TypeAlias = tvm.tirx.PrimFunc TESchedule: TypeAlias = Any # te.Schedule not available on tvm > 0.19 +TEParam: TypeAlias = te.Tensor | tvm.tirx.Var class TVMExprCompiler: - def __init__(self, expr: TVMBaseExpr, tir_schedule: bool = True): + def __init__(self, expr: TVMBaseExpr): self._expr = expr - self._tir_schedule = tir_schedule def generate(self) -> "TVMSchedulableExpr": if isinstance(self._expr, TVMGraph): @@ -48,10 +46,9 @@ def generate(self) -> "TVMSchedulableExpr": 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) + args = cast(list[TEParam], params) + prim_func = te.create_prim_func(args) + return TVMSchedulableExprTIR(self._expr, prim_func) class TVMSchedulableExpr(ABC): @@ -63,34 +60,6 @@ def schedule(self, schedule: Any = None) -> "TVMScheduledExpr": ... 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 @@ -106,10 +75,9 @@ 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) + sch = tvm.s_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(): @@ -127,23 +95,6 @@ def schedulable(self) -> TVMSchedulableExpr: ... 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 diff --git a/src/xtc/backends/tvm/TVMScheduler.py b/src/xtc/backends/tvm/TVMScheduler.py index 472d8a073..afe7702a3 100644 --- a/src/xtc/backends/tvm/TVMScheduler.py +++ b/src/xtc/backends/tvm/TVMScheduler.py @@ -33,228 +33,6 @@ class TVMScheduleEmitter(ABC): def emit(self, scheduler: "TVMScheduler"): ... -class TVMScheduleEmitterTE(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 _parallel_dims(self, sched: LoopNest) -> list[str]: - op_dims = self._op.operator.dims() - return [ - sched.abstract_dims[op_dims.index(d)] for d in self._op.operator.dims("P") - ] - - def _reduction_dims(self, sched: LoopNest) -> list[str]: - op_dims = self._op.operator.dims() - return [ - sched.abstract_dims[op_dims.index(d)] for d in self._op.operator.dims("R") - ] - - def _full_packs(self, sched: LoopNestNode) -> dict[str, tuple[int, int, int, int]]: - packs = {} - for axis, (input_idx, _, pad) in sched.pack_at.items(): - dim, factor, offset = tvm_cache_read_factor_offset(self._op, input_idx, pad) - packs[axis] = (input_idx, dim, factor, offset) - return packs - - def _full_fuses(self, sched: LoopNestNode) -> dict[str, tuple[int]]: - fuses = {} - for axis, input_idx in sched.fuse_producer_at.items(): - fuses[axis] = (input_idx,) - return fuses - - def _full_tilings(self, sched: LoopNestNode) -> dict[str, tuple[str, str, int]]: - order = sched.interchange - tiles = sched.tiles - tilings = {} - for dim, dim_tiles in tiles.items(): - prev_axis = dim - tilings[dim] = (dim, "", 0) - for idx, (axis, size) in enumerate(dim_tiles.items()): - tilings[axis] = (dim, prev_axis, size) - prev_axis = axis - tilings = {axis: tilings.get(axis, (axis, "", 0)) for axis in order} - return tilings - - def _write_buffer_tiling( - self, - sched: LoopNestNode, - axis: str, - parallel_dims: list[str], - tilings: dict[str, tuple[str, str, int]], - ) -> tuple[dict[str, tuple[str, str, int]], dict[str, tuple[str, str, int]]]: - child = { - parent: (dim, axis, factor) - for axis, (dim, parent, factor) in tilings.items() - } - tiles_axis = list(tilings) - tile_idx = tiles_axis.index(axis) - outer_tiles = dict(list(tilings.items())[: tile_idx + 1]) - inner_tiles = dict(list(tilings.items())[tile_idx + 1 :]) - outers_dims = set() - for axis, (dim, parent, factor) in list(outer_tiles.items()): - outers_dims.add(dim) - factor = child.get(axis, ("", "", 1))[2] - outer_tiles[f"{dim}_"] = (dim, axis, factor) - dims = set() - for axis, (dim, parent, factor) in list(inner_tiles.items()): - if outers_dims and dim not in outers_dims: - outers_dims.add(dim) - if dim in parallel_dims: - outer_tiles[f"{dim}_"] = (dim, dim, 0) - if dim not in dims: - dims.add(dim) - parent = "" - inner_tiles[axis] = (dim, parent, factor) - return outer_tiles, inner_tiles - - def _full_write_buffers( - self, sched: LoopNestNode, parallel_dims: list[str] - ) -> dict[tuple[str, str, str], dict[str, tuple[str, str, int]]]: - tilings = self._full_tilings(sched) - reorder_idx = {axis: idx for idx, axis in enumerate(tilings)} - write_axis = list(sched.buffer_at) - write_axis = sorted(write_axis, key=lambda axis: reorder_idx[axis]) - buffer_tilings = {} - out = ("O", "", "") - tiling = tilings - for idx, axis in enumerate(write_axis): - outer_tiling, inner_tiling = self._write_buffer_tiling( - sched, axis, parallel_dims, tiling - ) - buffer_tilings[out] = outer_tiling - out = (f"O_W{idx}", out[0], axis) - tiling = inner_tiling - buffer_tilings[out] = tiling - return buffer_tilings - - def _emit_assign_axis( - self, - sch: str, - tens: str, - parallel_dims: list[str], - reduction_dims: list[str], - outf: TextIO, - ) -> None: - if parallel_dims: - print(f"{', '.join(parallel_dims)}, = {tens}.op.axis", file=outf) - if reduction_dims: - print(f"{', '.join(reduction_dims)}, = {tens}.op.reduce_axis", file=outf) - - def _emit_assign_tilings( - self, - sch: str, - tens: str, - tilings: dict[str, tuple[str, str, int]], - outf: TextIO, - ) -> None: - for axis, (dim, parent, factor) in tilings.items(): - if not parent: - if axis != dim: - print(f"{axis} = {dim}", file=outf) - continue - if factor > 0: - print( - f"{parent}, {axis} = {sch}[{tens}].split({parent}, factor={factor})", - file=outf, - ) - else: - print(f"{axis} = {parent}", file=outf) - print(f"{sch}[{tens}].reorder({', '.join(tilings)})", file=outf) - - 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 not node.splits, "split not supported for TE Schedule" - - parallel_dims = self._parallel_dims(sched) - reduction_dims = self._reduction_dims(sched) - tilings = self._full_write_buffers(node, parallel_dims) - packings = self._full_packs(node) - fuses = self._full_fuses(node) - obj = self._obj_var - sch = self._sch_var - outf = self._outf - if packings: - print(f"INPS = list({obj}.values())[:-1]", file=outf) - for (tens, parent, axis), tiles in tilings.items(): - if not parent: - print(f"{tens} = {obj}['{self._op.name}']", file=outf) - else: - print(f'{tens} = {sch}.cache_write({parent}, "global")', file=outf) - for tile_axis in tiles: - if tile_axis in packings: - inp_idx, _, _, _ = packings[tile_axis] - print( - f'I_R{inp_idx} = {sch}.cache_read(INPS[{inp_idx}], "global", [{tens}])', - file=outf, - ) - if tile_axis in fuses: - (inp_idx,) = fuses[tile_axis] - print( - f"I_F{inp_idx} = {tens}.op.input_tensors[{inp_idx}]", file=outf - ) - for idx, ((tens, parent, axis), tiles) in enumerate(tilings.items()): - if parent: - print(f"{sch}[{tens}].compute_at({sch}[{parent}], {axis})", file=outf) - self._emit_assign_axis(sch, tens, parallel_dims, reduction_dims, outf) - self._emit_assign_tilings(sch, tens, tiles, outf) - for tile_axis in tiles: - if tile_axis in packings: - inp_idx, dim, factor, offset = packings[tile_axis] - print( - f"{sch}[I_R{inp_idx}].compute_at({sch}[{tens}], {tile_axis})", - file=outf, - ) - if factor != 0: - print( - f"{sch}[I_R{inp_idx}].storage_align(I_R{inp_idx}.op.axis[{dim}], factor={factor}, offset={offset})", - file=outf, - ) - if tile_axis in fuses: - (inp_idx,) = fuses[tile_axis] - print( - f"{sch}[I_F{inp_idx}].compute_at({sch}[{tens}], {tile_axis})", - file=outf, - ) - for u_axis in node.unroll: - if u_axis in tiles: - print(f"{sch}[{tens}].unroll({u_axis})", file=outf) - for v_axis in node.vectorize: - if v_axis in tiles: - print(f"{sch}[{tens}].vectorize({v_axis})", file=outf) - if node.parallelize: - if node.parallelize[0] in tiles: - if len(node.parallelize) > 1: - print( - f"{node.parallelize[-1]} = {sch}[{tens}].fuse({', '.join(node.parallelize)})", - file=outf, - ) - print( - f"{sch}[{tens}].parallel({node.parallelize[-1]})", - file=outf, - ) - - @override - def emit(self, scheduler: "TVMScheduler"): - sched = scheduler.get_loop_nest() - sched = tvm_update_loopnest_for_codegen(sched) - sched.check() - self._dump_schedule(sched) - - class TVMScheduleEmitterTIR(TVMScheduleEmitter): def __init__( self, @@ -281,7 +59,7 @@ def _dump_schedule_node(self, sched: LoopNest, node: LoopNestNode): outf = self._outf dims = sched.abstract_dims block = "O" - print(f'{block} = {sch}.get_block("{self._op.name}")', file=outf) + print(f'{block} = {sch}.get_sblock("{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) @@ -499,10 +277,7 @@ def backend(self) -> itf.back.Backend: @override def schedule(self) -> itf.schd.Schedule: io = StringIO() - if self._backend._tir_schedule: - emitter: TVMScheduleEmitter = TVMScheduleEmitterTIR(op=self._op, outf=io) - else: - emitter = TVMScheduleEmitterTE(op=self._op, outf=io) + emitter = TVMScheduleEmitterTIR(op=self._op, outf=io) emitter.emit(self) sched = io.getvalue() assert self._op.name is not None diff --git a/src/xtc/cli/explore.py b/src/xtc/cli/explore.py index 68cf34683..6c9ecbe2f 100644 --- a/src/xtc/cli/explore.py +++ b/src/xtc/cli/explore.py @@ -277,12 +277,6 @@ 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/csrcs/runtimes/host/evaluate_perf.c b/src/xtc/csrcs/runtimes/host/evaluate_perf.c index fe64d1aff..89004e341 100644 --- a/src/xtc/csrcs/runtimes/host/evaluate_perf.c +++ b/src/xtc/csrcs/runtimes/host/evaluate_perf.c @@ -18,14 +18,18 @@ typedef void (*func4_t)(void *,void *,void *,void *); typedef void (*func5_t)(void *,void *,void *,void *,void *); typedef void (*func6_t)(void *,void *,void *,void *,void *,void *); -typedef union { - int64_t v_int64; - double v_float64; - void *v_handle; - char *v_str; +typedef struct { + int32_t type_index; + int32_t zero_padding; + union { + int64_t v_int64; + double v_float64; + void *v_ptr; + char *v_str; + }; } PackedArg; -typedef int (*packed_func_t)(PackedArg *, int *, int, PackedArg *, int *); +typedef int (*packed_func_t)(void *, PackedArg *, int, PackedArg *); #define mem_barrier() asm("":::"memory") @@ -126,12 +130,13 @@ typedef int (*packed_func_t)(PackedArg *, int *, int, PackedArg *, int *); void evaluate_packed_perf(double *results, int events_num, const char *events_names[], int repeat, int number, int min_repeat_ms, - packed_func_t func, PackedArg *args, int *codes, int nargs) + packed_func_t func, PackedArg *args, int nargs) { PackedArg res; - int res_code = 0; + res.type_index = 0; + res.zero_padding = 0; res.v_int64 = 0; - define_evaluateN(func, args, codes, nargs, &res, &res_code); + define_evaluateN(func, NULL, args, nargs, &res); } void evaluate0_perf(double *results, int events_num, const char *events_names[], @@ -240,9 +245,9 @@ void evaluate(double *results, } void evaluate_packed(double *results, int repeat, int number, int min_repeat_ms, - packed_func_t func, PackedArg *args, int *codes, int nargs) + packed_func_t func, PackedArg *args, int nargs) { evaluate_packed_perf(results, 0, NULL, repeat, number, min_repeat_ms, - func, args, codes, nargs); + func, args, nargs); } diff --git a/src/xtc/csrcs/runtimes/host/tvm_runtime_init.c b/src/xtc/csrcs/runtimes/host/tvm_runtime_init.c new file mode 100644 index 000000000..c177ee304 --- /dev/null +++ b/src/xtc/csrcs/runtimes/host/tvm_runtime_init.c @@ -0,0 +1,30 @@ +/* + * SPDX-License-Identifier: BSD-3-Clause + * Copyright (c) 2024-2026 The XTC Project Authors + */ +/* + * Minimal TVM runtime init function. + * + * Initialize each TVM API call slot with the definition + * from tvm_runtime shared lib (or custom implementation). + */ +#include + +/* We omit actuall function prototype as we only assign the slots there */ +extern void TVMBackendAllocWorkspace(); +extern void TVMBackendFreeWorkspace(); +extern void TVMBackendParallelBarrier(); +extern void TVMBackendParallelLaunch(); + +void (*__TVMBackendAllocWorkspace)(); +void (*__TVMBackendFreeWorkspace)(); +void (*__TVMBackendParallelBarrier)(); +void (*__TVMBackendParallelLaunch)(); + +void xtc_tvm_init_runtime() +{ + __TVMBackendAllocWorkspace = TVMBackendAllocWorkspace; + __TVMBackendFreeWorkspace = TVMBackendFreeWorkspace; + __TVMBackendParallelBarrier = TVMBackendParallelBarrier; + __TVMBackendParallelLaunch = TVMBackendParallelLaunch; +} diff --git a/src/xtc/itf/runtime/common.py b/src/xtc/itf/runtime/common.py index f42b4c76d..31aade660 100644 --- a/src/xtc/itf/runtime/common.py +++ b/src/xtc/itf/runtime/common.py @@ -102,7 +102,6 @@ def evaluate_packed( min_repeat_ms: int, cfunc: CFunc, args: Any, - codes: Any, nargs: int, ) -> None: """Evaluate a packed function with timing measurements. @@ -114,7 +113,6 @@ def evaluate_packed( min_repeat_ms: Minimum time in milliseconds for each repeat. cfunc: Packed function pointer to evaluate. args: Pointer to array of packed arguments. - codes: Pointer to array of integers containing argument type codes. nargs: Number of arguments. """ ... @@ -129,7 +127,6 @@ def evaluate_packed_perf( min_repeat_ms: int, cfunc: CFunc, args: Any, - codes: Any, nargs: int, ) -> None: """Evaluate a packed function with performance counter measurements. @@ -142,7 +139,6 @@ def evaluate_packed_perf( min_repeat_ms: Minimum time in milliseconds for each repeat. cfunc: Packed function pointer to evaluate. args: Pointer to array of packed arguments. - codes: Pointer to array of integers containing argument type codes. nargs: Number of arguments. """ ... diff --git a/src/xtc/runtimes/accelerator/gpu/GPUDevice.py b/src/xtc/runtimes/accelerator/gpu/GPUDevice.py index 2fc3a4e1b..e84dfe62a 100644 --- a/src/xtc/runtimes/accelerator/gpu/GPUDevice.py +++ b/src/xtc/runtimes/accelerator/gpu/GPUDevice.py @@ -271,7 +271,6 @@ def evaluate_packed( min_repeat_ms: int, cfunc: CFunc, args: Any, - codes: Any, nargs: int, ) -> None: raise NotImplementedError("evaluate_packed is not implemented for GPU device") @@ -286,7 +285,6 @@ def evaluate_packed_perf( min_repeat_ms: int, cfunc: CFunc, args: Any, - codes: Any, nargs: int, ) -> None: raise NotImplementedError( diff --git a/src/xtc/runtimes/accelerator/mppa/MppaDevice.py b/src/xtc/runtimes/accelerator/mppa/MppaDevice.py index 33a33801e..714f5609d 100644 --- a/src/xtc/runtimes/accelerator/mppa/MppaDevice.py +++ b/src/xtc/runtimes/accelerator/mppa/MppaDevice.py @@ -726,7 +726,6 @@ def evaluate_packed( min_repeat_ms: int, cfunc: CFunc, args: Any, - codes: Any, nargs: int, ) -> None: raise NotImplementedError("evaluate_packed is not implemented for MPPA device") @@ -741,7 +740,6 @@ def evaluate_packed_perf( min_repeat_ms: int, cfunc: CFunc, args: Any, - codes: Any, nargs: int, ) -> None: raise NotImplementedError( diff --git a/src/xtc/runtimes/host/HostRuntime.py b/src/xtc/runtimes/host/HostRuntime.py index 36ff54307..cb43eb7dd 100644 --- a/src/xtc/runtimes/host/HostRuntime.py +++ b/src/xtc/runtimes/host/HostRuntime.py @@ -128,7 +128,6 @@ def evaluate_packed( min_repeat_ms: int, cfunc: CFunc, args: Any, - codes: Any, nargs: int, ) -> None: self.__get_runtime_func("evaluate_packed")( @@ -138,7 +137,6 @@ def evaluate_packed( ctypes.c_int(min_repeat_ms), ctypes.cast(cfunc.handle, ctypes.CFUNCTYPE(ctypes.c_voidp)), ctypes.cast(args, ctypes.POINTER(ctypes.c_voidp)), - ctypes.cast(codes, ctypes.POINTER(ctypes.c_int)), ctypes.c_int(nargs), ) @@ -152,7 +150,6 @@ def evaluate_packed_perf( min_repeat_ms: int, cfunc: CFunc, args: Any, - codes: Any, nargs: int, ) -> None: self.__get_runtime_func("evaluate_packed_perf")( @@ -164,7 +161,6 @@ def evaluate_packed_perf( ctypes.c_int(min_repeat_ms), ctypes.cast(cfunc.handle, ctypes.CFUNCTYPE(ctypes.c_voidp)), ctypes.cast(args, ctypes.POINTER(ctypes.c_voidp)), - ctypes.cast(codes, ctypes.POINTER(ctypes.c_int)), ctypes.c_int(nargs), ) diff --git a/src/xtc/runtimes/host/runtime.py b/src/xtc/runtimes/host/runtime.py index 4f10bf02d..b6300ffe0 100644 --- a/src/xtc/runtimes/host/runtime.py +++ b/src/xtc/runtimes/host/runtime.py @@ -74,7 +74,6 @@ def from_param(obj: str | bytes): ctypes.c_int, ctypes.CFUNCTYPE(ctypes.c_voidp), ctypes.POINTER(ctypes.c_voidp), - ctypes.POINTER(ctypes.c_int), ctypes.c_int, ], "restype": None, @@ -90,7 +89,6 @@ def from_param(obj: str | bytes): ctypes.c_int, ctypes.CFUNCTYPE(ctypes.c_voidp), ctypes.POINTER(ctypes.c_voidp), - ctypes.POINTER(ctypes.c_int), ctypes.c_int, ], "restype": None, diff --git a/src/xtc/search/explore.py b/src/xtc/search/explore.py index e47bcc7a8..e0cb9c577 100644 --- a/src/xtc/search/explore.py +++ b/src/xtc/search/explore.py @@ -111,7 +111,6 @@ 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): @@ -351,8 +350,6 @@ def compile_one( 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( diff --git a/src/xtc/targets/host/HostAREvaluator.py b/src/xtc/targets/host/HostAREvaluator.py index 324daf9ee..934e3644b 100644 --- a/src/xtc/targets/host/HostAREvaluator.py +++ b/src/xtc/targets/host/HostAREvaluator.py @@ -92,7 +92,7 @@ def __init__(self, module: "host.HostModule", **kwargs: Any) -> None: module=module, repeat=1, min_repeat_ms=0, - number=1, + number=0, **kwargs, ) diff --git a/src/xtc/targets/host/HostCEvaluator.py b/src/xtc/targets/host/HostCEvaluator.py index 9999cce94..792da122f 100644 --- a/src/xtc/targets/host/HostCEvaluator.py +++ b/src/xtc/targets/host/HostCEvaluator.py @@ -52,10 +52,13 @@ def _compile_to_shlib(self, shlib_base: str): hdrs_path += self._module.headers_path hdrs_path = list(dict.fromkeys(hdrs_path)) hdrs_opts = [f"-I{path}" for path in hdrs_path] - ext = ".dylib" if sys.platform == "darwin" else ".so" + ext = ".so" + if sys.platform == "darwin": + sh_opts += " -undefined dynamic_lookup" + ext = ".dylib" cmd = ( f"cc {sh_opts} {opts} {' '.join(hdrs_opts)} {' '.join(csrcs)} " - f"-o {shlib_name}{ext}" + f"-o {shlib_name}{ext} " ) p = subprocess.run( shlex.split(cmd), text=True, capture_output=True, cwd=cwd_dir @@ -88,7 +91,7 @@ def __init__(self, module: "host.HostModule", **kwargs: Any) -> None: module=module, repeat=1, min_repeat_ms=0, - number=1, + number=0, **kwargs, ) diff --git a/src/xtc/targets/host/HostEvaluator.py b/src/xtc/targets/host/HostEvaluator.py index a04ae4280..fd655fae5 100644 --- a/src/xtc/targets/host/HostEvaluator.py +++ b/src/xtc/targets/host/HostEvaluator.py @@ -105,7 +105,7 @@ def __init__(self, module: "host.HostModule", **kwargs: Any) -> None: module=module, repeat=1, min_repeat_ms=0, - number=1, + number=0, **kwargs, ) diff --git a/src/xtc/templates/tvm/packed_op_wrapper.c.jinja b/src/xtc/templates/tvm/packed_op_wrapper.c.jinja index 4e12103bc..dad29f33a 100644 --- a/src/xtc/templates/tvm/packed_op_wrapper.c.jinja +++ b/src/xtc/templates/tvm/packed_op_wrapper.c.jinja @@ -1,9 +1,31 @@ #include #include -/* Refer to tvm c_runtime_api.h */ -#define kTVMArgInt 0 -#define kTVMDLTensorHandle 7 +#ifdef __cplusplus +extern "C" { +#endif + +extern void xtc_tvm_init_runtime(); +static __attribute__((constructor)) void _init() { + xtc_tvm_init_runtime(); +} + + +/* From tvm/ffi/c_api.h (TVM 0.26) */ +#define kTVMFFINone 0 +#define kTVMFFIDLTensorPtr 7 + +typedef struct { + int32_t type_index; + union { + uint32_t zero_padding; + uint32_t small_str_len; + }; + union { + int64_t v_int64; + void* v_ptr; + }; +} TVMFFIAny; /* Refer to dlpack.h */ #define kDLCPU 1 @@ -31,16 +53,8 @@ typedef struct { } DLTensor; /* Bridge bare pointers to TVM packed function with DLTensor objects */ -#ifdef __cplusplus -extern "C" -#else -extern -#endif -void {{packed_func_name}}(void *args[], const int32_t *args_types, int32_t num_args, void *res, const int32_t *res_types, void *resource_manager); +extern int32_t {{tvm_ffi_func_name}}(void *ctx, TVMFFIAny *args, int32_t num_args, TVMFFIAny *result); -#ifdef __cplusplus -extern "C" -#endif void {{func_name}}({% for idx in range(inputs|length) %}{% if idx > 0 %}, {% endif %}const float *input{{idx}}{% endfor %}{% for idx in range(outputs|length) %}, float *output{{idx}}{% endfor %}) { static const DLDataType dtype = { kDLFloat, 32, 1 }; static const DLDevice dev = { kDLCPU, 0 }; @@ -52,9 +66,30 @@ void {{func_name}}({% for idx in range(inputs|length) %}{% if idx > 0 %}, {% end static const int64_t DL_output{{idx}}_shape[{{outputs[idx].shape|length}}] = { {% for dim in outputs[idx].shape %}{{dim}}, {% endfor %} }; DLTensor DL_output{{idx}} = { (void *)output{{idx}}, dev, {{outputs[idx].shape|length}}, dtype, (int64_t *)DL_output{{idx}}_shape, NULL, 0 }; {% endfor %} - void *args[{{inputs|length}} + {{outputs|length}}] = { {% for idx in range(inputs|length) %}&DL_input{{idx}}, {% endfor %}{% for idx in range(outputs|length) %}&DL_output{{idx}}, {% endfor %} }; - const int32_t types[{{inputs|length}} + {{outputs|length}}] = { {% for idx in range(inputs|length) %}kTVMDLTensorHandle, {% endfor %}{% for idx in range(outputs|length) %}kTVMDLTensorHandle, {% endfor %} }; - int64_t res; - int32_t res_type = kTVMArgInt; - {{packed_func_name}}(args, types, {{inputs|length}} + {{outputs|length}}, &res, &res_type, NULL); + TVMFFIAny args[{{inputs|length}} + {{outputs|length}}] = { + {%- for idx in range(inputs|length) %} + { + .type_index = kTVMFFIDLTensorPtr, + .zero_padding = 0, + .v_ptr = &DL_input{{idx}}, + }, + {%- endfor %} + {%- for idx in range(outputs|length) %} + { + .type_index = kTVMFFIDLTensorPtr, + .zero_padding = 0, + .v_ptr = &DL_output{{idx}}, + }, + {%- endfor %} + }; + TVMFFIAny res = { + .type_index = kTVMFFINone, + .zero_padding = 0, + .v_int64 = 0, + }; + (void){{tvm_ffi_func_name}}(NULL, args, {{inputs|length}} + {{outputs|length}}, &res); } + +#ifdef __cplusplus +} +#endif diff --git a/src/xtc/templates/tvm/tvm_ffi_op_wrapper.c.jinja b/src/xtc/templates/tvm/tvm_ffi_op_wrapper.c.jinja new file mode 100644 index 000000000..557f19b88 --- /dev/null +++ b/src/xtc/templates/tvm/tvm_ffi_op_wrapper.c.jinja @@ -0,0 +1,23 @@ +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +extern void xtc_tvm_init_runtime(); +static __attribute__((constructor)) void _init() { + xtc_tvm_init_runtime(); +} + + +/* Forward call to __tvm_ffi packed function */ +extern int32_t {{tvm_ffi_func_name}}(void *ctx, void *args, int32_t num_args, void *result); + +void {{func_name}}(void *ctx, void *args, int32_t num_args, void *result) { + (void){{tvm_ffi_func_name}}(ctx, args, num_args, result); +} + +#ifdef __cplusplus +} +#endif diff --git a/src/xtc/utils/cfunc.py b/src/xtc/utils/cfunc.py index bc5ee41be..828c79d81 100644 --- a/src/xtc/utils/cfunc.py +++ b/src/xtc/utils/cfunc.py @@ -7,106 +7,114 @@ __all__ = [ "CFunc", - "CArgValue", - "CArgCode", - "CRetValue", - "CPackedFunc", "_c_ascii_str", "_str_list_to_c", ] -class ArgTypeCode: - INT = 0 - HANDLE = 3 - NDARRAY_HANDLE = 13 - +# TVM 0.26 FFI ABI +class CTVMFFIAnyValue(ctypes.Union): + _fields_ = [ + ("v_int64", ctypes.c_int64), + ("v_float64", ctypes.c_double), + ("v_ptr", ctypes.c_void_p), + ("v_c_str", ctypes.c_char_p), + ("v_uint64", ctypes.c_uint64), + ] -CArgCode = ctypes.c_int +class CTVMFFIAny(ctypes.Structure): + _anonymous_ = ("value",) -class CArgValue(ctypes.Union): _fields_ = [ - ("v_int64", ctypes.c_int64), - ("v_float64", ctypes.c_double), - ("v_handle", ctypes.c_void_p), - ("v_str", ctypes.c_char_p), + ("type_index", ctypes.c_int32), + ("zero_padding", ctypes.c_uint32), + ("value", CTVMFFIAnyValue), ] -class CRetValue(CArgValue): - pass +class CTVMFFINDArrayArg(CTVMFFIAny): + def __init__(self, arg: Any): + assert arg.__class__.__name__ == "NDArray" + dl_tensor = arg.handle + super().__init__( + type_index=7, zero_padding=0, v_ptr=ctypes.cast(dl_tensor, ctypes.c_void_p) + ) + + +class CTVMFFIResult(CTVMFFIAny): + def __init__(self): + super().__init__() -CPackedFunc = ctypes.CFUNCTYPE( - ctypes.c_int, - ctypes.POINTER(CArgValue), - ctypes.POINTER(CArgCode), - ctypes.c_int, - ctypes.POINTER(CRetValue), - ctypes.POINTER(CArgCode), +CTVMFFIPackedFunc = ctypes.CFUNCTYPE( + ctypes.c_int32, + ctypes.c_void_p, + ctypes.POINTER(CTVMFFIAny), + ctypes.c_int32, + ctypes.POINTER(CTVMFFIAny), ) class CFunc: - def __init__(self, f: Any, packed: bool = False) -> None: - self.handle = f - self.is_packed = packed or ( - hasattr(self.handle, "packed") and self.handle.packed - ) + _supported_abis = ["bare", "tvm_ffi"] - def arg_tuple(self, arg: Any) -> Any: + def __init__(self, f: Any, abi: str | None = None) -> None: + self.handle = f + self.abi = abi + if self.abi is None: + if hasattr(self.handle, "packed") and self.handle.packed: + # TODO: for now infer tvm_ffi abi for packed + self.abi = "tvm_ffi" + if self.abi is None: + self.abi = "bare" + assert self.abi in self._supported_abis + + def _mangled_arg(self, arg: Any) -> Any: if arg.__class__.__name__ == "ndarray": # Numpy Array - assert not self.is_packed - return (arg.ctypes.data_as(ctypes.c_voidp), ArgTypeCode.HANDLE) + assert self.abi == "bare" + return arg.ctypes.data_as(ctypes.c_voidp) elif arg.__class__.__name__ == "NDArray": # TVM NDArray or our NDArray if ( hasattr(arg, "is_on_device") and arg.is_on_device() ): # Device living NDArray - if self.is_packed: - raise RuntimeError("TODO: device NDArray not supported yet") - else: - return ( - ctypes.cast(arg.data, ctypes.c_void_p), - ArgTypeCode.HANDLE, - ) - if self.is_packed: - return ( - CArgValue(v_handle=ctypes.cast(arg.handle, ctypes.c_void_p)), - ArgTypeCode.NDARRAY_HANDLE, - ) + assert self.abi == "bare", "TODO: device NDArray not supported yet" + if self.abi == "tvm_ffi": + assert self.abi == "tvm_ffi" + return CTVMFFINDArrayArg(arg) else: - return ( - ctypes.cast(arg.handle.contents.dl_tensor.data, ctypes.c_void_p), - ArgTypeCode.HANDLE, - ) + assert self.abi == "bare" + return ctypes.cast(arg.data, ctypes.c_voidp) else: assert 0, f"Unsupported argument class: {arg.__class__.__name__}" - def args_tuples(self, args: Any) -> list[Any]: - return [self.arg_tuple(arg) for arg in args] + def get_args_list(self, args: list[Any]) -> list[Any]: + return [self._mangled_arg(arg) for arg in args] + + def get_ctypes_args(self, args: list[Any]) -> Any: + args_list = self.get_args_list(args) + if self.abi == "tvm_ffi": + return (CTVMFFIAny * len(args_list))(*args_list) + else: + return (ctypes.c_voidp * len(args_list))(*args_list) def __call__(self, *args: Any): - args_tuples = self.args_tuples(args) - if self.is_packed: - args_array = (CArgValue * len(args_tuples))( - *[arg[0] for arg in args_tuples] - ) - args_codes = (CArgCode * len(args_tuples))(*[arg[1] for arg in args_tuples]) - result_val = CRetValue(0) - result_code = CArgCode(ArgTypeCode.INT) - res = CPackedFunc(self.handle)( - args_array, - args_codes, - len(args_tuples), - ctypes.byref(result_val), - ctypes.byref(result_code), - ctypes.c_int(len(args_tuples)), + func_addr = ctypes.cast(self.handle, ctypes.c_voidp).value + assert func_addr is not None + ctypes_args = self.get_ctypes_args(list(args)) + if self.abi == "tvm_ffi": + result = CTVMFFIResult() + ctx = ctypes.c_void_p() + CTVMFFIPackedFunc(func_addr)( + ctx, + ctypes_args, + len(ctypes_args), + ctypes.byref(result), ) - assert res == 0, f"error calling packed function" + assert result.v_int64 == 0, f"error calling packed function" else: - data_args = [arg[0] for arg in args_tuples] - self.handle(*data_args) + func_type = ctypes.CFUNCTYPE(None, *([ctypes.c_void_p] * len(ctypes_args))) + func_type(func_addr)(*ctypes_args) class _c_ascii_str: diff --git a/src/xtc/utils/evaluation.py b/src/xtc/utils/evaluation.py index d0fceb6f4..ad370793a 100644 --- a/src/xtc/utils/evaluation.py +++ b/src/xtc/utils/evaluation.py @@ -11,7 +11,7 @@ from xtc.graphs.xtc.graph import XTCGraph from xtc.graphs.xtc.expr import XTCTensorExpr from xtc.graphs.xtc.data import XTCTensor -from xtc.utils.cfunc import CFunc, CArgValue, CArgCode +from xtc.utils.cfunc import CFunc from xtc.itf.runtime.common import CommonRuntimeInterface from xtc.runtimes.host.HostRuntime import HostRuntime @@ -156,19 +156,13 @@ def evaluate_performance( ) -> tuple[list[float], int, str]: # TODO migrate host runtime to CommonRuntimeInterface cfunc = CFunc(func) - args_tuples = cfunc.args_tuples([*parameters[0], *parameters[1]]) + args_array = cfunc.get_ctypes_args([*parameters[0], *parameters[1]]) values_num = 1 if len(pmu_counters) > 0: values_num = len(pmu_counters) # FIXME check if the PMU counters are supported by the target results_array = (ctypes.c_double * (repeat * values_num))() - if cfunc.is_packed: - args_array_packed = (CArgValue * len(args_tuples))( - *[arg[0] for arg in args_tuples] - ) - args_codes_packed = (CArgCode * len(args_tuples))( - *[arg[1] for arg in args_tuples] - ) + if cfunc.abi == "tvm_ffi": runtime.evaluate_packed_perf( results_array, pmu_counters, @@ -176,15 +170,12 @@ def evaluate_performance( number, min_repeat_ms, cfunc, - args_array_packed, - args_codes_packed, - len(args_tuples), + args_array, + len(args_array), ) eval_results = [float(x) for x in results_array] else: - args_array = (ctypes.c_voidp * len(args_tuples))( - *[arg[0] for arg in args_tuples] - ) + assert cfunc.abi == "bare" eval_results = runtime.evaluate_perf( pmu_counters, repeat, diff --git a/src/xtc/utils/files.py b/src/xtc/utils/files.py new file mode 100644 index 000000000..8ea1cd093 --- /dev/null +++ b/src/xtc/utils/files.py @@ -0,0 +1,15 @@ +# +# SPDX-License-Identifier: BSD-3-Clause +# Copyright (c) 2024-2026 The XTC Project Authors +# +from pathlib import Path + + +def relative_to(file: str | Path, directory: str | Path) -> Path: + """Return `file` relative to `directory` if contained within it, otherwise absolute.""" + file = Path(file).resolve() + directory = Path(directory).resolve() + try: + return file.relative_to(directory) + except ValueError: + return file diff --git a/src/xtc/utils/host_tools.py b/src/xtc/utils/host_tools.py index bec4250ea..3a54698d4 100644 --- a/src/xtc/utils/host_tools.py +++ b/src/xtc/utils/host_tools.py @@ -60,6 +60,28 @@ def binutils_command(command: str, arch: str = "") -> str: return f"{prefix}{command}" +def cc_prefix(arch: str = "") -> str: + """ + Returns the cc prefix given the arch. + When arch is unspecified, assume no prefix. + """ + triple = target_triple(arch) + if not triple: + return "" + return f"{triple}-" + + +def cc_command(arch: str = "") -> str: + """ + Returns the cc compiler for the given arch. + For native prefer cc over gcc. + """ + prefix = cc_prefix(arch) + if not prefix: + return "cc" + return f"{prefix}gcc" + + def disassemble( obj_path: str | Path, function: str = "", @@ -70,22 +92,26 @@ def disassemble( ) -> str: """ Returns the disassembled multi-line string for the given obj_path. - Optionally disassembling only the given function. + Optionally disassembling only the given section or function. + Note that section name for Darwin are rewritten .text -> __text + and function name function -> _function. """ base_opts = [ "-dr", "--no-addresses", "--no-show-raw-insn", + "--disassemble", ] target = target_arch(arch) jumps_opts: list[str] = [] - disass_symbol_opt: list[str] = [] if target in ["x86_64", "aarch64"] or platform.system() == "Linux": jumps_opts = [] - disass_symbol_opt = [f"--disassemble={function}"] elif platform.system() == "Darwin": + if section and section.startswith("."): + section = f"__{section[1:]}" + if function: + function = f"_{function}" jumps_opts = [] - disass_symbol_opt = ["--disassemble-symbols=ltmp0"] color_opts = [ "--disassembler-color=on", ] @@ -93,7 +119,6 @@ def disassemble( obj_path = Path(obj_path) args = [ *base_opts, - *(disass_symbol_opt if function else ["--disassemble"]), *(jumps_opts if visualize_jumps else []), *(color_opts if color else []), str(obj_path), @@ -113,13 +138,19 @@ def disassemble( f" error code: {p.returncode}" ) output = "" - emit = False - # Filter out file header and optionally section + in_section = False + in_function = False + # Dump function or section when specified for line in p.stdout.splitlines(): if "section" in line: - emit = True - if section and not f"section {section}" in line: - emit = False - if emit: + in_function = False + in_section = True + if section and not section in line: + in_section = False + if in_section and line.startswith("<"): + in_function = True + if function and not line.startswith(f"<{function}>"): + in_function = False + if not function and in_section or in_function: output += f"{line}\n" return output diff --git a/tests/filecheck/backends/padding/test_gen_pad_dict_conv2d_tvm.py b/tests/filecheck/backends/padding/test_gen_pad_dict_conv2d_tvm.py index 440210eea..ee4e92204 100644 --- a/tests/filecheck/backends/padding/test_gen_pad_dict_conv2d_tvm.py +++ b/tests/filecheck/backends/padding/test_gen_pad_dict_conv2d_tvm.py @@ -39,54 +39,61 @@ # CHECK-NEXT: outputs: # CHECK-NEXT: - %3 : 1x4x4x16xfloat32 # CHECK-NEXT: nodes: -# CHECK-NEXT: - %2: pad(%0, padding={1: (2, 2), 2: (2, 2)}, constant_value=0) {name = 'pad'} : [1x8x8x3xfloat32] -> [1x12x12x3xfloat32] +# CHECK-NEXT: - %2: pad(%0, padding={1: (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: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] -# CHECK-NEXT: O = obj['conv'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(b, h, w, f, r, s, c) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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.sblock("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: O = sch.get_sblock("conv") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(b, h, w, f, r, s, c) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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.sblock("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: CODE: 0 diff --git a/tests/filecheck/backends/padding/test_gen_pad_int_matmul_unpad_tvm.py b/tests/filecheck/backends/padding/test_gen_pad_int_matmul_unpad_tvm.py index 9ba621c48..ad4976f88 100644 --- a/tests/filecheck/backends/padding/test_gen_pad_int_matmul_unpad_tvm.py +++ b/tests/filecheck/backends/padding/test_gen_pad_int_matmul_unpad_tvm.py @@ -45,68 +45,85 @@ # CHECK-NEXT: - %5: unpad(%4, padding=(2, 2)) {name = 'C'} : [18x18xfloat32] -> [14x14xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: A_pad = T.allocate([252], "float32", "global") -# CHECK-NEXT: B_pad = T.allocate([252], "float32", "global") -# CHECK-NEXT: matmul_padded = T.allocate([196], "float32", "global") -# CHECK-NEXT: A_pad_1 = T.Buffer((252,), data=A_pad) -# CHECK-NEXT: for i0, i1 in T.grid(14, 18): -# CHECK-NEXT: _0_1 = T.Buffer((196,), data=_0.data) -# CHECK-NEXT: A_pad_1[i0 * 18 + i1] = T.if_then_else(2 <= i1 and i1 < 16, _0_1[i0 * 14 + i1 - 2], T.float32(0.0)) -# CHECK-NEXT: B_pad_1 = T.Buffer((252,), data=B_pad) -# CHECK-NEXT: for i0, i1 in T.grid(18, 14): -# CHECK-NEXT: cse_var_1: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: _1_1 = T.Buffer((196,), data=_1.data) -# CHECK-NEXT: B_pad_1[cse_var_1] = T.if_then_else(2 <= i0 and i0 < 16, _1_1[cse_var_1 - 28], T.float32(0.0)) -# CHECK-NEXT: matmul_padded_1 = T.Buffer((196,), data=matmul_padded) -# CHECK-NEXT: for i, j in T.grid(14, 14): -# CHECK-NEXT: matmul_padded_1[i * 14 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(18): -# CHECK-NEXT: cse_var_2: T.int32 = i * 14 + j -# CHECK-NEXT: matmul_padded_1[cse_var_2] = matmul_padded_1[cse_var_2] + A_pad_1[i * 18 + k] * B_pad_1[k * 14 + j] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_matmul_unpad(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: A_pad = T.sblock_alloc_buffer((18, 18)) +# CHECK-NEXT: B_pad = T.sblock_alloc_buffer((18, 18)) +# CHECK-NEXT: matmul_padded = T.sblock_alloc_buffer((18, 18)) +# CHECK-NEXT: for i0, i1 in T.grid(18, 18): +# CHECK-NEXT: with T.sblock("A_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_0[v_i0 - 2, v_i1 - 2]) +# CHECK-NEXT: T.writes(A_pad[v_i0, v_i1]) +# CHECK-NEXT: A_pad[v_i0, v_i1] = T.if_then_else(2 <= v_i0 and v_i0 < 16 and 2 <= v_i1 and v_i1 < 16, _0[v_i0 - 2, v_i1 - 2], T.float32(0.0)) +# CHECK-NEXT: for i0, i1 in T.grid(18, 18): +# CHECK-NEXT: with T.sblock("B_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_1[v_i0 - 2, v_i1 - 2]) +# CHECK-NEXT: T.writes(B_pad[v_i0, v_i1]) +# CHECK-NEXT: B_pad[v_i0, v_i1] = T.if_then_else(2 <= v_i0 and v_i0 < 16 and 2 <= v_i1 and v_i1 < 16, _1[v_i0 - 2, v_i1 - 2], T.float32(0.0)) +# CHECK-NEXT: for i, j, k in T.grid(18, 18, 18): +# CHECK-NEXT: with T.sblock("matmul_padded"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(A_pad[v_i, v_k], B_pad[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul_padded[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul_padded[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul_padded[v_i, v_j] = matmul_padded[v_i, v_j] + A_pad[v_i, v_k] * B_pad[v_k, v_j] # CHECK-NEXT: for i0, i1 in T.grid(14, 14): -# CHECK-NEXT: cse_var_3: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: C_1 = T.Buffer((196,), data=C.data) -# CHECK-NEXT: C_1[cse_var_3] = matmul_padded_1[cse_var_3] -# CHECK-NEXT: O = obj['matmul_padded'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(i, j, k) +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(matmul_padded[v_i0 + 2, v_i1 + 2]) +# CHECK-NEXT: T.writes(C[v_i0, v_i1]) +# CHECK-NEXT: C[v_i0, v_i1] = matmul_padded[v_i0 + 2, v_i1 + 2] +# CHECK-NEXT: O = sch.get_sblock("matmul_padded") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(i, j, k) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: A_pad = T.allocate([252], "float32", "global") -# CHECK-NEXT: B_pad = T.allocate([252], "float32", "global") -# CHECK-NEXT: matmul_padded = T.allocate([196], "float32", "global") -# CHECK-NEXT: A_pad_1 = T.Buffer((252,), data=A_pad) -# CHECK-NEXT: for i0, i1 in T.grid(14, 18): -# CHECK-NEXT: _0_1 = T.Buffer((196,), data=_0.data) -# CHECK-NEXT: A_pad_1[i0 * 18 + i1] = T.if_then_else(2 <= i1 and i1 < 16, _0_1[i0 * 14 + i1 - 2], T.float32(0.0)) -# CHECK-NEXT: B_pad_1 = T.Buffer((252,), data=B_pad) -# CHECK-NEXT: for i0, i1 in T.grid(18, 14): -# CHECK-NEXT: cse_var_1: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: _1_1 = T.Buffer((196,), data=_1.data) -# CHECK-NEXT: B_pad_1[cse_var_1] = T.if_then_else(2 <= i0 and i0 < 16, _1_1[cse_var_1 - 28], T.float32(0.0)) -# CHECK-NEXT: matmul_padded_1 = T.Buffer((196,), data=matmul_padded) -# CHECK-NEXT: for i, j in T.grid(14, 14): -# CHECK-NEXT: matmul_padded_1[i * 14 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(18): -# CHECK-NEXT: cse_var_2: T.int32 = i * 14 + j -# CHECK-NEXT: matmul_padded_1[cse_var_2] = matmul_padded_1[cse_var_2] + A_pad_1[i * 18 + k] * B_pad_1[k * 14 + j] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_matmul_unpad(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: A_pad = T.sblock_alloc_buffer((18, 18)) +# CHECK-NEXT: B_pad = T.sblock_alloc_buffer((18, 18)) +# CHECK-NEXT: matmul_padded = T.sblock_alloc_buffer((18, 18)) +# CHECK-NEXT: for i0, i1 in T.grid(18, 18): +# CHECK-NEXT: with T.sblock("A_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_0[v_i0 - 2, v_i1 - 2]) +# CHECK-NEXT: T.writes(A_pad[v_i0, v_i1]) +# CHECK-NEXT: A_pad[v_i0, v_i1] = T.if_then_else(2 <= v_i0 and v_i0 < 16 and 2 <= v_i1 and v_i1 < 16, _0[v_i0 - 2, v_i1 - 2], T.float32(0.0)) +# CHECK-NEXT: for i0, i1 in T.grid(18, 18): +# CHECK-NEXT: with T.sblock("B_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_1[v_i0 - 2, v_i1 - 2]) +# CHECK-NEXT: T.writes(B_pad[v_i0, v_i1]) +# CHECK-NEXT: B_pad[v_i0, v_i1] = T.if_then_else(2 <= v_i0 and v_i0 < 16 and 2 <= v_i1 and v_i1 < 16, _1[v_i0 - 2, v_i1 - 2], T.float32(0.0)) +# CHECK-NEXT: for i, j, k in T.grid(18, 18, 18): +# CHECK-NEXT: with T.sblock("matmul_padded"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(A_pad[v_i, v_k], B_pad[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul_padded[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul_padded[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul_padded[v_i, v_j] = matmul_padded[v_i, v_j] + A_pad[v_i, v_k] * B_pad[v_k, v_j] # CHECK-NEXT: for i0, i1 in T.grid(14, 14): -# CHECK-NEXT: cse_var_3: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: C_1 = T.Buffer((196,), data=C.data) -# CHECK-NEXT: C_1[cse_var_3] = matmul_padded_1[cse_var_3] +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(matmul_padded[v_i0 + 2, v_i1 + 2]) +# CHECK-NEXT: T.writes(C[v_i0, v_i1]) +# CHECK-NEXT: C[v_i0, v_i1] = matmul_padded[v_i0 + 2, v_i1 + 2] # CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/padding/test_gen_pad_tuple_conv2d_tvm.py b/tests/filecheck/backends/padding/test_gen_pad_tuple_conv2d_tvm.py index 998bdf224..10c54b8b5 100644 --- a/tests/filecheck/backends/padding/test_gen_pad_tuple_conv2d_tvm.py +++ b/tests/filecheck/backends/padding/test_gen_pad_tuple_conv2d_tvm.py @@ -43,50 +43,57 @@ # CHECK-NEXT: - %3: conv2d(%2, %1, stride=(2, 2)) {name = 'conv'} : [1x12x12x3xfloat32, 5x5x3x16xfloat32] -> [1x4x4x16xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] -# CHECK-NEXT: O = obj['conv'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(b, h, w, f, r, s, c) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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.sblock("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: O = sch.get_sblock("conv") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(b, h, w, f, r, s, c) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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.sblock("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: CODE: 0 diff --git a/tests/filecheck/backends/padding/test_gen_pad_tuple_matmul_unpad_tvm.py b/tests/filecheck/backends/padding/test_gen_pad_tuple_matmul_unpad_tvm.py index 36b21bae1..fb967fdd3 100644 --- a/tests/filecheck/backends/padding/test_gen_pad_tuple_matmul_unpad_tvm.py +++ b/tests/filecheck/backends/padding/test_gen_pad_tuple_matmul_unpad_tvm.py @@ -45,68 +45,85 @@ # CHECK-NEXT: - %5: unpad(%4, padding=(0, 2)) {name = 'C'} : [16x16xfloat32] -> [14x14xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: A_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: B_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: matmul_padded = T.allocate([196], "float32", "global") -# CHECK-NEXT: A_pad_1 = T.Buffer((224,), data=A_pad) -# CHECK-NEXT: for i0, i1 in T.grid(14, 16): -# CHECK-NEXT: _0_1 = T.Buffer((196,), data=_0.data) -# CHECK-NEXT: A_pad_1[i0 * 16 + i1] = T.if_then_else(i1 < 14, _0_1[i0 * 14 + i1], T.float32(0.0)) -# CHECK-NEXT: B_pad_1 = T.Buffer((224,), data=B_pad) -# CHECK-NEXT: for i0, i1 in T.grid(16, 14): -# CHECK-NEXT: cse_var_1: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: _1_1 = T.Buffer((196,), data=_1.data) -# CHECK-NEXT: B_pad_1[cse_var_1] = T.if_then_else(i0 < 14, _1_1[cse_var_1], T.float32(0.0)) -# CHECK-NEXT: matmul_padded_1 = T.Buffer((196,), data=matmul_padded) -# CHECK-NEXT: for i, j in T.grid(14, 14): -# CHECK-NEXT: matmul_padded_1[i * 14 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(16): -# CHECK-NEXT: cse_var_2: T.int32 = i * 14 + j -# CHECK-NEXT: matmul_padded_1[cse_var_2] = matmul_padded_1[cse_var_2] + A_pad_1[i * 16 + k] * B_pad_1[k * 14 + j] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_matmul_unpad(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: A_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: B_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: matmul_padded = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("A_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_0[v_i0, v_i1]) +# CHECK-NEXT: T.writes(A_pad[v_i0, v_i1]) +# CHECK-NEXT: A_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _0[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("B_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_1[v_i0, v_i1]) +# CHECK-NEXT: T.writes(B_pad[v_i0, v_i1]) +# CHECK-NEXT: B_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _1[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i, j, k in T.grid(16, 16, 16): +# CHECK-NEXT: with T.sblock("matmul_padded"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(A_pad[v_i, v_k], B_pad[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul_padded[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul_padded[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul_padded[v_i, v_j] = matmul_padded[v_i, v_j] + A_pad[v_i, v_k] * B_pad[v_k, v_j] # CHECK-NEXT: for i0, i1 in T.grid(14, 14): -# CHECK-NEXT: cse_var_3: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: C_1 = T.Buffer((196,), data=C.data) -# CHECK-NEXT: C_1[cse_var_3] = matmul_padded_1[cse_var_3] -# CHECK-NEXT: O = obj['matmul_padded'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(i, j, k) +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(matmul_padded[v_i0, v_i1]) +# CHECK-NEXT: T.writes(C[v_i0, v_i1]) +# CHECK-NEXT: C[v_i0, v_i1] = matmul_padded[v_i0, v_i1] +# CHECK-NEXT: O = sch.get_sblock("matmul_padded") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(i, j, k) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: A_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: B_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: matmul_padded = T.allocate([196], "float32", "global") -# CHECK-NEXT: A_pad_1 = T.Buffer((224,), data=A_pad) -# CHECK-NEXT: for i0, i1 in T.grid(14, 16): -# CHECK-NEXT: _0_1 = T.Buffer((196,), data=_0.data) -# CHECK-NEXT: A_pad_1[i0 * 16 + i1] = T.if_then_else(i1 < 14, _0_1[i0 * 14 + i1], T.float32(0.0)) -# CHECK-NEXT: B_pad_1 = T.Buffer((224,), data=B_pad) -# CHECK-NEXT: for i0, i1 in T.grid(16, 14): -# CHECK-NEXT: cse_var_1: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: _1_1 = T.Buffer((196,), data=_1.data) -# CHECK-NEXT: B_pad_1[cse_var_1] = T.if_then_else(i0 < 14, _1_1[cse_var_1], T.float32(0.0)) -# CHECK-NEXT: matmul_padded_1 = T.Buffer((196,), data=matmul_padded) -# CHECK-NEXT: for i, j in T.grid(14, 14): -# CHECK-NEXT: matmul_padded_1[i * 14 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(16): -# CHECK-NEXT: cse_var_2: T.int32 = i * 14 + j -# CHECK-NEXT: matmul_padded_1[cse_var_2] = matmul_padded_1[cse_var_2] + A_pad_1[i * 16 + k] * B_pad_1[k * 14 + j] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_matmul_unpad(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: A_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: B_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: matmul_padded = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("A_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_0[v_i0, v_i1]) +# CHECK-NEXT: T.writes(A_pad[v_i0, v_i1]) +# CHECK-NEXT: A_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _0[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("B_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_1[v_i0, v_i1]) +# CHECK-NEXT: T.writes(B_pad[v_i0, v_i1]) +# CHECK-NEXT: B_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _1[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i, j, k in T.grid(16, 16, 16): +# CHECK-NEXT: with T.sblock("matmul_padded"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(A_pad[v_i, v_k], B_pad[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul_padded[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul_padded[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul_padded[v_i, v_j] = matmul_padded[v_i, v_j] + A_pad[v_i, v_k] * B_pad[v_k, v_j] # CHECK-NEXT: for i0, i1 in T.grid(14, 14): -# CHECK-NEXT: cse_var_3: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: C_1 = T.Buffer((196,), data=C.data) -# CHECK-NEXT: C_1[cse_var_3] = matmul_padded_1[cse_var_3] +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(matmul_padded[v_i0, v_i1]) +# CHECK-NEXT: T.writes(C[v_i0, v_i1]) +# CHECK-NEXT: C[v_i0, v_i1] = matmul_padded[v_i0, v_i1] # CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/padding/test_pad_constant_conv2d_tvm.py b/tests/filecheck/backends/padding/test_pad_constant_conv2d_tvm.py index 97dca91c5..8929a38ae 100644 --- a/tests/filecheck/backends/padding/test_pad_constant_conv2d_tvm.py +++ b/tests/filecheck/backends/padding/test_pad_constant_conv2d_tvm.py @@ -43,50 +43,57 @@ # CHECK-NEXT: - %3: conv2d(%2, %1, stride=(2, 2)) {name = 'conv'} : [1x12x12x3xfloat32, 5x5x3x16xfloat32] -> [1x4x4x16xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(3.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] -# CHECK-NEXT: O = obj['conv'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(b, h, w, f, r, s, c) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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(3.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.sblock("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: O = sch.get_sblock("conv") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(b, h, w, f, r, s, c) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(3.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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(3.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.sblock("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: CODE: 0 diff --git a/tests/filecheck/backends/padding/test_pad_conv2d_frsc_tvm.py b/tests/filecheck/backends/padding/test_pad_conv2d_frsc_tvm.py index 300095acc..95031e040 100644 --- a/tests/filecheck/backends/padding/test_pad_conv2d_frsc_tvm.py +++ b/tests/filecheck/backends/padding/test_pad_conv2d_frsc_tvm.py @@ -45,60 +45,85 @@ # CHECK-NEXT: - %4: conv2d(%2, %3, stride=(2, 2)) {name = 'conv'} : [1x12x12x3xfloat32, 5x5x3x16xfloat32] -> [1x4x4x16xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((16, 5, 5, 3), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: _3 = T.allocate([1200], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((16, 5, 5, 3), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: _3 = T.sblock_alloc_buffer((1200,)) +# CHECK-NEXT: T_reshape = T.sblock_alloc_buffer((5, 5, 3, 16)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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 i0 in range(1200): -# CHECK-NEXT: _3_1 = T.Buffer((1200,), data=_3) -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: _3_1[i0] = _1_1[i0 % 16 * 75 + i0 // 16] -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _3_1 = T.Buffer((1200,), data=_3) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _3_1[r * 240 + s * 48 + c * 16 + f] -# CHECK-NEXT: O = obj['conv'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(b, h, w, f, r, s, c) +# CHECK-NEXT: with T.sblock("%3"): +# CHECK-NEXT: v_i0 = T.axis.spatial(1200, i0) +# CHECK-NEXT: T.reads(_1[v_i0 % 16, v_i0 % 1200 // 240, v_i0 % 240 // 48, v_i0 % 48 // 16]) +# CHECK-NEXT: T.writes(_3[v_i0]) +# CHECK-NEXT: _3[v_i0] = _1[v_i0 % 16, v_i0 % 1200 // 240, v_i0 % 240 // 48, v_i0 % 48 // 16] +# CHECK-NEXT: for ax0, ax1, ax2, ax3 in T.grid(5, 5, 3, 16): +# CHECK-NEXT: with T.sblock("T_reshape"): +# CHECK-NEXT: v_ax0, v_ax1, v_ax2, v_ax3 = T.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) +# CHECK-NEXT: T.reads(_3[(v_ax0 * 240 + v_ax1 * 48 + v_ax2 * 16 + v_ax3) % 1200]) +# 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] = _3[(v_ax0 * 240 + v_ax1 * 48 + v_ax2 * 16 + v_ax3) % 1200] +# CHECK-NEXT: for b, h, w, f, r, s, c in T.grid(1, 4, 4, 16, 5, 5, 3): +# CHECK-NEXT: with T.sblock("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], T_reshape[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] * T_reshape[v_r, v_s, v_c, v_f] +# CHECK-NEXT: O = sch.get_sblock("conv") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(b, h, w, f, r, s, c) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((16, 5, 5, 3), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: _3 = T.allocate([1200], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((16, 5, 5, 3), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: _3 = T.sblock_alloc_buffer((1200,)) +# CHECK-NEXT: T_reshape = T.sblock_alloc_buffer((5, 5, 3, 16)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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 i0 in range(1200): -# CHECK-NEXT: _3_1 = T.Buffer((1200,), data=_3) -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: _3_1[i0] = _1_1[i0 % 16 * 75 + i0 // 16] -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _3_1 = T.Buffer((1200,), data=_3) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _3_1[r * 240 + s * 48 + c * 16 + f] +# CHECK-NEXT: with T.sblock("%3"): +# CHECK-NEXT: v_i0 = T.axis.spatial(1200, i0) +# CHECK-NEXT: T.reads(_1[v_i0 % 16, v_i0 % 1200 // 240, v_i0 % 240 // 48, v_i0 % 48 // 16]) +# CHECK-NEXT: T.writes(_3[v_i0]) +# CHECK-NEXT: _3[v_i0] = _1[v_i0 % 16, v_i0 % 1200 // 240, v_i0 % 240 // 48, v_i0 % 48 // 16] +# CHECK-NEXT: for ax0, ax1, ax2, ax3 in T.grid(5, 5, 3, 16): +# CHECK-NEXT: with T.sblock("T_reshape"): +# CHECK-NEXT: v_ax0, v_ax1, v_ax2, v_ax3 = T.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) +# CHECK-NEXT: T.reads(_3[(v_ax0 * 240 + v_ax1 * 48 + v_ax2 * 16 + v_ax3) % 1200]) +# 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] = _3[(v_ax0 * 240 + v_ax1 * 48 + v_ax2 * 16 + v_ax3) % 1200] +# CHECK-NEXT: for b, h, w, f, r, s, c in T.grid(1, 4, 4, 16, 5, 5, 3): +# CHECK-NEXT: with T.sblock("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], T_reshape[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] * T_reshape[v_r, v_s, v_c, v_f] # CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/padding/test_pad_conv2d_tvm.py b/tests/filecheck/backends/padding/test_pad_conv2d_tvm.py index cf693f2dc..ccc613fa5 100644 --- a/tests/filecheck/backends/padding/test_pad_conv2d_tvm.py +++ b/tests/filecheck/backends/padding/test_pad_conv2d_tvm.py @@ -43,50 +43,57 @@ # CHECK-NEXT: - %3: conv2d(%2, %1, stride=(2, 2)) {name = 'conv'} : [1x12x12x3xfloat32, 5x5x3x16xfloat32] -> [1x4x4x16xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] -# CHECK-NEXT: O = obj['conv'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(b, h, w, f, r, s, c) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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.sblock("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: O = sch.get_sblock("conv") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(b, h, w, f, r, s, c) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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.sblock("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: CODE: 0 diff --git a/tests/filecheck/backends/padding/test_pad_matmul_unpad_tvm.py b/tests/filecheck/backends/padding/test_pad_matmul_unpad_tvm.py index 73da33cd0..0b47b6564 100644 --- a/tests/filecheck/backends/padding/test_pad_matmul_unpad_tvm.py +++ b/tests/filecheck/backends/padding/test_pad_matmul_unpad_tvm.py @@ -45,68 +45,85 @@ # CHECK-NEXT: - %5: unpad(%4, padding={-2: (0, 2), -1: (0, 2)}) {name = 'C'} : [16x16xfloat32] -> [14x14xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: A_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: B_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: matmul_padded = T.allocate([196], "float32", "global") -# CHECK-NEXT: A_pad_1 = T.Buffer((224,), data=A_pad) -# CHECK-NEXT: for i0, i1 in T.grid(14, 16): -# CHECK-NEXT: _0_1 = T.Buffer((196,), data=_0.data) -# CHECK-NEXT: A_pad_1[i0 * 16 + i1] = T.if_then_else(i1 < 14, _0_1[i0 * 14 + i1], T.float32(0.0)) -# CHECK-NEXT: B_pad_1 = T.Buffer((224,), data=B_pad) -# CHECK-NEXT: for i0, i1 in T.grid(16, 14): -# CHECK-NEXT: cse_var_1: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: _1_1 = T.Buffer((196,), data=_1.data) -# CHECK-NEXT: B_pad_1[cse_var_1] = T.if_then_else(i0 < 14, _1_1[cse_var_1], T.float32(0.0)) -# CHECK-NEXT: matmul_padded_1 = T.Buffer((196,), data=matmul_padded) -# CHECK-NEXT: for i, j in T.grid(14, 14): -# CHECK-NEXT: matmul_padded_1[i * 14 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(16): -# CHECK-NEXT: cse_var_2: T.int32 = i * 14 + j -# CHECK-NEXT: matmul_padded_1[cse_var_2] = matmul_padded_1[cse_var_2] + A_pad_1[i * 16 + k] * B_pad_1[k * 14 + j] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_matmul_unpad(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: A_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: B_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: matmul_padded = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("A_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_0[v_i0, v_i1]) +# CHECK-NEXT: T.writes(A_pad[v_i0, v_i1]) +# CHECK-NEXT: A_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _0[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("B_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_1[v_i0, v_i1]) +# CHECK-NEXT: T.writes(B_pad[v_i0, v_i1]) +# CHECK-NEXT: B_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _1[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i, j, k in T.grid(16, 16, 16): +# CHECK-NEXT: with T.sblock("matmul_padded"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(A_pad[v_i, v_k], B_pad[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul_padded[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul_padded[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul_padded[v_i, v_j] = matmul_padded[v_i, v_j] + A_pad[v_i, v_k] * B_pad[v_k, v_j] # CHECK-NEXT: for i0, i1 in T.grid(14, 14): -# CHECK-NEXT: cse_var_3: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: C_1 = T.Buffer((196,), data=C.data) -# CHECK-NEXT: C_1[cse_var_3] = matmul_padded_1[cse_var_3] -# CHECK-NEXT: O = obj['matmul_padded'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(i, j, k) +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(matmul_padded[v_i0, v_i1]) +# CHECK-NEXT: T.writes(C[v_i0, v_i1]) +# CHECK-NEXT: C[v_i0, v_i1] = matmul_padded[v_i0, v_i1] +# CHECK-NEXT: O = sch.get_sblock("matmul_padded") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(i, j, k) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: A_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: B_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: matmul_padded = T.allocate([196], "float32", "global") -# CHECK-NEXT: A_pad_1 = T.Buffer((224,), data=A_pad) -# CHECK-NEXT: for i0, i1 in T.grid(14, 16): -# CHECK-NEXT: _0_1 = T.Buffer((196,), data=_0.data) -# CHECK-NEXT: A_pad_1[i0 * 16 + i1] = T.if_then_else(i1 < 14, _0_1[i0 * 14 + i1], T.float32(0.0)) -# CHECK-NEXT: B_pad_1 = T.Buffer((224,), data=B_pad) -# CHECK-NEXT: for i0, i1 in T.grid(16, 14): -# CHECK-NEXT: cse_var_1: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: _1_1 = T.Buffer((196,), data=_1.data) -# CHECK-NEXT: B_pad_1[cse_var_1] = T.if_then_else(i0 < 14, _1_1[cse_var_1], T.float32(0.0)) -# CHECK-NEXT: matmul_padded_1 = T.Buffer((196,), data=matmul_padded) -# CHECK-NEXT: for i, j in T.grid(14, 14): -# CHECK-NEXT: matmul_padded_1[i * 14 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(16): -# CHECK-NEXT: cse_var_2: T.int32 = i * 14 + j -# CHECK-NEXT: matmul_padded_1[cse_var_2] = matmul_padded_1[cse_var_2] + A_pad_1[i * 16 + k] * B_pad_1[k * 14 + j] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_matmul_unpad(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: A_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: B_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: matmul_padded = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("A_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_0[v_i0, v_i1]) +# CHECK-NEXT: T.writes(A_pad[v_i0, v_i1]) +# CHECK-NEXT: A_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _0[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("B_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_1[v_i0, v_i1]) +# CHECK-NEXT: T.writes(B_pad[v_i0, v_i1]) +# CHECK-NEXT: B_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _1[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i, j, k in T.grid(16, 16, 16): +# CHECK-NEXT: with T.sblock("matmul_padded"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(A_pad[v_i, v_k], B_pad[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul_padded[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul_padded[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul_padded[v_i, v_j] = matmul_padded[v_i, v_j] + A_pad[v_i, v_k] * B_pad[v_k, v_j] # CHECK-NEXT: for i0, i1 in T.grid(14, 14): -# CHECK-NEXT: cse_var_3: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: C_1 = T.Buffer((196,), data=C.data) -# CHECK-NEXT: C_1[cse_var_3] = matmul_padded_1[cse_var_3] +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(matmul_padded[v_i0, v_i1]) +# CHECK-NEXT: T.writes(C[v_i0, v_i1]) +# CHECK-NEXT: C[v_i0, v_i1] = matmul_padded[v_i0, v_i1] # CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/padding/test_pad_tuple_matmul_unpad_tvm.py b/tests/filecheck/backends/padding/test_pad_tuple_matmul_unpad_tvm.py index c02f44373..2a85eee02 100644 --- a/tests/filecheck/backends/padding/test_pad_tuple_matmul_unpad_tvm.py +++ b/tests/filecheck/backends/padding/test_pad_tuple_matmul_unpad_tvm.py @@ -45,68 +45,85 @@ # CHECK-NEXT: - %5: unpad(%4, padding={-2: (0, 2), -1: (0, 2)}) {name = 'C'} : [16x16xfloat32] -> [14x14xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: A_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: B_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: matmul_padded = T.allocate([196], "float32", "global") -# CHECK-NEXT: A_pad_1 = T.Buffer((224,), data=A_pad) -# CHECK-NEXT: for i0, i1 in T.grid(14, 16): -# CHECK-NEXT: _0_1 = T.Buffer((196,), data=_0.data) -# CHECK-NEXT: A_pad_1[i0 * 16 + i1] = T.if_then_else(i1 < 14, _0_1[i0 * 14 + i1], T.float32(0.0)) -# CHECK-NEXT: B_pad_1 = T.Buffer((224,), data=B_pad) -# CHECK-NEXT: for i0, i1 in T.grid(16, 14): -# CHECK-NEXT: cse_var_1: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: _1_1 = T.Buffer((196,), data=_1.data) -# CHECK-NEXT: B_pad_1[cse_var_1] = T.if_then_else(i0 < 14, _1_1[cse_var_1], T.float32(0.0)) -# CHECK-NEXT: matmul_padded_1 = T.Buffer((196,), data=matmul_padded) -# CHECK-NEXT: for i, j in T.grid(14, 14): -# CHECK-NEXT: matmul_padded_1[i * 14 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(16): -# CHECK-NEXT: cse_var_2: T.int32 = i * 14 + j -# CHECK-NEXT: matmul_padded_1[cse_var_2] = matmul_padded_1[cse_var_2] + A_pad_1[i * 16 + k] * B_pad_1[k * 14 + j] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_matmul_unpad(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: A_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: B_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: matmul_padded = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("A_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_0[v_i0, v_i1]) +# CHECK-NEXT: T.writes(A_pad[v_i0, v_i1]) +# CHECK-NEXT: A_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _0[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("B_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_1[v_i0, v_i1]) +# CHECK-NEXT: T.writes(B_pad[v_i0, v_i1]) +# CHECK-NEXT: B_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _1[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i, j, k in T.grid(16, 16, 16): +# CHECK-NEXT: with T.sblock("matmul_padded"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(A_pad[v_i, v_k], B_pad[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul_padded[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul_padded[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul_padded[v_i, v_j] = matmul_padded[v_i, v_j] + A_pad[v_i, v_k] * B_pad[v_k, v_j] # CHECK-NEXT: for i0, i1 in T.grid(14, 14): -# CHECK-NEXT: cse_var_3: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: C_1 = T.Buffer((196,), data=C.data) -# CHECK-NEXT: C_1[cse_var_3] = matmul_padded_1[cse_var_3] -# CHECK-NEXT: O = obj['matmul_padded'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(i, j, k) +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(matmul_padded[v_i0, v_i1]) +# CHECK-NEXT: T.writes(C[v_i0, v_i1]) +# CHECK-NEXT: C[v_i0, v_i1] = matmul_padded[v_i0, v_i1] +# CHECK-NEXT: O = sch.get_sblock("matmul_padded") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(i, j, k) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: A_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: B_pad = T.allocate([224], "float32", "global") -# CHECK-NEXT: matmul_padded = T.allocate([196], "float32", "global") -# CHECK-NEXT: A_pad_1 = T.Buffer((224,), data=A_pad) -# CHECK-NEXT: for i0, i1 in T.grid(14, 16): -# CHECK-NEXT: _0_1 = T.Buffer((196,), data=_0.data) -# CHECK-NEXT: A_pad_1[i0 * 16 + i1] = T.if_then_else(i1 < 14, _0_1[i0 * 14 + i1], T.float32(0.0)) -# CHECK-NEXT: B_pad_1 = T.Buffer((224,), data=B_pad) -# CHECK-NEXT: for i0, i1 in T.grid(16, 14): -# CHECK-NEXT: cse_var_1: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: _1_1 = T.Buffer((196,), data=_1.data) -# CHECK-NEXT: B_pad_1[cse_var_1] = T.if_then_else(i0 < 14, _1_1[cse_var_1], T.float32(0.0)) -# CHECK-NEXT: matmul_padded_1 = T.Buffer((196,), data=matmul_padded) -# CHECK-NEXT: for i, j in T.grid(14, 14): -# CHECK-NEXT: matmul_padded_1[i * 14 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(16): -# CHECK-NEXT: cse_var_2: T.int32 = i * 14 + j -# CHECK-NEXT: matmul_padded_1[cse_var_2] = matmul_padded_1[cse_var_2] + A_pad_1[i * 16 + k] * B_pad_1[k * 14 + j] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_matmul_unpad(_0: T.Buffer((14, 14), "float32"), _1: T.Buffer((14, 14), "float32"), C: T.Buffer((14, 14), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: A_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: B_pad = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: matmul_padded = T.sblock_alloc_buffer((16, 16)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("A_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_0[v_i0, v_i1]) +# CHECK-NEXT: T.writes(A_pad[v_i0, v_i1]) +# CHECK-NEXT: A_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _0[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i0, i1 in T.grid(16, 16): +# CHECK-NEXT: with T.sblock("B_pad"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(_1[v_i0, v_i1]) +# CHECK-NEXT: T.writes(B_pad[v_i0, v_i1]) +# CHECK-NEXT: B_pad[v_i0, v_i1] = T.if_then_else(0 <= v_i0 and v_i0 < 14 and 0 <= v_i1 and v_i1 < 14, _1[v_i0, v_i1], T.float32(0.0)) +# CHECK-NEXT: for i, j, k in T.grid(16, 16, 16): +# CHECK-NEXT: with T.sblock("matmul_padded"): +# CHECK-NEXT: v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) +# CHECK-NEXT: T.reads(A_pad[v_i, v_k], B_pad[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul_padded[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul_padded[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul_padded[v_i, v_j] = matmul_padded[v_i, v_j] + A_pad[v_i, v_k] * B_pad[v_k, v_j] # CHECK-NEXT: for i0, i1 in T.grid(14, 14): -# CHECK-NEXT: cse_var_3: T.int32 = i0 * 14 + i1 -# CHECK-NEXT: C_1 = T.Buffer((196,), data=C.data) -# CHECK-NEXT: C_1[cse_var_3] = matmul_padded_1[cse_var_3] +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) +# CHECK-NEXT: T.reads(matmul_padded[v_i0, v_i1]) +# CHECK-NEXT: T.writes(C[v_i0, v_i1]) +# CHECK-NEXT: C[v_i0, v_i1] = matmul_padded[v_i0, v_i1] # CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/test_conv2d_mini_tvm.py b/tests/filecheck/backends/test_conv2d_mini_tvm.py index 818006857..7ce1bc606 100644 --- a/tests/filecheck/backends/test_conv2d_mini_tvm.py +++ b/tests/filecheck/backends/test_conv2d_mini_tvm.py @@ -41,40 +41,43 @@ # CHECK-NEXT: - %2: conv2d(%0, %1, stride=(1, 1)) {name = 'O'} : [1x10x10x3xfloat32, 3x3x3x16xfloat32] -> [1x8x8x16xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 10, 10, 3), "float32"), _1: T.Buffer((3, 3, 3, 16), "float32"), O: T.Buffer((1, 8, 8, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for h, w, f in T.grid(8, 8, 16): -# CHECK-NEXT: O_1 = T.Buffer((1024,), data=O.data) -# CHECK-NEXT: O_1[h * 128 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(3, 3, 3): -# CHECK-NEXT: cse_var_1: T.int32 = h * 128 + w * 16 + f -# CHECK-NEXT: _0_1 = T.Buffer((300,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((432,), data=_1.data) -# CHECK-NEXT: O_1[cse_var_1] = O_1[cse_var_1] + _0_1[h * 30 + r * 30 + w * 3 + s * 3 + c] * _1_1[r * 144 + s * 48 + c * 16 + f] -# CHECK-NEXT: O = obj['O'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(b, h, w, f, r, s, c) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def conv2d_nhwc_mini(_0: T.Buffer((1, 10, 10, 3), "float32"), _1: T.Buffer((3, 3, 3, 16), "float32"), O: T.Buffer((1, 8, 8, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for b, h, w, f, r, s, c in T.grid(1, 8, 8, 16, 3, 3, 3): +# CHECK-NEXT: with T.sblock("O"): +# 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(_0[v_b, v_h + v_r, v_w + v_s, v_c], _1[v_r, v_s, v_c, v_f]) +# CHECK-NEXT: T.writes(O[v_b, v_h, v_w, v_f]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: O[v_b, v_h, v_w, v_f] = T.float32(0.0) +# CHECK-NEXT: O[v_b, v_h, v_w, v_f] = O[v_b, v_h, v_w, v_f] + _0[v_b, v_h + v_r, v_w + v_s, v_c] * _1[v_r, v_s, v_c, v_f] +# CHECK-NEXT: O = sch.get_sblock("O") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: sch.reorder(b, h, w, f, r, s, c) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 10, 10, 3), "float32"), _1: T.Buffer((3, 3, 3, 16), "float32"), O: T.Buffer((1, 8, 8, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for h, w, f in T.grid(8, 8, 16): -# CHECK-NEXT: O_1 = T.Buffer((1024,), data=O.data) -# CHECK-NEXT: O_1[h * 128 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(3, 3, 3): -# CHECK-NEXT: cse_var_1: T.int32 = h * 128 + w * 16 + f -# CHECK-NEXT: _0_1 = T.Buffer((300,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((432,), data=_1.data) -# CHECK-NEXT: O_1[cse_var_1] = O_1[cse_var_1] + _0_1[h * 30 + r * 30 + w * 3 + s * 3 + c] * _1_1[r * 144 + s * 48 + c * 16 + f] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def conv2d_nhwc_mini(_0: T.Buffer((1, 10, 10, 3), "float32"), _1: T.Buffer((3, 3, 3, 16), "float32"), O: T.Buffer((1, 8, 8, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for b, h, w, f, r, s, c in T.grid(1, 8, 8, 16, 3, 3, 3): +# CHECK-NEXT: with T.sblock("O"): +# 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(_0[v_b, v_h + v_r, v_w + v_s, v_c], _1[v_r, v_s, v_c, v_f]) +# CHECK-NEXT: T.writes(O[v_b, v_h, v_w, v_f]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: O[v_b, v_h, v_w, v_f] = T.float32(0.0) +# CHECK-NEXT: O[v_b, v_h, v_w, v_f] = O[v_b, v_h, v_w, v_f] + _0[v_b, v_h + v_r, v_w + v_s, v_c] * _1[v_r, v_s, v_c, v_f] # CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/test_conv2d_r181_tvm.py b/tests/filecheck/backends/test_conv2d_r181_tvm.py index b48de35f8..bba89d2ca 100644 --- a/tests/filecheck/backends/test_conv2d_r181_tvm.py +++ b/tests/filecheck/backends/test_conv2d_r181_tvm.py @@ -50,67 +50,56 @@ # CHECK-NEXT: - %2: conv2d(%0, %1, stride=(2, 2)) {name = 'O'} : [1x230x230x3xfloat32, 7x7x3x64xfloat32] -> [1x112x112x64xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 230, 230, 3), "float32"), _1: T.Buffer((7, 7, 3, 64), "float32"), O: T.Buffer((1, 112, 112, 64), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for h, w, f in T.grid(112, 112, 64): -# CHECK-NEXT: O_1 = T.Buffer((802816,), data=O.data) -# CHECK-NEXT: O_1[h * 7168 + w * 64 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(7, 7, 3): -# CHECK-NEXT: cse_var_1: T.int32 = h * 7168 + w * 64 + f -# CHECK-NEXT: _0_1 = T.Buffer((158700,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((9408,), data=_1.data) -# CHECK-NEXT: O_1[cse_var_1] = O_1[cse_var_1] + _0_1[h * 1380 + r * 690 + w * 6 + s * 3 + c] * _1_1[r * 1344 + s * 192 + c * 64 + f] -# CHECK-NEXT: O = obj['O'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: c, __u_c = sch[O].split(c, factor=3) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=4) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=16) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, r, s, c, __u_c, w1, f1) -# CHECK-NEXT: sch[O].unroll(w1) -# CHECK-NEXT: sch[O].unroll(__u_c) -# CHECK-NEXT: sch[O].vectorize(f1) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def conv2d_nhwc_r181(_0: T.Buffer((1, 230, 230, 3), "float32"), _1: T.Buffer((7, 7, 3, 64), "float32"), O: T.Buffer((1, 112, 112, 64), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for b, h, w, f, r, s, c in T.grid(1, 112, 112, 64, 7, 7, 3): +# CHECK-NEXT: with T.sblock("O"): +# 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(_0[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(O[v_b, v_h, v_w, v_f]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: O[v_b, v_h, v_w, v_f] = T.float32(0.0) +# CHECK-NEXT: O[v_b, v_h, v_w, v_f] = O[v_b, v_h, v_w, v_f] + _0[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: O = sch.get_sblock("O") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: w, w1, = sch.split(w, factors=[None, 4]) +# CHECK-NEXT: f, f1, = sch.split(f, factors=[None, 16]) +# CHECK-NEXT: c, __u_c, = sch.split(c, factors=[None, 3]) +# CHECK-NEXT: sch.reorder(b, h, w, f, r, s, c, __u_c, w1, f1) +# CHECK-NEXT: sch.unroll(w1) +# CHECK-NEXT: sch.unroll(__u_c) +# CHECK-NEXT: sch.vectorize(f1) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 230, 230, 3), "float32"), _1: T.Buffer((7, 7, 3, 64), "float32"), O: T.Buffer((1, 112, 112, 64), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for h, w_outer, f_outer in T.grid(112, 28, 4): -# CHECK-NEXT: cse_var_1: T.int32 = h * 7168 + w_outer * 256 + f_outer * 16 -# CHECK-NEXT: O_1 = T.Buffer((802816,), data=O.data) -# CHECK-NEXT: O_1[cse_var_1:cse_var_1 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: O_1[cse_var_1 + 64:cse_var_1 + 64 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: O_1[cse_var_1 + 128:cse_var_1 + 128 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: O_1[cse_var_1 + 192:cse_var_1 + 192 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: for r, s in T.grid(7, 7): -# CHECK-NEXT: cse_var_8: T.int32 = cse_var_1 + 64 -# CHECK-NEXT: cse_var_7: T.int32 = cse_var_1 + 192 -# CHECK-NEXT: cse_var_6: T.int32 = cse_var_1 + 128 -# CHECK-NEXT: cse_var_5: T.int32 = r * 1344 + s * 192 + f_outer * 16 -# CHECK-NEXT: cse_var_4: T.int32 = cse_var_5 + 64 -# CHECK-NEXT: cse_var_3: T.int32 = cse_var_5 + 128 -# CHECK-NEXT: cse_var_2: T.int32 = h * 1380 + r * 690 + w_outer * 24 + s * 3 -# CHECK-NEXT: _0_1 = T.Buffer((158700,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((9408,), data=_1.data) -# CHECK-NEXT: O_1[cse_var_1:cse_var_1 + 16] = O_1[cse_var_1:cse_var_1 + 16] + T.Broadcast(_0_1[cse_var_2], 16) * _1_1[cse_var_5:cse_var_5 + 16] -# CHECK-NEXT: O_1[cse_var_8:cse_var_8 + 16] = O_1[cse_var_8:cse_var_8 + 16] + T.Broadcast(_0_1[cse_var_2 + 6], 16) * _1_1[cse_var_5:cse_var_5 + 16] -# CHECK-NEXT: O_1[cse_var_6:cse_var_6 + 16] = O_1[cse_var_6:cse_var_6 + 16] + T.Broadcast(_0_1[cse_var_2 + 12], 16) * _1_1[cse_var_5:cse_var_5 + 16] -# CHECK-NEXT: O_1[cse_var_7:cse_var_7 + 16] = O_1[cse_var_7:cse_var_7 + 16] + T.Broadcast(_0_1[cse_var_2 + 18], 16) * _1_1[cse_var_5:cse_var_5 + 16] -# CHECK-NEXT: O_1[cse_var_1:cse_var_1 + 16] = O_1[cse_var_1:cse_var_1 + 16] + T.Broadcast(_0_1[cse_var_2 + 1], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: O_1[cse_var_8:cse_var_8 + 16] = O_1[cse_var_8:cse_var_8 + 16] + T.Broadcast(_0_1[cse_var_2 + 7], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: O_1[cse_var_6:cse_var_6 + 16] = O_1[cse_var_6:cse_var_6 + 16] + T.Broadcast(_0_1[cse_var_2 + 13], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: O_1[cse_var_7:cse_var_7 + 16] = O_1[cse_var_7:cse_var_7 + 16] + T.Broadcast(_0_1[cse_var_2 + 19], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: O_1[cse_var_1:cse_var_1 + 16] = O_1[cse_var_1:cse_var_1 + 16] + T.Broadcast(_0_1[cse_var_2 + 2], 16) * _1_1[cse_var_3:cse_var_3 + 16] -# CHECK-NEXT: O_1[cse_var_8:cse_var_8 + 16] = O_1[cse_var_8:cse_var_8 + 16] + T.Broadcast(_0_1[cse_var_2 + 8], 16) * _1_1[cse_var_3:cse_var_3 + 16] -# CHECK-NEXT: O_1[cse_var_6:cse_var_6 + 16] = O_1[cse_var_6:cse_var_6 + 16] + T.Broadcast(_0_1[cse_var_2 + 14], 16) * _1_1[cse_var_3:cse_var_3 + 16] -# CHECK-NEXT: O_1[cse_var_7:cse_var_7 + 16] = O_1[cse_var_7:cse_var_7 + 16] + T.Broadcast(_0_1[cse_var_2 + 20], 16) * _1_1[cse_var_3:cse_var_3 + 16] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def conv2d_nhwc_r181(_0: T.Buffer((1, 230, 230, 3), "float32"), _1: T.Buffer((7, 7, 3, 64), "float32"), O: T.Buffer((1, 112, 112, 64), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for b, h, w_0, f_0, r, s, c_0 in T.grid(1, 112, 28, 4, 7, 7, 1): +# CHECK-NEXT: for c_1 in T.unroll(3): +# CHECK-NEXT: for w_1 in T.unroll(4): +# CHECK-NEXT: for f_1 in T.vectorized(16): +# CHECK-NEXT: with T.sblock("O"): +# CHECK-NEXT: v_b, v_h = T.axis.remap("SS", [b, h]) +# CHECK-NEXT: v_w = T.axis.spatial(112, w_0 * 4 + w_1) +# CHECK-NEXT: v_f = T.axis.spatial(64, f_0 * 16 + f_1) +# CHECK-NEXT: v_r, v_s = T.axis.remap("RR", [r, s]) +# CHECK-NEXT: v_c = T.axis.reduce(3, c_0 * 3 + c_1) +# CHECK-NEXT: T.reads(_0[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(O[v_b, v_h, v_w, v_f]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: O[v_b, v_h, v_w, v_f] = T.float32(0.0) +# CHECK-NEXT: O[v_b, v_h, v_w, v_f] = O[v_b, v_h, v_w, v_f] + _0[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: CODE: 0 diff --git a/tests/filecheck/backends/test_matmul_ndiv_tvm.py b/tests/filecheck/backends/test_matmul_ndiv_tvm.py index aa0e5c22d..5aea2c6cf 100644 --- a/tests/filecheck/backends/test_matmul_ndiv_tvm.py +++ b/tests/filecheck/backends/test_matmul_ndiv_tvm.py @@ -46,56 +46,51 @@ # CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [4x512xfloat32, 512x32xfloat32] -> [4x32xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for i, j in T.grid(4, 32): -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: C_1[i * 32 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(512): -# CHECK-NEXT: cse_var_1: T.int32 = i * 32 + j -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_1] = C_1[cse_var_1] + _0_1[i * 512 + k] * _1_1[k * 32 + j] -# CHECK-NEXT: O = obj['C'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=3) -# CHECK-NEXT: i1, __u_i1 = sch[O].split(i1, factor=2) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: sch[O].reorder(k, i, j, i1, __u_i1, j1) -# CHECK-NEXT: sch[O].unroll(__u_i1) -# CHECK-NEXT: sch[O].vectorize(j1) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j, k in T.grid(4, 32, 512): +# CHECK-NEXT: with T.sblock("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_sblock("C") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, __u_i1, = sch.split(i, factors=[None, 1, 2]) +# CHECK-NEXT: j, j1, = sch.split(j, factors=[None, 16]) +# CHECK-NEXT: sch.reorder(k, i, j, i1, __u_i1, j1) +# CHECK-NEXT: sch.unroll(__u_i1) +# CHECK-NEXT: sch.vectorize(j1) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: for i_outer_init, j_outer_init, i_inner_outer_init in T.grid(2, 2, 2): -# CHECK-NEXT: if T.likely(i_outer_init * 3 + i_inner_outer_init * 2 < 4): -# CHECK-NEXT: C_1[i_outer_init * 96 + i_inner_outer_init * 64 + j_outer_init * 16:i_outer_init * 96 + i_inner_outer_init * 64 + j_outer_init * 16 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: if T.likely(i_outer_init * 3 + i_inner_outer_init * 2 < 3): -# CHECK-NEXT: if T.likely(i_inner_outer_init < 1): -# CHECK-NEXT: C_1[i_outer_init * 96 + i_inner_outer_init * 64 + j_outer_init * 16 + 32:i_outer_init * 96 + i_inner_outer_init * 64 + j_outer_init * 16 + 32 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: for k, i_outer, j_outer, i_inner_outer in T.grid(512, 2, 2, 2): -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: if T.likely(i_outer * 3 + i_inner_outer * 2 < 4): -# CHECK-NEXT: cse_var_2: T.int32 = j_outer * 16 -# CHECK-NEXT: cse_var_1: T.int32 = i_outer * 96 + i_inner_outer * 64 + cse_var_2 -# CHECK-NEXT: C_1[cse_var_1:cse_var_1 + 16] = C_1[cse_var_1:cse_var_1 + 16] + T.Broadcast(_0_1[i_outer * 1536 + i_inner_outer * 1024 + k], 16) * _1_1[k * 32 + cse_var_2:k * 32 + cse_var_2 + 16] -# CHECK-NEXT: if T.likely(i_outer * 3 + i_inner_outer * 2 < 3): -# CHECK-NEXT: if T.likely(i_inner_outer < 1): -# CHECK-NEXT: cse_var_4: T.int32 = j_outer * 16 -# CHECK-NEXT: cse_var_3: T.int32 = i_outer * 96 + i_inner_outer * 64 + cse_var_4 + 32 -# CHECK-NEXT: C_1[cse_var_3:cse_var_3 + 16] = C_1[cse_var_3:cse_var_3 + 16] + T.Broadcast(_0_1[i_outer * 1536 + i_inner_outer * 1024 + k + 512], 16) * _1_1[k * 32 + cse_var_4:k * 32 + cse_var_4 + 16] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for k, i_0, j_0, i_1 in T.grid(512, 2, 2, 1): +# CHECK-NEXT: for i_2 in T.unroll(2): +# CHECK-NEXT: for j_1 in T.vectorized(16): +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i = T.axis.spatial(4, i_0 * 2 + i_1 * 2 + i_2) +# CHECK-NEXT: v_j = T.axis.spatial(32, j_0 * 16 + j_1) +# CHECK-NEXT: v_k = T.axis.reduce(512, 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: CODE: 0 diff --git a/tests/filecheck/backends/test_matmul_tvm_tir.py b/tests/filecheck/backends/test_matmul_pack_parallel_tvm.py similarity index 82% rename from tests/filecheck/backends/test_matmul_tvm_tir.py rename to tests/filecheck/backends/test_matmul_pack_parallel_tvm.py index 401806b66..96111112f 100644 --- a/tests/filecheck/backends/test_matmul_tvm_tir.py +++ b/tests/filecheck/backends/test_matmul_pack_parallel_tvm.py @@ -14,7 +14,7 @@ graph = gb.graph print(graph) -impl = Backend(graph, tir_schedule=True) +impl = Backend(graph) sch = impl.get_scheduler() sch.tile("i", {"i1": 8, "i2": 4}) @@ -30,10 +30,12 @@ comp = impl.get_compiler( shared_lib=True, - dump_file="matmul_tvm_tir", + dump_file="matmul_pack_parallel_tvm", print_source_ir=True, print_transformed_ir=True, + bare_ptr=True, ) + module = comp.compile(sched) executor = module.get_executor(validate=True) res = executor.execute() @@ -50,23 +52,24 @@ # 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: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: @T.prim_func(s_tir=True) # 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: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): # CHECK-NEXT: for i, j, k in T.grid(64, 192, 256): -# CHECK-NEXT: with T.block("C"): +# CHECK-NEXT: with T.sblock("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: O = sch.get_sblock("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") @@ -84,31 +87,32 @@ # 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: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: @T.prim_func(s_tir=True) # 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: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: _1_global = T.sblock_alloc_buffer((256, 192)) +# CHECK-NEXT: C_global = T.sblock_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: with T.sblock("_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: T.sblock_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: with T.sblock("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) @@ -119,7 +123,7 @@ # 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: with T.sblock("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]) diff --git a/tests/filecheck/backends/test_matmul_pack_tvm.py b/tests/filecheck/backends/test_matmul_pack_tvm.py index 0a1f78ed0..4a02b162f 100644 --- a/tests/filecheck/backends/test_matmul_pack_tvm.py +++ b/tests/filecheck/backends/test_matmul_pack_tvm.py @@ -49,82 +49,76 @@ # CHECK-NEXT: - %2: matmul(%0, %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: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # 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: 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: _0_1 = T.Buffer((4096,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((4096,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_1] = C_1[cse_var_1] + _0_1[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, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: j, j_ = sch[O].split(j, factor=32) -# CHECK-NEXT: i_ = i -# CHECK-NEXT: sch[O].reorder(j, j_, i_) -# 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: j1 = j -# CHECK-NEXT: i, i1 = sch[O_W0].split(i, factor=8) -# 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, i, j1, i1, 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[O_W0].unroll(i2) -# CHECK-NEXT: sch[O_W0].vectorize(j2) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# 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({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j, k in T.grid(64, 64, 64): +# CHECK-NEXT: with T.sblock("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_sblock("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, = 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(j, k, i, j1, i1, 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.unroll(i2) +# CHECK-NEXT: sch.vectorize(j2) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # 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: C_global = T.allocate([2048], "float32", "global") -# CHECK-NEXT: _1_global = T.allocate([16640], "float32", "global") -# CHECK-NEXT: for j_outer in range(2): -# CHECK-NEXT: C_global_1 = T.Buffer((2048,), data=C_global) -# CHECK-NEXT: for i_c_outer_init, j_c_outer_init, i_c_inner_outer_init in T.grid(8, 2, 2): -# CHECK-NEXT: cse_var_1: T.int32 = i_c_outer_init * 256 + i_c_inner_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: _1_global_1 = T.Buffer((16640,), data=_1_global) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# 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({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: _1_global = T.sblock_alloc_buffer((64, 64)) +# CHECK-NEXT: C_global = T.sblock_alloc_buffer((64, 64)) +# CHECK-NEXT: for j_0 in range(2): +# CHECK-NEXT: for k_0 in range(4): # 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 + j_outer * 32 + ax1] -# CHECK-NEXT: for i_c_outer, j_c_outer, i_c_inner_outer, k_inner in T.grid(8, 2, 2, 16): -# CHECK-NEXT: cse_var_8: T.int32 = j_c_outer * 16 -# CHECK-NEXT: cse_var_7: T.int32 = k_inner * 1040 + cse_var_8 -# CHECK-NEXT: cse_var_6: T.int32 = i_c_outer * 256 + i_c_inner_outer * 128 + cse_var_8 -# CHECK-NEXT: cse_var_5: T.int32 = i_c_outer * 512 + i_c_inner_outer * 256 + k_outer * 16 + k_inner -# CHECK-NEXT: cse_var_4: T.int32 = cse_var_6 + 96 -# CHECK-NEXT: cse_var_3: T.int32 = cse_var_6 + 64 -# CHECK-NEXT: cse_var_2: T.int32 = cse_var_6 + 32 -# CHECK-NEXT: _0_1 = T.Buffer((4096,), data=_0.data) -# CHECK-NEXT: C_global_1[cse_var_6:cse_var_6 + 16] = C_global_1[cse_var_6:cse_var_6 + 16] + T.Broadcast(_0_1[cse_var_5], 16) * _1_global_1[cse_var_7:cse_var_7 + 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(_0_1[cse_var_5 + 64], 16) * _1_global_1[cse_var_7:cse_var_7 + 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(_0_1[cse_var_5 + 128], 16) * _1_global_1[cse_var_7:cse_var_7 + 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(_0_1[cse_var_5 + 192], 16) * _1_global_1[cse_var_7:cse_var_7 + 16] -# CHECK-NEXT: for j_inner, i in T.grid(32, 64): -# CHECK-NEXT: C_1 = T.Buffer((4096,), data=C.data) -# CHECK-NEXT: C_1[i * 64 + j_outer * 32 + j_inner] = C_global_1[i * 32 + j_inner] +# CHECK-NEXT: with T.sblock("_1_global"): +# CHECK-NEXT: v0 = T.axis.spatial(64, k_0 * 16 + ax0) +# CHECK-NEXT: v1 = T.axis.spatial(64, j_0 * 32 + ax1) +# CHECK-NEXT: T.reads(_1[v0, v1]) +# CHECK-NEXT: T.writes(_1_global[v0, v1]) +# CHECK-NEXT: T.sblock_attr({"buffer_dim_align": [[0, 0, 1024, 16]]}) +# CHECK-NEXT: _1_global[v0, v1] = _1[v0, v1] +# CHECK-NEXT: for i_0, j_1, i_1, k_1 in T.grid(8, 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.sblock("C"): +# CHECK-NEXT: v_i = T.axis.spatial(64, i_0 * 8 + i_1 * 4 + i_2) +# CHECK-NEXT: v_j = T.axis.spatial(64, j_0 * 32 + j_1 * 16 + j_2) +# CHECK-NEXT: v_k = T.axis.reduce(64, k_0 * 16 + k_1) +# 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(64, 32): +# CHECK-NEXT: with T.sblock("C_global"): +# CHECK-NEXT: v0 = T.axis.spatial(64, ax0) +# CHECK-NEXT: v1 = T.axis.spatial(64, j_0 * 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/backends/test_matmul_relu_fused_tvm.py b/tests/filecheck/backends/test_matmul_relu_fused_tvm.py new file mode 100644 index 000000000..bd7da771a --- /dev/null +++ b/tests/filecheck/backends/test_matmul_relu_fused_tvm.py @@ -0,0 +1,137 @@ +# 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) + +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", + 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_relu +# CHECK-NEXT: inputs: +# CHECK-NEXT: - %0 : 4x512xfloat32 +# CHECK-NEXT: - %1 : 512x32xfloat32 +# CHECK-NEXT: outputs: +# CHECK-NEXT: - %3 : 4x32xfloat32 +# CHECK-NEXT: nodes: +# CHECK-NEXT: - %2: matmul(%0, %1) {name = 'matmul'} : [4x512xfloat32, 512x32xfloat32] -> [4x32xfloat32] +# CHECK-NEXT: - %3: relu(%2) {name = 'relu'} : [4x32xfloat32] -> [4x32xfloat32] +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul_relu(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: matmul = T.sblock_alloc_buffer((4, 32)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: relu = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: for i, j, k in T.grid(4, 32, 512): +# CHECK-NEXT: with T.sblock("matmul"): +# 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(matmul[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul[v_i, v_j] = matmul[v_i, v_j] + _0[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: for ax0 in range(128): +# CHECK-NEXT: with T.sblock("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(128, ax0) +# CHECK-NEXT: T.reads(matmul[v_ax0 % 128 // 32, v_ax0 % 32]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0]) +# CHECK-NEXT: T_reshape_1[v_ax0] = matmul[v_ax0 % 128 // 32, v_ax0 % 32] +# CHECK-NEXT: for i in range(128): +# CHECK-NEXT: with T.sblock("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(128, 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 in T.grid(4, 32): +# CHECK-NEXT: with T.sblock("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) +# CHECK-NEXT: T.reads(relu[(v_ax0 * 32 + v_ax1) % 128]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0, v_ax1]) +# CHECK-NEXT: T_reshape[v_ax0, v_ax1] = relu[(v_ax0 * 32 + v_ax1) % 128] +# CHECK-NEXT: O = sch.get_sblock("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-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul_relu(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: matmul = T.sblock_alloc_buffer((4, 32)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: relu = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: for i_0, j_0, i_1, j_1 in T.grid(2, 2, 2, 16): +# CHECK-NEXT: for k in range(512): +# CHECK-NEXT: with T.sblock("matmul"): +# CHECK-NEXT: v_i = T.axis.spatial(4, i_0 * 2 + i_1) +# CHECK-NEXT: v_j = T.axis.spatial(32, j_0 * 16 + j_1) +# CHECK-NEXT: v_k = T.axis.reduce(512, k) +# CHECK-NEXT: T.reads(_0[v_i, v_k], _1[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul[v_i, v_j] = matmul[v_i, v_j] + _0[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: with T.sblock("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(128, i_0 * 64 + i_1 * 32 + j_0 * 16 + j_1) +# CHECK-NEXT: T.reads(matmul[v_ax0 % 128 // 32, v_ax0 % 32]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0]) +# CHECK-NEXT: T_reshape_1[v_ax0] = matmul[v_ax0 % 128 // 32, v_ax0 % 32] +# CHECK-NEXT: for i in range(128): +# CHECK-NEXT: with T.sblock("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(128, 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 in T.grid(4, 32): +# CHECK-NEXT: with T.sblock("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) +# CHECK-NEXT: T.reads(relu[(v_ax0 * 32 + v_ax1) % 128]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0, v_ax1]) +# CHECK-NEXT: T_reshape[v_ax0, v_ax1] = relu[(v_ax0 * 32 + v_ax1) % 128] +# CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/test_matmul_relu_fused_tvm_tir.py b/tests/filecheck/backends/test_matmul_relu_fused_tvm_tir.py deleted file mode 100644 index 86ce8cb10..000000000 --- a/tests/filecheck/backends/test_matmul_relu_fused_tvm_tir.py +++ /dev/null @@ -1,45 +0,0 @@ -# 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_relu_subset_tvm.py b/tests/filecheck/backends/test_matmul_relu_subset_tvm.py index 7c8886c13..5068dcfd9 100644 --- a/tests/filecheck/backends/test_matmul_relu_subset_tvm.py +++ b/tests/filecheck/backends/test_matmul_relu_subset_tvm.py @@ -47,69 +47,93 @@ # CHECK-NEXT: - %3: relu(%2) {name = 'relu'} : [4x32xfloat32] -> [4x32xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: matmul = T.allocate([128], "float32", "global") -# CHECK-NEXT: for i, j in T.grid(4, 32): -# CHECK-NEXT: matmul_1 = T.Buffer((128,), data=matmul) -# CHECK-NEXT: matmul_1[i * 32 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(512): -# CHECK-NEXT: cse_var_1: T.int32 = i * 32 + j -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: matmul_1[cse_var_1] = matmul_1[cse_var_1] + _0_1[i * 512 + k] * _1_1[k * 32 + j] -# CHECK-NEXT: matmul_1 = T.Buffer((128,), data=matmul) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul_relu(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: matmul = T.sblock_alloc_buffer((4, 32)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: relu = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: for i, j, k in T.grid(4, 32, 512): +# CHECK-NEXT: with T.sblock("matmul"): +# 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(matmul[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul[v_i, v_j] = matmul[v_i, v_j] + _0[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: for ax0 in range(128): +# CHECK-NEXT: with T.sblock("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(128, ax0) +# CHECK-NEXT: T.reads(matmul[v_ax0 % 128 // 32, v_ax0 % 32]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0]) +# CHECK-NEXT: T_reshape_1[v_ax0] = matmul[v_ax0 % 128 // 32, v_ax0 % 32] # CHECK-NEXT: for i in range(128): -# CHECK-NEXT: matmul_2 = T.Buffer((128,), data=matmul) -# CHECK-NEXT: matmul_1[i] = T.max(T.float32(0.0), matmul_2[i]) +# CHECK-NEXT: with T.sblock("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(128, 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 in T.grid(4, 32): -# CHECK-NEXT: cse_var_2: T.int32 = ax0 * 32 + ax1 -# CHECK-NEXT: T_reshape_1 = T.Buffer((128,), data=T_reshape.data) -# CHECK-NEXT: T_reshape_1[cse_var_2] = matmul_1[cse_var_2] -# CHECK-NEXT: O = obj['matmul'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=2) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: sch[O].reorder(k, i, j, i1, j1) -# CHECK-NEXT: sch[O].unroll(i1) -# CHECK-NEXT: sch[O].vectorize(j1) +# CHECK-NEXT: with T.sblock("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) +# CHECK-NEXT: T.reads(relu[(v_ax0 * 32 + v_ax1) % 128]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0, v_ax1]) +# CHECK-NEXT: T_reshape[v_ax0, v_ax1] = relu[(v_ax0 * 32 + v_ax1) % 128] +# CHECK-NEXT: O = sch.get_sblock("matmul") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, = sch.split(i, factors=[None, 2]) +# CHECK-NEXT: j, j1, = sch.split(j, factors=[None, 16]) +# CHECK-NEXT: sch.reorder(k, i, j, i1, j1) +# CHECK-NEXT: sch.unroll(i1) +# CHECK-NEXT: sch.vectorize(j1) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: matmul = T.allocate([128], "float32", "global") -# CHECK-NEXT: matmul_1 = T.Buffer((128,), data=matmul) -# CHECK-NEXT: for i_outer_init, j_outer_init in T.grid(2, 2): -# CHECK-NEXT: cse_var_1: T.int32 = i_outer_init * 64 + j_outer_init * 16 -# CHECK-NEXT: matmul_1[cse_var_1:cse_var_1 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: matmul_1[cse_var_1 + 32:cse_var_1 + 32 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: for k, i_outer, j_outer in T.grid(512, 2, 2): -# CHECK-NEXT: cse_var_6: T.int32 = j_outer * 16 -# CHECK-NEXT: cse_var_5: T.int32 = i_outer * 1024 + k -# CHECK-NEXT: cse_var_4: T.int32 = k * 32 + cse_var_6 -# CHECK-NEXT: cse_var_3: T.int32 = i_outer * 64 + cse_var_6 -# CHECK-NEXT: cse_var_2: T.int32 = cse_var_3 + 32 -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: matmul_1[cse_var_3:cse_var_3 + 16] = matmul_1[cse_var_3:cse_var_3 + 16] + T.Broadcast(_0_1[cse_var_5], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: matmul_1[cse_var_2:cse_var_2 + 16] = matmul_1[cse_var_2:cse_var_2 + 16] + T.Broadcast(_0_1[cse_var_5 + 512], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: matmul_2 = T.Buffer((128,), data=matmul) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul_relu(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: matmul = T.sblock_alloc_buffer((4, 32)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: relu = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: for k, i_0, j_0 in T.grid(512, 2, 2): +# CHECK-NEXT: for i_1 in T.unroll(2): +# CHECK-NEXT: for j_1 in T.vectorized(16): +# CHECK-NEXT: with T.sblock("matmul"): +# CHECK-NEXT: v_i = T.axis.spatial(4, i_0 * 2 + i_1) +# CHECK-NEXT: v_j = T.axis.spatial(32, j_0 * 16 + j_1) +# CHECK-NEXT: v_k = T.axis.reduce(512, k) +# CHECK-NEXT: T.reads(_0[v_i, v_k], _1[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul[v_i, v_j] = matmul[v_i, v_j] + _0[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: for ax0 in range(128): +# CHECK-NEXT: with T.sblock("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(128, ax0) +# CHECK-NEXT: T.reads(matmul[v_ax0 % 128 // 32, v_ax0 % 32]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0]) +# CHECK-NEXT: T_reshape_1[v_ax0] = matmul[v_ax0 % 128 // 32, v_ax0 % 32] # CHECK-NEXT: for i in range(128): -# CHECK-NEXT: matmul_3 = T.Buffer((128,), data=matmul) -# CHECK-NEXT: matmul_2[i] = T.max(T.float32(0.0), matmul_3[i]) +# CHECK-NEXT: with T.sblock("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(128, 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 in T.grid(4, 32): -# CHECK-NEXT: cse_var_7: T.int32 = ax0 * 32 + ax1 -# CHECK-NEXT: T_reshape_1 = T.Buffer((128,), data=T_reshape.data) -# CHECK-NEXT: T_reshape_1[cse_var_7] = matmul_2[cse_var_7] +# CHECK-NEXT: with T.sblock("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) +# CHECK-NEXT: T.reads(relu[(v_ax0 * 32 + v_ax1) % 128]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0, v_ax1]) +# CHECK-NEXT: T_reshape[v_ax0, v_ax1] = relu[(v_ax0 * 32 + v_ax1) % 128] # CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/test_matmul_relu_tvm.py b/tests/filecheck/backends/test_matmul_relu_tvm.py index 6438562f9..523730761 100644 --- a/tests/filecheck/backends/test_matmul_relu_tvm.py +++ b/tests/filecheck/backends/test_matmul_relu_tvm.py @@ -47,69 +47,93 @@ # CHECK-NEXT: - %3: relu(%2) {name = 'relu'} : [4x32xfloat32] -> [4x32xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: matmul = T.allocate([128], "float32", "global") -# CHECK-NEXT: for i, j in T.grid(4, 32): -# CHECK-NEXT: matmul_1 = T.Buffer((128,), data=matmul) -# CHECK-NEXT: matmul_1[i * 32 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(512): -# CHECK-NEXT: cse_var_1: T.int32 = i * 32 + j -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: matmul_1[cse_var_1] = matmul_1[cse_var_1] + _0_1[i * 512 + k] * _1_1[k * 32 + j] -# CHECK-NEXT: matmul_1 = T.Buffer((128,), data=matmul) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul_relu(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: matmul = T.sblock_alloc_buffer((4, 32)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: relu = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: for i, j, k in T.grid(4, 32, 512): +# CHECK-NEXT: with T.sblock("matmul"): +# 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(matmul[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul[v_i, v_j] = matmul[v_i, v_j] + _0[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: for ax0 in range(128): +# CHECK-NEXT: with T.sblock("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(128, ax0) +# CHECK-NEXT: T.reads(matmul[v_ax0 % 128 // 32, v_ax0 % 32]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0]) +# CHECK-NEXT: T_reshape_1[v_ax0] = matmul[v_ax0 % 128 // 32, v_ax0 % 32] # CHECK-NEXT: for i in range(128): -# CHECK-NEXT: matmul_2 = T.Buffer((128,), data=matmul) -# CHECK-NEXT: matmul_1[i] = T.max(T.float32(0.0), matmul_2[i]) +# CHECK-NEXT: with T.sblock("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(128, 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 in T.grid(4, 32): -# CHECK-NEXT: cse_var_2: T.int32 = ax0 * 32 + ax1 -# CHECK-NEXT: T_reshape_1 = T.Buffer((128,), data=T_reshape.data) -# CHECK-NEXT: T_reshape_1[cse_var_2] = matmul_1[cse_var_2] -# CHECK-NEXT: O = obj['matmul'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=2) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: sch[O].reorder(k, i, j, i1, j1) -# CHECK-NEXT: sch[O].unroll(i1) -# CHECK-NEXT: sch[O].vectorize(j1) +# CHECK-NEXT: with T.sblock("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) +# CHECK-NEXT: T.reads(relu[(v_ax0 * 32 + v_ax1) % 128]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0, v_ax1]) +# CHECK-NEXT: T_reshape[v_ax0, v_ax1] = relu[(v_ax0 * 32 + v_ax1) % 128] +# CHECK-NEXT: O = sch.get_sblock("matmul") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, = sch.split(i, factors=[None, 2]) +# CHECK-NEXT: j, j1, = sch.split(j, factors=[None, 16]) +# CHECK-NEXT: sch.reorder(k, i, j, i1, j1) +# CHECK-NEXT: sch.unroll(i1) +# CHECK-NEXT: sch.vectorize(j1) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: matmul = T.allocate([128], "float32", "global") -# CHECK-NEXT: matmul_1 = T.Buffer((128,), data=matmul) -# CHECK-NEXT: for i_outer_init, j_outer_init in T.grid(2, 2): -# CHECK-NEXT: cse_var_1: T.int32 = i_outer_init * 64 + j_outer_init * 16 -# CHECK-NEXT: matmul_1[cse_var_1:cse_var_1 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: matmul_1[cse_var_1 + 32:cse_var_1 + 32 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: for k, i_outer, j_outer in T.grid(512, 2, 2): -# CHECK-NEXT: cse_var_6: T.int32 = j_outer * 16 -# CHECK-NEXT: cse_var_5: T.int32 = i_outer * 1024 + k -# CHECK-NEXT: cse_var_4: T.int32 = k * 32 + cse_var_6 -# CHECK-NEXT: cse_var_3: T.int32 = i_outer * 64 + cse_var_6 -# CHECK-NEXT: cse_var_2: T.int32 = cse_var_3 + 32 -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: matmul_1[cse_var_3:cse_var_3 + 16] = matmul_1[cse_var_3:cse_var_3 + 16] + T.Broadcast(_0_1[cse_var_5], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: matmul_1[cse_var_2:cse_var_2 + 16] = matmul_1[cse_var_2:cse_var_2 + 16] + T.Broadcast(_0_1[cse_var_5 + 512], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: matmul_2 = T.Buffer((128,), data=matmul) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul_relu(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), T_reshape: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: matmul = T.sblock_alloc_buffer((4, 32)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: relu = T.sblock_alloc_buffer((128,)) +# CHECK-NEXT: for k, i_0, j_0 in T.grid(512, 2, 2): +# CHECK-NEXT: for i_1 in T.unroll(2): +# CHECK-NEXT: for j_1 in T.vectorized(16): +# CHECK-NEXT: with T.sblock("matmul"): +# CHECK-NEXT: v_i = T.axis.spatial(4, i_0 * 2 + i_1) +# CHECK-NEXT: v_j = T.axis.spatial(32, j_0 * 16 + j_1) +# CHECK-NEXT: v_k = T.axis.reduce(512, k) +# CHECK-NEXT: T.reads(_0[v_i, v_k], _1[v_k, v_j]) +# CHECK-NEXT: T.writes(matmul[v_i, v_j]) +# CHECK-NEXT: with T.init(): +# CHECK-NEXT: matmul[v_i, v_j] = T.float32(0.0) +# CHECK-NEXT: matmul[v_i, v_j] = matmul[v_i, v_j] + _0[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: for ax0 in range(128): +# CHECK-NEXT: with T.sblock("T_reshape"): +# CHECK-NEXT: v_ax0 = T.axis.spatial(128, ax0) +# CHECK-NEXT: T.reads(matmul[v_ax0 % 128 // 32, v_ax0 % 32]) +# CHECK-NEXT: T.writes(T_reshape_1[v_ax0]) +# CHECK-NEXT: T_reshape_1[v_ax0] = matmul[v_ax0 % 128 // 32, v_ax0 % 32] # CHECK-NEXT: for i in range(128): -# CHECK-NEXT: matmul_3 = T.Buffer((128,), data=matmul) -# CHECK-NEXT: matmul_2[i] = T.max(T.float32(0.0), matmul_3[i]) +# CHECK-NEXT: with T.sblock("relu"): +# CHECK-NEXT: v_i = T.axis.spatial(128, 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 in T.grid(4, 32): -# CHECK-NEXT: cse_var_7: T.int32 = ax0 * 32 + ax1 -# CHECK-NEXT: T_reshape_1 = T.Buffer((128,), data=T_reshape.data) -# CHECK-NEXT: T_reshape_1[cse_var_7] = matmul_2[cse_var_7] +# CHECK-NEXT: with T.sblock("T_reshape_1"): +# CHECK-NEXT: v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) +# CHECK-NEXT: T.reads(relu[(v_ax0 * 32 + v_ax1) % 128]) +# CHECK-NEXT: T.writes(T_reshape[v_ax0, v_ax1]) +# CHECK-NEXT: T_reshape[v_ax0, v_ax1] = relu[(v_ax0 * 32 + v_ax1) % 128] # CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/backends/test_matmul_split.py b/tests/filecheck/backends/test_matmul_split.py index ef7f07129..8ac5aae63 100644 --- a/tests/filecheck/backends/test_matmul_split.py +++ b/tests/filecheck/backends/test_matmul_split.py @@ -14,11 +14,6 @@ if len(sys.argv) > 2 and sys.argv[2] == "--descript": descript = True -if backend == "tvm": - backend_kwargs = {"tir_schedule": True} -else: - backend_kwargs = {} - backend = import_module(f"xtc.backends.{backend}") I, J, K, dtype = 64, 256, 256, "float32" @@ -29,7 +24,7 @@ O.matmul(a, b, name="C") graph = gb.graph -impl = backend.Backend(graph, **backend_kwargs) +impl = backend.Backend(graph) sch = impl.get_scheduler() if descript: diff --git a/tests/filecheck/backends/test_matmul_split_nested.py b/tests/filecheck/backends/test_matmul_split_nested.py index 719f54f71..c884fc0ea 100644 --- a/tests/filecheck/backends/test_matmul_split_nested.py +++ b/tests/filecheck/backends/test_matmul_split_nested.py @@ -14,11 +14,6 @@ if len(sys.argv) > 2 and sys.argv[2] == "--descript": descript = True -if backend == "tvm": - backend_kwargs = {"tir_schedule": True} -else: - backend_kwargs = {} - backend = import_module(f"xtc.backends.{backend}") I, J, K, dtype = 64, 192, 256, "float32" @@ -29,7 +24,7 @@ O.matmul(a, b, name="C") graph = gb.graph -impl = backend.Backend(graph, **backend_kwargs) +impl = backend.Backend(graph) sch = impl.get_scheduler() if descript: diff --git a/tests/filecheck/backends/test_matmul_tvm.py b/tests/filecheck/backends/test_matmul_tvm.py index bf7d22484..e7e18c7f8 100644 --- a/tests/filecheck/backends/test_matmul_tvm.py +++ b/tests/filecheck/backends/test_matmul_tvm.py @@ -4,7 +4,7 @@ import xtc.graphs.xtc.op as O from xtc.backends.tvm import Backend -I, J, K, dtype = 4, 32, 512, "float32" +I, J, K, dtype = 4, 32, 256, "float32" a = O.tensor((I, K), dtype, name="A") b = O.tensor((K, J), dtype, name="B") @@ -30,6 +30,7 @@ print_source_ir=True, print_transformed_ir=True, ) + module = comp.compile(sched) executor = module.get_executor(validate=True) res = executor.execute() @@ -38,59 +39,59 @@ # CHECK: graph: # CHECK-NEXT: name: matmul # CHECK-NEXT: inputs: -# CHECK-NEXT: - %0 : 4x512xfloat32 -# CHECK-NEXT: - %1 : 512x32xfloat32 +# CHECK-NEXT: - %0 : 4x256xfloat32 +# CHECK-NEXT: - %1 : 256x32xfloat32 # CHECK-NEXT: outputs: # CHECK-NEXT: - %2 : 4x32xfloat32 # CHECK-NEXT: nodes: -# CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [4x512xfloat32, 512x32xfloat32] -> [4x32xfloat32] +# CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [4x256xfloat32, 256x32xfloat32] -> [4x32xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for i, j in T.grid(4, 32): -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: C_1[i * 32 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(512): -# CHECK-NEXT: cse_var_1: T.int32 = i * 32 + j -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_1] = C_1[cse_var_1] + _0_1[i * 512 + k] * _1_1[k * 32 + j] -# CHECK-NEXT: O = obj['C'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=2) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: sch[O].reorder(k, i, j, i1, j1) -# CHECK-NEXT: sch[O].unroll(i1) -# CHECK-NEXT: sch[O].vectorize(j1) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j, k in T.grid(4, 32, 256): +# CHECK-NEXT: with T.sblock("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_sblock("C") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, = sch.split(i, factors=[None, 2]) +# CHECK-NEXT: j, j1, = sch.split(j, factors=[None, 16]) +# CHECK-NEXT: sch.reorder(k, i, j, i1, j1) +# CHECK-NEXT: sch.unroll(i1) +# CHECK-NEXT: sch.vectorize(j1) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: for i_outer_init, j_outer_init in T.grid(2, 2): -# CHECK-NEXT: cse_var_1: T.int32 = i_outer_init * 64 + j_outer_init * 16 -# CHECK-NEXT: C_1[cse_var_1:cse_var_1 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: C_1[cse_var_1 + 32:cse_var_1 + 32 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: for k, i_outer, j_outer in T.grid(512, 2, 2): -# CHECK-NEXT: cse_var_6: T.int32 = j_outer * 16 -# CHECK-NEXT: cse_var_5: T.int32 = i_outer * 1024 + k -# CHECK-NEXT: cse_var_4: T.int32 = k * 32 + cse_var_6 -# CHECK-NEXT: cse_var_3: T.int32 = i_outer * 64 + cse_var_6 -# CHECK-NEXT: cse_var_2: T.int32 = cse_var_3 + 32 -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_3:cse_var_3 + 16] = C_1[cse_var_3:cse_var_3 + 16] + T.Broadcast(_0_1[cse_var_5], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: C_1[cse_var_2:cse_var_2 + 16] = C_1[cse_var_2:cse_var_2 + 16] + T.Broadcast(_0_1[cse_var_5 + 512], 16) * _1_1[cse_var_4:cse_var_4 + 16] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for k, i_0, j_0 in T.grid(256, 2, 2): +# CHECK-NEXT: for i_1 in T.unroll(2): +# CHECK-NEXT: for j_1 in T.vectorized(16): +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i = T.axis.spatial(4, i_0 * 2 + i_1) +# CHECK-NEXT: v_j = T.axis.spatial(32, j_0 * 16 + j_1) +# CHECK-NEXT: v_k = T.axis.reduce(256, 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: CODE: 0 diff --git a/tests/filecheck/backends/test_matmul_unroll_tvm.py b/tests/filecheck/backends/test_matmul_unroll_tvm.py index 07cff71cb..fe2586082 100644 --- a/tests/filecheck/backends/test_matmul_unroll_tvm.py +++ b/tests/filecheck/backends/test_matmul_unroll_tvm.py @@ -43,47 +43,47 @@ # CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [4x256xfloat32, 256x32xfloat32] -> [4x32xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for i, j in T.grid(4, 32): -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: C_1[i * 32 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(256): -# CHECK-NEXT: cse_var_1: T.int32 = i * 32 + j -# CHECK-NEXT: _0_1 = T.Buffer((1024,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((8192,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_1] = C_1[cse_var_1] + _0_1[i * 256 + k] * _1_1[k * 32 + j] -# CHECK-NEXT: O = obj['C'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: k, __u_k = sch[O].split(k, factor=4) -# CHECK-NEXT: sch[O].reorder(i, j, k, __u_k) -# CHECK-NEXT: sch[O].unroll(__u_k) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j, k in T.grid(4, 32, 256): +# CHECK-NEXT: with T.sblock("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_sblock("C") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: k, __u_k, = sch.split(k, factors=[None, 4]) +# CHECK-NEXT: sch.reorder(i, j, k, __u_k) +# CHECK-NEXT: sch.unroll(__u_k) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for i, j in T.grid(4, 32): -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: C_1[i * 32 + j] = T.float32(0.0) -# CHECK-NEXT: for k_outer in range(64): -# CHECK-NEXT: cse_var_3: T.int32 = k_outer * 128 + j -# CHECK-NEXT: cse_var_2: T.int32 = i * 32 + j -# CHECK-NEXT: cse_var_1: T.int32 = i * 256 + k_outer * 4 -# CHECK-NEXT: _0_1 = T.Buffer((1024,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((8192,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_2] = C_1[cse_var_2] + _0_1[cse_var_1] * _1_1[cse_var_3] -# CHECK-NEXT: C_1[cse_var_2] = C_1[cse_var_2] + _0_1[cse_var_1 + 1] * _1_1[cse_var_3 + 32] -# CHECK-NEXT: C_1[cse_var_2] = C_1[cse_var_2] + _0_1[cse_var_1 + 2] * _1_1[cse_var_3 + 64] -# CHECK-NEXT: C_1[cse_var_2] = C_1[cse_var_2] + _0_1[cse_var_1 + 3] * _1_1[cse_var_3 + 96] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j, k_0 in T.grid(4, 32, 64): +# CHECK-NEXT: for k_1 in T.unroll(4): +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i, v_j = T.axis.remap("SS", [i, j]) +# CHECK-NEXT: v_k = T.axis.reduce(256, k_0 * 4 + k_1) +# 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: CODE: 0 diff --git a/tests/filecheck/backends/test_matmul_vectorize_tvm.py b/tests/filecheck/backends/test_matmul_vectorize_tvm.py index 4b09e674a..abcaef42d 100644 --- a/tests/filecheck/backends/test_matmul_vectorize_tvm.py +++ b/tests/filecheck/backends/test_matmul_vectorize_tvm.py @@ -45,55 +45,51 @@ # CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [4x256xfloat32, 256x32xfloat32] -> [4x32xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for i, j in T.grid(4, 32): -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: C_1[i * 32 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(256): -# CHECK-NEXT: cse_var_1: T.int32 = i * 32 + j -# CHECK-NEXT: _0_1 = T.Buffer((1024,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((8192,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_1] = C_1[cse_var_1] + _0_1[i * 256 + k] * _1_1[k * 32 + j] -# CHECK-NEXT: O = obj['C'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: j, j0 = sch[O].split(j, factor=24) -# CHECK-NEXT: j0, __v_j0 = sch[O].split(j0, factor=8) -# CHECK-NEXT: sch[O].reorder(i, j, k, j0, __v_j0) -# CHECK-NEXT: sch[O].unroll(j0) -# CHECK-NEXT: sch[O].vectorize(__v_j0) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j, k in T.grid(4, 32, 256): +# CHECK-NEXT: with T.sblock("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_sblock("C") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: j, j0, __v_j0, = sch.split(j, factors=[None, 3, 8]) +# CHECK-NEXT: sch.reorder(i, j, k, j0, __v_j0) +# CHECK-NEXT: sch.unroll(j0) +# CHECK-NEXT: sch.vectorize(__v_j0) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for i, j_outer in T.grid(4, 2): -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: C_1[i * 32 + j_outer * 24:i * 32 + j_outer * 24 + 8] = T.Broadcast(T.float32(0.0), 8) -# CHECK-NEXT: if T.likely(j_outer < 1): -# CHECK-NEXT: C_1[i * 32 + j_outer * 24 + 8:i * 32 + j_outer * 24 + 8 + 8] = T.Broadcast(T.float32(0.0), 8) -# CHECK-NEXT: if T.likely(j_outer < 1): -# CHECK-NEXT: C_1[i * 32 + j_outer * 24 + 16:i * 32 + j_outer * 24 + 16 + 8] = T.Broadcast(T.float32(0.0), 8) -# CHECK-NEXT: for k in range(256): -# CHECK-NEXT: cse_var_2: T.int32 = j_outer * 24 -# CHECK-NEXT: cse_var_1: T.int32 = i * 32 + cse_var_2 -# CHECK-NEXT: _0_1 = T.Buffer((1024,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((8192,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_1:cse_var_1 + 8] = C_1[cse_var_1:cse_var_1 + 8] + T.Broadcast(_0_1[i * 256 + k], 8) * _1_1[k * 32 + cse_var_2:k * 32 + cse_var_2 + 8] -# CHECK-NEXT: if T.likely(j_outer < 1): -# CHECK-NEXT: cse_var_3: T.int32 = cse_var_1 + 8 -# CHECK-NEXT: C_1[cse_var_3:cse_var_3 + 8] = C_1[cse_var_3:cse_var_3 + 8] + T.Broadcast(_0_1[i * 256 + k], 8) * _1_1[k * 32 + cse_var_2 + 8:k * 32 + cse_var_2 + 8 + 8] -# CHECK-NEXT: if T.likely(j_outer < 1): -# CHECK-NEXT: cse_var_4: T.int32 = cse_var_1 + 16 -# CHECK-NEXT: C_1[cse_var_4:cse_var_4 + 8] = C_1[cse_var_4:cse_var_4 + 8] + T.Broadcast(_0_1[i * 256 + k], 8) * _1_1[k * 32 + cse_var_2 + 16:k * 32 + cse_var_2 + 16 + 8] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 256), "float32"), _1: T.Buffer((256, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j_0, k in T.grid(4, 2, 256): +# CHECK-NEXT: for j_1 in T.unroll(3): +# CHECK-NEXT: for j_2 in T.vectorized(8): +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i = T.axis.spatial(4, i) +# CHECK-NEXT: v_j = T.axis.spatial(32, j_0 * 24 + j_1 * 8 + j_2) +# CHECK-NEXT: v_k = T.axis.reduce(256, k) +# CHECK-NEXT: T.where((j_0 * 3 + j_1) * 8 + j_2 < 32) +# 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: CODE: 0 diff --git a/tests/filecheck/backends/test_pad_conv2d_fused_tvm.py b/tests/filecheck/backends/test_pad_conv2d_fused_tvm.py index 2f6c5d1eb..1906e6c46 100644 --- a/tests/filecheck/backends/test_pad_conv2d_fused_tvm.py +++ b/tests/filecheck/backends/test_pad_conv2d_fused_tvm.py @@ -46,55 +46,65 @@ # CHECK-NEXT: - %3: conv2d(%2, %1, stride=(2, 2)) {name = 'conv'} : [1x12x12x3xfloat32, 5x5x3x16xfloat32] -> [1x4x4x16xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([363], "float32", "global") -# CHECK-NEXT: pad_1 = T.Buffer((363,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(11, 11, 3): -# CHECK-NEXT: cse_var_1: T.int32 = i2 * 3 -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 33 + cse_var_1 + i3] = T.if_then_else(2 <= i1 and i1 < 10 and 2 <= i2 and i2 < 10, _0_1[i1 * 24 + cse_var_1 + i3 - 54], T.float32(0.0)) -# CHECK-NEXT: for h, w, f in T.grid(4, 4, 16): -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16 + f] = T.float32(0.0) -# CHECK-NEXT: for r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_2: T.int32 = h * 64 + w * 16 + f -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_2] = conv_1[cse_var_2] + pad_1[h * 66 + r * 33 + w * 6 + s * 3 + c] * _1_1[r * 240 + s * 48 + c * 16 + f] -# CHECK-NEXT: O = obj['conv'] -# CHECK-NEXT: I_F0 = O.op.input_tensors[0] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: sch[O].reorder(b, h, w, r, s, c, f) -# CHECK-NEXT: sch[I_F0].compute_at(sch[O], w) -# CHECK-NEXT: sch[O].vectorize(f) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for i0, i1, i2, i3 in T.grid(1, 12, 12, 3): +# CHECK-NEXT: with T.sblock("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.sblock("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: O = sch.get_sblock("conv") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# 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, 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: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: pad = T.allocate([75], "float32", "global") -# CHECK-NEXT: for h, w in T.grid(4, 4): -# CHECK-NEXT: pad_1 = T.Buffer((75,), data=pad) -# CHECK-NEXT: for i1, i2, i3 in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_3: T.int32 = i2 * 3 -# CHECK-NEXT: cse_var_2: T.int32 = i1 // 2 + h -# CHECK-NEXT: cse_var_1: T.int32 = i2 // 2 + w -# CHECK-NEXT: _0_1 = T.Buffer((192,), data=_0.data) -# CHECK-NEXT: pad_1[i1 * 15 + cse_var_3 + i3] = T.if_then_else(1 <= cse_var_2 and cse_var_2 < 5 and 1 <= cse_var_1 and cse_var_1 < 5, _0_1[h * 48 + i1 * 24 + w * 6 + cse_var_3 + i3 - 54], T.float32(0.0)) -# CHECK-NEXT: conv_1 = T.Buffer((256,), data=conv.data) -# CHECK-NEXT: conv_1[h * 64 + w * 16:h * 64 + w * 16 + 16] = T.Broadcast(T.float32(0.0), 16) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def pad_conv2d_nhwc_mini(_0: T.Buffer((1, 8, 8, 3), "float32"), _1: T.Buffer((5, 5, 3, 16), "float32"), conv: T.Buffer((1, 4, 4, 16), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: for b, h, w in T.grid(1, 4, 4): +# CHECK-NEXT: for ax0, ax1, ax2 in T.grid(5, 5, 3): +# CHECK-NEXT: with T.sblock("pad"): +# CHECK-NEXT: v_i0 = T.axis.spatial(1, 0) +# CHECK-NEXT: v_i1 = T.axis.spatial(12, h * 2 + ax0) +# CHECK-NEXT: v_i2 = T.axis.spatial(12, w * 2 + ax1) +# CHECK-NEXT: v_i3 = T.axis.spatial(3, ax2) +# 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 r, s, c in T.grid(5, 5, 3): -# CHECK-NEXT: cse_var_4: T.int32 = h * 64 + w * 16 -# CHECK-NEXT: _1_1 = T.Buffer((1200,), data=_1.data) -# CHECK-NEXT: conv_1[cse_var_4:cse_var_4 + 16] = conv_1[cse_var_4:cse_var_4 + 16] + T.Broadcast(pad_1[r * 15 + s * 3 + c], 16) * _1_1[r * 240 + s * 48 + c * 16:r * 240 + s * 48 + c * 16 + 16] +# CHECK-NEXT: for f in T.vectorized(16): +# CHECK-NEXT: with T.sblock("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: 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.py similarity index 81% rename from tests/filecheck/backends/test_pad_conv2d_relu_fused_tvm_tir.py rename to tests/filecheck/backends/test_pad_conv2d_relu_fused_tvm.py index 7c587877c..ba4dd5c55 100644 --- a/tests/filecheck/backends/test_pad_conv2d_relu_fused_tvm_tir.py +++ b/tests/filecheck/backends/test_pad_conv2d_relu_fused_tvm.py @@ -17,7 +17,7 @@ graph = gb.graph print(graph) -impl = Backend(graph, tir_schedule=True) +impl = Backend(graph) sch = impl.get_scheduler(default_node="conv") sch.interchange(["b", "h", "w", "r", "s", "c", "f"]) @@ -27,7 +27,7 @@ sched = sch.schedule() comp = impl.get_compiler( shared_lib=True, - dump_file="pad_conv2d_relu_fused_tvm_tir", + dump_file="pad_conv2d_relu_fused_tvm", print_source_ir=True, print_transformed_ir=True, ) @@ -49,26 +49,27 @@ # 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: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: @T.prim_func(s_tir=True) # 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: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: conv = T.sblock_alloc_buffer((1, 4, 4, 16)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((256,)) +# CHECK-NEXT: relu = T.sblock_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: with T.sblock("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: with T.sblock("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]) @@ -76,24 +77,24 @@ # 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: with T.sblock("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: with T.sblock("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: with T.sblock("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: O = sch.get_sblock("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] @@ -103,22 +104,23 @@ # 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: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func +# CHECK-NEXT: @T.prim_func(s_tir=True) # 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: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: pad = T.sblock_alloc_buffer((1, 12, 12, 3)) +# CHECK-NEXT: conv = T.sblock_alloc_buffer((1, 4, 4, 16)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((256,)) +# CHECK-NEXT: relu = T.sblock_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: with T.sblock("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) @@ -128,7 +130,7 @@ # 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: with T.sblock("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]) @@ -136,19 +138,19 @@ # 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: with T.sblock("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: with T.sblock("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: with T.sblock("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]) diff --git a/tests/filecheck/backends/test_relu_matmul_fused_tvm.py b/tests/filecheck/backends/test_relu_matmul_fused_tvm.py index 6b072a4af..4ad285bfa 100644 --- a/tests/filecheck/backends/test_relu_matmul_fused_tvm.py +++ b/tests/filecheck/backends/test_relu_matmul_fused_tvm.py @@ -32,7 +32,7 @@ comp = impl.get_compiler( shared_lib=True, - dump_file="relu_matmul_tvm_fused", + dump_file="relu_matmul_fused_tvm", print_source_ir=True, print_transformed_ir=True, ) @@ -54,105 +54,123 @@ # 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: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # 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: @T.prim_func(s_tir=True) +# 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({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: T_reshape = T.sblock_alloc_buffer((4096,)) +# CHECK-NEXT: relu = T.sblock_alloc_buffer((4096,)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((64, 64)) # 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: with T.sblock("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: 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: with T.sblock("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.sblock("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.sblock("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_sblock("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: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # 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: @T.prim_func(s_tir=True) +# 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({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: T_reshape = T.sblock_alloc_buffer((4096,)) +# CHECK-NEXT: relu = T.sblock_alloc_buffer((4096,)) +# CHECK-NEXT: T_reshape_1 = T.sblock_alloc_buffer((64, 64)) +# CHECK-NEXT: _1_global = T.sblock_alloc_buffer((64, 64)) +# CHECK-NEXT: C_global = T.sblock_alloc_buffer((64, 64)) # 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: with T.sblock("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: 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: with T.sblock("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: _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: with T.sblock("_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.sblock_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.sblock("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.sblock("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.sblock("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/backends/test_relu_matmul_fused_tvm_tir.py b/tests/filecheck/backends/test_relu_matmul_fused_tvm_tir.py deleted file mode 100644 index 5afeea78c..000000000 --- a/tests/filecheck/backends/test_relu_matmul_fused_tvm_tir.py +++ /dev/null @@ -1,174 +0,0 @@ -# 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/mlir_loop/gen_assembly/skylake_generic_matmul.mlir b/tests/filecheck/mlir_loop/gen_assembly/skylake_generic_matmul.mlir index 39a0980bd..0152f0d0e 100644 --- a/tests/filecheck/mlir_loop/gen_assembly/skylake_generic_matmul.mlir +++ b/tests/filecheck/mlir_loop/gen_assembly/skylake_generic_matmul.mlir @@ -36,9 +36,7 @@ func.func @myfun( } return } -// CHECK: Disassembly of section .text: -// CHECK-NEXT: -// CHECK-NEXT: : +// CHECK: : // CHECK-NEXT: add $0x1c00,%rsi // CHECK-NEXT: add $0x1c,%rdi // CHECK-NEXT: xor %eax,%eax @@ -149,4 +147,4 @@ func.func @myfun( // CHECK-NEXT: cmp $0xff,%rax // CHECK-NEXT: lea 0x1(%rax),%rax // CHECK-NEXT: jb -// CHECK-NEXT: vzeroupper +// CHECK-NEXT: vzeroupper diff --git a/tests/filecheck/mlir_loop/gen_assembly/skylake_matmul.mlir b/tests/filecheck/mlir_loop/gen_assembly/skylake_matmul.mlir index 29b0910ac..ddb488d72 100644 --- a/tests/filecheck/mlir_loop/gen_assembly/skylake_matmul.mlir +++ b/tests/filecheck/mlir_loop/gen_assembly/skylake_matmul.mlir @@ -23,9 +23,7 @@ func.func @myfun( outs(%C : memref<256x256xf32>) return } -// CHECK: Disassembly of section .text: -// CHECK-NEXT: -// CHECK-NEXT: : +// CHECK: : // CHECK-NEXT: add $0x1c00,%rsi // CHECK-NEXT: add $0x1c,%rdi // CHECK-NEXT: xor %eax,%eax @@ -136,4 +134,4 @@ func.func @myfun( // CHECK-NEXT: cmp $0xff,%rax // CHECK-NEXT: lea 0x1(%rax),%rax // CHECK-NEXT: jb -// CHECK-NEXT: vzeroupper +// CHECK-NEXT: vzeroupper diff --git a/tests/filecheck/mlir_loop/gen_assembly/skylake_small_matmul.mlir b/tests/filecheck/mlir_loop/gen_assembly/skylake_small_matmul.mlir index 5f88b5ea6..37a959b6f 100644 --- a/tests/filecheck/mlir_loop/gen_assembly/skylake_small_matmul.mlir +++ b/tests/filecheck/mlir_loop/gen_assembly/skylake_small_matmul.mlir @@ -23,9 +23,7 @@ func.func @myfun( outs(%C : memref<8x8xf32>) return } -// CHECK: Disassembly of section .text: -// CHECK-NEXT: -// CHECK-NEXT: : +// CHECK: : // CHECK-NEXT: vmovups (%rsi),%ymm0 // CHECK-NEXT: vmovups 0x20(%rsi),%ymm1 // CHECK-NEXT: vmovups 0x40(%rsi),%ymm2 @@ -57,4 +55,4 @@ func.func @myfun( // CHECK-NEXT: add $0x20,%rcx // CHECK-NEXT: cmp $0x7,%rax // CHECK-NEXT: jb -// CHECK-NEXT: vzeroupper +// CHECK-NEXT: vzeroupper diff --git a/tests/filecheck/mlir_loop/gen_assembly/skylake_split_matmul.mlir b/tests/filecheck/mlir_loop/gen_assembly/skylake_split_matmul.mlir index cfc0cbdc1..8964ac109 100644 --- a/tests/filecheck/mlir_loop/gen_assembly/skylake_split_matmul.mlir +++ b/tests/filecheck/mlir_loop/gen_assembly/skylake_split_matmul.mlir @@ -28,9 +28,7 @@ func.func @myfun( outs(%C : memref<258x256xf32>) return } -// CHECK: Disassembly of section .text: -// CHECK-NEXT: -// CHECK-NEXT: : +// CHECK: : // CHECK-NEXT: xor %ecx,%ecx // CHECK-NEXT: mov %rsi,%rax // CHECK-NEXT: vmovss (%rdi,%rcx,4),%xmm0 diff --git a/tests/filecheck/mlir_loop/gen_assembly/tigerlake_matmul.mlir b/tests/filecheck/mlir_loop/gen_assembly/tigerlake_matmul.mlir index 40555a953..369dcbcd7 100644 --- a/tests/filecheck/mlir_loop/gen_assembly/tigerlake_matmul.mlir +++ b/tests/filecheck/mlir_loop/gen_assembly/tigerlake_matmul.mlir @@ -23,9 +23,7 @@ func.func @myfun( outs(%C : memref<256x256xf32>) return } -// CHECK: Disassembly of section .text: -// CHECK-NEXT: -// CHECK-NEXT: : +// CHECK: : // CHECK-NEXT: add $0x1c00,%rsi // CHECK-NEXT: add $0x1c,%rdi // CHECK-NEXT: xor %eax,%eax @@ -96,4 +94,4 @@ func.func @myfun( // CHECK-NEXT: cmp $0xff,%rax // CHECK-NEXT: lea 0x1(%rax),%rax // CHECK-NEXT: jb -// CHECK-NEXT: vzeroupper +// CHECK-NEXT: vzeroupper diff --git a/tests/filecheck/schedules/test_matmul_descript_extend_tvm_goto.py b/tests/filecheck/schedules/test_matmul_descript_extend_tvm_goto.py index e23336859..c26da306f 100644 --- a/tests/filecheck/schedules/test_matmul_descript_extend_tvm_goto.py +++ b/tests/filecheck/schedules/test_matmul_descript_extend_tvm_goto.py @@ -60,142 +60,81 @@ res = executor.execute() print(f"CODE: {res}") -#CHECK: graph: -#CHECK-NEXT: name: matmul -#CHECK-NEXT: inputs: -#CHECK-NEXT: - %0 : 512x512xfloat32 -#CHECK-NEXT: - %1 : 512x512xfloat32 -#CHECK-NEXT: outputs: -#CHECK-NEXT: - %2 : 512x512xfloat32 -#CHECK-NEXT: nodes: -#CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [512x512xfloat32, 512x512xfloat32] -> [512x512xfloat32] -#CHECK-EMPTY: -#CHECK-NEXT:# from tvm.script import ir as I -#CHECK-NEXT:# from tvm.script import tir as T -#CHECK-EMPTY: -#CHECK-NEXT:@I.ir_module -#CHECK-NEXT:class Module: -#CHECK-NEXT: @T.prim_func -#CHECK-NEXT: def main(_0: T.Buffer((512, 512), "float32"), _1: T.Buffer((512, 512), "float32"), C: T.Buffer((512, 512), "float32")): -#CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -#CHECK-NEXT: for i, j in T.grid(512, 512): -#CHECK-NEXT: C_1 = T.Buffer((262144,), data=C.data) -#CHECK-NEXT: C_1[i * 512 + j] = T.float32(0.0) -#CHECK-NEXT: for k in range(512): -#CHECK-NEXT: cse_var_2: T.int32 = i * 512 -#CHECK-NEXT: cse_var_1: T.int32 = cse_var_2 + j -#CHECK-NEXT: _0_1 = T.Buffer((262144,), data=_0.data) -#CHECK-NEXT: _1_1 = T.Buffer((262144,), data=_1.data) -#CHECK-NEXT: C_1[cse_var_1] = C_1[cse_var_1] + _0_1[cse_var_2 + k] * _1_1[k * 512 + j] -#CHECK-NEXT:INPS = list(obj.values())[:-1] -#CHECK-NEXT:O = obj['C'] -#CHECK-NEXT:I_R0 = sch.cache_read(INPS[0], "global", [O]) -#CHECK-NEXT:i, j, = O.op.axis -#CHECK-NEXT:k, = O.op.reduce_axis -#CHECK-NEXT:j, j0 = sch[O].split(j, factor=36) -#CHECK-NEXT:i, i0 = sch[O].split(i, factor=128) -#CHECK-NEXT:k, k0 = sch[O].split(k, factor=16) -#CHECK-NEXT:k0, __u_k0 = sch[O].split(k0, factor=2) -#CHECK-NEXT:i0, i1 = sch[O].split(i0, factor=2) -#CHECK-NEXT:j0, j1 = sch[O].split(j0, factor=6) -#CHECK-NEXT:j1, __v_j1 = sch[O].split(j1, factor=2) -#CHECK-NEXT:sch[O].reorder(j, k, i, j0, i0, k0, __u_k0, i1, j1, __v_j1) -#CHECK-NEXT:sch[I_R0].compute_at(sch[O], i) -#CHECK-NEXT:sch[O].unroll(__u_k0) -#CHECK-NEXT:sch[O].unroll(i1) -#CHECK-NEXT:sch[O].unroll(j1) -#CHECK-NEXT:sch[O].vectorize(__v_j1) -#CHECK-NEXT:sch[O].parallel(j) -#CHECK-EMPTY: -#CHECK-NEXT:# from tvm.script import ir as I -#CHECK-NEXT:# from tvm.script import tir as T -#CHECK-EMPTY: -#CHECK-NEXT:@I.ir_module -#CHECK-NEXT:class Module: -#CHECK-NEXT: @T.prim_func -#CHECK-NEXT: def main(_0: T.Buffer((512, 512), "float32"), _1: T.Buffer((512, 512), "float32"), C: T.Buffer((512, 512), "float32")): -#CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -#CHECK-NEXT: for j_outer in T.parallel(15): -#CHECK-NEXT: _0_global = T.allocate([2048], "float32", "global") -#CHECK-NEXT: C_1 = T.Buffer((262144,), data=C.data) -#CHECK-NEXT: for i_outer_init, j_inner_outer_init, i_inner_outer_init in T.grid(4, 6, 64): -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer_init * 3 // 2 < 128): -#CHECK-NEXT: C_1[i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6:i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 2] = T.Broadcast(T.float32(0.0), 2) -#CHECK-NEXT: if T.likely(j_outer * 9 + (j_inner_outer_init * 3 + 1) // 2 < 128): -#CHECK-NEXT: C_1[i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 2:i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 2 + 2] = T.Broadcast(T.float32(0.0), 2) -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer_init * 3 // 2 < 127): -#CHECK-NEXT: C_1[i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 4:i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 4 + 2] = T.Broadcast(T.float32(0.0), 2) -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer_init * 3 // 2 < 128): -#CHECK-NEXT: C_1[i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 512:i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 512 + 2] = T.Broadcast(T.float32(0.0), 2) -#CHECK-NEXT: if T.likely(j_outer * 9 + (j_inner_outer_init * 3 + 1) // 2 < 128): -#CHECK-NEXT: C_1[i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 514:i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 514 + 2] = T.Broadcast(T.float32(0.0), 2) -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer_init * 3 // 2 < 127): -#CHECK-NEXT: C_1[i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 516:i_outer_init * 65536 + i_inner_outer_init * 1024 + j_outer * 36 + j_inner_outer_init * 6 + 516 + 2] = T.Broadcast(T.float32(0.0), 2) -#CHECK-NEXT: for k_outer, i_outer in T.grid(32, 4): -#CHECK-NEXT: _0_global_1 = T.Buffer((2048,), data=_0_global) -#CHECK-NEXT: for ax0, ax1 in T.grid(128, 16): -#CHECK-NEXT: _0_1 = T.Buffer((262144,), data=_0.data) -#CHECK-NEXT: _0_global_1[ax0 * 16 + ax1] = _0_1[i_outer * 65536 + ax0 * 512 + k_outer * 16 + ax1] -#CHECK-NEXT: for j_inner_outer, i_inner_outer, k_inner_outer in T.grid(6, 64, 8): -#CHECK-NEXT: _1_1 = T.Buffer((262144,), data=_1.data) -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer * 3 // 2 < 128): -#CHECK-NEXT: cse_var_3: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_2: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_1: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_3 + cse_var_2 -#CHECK-NEXT: C_1[cse_var_1:cse_var_1 + 2] = C_1[cse_var_1:cse_var_1 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_3 + cse_var_2:k_outer * 8192 + k_inner_outer * 1024 + cse_var_3 + cse_var_2 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + (j_inner_outer * 3 + 1) // 2 < 128): -#CHECK-NEXT: cse_var_6: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_5: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_4: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_6 + cse_var_5 + 2 -#CHECK-NEXT: C_1[cse_var_4:cse_var_4 + 2] = C_1[cse_var_4:cse_var_4 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_6 + cse_var_5 + 2:k_outer * 8192 + k_inner_outer * 1024 + cse_var_6 + cse_var_5 + 2 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer * 3 // 2 < 127): -#CHECK-NEXT: cse_var_9: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_8: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_7: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_9 + cse_var_8 + 4 -#CHECK-NEXT: C_1[cse_var_7:cse_var_7 + 2] = C_1[cse_var_7:cse_var_7 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_9 + cse_var_8 + 4:k_outer * 8192 + k_inner_outer * 1024 + cse_var_9 + cse_var_8 + 4 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer * 3 // 2 < 128): -#CHECK-NEXT: cse_var_12: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_11: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_10: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_12 + cse_var_11 + 512 -#CHECK-NEXT: C_1[cse_var_10:cse_var_10 + 2] = C_1[cse_var_10:cse_var_10 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 16], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_12 + cse_var_11:k_outer * 8192 + k_inner_outer * 1024 + cse_var_12 + cse_var_11 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + (j_inner_outer * 3 + 1) // 2 < 128): -#CHECK-NEXT: cse_var_15: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_14: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_13: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_15 + cse_var_14 + 514 -#CHECK-NEXT: C_1[cse_var_13:cse_var_13 + 2] = C_1[cse_var_13:cse_var_13 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 16], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_15 + cse_var_14 + 2:k_outer * 8192 + k_inner_outer * 1024 + cse_var_15 + cse_var_14 + 2 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer * 3 // 2 < 127): -#CHECK-NEXT: cse_var_18: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_17: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_16: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_18 + cse_var_17 + 516 -#CHECK-NEXT: C_1[cse_var_16:cse_var_16 + 2] = C_1[cse_var_16:cse_var_16 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 16], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_18 + cse_var_17 + 4:k_outer * 8192 + k_inner_outer * 1024 + cse_var_18 + cse_var_17 + 4 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer * 3 // 2 < 128): -#CHECK-NEXT: cse_var_21: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_20: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_19: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_21 + cse_var_20 -#CHECK-NEXT: C_1[cse_var_19:cse_var_19 + 2] = C_1[cse_var_19:cse_var_19 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 1], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_21 + cse_var_20 + 512:k_outer * 8192 + k_inner_outer * 1024 + cse_var_21 + cse_var_20 + 512 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + (j_inner_outer * 3 + 1) // 2 < 128): -#CHECK-NEXT: cse_var_24: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_23: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_22: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_24 + cse_var_23 + 2 -#CHECK-NEXT: C_1[cse_var_22:cse_var_22 + 2] = C_1[cse_var_22:cse_var_22 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 1], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_24 + cse_var_23 + 514:k_outer * 8192 + k_inner_outer * 1024 + cse_var_24 + cse_var_23 + 514 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer * 3 // 2 < 127): -#CHECK-NEXT: cse_var_27: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_26: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_25: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_27 + cse_var_26 + 4 -#CHECK-NEXT: C_1[cse_var_25:cse_var_25 + 2] = C_1[cse_var_25:cse_var_25 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 1], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_27 + cse_var_26 + 516:k_outer * 8192 + k_inner_outer * 1024 + cse_var_27 + cse_var_26 + 516 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer * 3 // 2 < 128): -#CHECK-NEXT: cse_var_30: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_29: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_28: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_30 + cse_var_29 + 512 -#CHECK-NEXT: C_1[cse_var_28:cse_var_28 + 2] = C_1[cse_var_28:cse_var_28 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 17], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_30 + cse_var_29 + 512:k_outer * 8192 + k_inner_outer * 1024 + cse_var_30 + cse_var_29 + 512 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + (j_inner_outer * 3 + 1) // 2 < 128): -#CHECK-NEXT: cse_var_33: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_32: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_31: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_33 + cse_var_32 + 514 -#CHECK-NEXT: C_1[cse_var_31:cse_var_31 + 2] = C_1[cse_var_31:cse_var_31 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 17], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_33 + cse_var_32 + 514:k_outer * 8192 + k_inner_outer * 1024 + cse_var_33 + cse_var_32 + 514 + 2] -#CHECK-NEXT: if T.likely(j_outer * 9 + j_inner_outer * 3 // 2 < 127): -#CHECK-NEXT: cse_var_36: T.int32 = j_outer * 36 -#CHECK-NEXT: cse_var_35: T.int32 = j_inner_outer * 6 -#CHECK-NEXT: cse_var_34: T.int32 = i_outer * 65536 + i_inner_outer * 1024 + cse_var_36 + cse_var_35 + 516 -#CHECK-NEXT: C_1[cse_var_34:cse_var_34 + 2] = C_1[cse_var_34:cse_var_34 + 2] + T.Broadcast(_0_global_1[i_inner_outer * 32 + k_inner_outer * 2 + 17], 2) * _1_1[k_outer * 8192 + k_inner_outer * 1024 + cse_var_36 + cse_var_35 + 516:k_outer * 8192 + k_inner_outer * 1024 + cse_var_36 + cse_var_35 + 516 + 2] -#CHECK-NEXT:CODE: 0 +# CHECK: graph: +# CHECK-NEXT: name: matmul +# CHECK-NEXT: inputs: +# CHECK-NEXT: - %0 : 512x512xfloat32 +# CHECK-NEXT: - %1 : 512x512xfloat32 +# CHECK-NEXT: outputs: +# CHECK-NEXT: - %2 : 512x512xfloat32 +# CHECK-NEXT: nodes: +# CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [512x512xfloat32, 512x512xfloat32] -> [512x512xfloat32] +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((512, 512), "float32"), _1: T.Buffer((512, 512), "float32"), C: T.Buffer((512, 512), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j, k in T.grid(512, 512, 512): +# CHECK-NEXT: with T.sblock("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_sblock("C") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: I_R0 = sch.cache_read(O, 0, "global") +# CHECK-NEXT: i, i0, i1, = sch.split(i, factors=[None, 64, 2]) +# CHECK-NEXT: j, j0, j1, __v_j1, = sch.split(j, factors=[None, 12, 3, 2]) +# CHECK-NEXT: k, k0, __u_k0, = sch.split(k, factors=[None, 8, 2]) +# CHECK-NEXT: sch.reorder(j, k, i, j0, i0, k0, __u_k0, i1, j1, __v_j1) +# CHECK-NEXT: sch.compute_at(I_R0, i) +# CHECK-NEXT: sch.unroll(__u_k0) +# CHECK-NEXT: sch.unroll(i1) +# CHECK-NEXT: sch.unroll(j1) +# CHECK-NEXT: sch.vectorize(__v_j1) +# CHECK-NEXT: sch.parallel(j) +# CHECK-NEXT: +# CHECK-NEXT: # from tvm.script import ir as I +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis +# CHECK-NEXT: +# CHECK-NEXT: @I.ir_module +# CHECK-NEXT: class Module: +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((512, 512), "float32"), _1: T.Buffer((512, 512), "float32"), C: T.Buffer((512, 512), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: _0_global = T.sblock_alloc_buffer((512, 512)) +# CHECK-NEXT: for j_0 in T.parallel(8): +# CHECK-NEXT: for k_0, i_0 in T.grid(32, 4): +# CHECK-NEXT: for ax0, ax1 in T.grid(128, 16): +# CHECK-NEXT: with T.sblock("_0_global"): +# CHECK-NEXT: v0 = T.axis.spatial(512, i_0 * 128 + ax0) +# CHECK-NEXT: v1 = T.axis.spatial(512, k_0 * 16 + ax1) +# CHECK-NEXT: T.reads(_0[v0, v1]) +# CHECK-NEXT: T.writes(_0_global[v0, v1]) +# CHECK-NEXT: _0_global[v0, v1] = _0[v0, v1] +# CHECK-NEXT: for j_1, i_1, k_1 in T.grid(12, 64, 8): +# CHECK-NEXT: for k_2 in T.unroll(2): +# CHECK-NEXT: for i_2 in T.unroll(2): +# CHECK-NEXT: for j_2 in T.unroll(3): +# CHECK-NEXT: for j_3 in T.vectorized(2): +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i = T.axis.spatial(512, i_0 * 128 + i_1 * 2 + i_2) +# CHECK-NEXT: v_j = T.axis.spatial(512, j_0 * 72 + j_1 * 6 + j_2 * 2 + j_3) +# CHECK-NEXT: v_k = T.axis.reduce(512, k_0 * 16 + k_1 * 2 + k_2) +# CHECK-NEXT: T.where(((j_0 * 12 + j_1) * 3 + j_2) * 2 + j_3 < 512) +# CHECK-NEXT: T.reads(_0_global[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_global[v_i, v_k] * _1[v_k, v_j] +# CHECK-NEXT: CODE: 0 diff --git a/tests/filecheck/schedules/test_matmul_descript_tvm.py b/tests/filecheck/schedules/test_matmul_descript_tvm.py index c484775ae..75fb65ed9 100644 --- a/tests/filecheck/schedules/test_matmul_descript_tvm.py +++ b/tests/filecheck/schedules/test_matmul_descript_tvm.py @@ -55,51 +55,51 @@ # CHECK-NEXT: - %2: matmul(%0, %1) {name = 'C'} : [4x512xfloat32, 512x32xfloat32] -> [4x32xfloat32] # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: for i, j in T.grid(4, 32): -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: C_1[i * 32 + j] = T.float32(0.0) -# CHECK-NEXT: for k in range(512): -# CHECK-NEXT: cse_var_1: T.int32 = i * 32 + j -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_1] = C_1[cse_var_1] + _0_1[i * 512 + k] * _1_1[k * 32 + j] -# CHECK-NEXT: O = obj['C'] -# CHECK-NEXT: I, J, = O.op.axis -# CHECK-NEXT: K, = O.op.reduce_axis -# CHECK-NEXT: I, I0 = sch[O].split(I, factor=2) -# CHECK-NEXT: J, J0 = sch[O].split(J, factor=16) -# CHECK-NEXT: sch[O].reorder(K, I, J, I0, J0) -# CHECK-NEXT: sch[O].unroll(I0) -# CHECK-NEXT: sch[O].vectorize(J0) +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for i, j, k in T.grid(4, 32, 512): +# CHECK-NEXT: with T.sblock("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_sblock("C") +# CHECK-NEXT: I, J, K, = sch.get_loops(O) +# CHECK-NEXT: I, I0, = sch.split(I, factors=[None, 2]) +# CHECK-NEXT: J, J0, = sch.split(J, factors=[None, 16]) +# CHECK-NEXT: sch.reorder(K, I, J, I0, J0) +# CHECK-NEXT: sch.unroll(I0) +# CHECK-NEXT: sch.vectorize(J0) # CHECK-NEXT: # CHECK-NEXT: # from tvm.script import ir as I -# CHECK-NEXT: # from tvm.script import tir as T +# CHECK-NEXT: # from tvm.script import tirx as T +# CHECK-NEXT: # from tvm.tirx.layout import Axis # CHECK-NEXT: # CHECK-NEXT: @I.ir_module # CHECK-NEXT: class Module: -# CHECK-NEXT: @T.prim_func -# CHECK-NEXT: def main(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): -# CHECK-NEXT: T.func_attr({"from_legacy_te_schedule": T.bool(True), "tir.noalias": T.bool(True)}) -# CHECK-NEXT: C_1 = T.Buffer((128,), data=C.data) -# CHECK-NEXT: for i_outer_init, j_outer_init in T.grid(2, 2): -# CHECK-NEXT: cse_var_1: T.int32 = i_outer_init * 64 + j_outer_init * 16 -# CHECK-NEXT: C_1[cse_var_1:cse_var_1 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: C_1[cse_var_1 + 32:cse_var_1 + 32 + 16] = T.Broadcast(T.float32(0.0), 16) -# CHECK-NEXT: for k, i_outer, j_outer in T.grid(512, 2, 2): -# CHECK-NEXT: cse_var_6: T.int32 = j_outer * 16 -# CHECK-NEXT: cse_var_5: T.int32 = i_outer * 1024 + k -# CHECK-NEXT: cse_var_4: T.int32 = k * 32 + cse_var_6 -# CHECK-NEXT: cse_var_3: T.int32 = i_outer * 64 + cse_var_6 -# CHECK-NEXT: cse_var_2: T.int32 = cse_var_3 + 32 -# CHECK-NEXT: _0_1 = T.Buffer((2048,), data=_0.data) -# CHECK-NEXT: _1_1 = T.Buffer((16384,), data=_1.data) -# CHECK-NEXT: C_1[cse_var_3:cse_var_3 + 16] = C_1[cse_var_3:cse_var_3 + 16] + T.Broadcast(_0_1[cse_var_5], 16) * _1_1[cse_var_4:cse_var_4 + 16] -# CHECK-NEXT: C_1[cse_var_2:cse_var_2 + 16] = C_1[cse_var_2:cse_var_2 + 16] + T.Broadcast(_0_1[cse_var_5 + 512], 16) * _1_1[cse_var_4:cse_var_4 + 16] +# CHECK-NEXT: @T.prim_func(s_tir=True) +# CHECK-NEXT: def matmul(_0: T.Buffer((4, 512), "float32"), _1: T.Buffer((512, 32), "float32"), C: T.Buffer((4, 32), "float32")): +# CHECK-NEXT: T.func_attr({"tirx.noalias": True}) +# CHECK-NEXT: # with T.sblock("root"): +# CHECK-NEXT: for k, i_0, j_0 in T.grid(512, 2, 2): +# CHECK-NEXT: for i_1 in T.unroll(2): +# CHECK-NEXT: for j_1 in T.vectorized(16): +# CHECK-NEXT: with T.sblock("C"): +# CHECK-NEXT: v_i = T.axis.spatial(4, i_0 * 2 + i_1) +# CHECK-NEXT: v_j = T.axis.spatial(32, j_0 * 16 + j_1) +# CHECK-NEXT: v_k = T.axis.reduce(512, 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: CODE: 0 diff --git a/tests/filecheck/search/test_conv_ppwrprp.py b/tests/filecheck/search/test_conv_ppwrprp.py index 108f9a436..6db391bd5 100644 --- a/tests/filecheck/search/test_conv_ppwrprp.py +++ b/tests/filecheck/search/test_conv_ppwrprp.py @@ -14,142 +14,92 @@ utils.print_exhaustive_samples(backend, strategy, 200) # CHECK: schedule O0: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=1) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=1) -# CHECK-NEXT: b1, b2 = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h2 = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w2 = sch[O].split(w1, factor=1) -# CHECK-NEXT: f1, f2 = sch[O].split(f1, factor=1) -# CHECK-NEXT: r, r1 = sch[O].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O].split(f2, factor=1) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O].unroll(w3) -# CHECK-NEXT: sch[O].unroll(h3) -# CHECK-NEXT: sch[O].unroll(b3) -# CHECK-NEXT: sch[O].unroll(c1) -# CHECK-NEXT: sch[O].unroll(s1) -# CHECK-NEXT: sch[O].unroll(r1) -# CHECK-NEXT: sch[O].vectorize(f3) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 1, 1, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 1, 1, 1]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O1: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=1) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=1) -# CHECK-NEXT: b1, b2 = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h2 = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w2 = sch[O].split(w1, factor=1) -# CHECK-NEXT: f1, f2 = sch[O].split(f1, factor=1) -# CHECK-NEXT: r, r1 = sch[O].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O].split(f2, factor=1) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O].unroll(w3) -# CHECK-NEXT: sch[O].unroll(h3) -# CHECK-NEXT: sch[O].unroll(b3) -# CHECK-NEXT: sch[O].unroll(c1) -# CHECK-NEXT: sch[O].unroll(s1) -# CHECK-NEXT: sch[O].unroll(r1) -# CHECK-NEXT: sch[O].vectorize(f3) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 1, 1, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 1, 1, 1]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O2: [1, 1, 1, 1, 1, 1, 1, 1, 2, 1, 1, 16, 1, 1, 1, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=2) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=16) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=2) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=16) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=2) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 2, 1, 2]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 16, 1, 16]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O3: [1, 1, 1, 1, 1, 1, 1, 1, 2, 1, 1, 16, 1, 1, 3, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=2) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=16) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=2) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=16) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=3) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=2) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 2, 1, 2]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 16, 1, 16]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 3]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: sample 0: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0] # CHECK-NEXT: sample 1: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1] @@ -352,40 +302,24 @@ # CHECK-NEXT: sample 198: [1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 1, 16, 1, 1, 3, 0] # CHECK-NEXT: sample 199: [1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 1, 16, 1, 1, 3, 1] # CHECK-NEXT: stats {'filtered': 200, 'all': 404} -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=1) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=32) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=1) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=16) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=3) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 1, 1, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 32, 1, 16]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 3]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) diff --git a/tests/filecheck/search/test_conv_ppwrprpv.py b/tests/filecheck/search/test_conv_ppwrprpv.py index c37c68096..e437eaca0 100644 --- a/tests/filecheck/search/test_conv_ppwrprpv.py +++ b/tests/filecheck/search/test_conv_ppwrprpv.py @@ -14,142 +14,92 @@ utils.print_exhaustive_samples(backend, strategy, 200) # CHECK: schedule O0: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=1) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=1) -# CHECK-NEXT: b1, b2 = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h2 = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w2 = sch[O].split(w1, factor=1) -# CHECK-NEXT: f1, f2 = sch[O].split(f1, factor=1) -# CHECK-NEXT: r, r1 = sch[O].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O].split(f2, factor=1) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O].unroll(w3) -# CHECK-NEXT: sch[O].unroll(h3) -# CHECK-NEXT: sch[O].unroll(b3) -# CHECK-NEXT: sch[O].unroll(c1) -# CHECK-NEXT: sch[O].unroll(s1) -# CHECK-NEXT: sch[O].unroll(r1) -# CHECK-NEXT: sch[O].vectorize(f3) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 1, 1, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 1, 1, 1]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O1: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=1) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=1) -# CHECK-NEXT: b1, b2 = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h2 = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w2 = sch[O].split(w1, factor=1) -# CHECK-NEXT: f1, f2 = sch[O].split(f1, factor=1) -# CHECK-NEXT: r, r1 = sch[O].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O].split(f2, factor=1) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O].unroll(w3) -# CHECK-NEXT: sch[O].unroll(h3) -# CHECK-NEXT: sch[O].unroll(b3) -# CHECK-NEXT: sch[O].unroll(c1) -# CHECK-NEXT: sch[O].unroll(s1) -# CHECK-NEXT: sch[O].unroll(r1) -# CHECK-NEXT: sch[O].vectorize(f3) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 1, 1, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 1, 1, 1]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O2: [1, 1, 1, 1, 1, 1, 1, 1, 2, 1, 1, 16, 1, 1, 1, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=2) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=16) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=2) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=16) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=2) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 2, 1, 2]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 16, 1, 16]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O3: [1, 1, 1, 1, 1, 1, 1, 1, 2, 1, 1, 16, 1, 1, 3, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=2) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=16) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=2) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=16) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=3) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=2) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 2, 1, 2]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 16, 1, 16]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 3]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: sample 0: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 16, 1, 1, 1, 0] # CHECK-NEXT: sample 1: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 16, 1, 1, 1, 1] @@ -352,40 +302,24 @@ # CHECK-NEXT: sample 198: [1, 1, 1, 1, 2, 1, 1, 2, 1, 1, 1, 32, 1, 1, 1, 0] # CHECK-NEXT: sample 199: [1, 1, 1, 1, 2, 1, 1, 2, 1, 1, 1, 32, 1, 1, 1, 1] # CHECK-NEXT: stats {'filtered_vec': 200, 'filtered': 3040, 'all': 9042} -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=2) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=2) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=32) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=2) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=2) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=32) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=32) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 2, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 1, 2, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 32, 1, 32]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) diff --git a/tests/filecheck/search/test_conv_ppwrprpvr.py b/tests/filecheck/search/test_conv_ppwrprpvr.py index b69a39474..8f6646df7 100644 --- a/tests/filecheck/search/test_conv_ppwrprpvr.py +++ b/tests/filecheck/search/test_conv_ppwrprpvr.py @@ -21,142 +21,92 @@ utils.print_exhaustive_samples(backend, strategy, 200) # CHECK: schedule O0: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=1) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=1) -# CHECK-NEXT: b1, b2 = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h2 = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w2 = sch[O].split(w1, factor=1) -# CHECK-NEXT: f1, f2 = sch[O].split(f1, factor=1) -# CHECK-NEXT: r, r1 = sch[O].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O].split(f2, factor=1) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O].unroll(w3) -# CHECK-NEXT: sch[O].unroll(h3) -# CHECK-NEXT: sch[O].unroll(b3) -# CHECK-NEXT: sch[O].unroll(c1) -# CHECK-NEXT: sch[O].unroll(s1) -# CHECK-NEXT: sch[O].unroll(r1) -# CHECK-NEXT: sch[O].vectorize(f3) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 1, 1, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 1, 1, 1]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O1: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=1) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=1) -# CHECK-NEXT: b1, b2 = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h2 = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w2 = sch[O].split(w1, factor=1) -# CHECK-NEXT: f1, f2 = sch[O].split(f1, factor=1) -# CHECK-NEXT: r, r1 = sch[O].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O].split(f2, factor=1) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O].unroll(w3) -# CHECK-NEXT: sch[O].unroll(h3) -# CHECK-NEXT: sch[O].unroll(b3) -# CHECK-NEXT: sch[O].unroll(c1) -# CHECK-NEXT: sch[O].unroll(s1) -# CHECK-NEXT: sch[O].unroll(r1) -# CHECK-NEXT: sch[O].vectorize(f3) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 1, 1, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 1, 1, 1]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O2: [1, 1, 1, 1, 1, 1, 1, 1, 2, 1, 1, 16, 1, 1, 1, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=2) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=16) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=2) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=16) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=1) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=2) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 2, 1, 2]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 16, 1, 16]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: schedule O3: [1, 1, 1, 1, 1, 1, 1, 1, 2, 1, 1, 16, 1, 1, 3, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=1) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=2) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=16) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=1) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=2) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=16) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=3) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=1) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=2) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 1, 1, 1]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 2, 1, 2]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 16, 1, 16]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 3]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) # CHECK-NEXT: # CHECK-NEXT: sample 0: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 16, 1, 1, 1, 0] # CHECK-NEXT: sample 1: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 16, 1, 1, 1, 1] @@ -359,40 +309,24 @@ # CHECK-NEXT: sample 198: [1, 1, 1, 1, 1, 2, 2, 1, 1, 1, 1, 32, 1, 1, 3, 0] # CHECK-NEXT: sample 199: [1, 1, 1, 1, 1, 2, 2, 1, 1, 1, 1, 32, 1, 1, 3, 1] # CHECK-NEXT: stats {'filtered_l2': 200, 'filtered_l1': 204, 'filtered_reg': 264, 'filtered_vec': 268, 'filtered': 3836, 'all': 6356} -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: b, h, w, f, = O.op.axis -# CHECK-NEXT: r, s, c, = O.op.reduce_axis -# CHECK-NEXT: b, b1 = sch[O].split(b, factor=1) -# CHECK-NEXT: h, h1 = sch[O].split(h, factor=2) -# CHECK-NEXT: w, w1 = sch[O].split(w, factor=2) -# CHECK-NEXT: f, f1 = sch[O].split(f, factor=32) -# CHECK-NEXT: b1, b_ = sch[O].split(b1, factor=1) -# CHECK-NEXT: h1, h_ = sch[O].split(h1, factor=2) -# CHECK-NEXT: w1, w_ = sch[O].split(w1, factor=1) -# CHECK-NEXT: f1, f_ = sch[O].split(f1, factor=32) -# CHECK-NEXT: sch[O].reorder(b, h, w, f, b1, h1, w1, f1, b_, h_, w_, f_) -# CHECK-NEXT: f = sch[O].fuse(b, h, w, f) -# CHECK-NEXT: sch[O].parallel(f) -# CHECK-NEXT: sch[O_W0].compute_at(sch[O], f1) -# CHECK-NEXT: b, h, w, f, = O_W0.op.axis -# CHECK-NEXT: r, s, c, = O_W0.op.reduce_axis -# CHECK-NEXT: b2 = b -# CHECK-NEXT: h2 = h -# CHECK-NEXT: w2 = w -# CHECK-NEXT: f2 = f -# CHECK-NEXT: r, r1 = sch[O_W0].split(r, factor=1) -# CHECK-NEXT: s, s1 = sch[O_W0].split(s, factor=1) -# CHECK-NEXT: c, c1 = sch[O_W0].split(c, factor=3) -# CHECK-NEXT: b2, b3 = sch[O_W0].split(b2, factor=1) -# CHECK-NEXT: h2, h3 = sch[O_W0].split(h2, factor=2) -# CHECK-NEXT: w2, w3 = sch[O_W0].split(w2, factor=1) -# CHECK-NEXT: f2, f3 = sch[O_W0].split(f2, factor=32) -# CHECK-NEXT: sch[O_W0].reorder(r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) -# CHECK-NEXT: sch[O_W0].unroll(w3) -# CHECK-NEXT: sch[O_W0].unroll(h3) -# CHECK-NEXT: sch[O_W0].unroll(b3) -# CHECK-NEXT: sch[O_W0].unroll(c1) -# CHECK-NEXT: sch[O_W0].unroll(s1) -# CHECK-NEXT: sch[O_W0].unroll(r1) -# CHECK-NEXT: sch[O_W0].vectorize(f3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: b, h, w, f, r, s, c, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: b, b1, b2, b3, = sch.split(b, factors=[None, 1, 1, 1]) +# CHECK-NEXT: h, h1, h2, h3, = sch.split(h, factors=[None, 2, 1, 2]) +# CHECK-NEXT: w, w1, w2, w3, = sch.split(w, factors=[None, 2, 1, 1]) +# CHECK-NEXT: f, f1, f2, f3, = sch.split(f, factors=[None, 32, 1, 32]) +# CHECK-NEXT: r, r1, = sch.split(r, factors=[None, 1]) +# CHECK-NEXT: s, s1, = sch.split(s, factors=[None, 1]) +# CHECK-NEXT: c, c1, = sch.split(c, factors=[None, 3]) +# CHECK-NEXT: sch.reorder(b, h, w, f, b1, h1, w1, f1, r, s, c, b2, h2, w2, f2, r1, s1, c1, b3, h3, w3, f3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, f1) +# CHECK-NEXT: sch.unroll(w3) +# CHECK-NEXT: sch.unroll(h3) +# CHECK-NEXT: sch.unroll(b3) +# CHECK-NEXT: sch.unroll(c1) +# CHECK-NEXT: sch.unroll(s1) +# CHECK-NEXT: sch.unroll(r1) +# CHECK-NEXT: sch.vectorize(f3) +# CHECK-NEXT: f = sch.fuse(b, h, w, f) +# CHECK-NEXT: sch.parallel(f) diff --git a/tests/filecheck/search/test_matmul_ppwrprp.py b/tests/filecheck/search/test_matmul_ppwrprp.py index 0ac408ef4..669b96188 100644 --- a/tests/filecheck/search/test_matmul_ppwrprp.py +++ b/tests/filecheck/search/test_matmul_ppwrprp.py @@ -14,90 +14,60 @@ utils.print_exhaustive_samples(backend, strategy, 200) # CHECK: schedule O0: [1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=1) -# CHECK-NEXT: i1, i2 = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j2 = sch[O].split(j1, factor=1) -# CHECK-NEXT: k, k1 = sch[O].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O].split(j2, factor=1) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O].unroll(i3) -# CHECK-NEXT: sch[O].unroll(k1) -# CHECK-NEXT: sch[O].vectorize(j3) -# CHECK-NEXT: j = sch[O].fuse(i, j) -# CHECK-NEXT: sch[O].parallel(j) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 1, 1, 1]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O1: [1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=1) -# CHECK-NEXT: i1, i2 = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j2 = sch[O].split(j1, factor=1) -# CHECK-NEXT: k, k1 = sch[O].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O].split(j2, factor=1) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O].unroll(i3) -# CHECK-NEXT: sch[O].unroll(k1) -# CHECK-NEXT: sch[O].vectorize(j3) -# CHECK-NEXT: j = sch[O].fuse(i, j) -# CHECK-NEXT: sch[O].parallel(j) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 1, 1, 1]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O2: [1, 1, 1, 1, 1, 16, 1, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=16) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 16, 1, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O3: [1, 1, 3, 1, 1, 16, 12, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=3) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=3) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=16) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=12) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=3) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 3, 1, 3]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 16, 1, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 12]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: sample 0: [1, 1, 1, 1, 1, 1, 1, 0] # CHECK-NEXT: sample 1: [1, 1, 1, 1, 1, 1, 1, 1] @@ -300,26 +270,16 @@ # CHECK-NEXT: sample 198: [1, 1, 1, 1, 32, 1, 1, 0] # CHECK-NEXT: sample 199: [1, 1, 1, 1, 32, 1, 1, 1] # CHECK-NEXT: stats {'filtered': 200, 'all': 242} -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=32) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=32) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=1) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 1, 32, 1]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) diff --git a/tests/filecheck/search/test_matmul_ppwrprpv.py b/tests/filecheck/search/test_matmul_ppwrprpv.py index 6c73f103c..8afff51e0 100644 --- a/tests/filecheck/search/test_matmul_ppwrprpv.py +++ b/tests/filecheck/search/test_matmul_ppwrprpv.py @@ -14,90 +14,60 @@ utils.print_exhaustive_samples(backend, strategy, 200) # CHECK: schedule O0: [1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=1) -# CHECK-NEXT: i1, i2 = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j2 = sch[O].split(j1, factor=1) -# CHECK-NEXT: k, k1 = sch[O].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O].split(j2, factor=1) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O].unroll(i3) -# CHECK-NEXT: sch[O].unroll(k1) -# CHECK-NEXT: sch[O].vectorize(j3) -# CHECK-NEXT: j = sch[O].fuse(i, j) -# CHECK-NEXT: sch[O].parallel(j) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 1, 1, 1]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O1: [1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=1) -# CHECK-NEXT: i1, i2 = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j2 = sch[O].split(j1, factor=1) -# CHECK-NEXT: k, k1 = sch[O].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O].split(j2, factor=1) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O].unroll(i3) -# CHECK-NEXT: sch[O].unroll(k1) -# CHECK-NEXT: sch[O].vectorize(j3) -# CHECK-NEXT: j = sch[O].fuse(i, j) -# CHECK-NEXT: sch[O].parallel(j) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 1, 1, 1]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O2: [1, 1, 1, 1, 1, 16, 1, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=16) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 16, 1, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O3: [1, 1, 3, 1, 1, 16, 12, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=3) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=3) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=16) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=12) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=3) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 3, 1, 3]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 16, 1, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 12]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: sample 0: [1, 1, 1, 1, 1, 16, 1, 0] # CHECK-NEXT: sample 1: [1, 1, 1, 1, 1, 16, 1, 1] @@ -300,26 +270,16 @@ # CHECK-NEXT: sample 198: [3, 1, 1, 1, 1, 16, 4, 0] # CHECK-NEXT: sample 199: [3, 1, 1, 1, 1, 16, 4, 1] # CHECK-NEXT: stats {'filtered_vec': 200, 'filtered': 2944, 'all': 6104} -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=3) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=16) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=4) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 3, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 16, 1, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 4]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) diff --git a/tests/filecheck/search/test_matmul_ppwrprpvr.py b/tests/filecheck/search/test_matmul_ppwrprpvr.py index 39aa0b22a..7f1d1a6ae 100644 --- a/tests/filecheck/search/test_matmul_ppwrprpvr.py +++ b/tests/filecheck/search/test_matmul_ppwrprpvr.py @@ -21,90 +21,60 @@ utils.print_exhaustive_samples(backend, strategy, 200) # CHECK: schedule O0: [1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=1) -# CHECK-NEXT: i1, i2 = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j2 = sch[O].split(j1, factor=1) -# CHECK-NEXT: k, k1 = sch[O].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O].split(j2, factor=1) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O].unroll(i3) -# CHECK-NEXT: sch[O].unroll(k1) -# CHECK-NEXT: sch[O].vectorize(j3) -# CHECK-NEXT: j = sch[O].fuse(i, j) -# CHECK-NEXT: sch[O].parallel(j) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 1, 1, 1]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O1: [1, 1, 1, 1, 1, 1, 1, 0] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=1) -# CHECK-NEXT: i1, i2 = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j2 = sch[O].split(j1, factor=1) -# CHECK-NEXT: k, k1 = sch[O].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O].split(j2, factor=1) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O].unroll(i3) -# CHECK-NEXT: sch[O].unroll(k1) -# CHECK-NEXT: sch[O].vectorize(j3) -# CHECK-NEXT: j = sch[O].fuse(i, j) -# CHECK-NEXT: sch[O].parallel(j) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 1, 1, 1]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O2: [1, 1, 1, 1, 1, 16, 1, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=1) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=1) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=16) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 1, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 16, 1, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: schedule O3: [1, 1, 3, 1, 1, 16, 12, 1] -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=3) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=16) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=3) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=16) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=12) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=3) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 3, 1, 3]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 16, 1, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 12]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) # CHECK-NEXT: # CHECK-NEXT: sample 0: [1, 1, 1, 1, 1, 16, 1, 0] # CHECK-NEXT: sample 1: [1, 1, 1, 1, 1, 16, 1, 1] @@ -307,26 +277,16 @@ # CHECK-NEXT: sample 198: [1, 21, 1, 1, 2, 16, 1, 0] # CHECK-NEXT: sample 199: [1, 21, 1, 1, 2, 16, 1, 1] # CHECK-NEXT: stats {'filtered_l2': 200, 'filtered_l1': 210, 'filtered_reg': 230, 'filtered_vec': 308, 'filtered': 4252, 'all': 5498} -# CHECK-NEXT: O = obj['%2'] -# CHECK-NEXT: O_W0 = sch.cache_write(O, "global") -# CHECK-NEXT: i, j, = O.op.axis -# CHECK-NEXT: k, = O.op.reduce_axis -# CHECK-NEXT: i, i1 = sch[O].split(i, factor=21) -# CHECK-NEXT: j, j1 = sch[O].split(j, factor=32) -# CHECK-NEXT: i1, i_ = sch[O].split(i1, factor=21) -# CHECK-NEXT: j1, j_ = sch[O].split(j1, factor=32) -# CHECK-NEXT: sch[O].reorder(i, j, i1, j1, 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], j1) -# CHECK-NEXT: i, j, = O_W0.op.axis -# CHECK-NEXT: k, = O_W0.op.reduce_axis -# CHECK-NEXT: i2 = i -# CHECK-NEXT: j2 = j -# CHECK-NEXT: k, k1 = sch[O_W0].split(k, factor=1) -# CHECK-NEXT: i2, i3 = sch[O_W0].split(i2, factor=1) -# CHECK-NEXT: j2, j3 = sch[O_W0].split(j2, factor=16) -# CHECK-NEXT: sch[O_W0].reorder(k, i2, j2, k1, i3, j3) -# CHECK-NEXT: sch[O_W0].unroll(i3) -# CHECK-NEXT: sch[O_W0].unroll(k1) -# CHECK-NEXT: sch[O_W0].vectorize(j3) +# CHECK-NEXT: O = sch.get_sblock("%2") +# CHECK-NEXT: i, j, k, = sch.get_loops(O) +# CHECK-NEXT: O_W0 = sch.cache_write(O, 0, "global") +# CHECK-NEXT: i, i1, i2, i3, = sch.split(i, factors=[None, 1, 21, 1]) +# CHECK-NEXT: j, j1, j2, j3, = sch.split(j, factors=[None, 16, 2, 16]) +# CHECK-NEXT: k, k1, = sch.split(k, factors=[None, 1]) +# CHECK-NEXT: sch.reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3) +# CHECK-NEXT: sch.reverse_compute_at(O_W0, j1) +# CHECK-NEXT: sch.unroll(i3) +# CHECK-NEXT: sch.unroll(k1) +# CHECK-NEXT: sch.vectorize(j3) +# CHECK-NEXT: j = sch.fuse(i, j) +# CHECK-NEXT: sch.parallel(j) diff --git a/tests/pytest/tvm/test_tvm_impl.py b/tests/pytest/tvm/test_tvm_impl.py index 50f2edabd..d79b51cee 100644 --- a/tests/pytest/tvm/test_tvm_impl.py +++ b/tests/pytest/tvm/test_tvm_impl.py @@ -20,8 +20,8 @@ def sched_tile2(sch): print(sch) return [ "reorder(i, i1, i2, j, j1, j2, k, k1)", - "split(j, factor=64)", - "split(j1, factor=64)", + "split(i, factors=[None, 16, 4])", + "split(j, factors=[None, 1, 64])", ] def sched_tile2p(sch): @@ -54,11 +54,9 @@ def sched_tile3wc(sch): print(sch) # Expected in TVM schedule return [ - "sch[O].reorder(i, j, i_, j_)", - "sch[O_W0].compute_at(sch[O], j)", - "sch[O_W0].reorder(i1, j1, i_, j_)", - "sch[O_W1].compute_at(sch[O_W0], j1)", - "sch[O_W1].reorder(k, i2, j2, k1, i3, j3)", + "reorder(i, j, i1, j1, k, i2, j2, k1, i3, j3)", + "reverse_compute_at(O_W0, j)", + "reverse_compute_at(O_W0, j1)", ] def sched_tile_unroll_vec(sch): @@ -72,11 +70,11 @@ def sched_tile_unroll_vec(sch): print(sch) # Expected in TVM schedule return [ - "sch[O].reorder(j, k, i, k1, __u_k1, i1, j1, __v_j1)", - "sch[O].unroll(__u_k1)", - "sch[O].unroll(i1)", - "sch[O].unroll(j1)", - "sch[O].vectorize(__v_j1)", + "reorder(j, k, i, k1, __u_k1, i1, j1, __v_j1)", + "unroll(__u_k1)", + "unroll(i1)", + "unroll(j1)", + "vectorize(__v_j1)", ] def check_schedule(impl, sched_func): @@ -144,15 +142,22 @@ def check_compile_evaluate(imp, schedule, compiler_args, evaluate_args): @requires_tvm @pytest.mark.parametrize( - "compiler_args", + "module_type", ( - {"shared_lib": True}, - {"emit_c": True}, - {"ar_lib": True}, - ) + "shared_lib", + "emit_c", + "ar_lib", + ), ) -def test_backend_variant(tmpdir, compiler_args): - impl = matmul_impl(*MATMUL_ARGS, "matmul", emit_c=True) +@pytest.mark.parametrize( + "bare_ptr", + ( + False, + True, + ), +) +def test_backend_variant(tmpdir, module_type, bare_ptr): + impl = matmul_impl(*MATMUL_ARGS, "matmul") print(impl.graph) libpath = Path(tmpdir) / impl.graph.name schedule = check_schedule(impl, sched_tile2p) @@ -161,7 +166,8 @@ def test_backend_variant(tmpdir, compiler_args): schedule, { "dump_file": str(libpath), - **compiler_args, + module_type: True, + "bare_ptr": bare_ptr, }, { "validate": True,