diff --git a/.github/workflows/wafer-isolation.yml b/.github/workflows/wafer-isolation.yml new file mode 100644 index 00000000..41fc0ac1 --- /dev/null +++ b/.github/workflows/wafer-isolation.yml @@ -0,0 +1,46 @@ +name: Wafer isolated build and offline checks + +# Separate, opt-in job: the Ascend workflow and its runner stay independent. +# The configured runner needs the pinned LLVM, Wafer SDK, Python build tools +# and pytest. This workflow submits no work to a Wafer device. +on: + workflow_dispatch: + +jobs: + WaferOffline: + if: vars.WAFER_CI_RUNNER != '' + runs-on: ${{ vars.WAFER_CI_RUNNER }} + env: + LLVM_SYSPATH: ${{ vars.WAFER_CI_LLVM_SYSPATH }} + WAFER_DEPS_ROOT: ${{ vars.WAFER_CI_DEPS_ROOT }} + WAFER_BUILD_DIR: ${{ github.workspace }}/.wafer-build + CMAKE_BUILD_PARALLEL_LEVEL: '4' + DICP_BACKEND: wafer + USE_SIM_MODE: '1' + steps: + - uses: actions/checkout@v4 + with: + submodules: recursive + - name: Build both compiler trees + run: bash scripts/wafer/compile_wafer.sh + - name: Assemble and install an isolated wheel + run: | + python setup_on_wafer.py --build-dir "$WAFER_BUILD_DIR" --wheel-dir "$WAFER_BUILD_DIR/wheel" + python -m venv --system-site-packages "$WAFER_BUILD_DIR/venv" + "$WAFER_BUILD_DIR/venv/bin/python" -m pip install --no-deps --force-reinstall "$WAFER_BUILD_DIR"/wheel/triton-*.whl + - name: Check frontend, patch profiles, tools and loader contracts + run: | + cd "$RUNNER_TEMP" + "$WAFER_BUILD_DIR/venv/bin/python" -m pytest -q \ + "$GITHUB_WORKSPACE/test/wafer/test_frontend_isolation.py" \ + "$GITHUB_WORKSPACE/test/wafer/test_tle_frontend.py" \ + "$GITHUB_WORKSPACE/test/wafer/test_external_interfaces.py" \ + "$GITHUB_WORKSPACE/test/wafer/test_scalar_copy.py" \ + "$GITHUB_WORKSPACE/test/wafer/test_loader_isolation.py" \ + "$GITHUB_WORKSPACE/test/wafer/test_commonir_abi.py" \ + "$GITHUB_WORKSPACE/test/wafer/test_patch_profiles.py" + - uses: actions/upload-artifact@v4 + if: always() + with: + name: wafer-build-identity + path: .wafer-build/wafer-build.json diff --git a/backend/compiler.py b/backend/compiler.py index 0390763e..e438068a 100644 --- a/backend/compiler.py +++ b/backend/compiler.py @@ -126,12 +126,17 @@ def __init__(self, target: str) -> None: elif self.driver.target == "maca": self.capability = 80 self.binary_ext = "mcfatbin" + elif self.driver.target == "wafer": + from triton.backends.dicp_triton.wafer import WaferBackend + + self._wafer_backend = WaferBackend(target) + self.binary_ext = self._wafer_backend.binary_ext else: raise RuntimeError(f"Target '{self.driver.target}' is not supported.") @staticmethod def supports_target(target: GPUTarget): - return target.backend in ["ascend", "mlu", "maca", "cpu"] + return target.backend in ["ascend", "mlu", "maca", "cpu", "wafer"] @staticmethod def make_ttir(mod, metadata, opt): @@ -152,7 +157,9 @@ def make_ttir(mod, metadata, opt): def add_stages(self, stages, options, language=None): if self.driver.is_cpu_verify: return self._cpu_backend.add_stages(stages, options, language) - if self.driver.target == "ascend": + if self.driver.target == "wafer": + return self._wafer_backend.add_stages(stages, options, language) + elif self.driver.target == "ascend": from triton.backends.dicp_triton.npu import ( make_ttir, ttir_to_linalg_dicp, @@ -245,7 +252,9 @@ def add_stages(self, stages, options, language=None): def load_dialects(self, ctx): if self.driver.is_cpu_verify: return self._cpu_backend.load_dialects(ctx) - if self.driver.target == "mlu": + if self.driver.target == "wafer": + return self._wafer_backend.load_dialects(ctx) + elif self.driver.target == "mlu": from triton._C.libtriton import mlu mlu.load_dialects(ctx) @@ -263,7 +272,9 @@ def get_driver(self): def parse_options(self, options: dict) -> Any: if self.driver.is_cpu_verify: return self._cpu_backend.parse_options(options) - if self.target.backend == "ascend": + if self.target.backend == "wafer": + return self._wafer_backend.parse_options(options) + elif self.target.backend == "ascend": from triton.backends.dicp_triton.npu import NPUOptions args = { @@ -340,6 +351,8 @@ def get_codegen_implementation(self, options=None): codegen_fns = dict() if self.driver.is_cpu_verify: return self._cpu_backend.get_codegen_implementation(options) + elif self.target.backend == "wafer": + return self._wafer_backend.get_codegen_implementation(options) elif self.target.backend == "ascend": from triton.backends.dicp_triton.npu import min_dot_size @@ -366,7 +379,9 @@ def get_codegen_implementation(self, options=None): def pack_metadata(self, metadata): if self.driver.is_cpu_verify: return self._cpu_backend.pack_metadata(metadata) - if self.target.backend == "ascend": + if self.target.backend == "wafer": + return self._wafer_backend.pack_metadata(metadata) + elif self.target.backend == "ascend": KERNEL_NAME_MAX_LEN = 49 kernel_name_orig = metadata.kernel_name @@ -395,7 +410,9 @@ def pack_metadata(self, metadata): def hash(self): if self.driver.is_cpu_verify: return self._cpu_backend.hash() - if self.target.backend == "mlu": + if self.target.backend == "wafer": + return self._wafer_backend.hash() + elif self.target.backend == "mlu": from triton.backends.dicp_triton.mlu import get_cnas_version version = get_cnas_version() @@ -405,7 +422,9 @@ def hash(self): return str(version_key) def get_module_map(self) -> Dict[str, ModuleType]: - if self.target.backend == "mlu": + if self.target.backend == "wafer": + return self._wafer_backend.get_module_map() + elif self.target.backend == "mlu": from triton.language.extra.mlu import libdevice return {"triton.language.extra.libdevice": libdevice} diff --git a/backend/driver.py b/backend/driver.py index 098b75ce..48fe7950 100644 --- a/backend/driver.py +++ b/backend/driver.py @@ -170,6 +170,26 @@ def __init__(self, target=None): from .ascend_autotune_hooks import hook_autotune_for_ascend hook_autotune_for_ascend() + elif backend == "wafer": + from .wafer_runtime import ( + SimulatorUtils, + WaferLauncher, + WaferUtils, + get_runtime, + ) + + self.target = "wafer" + if os.getenv("USE_SIM_MODE", "0").lower() in ("1", "true", "yes"): + self.utils = SimulatorUtils() + self.get_current_device = lambda: 0 + self.set_current_device = lambda device: None + else: + self.utils = WaferUtils() + self.get_current_device = lambda: get_runtime().current_device() + self.set_current_device = lambda device: get_runtime().set_device( + device + ) + self.launcher_cls = WaferLauncher elif backend == "nvidia": from triton.backends.nvidia.driver import CudaLauncher, CudaUtils @@ -197,6 +217,8 @@ def is_active(): @classmethod def is_active(self): + if get_current_backend() == "wafer": + return True try: current_backend = get_current_backend() if current_backend == "ascend": @@ -257,12 +279,21 @@ def get_device_capability(self): return ("maca", 0) elif self.target == "ascend": return ("ascend", 0) + elif self.target == "wafer": + return ("wafer", 0) elif self.target == "nvidia": capability = torch.cuda.get_device_capability(self.get_current_device()) return ("cuda", capability) return ("dicp", 0) def get_current_stream(self, device): + if self.target == "wafer": + if os.getenv("USE_SIM_MODE", "0").lower() in ("1", "true", "yes"): + return None + from .wafer_runtime import get_runtime + + stream = get_runtime().current_stream(device) + return None if stream is None else stream.txda_stream import torch if self.target == "mlu": @@ -286,6 +317,8 @@ def get_current_stream(self, device): return None def get_current_device(self): + if self.target == "wafer": + return 0 import torch # dicp doesn't have a device to return. Return something. @@ -338,6 +371,8 @@ def get_current_target(self): arch = self.utils.get_arch() warp_size = 0 return GPUTarget(backend, arch, warp_size) + elif self.target == "wafer": + return GPUTarget("wafer", "wafer", 32) elif self.target == "nvidia": device = self.get_current_device() capability = torch.cuda.get_device_capability(device) @@ -350,7 +385,14 @@ def assemble_tensormap_to_arg(self, tensormaps_info, args): return args def get_device_interface(self): - if self.target == "ascend": + if self.target == "wafer": + from .wafer_runtime import get_runtime + + runtime = get_runtime() + if not hasattr(runtime, "Event") or not hasattr(runtime, "synchronize"): + raise RuntimeError("Wafer benchmarking requires torch_txda Event and synchronize support") + return runtime + elif self.target == "ascend": import torch return torch.npu @@ -362,7 +404,11 @@ def get_device_interface(self): assert False, f"Not implemented for {self.target}" def get_empty_cache_for_benchmark(self): - if self.target == "ascend": + if self.target == "wafer": + # no device cache flush. + # None distinguishes this policy from other backends' zeroable tensor. + return None + elif self.target == "ascend": import torch cache_size = 192 * 1024 * 1024 @@ -383,6 +429,12 @@ def get_active_torch_device(self): if self.is_cpu_verify: return self._cpu_driver.get_active_torch_device() + if self.target == "wafer" and os.getenv("USE_SIM_MODE", "0").lower() not in ( + "1", + "true", + "yes", + ): + return torch.device("txda", self.get_current_device()) return torch.device("cpu") def map_python_to_cpp_type(self, ty: str) -> str: @@ -407,4 +459,5 @@ def map_python_to_cpp_type(self, ty: str) -> str: @classmethod def clear_cache(self, cache): - cache.zero_() + if cache is not None: + cache.zero_() diff --git a/backend/utils.py b/backend/utils.py index 189ad7c2..d6599d6a 100644 --- a/backend/utils.py +++ b/backend/utils.py @@ -157,6 +157,11 @@ def get_current_backend(): global backend if backend is not None: return backend + override = os.getenv("DICP_BACKEND", "").lower() + if override: + if override not in {"ascend", "mlu", "maca", "nvidia", "wafer"}: + raise RuntimeError(f"Unsupported DICP_BACKEND '{override}'.") + backend = override elif command_exists("npu-smi"): backend = "ascend" elif command_exists("cnmon"): @@ -165,6 +170,8 @@ def get_current_backend(): backend = "maca" elif command_exists("nvidia-smi"): backend = "nvidia" + elif command_exists("tsm_smi"): + backend = "wafer" else: backend = None return backend diff --git a/backend/wafer.py b/backend/wafer.py new file mode 100644 index 00000000..697561f4 --- /dev/null +++ b/backend/wafer.py @@ -0,0 +1,619 @@ +import hashlib +import json +import os +import re +import shutil +import subprocess +import tempfile +from dataclasses import dataclass +from pathlib import Path +from types import ModuleType +from typing import Any, Dict, Tuple + +from triton._C.libtriton import ir, passes +from triton.backends.compiler import BaseBackend, GPUTarget + +from .wafer_cache import cache_digest, file_fingerprint + + +@dataclass(frozen=True) +class WaferOptions: + debug: bool = False + arch: str = None + num_warps: int = 0 + num_ctas: int = 0 + num_stages: int = 1 + precision_mode: int = 0 + enable_pipeline: bool = False + num_buffers_warp_spec: int = 0 + num_consumer_groups: int = 0 + reg_dec_producer: int = 0 + reg_inc_consumer: int = 0 + enable_warp_specialization: bool = False + enable_fp_fusion: bool = False + extern_libs: tuple = None + cluster_dims: tuple = (1, 1, 1) + launch_mode: str = "simt" + shared: bool = False + allow_fp8e4nv: bool = False + allowed_dot_input_precisions: Tuple[str, ...] = ("ieee",) + sanitize_overflow: bool = True + max_num_imprecise_acc_default: int = 0 + supported_fp8_dtypes: Tuple[str, ...] = ("fp8e5", "fp8e4b15", "fp8e4nv") + deprecated_fp8_dtypes: Tuple[str, ...] = () + + def __post_init__(self): + if type(self.precision_mode) is not int or self.precision_mode not in (0, 1, 2): + raise ValueError("Wafer precision_mode must be 0, 1 or 2") + if self.launch_mode not in ("simt", "cluster"): + raise ValueError("Wafer launch_mode must be 'simt' or 'cluster'") + if self.launch_mode == "cluster" and tuple(self.cluster_dims) != (1, 1, 1): + raise ValueError("Wafer cluster launch uses one cluster: cluster_dims=(1, 1, 1)") + + def hash(self): + key = "_".join(f"{name}-{value}" for name, value in self.__dict__.items()) + return hashlib.sha256(key.encode("utf-8")).hexdigest() + + +def _run_tool(arguments): + subprocess.check_call( + arguments, + stdout=None if os.getenv("MLIR_ENABLE_DUMP") == "1" else subprocess.DEVNULL, + ) + + +def _dump_file(path): + dump_dir = os.getenv("TRITON_DUMP_PATH") + if dump_dir: + Path(dump_dir).mkdir(parents=True, exist_ok=True) + shutil.copy(path, Path(dump_dir) / Path(path).name) + + +def _find_wafer_opt(): + override = os.getenv("WAFER_OPT_PATH") + if override: + path = Path(override) + if path.is_file(): + return str(path) + raise RuntimeError(f"WAFER_OPT_PATH does not name a file: {path}") + + backend_dir = Path(__file__).resolve().parent + candidates = ( + backend_dir / "bin" / "wafer-opt", + backend_dir.parent + / "third_party" + / "wafer" + / "build_manual" + / "third_party" + / "wafer" + / "bin" + / "wafer-opt", + backend_dir.parent + / "third_party" + / "wafer" + / "build_manual" + / "install" + / "triton" + / "backends" + / "wafer" + / "bin" + / "wafer-opt", + ) + for candidate in candidates: + if candidate.is_file(): + return str(candidate) + path = shutil.which("wafer-opt") + if path: + return path + raise RuntimeError( + "wafer-opt not found; run scripts/wafer/compile_wafer.sh or set WAFER_OPT_PATH" + ) + + +def _find_llvm_tool(name): + llvm_bin = os.getenv("LLVM_BINARY_DIR") + if llvm_bin: + candidate = Path(llvm_bin) / name + if candidate.is_file(): + return str(candidate) + path = shutil.which(name) + if path: + return path + raise RuntimeError(f"{name} not found; set LLVM_BINARY_DIR") + + +def _run_wafer_stage(source, arguments, source_name, output_name): + with tempfile.TemporaryDirectory() as tmpdir: + source_path = Path(tmpdir) / source_name + output_path = Path(tmpdir) / output_name + source_path.write_text(str(source), encoding="utf-8") + command = [ + _find_wafer_opt(), + str(source_path), + *arguments, + "-o", + str(output_path), + ] + _run_tool(command) + _dump_file(source_path) + _dump_file(output_path) + return output_path.read_text(encoding="utf-8") + + +def _precision_mode_from_env(): + # Keep the old switch as mode 2; explicit FlagTree-style mode takes priority. + value = os.getenv("PRECISION_MODE") + if value is None: + return 2 if os.getenv("PRECISION_PRIORITY", "0").lower() in ("1", "true", "yes") else 0 + if value not in ("0", "1", "2"): + raise ValueError("PRECISION_MODE must be 0, 1 or 2") + return int(value) + + +def ttir_to_coreir(module, precision_mode=None, enable_pipeline=False, num_stages=1): + if precision_mode is None: + precision_mode = _precision_mode_from_env() + core_to_mk = f"--core-dialects-to-mk=precision-mode={precision_mode}" + return _run_wafer_stage( + module, + [ + "--triton-to-core-dialects", + "--tle-to-mk", + "--dsa-memory-to-core", + "--linalg-tiling", + core_to_mk, + "--linalg-fusion", + "--legalize-tensor-form-loops", + "--one-shot-bufferize", + "--convert-bufferization-to-memref", + "--materialize-strided-linalg-inputs", + *([f"--mk-pipeline=num-stages={num_stages} max-stages=2", + "--mk-loop-bound-canonicalize"] if enable_pipeline else []), + "--cse", + "--canonicalize", + ], + "ttir.mlir", + "coreir.mlir", + ) + + +def coreir_to_wafer_ir(module, enable_pipeline=False): + return _run_wafer_stage( + module, + [ + "--spmd-allocate-shared-memory", + "--expand-strided-metadata", + "--lower-affine", + "--mk-to-wafer", + *(["--wafer-insert-barrier"] if enable_pipeline else []), + "--cse", + ], + "coreir.mlir", + "wafer_ir.mlir", + ) + + +def wafer_ir_to_llir(module, metadata): + with tempfile.TemporaryDirectory() as tmpdir: + source_path = Path(tmpdir) / "wafer_ir.mlir" + llvm_mlir_path = Path(tmpdir) / "llvm.mlir" + llvm_ir_path = Path(tmpdir) / "kernel.ll" + source_path.write_text(str(module), encoding="utf-8") + wafer_arguments = [ + _find_wafer_opt(), + str(source_path), + "--wafer-memref-to-llvm", + "--addr-to-llvm", + "--convert-scf-to-cf", + # Keep log1p for libm: log(1+x) loses tiny inputs and signed zero. + "--convert-math-to-llvm=approximate-log1p=false", + "--convert-math-to-libm", + "--convert-cf-to-llvm", + "--convert-func-to-llvm", + "--expand-strided-metadata", + "--finalize-memref-to-llvm", + "--kernel-arg-buffer", + "--wafer-to-llvm", + "--convert-arith-to-llvm", + "--reconcile-unrealized-casts", + "--canonicalize", + "--export-kernel-symbols", + "-o", + str(llvm_mlir_path), + ] + _run_tool(wafer_arguments) + _run_tool( + [ + _find_llvm_tool("mlir-translate"), + str(llvm_mlir_path), + "--mlir-to-llvmir", + "-o", + str(llvm_ir_path), + ] + ) + llvm_ir = llvm_ir_path.read_text(encoding="utf-8") + names = re.findall(r"define\s+(?:\w+\s+)*@([\w.$]+)\(", llvm_ir) + if names: + metadata["name"] = names[0] + metadata.setdefault("shared", 0) + _dump_file(llvm_mlir_path) + _dump_file(llvm_ir_path) + return llvm_ir + + +def llir_to_object(llvm_ir, metadata, simulator=None): + if simulator is None: + simulator = simulator_enabled() + with tempfile.TemporaryDirectory() as tmpdir: + source_path = Path(tmpdir) / "kernel.ll" + object_path = Path(tmpdir) / "kernel.o" + source_path.write_text(llvm_ir, encoding="utf-8") + compiler = _find_llvm_tool("clang++") + arguments = [ + compiler, + str(source_path), + "-O2", + "-c", + "-fPIC", + "-o", + str(object_path), + ] + if not simulator: + arguments.extend( + ["--target=riscv64-unknown-elf", "-march=rv64imfdc", "-mabi=lp64d"] + ) + _run_tool(arguments) + _dump_file(object_path) + return object_path.read_bytes() + + +def _find_linker_library(linker, name): + output = subprocess.check_output( + [str(linker), "-march=rv64imfdc", "-mabi=lp64d", f"-print-file-name={name}"], + text=True, + ).strip() + path = Path(output) + if output == name or not path.is_file(): + raise RuntimeError(f"{linker} could not locate {name}: {output}") + return path + + +LINK_FLAGS = ( + "-shared", + "-march=rv64imfdc", + "-mabi=lp64d", + "-O2", + "-nostartfiles", + # All libraries are supplied explicitly and included in the cache key. + # Do not let GCC append the original libc after our firmware-adapted copy. + "-nodefaultlibs", + "-Wl,--allow-shlib-undefined", + "-Wl,--no-dynamic-linker", + "-Wl,--gc-sections", + "-Wl,--unique=.rodata.name", +) + +# Kuiper 1.4 firmware renamed the device logging API. Rename references in +# private archive/object copies; preserve the vendor implementation and varargs +# ABI instead of supplying empty logging stubs or modifying the installed SDK. +RCS_LOG_SYMBOLS = { + "tx8_kernel_printf": "rcs_kernel_printf", + "tx8_kernel_vprintf": "rcs_kernel_vprintf", + "tx8_kernel_vsnprintf": "rcs_kernel_vsnprintf", + "tsm_ep_log": "rcs_ep_log", + "_tsm_ep_log": "_rcs_ep_log", +} + +# Keep the CRT assertion bound to the firmware's newlib service. Pulling the +# toolchain's static newlib implementation also pulls unsupported POSIX syscalls. +# Rename only its private archive copy; internal libc references remain paired +# with that implementation, while the CRT's __assert_func stays a firmware import. +FIRMWARE_LIBC_SYMBOLS = {"__assert_func": "__wafer_newlib_assert_func"} + + +def device_log_abi(): + abi = os.getenv("WAFER_DEVICE_LOG_ABI", "wafer") + if abi not in ("wafer", "rcs"): + raise ValueError( + f"Unsupported WAFER_DEVICE_LOG_ABI={abi!r}; expected wafer or rcs" + ) + return abi + + +def _runtime_link_inputs(): + deps_root = os.getenv("WAFER_DEPS_ROOT") + if not deps_root: + raise RuntimeError("WAFER_DEPS_ROOT is not set; source init_wafer_env.sh first.") + wafer_deps_root = Path(deps_root) + toolchain_root = Path( + os.getenv( + "XUANTIE_NAME", wafer_deps_root / "Xuantie-900-gcc-elf-newlib-x86_64-V2.10.2" + ) + ) + linker = toolchain_root / "bin" / "riscv64-unknown-elf-gcc" + wafer_lib_dir = Path( + os.getenv("WAFER_RUNTIME_LIB_DIR", Path(__file__).resolve().parent / "lib") + ) + if not wafer_lib_dir.is_dir(): + wafer_lib_dir = ( + Path(__file__).resolve().parent.parent / "third_party" / "wafer" / "lib" + ) + + required = [linker, wafer_lib_dir, wafer_deps_root / "lib"] + missing = [str(path) for path in required if not path.exists()] + if missing: + raise RuntimeError( + "Wafer runtime link dependencies are missing: " + ", ".join(missing) + ) + libraries = [ + wafer_deps_root / "lib" / name + for name in ("libcommon_util.a", "libinstr_tx81.a", "liblibc_stub.a") + ] + libraries.append(wafer_lib_dir / "libvr.a") + # Xuantie GCC normally supplies libgloss along with libc. Keep that + # existing dependency explicit when using -nodefaultlibs, including it + # in the cache identity instead of relying on GCC's hidden defaults. + libraries.extend( + _find_linker_library(linker, name) + for name in ("libm.a", "libc.a", "libgcc.a", "libgloss.a") + ) + for library in libraries: + if not library.is_file(): + raise RuntimeError(f"Wafer runtime link library is missing: {library}") + return linker, libraries + + +def _link_fingerprint(linker, libraries, log_abi=None): + log_abi = device_log_abi() if log_abi is None else log_abi + result = { + "linker": file_fingerprint(linker), + "ld": file_fingerprint(linker.parent / "riscv64-unknown-elf-ld"), + "flags": LINK_FLAGS, + "archive_groups": [4, len(libraries) - 4], + "libraries": [file_fingerprint(path) for path in libraries], + "device_log_abi": log_abi, + "firmware_libc_symbols": FIRMWARE_LIBC_SYMBOLS, + "objcopy": file_fingerprint(_find_llvm_tool("llvm-objcopy")), + } + if log_abi == "rcs": + result["log_symbols"] = RCS_LOG_SYMBOLS + return result + + +def _adapt_logging_file(source, destination): + _run_tool( + [ + _find_llvm_tool("llvm-objcopy"), + *(f"--redefine-sym={old}={new}" for old, new in RCS_LOG_SYMBOLS.items()), + str(source), + str(destination), + ] + ) + + +def _adapt_logging_libraries(libraries, link_fingerprint): + from triton.runtime.cache import get_cache_manager + + cache = get_cache_manager(cache_digest({"logging_archives": link_fingerprint})) + adapted = [] + # Only Wafer and Wafer archives contain the renamed device APIs. + for index, library in enumerate(libraries[:4]): + name = f"{index}-{library.name}" + path = cache.get_file(name) + if path is None: + with tempfile.TemporaryDirectory() as tmpdir: + output = Path(tmpdir) / name + _adapt_logging_file(library, output) + path = cache.put(output.read_bytes(), name, binary=True) + adapted.append(Path(path)) + return adapted + libraries[4:] + + +def _adapt_firmware_libc(libraries): + from triton.runtime.cache import get_cache_manager + + adapted = list(libraries) + for index, library in enumerate(libraries): + if library.name != "libc.a": + continue + objcopy = _find_llvm_tool("llvm-objcopy") + cache = get_cache_manager(cache_digest({ + "firmware_libc": file_fingerprint(library), + "symbols": FIRMWARE_LIBC_SYMBOLS, + "objcopy": file_fingerprint(objcopy), + })) + path = cache.get_file("libc.a") + if path is None: + with tempfile.TemporaryDirectory() as tmpdir: + output = Path(tmpdir) / "libc.a" + _run_tool([ + objcopy, + *(f"--redefine-sym={old}={new}" for old, new in FIRMWARE_LIBC_SYMBOLS.items()), + str(library), str(output), + ]) + path = cache.put(output.read_bytes(), "libc.a", binary=True) + adapted[index] = Path(path) + return adapted + + +def object_to_binary(obj, metadata, simulator=None, log_abi=None): + if simulator is None: + simulator = simulator_enabled() + if simulator: + raise RuntimeError( + "Wafer simulator linking requires libvr, libtriton_cmodel, " + "libtx8be_op_cmodel, and libneuralcore_qemu; they are not part of the current SDK." + ) + linker, libraries = _runtime_link_inputs() + log_abi = device_log_abi() if log_abi is None else log_abi + link_fingerprint = _link_fingerprint(linker, libraries, log_abi) + + key = cache_digest( + {"object": hashlib.sha256(obj).hexdigest(), "link": link_fingerprint} + ) + from triton.runtime.cache import get_cache_manager + + cache = get_cache_manager(key) + cache_path = cache.get_file("kernel.so") + if cache_path is None: + with tempfile.TemporaryDirectory() as tmpdir: + object_path = Path(tmpdir) / "kernel.o" + binary_path = Path(tmpdir) / "kernel.so" + object_path.write_bytes(obj) + if log_abi == "rcs": + libraries = _adapt_logging_libraries(libraries, link_fingerprint) + adapted_object = Path(tmpdir) / "kernel-rcs.o" + _adapt_logging_file(object_path, adapted_object) + object_path = adapted_object + libraries = _adapt_firmware_libc(libraries) + command = [ + str(linker), + *LINK_FLAGS, + str(object_path), + "-Wl,--start-group", + *(str(path) for path in libraries[:4]), + "-Wl,--end-group", + "-Wl,--start-group", + *(str(path) for path in libraries[4:]), + "-Wl,--end-group", + "-o", + str(binary_path), + ] + cache.put(obj, "kernel.o", binary=True) + cache.put( + json.dumps({"command": command, "inputs": link_fingerprint}, indent=2), + "link.json", + binary=False, + ) + _run_tool(command) + _dump_file(binary_path) + cache_path = cache.put(binary_path.read_bytes(), "kernel.so", binary=True) + + metadata["kernel_path"] = cache_path + metadata["so_key"] = Path(cache_path).parent.name + metadata["device_log_abi"] = log_abi + return Path(cache_path).read_bytes() + + +def runtime_binary_enabled(): + return os.getenv("WAFER_ENABLE_RUNTIME", "0").lower() in ("1", "true", "yes") + + +def simulator_enabled(): + return os.getenv("USE_SIM_MODE", "0").lower() in ("1", "true", "yes") + + +def __getattr__(name): + # Keep explicit legacy imports without advertising duplicate backend classes. + aliases = {"TXDAOptions": WaferOptions, "TXDABackend": WaferBackend} + if name in aliases: + return aliases[name] + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + +class WaferBackend(BaseBackend): + def __init__(self, target): + super().__init__(target) + self.simulator = simulator_enabled() + self.runtime = runtime_binary_enabled() + self.device_log_abi = device_log_abi() + self.precision_mode = _precision_mode_from_env() + self.enable_pipeline = os.getenv("TRITON_PIPELINE", "0").lower() in ("1", "true", "yes") + self.binary_ext = "so" if self.runtime else "o" + + @staticmethod + def supports_target(target: GPUTarget): + return target.backend in ("wafer", "txda") + + def parse_options(self, options: dict) -> Any: + arguments = { + name: options[name] + for name in WaferOptions.__dataclass_fields__ + if name in options + } + arguments.setdefault("arch", self.target.arch) + arguments.setdefault("precision_mode", self.precision_mode) + arguments.setdefault("enable_pipeline", self.enable_pipeline) + return WaferOptions(**arguments) + + def hash(self): + inputs = { + "target": [self.target.backend, self.target.arch, self.target.warp_size], + "simulator": self.simulator, + "runtime": self.runtime, + "precision_mode": self.precision_mode, + "enable_pipeline": self.enable_pipeline, + "tools": [ + file_fingerprint(_find_wafer_opt()), + file_fingerprint(_find_llvm_tool("mlir-translate")), + file_fingerprint(_find_llvm_tool("clang++")), + ], + "source": file_fingerprint(__file__), + } + if self.runtime and not self.simulator: + inputs["link"] = _link_fingerprint( + *_runtime_link_inputs(), self.device_log_abi + ) + return cache_digest(inputs) + + def get_codegen_implementation(self, options): + return {"min_dot_size": lambda lhs_type, rhs_type: (1, 1, 1)} + + def pack_metadata(self, metadata): + return ( + metadata.num_warps, + metadata.num_ctas, + metadata.shared, + metadata.cluster_dims[0], + metadata.cluster_dims[1], + metadata.cluster_dims[2], + ) + + def load_dialects(self, context): + from triton._C.libtriton import wafer + + wafer.load_dialects(context) + wafer.tle.load_dialects(context) + + @staticmethod + def make_ttir(module, metadata, options): + pass_manager = ir.pass_manager(module.context) + pass_manager.enable_debug() + passes.common.add_inliner(pass_manager) + passes.ttir.add_combine(pass_manager) + passes.common.add_canonicalizer(pass_manager) + passes.ttir.add_reorder_broadcast(pass_manager) + passes.common.add_cse(pass_manager) + passes.common.add_licm(pass_manager) + passes.common.add_symbol_dce(pass_manager) + pass_manager.run(module) + metadata.setdefault("shared", 0) + return module + + def add_stages(self, stages, options, language=None): + stages["ttir"] = lambda source, metadata: self.make_ttir( + source, metadata, options + ) + stages["coreir"] = lambda source, metadata: ttir_to_coreir( + source, options.precision_mode, options.enable_pipeline, options.num_stages) + stages["wafer_ir"] = lambda source, metadata: coreir_to_wafer_ir(source, options.enable_pipeline) + stages["llir"] = lambda source, metadata: wafer_ir_to_llir(source, metadata) + if self.runtime: + stages["so"] = lambda source, metadata: object_to_binary( + llir_to_object(source, metadata, self.simulator), + metadata, + self.simulator, + self.device_log_abi, + ) + else: + stages["o"] = lambda source, metadata: llir_to_object( + source, metadata, self.simulator + ) + + def get_module_map(self) -> Dict[str, ModuleType]: + try: + from triton.language.extra.wafer import libdevice + + return {"triton.language.extra.libdevice": libdevice} + except ImportError: + return {} diff --git a/backend/wafer_cache.py b/backend/wafer_cache.py new file mode 100644 index 00000000..63622240 --- /dev/null +++ b/backend/wafer_cache.py @@ -0,0 +1,28 @@ +"""Content fingerprints shared by Wafer's compiler and native launcher caches.""" + +import functools +import hashlib +import json +from pathlib import Path + + +def cache_digest(value): + return hashlib.sha256(json.dumps(value, sort_keys=True).encode()).hexdigest() + + +@functools.lru_cache(maxsize=256) +def _content_digest(path, size, mtime_ns, ctime_ns): + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def file_fingerprint(path): + path = Path(path).resolve() + stat = path.stat() + return ( + str(path), + _content_digest(str(path), stat.st_size, stat.st_mtime_ns, stat.st_ctime_ns), + ) diff --git a/backend/wafer_runtime.py b/backend/wafer_runtime.py new file mode 100644 index 00000000..afecc26a --- /dev/null +++ b/backend/wafer_runtime.py @@ -0,0 +1,486 @@ +import ctypes +import importlib.util +import os +import shutil +import subprocess +import sysconfig +import tempfile +import weakref +from functools import lru_cache +from pathlib import Path +from types import SimpleNamespace + +from triton.runtime.cache import get_cache_manager + +from .wafer_cache import cache_digest, file_fingerprint + + +def _sdk_path(name): + root = os.getenv("KUIPER_ROOT") + if not root: + raise RuntimeError("KUIPER_ROOT is not set; source init_wafer_env.sh first.") + return os.path.join(root, name) + + +class _KuiperRuntime: + def __init__(self): + self.library = ctypes.CDLL(os.path.join(_sdk_path("lib"), "libhpgr.so")) + self.library.txGetDevice.argtypes = [ctypes.POINTER(ctypes.c_uint32)] + self.library.txGetDevice.restype = ctypes.c_int + self.library.txSetDevice.argtypes = [ctypes.c_uint32] + self.library.txSetDevice.restype = ctypes.c_int + + def current_device(self): + device = ctypes.c_uint32() + status = self.library.txGetDevice(ctypes.byref(device)) + if status != 0: + raise RuntimeError(f"txGetDevice failed with status 0x{status:x}") + return device.value + + def set_device(self, device): + status = self.library.txSetDevice(device) + if status != 0: + raise RuntimeError(f"txSetDevice failed with status 0x{status:x}") + + def current_stream(self, device=None): + # Without torch_txda there is no framework stream context. + return None + + +def get_runtime(): + try: + import torch + import torch_txda # noqa: F401 + + if hasattr(torch, "txda"): + return torch.txda + except (ImportError, AttributeError): + pass + return _KuiperRuntime() + + +def _launcher_compiler(): + compiler = os.getenv("CXX") or shutil.which("clang++") or shutil.which("g++") + if compiler is None: + raise RuntimeError("Failed to find a C++ compiler; set CXX.") + return shutil.which(compiler) or compiler + + +def _build_launcher(name, source, directory): + suffix = sysconfig.get_config_var("EXT_SUFFIX") + output = os.path.join(directory, f"{name}{suffix}") + compiler = _launcher_compiler() + include_dirs = [_sdk_path("include"), sysconfig.get_path("include")] + library_dirs = [_sdk_path("lib")] + libraries = ["hpgr"] + command = [ + compiler, + source, + "-O3", + "-shared", + "-fPIC", + "-std=c++17", + "-Wno-psabi", + "-o", + output, + ] + command += [f"-I{path}" for path in include_dirs] + command += [f"-L{path}" for path in library_dirs] + command += [f"-l{library}" for library in libraries] + subprocess.check_call(command) + return output + + +def _launcher_cache_key(source): + headers = sorted(Path(_sdk_path("include")).rglob("*.h")) + return cache_digest( + { + "source": source, + "compiler": file_fingerprint(_launcher_compiler()), + "python": [ + sysconfig.get_config_var("SOABI"), + sysconfig.get_config_var("EXT_SUFFIX"), + file_fingerprint(Path(sysconfig.get_path("include")) / "Python.h"), + file_fingerprint(sysconfig.get_config_h_filename()), + ], + "sdk_headers": [file_fingerprint(path) for path in headers], + "runtime": file_fingerprint(_sdk_path("lib/libhpgr.so")), + } + ) + + +def compile_launcher(source): + name = "__triton_launcher" + cache = get_cache_manager(_launcher_cache_key(source)) + cache_path = cache.get_file(f"{name}.so") + if cache_path is None: + with tempfile.TemporaryDirectory() as directory: + source_path = os.path.join(directory, f"{name}.cpp") + Path(source_path).write_text(source, encoding="utf-8") + shared_object = _build_launcher(name, source_path, directory) + cache_path = cache.put( + Path(shared_object).read_bytes(), f"{name}.so", binary=True + ) + spec = importlib.util.spec_from_file_location(name, cache_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def _cpp_type(type_name): + if type_name.startswith("*"): + return "PyObject*" + if type_name in ("fp16", "bf16"): + raise NotImplementedError( + "Wafer fp16/bf16 scalar packing is not implemented; pass an fp32 scalar " + "and cast inside the kernel. fp16/bf16 tensor pointers are supported." + ) + return { + "i1": "int32_t", + # Parse signed narrow values as C int, then copy their low bytes into + # the 64-bit slot. KernelArgBufferPass loads the declared scalar width. + "i8": "int32_t", + "i16": "int32_t", + "i32": "int32_t", + "i64": "int64_t", + "u1": "uint32_t", + "u8": "uint8_t", + "u16": "uint16_t", + "u32": "uint32_t", + "u64": "uint64_t", + "fp32": "float", + "f32": "float", + "fp64": "double", + }[type_name] + + +def _parse_format(type_name): + if type_name.startswith("*"): + return "O" + return { + "int8_t": "b", + "int16_t": "h", + "int32_t": "i", + "int64_t": "L", + "uint8_t": "B", + "uint16_t": "H", + "uint32_t": "I", + "uint64_t": "K", + "float": "f", + "double": "d", + }[_cpp_type(type_name)] + + +def make_launcher(signature, launch_mode="simt"): + if launch_mode not in ("simt", "cluster"): + raise ValueError(f"Unknown Wafer launch mode: {launch_mode}") + cluster_check = "" + launch_function = "txLaunchKernelGGL" + cluster_argument = "" + if launch_mode == "cluster": + launch_function = "txLaunchClusterKernelGGL" + cluster_argument = "dim3({1, 1, 1}), " + cluster_check = ''' + if (grid_y != 1 || grid_z != 1 || grid_x > 16) { + PyErr_SetString(PyExc_ValueError, "Wafer cluster grid must be (1..16, 1, 1)"); return NULL; + } +''' + declarations = " ".join( + f"{_cpp_type(type_name)} arg{index};" for index, type_name in signature.items() + ) + parse_format = "iiiOKOOOO" + "".join( + _parse_format(type_name) for type_name in signature.values() + ) + parse_args = "".join(f", &arg{index}" for index in signature) + pointer_setup = "\n".join( + f"void *ptr{index} = get_pointer(arg{index}); if (PyErr_Occurred()) return NULL;" + for index, type_name in signature.items() + if type_name.startswith("*") + ) + kernel_args = "\n".join( + ( + f"runtime_args.push_back(1); runtime_args.push_back((uint64_t)ptr{index});" + if type_name.startswith("*") + else f"uint64_t scalar{index} = 0; memcpy(&scalar{index}, &arg{index}, sizeof(arg{index})); runtime_args.push_back(scalar{index});" + ) + for index, type_name in signature.items() + ) + return f""" +#define PY_SSIZE_T_CLEAN +#include +#include +#include +#include +#include +#include +#include "tx_runtime.h" + +static void *get_pointer(PyObject *object) {{ + if (object == Py_None) return nullptr; + if (PyLong_Check(object)) return PyLong_AsVoidPtr(object); + PyObject *value = PyObject_CallMethod(object, "data_ptr", nullptr); + if (!value) return nullptr; + void *pointer = PyLong_AsVoidPtr(value); + Py_DECREF(value); + return pointer; +}} + +static PyObject *launch(PyObject *, PyObject *args) {{ + int grid_x, grid_y, grid_z; + unsigned long long function; + PyObject *stream_object, *kernel_metadata, *launch_metadata, *enter_hook, *exit_hook; + {declarations} + if (!PyArg_ParseTuple(args, "{parse_format}", &grid_x, &grid_y, &grid_z, &stream_object, &function, + &kernel_metadata, &launch_metadata, &enter_hook, &exit_hook{parse_args})) return NULL; + if (grid_x < 0 || grid_y < 0 || grid_z < 0) {{ + PyErr_SetString(PyExc_ValueError, "Wafer grid dimensions must be nonnegative"); return NULL; + }} + if (grid_x == 0 || grid_y == 0 || grid_z == 0) Py_RETURN_NONE; + {cluster_check} + txStream_t stream = stream_object == Py_None ? nullptr : (txStream_t)PyLong_AsVoidPtr(stream_object); + if (PyErr_Occurred()) return NULL; + if (enter_hook != Py_None) {{ + PyObject *result = PyObject_CallFunctionObjArgs(enter_hook, launch_metadata, NULL); + if (!result) return NULL; + Py_DECREF(result); + }} + {pointer_setup} + std::vector runtime_args; + {kernel_args} + runtime_args.insert(runtime_args.end(), {{(uint64_t)grid_x, (uint64_t)grid_y, (uint64_t)grid_z, 0, 0, 0}}); + PyObject *path_object = PyObject_GetAttrString(kernel_metadata, "kernel_path"); + if (!path_object) return NULL; + PyObject *name_object = PyObject_GetAttrString(kernel_metadata, "name"); + if (!name_object) {{ Py_DECREF(path_object); return NULL; }} + const char *kernel_path = PyUnicode_AsUTF8(path_object); + if (!kernel_path) {{ Py_DECREF(path_object); Py_DECREF(name_object); return NULL; }} + const char *kernel_name = PyUnicode_AsUTF8(name_object); + if (!kernel_name) {{ Py_DECREF(path_object); Py_DECREF(name_object); return NULL; }} + void *binary = nullptr; + size_t size = 0; + if (!function) {{ + FILE *file = fopen(kernel_path, "rb"); + if (!file) {{ PyErr_SetFromErrnoWithFilename(PyExc_OSError, kernel_path); Py_DECREF(path_object); Py_DECREF(name_object); return NULL; }} + long length = -1; + if (fseek(file, 0, SEEK_END) == 0) length = ftell(file); + if (length <= 0 || fseek(file, 0, SEEK_SET) != 0) {{ + PyErr_Format(PyExc_OSError, "Invalid or empty Wafer kernel: %s", kernel_path); + fclose(file); Py_DECREF(path_object); Py_DECREF(name_object); return NULL; + }} + size = (size_t)length; + binary = malloc(size); + if (!binary || fread(binary, 1, size, file) != size) {{ fclose(file); free(binary); Py_DECREF(path_object); Py_DECREF(name_object); PyErr_SetString(PyExc_RuntimeError, "Failed to read Wafer kernel"); return NULL; }} + fclose(file); + }} + txError_t status; + // The argument tuple keeps tensor objects alive while the GIL is released. + // Both paths remain synchronous, including stream errors and module lifetime. + Py_BEGIN_ALLOW_THREADS + if (function) {{ + status = {"txLaunchClusterKernel" if launch_mode == "cluster" else "txLaunchKernel"}( + (txFunction_t)function, {cluster_argument} + dim3({{(uint32_t)grid_x, (uint32_t)grid_y, (uint32_t)grid_z}}), dim3({{1, 1, 1}}), + runtime_args.data(), runtime_args.size() * sizeof(uint64_t), 0, stream); + }} else {{ + status = {launch_function}(kernel_name, (uint64_t)binary, size, {cluster_argument} + dim3({{(uint32_t)grid_x, (uint32_t)grid_y, (uint32_t)grid_z}}), dim3({{1, 1, 1}}), + runtime_args.data(), runtime_args.size() * sizeof(uint64_t), 0, stream); + }} + if (status == TX_SUCCESS) status = txStreamSynchronize(stream); + Py_END_ALLOW_THREADS + free(binary); + if (status != TX_SUCCESS) {{ + PyErr_Format(PyExc_RuntimeError, "Wafer kernel %s (%s) failed with Kuiper status 0x%x, stream=%p", + kernel_name, kernel_path, (unsigned int)status, (void*)stream); + }} + Py_DECREF(path_object); + Py_DECREF(name_object); + if (status != TX_SUCCESS) return NULL; + if (exit_hook != Py_None) {{ + PyObject *result = PyObject_CallFunctionObjArgs(exit_hook, launch_metadata, NULL); + if (!result) return NULL; + Py_DECREF(result); + }} + Py_RETURN_NONE; +}} +static PyMethodDef methods[] = {{{{"launch", launch, METH_VARARGS, "Launch a Wafer kernel"}}, {{NULL, NULL, 0, NULL}}}}; +static struct PyModuleDef module = {{PyModuleDef_HEAD_INIT, "__triton_launcher", NULL, -1, methods}}; +PyMODINIT_FUNC PyInit___triton_launcher(void) {{ return PyModule_Create(&module); }} +""" + + +def __getattr__(name): + aliases = {"TXDAUtils": WaferUtils, "TXDALauncher": WaferLauncher} + if name in aliases: + return aliases[name] + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + +class WaferUtils: + def load_binary(self, name, kernel, shared_mem, device): + api = os.getenv("WAFER_LAUNCH_API", "ggl") + if api == "module": + owner = _LoadedModule(name, kernel, device) + return owner, owner.function, 0, 0, 1024 + if api != "ggl": + raise ValueError("WAFER_LAUNCH_API must be 'ggl' or 'module'") + # Kuiper loads the ELF during launch. Retain the binary as an opaque, + # non-null lifetime token so CompiledKernel initializes only once. + return kernel, 0, 0, 0, 1024 + + def get_device_properties(self, device=None): + return {"max_shared_mem": 3 * 1024 * 1024 - 2 * 0x10000} + + +class _LoadedModule: + """Own one SDK module for CompiledKernel's lifetime, including ELF storage. + + Device launches synchronize before returning, so normal destruction cannot + unload a module with an outstanding launch. The SDK ABI stays unchanged. + """ + def __init__(self, name, binary, device): + runtime = _KuiperRuntime() + library = runtime.library + library.txModuleLoad.argtypes = [ctypes.POINTER(ctypes.c_void_p), ctypes.c_void_p, ctypes.c_uint32] + library.txModuleGetFunction.argtypes = [ctypes.POINTER(ctypes.c_void_p), ctypes.c_void_p, ctypes.c_char_p] + library.txModuleUnload.argtypes = [ctypes.c_void_p] + for operation in (library.txModuleLoad, library.txModuleGetFunction, library.txModuleUnload): + operation.restype = ctypes.c_int + if not binary or len(binary) > 0xFFFFFFFF: + raise ValueError("Wafer module ELF size must fit a nonzero uint32") + self.binary = ctypes.create_string_buffer(binary) + module, function = ctypes.c_void_p(), ctypes.c_void_p() + previous = runtime.current_device() + runtime.set_device(device) + try: + status = library.txModuleLoad(ctypes.byref(module), self.binary, len(binary)) + if status: + raise RuntimeError(f"txModuleLoad({name}) failed with status 0x{status:x}") + self._release = weakref.finalize(self, self._unload, runtime, device, module) + status = library.txModuleGetFunction(ctypes.byref(function), module, name.encode()) + if status or not function.value: + self._release() + raise RuntimeError(f"txModuleGetFunction({name}) failed with status 0x{status:x}") + self.function = function.value + finally: + runtime.set_device(previous) + + @staticmethod + def _unload(runtime, device, module): + previous = runtime.current_device() + runtime.set_device(device) + try: + status = runtime.library.txModuleUnload(module) + if status: + raise RuntimeError(f"txModuleUnload failed with status 0x{status:x}") + finally: + runtime.set_device(previous) + + def close(self): + self._release() + + +class SimulatorUtils: + def load_binary(self, name, kernel, shared_mem, device): + with tempfile.NamedTemporaryFile(mode="wb", suffix=".so", delete=False) as file: + file.write(kernel) + path = file.name + import ctypes + + module = ctypes.CDLL(path) + os.unlink(path) + function = ctypes.cast(getattr(module, name), ctypes.c_void_p).value + return module, function, 0, 0, 1024 + + def get_device_properties(self, device=None): + return {"max_shared_mem": 3 * 1024 * 1024 - 2 * 0x10000} + + +class WaferLauncher: + def __init__(self, src, metadata): + argument_names = getattr(getattr(src, "fn", None), "arg_names", ()) + + def argument_index(key): + key = key[0] if isinstance(key, tuple) else key + return argument_names.index(key) if isinstance(key, str) else int(key) + + signature = dict( + sorted( + (argument_index(index), type_name) + for index, type_name in src.signature.items() + ) + ) + constants = {argument_index(index) for index in getattr(src, "constants", {})} + self.source_argument_count = len(argument_names) or ( + max(signature, default=-1) + 1 + ) + signature = { + index: type_name + for index, type_name in signature.items() + if index not in constants + } + self.runtime_argument_indices = tuple(signature) + self.metadata = metadata + self.launch = compile_launcher( + make_launcher(signature, getattr(metadata, "launch_mode", "simt")) + ).launch + + def __call__(self, *args, **kwargs): + arguments = list(args) + count = len(arguments) - 9 + if count == self.source_argument_count: + # JITFunction passes all bound arguments, including constexpr and + # specialized values. CompiledKernel also accepts runtime-only args. + arguments = arguments[:9] + [ + arguments[9 + index] for index in self.runtime_argument_indices + ] + elif count != len(self.runtime_argument_indices): + raise TypeError( + f"Wafer launcher expected {len(self.runtime_argument_indices)} runtime arguments " + f"or {self.source_argument_count} source arguments, got {count}" + ) + arguments[5] = self.metadata + return self.launch(*arguments, **kwargs) + + +@lru_cache(maxsize=8) +def _noc_initializer(toolchain_key): + from . import wafer + + declarations = { + "module_init": "void @module_init(ptr)", + "module_cleanup": "void @module_cleanup(ptr)", + "__NoCRingInit": "void @__NoCRingInit()", + } + entry = "__wafer_noc_init" + source = "" + for name in (entry, *declarations): + source += (f'@name_{name} = weak constant [{len(name) + 1} x i8] ' + f'c"{name}\\00", section ".rodata.name", align 1\n') + source += (f'@export_{name} = constant {{ptr, ptr}} ' + f'{{ptr @{name}, ptr @name_{name}}}, section "ExportedDYNSYMTab", align 8\n') + source += "\n".join("declare " + declaration for declaration in declarations.values()) + source += (f"\ndefine void @{entry}(ptr %args) {{\n" + " call void @__NoCRingInit()\n ret void\n}\n") + metadata = {"name": entry, "launch_mode": "cluster"} + wafer.object_to_binary(wafer.llir_to_object(source, metadata, simulator=False), + metadata, simulator=False) + src = SimpleNamespace(signature={}) + return WaferLauncher(src, SimpleNamespace(**metadata)) + + +def initialize_noc(stream=None): + """Clear the ring's two sync words on all 16 tiles and wait for completion. + + Call before a NoC collective, after previous device work has completed. + The separate cluster kernel prevents late tile initialization from erasing + a peer's request. Firmware recovery alone does not clear these SPM words. + """ + from triton.backends.compiler import GPUTarget + from .wafer import WaferBackend, runtime_binary_enabled, simulator_enabled + + if simulator_enabled() or not runtime_binary_enabled(): + raise RuntimeError("NoC initialization requires the Wafer hardware runtime") + key = WaferBackend(GPUTarget("wafer", "wafer", 32)).hash() + launcher = _noc_initializer(key) + launcher(16, 1, 1, stream, 0, None, None, None, None) diff --git a/scripts/wafer/apply_triton_profile.py b/scripts/wafer/apply_triton_profile.py new file mode 100644 index 00000000..3e6e572d --- /dev/null +++ b/scripts/wafer/apply_triton_profile.py @@ -0,0 +1,130 @@ +#!/usr/bin/env python3 +"""Apply one pinned patch profile; replacing tracked edits requires --force.""" + +import argparse +import hashlib +import json +import os +from pathlib import Path +import shlex +import subprocess +import sys +import tempfile + +ROOT = Path(__file__).resolve().parents[2] +CATALOG = ROOT / "third_party/wafer/patches/triton/profiles.json" +FORCE_EFFECTS = ( + "--force discards ALL uncommitted changes to tracked files in the selected " + "Triton source, including staged edits, edits outside the patch set and any " + "previously applied profile. No automatic backup is made. Untracked and " + "ignored files are preserved; conflicting files cause an error. " + "The pinned Triton commit is still required." +) + + +def git(source, *args, env=None, data=None): + return subprocess.check_output( + ["git", "-C", str(source), *args], env=env, input=data, stderr=subprocess.PIPE + ) + + +def restore_tracked_source(source, base, tree): + """Reset tracked files only after checking for untracked obstructions.""" + if Path(git(source, "rev-parse", "--show-toplevel").decode().strip()).resolve() != source: + raise RuntimeError("--force requires --source to name the Triton repository root") + # Restore may otherwise overwrite an untracked file left by a staged + # deletion. Check both the base tree and the new profile before any reset. + targets = set() + for revision in (base, tree): + targets.update(os.fsdecode(p) for p in git( + source, "ls-tree", "-r", "--name-only", "-z", revision + ).split(b"\0") if p) + parents = {str(parent) for p in targets for parent in Path(p).parents} + # No --exclude-standard: ignored build files need the same protection. + untracked = [os.fsdecode(p) for p in git(source, "ls-files", "--others", "-z").split(b"\0") if p] + conflicts = [p for p in untracked if p in targets or p in parents or + any(str(parent) in targets for parent in Path(p).parents)] + if conflicts: + raise RuntimeError( + "--force would overwrite untracked or ignored files; no tracked files were reset. " + "Move these files aside before retrying: " + ", ".join(sorted(conflicts)) + ) + print(f"WARNING: {FORCE_EFFECTS}\nSource: {source}", file=sys.stderr) + git(source, "restore", f"--source={base}", "--staged", "--worktree", "--", ".") + + +def apply_profile(source, profile, check=False, force=False): + if check and force: + raise RuntimeError("--check and --force cannot be used together") + source = Path(source).resolve() + catalog = json.loads(CATALOG.read_text()) + base = catalog["triton_commit"] + if git(source, "rev-parse", "HEAD").decode().strip() != base: + raise RuntimeError(f"{profile} requires Triton {base}: {source}") + patches = [ROOT / p for p in catalog["profiles"][profile]] + identity = hashlib.sha256() + for path in patches: + identity.update(path.relative_to(ROOT).as_posix().encode() + b"\0") + identity.update(path.read_bytes()) + + # Validate the entire profile in a temporary index before any destructive + # operation. A bad patch must not discard the caller's existing edits. + with tempfile.TemporaryDirectory(prefix="triton-profile-") as tmp: + env = dict(os.environ, GIT_INDEX_FILE=str(Path(tmp) / "index")) + git(source, "read-tree", base, env=env) + for path in patches: + git(source, "apply", "--cached", "--whitespace=nowarn", str(path), env=env) + tree = git(source, "write-tree", env=env).decode().strip() + paths = git(source, "diff", "--name-only", base, tree).decode().splitlines() + expected = {p: git(source, "show", f"{tree}:{p}") for p in paths} + dirty = set(git(source, "diff", "--name-only", "HEAD").decode().splitlines()) + matches = dirty <= set(paths) and all( + (source / p).is_file() and (source / p).read_bytes() == data + for p, data in expected.items() + ) + if force or not matches: + if check: + raise RuntimeError(f"Source does not match {profile}: {source}") + if dirty and not force: + command = shlex.join([ + sys.executable, str(Path(__file__).resolve()), "--source", str(source), + "--profile", profile, "--force", + ]) + raise RuntimeError( + f"Refusing to replace modified Triton source at {source}; no reset was performed.\n" + f"To replace these changes explicitly, run:\n {command}\n" + f"WARNING: {FORCE_EFFECTS}" + ) + delta = git(source, "diff", "--binary", base, tree) + if force: + restore_tracked_source(source, base, tree) + git(source, "apply", "--check", "-", data=delta) + git(source, "apply", "-", data=delta) + return {"profile": profile, "triton_commit": base, "patch_sha256": identity.hexdigest(), + "source": str(source), "files": { + p: hashlib.sha256(data).hexdigest() for p, data in expected.items() + }} + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--source", type=Path, default=ROOT / "third_party/triton") + parser.add_argument("--profile", choices=json.loads(CATALOG.read_text())["profiles"], required=True) + mode = parser.add_mutually_exclusive_group() + mode.add_argument("--check", action="store_true", help="Verify the profile without changing source files") + mode.add_argument("--force", action="store_true", help=FORCE_EFFECTS) + parser.add_argument("--record", type=Path) + args = parser.parse_args() + try: + result = apply_profile(args.source, args.profile, check=args.check, force=args.force) + except (RuntimeError, subprocess.CalledProcessError) as exc: + detail = exc.stderr.decode(errors="replace") if isinstance(exc, subprocess.CalledProcessError) else str(exc) + raise SystemExit(detail) from None + if args.record: + args.record.parent.mkdir(parents=True, exist_ok=True) + args.record.write_text(json.dumps(result, indent=2) + "\n") + print(f"Triton profile verified: {args.profile} ({result['patch_sha256'][:12]})") + + +if __name__ == "__main__": + main() diff --git a/scripts/wafer/apply_wafer_triton_patches.sh b/scripts/wafer/apply_wafer_triton_patches.sh new file mode 100755 index 00000000..7b25db92 --- /dev/null +++ b/scripts/wafer/apply_wafer_triton_patches.sh @@ -0,0 +1,4 @@ +#!/usr/bin/env bash +set -euo pipefail +SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd) +exec "${PYTHON:-python3}" "$SCRIPT_DIR/apply_triton_profile.py" --profile wafer-tools "$@" diff --git a/scripts/wafer/audit_wafer_elf.py b/scripts/wafer/audit_wafer_elf.py new file mode 100755 index 00000000..61869c03 --- /dev/null +++ b/scripts/wafer/audit_wafer_elf.py @@ -0,0 +1,110 @@ +#!/usr/bin/env python3 +"""Audit a DLCompiler runtime-linked Wafer kernel before submitting device work.""" + +import argparse +import hashlib +import json +import os +from pathlib import Path +import struct +import subprocess + + +NOC_IMPORTS = frozenset({ + "direct_dte_attach", "direct_dte_release", "direct_dte_send_async", + "direct_dte_wait_done", "direct_fsm_monitor_deinit", "direct_fsm_monitor_init", + "direct_fsm_monitor_receive", "get_spm_memory_mapping", "get_tile_spm_addr_base", + "set_direct_fsm_monitor_dst_addr", +}) + + +def audit_kernel(path, log_abi="rcs", noc_firmware_elf=None): + path = Path(path) + data = path.read_bytes() + if len(data) < 64 or data[:6] != b"\x7fELF\x02\x01": + raise ValueError(f"Expected a little-endian ELF64 kernel: {path}") + elf_type, machine = struct.unpack_from("&2 + echo " source ${BASH_SOURCE[0]}" >&2 + exit 1 +fi + +WAFER_WORKSPACE_ROOT=${WAFER_WORKSPACE_ROOT:-$(cd "$(dirname "${BASH_SOURCE[0]}")/../../.." && pwd)} +CONDA_ROOT=${CONDA_ROOT:-$WAFER_WORKSPACE_ROOT/miniconda3} +ENV_NAME=${ENV_NAME:-wafer310} + +if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then + echo "ERROR: Conda not found at $CONDA_ROOT" >&2 + return 1 +fi +if [[ ! -d "$CONDA_ROOT/envs/$ENV_NAME" ]]; then + echo "ERROR: Conda environment not found: $ENV_NAME" >&2 + return 1 +fi + +source "$CONDA_ROOT/etc/profile.d/conda.sh" +conda activate "$ENV_NAME" + +export WAFER_WORKSPACE_ROOT +export CONDA_ROOT +export ENV_NAME +export REPO_DIR=${REPO_DIR:-$WAFER_WORKSPACE_ROOT/DLCompiler} +export DEPS_ROOT=${DEPS_ROOT:-$WAFER_WORKSPACE_ROOT/deps} +export PACKAGE_ROOT=${PACKAGE_ROOT:-$WAFER_WORKSPACE_ROOT/packages} +export PYTHON="$CONDA_PREFIX/bin/python" +export LLVM_COMMIT=${LLVM_COMMIT:-7d5de3033187c8a3bb4d2e322f5462cdaf49808f} +export LLVM_SYSPATH=${LLVM_SYSPATH:-$DEPS_ROOT/llvm-7d5de303-ubuntu-x64} +export LLVM_BINARY_DIR="$LLVM_SYSPATH/bin" +export LLVM_DIR="$LLVM_SYSPATH/lib/cmake/llvm" +export MLIR_DIR="$LLVM_SYSPATH/lib/cmake/mlir" +export WAFER_DEPS_ROOT=${WAFER_DEPS_ROOT:-$DEPS_ROOT/wafer_deps} +export KUIPER_ROOT=${KUIPER_ROOT:-/usr/local/kuiper} +export WAFER_SDK_INCLUDE_DIR=${WAFER_SDK_INCLUDE_DIR:-$WAFER_DEPS_ROOT/include} +export WAFER_RT_THREAD_SMP_ROOT=${WAFER_RT_THREAD_SMP_ROOT:-$WAFER_DEPS_ROOT/tx8-yoc-rt-thread-smp} +export XUANTIE_NAME=${XUANTIE_NAME:-$WAFER_DEPS_ROOT/Xuantie-900-gcc-elf-newlib-x86_64-V2.10.2} +export WAFER_BUILD_DIR=${WAFER_BUILD_DIR:-$WAFER_WORKSPACE_ROOT/build/wafer} +# Installed wheels include libvr.a; only override that default when this +# workspace also has a freshly built hardware CRT. +if [[ -z ${WAFER_RUNTIME_LIB_DIR:-} ]]; then + for runtime_dir in "$WAFER_BUILD_DIR/tools/third_party/wafer/crt/lib" "$WAFER_BUILD_DIR/third_party/wafer/crt/lib"; do + if [[ -f "$runtime_dir/libvr.a" ]]; then + export WAFER_RUNTIME_LIB_DIR="$runtime_dir" + break + fi + done +fi +export DICP_BACKEND=wafer +export USE_SIM_MODE=${USE_SIM_MODE:-1} +# This workspace uses Kuiper 1.4 firmware with the RCS device logging API. +export WAFER_DEVICE_LOG_ABI=${WAFER_DEVICE_LOG_ABI:-rcs} +export PATH="$LLVM_BINARY_DIR:$PATH" +export LD_LIBRARY_PATH="$KUIPER_ROOT/lib${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}" + +hash -r 2>/dev/null || true + +echo "Wafer build environment activated" +echo " Conda: $CONDA_DEFAULT_ENV" +echo " Python: $PYTHON" +echo " LLVM: $LLVM_SYSPATH" +echo " SDK: $WAFER_DEPS_ROOT" +echo " Kuiper: $KUIPER_ROOT" diff --git a/scripts/wafer/install_wafer.sh b/scripts/wafer/install_wafer.sh new file mode 100755 index 00000000..a6bf545b --- /dev/null +++ b/scripts/wafer/install_wafer.sh @@ -0,0 +1,88 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd) +BUILD_DIR=${WAFER_BUILD_DIR:-$SCRIPT_DIR/third_party/wafer/build_manual} +WHEEL_DIR=${WAFER_WHEEL_DIR:-$BUILD_DIR/wheel} +PYTHON=${PYTHON:-python3} +SKIP_BUILD=0 + +if [[ ${1:-} == "--skip-build" ]]; then + SKIP_BUILD=1 +elif [[ -n ${1:-} ]]; then + echo "Usage: $0 [--skip-build]" >&2 + exit 2 +fi + +if [[ -f "$SCRIPT_DIR/wafer_env.sh" ]]; then + source "$SCRIPT_DIR/wafer_env.sh" +fi + +: "${LLVM_SYSPATH:?Run 'bash scripts/wafer/setup_wafer_env.sh' before installing Wafer}" +: "${LLVM_BINARY_DIR:?LLVM_BINARY_DIR is not set}" + +if [[ $SKIP_BUILD == 0 ]]; then + bash "$SCRIPT_DIR/scripts/wafer/compile_wafer.sh" +fi + +for artifact in "$BUILD_DIR/wafer-build.json"; do + if [[ ! -f "$artifact" ]]; then + echo "ERROR: required build artifact not found: $artifact" >&2 + exit 1 + fi +done + +rm -rf "$WHEEL_DIR" +"$PYTHON" "$SCRIPT_DIR/setup_on_wafer.py" \ + --build-dir "$BUILD_DIR" \ + --wheel-dir "$WHEEL_DIR" + +wheel=$(find "$WHEEL_DIR" -maxdepth 1 -name 'triton-*.whl' -type f -print -quit) +if [[ -z "$wheel" ]]; then + echo "ERROR: Wafer wheel was not produced under $WHEEL_DIR" >&2 + exit 1 +fi +"$PYTHON" -m pip install --no-index --no-deps --force-reinstall "$wheel" + +( + cd /tmp + DICP_BACKEND=wafer USE_SIM_MODE=1 LLVM_BINARY_DIR="$LLVM_BINARY_DIR" \ + "$PYTHON" - <<'PY' +import tempfile +from pathlib import Path + +import triton +from triton._C import libtriton +from triton.backends import backends +from triton.backends.compiler import GPUTarget + +if list(backends) != ["dicp_triton"]: + raise RuntimeError(f"Unexpected Triton backends: {list(backends)}") + +target = GPUTarget("wafer", "wafer", 32) +backend = backends["dicp_triton"].compiler(target) +backend.load_dialects(libtriton.ir.context()) + +import triton.language.extra.wafer # noqa: F401, E402 +import triton.experimental.tle.language # noqa: F401, E402 + +if hasattr(libtriton, "dicp_triton"): + raise RuntimeError("The Wafer-only package unexpectedly contains the original DICP C++ binding") + +with tempfile.TemporaryDirectory() as tmpdir: + source = Path(tmpdir) / "wafer_install_check.ttir" + source.write_text( + 'module { tt.func public @wafer_install_check() ' + 'attributes {tt.kernel = 1 : i1} { tt.return } }' + ) + kernel = triton.compile(str(source), target=target) + if not kernel.asm["o"].startswith(b"\x7fELF"): + raise RuntimeError("Wafer compiler did not produce an ELF object") + print("Wafer compiler installation verified:") + print(" triton:", triton.__file__) + print(" libtriton:", libtriton.__file__) + print(" stages:", list(kernel.asm)) + print(" object bytes:", len(kernel.asm["o"])) +PY +) diff --git a/scripts/wafer/inventory_wafer_tests.py b/scripts/wafer/inventory_wafer_tests.py new file mode 100644 index 00000000..88d1baef --- /dev/null +++ b/scripts/wafer/inventory_wafer_tests.py @@ -0,0 +1,134 @@ +#!/usr/bin/env python3 +"""Inventory repository tests statically, without importing or running them.""" + +import argparse +import ast +import csv +import json +from pathlib import Path +import subprocess + + +REPO = Path(__file__).resolve().parents[2] + + +def test_functions(tree): + functions = [] + for node in tree.body: + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name.startswith("test"): + functions.append(node) + elif isinstance(node, ast.ClassDef) and node.name.startswith("Test"): + functions.extend( + child for child in node.body + if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef)) and child.name.startswith("test") + ) + return functions + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output-dir", required=True, type=Path) + parser.add_argument("--execution-summary", type=Path, + help="Optional run_wafer_example_suite.py summary to join with the static inventory") + args = parser.parse_args() + args.output_dir.mkdir(parents=True, exist_ok=True) + # Include new files before the final stage is committed and all project test + # trees (including python/dlBLAS), while respecting ignored build outputs. + files = sorted(set(subprocess.check_output( + ["git", "ls-files", "--cached", "--others", "--exclude-standard"], + cwd=REPO, text=True, + ).splitlines())) + config = ast.parse((REPO / "third_party/wafer/examples/conftest.py").read_text()) + ignored = set() + for node in config.body: + if isinstance(node, ast.Assign) and any(isinstance(t, ast.Name) and t.id == "collect_ignore" for t in node.targets): + ignored.update(ast.literal_eval(node.value)) + executions = {} + if args.execution_summary: + run = json.loads(args.execution_summary.read_text()) + for result in run['files']: + name = result['file'] + if not name.startswith(('test/', 'third_party/')): + name = 'third_party/wafer/examples/' + name + executions[name] = run.get('execution', run.get('suite', 'unknown')) + ':' + result['status'] + rows = [] + for name in sorted(files): + path = Path(name) + if path.suffix != ".py" or not (path.name.startswith("test_") or path.name.endswith("_test.py")): + continue + text = (REPO / path).read_text() + tree = ast.parse(text) + tests = test_functions(tree) + if name.startswith("test/wafer/"): + group, status = "wafer_regression", executions.get(name, "execution_not_inferred_by_static_scan") + elif name.startswith("third_party/wafer/examples/"): + group, status = "wafer_examples", executions.get(name, "execution_not_inferred_by_static_scan") + elif name.startswith("third_party/wafer/third_party/"): + group, status = "bundled_dependency_examples", "not_executed" + else: + group, status = "/".join(path.parts[:2]), "not_executed_for_wafer" + imports = sorted({ + node.module for node in ast.walk(tree) + if isinstance(node, ast.ImportFrom) and node.module and + any(part in node.module for part in ("ztc", "cuda", "npu", "testing", "benchmark")) + }) + assertions = sum( + isinstance(node, ast.Assert) or ( + isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute) + and node.func.attr.startswith("assert") + ) for node in ast.walk(tree) + ) + rows.append({ + "file": name, + "group": group, + "execution_scope": status, + "static_test_functions": len(tests), + "test_names": ",".join(node.name for node in tests), + "ignored_by_wafer_examples_conftest": group == "wafer_examples" and path.name in ignored, + "assertion_sites_in_file": assertions, + "special_imports": ",".join(imports), + "has_main_entry": any(isinstance(node, ast.If) and "__name__" in ast.unparse(node.test) for node in tree.body), + }) + upstream = {} + for directory, paths in { + "third_party/triton": ("test", "python/test"), + "third_party/ascendnpu-ir": ("bishengir/test",), + }.items(): + tracked = subprocess.check_output( + ["git", "ls-files", "--", *paths], cwd=REPO / directory, text=True + ).splitlines() + upstream[directory] = { + "tracked_files_in_test_trees": len(tracked), + "mlir_fixtures": sum(name.endswith(".mlir") for name in tracked), + "python_test_files": sum(Path(name).name.startswith("test_") and name.endswith(".py") for name in tracked), + "execution_scope": "not_executed_as_upstream_suites", + } + summary = { + "source_commit": subprocess.check_output(["git", "rev-parse", "HEAD"], cwd=REPO, text=True).strip(), + "method": "Static AST inventory; function counts do not expand parametrization or prove execution.", + "file_scope": "Project Python test_*.py/*_test.py files from git ls-files; upstream submodules counted separately.", + "execution_summary": str(args.execution_summary) if args.execution_summary else None, + "groups": { + group: { + "files": sum(row["group"] == group for row in rows), + "static_test_functions": sum(row["static_test_functions"] for row in rows if row["group"] == group), + } for group in sorted({row["group"] for row in rows}) + }, + "wafer_examples_collect_ignore": sorted(ignored), + "wafer_examples_existing_ignored_files": sum(row["ignored_by_wafer_examples_conftest"] for row in rows), + "wafer_examples_files_without_detected_assertions": [ + row["file"] for row in rows if row["group"] == "wafer_examples" and not row["assertion_sites_in_file"] + ], + "wafer_mlir_fixtures": [name for name in files if name.startswith("third_party/wafer/") and name.endswith(".mlir")], + "upstream_submodules": upstream, + } + with (args.output_dir / "repository-tests.tsv").open("w", newline="") as stream: + writer = csv.DictWriter(stream, fieldnames=list(rows[0]), delimiter="\t", lineterminator="\n") + writer.writeheader() + writer.writerows(rows) + (args.output_dir / "test-inventory.json").write_text(json.dumps(summary, indent=2) + "\n") + print(json.dumps(summary, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/wafer/migrate_wafer_env.sh b/scripts/wafer/migrate_wafer_env.sh new file mode 100755 index 00000000..61d4a948 --- /dev/null +++ b/scripts/wafer/migrate_wafer_env.sh @@ -0,0 +1,26 @@ +#!/usr/bin/env bash +set -euo pipefail + +SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd) +COMPILER_ONLY=0 +case "${1:-}" in + --compiler-only) COMPILER_ONLY=1 ;; + -h|--help) + printf '%s\n' 'Usage: bash scripts/wafer/migrate_wafer_env.sh [--compiler-only]' \ + 'Prepare LLVM_SYSPATH and WAFER_DEPS_ROOT locally before running.' \ + 'Full mode also requires an existing matching Torch and Kuiper environment.' \ + 'This command checks dependencies and writes wafer_env.sh; it does not install dependencies.' + exit 0 ;; + '') ;; + *) printf 'ERROR: unknown argument: %s\n' "$1" >&2; exit 2 ;; +esac +if [[ $# -gt 1 ]]; then + printf '%s\n' 'ERROR: too many arguments' >&2 + exit 2 +fi +: "${LLVM_SYSPATH:?Provide the local LLVM directory}" +: "${WAFER_DEPS_ROOT:?Provide the local Wafer SDK directory}" +if [[ $COMPILER_ONLY == 0 ]]; then + "${PYTHON:-python3}" "$SCRIPT_DIR/test/wafer/verify_wafer_torch_stack.py" +fi +bash "$SCRIPT_DIR/scripts/wafer/setup_wafer_env.sh" \ No newline at end of file diff --git a/scripts/wafer/package_wafer.py b/scripts/wafer/package_wafer.py new file mode 100644 index 00000000..d5b538e4 --- /dev/null +++ b/scripts/wafer/package_wafer.py @@ -0,0 +1,78 @@ +"""Assemble a Wafer-only wheel from pinned Python sources and audited binaries.""" + +import hashlib +import json +from pathlib import Path +import shutil +import sysconfig + + +def prepare_package(repo, staging, manifest_path, manifest, version, revision): + expected_abi = manifest.get("python_soabi") + if expected_abi != sysconfig.get_config_var("SOABI"): + raise RuntimeError(f"Frontend Python ABI {expected_abi!r} differs from packager ABI") + source = Path(manifest["frontend"]["source"]) + package = staging / "triton" + ignored = shutil.ignore_patterns("__pycache__", "*.pyc", "*.so", "*.a", "*.o") + shutil.copytree(source / "python/triton", package, ignore=ignored) + # Keep the Python DICP dispatch entry point, but no Ascend language package + # or original DICP C++ plugin. Vendor imports remain behind target branches. + backend = package / "backends/dicp_triton" + shutil.copytree(repo / "backend", backend, ignore=shutil.ignore_patterns( + "__pycache__", "*.pyc", "*.so", "*.a", "*.o", "dicp_opt", "bin")) + language = package / "language/extra" + for name in ("wafer", "txda"): + # txda only re-exports the Wafer language API for existing callers. + shutil.copytree(repo / "third_party/wafer/language" / name, language / name, ignore=ignored) + shutil.copytree(repo / "third_party/wafer/experimental/tle", package / "experimental/tle", ignore=ignored) + (package / "_C").mkdir(exist_ok=True) + (package / "_C/__init__.py").touch() + destinations = { + "libtriton.so": package / "_C/libtriton.so", + "FileCheck": package / "_C/FileCheck", + "wafer-opt": backend / "bin/wafer-opt", + "libvr.a": backend / "lib/libvr.a", + } + for name, destination in destinations.items(): + destination.parent.mkdir(parents=True, exist_ok=True) + shutil.copy2(manifest["artifacts"][name]["path"], destination) + init = package / "__init__.py" + init.write_text(init.read_text().replace("__version__ = '3.5.0'", f"__version__ = {version!r}")) + identity = { + "schema": 1, "version": version, "dlcompiler_commit": revision, + "package_backends": ["wafer"], "python_entry_point": "dicp_triton", + "build": manifest, + "python_sources": { + str(p.relative_to(package)): hashlib.sha256(p.read_bytes()).hexdigest() + for p in sorted(package.rglob("*.py")) + }, + } + (backend / "wafer-package.json").write_text(json.dumps(identity, indent=2) + "\n") + shutil.copy2(manifest_path, backend / "bin/wafer-build.json") + shutil.copy2(source / "LICENSE", staging / "LICENSE.triton") + for name in ("LICENSE", "LICENSE.txt"): + if (repo / name).exists(): + shutil.copy2(repo / name, staging / "LICENSE.dlcompiler") + break + (staging / "setup.py").write_text('''from pathlib import Path +from setuptools import Distribution, find_namespace_packages, setup + +class BinaryDistribution(Distribution): + def has_ext_modules(self): + return True + +root = Path("triton") +setup( + name="triton", version=VERSION, + description="DLCompiler Wafer-only Triton compiler", + packages=find_namespace_packages(include=["triton", "triton.*"]), + package_data={"triton": [str(p.relative_to(root)) for p in root.rglob("*") if p.is_file()]}, + include_package_data=False, + distclass=BinaryDistribution, + python_requires=">=3.10", + install_requires=["setuptools>=40.8.0", "pybind11>=2.13.1"], + entry_points={"triton.backends": ["dicp_triton = triton.backends.dicp_triton"]}, + license_files=["LICENSE.*"], +) +'''.replace('version=VERSION', 'version=' + repr(version))) + return identity diff --git a/scripts/wafer/run_wafer_example_suite.py b/scripts/wafer/run_wafer_example_suite.py new file mode 100644 index 00000000..67336952 --- /dev/null +++ b/scripts/wafer/run_wafer_example_suite.py @@ -0,0 +1,131 @@ +#!/usr/bin/env python3 +"""Run native TXDA tests in isolated processes, keeping exact nodeids and launch evidence.""" +import argparse +from collections import Counter, defaultdict +import hashlib +import json +import os +from pathlib import Path +import signal +import subprocess +import sys +import time + +REPO = Path(__file__).resolve().parents[2] +EXAMPLES = REPO / 'third_party/wafer/examples' +MANIFEST = REPO / 'test/wafer/suites/accepted.txt' + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--output-dir', required=True, type=Path) + parser.add_argument('--suite', choices=('accepted', 'examples', 'ops', 'runtime', 'native_math', 'host'), default='accepted') + parser.add_argument('--select', nargs='+', help='Paths relative to the suite root (repository root for accepted)') + parser.add_argument('--nodeids-file', type=Path, + help='With --suite accepted, use an explicit nodeid list (e.g. remaining parameters)') + parser.add_argument('--timeout', type=int, default=1800, help='Per-file timeout; any timeout stops device scheduling') + args = parser.parse_args() + if args.nodeids_file and args.suite != 'accepted': + parser.error('--nodeids-file requires --suite accepted') + manifest = args.nodeids_file.resolve() if args.nodeids_file else MANIFEST + output = args.output_dir.resolve() + output.mkdir(parents=True, exist_ok=True) + # Never merge stale evidence from another source state into a fresh run. + if (output / 'summary.json').exists(): + parser.error('Output already contains a run; choose a new directory') + accepted = defaultdict(list) + for node in manifest.read_text().splitlines(): + accepted[node.split('::', 1)[0]].append(node) + root = {'accepted': REPO, 'examples': EXAMPLES, 'host': REPO / 'test/wafer'}.get(args.suite, + REPO / 'test/wafer' / args.suite) + if args.suite == 'accepted': + files = [REPO / name for name in accepted] + else: + files = sorted(root.glob('test_*.py') if args.suite == 'host' else root.rglob('test_*.py')) + files = [p for p in files if p.name != 'test_common.py'] + if args.select: + requested = set(args.select) + found = {str(p.relative_to(root)) for p in files} + if requested - found: + parser.error(f'Unknown files: {sorted(requested - found)}') + files = [p for p in files if str(p.relative_to(root)) in requested] + files.sort(key=lambda p: (2 if '/tle/' in str(p) else 1 if p.name == 'test_dot_scaled.py' else 0, str(p))) + environment = os.environ.copy() + environment['PYTHONPATH'] = os.pathsep.join(filter(None, (str(REPO / 'scripts/wafer'), environment.get('PYTHONPATH')))) + if args.suite == 'host': + # Host link tests import the source backend but link the delivered CRT. + from triton.backends.dicp_triton import wafer + environment.setdefault('WAFER_RUNTIME_LIB_DIR', str(Path(wafer.__file__).parent / 'lib')) + if args.suite != 'host': + for key, expected in dict(DICP_BACKEND='wafer', USE_SIM_MODE='0', WAFER_ENABLE_RUNTIME='1').items(): + if environment.get(key) != expected: + parser.error(f'Activate the Wafer environment first: requires {key}={expected}') + try: + revision = subprocess.check_output(['git', 'rev-parse', 'HEAD'], cwd=REPO, text=True).strip() + except subprocess.CalledProcessError: + revision = None + summary = dict(suite=args.suite, revision=revision, python=sys.executable, + selection=str(manifest), selection_sha256=hashlib.sha256(manifest.read_bytes()).hexdigest(), + files=[], blocked_reason=None) + for index, source in enumerate(files, 1): + relative = str(source.relative_to(REPO)) + directory = output / relative.removesuffix('.py') + directory.mkdir(parents=True, exist_ok=True) + record = dict(file=relative, source_sha256=hashlib.sha256(source.read_bytes()).hexdigest()) + if summary['blocked_reason']: + record.update(status='not_run_after_device_error', reason=summary['blocked_reason']) + else: + events_path = directory / 'events.jsonl' + events_path.write_text('') + environment['WAFER_TEST_EVENTS'] = str(events_path) + # These two accepted files have mode-2 evidence, not mode-0 evidence. + precision = 2 if source.name in ('test_mod.py', 'test_device_print.py') else 0 + environment['PRECISION_MODE'] = str(precision) + command = [sys.executable, '-m', 'pytest', '-q', '-p', 'wafer_pytest', str(source), '--tb=short', '-ra', + f'--junitxml={directory / "junit.xml"}'] + if args.suite != 'host': + command.append('--wafer-hardware') + if args.suite == 'accepted': + selection = directory / 'nodeids.txt' + selection.write_text('\n'.join(accepted[relative]) + '\n') + command.append(f'--wafer-nodeids={selection}') + print(f'[{index}/{len(files)}] {relative}', flush=True) + started = time.monotonic() + timed_out = False + with (directory / 'pytest.log').open('w') as log: + process = subprocess.Popen(command, cwd=REPO.parent, env=environment, + stdout=log, stderr=subprocess.STDOUT, start_new_session=True) + try: + code = process.wait(timeout=args.timeout) + except subprocess.TimeoutExpired: + timed_out = True + os.killpg(process.pid, signal.SIGTERM) + try: + process.wait(timeout=5) + except subprocess.TimeoutExpired: + os.killpg(process.pid, signal.SIGKILL) + process.wait() + code = process.returncode + events = [json.loads(line) for line in events_path.read_text().splitlines() if line.strip()] + tests = [e for e in events if e['event'] == 'test_result'] + counts = Counter(e['outcome'] for e in tests) + launches = Counter(e['event'] for e in events) + collected = next((e['nodeids'] for e in events if e['event'] == 'collection'), []) + completed = {e['nodeid'] for e in tests} + missing = sorted(set(collected) - completed) + record.update(status='passed' if code == 0 and not missing else 'failed', returncode=code, + seconds=round(time.monotonic() - started, 3), precision_mode=precision, + collected=len(collected), outcomes=dict(counts), completed_launches=launches['launch_complete'], + passed_without_launch=[e['nodeid'] for e in tests if e['outcome'] == 'passed' and not e['launches']], + not_completed=missing, command=command, log=str(directory / 'pytest.log')) + if timed_out or code < 0 or code == 3 or launches['device_error'] or launches['launch_start'] != launches['launch_complete']: + summary['blocked_reason'] = f'Incomplete launch, device error or process timeout/crash in {relative}; inspect its log before continuing.' + print(f" {record['status']}: {dict(counts)}, launches={launches['launch_complete']}", flush=True) + summary['files'].append(record) + (directory / 'result.json').write_text(json.dumps(record, indent=2) + '\n') + (output / 'summary.json').write_text(json.dumps(summary, indent=2) + '\n') + return int(any(row['status'] != 'passed' for row in summary['files'])) + + +if __name__ == '__main__': + raise SystemExit(main()) diff --git a/scripts/wafer/setup_llvm22_env.sh b/scripts/wafer/setup_llvm22_env.sh new file mode 100755 index 00000000..cd6736f8 --- /dev/null +++ b/scripts/wafer/setup_llvm22_env.sh @@ -0,0 +1,109 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd) + +LLVM_COMMIT=7d5de3033187c8a3bb4d2e322f5462cdaf49808f +: "${LLVM_SYSPATH:?Provide the local LLVM directory}" +LLVM_ENV_FILE=${LLVM_ENV_FILE:-$SCRIPT_DIR/llvm22_env.sh} +TRITON_LLVM_HASH_FILE=$SCRIPT_DIR/third_party/triton/cmake/llvm-hash.txt + +required_tools=( + FileCheck + ld.lld + llvm-config + llvm-tblgen + mlir-opt + mlir-tblgen + mlir-translate +) + +if [[ $(uname -s) != "Linux" || $(uname -m) != "x86_64" ]]; then + echo "ERROR: the pinned prebuilt package supports Linux x86_64 only" >&2 + exit 1 +fi + +test -f "$TRITON_LLVM_HASH_FILE" || { + echo "ERROR: initialize third_party/triton before preparing LLVM" >&2 + exit 1 +} +TRITON_LLVM_COMMIT=$(tr -d '[:space:]' < "$TRITON_LLVM_HASH_FILE") +if [[ "$TRITON_LLVM_COMMIT" != "$LLVM_COMMIT" ]]; then + echo "ERROR: LLVM pin does not match third_party/triton" >&2 + echo " script: $LLVM_COMMIT" >&2 + echo " triton: $TRITON_LLVM_COMMIT" >&2 + exit 1 +fi + +verify_toolchain() { + local llvm_root=$1 + local tool + + test -d "$llvm_root" || { + echo "ERROR: LLVM directory does not exist: $llvm_root" >&2 + return 1 + } + + for tool in "${required_tools[@]}"; do + test -x "$llvm_root/bin/$tool" || { + echo "ERROR: required LLVM tool is missing: $llvm_root/bin/$tool" >&2 + return 1 + } + done + + local llvm_version + local mlir_version + llvm_version=$($llvm_root/bin/llvm-config --version) + mlir_version=$($llvm_root/bin/mlir-opt --version | sed -n 's/^.*version //p' | head -n 1) + if [[ "$llvm_version" != "22.0.0git" || "$mlir_version" != "22.0.0git" ]]; then + echo "ERROR: expected LLVM/MLIR 22.0.0git, got LLVM=$llvm_version MLIR=$mlir_version" >&2 + return 1 + fi + + grep -q 'PACKAGE_VERSION "22.0.0git"' \ + "$llvm_root/lib/cmake/llvm/LLVMConfigVersion.cmake" || { + echo "ERROR: LLVM CMake package is not version 22.0.0git" >&2 + return 1 + } + grep -q 'PACKAGE_VERSION "22.0.0git"' \ + "$llvm_root/lib/cmake/mlir/MLIRConfigVersion.cmake" || { + echo "ERROR: MLIR CMake package is not version 22.0.0git" >&2 + return 1 + } +} + +verify_toolchain "$LLVM_SYSPATH" + +cat > "$LLVM_ENV_FILE" </dev/null || true + +for tool in llvm-config llvm-tblgen mlir-opt mlir-tblgen mlir-translate; do + resolved=\$(command -v "\$tool" || true) + expected="\$LLVM_BINARY_DIR/\$tool" + if [[ "\$resolved" != "\$expected" ]]; then + echo "ERROR: mixed LLVM toolchain: \$tool resolves to \$resolved, expected \$expected" >&2 + return 1 2>/dev/null || exit 1 + fi +done + +if [[ \$(llvm-config --version) != "22.0.0git" ]]; then + echo "ERROR: LLVM 22 environment activation failed" >&2 + return 1 2>/dev/null || exit 1 +fi +EOF +chmod +x "$LLVM_ENV_FILE" + +printf '\nLLVM/MLIR environment is ready.\n' +printf ' commit: %s\n' "$LLVM_COMMIT" +printf ' root: %s\n' "$LLVM_SYSPATH" +printf ' env: %s\n' "$LLVM_ENV_FILE" +printf '\nActivate with:\n source %q\n' "$LLVM_ENV_FILE" diff --git a/scripts/wafer/setup_wafer_env.sh b/scripts/wafer/setup_wafer_env.sh new file mode 100755 index 00000000..87074a8c --- /dev/null +++ b/scripts/wafer/setup_wafer_env.sh @@ -0,0 +1,44 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd) +LLVM_COMMIT=7d5de3033187c8a3bb4d2e322f5462cdaf49808f +: "${LLVM_SYSPATH:?Provide the local LLVM directory}" +: "${WAFER_DEPS_ROOT:?Provide the local Wafer SDK directory}" +WAFER_SDK_ROOT=$WAFER_DEPS_ROOT +ENV_FILE=${WAFER_ENV_FILE:-$SCRIPT_DIR/wafer_env.sh} + +if [[ ! -x "$LLVM_SYSPATH/bin/llvm-config" ]]; then + echo "ERROR: LLVM 22 package not found at $LLVM_SYSPATH" >&2 + exit 1 +fi +if [[ $("$LLVM_SYSPATH/bin/llvm-config" --version) != "22.0.0git" ]]; then + echo "ERROR: $LLVM_SYSPATH is not the required LLVM 22 package" >&2 + exit 1 +fi +if [[ ! -f "$WAFER_SDK_ROOT/include/instr_def.h" ]]; then + echo "ERROR: Wafer compiler header not found under $WAFER_SDK_ROOT/include" >&2 + exit 1 +fi + +cat >"$ENV_FILE" < {tt.divisibility = 16 : i32}, %arg1: !tt.ptr {tt.divisibility = 16 : i32}) attributes {noinline = false} { + %cst = arith.constant dense<16> : tensor<16x1xi32> + %c16_i32 = arith.constant 16 : i32 + %0 = tt.get_program_id x : i32 + %1 = arith.muli %0, %c16_i32 : i32 + %2 = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> + %3 = tt.splat %1 : i32 -> tensor<16xi32> + %4 = arith.addi %3, %2 : tensor<16xi32> + %5 = tt.expand_dims %4 {axis = 1 : i32} : tensor<16xi32> -> tensor<16x1xi32> + %6 = arith.muli %5, %cst : tensor<16x1xi32> + %7 = tt.splat %arg0 : !tt.ptr -> tensor<16x1x!tt.ptr> + %8 = tt.addptr %7, %6 : tensor<16x1x!tt.ptr>, tensor<16x1xi32> + %9 = tt.expand_dims %2 {axis = 0 : i32} : tensor<16xi32> -> tensor<1x16xi32> + %10 = tt.broadcast %8 : tensor<16x1x!tt.ptr> -> tensor<16x16x!tt.ptr> + %11 = tt.broadcast %9 : tensor<1x16xi32> -> tensor<16x16xi32> + %12 = tt.addptr %10, %11 : tensor<16x16x!tt.ptr>, tensor<16x16xi32> + %13 = tt.load %12 : tensor<16x16x!tt.ptr> + %14:2 = "tt.reduce"(%13, %11) <{axis = 1 : i32}> ({ + ^bb0(%arg2: f32, %arg3: i32, %arg4: f32, %arg5: i32): + %17 = arith.cmpf oeq, %arg2, %arg4 : f32 + %18 = arith.cmpi slt, %arg3, %arg5 : i32 + %19 = arith.andi %17, %18 : i1 + %20 = arith.cmpf ogt, %arg2, %arg4 : f32 + %21 = arith.ori %20, %19 : i1 + %22 = arith.select %21, %arg2, %arg4 : f32 + %23 = arith.select %21, %arg3, %arg5 : i32 + tt.reduce.return %22, %23 : f32, i32 + }) : (tensor<16x16xf32>, tensor<16x16xi32>) -> (tensor<16xf32>, tensor<16xi32>) + %15 = tt.splat %arg1 : !tt.ptr -> tensor<16x!tt.ptr> + %16 = tt.addptr %15, %4 : tensor<16x!tt.ptr>, tensor<16xi32> + tt.store %16, %14#1 : tensor<16x!tt.ptr> + tt.return + } +} diff --git a/test/wafer/ir/interfaces-flip.mlir b/test/wafer/ir/interfaces-flip.mlir new file mode 100644 index 00000000..453d5c44 --- /dev/null +++ b/test/wafer/ir/interfaces-flip.mlir @@ -0,0 +1,48 @@ +module { + tt.func public @flip_kernel(%arg0: !tt.ptr {tt.divisibility = 16 : i32}, %arg1: !tt.ptr {tt.divisibility = 16 : i32}) attributes {noinline = false} { + %cst = arith.constant dense<8> : tensor<64xi32> + %0 = tt.make_range {end = 8 : i32, start = 0 : i32} : tensor<8xi32> + %1 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %2 = arith.muli %1, %cst : tensor<64xi32> + %3 = tt.expand_dims %0 {axis = 0 : i32} : tensor<8xi32> -> tensor<1x8xi32> + %4 = tt.expand_dims %2 {axis = 1 : i32} : tensor<64xi32> -> tensor<64x1xi32> + %5 = tt.broadcast %3 : tensor<1x8xi32> -> tensor<64x8xi32> + %6 = tt.broadcast %4 : tensor<64x1xi32> -> tensor<64x8xi32> + %7 = arith.addi %5, %6 : tensor<64x8xi32> + %8 = tt.splat %arg0 : !tt.ptr -> tensor<64x8x!tt.ptr> + %9 = tt.addptr %8, %7 : tensor<64x8x!tt.ptr>, tensor<64x8xi32> + %10 = tt.load %9 : tensor<64x8x!tt.ptr> + %11 = tt.bitcast %10 : tensor<64x8xf32> -> tensor<64x8xi32> + %12 = tt.reshape %11 : tensor<64x8xi32> -> tensor<64x2x2x2xi32> + %13 = "tt.reduce"(%12) <{axis = 1 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %29 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %29 : i32 + }) : (tensor<64x2x2x2xi32>) -> tensor<64x2x2xi32> + %14 = tt.expand_dims %13 {axis = 1 : i32} : tensor<64x2x2xi32> -> tensor<64x1x2x2xi32> + %15 = tt.broadcast %14 : tensor<64x1x2x2xi32> -> tensor<64x2x2x2xi32> + %16 = arith.xori %12, %15 : tensor<64x2x2x2xi32> + %17 = "tt.reduce"(%16) <{axis = 2 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %29 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %29 : i32 + }) : (tensor<64x2x2x2xi32>) -> tensor<64x2x2xi32> + %18 = tt.expand_dims %17 {axis = 2 : i32} : tensor<64x2x2xi32> -> tensor<64x2x1x2xi32> + %19 = tt.broadcast %18 : tensor<64x2x1x2xi32> -> tensor<64x2x2x2xi32> + %20 = arith.xori %16, %19 : tensor<64x2x2x2xi32> + %21 = "tt.reduce"(%20) <{axis = 3 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %29 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %29 : i32 + }) : (tensor<64x2x2x2xi32>) -> tensor<64x2x2xi32> + %22 = tt.expand_dims %21 {axis = 3 : i32} : tensor<64x2x2xi32> -> tensor<64x2x2x1xi32> + %23 = tt.broadcast %22 : tensor<64x2x2x1xi32> -> tensor<64x2x2x2xi32> + %24 = arith.xori %20, %23 : tensor<64x2x2x2xi32> + %25 = tt.reshape %24 : tensor<64x2x2x2xi32> -> tensor<64x8xi32> + %26 = tt.bitcast %25 : tensor<64x8xi32> -> tensor<64x8xf32> + %27 = tt.splat %arg1 : !tt.ptr -> tensor<64x8x!tt.ptr> + %28 = tt.addptr %27, %7 : tensor<64x8x!tt.ptr>, tensor<64x8xi32> + tt.store %28, %26 : tensor<64x8x!tt.ptr> + tt.return + } +} diff --git a/test/wafer/ir/interfaces-sort.mlir b/test/wafer/ir/interfaces-sort.mlir new file mode 100644 index 00000000..25c384aa --- /dev/null +++ b/test/wafer/ir/interfaces-sort.mlir @@ -0,0 +1,125 @@ +module { + tt.func public @sort_kernel(%arg0: !tt.ptr {tt.divisibility = 16 : i32}, %arg1: !tt.ptr {tt.divisibility = 16 : i32}) attributes {noinline = false} { + %cst = arith.constant dense<8> : tensor<64xi32> + %0 = tt.make_range {end = 8 : i32, start = 0 : i32} : tensor<8xi32> + %1 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %2 = arith.muli %1, %cst : tensor<64xi32> + %3 = tt.expand_dims %0 {axis = 0 : i32} : tensor<8xi32> -> tensor<1x8xi32> + %4 = tt.expand_dims %2 {axis = 1 : i32} : tensor<64xi32> -> tensor<64x1xi32> + %5 = tt.broadcast %3 : tensor<1x8xi32> -> tensor<64x8xi32> + %6 = tt.broadcast %4 : tensor<64x1xi32> -> tensor<64x8xi32> + %7 = arith.addi %5, %6 : tensor<64x8xi32> + %8 = tt.splat %arg0 : !tt.ptr -> tensor<64x8x!tt.ptr> + %9 = tt.addptr %8, %7 : tensor<64x8x!tt.ptr>, tensor<64x8xi32> + %10 = tt.load %9 : tensor<64x8x!tt.ptr> + %11 = tt.reshape %10 : tensor<64x8xf32> -> tensor<2x2x2x2x2x2x2x2x2xf32> + %12 = tt.make_range {end = 2 : i32, start = 0 : i32} : tensor<2xi32> + %13 = tt.reshape %12 : tensor<2xi32> -> tensor<1x1x1x1x1x1x1x2x1xi32> + %14 = tt.bitcast %11 : tensor<2x2x2x2x2x2x2x2x2xf32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %15 = "tt.reduce"(%14) <{axis = 8 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %94 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %94 : i32 + }) : (tensor<2x2x2x2x2x2x2x2x2xi32>) -> tensor<2x2x2x2x2x2x2x2xi32> + %16 = tt.expand_dims %15 {axis = 8 : i32} : tensor<2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x1xi32> + %17 = tt.broadcast %16 : tensor<2x2x2x2x2x2x2x2x1xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %18 = arith.xori %14, %17 : tensor<2x2x2x2x2x2x2x2x2xi32> + %19 = tt.bitcast %18 : tensor<2x2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xf32> + %20 = tt.reshape %12 : tensor<2xi32> -> tensor<1x1x1x1x1x1x1x1x2xi32> + %21 = arith.cmpf ogt, %11, %19 : tensor<2x2x2x2x2x2x2x2x2xf32> + %22 = tt.broadcast %13 : tensor<1x1x1x1x1x1x1x2x1xi32> -> tensor<1x1x1x1x1x1x1x2x2xi32> + %23 = tt.broadcast %20 : tensor<1x1x1x1x1x1x1x1x2xi32> -> tensor<1x1x1x1x1x1x1x2x2xi32> + %24 = arith.xori %22, %23 : tensor<1x1x1x1x1x1x1x2x2xi32> + %25 = arith.extui %21 : tensor<2x2x2x2x2x2x2x2x2xi1> to tensor<2x2x2x2x2x2x2x2x2xi32> + %26 = tt.broadcast %24 : tensor<1x1x1x1x1x1x1x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %27 = arith.cmpi ne, %25, %26 : tensor<2x2x2x2x2x2x2x2x2xi32> + %28 = arith.select %27, %19, %11 : tensor<2x2x2x2x2x2x2x2x2xi1>, tensor<2x2x2x2x2x2x2x2x2xf32> + %29 = tt.reshape %12 : tensor<2xi32> -> tensor<1x1x1x1x1x1x2x1x1xi32> + %30 = tt.bitcast %28 : tensor<2x2x2x2x2x2x2x2x2xf32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %31 = "tt.reduce"(%30) <{axis = 7 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %94 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %94 : i32 + }) : (tensor<2x2x2x2x2x2x2x2x2xi32>) -> tensor<2x2x2x2x2x2x2x2xi32> + %32 = tt.expand_dims %31 {axis = 7 : i32} : tensor<2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x1x2xi32> + %33 = tt.broadcast %32 : tensor<2x2x2x2x2x2x2x1x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %34 = arith.xori %30, %33 : tensor<2x2x2x2x2x2x2x2x2xi32> + %35 = tt.bitcast %34 : tensor<2x2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xf32> + %36 = arith.cmpf ogt, %28, %35 : tensor<2x2x2x2x2x2x2x2x2xf32> + %37 = tt.broadcast %29 : tensor<1x1x1x1x1x1x2x1x1xi32> -> tensor<1x1x1x1x1x1x2x2x1xi32> + %38 = tt.broadcast %13 : tensor<1x1x1x1x1x1x1x2x1xi32> -> tensor<1x1x1x1x1x1x2x2x1xi32> + %39 = arith.xori %37, %38 : tensor<1x1x1x1x1x1x2x2x1xi32> + %40 = arith.extui %36 : tensor<2x2x2x2x2x2x2x2x2xi1> to tensor<2x2x2x2x2x2x2x2x2xi32> + %41 = tt.broadcast %39 : tensor<1x1x1x1x1x1x2x2x1xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %42 = arith.cmpi ne, %40, %41 : tensor<2x2x2x2x2x2x2x2x2xi32> + %43 = arith.select %42, %35, %28 : tensor<2x2x2x2x2x2x2x2x2xi1>, tensor<2x2x2x2x2x2x2x2x2xf32> + %44 = tt.bitcast %43 : tensor<2x2x2x2x2x2x2x2x2xf32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %45 = "tt.reduce"(%44) <{axis = 8 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %94 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %94 : i32 + }) : (tensor<2x2x2x2x2x2x2x2x2xi32>) -> tensor<2x2x2x2x2x2x2x2xi32> + %46 = tt.expand_dims %45 {axis = 8 : i32} : tensor<2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x1xi32> + %47 = tt.broadcast %46 : tensor<2x2x2x2x2x2x2x2x1xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %48 = arith.xori %44, %47 : tensor<2x2x2x2x2x2x2x2x2xi32> + %49 = tt.bitcast %48 : tensor<2x2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xf32> + %50 = arith.cmpf ogt, %43, %49 : tensor<2x2x2x2x2x2x2x2x2xf32> + %51 = tt.broadcast %29 : tensor<1x1x1x1x1x1x2x1x1xi32> -> tensor<1x1x1x1x1x1x2x1x2xi32> + %52 = tt.broadcast %20 : tensor<1x1x1x1x1x1x1x1x2xi32> -> tensor<1x1x1x1x1x1x2x1x2xi32> + %53 = arith.xori %51, %52 : tensor<1x1x1x1x1x1x2x1x2xi32> + %54 = arith.extui %50 : tensor<2x2x2x2x2x2x2x2x2xi1> to tensor<2x2x2x2x2x2x2x2x2xi32> + %55 = tt.broadcast %53 : tensor<1x1x1x1x1x1x2x1x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %56 = arith.cmpi ne, %54, %55 : tensor<2x2x2x2x2x2x2x2x2xi32> + %57 = arith.select %56, %49, %43 : tensor<2x2x2x2x2x2x2x2x2xi1>, tensor<2x2x2x2x2x2x2x2x2xf32> + %58 = tt.bitcast %57 : tensor<2x2x2x2x2x2x2x2x2xf32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %59 = "tt.reduce"(%58) <{axis = 6 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %94 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %94 : i32 + }) : (tensor<2x2x2x2x2x2x2x2x2xi32>) -> tensor<2x2x2x2x2x2x2x2xi32> + %60 = tt.expand_dims %59 {axis = 6 : i32} : tensor<2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x1x2x2xi32> + %61 = tt.broadcast %60 : tensor<2x2x2x2x2x2x1x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %62 = arith.xori %58, %61 : tensor<2x2x2x2x2x2x2x2x2xi32> + %63 = tt.bitcast %62 : tensor<2x2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xf32> + %64 = arith.cmpf ogt, %57, %63 : tensor<2x2x2x2x2x2x2x2x2xf32> + %65 = arith.extui %64 : tensor<2x2x2x2x2x2x2x2x2xi1> to tensor<2x2x2x2x2x2x2x2x2xi32> + %66 = tt.broadcast %29 : tensor<1x1x1x1x1x1x2x1x1xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %67 = arith.cmpi ne, %65, %66 : tensor<2x2x2x2x2x2x2x2x2xi32> + %68 = arith.select %67, %63, %57 : tensor<2x2x2x2x2x2x2x2x2xi1>, tensor<2x2x2x2x2x2x2x2x2xf32> + %69 = tt.bitcast %68 : tensor<2x2x2x2x2x2x2x2x2xf32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %70 = "tt.reduce"(%69) <{axis = 7 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %94 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %94 : i32 + }) : (tensor<2x2x2x2x2x2x2x2x2xi32>) -> tensor<2x2x2x2x2x2x2x2xi32> + %71 = tt.expand_dims %70 {axis = 7 : i32} : tensor<2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x1x2xi32> + %72 = tt.broadcast %71 : tensor<2x2x2x2x2x2x2x1x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %73 = arith.xori %69, %72 : tensor<2x2x2x2x2x2x2x2x2xi32> + %74 = tt.bitcast %73 : tensor<2x2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xf32> + %75 = arith.cmpf ogt, %68, %74 : tensor<2x2x2x2x2x2x2x2x2xf32> + %76 = arith.extui %75 : tensor<2x2x2x2x2x2x2x2x2xi1> to tensor<2x2x2x2x2x2x2x2x2xi32> + %77 = tt.broadcast %13 : tensor<1x1x1x1x1x1x1x2x1xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %78 = arith.cmpi ne, %76, %77 : tensor<2x2x2x2x2x2x2x2x2xi32> + %79 = arith.select %78, %74, %68 : tensor<2x2x2x2x2x2x2x2x2xi1>, tensor<2x2x2x2x2x2x2x2x2xf32> + %80 = tt.bitcast %79 : tensor<2x2x2x2x2x2x2x2x2xf32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %81 = "tt.reduce"(%80) <{axis = 8 : i32}> ({ + ^bb0(%arg2: i32, %arg3: i32): + %94 = arith.xori %arg2, %arg3 : i32 + tt.reduce.return %94 : i32 + }) : (tensor<2x2x2x2x2x2x2x2x2xi32>) -> tensor<2x2x2x2x2x2x2x2xi32> + %82 = tt.expand_dims %81 {axis = 8 : i32} : tensor<2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x1xi32> + %83 = tt.broadcast %82 : tensor<2x2x2x2x2x2x2x2x1xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %84 = arith.xori %80, %83 : tensor<2x2x2x2x2x2x2x2x2xi32> + %85 = tt.bitcast %84 : tensor<2x2x2x2x2x2x2x2x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xf32> + %86 = arith.cmpf ogt, %79, %85 : tensor<2x2x2x2x2x2x2x2x2xf32> + %87 = arith.extui %86 : tensor<2x2x2x2x2x2x2x2x2xi1> to tensor<2x2x2x2x2x2x2x2x2xi32> + %88 = tt.broadcast %20 : tensor<1x1x1x1x1x1x1x1x2xi32> -> tensor<2x2x2x2x2x2x2x2x2xi32> + %89 = arith.cmpi ne, %87, %88 : tensor<2x2x2x2x2x2x2x2x2xi32> + %90 = arith.select %89, %85, %79 : tensor<2x2x2x2x2x2x2x2x2xi1>, tensor<2x2x2x2x2x2x2x2x2xf32> + %91 = tt.reshape %90 : tensor<2x2x2x2x2x2x2x2x2xf32> -> tensor<64x8xf32> + %92 = tt.splat %arg1 : !tt.ptr -> tensor<64x8x!tt.ptr> + %93 = tt.addptr %92, %7 : tensor<64x8x!tt.ptr>, tensor<64x8xi32> + tt.store %93, %91 : tensor<64x8x!tt.ptr> + tt.return + } +} diff --git a/test/wafer/ir/pointer-state-modulo.mlir b/test/wafer/ir/pointer-state-modulo.mlir new file mode 100644 index 00000000..16a688ca --- /dev/null +++ b/test/wafer/ir/pointer-state-modulo.mlir @@ -0,0 +1,81 @@ +#loc = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":10:0) +#loc21 = loc("a_ptr"(#loc)) +#loc22 = loc("c_ptr"(#loc)) +#loc23 = loc("M"(#loc)) +#loc24 = loc("N"(#loc)) +#loc25 = loc("stride_am"(#loc)) +#loc26 = loc("stride_cm"(#loc)) +module { + tt.func public @wrap_stacked(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 loc("M"(#loc)), %N: i32 loc("N"(#loc)), %stride_am: i32 loc("stride_am"(#loc)), %stride_cm: i32 loc("stride_cm"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c2_i32 = arith.constant 2 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant dense<4> : tensor<4x4xi32> loc(#loc2) + %offs_am = arith.constant dense<2> : tensor<4xi32> loc(#loc27) + %offs_am_0 = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc28) + %offs_am_1 = arith.addi %offs_am_0, %offs_am : tensor<4xi32> loc(#loc27) + %offs_am_2 = tt.splat %M : i32 -> tensor<4xi32> loc(#loc29) + %offs_am_3 = arith.remsi %offs_am_1, %offs_am_2 : tensor<4xi32> loc(#loc29) + %a_ptrs = tt.expand_dims %offs_am_3 {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc30) + %a_ptrs_4 = tt.splat %stride_am : i32 -> tensor<4x1xi32> loc(#loc31) + %a_ptrs_5 = arith.muli %a_ptrs, %a_ptrs_4 : tensor<4x1xi32> loc(#loc31) + %a_ptrs_6 = tt.expand_dims %offs_am_0 {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc32) + %a_ptrs_7 = tt.broadcast %a_ptrs_5 : tensor<4x1xi32> -> tensor<4x4xi32> loc(#loc33) + %a_ptrs_8 = tt.broadcast %a_ptrs_6 : tensor<1x4xi32> -> tensor<4x4xi32> loc(#loc33) + %a_ptrs_9 = arith.addi %a_ptrs_7, %a_ptrs_8 : tensor<4x4xi32> loc(#loc33) + %a_ptrs_10 = tt.splat %a_ptr : !tt.ptr -> tensor<4x4x!tt.ptr> loc(#loc34) + %a_ptrs_11 = tt.addptr %a_ptrs_10, %a_ptrs_9 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc34) + %c_ptrs = tt.expand_dims %offs_am_0 {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc35) + %c_ptrs_12 = tt.splat %stride_cm : i32 -> tensor<4x1xi32> loc(#loc36) + %c_ptrs_13 = arith.muli %c_ptrs_12, %c_ptrs : tensor<4x1xi32> loc(#loc36) + %c_ptrs_14 = tt.splat %c_ptr : !tt.ptr -> tensor<4x1x!tt.ptr> loc(#loc37) + %c_ptrs_15 = tt.addptr %c_ptrs_14, %c_ptrs_13 : tensor<4x1x!tt.ptr>, tensor<4x1xi32> loc(#loc37) + %c_ptrs_16 = tt.broadcast %c_ptrs_15 : tensor<4x1x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc38) + %c_ptrs_17 = tt.addptr %c_ptrs_16, %a_ptrs_8 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc38) + %c_ptrs_18:2 = scf.for %k = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%a_ptrs_19 = %a_ptrs_11, %c_ptrs_20 = %c_ptrs_17) -> (tensor<4x4x!tt.ptr>, tensor<4x4x!tt.ptr>) : i32 { + %a = tt.load %a_ptrs_19 : tensor<4x4x!tt.ptr> loc(#loc40) + tt.store %c_ptrs_20, %a : tensor<4x4x!tt.ptr> loc(#loc16) + %a_ptrs_21 = tt.addptr %a_ptrs_19, %cst : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc41) + %c_ptrs_22 = tt.addptr %c_ptrs_20, %cst : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc42) + scf.yield %a_ptrs_21, %c_ptrs_22 : tensor<4x4x!tt.ptr>, tensor<4x4x!tt.ptr> loc(#loc19) + } loc(#loc43) + tt.return loc(#loc20) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":19:22) +#loc2 = loc(unknown) +#loc3 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":11:19) +#loc4 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":11:32) +#loc5 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":11:38) +#loc6 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":13:30) +#loc7 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":13:41) +#loc8 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":13:61) +#loc9 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":13:53) +#loc10 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":13:22) +#loc11 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":17:41) +#loc12 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":17:33) +#loc13 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":17:21) +#loc14 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":17:52) +#loc15 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":20:20) +#loc16 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":21:25) +#loc17 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":22:18) +#loc18 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":23:18) +#loc19 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":23:8) +#loc20 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_modulo.py":19:4) +#loc27 = loc("offs_am"(#loc3)) +#loc28 = loc("offs_am"(#loc4)) +#loc29 = loc("offs_am"(#loc5)) +#loc30 = loc("a_ptrs"(#loc6)) +#loc31 = loc("a_ptrs"(#loc7)) +#loc32 = loc("a_ptrs"(#loc8)) +#loc33 = loc("a_ptrs"(#loc9)) +#loc34 = loc("a_ptrs"(#loc10)) +#loc35 = loc("c_ptrs"(#loc11)) +#loc36 = loc("c_ptrs"(#loc12)) +#loc37 = loc("c_ptrs"(#loc13)) +#loc38 = loc("c_ptrs"(#loc14)) +#loc39 = loc("a_ptrs"(#loc1)) +#loc40 = loc("a"(#loc15)) +#loc41 = loc("a_ptrs"(#loc17)) +#loc42 = loc("c_ptrs"(#loc18)) +#loc43 = loc("c_ptrs"(#loc39)) diff --git a/test/wafer/ir/pointer-state-nested_loops.mlir b/test/wafer/ir/pointer-state-nested_loops.mlir new file mode 100644 index 00000000..ca61f2c9 --- /dev/null +++ b/test/wafer/ir/pointer-state-nested_loops.mlir @@ -0,0 +1,103 @@ +#loc = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":217:0) +#loc29 = loc("in_ptr"(#loc)) +#loc30 = loc("out_ptr"(#loc)) +#loc31 = loc("stride_m"(#loc)) +module { + tt.func public @nested3(%in_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("in_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %stride_m: i32 {tt.divisibility = 16 : i32} loc("stride_m"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<6> : tensor<2x2xi32> loc(#loc1) + %cst_0 = arith.constant dense<4> : tensor<2x2xi32> loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c2_i32 = arith.constant 2 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst_1 = arith.constant dense<2> : tensor<2x2xi32> loc(#loc1) + %offs_am = tt.make_range {end = 2 : i32, start = 0 : i32} : tensor<2xi32> loc(#loc32) + %a_ptrs = tt.expand_dims %offs_am {axis = 1 : i32} : tensor<2xi32> -> tensor<2x1xi32> loc(#loc33) + %a_ptrs_2 = tt.splat %stride_m : i32 -> tensor<2x1xi32> loc(#loc34) + %a_ptrs_3 = arith.muli %a_ptrs, %a_ptrs_2 : tensor<2x1xi32> loc(#loc34) + %a_ptrs_4 = tt.expand_dims %offs_am {axis = 0 : i32} : tensor<2xi32> -> tensor<1x2xi32> loc(#loc35) + %a_ptrs_5 = tt.broadcast %a_ptrs_3 : tensor<2x1xi32> -> tensor<2x2xi32> loc(#loc36) + %a_ptrs_6 = tt.broadcast %a_ptrs_4 : tensor<1x2xi32> -> tensor<2x2xi32> loc(#loc36) + %a_ptrs_7 = arith.addi %a_ptrs_5, %a_ptrs_6 : tensor<2x2xi32> loc(#loc36) + %a_ptrs_8 = tt.splat %in_ptr : !tt.ptr -> tensor<2x2x!tt.ptr> loc(#loc37) + %a_ptrs_9 = tt.addptr %a_ptrs_8, %a_ptrs_7 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> loc(#loc37) + %c_ptrs = tt.splat %out_ptr : !tt.ptr -> tensor<2x1x!tt.ptr> loc(#loc38) + %c_ptrs_10 = tt.addptr %c_ptrs, %a_ptrs_3 : tensor<2x1x!tt.ptr>, tensor<2x1xi32> loc(#loc38) + %c_ptrs_11 = tt.broadcast %c_ptrs_10 : tensor<2x1x!tt.ptr> -> tensor<2x2x!tt.ptr> loc(#loc39) + %c_ptrs_12 = tt.addptr %c_ptrs_11, %a_ptrs_6 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> loc(#loc39) + %c_ptrs_13:2 = scf.for %i = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%a_ptrs_14 = %a_ptrs_9, %c_ptrs_15 = %c_ptrs_12) -> (tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr>) : i32 { + %a1 = tt.load %a_ptrs_14 : tensor<2x2x!tt.ptr> loc(#loc41) + %c_ptrs_16:2 = scf.for %j = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%a_ptrs_18 = %a_ptrs_14, %c_ptrs_19 = %c_ptrs_15) -> (tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr>) : i32 { + %a_ptrs_20 = tt.addptr %a_ptrs_18, %cst_1 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> loc(#loc43) + %a2 = tt.load %a_ptrs_20 : tensor<2x2x!tt.ptr> loc(#loc44) + %c_ptrs_21:2 = scf.for %k = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%a_ptrs_22 = %a_ptrs_20, %c_ptrs_23 = %c_ptrs_19) -> (tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr>) : i32 { + %a_ptrs_24 = tt.addptr %a_ptrs_22, %cst_1 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> loc(#loc46) + %a3 = tt.load %a_ptrs_24 : tensor<2x2x!tt.ptr> loc(#loc47) + tt.store %c_ptrs_23, %a1 : tensor<2x2x!tt.ptr> loc(#loc18) + %c_ptrs_25 = tt.addptr %c_ptrs_23, %cst_1 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> loc(#loc48) + tt.store %c_ptrs_25, %a2 : tensor<2x2x!tt.ptr> loc(#loc20) + %c_ptrs_26 = tt.addptr %c_ptrs_23, %cst_0 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> loc(#loc55) + tt.store %c_ptrs_26, %a3 : tensor<2x2x!tt.ptr> loc(#loc22) + %c_ptrs_27 = tt.addptr %c_ptrs_23, %cst : tensor<2x2x!tt.ptr>, tensor<2x2xi32> loc(#loc56) + scf.yield %a_ptrs_24, %c_ptrs_27 : tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr> loc(#loc24) + } loc(#loc54) + scf.yield %c_ptrs_21#0, %c_ptrs_21#1 : tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr> loc(#loc25) + } loc(#loc53) + %a_ptrs_17 = tt.addptr %c_ptrs_16#0, %cst_1 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> loc(#loc51) + scf.yield %a_ptrs_17, %c_ptrs_16#1 : tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr> loc(#loc27) + } loc(#loc52) + tt.return loc(#loc28) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":218:27) +#loc3 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":220:31) +#loc4 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":220:42) +#loc5 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":220:61) +#loc6 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":220:53) +#loc7 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":220:23) +#loc8 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":224:23) +#loc9 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":224:53) +#loc10 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":226:22) +#loc11 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":227:21) +#loc12 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":229:26) +#loc13 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":230:22) +#loc14 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":231:25) +#loc15 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":233:30) +#loc16 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":234:26) +#loc17 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":235:29) +#loc18 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":236:33) +#loc19 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":237:26) +#loc20 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":239:33) +#loc21 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":240:26) +#loc22 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":241:33) +#loc23 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":242:26) +#loc24 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":242:16) +#loc25 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":233:12) +#loc26 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":244:18) +#loc27 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":244:8) +#loc28 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_nested_loops.py":226:4) +#loc32 = loc("offs_am"(#loc2)) +#loc33 = loc("a_ptrs"(#loc3)) +#loc34 = loc("a_ptrs"(#loc4)) +#loc35 = loc("a_ptrs"(#loc5)) +#loc36 = loc("a_ptrs"(#loc6)) +#loc37 = loc("a_ptrs"(#loc7)) +#loc38 = loc("c_ptrs"(#loc8)) +#loc39 = loc("c_ptrs"(#loc9)) +#loc40 = loc("a_ptrs"(#loc10)) +#loc41 = loc("a1"(#loc11)) +#loc42 = loc("a_ptrs"(#loc12)) +#loc43 = loc("a_ptrs"(#loc13)) +#loc44 = loc("a2"(#loc14)) +#loc45 = loc("a_ptrs"(#loc15)) +#loc46 = loc("a_ptrs"(#loc16)) +#loc47 = loc("a3"(#loc17)) +#loc48 = loc("c_ptrs"(#loc19)) +#loc49 = loc("c_ptrs"(#loc21)) +#loc50 = loc("c_ptrs"(#loc23)) +#loc51 = loc("a_ptrs"(#loc26)) +#loc52 = loc("c_ptrs"(#loc40)) +#loc53 = loc("c_ptrs"(#loc42)) +#loc54 = loc("c_ptrs"(#loc45)) +#loc55 = loc(fused[#loc49, #loc48]) +#loc56 = loc(fused[#loc50, #loc49, #loc48]) diff --git a/test/wafer/ir/pointer-state-scalar_store.mlir b/test/wafer/ir/pointer-state-scalar_store.mlir new file mode 100644 index 00000000..d0369972 --- /dev/null +++ b/test/wafer/ir/pointer-state-scalar_store.mlir @@ -0,0 +1,30 @@ +#loc = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_scalar_store.py":8:0) +#loc9 = loc("output_ptr"(#loc)) +module { + tt.func public @reduce_kernel_2d(%output_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("output_ptr"(#loc))) attributes {noinline = false} { + %c8_i32 = arith.constant 8 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %pid0 = tt.get_program_id x : i32 loc(#loc10) + %base_ptr = tt.addptr %output_ptr, %pid0 : !tt.ptr, i32 loc(#loc11) + %base_ptr_0 = scf.for %i = %c0_i32 to %c8_i32 step %c1_i32 iter_args(%base_ptr_1 = %base_ptr) -> (!tt.ptr) : i32 { + %0 = arith.sitofp %i : i32 to f32 loc(#loc5) + tt.store %base_ptr_1, %0 : !tt.ptr loc(#loc5) + %base_ptr_2 = tt.addptr %base_ptr_1, %c1_i32 : !tt.ptr, i32 loc(#loc13) + scf.yield %base_ptr_2 : !tt.ptr loc(#loc7) + } loc(#loc12) + tt.return loc(#loc8) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_scalar_store.py":14:22) +#loc2 = loc(unknown) +#loc3 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_scalar_store.py":12:25) +#loc4 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_scalar_store.py":13:28) +#loc5 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_scalar_store.py":16:27) +#loc6 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_scalar_store.py":17:20) +#loc7 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_scalar_store.py":17:8) +#loc8 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_scalar_store.py":14:4) +#loc10 = loc("pid0"(#loc3)) +#loc11 = loc("base_ptr"(#loc4)) +#loc12 = loc("base_ptr"(#loc1)) +#loc13 = loc("base_ptr"(#loc6)) diff --git a/test/wafer/ir/pointer-state-tensor_index_iterargs.mlir b/test/wafer/ir/pointer-state-tensor_index_iterargs.mlir new file mode 100644 index 00000000..80d7bcce --- /dev/null +++ b/test/wafer/ir/pointer-state-tensor_index_iterargs.mlir @@ -0,0 +1,48 @@ +#loc = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":11:0) +#loc13 = loc("in0"(#loc)) +#loc14 = loc("out0"(#loc)) +#loc15 = loc("mask_bound"(#loc)) +module { + tt.func public @addptr_with_masks(%in0: !tt.ptr {tt.divisibility = 16 : i32} loc("in0"(#loc)), %out0: !tt.ptr {tt.divisibility = 16 : i32} loc("out0"(#loc)), %mask_bound: i32 loc("mask_bound"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c4_i32 = arith.constant 4 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant dense<4> : tensor<4xi32> loc(#loc2) + %cst_0 = arith.constant dense<-11> : tensor<4xi32> loc(#loc2) + %offs = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc16) + %mask = tt.splat %mask_bound : i32 -> tensor<4xi32> loc(#loc17) + %a = tt.splat %in0 : !tt.ptr -> tensor<4x!tt.ptr> loc(#loc18) + %0 = tt.splat %out0 : !tt.ptr -> tensor<4x!tt.ptr> loc(#loc6) + %out_offs:2 = scf.for %i = %c0_i32 to %c4_i32 step %c1_i32 iter_args(%offs_1 = %offs, %out_offs_2 = %offs) -> (tensor<4xi32>, tensor<4xi32>) : i32 { + %mask_3 = arith.cmpi slt, %offs_1, %mask : tensor<4xi32> loc(#loc17) + %a_4 = tt.addptr %a, %offs_1 : tensor<4x!tt.ptr>, tensor<4xi32> loc(#loc18) + %a_5 = tt.load %a_4, %mask_3, %cst_0 : tensor<4x!tt.ptr> loc(#loc20) + %1 = tt.addptr %0, %out_offs_2 : tensor<4x!tt.ptr>, tensor<4xi32> loc(#loc6) + tt.store %1, %a_5 : tensor<4x!tt.ptr> loc(#loc8) + %offs_6 = arith.addi %offs_1, %cst : tensor<4xi32> loc(#loc21) + %out_offs_7 = arith.addi %out_offs_2, %cst : tensor<4xi32> loc(#loc22) + scf.yield %offs_6, %out_offs_7 : tensor<4xi32>, tensor<4xi32> loc(#loc11) + } loc(#loc23) + tt.return loc(#loc12) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":19:22) +#loc2 = loc(unknown) +#loc3 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":12:24) +#loc4 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":20:22) +#loc5 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":21:26) +#loc6 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":22:24) +#loc7 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":21:20) +#loc8 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":22:34) +#loc9 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":23:16) +#loc10 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":24:20) +#loc11 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":24:8) +#loc12 = loc("/home/ubuntu/zwr/DLCompiler/third_party/wafer/examples/test_tensor_index_iterargs.py":19:4) +#loc16 = loc("offs"(#loc3)) +#loc17 = loc("mask"(#loc4)) +#loc18 = loc("a"(#loc5)) +#loc19 = loc("offs"(#loc1)) +#loc20 = loc("a"(#loc7)) +#loc21 = loc("offs"(#loc9)) +#loc22 = loc("out_offs"(#loc10)) +#loc23 = loc("out_offs"(#loc19)) diff --git a/test/wafer/native_math/__init__.py b/test/wafer/native_math/__init__.py new file mode 100644 index 00000000..7f6a332c --- /dev/null +++ b/test/wafer/native_math/__init__.py @@ -0,0 +1 @@ +"""Native Wafer math coverage; explicit hardware opt-in is required.""" diff --git a/test/wafer/native_math/conftest.py b/test/wafer/native_math/conftest.py new file mode 100644 index 00000000..1b235d15 --- /dev/null +++ b/test/wafer/native_math/conftest.py @@ -0,0 +1,7 @@ +"""Native math tests use ordinary TXDA tensors and the production launcher.""" +import pytest + + +@pytest.fixture(autouse=True) +def require_device(wafer_device): + return wafer_device diff --git a/test/wafer/native_math/test_log1p.py b/test/wafer/native_math/test_log1p.py new file mode 100644 index 00000000..a9629a46 --- /dev/null +++ b/test/wafer/native_math/test_log1p.py @@ -0,0 +1,70 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +"""Ascend log1p algorithm and original parameter matrix, using native Wafer tensors.""" +import pytest +import torch +import torch_txda # noqa: F401 -- registers the TXDA device +import triton +import triton.language as tl +from triton.language.extra import libdevice + +@triton.jit +def triton_log1p( + in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = x_index < xnumel + tmp0 = tl.load(in_ptr0 + x_index, xmask) + tmp1 = tl.load(in_ptr1 + x_index, xmask) + tmp2 = tmp0 + libdevice.log1p(tmp1) + tl.store(out_ptr0 + x_index, tmp2, xmask) + + +@pytest.mark.parametrize("param_list", [["float32", (2, 4096, 8), 2, 32768, 1024]]) +def test_log1p(param_list): + sigtype, shape, ncore, xblock, xblock_sub = param_list + torch.manual_seed(0) + a = torch.randn(shape, dtype=getattr(torch, sigtype)) + b = torch.randn_like(a) + out = (torch.zeros_like(a)).to("txda") + triton_log1p[(ncore,)]((a).to("txda"), (b).to("txda"), out, a.numel(), xblock, xblock_sub) + torch.testing.assert_close(out.cpu(), a + torch.log1p(b), rtol=1e-4, atol=1e-4, equal_nan=True) + + +@triton.jit +def log1p_kernel(X, Y, N: tl.constexpr): + i = tl.arange(0, N) + tl.store(Y + i, libdevice.log1p(tl.load(X + i))) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) +def test_log1p_special_values(dtype): + host = torch.tensor([-2, -1, -1 + torch.finfo(dtype).eps, -1e-7, -1e-8, -0.0, 0.0, 1e-8, + 1e-7, 1e-4, 1, 2, 10, float("inf"), -float("inf"), float("nan")], dtype=dtype) + out = (torch.zeros_like(host)).to("txda") + log1p_kernel[(1,)]((host).to("txda"), out, host.numel()) + # No absolute tolerance: log(1+x) incorrectly rounds tiny x to zero. + expected, actual = torch.log1p(host), out.cpu() + tolerance = 1e-6 if dtype == torch.float32 else 1e-3 + torch.testing.assert_close(actual, expected, rtol=tolerance, atol=0, equal_nan=True) + torch.testing.assert_close(torch.signbit(actual[5:7]), torch.signbit(expected[5:7])) diff --git a/test/wafer/native_math/test_multi_return.py b/test/wafer/native_math/test_multi_return.py new file mode 100644 index 00000000..3454414c --- /dev/null +++ b/test/wafer/native_math/test_multi_return.py @@ -0,0 +1,338 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +"""Original Ascend cross-entropy/gradient kernels with native Wafer tanh. + +Host reductions and the independent autograd reference run on CPU; all Triton +kernels receive native TXDA tensors. +""" +import pytest +import torch +import torch_txda # noqa: F401 -- registers the TXDA device +import triton +import triton.language as tl +from triton.language.extra import libdevice + +@triton.jit +def liger_cross_entropy_kernel( + X_ptr, + X_stride, + Y_ptr, + Y_stride, + weight_ptr, + loss_ptr, + z_loss_ptr, + loss_stride, + n_cols, + n_non_ignore, + sum_non_ignore_weight, + weight_sum, + ignore_index, + lse_square_scale: tl.constexpr, + label_smoothing: tl.constexpr, + reduction: tl.constexpr, # set it as constexpr since reduction is always known at compile time + softcap, + RETURN_Z_LOSS: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + HAS_WEIGHT: tl.constexpr, + HAS_SOFTCAPPING: tl.constexpr, +): + """ + This kernel computes both cross entropy loss and the gradient of the input. + + Parameters: + X_ptr: Pointer to input tensor. + X_stride (int): The stride of the input tensor. + Y_ptr: Pointer to target tensor. + Y_stride (int): The stride of the target tensor. + weight_ptr: Pointer to weight tensor. + loss_ptr: Pointer to tensor to store the loss. + z_loss_ptr: Pointer to tensor to store the z loss. No operation if RETURN_Z_LOSS is 0. + loss_stride (int): The stride of the loss tensor. + n_cols (int): The number of columns in the input tensor. + n_non_ignore (float): The number of non-ignored elements in the batch. + sum_non_ignore_weight (float): The sum of non-ignored target's weights in the batch. + weight_sum (float): The sum of weight tensor. + ignore_index (int): The index to ignore in the target. + label_smoothing (float): The amount of smoothing when computing the loss, where 0.0 means no smoothing. + lse_square_scale (float): The scaler of (logsumexp(_input)) ^ 2 adding to the loss for the stability of training. + reduction (str): The string for the reduction to apply + softcap (float): The upper threshold for scaling logits to the range (-softcap, +softcap). + RETURN_Z_LOSS (int): The boolean value to decide whether storing z loss to z_loss_ptr or not. It must be 0 or 1. + BLOCK_SIZE (int): The block size for Triton operations. + HAS_WEIGHT (bool): The boolean value to determine whether assigning weight to each of the classes. + HAS_SOFTCAPPING (bool): The boolean value to determine whether applying soft-capping or not. + """ + + # If B*T*V is too large, program_id * stride will overflow out of int32, so we convert to int64 + program_id = tl.program_id(0).to(tl.int64) + + # 1. Load Y_ptr first because if the target is ignore_index, we can return right away + Y_ptr += program_id * Y_stride + y = tl.load(Y_ptr) + + # 2. locate the start index + X_ptr += program_id * X_stride + + if y == ignore_index: + # set all X_ptr as 0 + for i in range(0, n_cols, BLOCK_SIZE): + X_offsets = i + tl.arange(0, BLOCK_SIZE) + tl.store(X_ptr + X_offsets, 0.0, mask=X_offsets < n_cols) + return + + loss_ptr += program_id * loss_stride + if RETURN_Z_LOSS: + z_loss_ptr += program_id * loss_stride + + if HAS_WEIGHT: + weight_y = tl.load(weight_ptr + y).cast(tl.float32) + + # Online softmax: 2 loads + 1 store (compared with 3 loads + 1 store for the safe softmax) + + # 3. [Online softmax] first pass: find max + sum + m = float("-inf") # m is the max value. use the notation from the paper + d = 0.0 # d is the sum. use the notation from the paper + ori_X_y = tl.load(X_ptr + y).cast( + tl.float32 + ) # we need to store the original value of X_y for the loss calculation + if HAS_SOFTCAPPING: + ori_X_y = softcap * libdevice.tanh(ori_X_y / softcap) + + # Label smoothing is a general case of normal cross entropy + scaled_x_sum = 0.0 + eps = label_smoothing / n_cols + + for i in range(0, n_cols, BLOCK_SIZE): + X_offsets = i + tl.arange(0, BLOCK_SIZE) + X_block = tl.load( + X_ptr + X_offsets, + mask=X_offsets < n_cols, + other=float("-inf"), + # Ensure float32 precision for softmax calculation + ).cast(tl.float32) + if HAS_SOFTCAPPING: + X_block = softcap * libdevice.tanh(X_block / softcap) + block_max = tl.max(X_block) + if label_smoothing > 0: + # scale X beforehand to avoid overflow + if HAS_WEIGHT: + weight_block = tl.load(weight_ptr + X_offsets, mask=X_offsets < n_cols) + scaled_x_sum += tl.sum( + tl.where(X_offsets < n_cols, -eps * X_block * weight_block, 0.0) + ) + else: + scaled_x_sum += tl.sum( + tl.where(X_offsets < n_cols, -eps * X_block, 0.0) + ) + m_new = tl.maximum(m, block_max) + d = d * tl.exp(m - m_new) + tl.sum(tl.exp(X_block - m_new)) + m = m_new + + # log (sum(e^(X_i))) = log (sum(e ^ (max(X) * e ^ (X_i - max(X))))) + # = log (e^(max(X)) * sum(e ^ (X_i - max(X)))) + # = max(X) + log (sum(e ^ (X_i - max(X)))) = m + log d + lse = m + tl.log(d) + + # 4. [Online Softmax] Second pass: compute gradients + for i in range(0, n_cols, BLOCK_SIZE): + X_offsets = i + tl.arange(0, BLOCK_SIZE) + X_block = tl.load( + X_ptr + X_offsets, + mask=X_offsets < n_cols, + other=float("-inf"), + # Ensure float32 precision for softmax calculation + ).cast(tl.float32) + if HAS_SOFTCAPPING: + intermediate = libdevice.tanh(X_block / softcap) + X_block = softcap * intermediate + + if not HAS_WEIGHT: + X_block = tl.exp(X_block - m) / d + # derivative of z-loss: 2 * lse_square_scale * lse * softmax(x_i) + X_block += 2 * lse_square_scale * lse * X_block + # smoothing term + X_block += -eps + # special handle dx_y + X_block = tl.where(X_offsets != y, X_block, X_block - (1 - label_smoothing)) + # reduction scale + if reduction == "mean": + X_block = X_block / n_non_ignore + else: + weight_block = tl.load(weight_ptr + X_offsets, mask=X_offsets < n_cols) + softmax_X = tl.exp(X_block - m) / d + # derivative of original_loss + dloss_ori = (1 - label_smoothing) * softmax_X + # specially handle dx_y + dloss_ori = tl.where( + X_offsets != y, dloss_ori, dloss_ori - (1 - label_smoothing) + ) + dloss_ori = dloss_ori * weight_y + # derivative of smooth_loss + dloss_smooth = eps * (-weight_block + softmax_X * weight_sum) + # derivative of z-loss + dz_loss = 2 * lse_square_scale * lse * softmax_X + # reduction scale + if reduction == "mean": + dloss_ori = dloss_ori / sum_non_ignore_weight + dloss_smooth = dloss_smooth / sum_non_ignore_weight + dz_loss = dz_loss / n_non_ignore + # derivative of total_loss + X_block = dloss_ori + dloss_smooth + dz_loss + + # chain rule softcapping + # d(softcap * libdevice.tanh(x / softcap)) = (1 - tanh^2(x / softcap)) + if HAS_SOFTCAPPING: + X_block = X_block * (1 - intermediate * intermediate) + + tl.store(X_ptr + X_offsets, X_block, mask=X_offsets < n_cols) + + # We need tl.debug_barrier() to ensure the new result of X_ptr is written as mentioned in + tl.debug_barrier() + + # 5. Calculate the loss + + # loss = log (softmax(X_y)) = log ((e ^ (X_y - max(X)) / sum(e ^ (X - max(X)))) + # = (X_y - max(X)) - log(sum(e ^ (X - max(X)))) + # = X_y - m - log d = X_y - lse + # sum(e ^ (X - max(X))) must >= 1 because the max term is e ^ 0 = 1 + # So we can safely calculate log (softmax(X_y)) without overflow + loss = lse - ori_X_y + if HAS_WEIGHT: + loss = weight_y * loss + + # Original loss = H(q, p), with label smoothing regularization = H(q', p) and (label_smoothing / V) = eps + # H(q', p) = (1 - label_smoothing) * H(q, p) + label_smoothing * H(u, p) + # = (1 - label_smoothing) * H(q, p) + eps * sum(logsoftmax(x_i)) + # By using m (global max of xi) and d (sum of e^(xi-m)), we can simplify as: + # = (1 - label_smoothing) * H(q, p) + (sum(-eps * x_i) + label_smoothing * (m + logd)) + if label_smoothing > 0: + if HAS_WEIGHT: + smooth_loss = scaled_x_sum + eps * lse * weight_sum + else: + smooth_loss = scaled_x_sum + label_smoothing * lse + loss = loss * (1 - label_smoothing) + smooth_loss + + # An auxiliary loss, z_loss + z_loss = lse_square_scale * lse * lse + # Normalize the loss by the number of non-ignored elements if reduction is "mean" + if reduction == "mean": + if HAS_WEIGHT: + loss = loss / sum_non_ignore_weight + else: + loss = loss / n_non_ignore + z_loss = z_loss / n_non_ignore + loss += z_loss + + tl.store(loss_ptr, loss) + if RETURN_Z_LOSS: + tl.store(z_loss_ptr, z_loss) + +@triton.jit +def element_mul_kernel( + X_ptr, + X_stride, + grad_output_ptr, + n_cols, + BLOCK_SIZE: tl.constexpr, +): + """ + This function multiplies each element of the tensor pointed by X_ptr with the value pointed by grad_output_ptr. + The multiplication is performed in-place on the tensor pointed by X_ptr. + + Parameters: + X_ptr: Pointer to the input tensor. + X_stride (int): The stride of the input tensor. + grad_output_ptr: Pointer to the gradient output value. + n_cols (int): The number of columns in the input tensor. + BLOCK_SIZE (int): The block size for Triton operations. + """ + + # Get the program ID and convert it to int64 to avoid overflow + program_id = tl.program_id(0).to(tl.int64) + + # Locate the start index + X_ptr += program_id * X_stride + + # Load the gradient output value + grad_output = tl.load(grad_output_ptr) + + # Perform the element-wise multiplication + for i in range(0, n_cols, BLOCK_SIZE): + X_offsets = i + tl.arange(0, BLOCK_SIZE) + X_block = tl.load(X_ptr + X_offsets, mask=X_offsets < n_cols) + tl.store(X_ptr + X_offsets, X_block * grad_output, mask=X_offsets < n_cols) + + +def run_cross_entropy(host, target, grad): + rows, cols = host.shape + x, y = (host).to("txda"), (target).to("txda") + loss = (torch.zeros(rows, dtype=host.dtype)).to("txda") + z_loss = (torch.zeros(rows, dtype=host.dtype)).to("txda") + non_ignore = int((target != 0).sum()) + # Preserve the original kernel, grid, softcap, smoothing and reduction. + # Framework bookkeeping/reductions are CPU-side, not another NPU OP. + liger_cross_entropy_kernel[(rows,)]( + x, x.stride(0), y, y.stride(0), None, loss, z_loss, 1, cols, + non_ignore, non_ignore, 0.0, 0, 1e-4, 0.1, "mean", 30.0, True, + min(32768, triton.next_power_of_2(cols)), False, True, + ) + # The original backward multiplies its in-place gradient by grad_output. + element_mul_kernel[(rows,)](x, x.stride(0), (grad).to("txda"), cols, + min(32768, triton.next_power_of_2(cols))) + return loss.cpu().sum(), z_loss.cpu().sum(), x.cpu() + + +def reference_cross_entropy(host, target, grad): + # Independent CPU autograd oracle, including softcap, label smoothing, + # ignored labels and z-loss. Round per-row stores as in the original ABI. + x = host.float().requires_grad_(True) + capped = 30.0 * torch.tanh(x / 30.0) + lse = torch.logsumexp(capped, dim=-1) + loss = torch.nn.functional.cross_entropy(capped, target, ignore_index=0, + label_smoothing=0.1, reduction="none") + mask = target != 0 + z = torch.where(mask, 1e-4 * lse.square(), 0.0) + losses = (loss + z) / mask.sum() + losses.sum().backward() + # The kernel first stores the unscaled gradient in the input's dtype; + # the separate backward kernel then multiplies it by the scalar gradient. + expected_grad = (x.grad.to(host.dtype) * grad).to(host.dtype) + return losses.to(host.dtype).sum(), (z / mask.sum()).to(host.dtype).sum(), expected_grad + + +@pytest.mark.parametrize("B,T,V", [(2, 512, 4096)]) +@pytest.mark.parametrize("scalar,dtype,atol,rtol", [ + (1.0, torch.bfloat16, 1e-8, 5e-2), + (1.0, torch.float32, 1e-8, 1e-6), +]) +def test_correctness_functional(B, T, V, scalar, dtype, atol, rtol): + torch.manual_seed(0) + host = torch.randn(B * T, V, dtype=dtype) * scalar + target = torch.randint(0, V, (B * T,), dtype=torch.long) + grad = torch.randn((), dtype=dtype) + first = run_cross_entropy(host, target, grad) + second = run_cross_entropy(host, target, grad) + reference = reference_cross_entropy(host, target, grad) + for actual, repeat, expected in zip(first, second, reference): + # Keep the source's two-execution comparisons and add independent truth. + assert torch.allclose(actual, repeat, atol=atol, rtol=rtol) + torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol) diff --git a/test/wafer/native_math/test_relu.py b/test/wafer/native_math/test_relu.py new file mode 100644 index 00000000..a3611a8f --- /dev/null +++ b/test/wafer/native_math/test_relu.py @@ -0,0 +1,78 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +"""Ascend relu algorithm and original parameter matrix, using native Wafer tensors.""" +import pytest +import torch +import torch_txda # noqa: F401 -- registers the TXDA device +import triton +import triton.language as tl + + +@triton.jit +def wafer_relu(x): + # Ordered comparison keeps NaN and -0, matching torch.relu. The current + # Wafer maximum instruction discards NaN, so use its compare/select path. + return tl.where(x < 0, 0, x) + +@triton.jit +def triton_relu( + in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = x_index < xnumel + tmp0 = tl.load(in_ptr0 + x_index, xmask) + tmp1 = tl.load(in_ptr1 + x_index, xmask) + tmp2 = tmp0 + wafer_relu(tmp1) + tl.store(out_ptr0 + x_index, tmp2, xmask) + + +@pytest.mark.parametrize("param_list", [ + ["float32", (2, 4096, 8), 2, 32768, 512], + ["float16", (2, 4096, 8), 2, 32768, 512], +]) +def test_relu(param_list): + sigtype, shape, ncore, xblock, xblock_sub = param_list + torch.manual_seed(0) + a = torch.randn(shape, dtype=getattr(torch, sigtype)) + b = torch.randn_like(a) + out = (torch.zeros_like(a)).to("txda") + triton_relu[(ncore,)]((a).to("txda"), (b).to("txda"), out, a.numel(), xblock, xblock_sub) + tolerance = 1e-3 if sigtype == "float16" else 1e-4 + torch.testing.assert_close(out.cpu(), a + torch.relu(b), rtol=tolerance, atol=tolerance, equal_nan=True) + + +@triton.jit +def relu_kernel(X, Y, N: tl.constexpr): + i = tl.arange(0, N) + x = tl.load(X + i) + tl.store(Y + i, wafer_relu(x)) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) +def test_relu_special_values(dtype): + host = torch.tensor([-float("inf"), -1, -0.0, 0.0, 1e-5, 1, float("inf"), float("nan")], dtype=dtype) + out = (torch.zeros_like(host)).to("txda") + relu_kernel[(1,)]((host).to("txda"), out, host.numel()) + expected, actual = torch.relu(host), out.cpu() + torch.testing.assert_close(actual, expected, rtol=0, atol=0, equal_nan=True) + torch.testing.assert_close(torch.signbit(actual[:-1]), torch.signbit(expected[:-1])) diff --git a/test/wafer/native_math/test_unary.py b/test/wafer/native_math/test_unary.py new file mode 100644 index 00000000..5d25cd3a --- /dev/null +++ b/test/wafer/native_math/test_unary.py @@ -0,0 +1,98 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +"""Ascend atan/scalar coverage and FlagTree isnan coverage on native Wafer APIs. + +Sources: test/ascend/passed_tests/test_{atan,isnan,scalar_calc}.py at b1991f1; +FlagTree third_party/tsingmicro/examples/test_libdevice.py at 22f4ff0. +The original license notice is retained above. Shapes, dtypes and comparison +tolerances are preserved; only device transport, references and OP entry change. +""" +import pytest +import torch +import torch_txda # noqa: F401 -- registers the TXDA device +import triton +import triton.language as tl +from triton.language.extra import libdevice + + +@triton.jit +def unary_kernel(X, Y, N: tl.constexpr, BLOCK: tl.constexpr, OP: tl.constexpr): + offsets = tl.arange(0, BLOCK) + x = tl.load(X + offsets, offsets < N, 0) + result = getattr(libdevice, OP)(x) + tl.store(Y + offsets, result, offsets < N) + + +@triton.jit +def scalar_tanh_kernel(X, Y): + tl.store(Y, libdevice.tanh(tl.load(X))) + + +@pytest.mark.parametrize("dtype,sigtype", [(torch.float32, "float32"), (torch.float16, "float16")]) +@pytest.mark.parametrize("N,NUMEL", [(3, 32), (-32, 32), (37, 64), (-256, 256), (781, 1024)]) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + torch.manual_seed(0) + host = torch.randn((N,), dtype=dtype) + x, out = (host).to("txda"), (torch.zeros_like(host)).to("txda") + unary_kernel[(1,)](x, out, N, NUMEL, "atan", debug=True) + tolerance = 1e-3 if sigtype == "float16" else 1e-4 + torch.testing.assert_close(out.cpu(), torch.atan(host), rtol=tolerance, atol=tolerance, equal_nan=True) + + +@pytest.mark.parametrize("sigtype", ["float32", "float16", "bfloat16"]) +@pytest.mark.parametrize("N", [256]) +def test_isnan(sigtype, N): + # Extend FlagTree's FP32/4/128 cases with Ascend's three dtype/256 cases. + torch.manual_seed(0) + host = torch.randn((N,), dtype=getattr(torch, sigtype)) + host[1] = float("nan") + host[N // 4] = float("inf") + host[N // 2] = -float("inf") + out = (torch.zeros(N, dtype=torch.bool)).to("txda") + unary_kernel[(1,)]((host).to("txda"), out, N, N, "isnan") + torch.testing.assert_close(out.cpu(), torch.isnan(host), rtol=0, atol=0) + assert out.cpu()[1].item() is True + + +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_tanh_calc(param_list): + sigtype, N = param_list + torch.manual_seed(0) + host = torch.randn(N, dtype=getattr(torch, sigtype)) + out = (torch.zeros(1, dtype=host.dtype)).to("txda") + scalar_tanh_kernel[(1,)]((host).to("txda"), out) + torch.testing.assert_close(out.cpu()[0], torch.tanh(host[0]), rtol=1e-4, atol=1e-4) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("op", ["atan", "tanh", "isnan"]) +def test_unary_special_values(dtype, op): + host = torch.tensor([-float("inf"), -10, -1, -0.0, 0.0, 1e-5, 1, 10, + float("inf"), float("nan")], dtype=dtype) + expected = getattr(torch, op)(host) + out = (torch.zeros_like(expected)).to("txda") + unary_kernel[(1,)]((host).to("txda"), out, host.numel(), 16, op) + actual = out.cpu() + tolerance = {torch.float16: 1e-3, torch.bfloat16: 1e-3, torch.float32: 1e-4}[dtype] + torch.testing.assert_close(actual, expected, rtol=tolerance, atol=tolerance, equal_nan=True) + if op != "isnan": + torch.testing.assert_close(torch.signbit(actual[3:5]), torch.signbit(expected[3:5])) diff --git a/test/wafer/ops/conftest.py b/test/wafer/ops/conftest.py new file mode 100644 index 00000000..6bc34a3a --- /dev/null +++ b/test/wafer/ops/conftest.py @@ -0,0 +1,13 @@ +"""Ascend-derived cases use TXDA tensors and their original deterministic inputs.""" +import pytest + + +@pytest.fixture(autouse=True) +def require_device(wafer_device): + import numpy as np + import torch + # The removed harness reset both generators before every test. Keep its + # input baseline explicitly; test-local seeds can still override it. + np.random.seed(0) + torch.manual_seed(0) + return wafer_device diff --git a/test/wafer/ops/test_2d_permute.py b/test/wafer/ops/test_2d_permute.py new file mode 100644 index 00000000..958ecaa1 --- /dev/null +++ b/test/wafer/ops/test_2d_permute.py @@ -0,0 +1,59 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 + + +def fn(x): + return x.t() + + +@triton.jit +def triton_2d_permute(output_ptr, input_ptr, X: tl.constexpr, Y: tl.constexpr): + xindex = tl.arange(0, X * Y) + input_local = tl.load(input_ptr + xindex) + output_local = input_local.reshape(X, Y).trans().reshape(X * Y) + tl.store(output_ptr + xindex, output_local) + + +@pytest.mark.parametrize("X", [32, 64, 256]) +@pytest.mark.parametrize("Y", [16, 32]) +def test_cases(X, Y): + + x = torch.randn((X, Y)).cpu() + output1 = fn(x) + output2 = torch.randn(output1.shape, dtype=output1.dtype).cpu() + + output2_txda = output2.to("txda") + x_txda = x.to("txda") + triton_2d_permute[1, 1, 1](output2_txda, x_txda, X, Y, debug=True) + with torch.no_grad(): + output2.copy_(output2_txda.cpu()) + print(output1) + print(output2) + + torch.testing.assert_close(output1, output2, rtol=1e-3, atol=1e-3) diff --git a/test/wafer/ops/test_3Dgrid.py b/test/wafer/ops/test_3Dgrid.py new file mode 100644 index 00000000..b6b8e410 --- /dev/null +++ b/test/wafer/ops/test_3Dgrid.py @@ -0,0 +1,118 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import pytest + +BLOCK: tl.constexpr = 32 + + +@triton.jit +def triton_( + in_ptr0, + out_ptr0, + x0_numel, + r1_numel, + XBLOCK: tl.constexpr, + XBLOCK_SUB: tl.constexpr, + block_id_threshold: tl.constexpr, + XBLOCK1: tl.constexpr, + num_core: tl.constexpr, +): + RBLOCK: tl.constexpr = 64 + + block_idx = ( + tl.program_id(0) * tl.num_programs(1) * tl.num_programs(2) + + tl.program_id(1) * tl.num_programs(2) + + tl.program_id(2) + ) + if block_idx < block_id_threshold: + offset = block_idx * XBLOCK + loops1 = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB # 32+23 / 24 = 2 + upper = offset + XBLOCK + else: + offset = ( + block_id_threshold * XBLOCK + (block_idx - block_id_threshold) * XBLOCK1 + ) # pid=34 offset = 9*32 + (34-9)*24 = 888 + loops1 = (XBLOCK1 + XBLOCK_SUB - 1) // XBLOCK_SUB # 1 + if block_idx == num_core - 1: + upper = x0_numel + else: + upper = offset + XBLOCK1 # 912 + + base1 = tl.arange(0, XBLOCK_SUB) + base2 = tl.arange(0, RBLOCK) + loops2: tl.constexpr = (r1_numel + RBLOCK - 1) // RBLOCK + for loop1 in range(loops1): + x = offset + (loop1 * XBLOCK_SUB) + base1 + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1[None, :] + x0 = offset + (loop1 * XBLOCK_SUB) + base1[:, None] + xmask = x0 < upper + r1_prime = base2[:, None] + rindex = base2 + r1 = base2[None, :] + rmask = r1 < r1_numel + tmp0 = tl.load(in_ptr0 + (r1 + (64 * x0)), rmask & xmask, other=0.0) + + tmp1 = tl.reshape(tmp0, [XBLOCK_SUB, RBLOCK]) + tmp2_tmp = tl.sum(tmp1, 1) + tmp2 = tmp2_tmp.reshape(XBLOCK_SUB, 1) + + tl.store(out_ptr0 + (x0), tmp2, xmask) + + +guards = {"dummy": None} + + +# @pytest.mark.skip(reason="multi-process error, to be fixed.") +@pytest.mark.parametrize("size", [(1025, 64)]) +def test_3dgrid(size): + b = torch.randn((size), dtype=torch.float32).cpu() + c = torch.sum(b, dim=1) + + ret = torch.randn((size[0]), dtype=torch.float32).cpu() + + b_txda = b.to("txda") + ret_txda = ret.to("txda") + triton_[5, 2, 4]( + b_txda, + ret_txda, + size[0], + size[1], + XBLOCK=32, + XBLOCK_SUB=16, + block_id_threshold=9, + XBLOCK1=24, + num_core=40, + debug=True, + ) + with torch.no_grad(): + ret.copy_(ret_txda.cpu()) + print(c[0:8]) + print(ret[0:8]) + torch.testing.assert_close(c, ret) + print("test 3D launch passed") + + +if __name__ == "__main__": + pytest.main([__file__]) diff --git a/test/wafer/ops/test_abs.py b/test/wafer/ops/test_abs.py new file mode 100644 index 00000000..c966cba6 --- /dev/null +++ b/test/wafer/ops/test_abs.py @@ -0,0 +1,65 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0): + res = torch.abs(x0) + return res + + +@triton.jit +def triton_abs(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp2 = tl.abs(tmp0) + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["int32", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_abs[ncore, 1, 1](x0_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_abs_2.py b/test/wafer/ops/test_abs_2.py new file mode 100644 index 00000000..bac6cc96 --- /dev/null +++ b/test/wafer/ops/test_abs_2.py @@ -0,0 +1,70 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_abs(x0): + res = torch.abs(x0) + return res + + +@triton.jit +def triton_abs(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.abs(tmp0) + tl.store(out_ptr0 + (x0), tmp1, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float16", (4, 4), 4, 4, 4], + ["float32", (4, 4), 4, 4, 4], + ], +) +def test_abs(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype) + y_ref = torch_abs(x0) + tyname = test_common.get_triton_sig_typename(dtype) + + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0 = x0.cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_abs[ncore, 1, 1](x0_txda, y_cal_txda, xblock, xblock_sub, debug=True) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_add.py b/test/wafer/ops/test_add.py new file mode 100644 index 00000000..3530dd06 --- /dev/null +++ b/test/wafer/ops/test_add.py @@ -0,0 +1,97 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0, x1): + res = x0 + x1 + return res + + +@triton.jit +def triton_add( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 + tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_add[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float32", (128, 4096, 160), 1310720, 64, 8], + ["float16", (128, 4096, 160), 1310720, 64, 8], + ["int8", (128, 4096, 160), 1310720, 64, 8], + ], +) +def test_all_blocks_parallel(param_list, monkeypatch): + monkeypatch.setenv("TRITON_ALL_BLOCKS_PARALLEL", "1") + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_add[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) + monkeypatch.delenv("TRITON_ALL_BLOCKS_PARALLEL") diff --git a/test/wafer/ops/test_add_multi_return.py b/test/wafer/ops/test_add_multi_return.py new file mode 100644 index 00000000..f67f0c0b --- /dev/null +++ b/test/wafer/ops/test_add_multi_return.py @@ -0,0 +1,99 @@ +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise_even_blocks(x0, x1, xblock): + """ + 参考实现:只有第偶数个 block 的元素会相加(block id 从 0 开始计数),其他 block 输出保持为 0(或 y_cal 初始值)。 + xblock: 一个 block 的大小(线性元素数),与 Triton kernel 中的 XBLOCK 对应 + """ + # 展平为一维线性内存,保持 dtype & device + x0_flat = x0.reshape(-1) + x1_flat = x1.reshape(-1) + n = x0_flat.numel() + idx = torch.arange(n, device=x0.device) + block_id = (idx // xblock) % 2 + mask = block_id == 0 + res_flat = torch.zeros_like(x0_flat) + # 只有偶数 block 做加法 + res_flat[mask] = x0_flat[mask] + x1_flat[mask] + return res_flat.reshape(x0.shape) + + +@triton.jit +def triton_add( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + """ + Triton kernel:只有 program_id(0) 为偶数的 block 才会执行 load/add/store;奇数 block 跳过(不修改输出) + """ + bid = tl.program_id(0) # block id + # 如果此 block 是奇数,直接返回(不做任何 store) + if bid % 2 != 0: + return + + offset = bid * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = x0_prime + tmp0 = tl.load(in_ptr0 + x0, None) + tmp1 = tl.load(in_ptr1 + x0, None) + tmp2 = tmp0 + tmp1 + tl.store(out_ptr0 + x0, tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # 参考结果:只在偶数 block 做加法 + y_ref = torch_pointwise_even_blocks(x0, x1, xblock) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_add[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float32", (128, 2048, 1280), 1310720, 256, 32], + ["float16", (128, 4096, 1280), 1310720, 512, 64], + ["int8", (128, 4096, 1280), 1310720, 512, 64], + ], +) +def test_all_blocks_parallel(param_list, monkeypatch): + monkeypatch.setenv("TRITON_ALL_BLOCKS_PARALLEL", "1") + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise_even_blocks(x0, x1, xblock) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_add[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) + monkeypatch.delenv("TRITON_ALL_BLOCKS_PARALLEL") diff --git a/test/wafer/ops/test_advance.py b/test/wafer/ops/test_advance.py new file mode 100644 index 00000000..83cfae1b --- /dev/null +++ b/test/wafer/ops/test_advance.py @@ -0,0 +1,231 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest + + +@triton.jit +def fn_npu_( + output_ptr, + x_ptr, + y_ptr, + z_ptr, + output_ptr1, + XB: tl.constexpr, + YB: tl.constexpr, + ZB: tl.constexpr, +): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + # idx = tl.arange(0,XB*YB*ZB) + block_ptr_in = tl.make_block_ptr( + base=x_ptr, + shape=(XB, YB, ZB), + strides=(YB * ZB, ZB, 1), + offsets=(9, 6, 5), + block_shape=(XB, YB, ZB), + order=(2, 1, 0), + ) + bbptr = tl.advance(block_ptr_in, (-9, -6, -5)) + # XB,YB,1 + X = tl.load(bbptr) + # X = tl.load(x_ptr + idx) + # Y = tl.load(y_ptr + idx) + + # xx=tl.view(X,(ZB*YB,XB)) + + oidx = ( + xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + ) + + block_ptr_out = tl.make_block_ptr( + base=output_ptr, + shape=(XB, YB, ZB), + strides=(YB * ZB, ZB, 1), + offsets=(0, 0, 0), + block_shape=(XB, YB, ZB), + order=(2, 1, 0), + ) + tl.store(block_ptr_out, X) + # tl.store(output_ptr + tl.arange(0,ZB*YB)[:,None]*XB+xidx[None,:], xx) + # tl.store(output_ptr + xidx[:,None]*YB+yidx[None,:], yy) + + +@triton.jit +def fn_npu_2d( + output_ptr, + x_ptr, + y_ptr, + z_ptr, + output_ptr1, + XB: tl.constexpr, + YB: tl.constexpr, + ZB: tl.constexpr, +): + xoffset = tl.program_id(0) + block_ptr_in = tl.make_block_ptr( + base=x_ptr, + shape=(XB, YB), + strides=(YB, 1), + offsets=(6 + xoffset, 5), + block_shape=(XB, YB), + order=(1, 0), + ) + bbptr = tl.advance(block_ptr_in, (-6, -5)) + # XB,YB,1 + X = tl.load(bbptr) + + block_ptr_out = tl.make_block_ptr( + base=output_ptr, + shape=(XB, YB), + strides=(YB, 1), + offsets=(xoffset, 0), + block_shape=(XB, YB), + order=(1, 0), + ) + tl.store(block_ptr_out, X) + + +@triton.jit +def fn_npu_3d(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + block_ptr_in = tl.make_block_ptr( + base=x_ptr, + shape=(XB, YB, ZB), + strides=(YB * ZB, ZB, 1), + offsets=(0, 0, 0), + block_shape=(XB, YB, 2), + order=(2, 1, 0), + ) + + block_ptr_out = tl.make_block_ptr( + base=output_ptr, + shape=(XB, YB, ZB), + strides=(YB * ZB, ZB, 1), + offsets=(0, 0, 0), + block_shape=(XB, YB, 2), + order=(2, 1, 0), + ) + + for _ in range(ZB // 2): + X = tl.load(block_ptr_in, boundary_check=(0, 1, 2)) + tl.store(block_ptr_out, X, boundary_check=(0, 1, 2)) + block_ptr_in = tl.advance(block_ptr_in, (0, 0, 2)) + block_ptr_out = tl.advance(block_ptr_out, (0, 0, 2)) + + +@pytest.mark.parametrize("dtype", ["int32", "float32", "int16"]) +@pytest.mark.parametrize("shape", [(32, 8, 8), (8, 8, 4)]) +def test_advance_with_boundary_check(dtype, shape): + x = torch.randint( + low=-128, high=128, size=shape, dtype=eval("torch." + dtype) + ).cpu() + + output = torch.randint(1, shape, dtype=eval("torch." + dtype)).cpu() + a = x + + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_3d[1, 1, 1](output_txda, x_txda, XB=shape[0], YB=shape[1], ZB=shape[2]) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + torch.testing.assert_close(output, a) + + +@pytest.mark.parametrize("dtype", ["int32", "float32", "int16"]) +@pytest.mark.parametrize("shape", [(2, 4), (4, 2), (2, 16), (16, 1)]) +def test_advance_supplement(dtype, shape): + x = torch.randint( + low=-128, high=128, size=shape, dtype=eval("torch." + dtype) + ).cpu() + y = torch.randint( + low=-128, high=128, size=shape, dtype=eval("torch." + dtype) + ).cpu() + z = torch.randint( + low=-128, high=128, size=shape, dtype=eval("torch." + dtype) + ).cpu() + + output = torch.randint(1, shape, dtype=eval("torch." + dtype)).cpu() + output1 = output + + a = x + + output_txda = output.to("txda") + x_txda = x.to("txda") + y_txda = y.to("txda") + z_txda = z.to("txda") + output1_txda = output_txda + fn_npu_2d[1, 1, 1](output_txda, x_txda, y_txda, z_txda, output1_txda, XB=shape[0], YB=shape[1], ZB=1) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output1.copy_(output1_txda.cpu()) + + torch.testing.assert_close(output, a) + + +paras = [ + ("*fp32", eval("torch.float32"), 2, 256, 16), + ("*fp32", eval("torch.float32"), 8, 8, 4), + ("*fp16", eval("torch.float16"), 2, 256, 16), + ("*fp16", eval("torch.float16"), 8, 8, 4), + ("*i8", eval("torch.int8"), 2, 256, 16), + ("*i8", eval("torch.int8"), 8, 8, 4), +] + + +@pytest.mark.parametrize("para_type,data_type,XB,YB,ZB", paras) +def test_npu(para_type, data_type, XB, YB, ZB): + + x = torch.randint(low=-128, high=128, size=(XB, YB, ZB), dtype=data_type).cpu() + y = torch.randint(low=-128, high=128, size=(XB, YB, ZB), dtype=data_type).cpu() + z = torch.randint(low=-128, high=128, size=(XB, YB, ZB), dtype=data_type).cpu() + + print(f"shape = {x.shape}") + print(x.dtype) + + output = torch.randint(1, (XB, YB, ZB), dtype=data_type).cpu() + output1 = output + print(f"output.dtype={output.dtype}") + + a = x + print(a) + output_txda = output.to("txda") + x_txda = x.to("txda") + y_txda = y.to("txda") + z_txda = z.to("txda") + output1_txda = output_txda + fn_npu_[1, 1, 1](output_txda, x_txda, y_txda, z_txda, output1_txda, XB=XB, YB=YB, ZB=ZB, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output1.copy_(output1_txda.cpu()) + print(output) + torch.testing.assert_close(output, a) + + +if __name__ == "__main__": + pytest.main([__file__]) diff --git a/test/wafer/ops/test_and.py b/test/wafer/ops/test_and.py new file mode 100644 index 00000000..c69d8cbb --- /dev/null +++ b/test/wafer/ops/test_and.py @@ -0,0 +1,67 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_and(x0, x1): + res = x0 & x1 + return res + + +@triton.jit +def triton_and(in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index) + tmp1 = tl.load(in_ptr1 + x_index) + tmp2 = tmp0 & tmp1 + tl.store(out_ptr0 + x_index, tmp2) + + +@pytest.mark.parametrize('param_list', + [ + ['int32', (2, 4096, 8), 2, 32768, 1024], + ]) +def test_and(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_and(x0, x1) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval('torch.' + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_and[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_arange.py b/test/wafer/ops/test_arange.py new file mode 100644 index 00000000..5c1e029f --- /dev/null +++ b/test/wafer/ops/test_arange.py @@ -0,0 +1,159 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import math +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import test_common + + +def torch_arange(start, end): + TRITON_MAX_TENSOR_NUMEL = 1048576 + if end < start: + raise ValueError( + "arange's end argument must be greater than the start argument" + ) + if end - start > TRITON_MAX_TENSOR_NUMEL: + raise ValueError( + f"end - start must be less than or equal to TRITON_MAX_TENSOR_NUMEL = {TRITON_MAX_TENSOR_NUMEL}" + ) + return torch.arange(start, end) + + +def torch_arange_access(start, end): + z = torch.zeros([end], dtype=torch.int32).cpu() + v = torch.arange(start, end).cpu() + z[start:end] = v + return z + + +@triton.jit +def triton_arange(z, BLOCK: tl.constexpr, START: tl.constexpr, END: tl.constexpr): + off = tl.arange(0, BLOCK) + val = tl.arange(START, END) + tl.store(z + off, val) + + +@triton.jit +def triton_arange_access( + z, BLOCK: tl.constexpr, START: tl.constexpr, END: tl.constexpr +): + off = tl.arange(START, END) + val = tl.arange(START, END) + tl.store(z + off, val) + + +@pytest.mark.parametrize( + "param_list", + [ + [0, 103], + [0, 1024], + ], +) +def test_case(param_list): + start, end = param_list + shape = [end - start] + block = end - start + dtype = "int32" + + y_ref = torch_arange(start, end) + y_cal = torch.zeros(shape, dtype=torch.int32).cpu() + + y_cal_txda = y_cal.to("txda") + triton_arange[(1,)](y_cal_txda, START=start, END=end, BLOCK=block) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + + test_common.validate_cmp(dtype, y_cal, y_ref) + + +@pytest.mark.parametrize( + "param_list", + [ + [0, 103], + [0, 1024], + ], +) +def test_case_access(param_list): + start, end = param_list + shape = [end] + block = end - start + dtype = "int32" + + y_ref = torch_arange_access(start, end) + y_cal = torch.zeros(shape, dtype=torch.int32).cpu() + + y_cal_txda = y_cal.to("txda") + triton_arange_access[(1,)](y_cal_txda, START=start, END=end, BLOCK=block) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + + test_common.validate_cmp(dtype, y_cal, y_ref) + + +@pytest.mark.parametrize( + "invalid_param_list", + [ + [0, 10000000], + # [8, 128], + ], +) +@test_common.raises_with_match( + triton.compiler.errors.CompilationError, + r"arange's range must be a power of 2|at \d+:\d+:", +) +def test_arange_invalid_range(invalid_param_list): + start, end = invalid_param_list + shape = [end - start] + block = end - start + + y_cal = torch.zeros(shape, dtype=torch.int32).cpu() + + y_cal_txda = y_cal.to("txda") + triton_arange[(1,)](y_cal_txda, START=start, END=end, BLOCK=block) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + + +@pytest.mark.parametrize( + "invalid_param_list", + [ + [1024, 128], + ], +) +@test_common.raises_with_match( + triton.compiler.errors.CompilationError, + r"arange's range must be a power of 2|at \d+:\d+:", +) +def test_arange_invalid_revinput(invalid_param_list): + start, end = invalid_param_list + range = abs(end - start) + shape = [range] + block = range + + y_cal = torch.zeros(shape, dtype=torch.int32).cpu() + + y_cal_txda = y_cal.to("txda") + triton_arange[(1,)](y_cal_txda, START=start, END=end, BLOCK=block) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) diff --git a/test/wafer/ops/test_associative_scan.py b/test/wafer/ops/test_associative_scan.py new file mode 100644 index 00000000..938c3fa6 --- /dev/null +++ b/test/wafer/ops/test_associative_scan.py @@ -0,0 +1,222 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import math +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + +import test_common + + +def torch_func(x, dim, reverse): + if reverse: + x = torch.flip(x, [dim]) + res = torch.cumsum(x, dim=dim) + return res + + +def combine_fn_test_torch(a, b, combine_fn): + return torch.maximum(a, b) + + +def torch_func_scan(x: torch.Tensor, dim: int, combine_fn="maximum", reverse=False): + """ + PyTorch implements associative_scan, with semantics fully aligned with Triton. + """ + dim = dim % x.ndim + + if reverse: + x = x.flip(dim) + + N = x.size(dim) + tensors = torch.unbind(x, dim=dim) + + outputs = [] + carry = tensors[0] + outputs.append(carry) + + for i in range(1, N): + carry = combine_fn_test_torch(tensors[i], carry, combine_fn) + outputs.append(carry) + + output = torch.stack(outputs, dim=dim) + + if reverse: + output = output.flip(dim) + + return output + + +@triton.jit +def combine_fn_test(a, b): + return tl.maximum(a, b) + + +@triton.jit +def triton_kernel_1d_scan( + out_ptr0, + in_ptr0, + dim: tl.constexpr, + reverse: tl.constexpr, + numel_x: tl.constexpr, + XBLOCK: tl.constexpr, +): + tl.static_assert( + numel_x == XBLOCK, "numel_x must be equal to XBLOCK in this kernel" + ) + idx = tl.arange(0, XBLOCK) + x = tl.load(in_ptr0 + idx) + ret = tl.associative_scan(x, axis=dim, reverse=reverse, combine_fn=combine_fn_test) + tl.store(out_ptr0 + idx, ret) + + +@triton.jit +def triton_kernel_2d_scan( + out_ptr0, + in_ptr0, + dim: tl.constexpr, + reverse: tl.constexpr, + numel_x: tl.constexpr, + numel_r: tl.constexpr, + XBLOCK: tl.constexpr, + RBLOCK: tl.constexpr, +): + tl.static_assert( + numel_x == XBLOCK, "numel_x must be equal to XBLOCK in this kernel" + ) + tl.static_assert( + numel_r == RBLOCK, "numel_r must be equal to RBLOCK in this kernel" + ) + idx_x = tl.arange(0, XBLOCK) + idx_r = tl.arange(0, RBLOCK) + idx = idx_x[:, None] * numel_r + idx_r[None, :] + x = tl.load(in_ptr0 + idx) + ret = tl.associative_scan(x, axis=dim, reverse=reverse, combine_fn=combine_fn_test) + tl.store(out_ptr0 + idx, ret) + + +@triton.jit +def triton_kernel_3d_scan( + out_ptr0, + in_ptr0, + dim: tl.constexpr, + reverse: tl.constexpr, + numel_x: tl.constexpr, + numel_r: tl.constexpr, + numel_z: tl.constexpr, + XBLOCK: tl.constexpr, + RBLOCK: tl.constexpr, + ZBLOCK: tl.constexpr, +): + tl.static_assert( + numel_x == XBLOCK, "numel_x must be equal to XBLOCK in this kernel" + ) + tl.static_assert( + numel_r == RBLOCK, "numel_r must be equal to RBLOCK in this kernel" + ) + tl.static_assert( + numel_z == ZBLOCK, "numel_z must be equal to ZBLOCK in this kernel" + ) + idx_x = tl.arange(0, XBLOCK) + idx_r = tl.arange(0, RBLOCK) + idx_z = tl.arange(0, ZBLOCK) + idx = ( + idx_x[:, None, None] * numel_r * numel_z + + idx_r[None, :, None] * numel_z + + idx_z[None, None, :] + ) + x = tl.load(in_ptr0 + idx) + ret = tl.associative_scan(x, axis=dim, reverse=reverse, combine_fn=combine_fn_test) + tl.store(out_ptr0 + idx, ret) + + +def triton_func_scan(x, dim, reverse): + res = torch.empty_like(x) + print(f"res.dtype = {res.dtype}") + shape = x.size() + if len(shape) == 1: + if dim >= 1: + pytest.skip("dim >= 1 for 1D tensor, skipping.") + res_txda = res.to("txda") + x_txda = x.to("txda") + triton_kernel_1d_scan[1, 1, 1](res_txda, x_txda, dim, reverse, x_txda.shape[0], x_txda.shape[0]) + with torch.no_grad(): + res.copy_(res_txda.cpu()) + elif len(shape) == 2: + if dim >= 2: + pytest.skip("dim >= 2 for 2D tensor, skipping.") + res_txda = res.to("txda") + x_txda = x.to("txda") + triton_kernel_2d_scan[1, 1, 1]( + res_txda, x_txda, dim, reverse, x_txda.shape[0], x_txda.shape[1], x_txda.shape[0], x_txda.shape[1] + ) + with torch.no_grad(): + res.copy_(res_txda.cpu()) + elif len(shape) == 3: + if dim >= 3: + pytest.skip("dim >= 3 for 3D tensor, skipping.") + res_txda = res.to("txda") + x_txda = x.to("txda") + triton_kernel_3d_scan[1, 1, 1]( + res_txda, + x_txda, + dim, + reverse, + x_txda.shape[0], + x_txda.shape[1], + x_txda.shape[2], + x_txda.shape[0], + x_txda.shape[1], + x_txda.shape[2], + ) + with torch.no_grad(): + res.copy_(res_txda.cpu()) + else: + pytest.skip(f"This testcase unsupported tensor dimension: {len(shape)}") + + return res + + +@pytest.mark.parametrize("dtype", ["int32", "float32"]) +@pytest.mark.parametrize("shape", [(128,), (8, 4), (128, 4, 16)]) +@pytest.mark.parametrize("dim", [0, 1, 2]) +@pytest.mark.parametrize( + "combine_fn", + [ + "maximum", + ], +) +@pytest.mark.parametrize("reverse", [False]) +def test_scan(dtype, shape, dim, combine_fn, reverse): + torch.manual_seed(0) + x = test_common.generate_tensor(shape=shape, dtype=dtype) + x_gold = x + cpu_res = torch_func_scan(x_gold, dim, combine_fn, reverse) + print(f"cpu_res: {cpu_res}") + + x_npu = x.cpu() + triton_res = triton_func_scan(x_npu, dim, reverse) + print(f"triton_res: {triton_res}") + + test_common.validate_cmp(dtype, triton_res, cpu_res) + print(f"Validate PASS") diff --git a/test/wafer/ops/test_associative_scan_multi_input.py b/test/wafer/ops/test_associative_scan_multi_input.py new file mode 100644 index 00000000..2ba9c740 --- /dev/null +++ b/test/wafer/ops/test_associative_scan_multi_input.py @@ -0,0 +1,143 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import pytest + + +@triton.jit +def sum_combine_fn(a_value, a_index, b_value, b_index): + new_val = a_value + b_value + new_idx = a_index + b_index + return (new_val, new_idx) + + +@triton.jit +def prefix_scan_last_dim_kernel( + vals_ptr, + idx_ptr, + out_vals_ptr, + out_idxs_ptr, + axis_size, + total_slices, + BLOCK_SIZE: tl.constexpr, +): + slice_id = tl.program_id(0) + if slice_id >= total_slices: + return + + offs = tl.arange(0, BLOCK_SIZE) + mask = offs < axis_size + base = slice_id * axis_size + + vals = tl.load(vals_ptr + base + offs, mask=mask, other=0.0) + idxs = tl.load(idx_ptr + base + offs, mask=mask, other=0) + + pre_vals, pre_idxs = tl.associative_scan( + (vals, idxs), axis=0, combine_fn=sum_combine_fn + ) + + tl.store(out_vals_ptr + base + offs, pre_vals, mask=mask) + tl.store(out_idxs_ptr + base + offs, pre_idxs, mask=mask) + + +def multi_input_prefix_sum(values: torch.Tensor, index: torch.Tensor, axis=0): + assert values.shape == index.shape + assert values.device == index.device + rank = values.ndim + if axis < 0: + axis += rank + assert 0 <= axis < rank + + # 1. change to make axis the last dim + order = list(range(rank)) + if axis != rank - 1: + order[axis], order[-1] = order[-1], order[axis] + inv_order = [0] * rank + for i, o in enumerate(order): + inv_order[o] = i + + vals_p = values.permute(order).contiguous() + idxs_p = index.permute(order).contiguous() + + shape_p = vals_p.shape + axis_size = shape_p[-1] + total_slices = 1 + for d in shape_p[:-1]: + total_slices *= d + + out_vals_p = torch.empty_like(vals_p) + out_idxs_p = torch.empty_like(idxs_p) + + BLOCK_SIZE = 1 << (axis_size - 1).bit_length() + vals_p_txda = vals_p.to("txda") + idxs_p_txda = idxs_p.to("txda") + out_vals_p_txda = out_vals_p.to("txda") + out_idxs_p_txda = out_idxs_p.to("txda") + prefix_scan_last_dim_kernel[(total_slices,)]( + vals_p_txda, + idxs_p_txda, + out_vals_p_txda, + out_idxs_p_txda, + axis_size, + total_slices, + BLOCK_SIZE=BLOCK_SIZE, + ) + with torch.no_grad(): + out_vals_p.copy_(out_vals_p_txda.cpu()) + out_idxs_p.copy_(out_idxs_p_txda.cpu()) + + # 2. permute back + out_vals = out_vals_p.permute(inv_order) + out_idxs = out_idxs_p.permute(inv_order) + return out_vals, out_idxs + + +@pytest.mark.parametrize( + "shape, axis", + [ + ((10,), 0), + ((4, 4), 0), + ((2, 10, 5), 1), + ], +) +def test_multi_input_prefix_sum(shape, axis): + torch.manual_seed(0) + device = "cpu" + + values = torch.randn(shape, device=device, dtype=torch.float32) + index = torch.arange(values.numel(), device=device, dtype=torch.int32).reshape( + shape + ) + + triton_vals, triton_idxs = multi_input_prefix_sum(values, index, axis=axis) + + torch_vals = values.cumsum(dim=axis) + torch_idxs = index.cumsum(dim=axis) + + assert torch.allclose( + triton_vals, torch_vals, rtol=1e-5, atol=1e-8 + ), f"数值不匹配!shape={shape}, axis={axis}\nTriton: {triton_vals}\nPyTorch: {torch_vals}" + assert torch.equal( + triton_idxs, torch_idxs + ), f"索引不匹配!shape={shape}, axis={axis}\nTriton: {triton_idxs}\nPyTorch: {torch_idxs}" diff --git a/test/wafer/ops/test_block_ptr.py b/test/wafer/ops/test_block_ptr.py new file mode 100644 index 00000000..19e93dda --- /dev/null +++ b/test/wafer/ops/test_block_ptr.py @@ -0,0 +1,93 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest + +@triton.jit +def fn_npu_(output_ptr, x_ptr,y_ptr,z_ptr,output_ptr1,XB : tl.constexpr,YB : tl.constexpr,ZB : tl.constexpr): + xidx=tl.arange(0,XB) + yidx=tl.arange(0,YB) + zidx=tl.arange(0,ZB) + idx=xidx[:,None,None]*YB*ZB+yidx[None,:,None]*ZB+zidx[None,None,:] + block_ptr_in=tl.make_block_ptr( + base = x_ptr, + shape = (XB,YB,ZB), + strides = (YB*ZB,ZB,1), + offsets = (0,0,0), + block_shape = (XB,YB,ZB), + order = (2,1,0), + ) + X = tl.load(block_ptr_in) + + oidx=xidx[:,None,None]*YB*ZB+yidx[None,:,None]*ZB+zidx[None,None,:] + + block_ptr_out=tl.make_block_ptr( + base = output_ptr, + shape = (XB,YB,ZB), + strides = (YB*ZB,ZB,1), + offsets = (0,0,0), + block_shape = (XB,YB,ZB), + order = (2,1,0), + ) + tl.store(block_ptr_out,X) + +paras = [ + ('*fp32',eval('torch.float32'),2,256,16), + ('*fp32',eval('torch.float32'),8,8,4), + ('*fp16',eval('torch.float16'),2,256,16), + ('*fp16',eval('torch.float16'),8,8,4), + ('*i8',eval('torch.int8'),2,256,16), + ('*i8',eval('torch.int8'),8,8,4), +] + +@pytest.mark.parametrize('para_type,data_type,XB,YB,ZB', paras) +def test_npu(para_type,data_type,XB,YB,ZB): + + x = torch.randint(low=-128,high=128,size=(XB,YB,ZB),dtype=data_type).cpu() + y = torch.randint(low=-128,high=128,size=(XB,YB,ZB),dtype=data_type).cpu() + z = torch.randint(low=-128,high=128,size=(XB,YB,ZB),dtype=data_type).cpu() + + print(f"shape = {x.shape}") + print(x.dtype) + + output = torch.randint(1, (XB,YB,ZB), dtype=data_type).cpu() + output1 = output + print(f"output.dtype={output.dtype}") + + a = x + print(a) + output_txda = output.to("txda") + x_txda = x.to("txda") + y_txda = y.to("txda") + z_txda = z.to("txda") + output1_txda = output_txda + fn_npu_[1,1,1](output_txda,x_txda,y_txda,z_txda,output1_txda, XB=XB, YB=YB, ZB=ZB, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output1.copy_(output1_txda.cpu()) + print(output) + torch.testing.assert_close(output,a) diff --git a/test/wafer/ops/test_broadcast_op.py b/test/wafer/ops/test_broadcast_op.py new file mode 100644 index 00000000..c84606e4 --- /dev/null +++ b/test/wafer/ops/test_broadcast_op.py @@ -0,0 +1,60 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 + +NBLOCKS = 1 +XS = tl.constexpr(128) +YS = tl.constexpr(4) +ZS = tl.constexpr(8) +NUMEL = tl.constexpr(XS.value * ZS.value) + + +@triton.jit +def fn_broadcast(output_ptr, x_ptr, length): + col_offsets = tl.arange(0, NUMEL) + input = tl.load(x_ptr + col_offsets) + result = ( + input.reshape((XS, 1, ZS)).broadcast_to((XS, YS, ZS)).reshape((XS * YS * ZS)) + ) + brc_col_offsets = tl.arange(0, NUMEL * YS) + tl.store(output_ptr + brc_col_offsets, result) + + +def test_broadcast(): + length = NUMEL + + x = torch.randn((XS, 1, ZS), dtype=torch.float32).cpu() + output = torch.randn((XS, YS, ZS), dtype=torch.float32).cpu() + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_broadcast[NBLOCKS, 1, 1](output_txda, x_txda, length, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + assert torch.equal(output, x.repeat(1, YS, 1)) + + +if __name__ == "__main__": + test_broadcast() diff --git a/test/wafer/ops/test_cat_dim.py b/test/wafer/ops/test_cat_dim.py new file mode 100644 index 00000000..80168a39 --- /dev/null +++ b/test/wafer/ops/test_cat_dim.py @@ -0,0 +1,123 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest + +@triton.jit +def fn3_dim0(output_ptr, x1_ptr, x2_ptr, x3_ptr, x1_shape: tl.constexpr, x2_shape: tl.constexpr, x3_shape: tl.constexpr): + idx_start = 0 + x1_idx = tl.arange(0,x1_shape) + X1 = tl.load(x1_ptr + x1_idx) + tl.store(output_ptr + x1_idx, X1) + + idx_start += x1_shape + x2_idx = tl.arange(0,x2_shape) + X2 = tl.load(x2_ptr + x2_idx) + tl.store(output_ptr + idx_start + x2_idx, X2) + + idx_start += x2_shape + x3_idx = tl.arange(0,x3_shape) + X3 = tl.load(x3_ptr + x3_idx) + tl.store(output_ptr + idx_start + x3_idx, X3) + + +@triton.jit +def fn4_dim1(output_ptr, x0_ptr, x1_ptr, x2_ptr, x3_ptr, dim0_len: tl.constexpr, + x0_len: tl.constexpr, x1_len: tl.constexpr, x2_len: tl.constexpr, x3_len: tl.constexpr): + + total_dim1_len = x0_len + x1_len + x2_len + x3_len + x0 = tl.load(x0_ptr + tl.arange(0, dim0_len * x0_len)) + x0 = x0.reshape(dim0_len, x0_len) + x1 = tl.load(x1_ptr + tl.arange(0, dim0_len * x1_len)) + x1 = x1.reshape(dim0_len, x1_len) + x2 = tl.load(x2_ptr + tl.arange(0, dim0_len * x2_len)) + x2 = x2.reshape(dim0_len, x2_len) + x3 = tl.load(x3_ptr + tl.arange(0, dim0_len * x3_len)) + x3 = x3.reshape(dim0_len, x3_len) + idx_start = 0 + nidx0 = (tl.arange(0, dim0_len)[:, None] * total_dim1_len + idx_start) + tl.arange(0, x0_len) + tl.store(output_ptr + nidx0, x0) + + idx_start += x0_len + nidx1 = (tl.arange(0, dim0_len)[:, None] * total_dim1_len + idx_start) + tl.arange(0, x1_len) + tl.store(output_ptr + nidx1, x1) + + idx_start += x1_len + nidx2 = (tl.arange(0, dim0_len)[:, None] * total_dim1_len + idx_start) + tl.arange(0, x2_len) + tl.store(output_ptr + nidx2, x2) + + idx_start += x2_len + nidx3 = (tl.arange(0, dim0_len)[:, None] * total_dim1_len + idx_start) + tl.arange(0, x3_len) + tl.store(output_ptr + nidx3, x3) + + +def test_cat_dim0(): + data_type = torch.float16 + x1_shape = (64, 64) + x2_shape = (16, 64) + x3_shape = (32, 64) + + x1 = torch.rand(x1_shape, dtype=data_type).cpu() + x2 = torch.rand(x2_shape, dtype=data_type).cpu() + x3 = torch.rand(x3_shape, dtype=data_type).cpu() + res = torch.zeros((x1_shape[0] + x2_shape[0] + x3_shape[0], 64), dtype=data_type).cpu() + res_txda = res.to("txda") + x1_txda = x1.to("txda") + x2_txda = x2.to("txda") + x3_txda = x3.to("txda") + fn3_dim0[(1,1,1)](res_txda, x1_txda, x2_txda, x3_txda, x1_shape[0]*x1_shape[1], x2_shape[0]*x2_shape[1], x3_shape[0]*x3_shape[1]) + with torch.no_grad(): + res.copy_(res_txda.cpu()) + + res_ref = torch.cat((x1, x2, x3), dim=0) + assert torch.allclose(res_ref, res, rtol=1e-03, atol=1e-03, equal_nan=True) + + +def test_cat_dim1(): + data_type = torch.float16 + x0_shape = (64,4) + x1_shape = (64,64) + x2_shape = (64,8) + x3_shape = (64,16) + + dim = 1 + x0 = torch.rand(x0_shape, dtype=data_type).cpu() + x1 = torch.rand(x1_shape, dtype=data_type).cpu() + x2 = torch.rand(x2_shape, dtype=data_type).cpu() + x3 = torch.rand(x3_shape, dtype=data_type).cpu() + res = torch.zeros((64, x0_shape[dim] + x1_shape[dim] + x2_shape[dim] + x3_shape[dim]), dtype=data_type).cpu() + res_txda = res.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + x2_txda = x2.to("txda") + x3_txda = x3.to("txda") + fn4_dim1[(1,1,1)](res_txda, x0_txda, x1_txda, x2_txda, x3_txda, 64, x0_shape[1], x1_shape[1], x2_shape[1], x3_shape[1]) + with torch.no_grad(): + res.copy_(res_txda.cpu()) + + res_ref = torch.cat((x0, x1, x2, x3), dim = 1) + assert torch.allclose(res_ref, res, rtol=1e-03, atol=1e-03, equal_nan=True) \ No newline at end of file diff --git a/test/wafer/ops/test_cdiv.py b/test/wafer/ops/test_cdiv.py new file mode 100644 index 00000000..390a7522 --- /dev/null +++ b/test/wafer/ops/test_cdiv.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_cdiv(x0, x1): + return torch.div(x0, x1, rounding_mode='trunc') + (x0 % x1 > 0).to(torch.int) + + +@triton.jit +def triton_cdiv(in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = tl.cdiv(XBLOCK, XBLOCK_SUB) + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index, None) + tmp1 = tl.load(in_ptr1 + x_index, None) + tmp2 = tl.cdiv(tmp0, tmp1) + tl.store(out_ptr0 + x_index, tmp2, None) + + +@pytest.mark.parametrize('param_list', + [ + ['int32', (4096,), 1, 4096, 4096], + ]) +def test_cdiv(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + 1 + # torch结果 + torch_res = torch_cdiv(x0, x1) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval('torch.' + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_cdiv[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_ceil.py b/test/wafer/ops/test_ceil.py new file mode 100644 index 00000000..4d5feba1 --- /dev/null +++ b/test/wafer/ops/test_ceil.py @@ -0,0 +1,68 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + +def torch_ceil(x0): + res = torch.ceil(x0) + return res + +@triton.jit +def triton_ceil(in_ptr0, out_ptr0, XBLOCK : tl.constexpr, XBLOCK_SUB : tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.ceil(tmp0) + tl.store(out_ptr0 + (x0), tmp1, None) + + +@pytest.mark.parametrize('param_list', + [ + # ['float16', (2, 4096, 8), 32, 2048, 64], + ['float32', (2, 4096, 8), 32, 2048, 64], + # ['int8', (2, 4096, 8), 32, 2048, 64], + ]) +def test_ceil(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype) + y_ref = torch_ceil(x0) + tyname = test_common.get_triton_sig_typename(dtype) + + y_cal = torch.zeros(shape, dtype=eval('torch.' + dtype)).cpu() + x0 = x0.cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_ceil[ncore, 1, 1](x0_txda, y_cal_txda, xblock, xblock_sub, debug=True) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + y_ref = y_ref.cpu() + test_common.validate_cmp_with_expection(dtype, y_cal, y_ref, True) diff --git a/test/wafer/ops/test_clamp.py b/test/wafer/ops/test_clamp.py new file mode 100644 index 00000000..91d745d4 --- /dev/null +++ b/test/wafer/ops/test_clamp.py @@ -0,0 +1,64 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time +import torch +import torch_txda # noqa: F401 +import test_common +def torch_clamp_float(x0): + res = torch.clamp(x0, 0.0, 100.0) + return res + +@triton.jit +def triton_clamp_float(in_ptr0, out_ptr0, XBLOCK : tl.constexpr, XBLOCK_SUB : tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.clamp(tmp0, 0.0, 100.0) + tl.store(out_ptr0 + (x0), tmp1, None) + +@pytest.mark.parametrize('param_list', + [ + # int原生不支持 + ['float16', (4, 4), 4, 4, 4], + ['float32', (4, 4), 4, 4, 4], + ]) +def test_clamp(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype) + y_ref = torch_clamp_float(x0) + tyname = test_common.get_triton_sig_typename(dtype) + + y_cal = torch.zeros(shape, dtype = eval('torch.' + dtype)).cpu() + x0 = x0.cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_clamp_float[ncore, 1, 1](x0_txda, y_cal_txda, xblock, xblock_sub, debug=True) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_common.py b/test/wafer/ops/test_common.py new file mode 100644 index 00000000..ce564caa --- /dev/null +++ b/test/wafer/ops/test_common.py @@ -0,0 +1,247 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +from typing import Optional +import torch +import torch_txda # noqa: F401 +import pytest +import functools +import re +import numpy as np + +_float_dtypes = ["float32", "float16", "bfloat16"] +_int_dtypes = ["int32", "int64", "int16", "int8"] +_uint_dtypes = ["uint8", "uint16", "uint32", "uint64"] +_all_dtypes_no_bool = _float_dtypes + _int_dtypes +_all_dtypes = _all_dtypes_no_bool + ["bool"] +_32bit_dtypes = ["float32", "int32"] +_16bit_dtypes = ["float16", "bfloat16", "int16"] + + +def generate_numpy(shape, dtype, low=None, high=None): + if dtype in _int_dtypes + _uint_dtypes: + iinfo = np.iinfo(getattr(np, dtype)) + low = iinfo.min if low is None else max(low, iinfo.min) + high = iinfo.max if high is None else min(high, iinfo.max) + dty = getattr(np, dtype) + return np.random.randint(low, high, shape, dtype=dty) + elif dtype == "float16" or dtype == "float32": + return np.random.normal(0, 1, shape).astype(dtype) + elif dtype == "bfloat16": + return ( + np.random.normal(0, 1, shape).astype("float32").view("uint32") + & np.uint32(0xFFFF0000) + ).view("float32") + elif dtype == "bool": + return np.random.randint(low=0, high=2, size=shape).astype(bool) + else: + raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype)) + + +def generate_tensor(shape, dtype): + if dtype == "float32" or dtype == "float16" or dtype == "bfloat16": + return torch.randn(size=shape, dtype=eval("torch." + dtype)) + elif dtype == "int32" or dtype == "int64" or dtype == "int16": + return torch.randint(low=0, high=2000, size=shape, dtype=eval("torch." + dtype)) + elif dtype == "int8": + return torch.randint(low=0, high=127, size=shape, dtype=eval("torch." + dtype)) + elif dtype == "bool": + return torch.randint(low=0, high=2, size=shape).bool() + elif dtype == "uint8": + return torch.randint(low=0, high=255, size=shape, dtype=torch.uint8) + else: + raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype)) + + +def get_triton_sig_typename(dtype): + if dtype == "float32": + tyname = "*fp32" + elif dtype == "int32": + tyname = "*i32" + elif dtype == "int64": + tyname = "*i64" + elif dtype == "float16": + tyname = "*fp16" + elif dtype == "int16": + tyname = "*i16" + elif dtype == "int8": + tyname = "*i8" + elif dtype == "bool": + tyname = "*i1" + else: + raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype)) + return tyname + + +# Relative error: abs(x_ref - x_cal) / abs(x_ref) +# Absolute error: abs(x_ref - x_cal) + + +# calculation type operators require different error range +# It is a stricter verification and not satisfied now, save it here +def validate_cal(dtype, y_cal, y_ref): + if dtype == "float16": + if torch.mean(y_ref) < 0.001: + assert ( + torch.abs(y_cal - y_ref) < 0.001 + ), "|y_cal - y_ref| < 0.001 is required !" + else: + diff = torch.div(torch.abs(y_cal - y_ref), torch.abs(y_cal)) < 0.001 + # all true + assert diff.all(), "Relative error is less than 0.001 !" + if dtype == "float32": + if torch.mean(y_ref) < 0.0001: + assert ( + torch.abs(y_cal - y_ref) < 0.0001 + ), "|y_cal - y_ref| < 0.0001 is required !" + else: + diff = torch.div(torch.abs(y_cal - y_ref), torch.abs(y_cal)) < 0.0001 + assert diff.all(), "Relative error is less than 0.001 !" + elif dtype == "bfloat16": + diff = torch.div(torch.abs(y_cal - y_ref), torch.abs(y_cal)) < 0.001 + assert diff.all(), "Relative error is less than 0.001 !" + elif dtype == "int32" or dtype == "int64" or dtype == "int16" or dtype == "int8": + assert torch.equal(y_cal, y_ref) + elif ( + dtype == "uint8" or dtype == "uint16" or dtype == "uint32" or dtype == "uint64" + ): + assert torch.equal(y_cal, y_ref) + elif dtype == "bool": + assert torch.equal(y_cal, y_ref) + else: + raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype)) + + +# moving and comparison ops require no precision error +def validate_cmp(dtype, y_cal, y_ref, overflow_mode: Optional[str] = None): + y_cal = y_cal.cpu() + y_ref = y_ref.cpu() + if overflow_mode == "saturate": + if dtype in ["float32", "float16"]: + min_value = -torch.finfo(dtype).min + max_value = torch.finfo(dtype).max + elif dtype in ["int32", "int16", "int8"]: + min_value = torch.iinfo(dtype).min + max_value = torch.iinfo(dtype).max + elif dtype == "bool": + min_value = 0 + max_value = 1 + else: + raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype)) + y_ref = torch.clamp(y_ref, min=min_value, max=max_value) + if dtype == "float16": + torch.testing.assert_close(y_ref, y_cal, rtol=1e-03, atol=1e-03, equal_nan=True) + elif dtype == "bfloat16": + torch.testing.assert_close( + y_ref.to(torch.float32), + y_cal.to(torch.float32), + rtol=1e-03, + atol=1e-03, + equal_nan=True, + ) + elif dtype == "float32": + torch.testing.assert_close(y_ref, y_cal, rtol=1e-04, atol=1e-04, equal_nan=True) + elif dtype == "int32" or dtype == "int64" or dtype == "int16" or dtype == "int8": + assert torch.equal(y_cal, y_ref) + elif ( + dtype == "uint8" or dtype == "uint16" or dtype == "uint32" or dtype == "uint64" + ): + assert torch.equal(y_cal, y_ref) + elif dtype == "bool": + assert torch.equal(y_cal, y_ref) + else: + raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype)) + + +def validate_cmp_with_expection(dtype, y_cal, y_ref, expect): + if dtype == "float32" or dtype == "float16" or dtype == "bfloat16": + if expect: + assert torch.allclose(y_ref, y_cal, rtol=1e-03, atol=1e-03, equal_nan=True) + else: + assert not torch.allclose( + y_ref, y_cal, rtol=1e-03, atol=1e-03, equal_nan=True + ) + elif ( + dtype == "int32" + or dtype == "int64" + or dtype == "int16" + or dtype == "int8" + or dtype == "uint8" + or dtype == "uint16" + or dtype == "uint32" + or dtype == "uint64" + ): + if expect: + assert torch.equal(y_cal, y_ref) + else: + assert not torch.equal(y_cal, y_ref) + else: + raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype)) + + +# Use the following pytest fixture to run one test case by only single worker. +# Refer to https://pytest-xdist.readthedocs.io/en/stable/how-to.html#making-session-scoped-fixtures-execute-only-once +@pytest.fixture(scope="function") +def pytest_runonce(worker_id, request, cache): + if (cache.get(request.node.nodeid, "none")) == "none": + cache.set(request.node.nodeid, worker_id) + else: + file_name = f"pytest_{worker_id}.txt" + with open(file_name, "a") as file: + file.write(f"{request.node.nodeid} is already processed by {worker_id}") + return True + yield True + cache.set(request.node.nodeid, "none") + + +def raises_with_match(expected_exception, match_pattern): + def decorator(test_func): + @functools.wraps(test_func) + def wrapper(*args, **kwargs): + with pytest.raises(expected_exception, match=match_pattern): + return test_func(*args, **kwargs) + + return wrapper + + return decorator + + +def capture_output(expected_output): + def decorator(test_func): + @functools.wraps(test_func) + def wrapper(*args, **kwargs): + capsys = kwargs.pop("capsys", None) + if capsys is None: + try: + capsys = pytest.fixture(capsys)() + except: + raise RuntimeError( + "This decorator requires pytest's capsys fixture" + ) + test_func(capsys, *args, **kwargs) + captured = capsys.readouterr() + # pybind11::scoped_ostream_redirect captures std::cout with \x00 inserted + # for now, no idea how to eliminate \x00 from C++ side. + cleaned = re.sub(r"\x00", "", captured.out) + assert expected_output in cleaned + + return wrapper + + return decorator diff --git a/test/wafer/ops/test_conv.py b/test/wafer/ops/test_conv.py new file mode 100644 index 00000000..9af1c203 --- /dev/null +++ b/test/wafer/ops/test_conv.py @@ -0,0 +1,257 @@ +import math +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +@triton.jit +def _conv_transpose2d_kernel( + x_ptr, + w_ptr, + output_ptr, + bias_ptr, + stride_h, + stride_w, + padding_h, + padding_w, + dilation_h, + dilation_w, + in_channels, + out_channels, + kernel_h, + kernel_w, + h_in, + w_in, + h_out, + w_out, + stride_xn, + stride_xc, + stride_xh, + stride_xw, + stride_wn, + stride_wc, + stride_wh, + stride_ww, + stride_on, + stride_oc, + stride_oh, + stride_ow, + BLOCK_SIZE: tl.constexpr, +): + pid = tl.program_id(0) + num_pid_n = tl.cdiv(h_out * w_out, BLOCK_SIZE) + pid_b = pid // (num_pid_n * out_channels) + pid_c = (pid // num_pid_n) % out_channels + pid_n = pid % num_pid_n + + offs_n = pid_n * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + offs_h = offs_n // w_out + offs_w = offs_n % w_out + mask = offs_n < (h_out * w_out) + + bias = tl.load(bias_ptr + pid_c) if bias_ptr is not None else 0.0 + acc = tl.zeros((BLOCK_SIZE,), dtype=tl.float32) + bias + + x_offset = pid_b * stride_xn + w_offset = pid_c * stride_wc + + for c_in in range(in_channels): + w_c_offset = w_offset + c_in * stride_wn + for kh in range(kernel_h): + for kw in range(kernel_w): + h_in_val = offs_h + padding_h - kh * dilation_h + w_in_val = offs_w + padding_w - kw * dilation_w + + cond_h = (h_in_val >= 0) & (h_in_val < h_in * stride_h) + cond_w = (w_in_val >= 0) & (w_in_val < w_in * stride_w) + cond = cond_h & cond_w & mask + + cond_h_stride = (h_in_val % stride_h) == 0 + cond_w_stride = (w_in_val % stride_w) == 0 + final_cond = cond & cond_h_stride & cond_w_stride + + h_in_idx = tl.where(final_cond, h_in_val // stride_h, 0) + w_in_idx = tl.where(final_cond, w_in_val // stride_w, 0) + + x_offsets = ( + x_offset + + c_in * stride_xc + + h_in_idx * stride_xh + + w_in_idx * stride_xw + ) + + w_val = tl.load(w_ptr + w_c_offset + kh * stride_wh + kw * stride_ww) + x_vals = tl.load(x_ptr + x_offsets, mask=final_cond, other=0.0) + acc += x_vals * w_val + + out_offset = ( + pid_b * stride_on + pid_c * stride_oc + offs_h * stride_oh + offs_w * stride_ow + ) + tl.store(output_ptr + out_offset, acc, mask=mask) + + +class ModelNew(torch.nn.Module): + def __init__( + self, + in_channels, + out_channels, + kernel_size, + stride=1, + padding=0, + dilation=1, + bias=True, + ): + super().__init__() + self.in_channels = in_channels + self.out_channels = out_channels + self.kernel_size = kernel_size + self.stride = stride + self.padding = padding + self.dilation = dilation + + self.weight = torch.nn.Parameter( + torch.empty(in_channels, out_channels, kernel_size, kernel_size) + ) + if bias: + self.bias = torch.nn.Parameter(torch.empty(out_channels)) + else: + self.register_parameter("bias", None) + + torch.nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5)) + if self.bias is not None: + fan_in, _ = torch.nn.init._calculate_fan_in_and_fan_out(self.weight) + bound = 1 / math.sqrt(fan_in) if fan_in > 0 else 0 + torch.nn.init.uniform_(self.bias, -bound, bound) + + def forward(self, x): + batch, in_channels, h_in, w_in = x.shape + h_out = ( + (h_in - 1) * self.stride + - 2 * self.padding + + self.dilation * (self.kernel_size - 1) + + 1 + ) + w_out = ( + (w_in - 1) * self.stride + - 2 * self.padding + + self.dilation * (self.kernel_size - 1) + + 1 + ) + + out = torch.empty( + (batch, self.out_channels, h_out, w_out), device=x.device, dtype=x.dtype + ) + + stride_xn, stride_xc, stride_xh, stride_xw = x.stride() + stride_wn, stride_wc, stride_wh, stride_ww = self.weight.stride() + stride_on, stride_oc, stride_oh, stride_ow = out.stride() + + total_blocks = batch * self.out_channels * math.ceil((h_out * w_out) / 64) + + x_txda = x.to("txda") + self_weight_txda = self.weight.to("txda") + out_txda = out.to("txda") + self_bias_txda = self.bias.to("txda") + _conv_transpose2d_kernel[(total_blocks,)]( + x_txda, + self_weight_txda, + out_txda, + self_bias_txda, + self.stride, + self.stride, + self.padding, + self.padding, + self.dilation, + self.dilation, + self.in_channels, + self.out_channels, + self.kernel_size, + self.kernel_size, + h_in, + w_in, + h_out, + w_out, + stride_xn, + stride_xc, + stride_xh, + stride_xw, + stride_wn, + stride_wc, + stride_wh, + stride_ww, + stride_on, + stride_oc, + stride_oh, + stride_ow, + 64, + ) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + return out + + +# =========================================================== +# Parameterize ALL inputs via pytest.mark.parametrize +# =========================================================== +@pytest.mark.parametrize( + "batch,in_ch,out_ch,kernel,H,W,stride,padding,dilation,dtype", + [ + # your requested single-configuration repeated for two dtypes + # (2, 4, 4, 3, 8, 8, 2, 1, 1, "float32"), + (2, 4, 4, 3, 8, 8, 2, 1, 1, "float16"), + (2, 32, 32, 3, 32, 32, 5, 1, 2, "float16"), + ], +) +def test_conv_transpose2d_param( + batch, in_ch, out_ch, kernel, H, W, stride, padding, dilation, dtype +): + device = torch.device("cpu") + + # generate input on npu using test_common helper + x = test_common.generate_tensor((batch, in_ch, H, W), dtype).cpu() + + model = ModelNew( + in_channels=in_ch, + out_channels=out_ch, + kernel_size=kernel, + stride=stride, + padding=padding, + dilation=dilation, + bias=True, + ) + + model.to(device) + model.weight.data = ( + model.weight.data.to(device).to(eval("torch." + dtype)).contiguous() + ) + if model.bias is not None: + model.bias.data = ( + model.bias.data.to(device).to(eval("torch." + dtype)).contiguous() + ) + + ref = ( + torch.nn.ConvTranspose2d( + in_channels=in_ch, + out_channels=out_ch, + kernel_size=kernel, + stride=stride, + padding=padding, + dilation=dilation, + bias=True, + ) + .to(device) + .to(eval("torch." + dtype)) + ) + + with torch.no_grad(): + ref.weight.data.copy_(model.weight.data) + ref.bias.data.copy_(model.bias.data) + + with torch.no_grad(): + y_cal = model(x) + y_ref = ref(x) + + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_cos.py b/test/wafer/ops/test_cos.py new file mode 100644 index 00000000..87fdadab --- /dev/null +++ b/test/wafer/ops/test_cos.py @@ -0,0 +1,102 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + +def standard_unary(x0, dtype): + res = torch.cos(x0) + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + ret = tl.cos(x) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary(in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, 'float32'), + # (torch.float16, 'float16'), + # (torch.bfloat16, 'bfloat16'), + # (torch.int8, 'int8'), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize('dtype,sigtype', types) +@pytest.mark.parametrize('N,NUMEL', shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + print(out) + + test_common.validate_cmp(sigtype, out, ans) \ No newline at end of file diff --git a/test/wafer/ops/test_cos_2.py b/test/wafer/ops/test_cos_2.py new file mode 100644 index 00000000..be0b0967 --- /dev/null +++ b/test/wafer/ops/test_cos_2.py @@ -0,0 +1,64 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + +def torch_cos(x0): + res = torch.cos(x0) + return res + +@triton.jit +def triton_cos(in_ptr0, out_ptr0, XBLOCK : tl.constexpr, XBLOCK_SUB : tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.cos(tmp0) + tl.store(out_ptr0 + (x0), tmp1, None) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 32, 2048, 64], + ]) +def test_cos(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype) + y_ref = torch_cos(x0) + tyname = test_common.get_triton_sig_typename(dtype) + y_cal = torch.zeros(shape, dtype = eval('torch.' + dtype)).cpu() + x0 = x0.cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_cos[ncore, 1, 1](x0_txda, y_cal_txda, xblock, xblock_sub, debug=True) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_count_dim0.py b/test/wafer/ops/test_count_dim0.py new file mode 100644 index 00000000..384bd570 --- /dev/null +++ b/test/wafer/ops/test_count_dim0.py @@ -0,0 +1,214 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +# Copyright (c) Huawei Technologies Co., Ltd. 2025-2025. All rights reserved. +import pytest +import triton +import triton.language as tl +import time +import torch +import torch_txda # noqa: F401 +import test_common + +def standard_count(x0, cmp_val, dim, dtype): + res = (x0 == cmp_val).sum(dim=dim) + return res + +def standard_count_gt(x0, cmp_val, dim, dtype): + res = (x0 > cmp_val).sum(dim=dim) + return res + +def standard_count_lt(x0, cmp_val, dim, dtype): + res = (x0 < cmp_val).sum(dim=dim) + return res + +@triton.jit +def count(in_ptr0, out_ptr0, cmp_val, dim : tl.constexpr, M : tl.constexpr, N : tl.constexpr, MNUMEL: tl.constexpr, NNUMEL: tl.constexpr): + mblk_idx = tl.arange(0,MNUMEL) + nblk_idx = tl.arange(0,NNUMEL) + mmask = mblk_idx < M + nmask = nblk_idx < N + mask = (mmask[:,None]) & (nmask[None,:]) + idx = mblk_idx[:,None]*N + nblk_idx[None,:] + x = tl.load(in_ptr0+idx, mask = mask, other = 0) + tmp1 = (x == cmp_val) + tmp2 = tmp1.to(tl.float32) + ret = tl.sum(tmp2, dim) + tl.store(out_ptr0 + nblk_idx, ret, mask = nmask) + +@triton.jit +def count_gt(in_ptr0, out_ptr0, cmp_val, dim : tl.constexpr, M : tl.constexpr, N : tl.constexpr, MNUMEL: tl.constexpr, NNUMEL: tl.constexpr): + mblk_idx = tl.arange(0,MNUMEL) + nblk_idx = tl.arange(0,NNUMEL) + mmask = mblk_idx < M + nmask = nblk_idx < N + mask = (mmask[:,None]) & (nmask[None,:]) + idx = mblk_idx[:,None]*N + nblk_idx[None,:] + x = tl.load(in_ptr0+idx, mask = mask, other = 0) + tmp1 = (x > cmp_val) + tmp2 = tmp1.to(tl.float32) + ret = tl.sum(tmp2, dim) + tl.store(out_ptr0 + nblk_idx, ret, mask = nmask) + +@triton.jit +def count_lt(in_ptr0, out_ptr0, cmp_val, dim : tl.constexpr, M : tl.constexpr, N : tl.constexpr, MNUMEL: tl.constexpr, NNUMEL: tl.constexpr): + mblk_idx = tl.arange(0,MNUMEL) + nblk_idx = tl.arange(0,NNUMEL) + mmask = mblk_idx < M + nmask = nblk_idx < N + mask = (mmask[:,None]) & (nmask[None,:]) + idx = mblk_idx[:,None]*N + nblk_idx[None,:] + x = tl.load(in_ptr0+idx, mask = mask, other = 0) + tmp1 = (x < cmp_val) + tmp2 = tmp1.to(tl.float32) + ret = tl.sum(tmp2, dim) + tl.store(out_ptr0 + nblk_idx, ret, mask = nmask) + + +shapes=[ + (57,3,64,16), (57,-32,64,32), (57,37,64,64), + (64,3,64,16), (64,-32,64,32), (64,37,64,64), + (3,3,8,8), (-32,3,32,8), (37,3,64,8), + (3,1,8,8), (-32,1,32,8), (37,1,64,8) +] + +map_for_64_t = {37:(31,32),263:(107,128)} +map_for_32_t = {263:(137,256)} + + +types0 = [ + (torch.int8,'int8'), +] +@pytest.mark.parametrize('dtype, sigtype',types0) +@pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_count_eq_dim0_common(dtype, sigtype, M, N, MNUMEL, NNUMEL): + M = (-M)//torch.tensor(0,dtype=dtype).element_size() if M<0 else M + N = (-N)//torch.tensor(0,dtype=dtype).element_size() if N<0 else N + + if sigtype == 'int64': + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == 'float32' or sigtype == 'bfloat16' or sigtype == 'int32': + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + cmp_val = 8 + x0 = test_common.generate_tensor(shape = (M,N),dtype = sigtype) + ans = standard_count(x0, cmp_val,0, dtype) + x0 = x0.cpu() + print(ans) + output = torch.zeros((N,), dtype = torch.float32).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + count[1,1,1](x0_txda, output_txda, cmp_val, 0, M = M, N = N,MNUMEL = MNUMEL, NNUMEL = NNUMEL, debug = True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + test_common.validate_cmp('float32', output, ans.to(torch.float32)) + +#------------------------------------------------------------------------------------- + +types1 = [ + (torch.float32,'float32'), + (torch.float32,'float16'), + (torch.int8,'int8'), +] +@pytest.mark.parametrize('dtype, sigtype',types1) +@pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_count_gt_dim0_common(dtype, sigtype, M, N, MNUMEL, NNUMEL): + M = (-M)//torch.tensor(0,dtype=dtype).element_size() if M<0 else M + N = (-N)//torch.tensor(0,dtype=dtype).element_size() if N<0 else N + + if sigtype == 'int64': + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == 'float32' or sigtype == 'bfloat16' or sigtype == 'int32': + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + if dtype == torch.int8: + cmp_val = 8 + else: + cmp_val = 0.5 + x0 = test_common.generate_tensor(shape = (M,N),dtype = sigtype) + ans = standard_count_gt(x0, cmp_val,0, dtype) + x0 = x0.cpu() + print(ans) + output = torch.zeros((N,), dtype = torch.float32).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + count_gt[1,1,1](x0_txda, output_txda, cmp_val, 0, M = M, N = N,MNUMEL = MNUMEL, NNUMEL = NNUMEL, debug = True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + test_common.validate_cmp("float32", output, ans.to(torch.float32)) + + +shapes1=[ + (64,3,64,16), (64,-32,64,32), (64,37,64,64) +] +@pytest.mark.parametrize('dtype, sigtype',types1) +@pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes1) +def test_count_lt_dim0_common(dtype, sigtype, M, N, MNUMEL, NNUMEL): + M = (-M)//torch.tensor(0,dtype=dtype).element_size() if M<0 else M + N = (-N)//torch.tensor(0,dtype=dtype).element_size() if N<0 else N + + if sigtype == 'int64': + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == 'float32' or sigtype == 'bfloat16' or sigtype == 'int32': + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + if dtype == torch.int8: + cmp_val = 8 + else: + cmp_val = 0.5 + x0 = test_common.generate_tensor(shape = (M,N),dtype = sigtype) + ans = standard_count_lt(x0, cmp_val,0, dtype) + x0 = x0.cpu() + print(ans) + output = torch.zeros((N,), dtype = torch.float32).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + count_lt[1,1,1](x0_txda, output_txda, cmp_val, 0, M = M, N = N,MNUMEL = MNUMEL, NNUMEL = NNUMEL, debug = True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + test_common.validate_cmp("float32", output, ans.to(torch.float32)) diff --git a/test/wafer/ops/test_count_dim1.py b/test/wafer/ops/test_count_dim1.py new file mode 100644 index 00000000..7ea789da --- /dev/null +++ b/test/wafer/ops/test_count_dim1.py @@ -0,0 +1,222 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +# Copyright (c) Huawei Technologies Co., Ltd. 2025-2025. All rights reserved. +import pytest +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + +def standard_count(x0, cmp_val, dim, dtype): + res = (x0 == cmp_val).sum(dim=dim) + return res + +def standard_count_gt(x0, cmp_val, dim, dtype): + res = (x0 > cmp_val).sum(dim=dim) + return res + +def standard_count_lt(x0, cmp_val, dim, dtype): + res = (x0 < cmp_val).sum(dim=dim) + return res + +@triton.jit +def count(in_ptr0, out_ptr0, cmp_val, dim : tl.constexpr, M : tl.constexpr, N : tl.constexpr, MNUMEL: tl.constexpr, NNUMEL: tl.constexpr): + mblk_idx = tl.arange(0,MNUMEL) + nblk_idx = tl.arange(0,NNUMEL) + mmask = mblk_idx < M + nmask = nblk_idx < N + mask = (mmask[:,None]) & (nmask[None,:]) + idx = mblk_idx[:,None]*N + nblk_idx[None,:] + x = tl.load(in_ptr0+idx, mask = mask, other = 0) + tmp1 = (x == cmp_val) + tmp2 = tmp1.to(tl.float32) + ret = tl.sum(tmp2, dim) + tl.store(out_ptr0 + mblk_idx, ret, mask = mmask) + +@triton.jit +def count_gt(in_ptr0, out_ptr0, cmp_val, dim : tl.constexpr, M : tl.constexpr, N : tl.constexpr, MNUMEL: tl.constexpr, NNUMEL: tl.constexpr): + mblk_idx = tl.arange(0,MNUMEL) + nblk_idx = tl.arange(0,NNUMEL) + mmask = mblk_idx < M + nmask = nblk_idx < N + mask = (mmask[:,None]) & (nmask[None,:]) + idx = mblk_idx[:,None]*N + nblk_idx[None,:] + x = tl.load(in_ptr0+idx, mask = mask, other = 0) + tmp1 = (x > cmp_val) + tmp2 = tmp1.to(tl.float32) + ret = tl.sum(tmp2, dim) + tl.store(out_ptr0 + mblk_idx, ret, mask = mmask) + +@triton.jit +def count_lt(in_ptr0, out_ptr0, cmp_val, dim : tl.constexpr, M : tl.constexpr, N : tl.constexpr, MNUMEL: tl.constexpr, NNUMEL: tl.constexpr): + mblk_idx = tl.arange(0,MNUMEL) + nblk_idx = tl.arange(0,NNUMEL) + mmask = mblk_idx < M + nmask = nblk_idx < N + mask = (mmask[:,None]) & (nmask[None,:]) + idx = mblk_idx[:,None]*N + nblk_idx[None,:] + x = tl.load(in_ptr0+idx, mask = mask, other = 0) + tmp1 = (x < cmp_val) + tmp2 = tmp1.to(tl.float32) + ret = tl.sum(tmp2, dim) + tl.store(out_ptr0 + mblk_idx, ret, mask = mmask) + + +# if shape axis = 32/256 , then actual shape = axis/element_size() + +shapes=[ + (57,3,64,16), (57,-32,64,32), + (64,3,64,16), (64,-32,64,32), + (3,3,8,8), (-32,3,32,8), (37,3,64,8), + (3,1,8,8), (-32,1,32,8), (37,1,64,8) +] + +map_for_64_t = {37:(31,32),263:(107,128)} +map_for_32_t = {263:(137,256)} + + +types0 = [ + (torch.int8,'int8'), +] +@pytest.mark.parametrize('dtype, sigtype',types0) +@pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_count_eq_dim0_common(dtype, sigtype, M, N, MNUMEL, NNUMEL): + M = (-M)//torch.tensor(0,dtype=dtype).element_size() if M<0 else M + N = (-N)//torch.tensor(0,dtype=dtype).element_size() if N<0 else N + + if sigtype == 'int64': + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == 'float32' or sigtype == 'bfloat16' or sigtype == 'int32': + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + cmp_val = 8 + x0 = test_common.generate_tensor(shape = (M,N),dtype = sigtype) + ans = standard_count(x0, cmp_val, 1, dtype) + x0 = x0.cpu() + print(ans) + output = torch.zeros((M,), dtype = torch.float32).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + count[1,1,1](x0_txda, output_txda, cmp_val, 1, M = M, N = N,MNUMEL = MNUMEL, NNUMEL = NNUMEL, debug = True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + test_common.validate_cmp('float32',output,ans.to(torch.float32)) + +#------------------------------------------------------------------------------------- + +types1 = [ + (torch.float32,'float32'), + (torch.float32,'float16'), + (torch.int8,'int8'), +] +@pytest.mark.parametrize('dtype, sigtype',types1) +@pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_count_gt_dim0_common(dtype, sigtype, M, N, MNUMEL, NNUMEL): + M = (-M)//torch.tensor(0,dtype=dtype).element_size() if M<0 else M + N = (-N)//torch.tensor(0,dtype=dtype).element_size() if N<0 else N + + if sigtype == 'int64': + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == 'float32' or sigtype == 'bfloat16' or sigtype == 'int32': + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + if dtype == torch.int8: + cmp_val = 8 + else: + cmp_val = 0.5 + x0 = test_common.generate_tensor(shape = (M,N),dtype = sigtype) + ans = standard_count_gt(x0, cmp_val, 1, dtype) + x0 = x0.cpu() + print(ans) + output = torch.zeros((M,), dtype = torch.float32).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + count_gt[1,1,1](x0_txda, output_txda, cmp_val, 1, M = M, N = N,MNUMEL = MNUMEL, NNUMEL = NNUMEL, debug = True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + test_common.validate_cmp('float32',output,ans.to(torch.float32)) + + + +types2 = [ + (torch.int8,'int8') +] +shapes2=[ + (57,-32,64,32), (64,-32,64,32) +] + +@pytest.mark.parametrize('dtype, sigtype',types2) +@pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes2) +def test_count_lt_dim0_common(dtype, sigtype, M, N, MNUMEL, NNUMEL): + M = (-M)//torch.tensor(0,dtype=dtype).element_size() if M<0 else M + N = (-N)//torch.tensor(0,dtype=dtype).element_size() if N<0 else N + + if sigtype == 'int64': + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == 'float32' or sigtype == 'bfloat16' or sigtype == 'int32': + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + if dtype == torch.int8: + cmp_val = 8 + else: + cmp_val = 0.5 + x0 = test_common.generate_tensor(shape = (M,N),dtype = sigtype) + ans = standard_count_lt(x0, cmp_val, 1, dtype) + x0 = x0.cpu() + print(ans) + output = torch.zeros((M,), dtype = torch.float32).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + count_lt[1,1,1](x0_txda, output_txda, cmp_val, 1, M = M, N = N,MNUMEL = MNUMEL, NNUMEL = NNUMEL, debug = True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + test_common.validate_cmp('float32',output,ans.to(torch.float32)) diff --git a/test/wafer/ops/test_cumprod.py b/test/wafer/ops/test_cumprod.py new file mode 100644 index 00000000..765bfdf1 --- /dev/null +++ b/test/wafer/ops/test_cumprod.py @@ -0,0 +1,108 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + +# from dlblas.utils.libentry import libentry + +from test_common import _all_dtypes_no_bool, validate_cmp + + +def torch_func(x, dim, reverse): + is_bf16 = x.dtype == torch.bfloat16 + if is_bf16: + x = x.to(torch.float32) + if reverse: + x = torch.flip(x, [dim]) + res = torch.cumprod(x, dim=dim) + if is_bf16: + res = res.to(torch.bfloat16) + return res + + +# @libentry() +@triton.jit +def triton_kernel( + out_ptr0, + in_ptr0, + dim: tl.constexpr, + reverse: tl.constexpr, + numel_x: tl.constexpr, + numel_r: tl.constexpr, + XBLOCK: tl.constexpr, + RBLOCK: tl.constexpr, +): + tl.static_assert( + numel_x == XBLOCK, "numel_x must be equal to XBLOCK in this kernel" + ) + tl.static_assert( + numel_r == RBLOCK, "numel_r must be equal to RBLOCK in this kernel" + ) + idx_x = tl.arange(0, XBLOCK) + idx_r = tl.arange(0, RBLOCK) + idx = idx_x[:, None] * numel_r + idx_r[None, :] + x = tl.load(in_ptr0 + idx) + ret = tl.cumprod(x, axis=dim, reverse=reverse) + tl.store(out_ptr0 + idx, ret) + + +def triton_func(x, dim, reverse): + res = torch.empty_like(x) + res_txda = res.to("txda") + x_txda = x.to("txda") + triton_kernel[1, 1, 1]( + res_txda, x_txda, dim, reverse, x_txda.shape[0], x_txda.shape[1], x_txda.shape[0], x_txda.shape[1] + ) + with torch.no_grad(): + res.copy_(res_txda.cpu()) + return res + + +def cumprod_generate_tensor(shape, dtype): + if dtype == "float32" or dtype == "float16" or dtype == "bfloat16": + return torch.rand(size=shape, dtype=eval("torch." + dtype)) + elif dtype == "int32" or dtype == "int64" or dtype == "int16": + return torch.randint(low=0, high=3, size=shape, dtype=eval("torch." + dtype)) + elif dtype == "int8": + return torch.randint(low=0, high=3, size=shape, dtype=eval("torch." + dtype)) + else: + raise ValueError(f"Unsupported dtype: {dtype}") + + +# dtype=int8, reverse=True not support; +not_support_dtype = {"int8", "bool"} +support_dtypes = [ + dtype for dtype in _all_dtypes_no_bool if dtype not in not_support_dtype +] + + +@pytest.mark.parametrize("dtype", support_dtypes) +@pytest.mark.parametrize("shape", [(8, 32)]) +@pytest.mark.parametrize("dim", [0, 1]) +@pytest.mark.parametrize("reverse", [False]) +def test_cumprod(dtype, shape, dim, reverse): + x0 = cumprod_generate_tensor(shape=shape, dtype=dtype).cpu() + triton_cal = triton_func(x0, dim, reverse) + torch_ref = torch_func(x0, dim, reverse) + validate_cmp(dtype, torch_ref, triton_cal) diff --git a/test/wafer/ops/test_cumsum.py b/test/wafer/ops/test_cumsum.py new file mode 100644 index 00000000..f61b4c45 --- /dev/null +++ b/test/wafer/ops/test_cumsum.py @@ -0,0 +1,95 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + +# from dlblas.utils.libentry import libentry + +from test_common import _all_dtypes_no_bool, generate_tensor, validate_cmp + + +def torch_func(x, dim, reverse): + if reverse: + x = torch.flip(x, [dim]) + res = torch.cumsum(x, dim=dim) + return res + + +# @libentry() +@triton.jit +def triton_kernel( + out_ptr0, + in_ptr0, + dim: tl.constexpr, + reverse: tl.constexpr, + numel_x: tl.constexpr, + numel_r: tl.constexpr, + XBLOCK: tl.constexpr, + RBLOCK: tl.constexpr, +): + tl.static_assert( + numel_x == XBLOCK, "numel_x must be equal to XBLOCK in this kernel" + ) + tl.static_assert( + numel_r == RBLOCK, "numel_r must be equal to RBLOCK in this kernel" + ) + idx_x = tl.arange(0, XBLOCK) + idx_r = tl.arange(0, RBLOCK) + idx = idx_x[:, None] * numel_r + idx_r[None, :] + x = tl.load(in_ptr0 + idx) + ret = tl.cumsum(x, axis=dim, reverse=reverse) + tl.store(out_ptr0 + idx, ret) + + +def triton_func(x, dim, reverse): + res = torch.empty_like(x) + res_txda = res.to("txda") + x_txda = x.to("txda") + triton_kernel[1, 1, 1]( + res_txda, x_txda, dim, reverse, x_txda.shape[0], x_txda.shape[1], x_txda.shape[0], x_txda.shape[1] + ) + with torch.no_grad(): + res.copy_(res_txda.cpu()) + return res + + +# dtype=int8, reverse=True not support; +not_support_dtype = {"int8", "bool"} +support_dtypes = [ + dtype for dtype in _all_dtypes_no_bool if dtype not in not_support_dtype +] + + +@pytest.mark.parametrize("dtype", support_dtypes) +@pytest.mark.parametrize("shape", [(8, 32)]) +@pytest.mark.parametrize("dim", [0, 1]) +@pytest.mark.parametrize("reverse", [False]) +def test_cumsum(dtype, shape, dim, reverse): + x0 = generate_tensor(shape=shape, dtype=dtype).cpu() + triton_cal = triton_func(x0, dim, reverse) + torch_dtype = eval("torch." + dtype) + if torch_dtype == torch.float16 or torch_dtype == torch.float32: + x0 = x0.to(torch.float32) + torch_ref = torch_func(x0, dim, reverse).to(torch_dtype) + validate_cmp(dtype, torch_ref, triton_cal) diff --git a/test/wafer/ops/test_debug_barrier.py b/test/wafer/ops/test_debug_barrier.py new file mode 100644 index 00000000..a5e49cdd --- /dev/null +++ b/test/wafer/ops/test_debug_barrier.py @@ -0,0 +1,67 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + +def torch_pointwise(x0, x1): + res = x0 - x1 + return res + + +@triton.jit +def triton_sub(in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 - tmp1 + tl.debug_barrier() + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 2, 32768, 1024], + ] + ) + +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype = eval('torch.' + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_sub[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_device_print.py b/test/wafer/ops/test_device_print.py new file mode 100644 index 00000000..14dcf53f --- /dev/null +++ b/test/wafer/ops/test_device_print.py @@ -0,0 +1,155 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import pytest +import test_common +import os + + +os.environ["TRITON_DEVICE_PRINT"] = "1" +os.environ["TRITON_ENABLE_TASKQUEUE"] = "0" +shape = (8,) +XS = 8 + + +def _device_print_values(dtype): + # Exercise printing, not overflow conversion: scalar assignment rejects + # values outside the destination dtype before the device kernel can run. + if dtype.is_floating_point: + limits = torch.finfo(dtype) + return [0.0, limits.eps, limits.tiny, 1.0, -1.0, limits.min, limits.max, 0.5] + limits = torch.iinfo(dtype) + return [0, limits.min, limits.min + 1, -1, 1, limits.max - 1, limits.max, 2] + + +def torch_func(x0, x1): + res = x0 + x1 + return res + + +@triton.jit +def triton_kernel(out_ptr0, in_ptr0, in_ptr1, XBLOCK: tl.constexpr): + idx = tl.arange(0, XBLOCK) + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.load(in_ptr1 + idx) + tmp2 = tmp0 + tmp1 + tl.device_print("OUTPUT = ", tmp2) + tl.store(out_ptr0 + idx, tmp2) + + +def triton_func(x0, x1, XS): + out = torch.empty_like(x0) + out_txda = out.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_kernel[1, 1, 1](out_txda, x0_txda, x1_txda, XS) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + return out + + +@pytest.mark.skip(reason="waiting for bishengir-compile to support") +@pytest.mark.parametrize("sigtype", ["int64"]) +def test_device_print_int64(capsys, sigtype): + dtype = eval(f"torch.{sigtype}") + x0 = torch.zeros(shape, dtype=dtype).cpu() + x1 = torch.ones(shape, dtype=dtype).cpu() + for i, value in enumerate(_device_print_values(dtype)): + x1[i] = value + torch_ref = torch_func(x0, x1) + triton_cal = triton_func(x0, x1, XS) + test_common.validate_cmp(sigtype, triton_cal, torch_ref) + + +@pytest.mark.parametrize("sigtype", ["int32"]) +def test_device_print_int32(capsys, sigtype): + dtype = eval(f"torch.{sigtype}") + x0 = torch.zeros(shape, dtype=dtype).cpu() + x1 = torch.ones(shape, dtype=dtype).cpu() + for i, value in enumerate(_device_print_values(dtype)): + x1[i] = value + torch_ref = torch_func(x0, x1) + triton_cal = triton_func(x0, x1, XS) + test_common.validate_cmp(sigtype, triton_cal, torch_ref) + + +@pytest.mark.parametrize("sigtype", ["int16"]) +def test_device_print_int16(capsys, sigtype): + dtype = eval(f"torch.{sigtype}") + x0 = torch.zeros(shape, dtype=dtype).cpu() + x1 = torch.ones(shape, dtype=dtype).cpu() + for i, value in enumerate(_device_print_values(dtype)): + x1[i] = value + torch_ref = torch_func(x0, x1) + triton_cal = triton_func(x0, x1, XS) + test_common.validate_cmp(sigtype, triton_cal, torch_ref) + + +@pytest.mark.parametrize("sigtype", ["int8"]) +def test_device_print_int8(capsys, sigtype): + dtype = eval(f"torch.{sigtype}") + x0 = torch.zeros(shape, dtype=dtype).cpu() + x1 = torch.ones(shape, dtype=dtype).cpu() + for i, value in enumerate(_device_print_values(dtype)): + x1[i] = value + torch_ref = torch_func(x0, x1) + triton_cal = triton_func(x0, x1, XS) + test_common.validate_cmp(sigtype, triton_cal, torch_ref) + + +@pytest.mark.parametrize("sigtype", ["float32"]) +def test_device_print_fp32(capsys, sigtype): + dtype = eval(f"torch.{sigtype}") + x0 = torch.zeros(shape, dtype=dtype).cpu() + x1 = torch.ones(shape, dtype=dtype).cpu() + for i, value in enumerate(_device_print_values(dtype)): + x1[i] = value + torch_ref = torch_func(x0, x1) + triton_cal = triton_func(x0, x1, XS) + test_common.validate_cmp(sigtype, triton_cal, torch_ref) + + +@pytest.mark.parametrize("sigtype", ["float16"]) +def test_device_print_fp16(capsys, sigtype): + dtype = eval(f"torch.{sigtype}") + x0 = torch.zeros(shape, dtype=dtype).cpu() + x1 = torch.ones(shape, dtype=dtype).cpu() + for i, value in enumerate(_device_print_values(dtype)): + x1[i] = value + torch_ref = torch_func(x0, x1) + triton_cal = triton_func(x0, x1, XS) + test_common.validate_cmp(sigtype, triton_cal, torch_ref) + + +@pytest.mark.skip(reason="waiting for bishengir-compile to support") +@pytest.mark.parametrize("sigtype", ["bfloat16"]) +def test_device_print_bf16(capsys, sigtype): + dtype = eval(f"torch.{sigtype}") + x0 = torch.zeros(shape, dtype=dtype).cpu() + x1 = torch.ones(shape, dtype=dtype).cpu() + for i, value in enumerate(_device_print_values(dtype)): + x1[i] = value + torch_ref = torch_func(x0, x1) + triton_cal = triton_func(x0, x1, XS) + test_common.validate_cmp(sigtype, triton_cal, torch_ref) diff --git a/test/wafer/ops/test_div.py b/test/wafer/ops/test_div.py new file mode 100644 index 00000000..2c87e48d --- /dev/null +++ b/test/wafer/ops/test_div.py @@ -0,0 +1,71 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + +def torch_pointwise(x0, x1): + res = x0 / x1 + return res + + +@triton.jit +def triton_div(in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 / tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 2, 32768, 1024], + ['float16', (2, 4096, 8), 2, 32768, 1024], + ['int8', (2, 4096, 8), 2, 32768, 1024], + ] + ) + +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + if dtype == 'int8': + dtype = 'float32' + x1 = x1.masked_fill(x1 == 0, 1) + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype = eval('torch.' + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_div[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_elementwise_ceil.py b/test/wafer/ops/test_elementwise_ceil.py new file mode 100644 index 00000000..2ee218c9 --- /dev/null +++ b/test/wafer/ops/test_elementwise_ceil.py @@ -0,0 +1,91 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time +import test_common +import os +import shutil + +import torch +import torch_txda # noqa: F401 + + +def standard_ceil(x0): + res = torch.ceil(x0) + return res + + +@triton.jit +def triton_ceil(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.ceil(x) + tl.store(out_ptr0 + idx_block, res, mask=mask) + + +types = [ + (torch.float32, 'float32'), +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + +ops = [ + ('ceil', triton_ceil, standard_ceil), +] + + +@pytest.mark.parametrize('opName, tritonOp, standOp', ops) +@pytest.mark.parametrize('dtype, sigtype', types) +@pytest.mark.parametrize('N, NUMEL', shapes) +def test_elementwise_common(opName, tritonOp, standOp, dtype, sigtype, N, NUMEL): + torch.manual_seed(0) + torch.txda.set_device(0) + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == 'int64': + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standOp(x0) + x0 = x0.cpu() + + output = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + tritonOp[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_elementwise_clip.py b/test/wafer/ops/test_elementwise_clip.py new file mode 100644 index 00000000..e80e8177 --- /dev/null +++ b/test/wafer/ops/test_elementwise_clip.py @@ -0,0 +1,89 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time +import test_common +import os +import shutil + +import torch +import torch_txda # noqa: F401 + +def standard_clamp(x0): + res = torch.clamp(x0, min=-10, max=10) + return res +@triton.jit +def triton_clamp(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.clamp(x, -10, 10) + tl.store(out_ptr0 + idx_block, res, mask=mask) + +types = [ + (torch.float32, 'float32'), + (torch.float16, 'float16'), + (torch.bfloat16, 'bfloat16'), +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + +ops = [ + ('clamp', triton_clamp, standard_clamp), +] + + +@pytest.mark.parametrize('opName, tritonOp, standOp', ops) +@pytest.mark.parametrize('dtype, sigtype', types) +@pytest.mark.parametrize('N, NUMEL', shapes) +def test_elementwise_common(opName, tritonOp, standOp, dtype, sigtype, N, NUMEL): + torch.manual_seed(0) + torch.txda.set_device(0) + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == 'int64': + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standOp(x0) + x0 = x0.cpu() + + output = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + tritonOp[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_elementwise_f2i.py b/test/wafer/ops/test_elementwise_f2i.py new file mode 100644 index 00000000..3b218ae1 --- /dev/null +++ b/test/wafer/ops/test_elementwise_f2i.py @@ -0,0 +1,133 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time +import test_common +import os +import shutil + +import torch +import torch_txda # noqa: F401 + +def standard_f2i32(x0): + res = x0.to(torch.int32) + return res + +def standard_f2i8(x0): + res = x0.to(torch.int8) + return res + +def standard_f2i16(x0): + res = x0.to(torch.int16) + return res + +def standard_f2i64(x0): + res = x0.to(torch.int64) + return res + +@triton.jit +def triton_f2i8(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.cast(x, tl.int8) + tl.store(out_ptr0 + idx_block, res, mask=mask) + +@triton.jit +def triton_f2i16(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.cast(x, tl.int16) + tl.store(out_ptr0 + idx_block, res, mask=mask) + +@triton.jit +def triton_f2i32(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.cast(x, tl.int32) + tl.store(out_ptr0 + idx_block, res, mask=mask) + +@triton.jit +def triton_f2i64(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.cast(x, tl.int64) + tl.store(out_ptr0 + idx_block, res, mask=mask) + +types = [ + (torch.float32, 'float32'), + (torch.float16, 'float16'), + (torch.bfloat16, 'bfloat16'), +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), +] + +map_for_64_t = {37: 31} + +ops = [ + ('f2i8', triton_f2i8, standard_f2i8, 'int8'), + ('f2i16', triton_f2i16, standard_f2i16, 'int16'), + ('f2i32', triton_f2i32, standard_f2i32, 'int32'), + ('f2i64', triton_f2i64, standard_f2i64, 'int64'), +] + + +def continue_func(opName, d_type): + if 'f2i' in opName and 'int' in d_type: + return True + + +@pytest.mark.parametrize('opName, tritonOp, standOp, dst_sigtype', ops) +@pytest.mark.parametrize('dtype, sigtype', types) +@pytest.mark.parametrize('N, NUMEL', shapes) +def test_elementwise_common(opName, tritonOp, standOp, dst_sigtype, dtype, sigtype, N, NUMEL): + if continue_func(opName, sigtype): + return + + torch.txda.set_device(0) + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == 'int64': + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standOp(x0) + x0 = x0.cpu() + + output = test_common.generate_tensor(shape=(N,), dtype=dst_sigtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + tritonOp[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + test_common.validate_cmp(dst_sigtype, output, ans) diff --git a/test/wafer/ops/test_elementwise_floor.py b/test/wafer/ops/test_elementwise_floor.py new file mode 100644 index 00000000..40f9230d --- /dev/null +++ b/test/wafer/ops/test_elementwise_floor.py @@ -0,0 +1,91 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time +import test_common +import os +import shutil + +import torch +import torch_txda # noqa: F401 + + +def standard_floor(x0): + res = torch.floor(x0) + return res + + +@triton.jit +def triton_floor(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.floor(x) + tl.store(out_ptr0 + idx_block, res, mask=mask) + + +types = [ + (torch.float32, 'float32'), +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + +ops = [ + ('floor', triton_floor, standard_floor), +] + + +@pytest.mark.parametrize('opName, tritonOp, standOp', ops) +@pytest.mark.parametrize('dtype, sigtype', types) +@pytest.mark.parametrize('N, NUMEL', shapes) +def test_elementwise_common(opName, tritonOp, standOp, dtype, sigtype, N, NUMEL): + torch.manual_seed(0) + torch.txda.set_device(0) + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == 'int64': + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standOp(x0) + x0 = x0.cpu() + + output = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + tritonOp[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_elementwise_i2f.py b/test/wafer/ops/test_elementwise_i2f.py new file mode 100644 index 00000000..eb8e3b2e --- /dev/null +++ b/test/wafer/ops/test_elementwise_i2f.py @@ -0,0 +1,128 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time +import test_common +import os +import shutil + +import torch +import torch_txda # noqa: F401 + + +def standard_i2f_float32(x0): + res = x0.to(torch.float32) + return res + + +def standard_i2f_float16(x0): + res = x0.to(torch.float16) + return res + + +def standard_i2f_bfloat16(x0): + res = x0.to(torch.bfloat16) + return res + + +@triton.jit +def triton_i2f_float32(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.cast(x, tl.float32) + tl.store(out_ptr0 + idx_block, res, mask=mask) + + +@triton.jit +def triton_i2f_float16(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.cast(x, tl.float16) + tl.store(out_ptr0 + idx_block, res, mask=mask) + + +@triton.jit +def triton_i2f_bfloat16(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.cast(x, tl.bfloat16) + tl.store(out_ptr0 + idx_block, res, mask=mask) + + +types = [ + # (torch.int8, 'int8'), # TO BE FIXED i8 -> f16、bf16 + # (torch.int16, 'int16'), # TO BE FIXED i16 -> f32、bf16 + (torch.int32, 'int32'), # TO BE FIXED i32 -> f16、bf16 + # (torch.int64, 'int64'), # TO BE FIXED i64 -> bf16 +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), +] + +map_for_64_t = {37: 31} + +ops = [ + # ('i2f16', triton_i2f_float16, standard_i2f_float16, 'float16'), + ('i2f32', triton_i2f_float32, standard_i2f_float32, 'float32'), + # ('i2fbf16', triton_i2f_bfloat16, standard_i2f_bfloat16, 'bfloat16'), +] + + +def continue_func(opName, d_type): + if 'i2f' in opName and 'float' in d_type: + return True + + +@pytest.mark.parametrize('opName, tritonOp, standOp, dst_sigtype', ops) +@pytest.mark.parametrize('dtype, sigtype', types) +@pytest.mark.parametrize('N, NUMEL', shapes) +def test_elementwise_common(opName, tritonOp, standOp, dst_sigtype, dtype, sigtype, N, NUMEL): + if continue_func(opName, sigtype): + return + + torch.txda.set_device(0) + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == 'int64': + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standOp(x0) + x0 = x0.cpu() + + output = test_common.generate_tensor(shape=(N,), dtype=dst_sigtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + tritonOp[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + test_common.validate_cmp(dst_sigtype, output, ans) diff --git a/test/wafer/ops/test_elementwise_round.py b/test/wafer/ops/test_elementwise_round.py new file mode 100644 index 00000000..3a510d2f --- /dev/null +++ b/test/wafer/ops/test_elementwise_round.py @@ -0,0 +1,95 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import test_common +import os +import shutil + +import torch +import torch_txda # noqa: F401 + + +def standard_round_to_nearest_neighbor_even(x0): # TO BE FIXED round to nearest neighbor even + res = torch.round(x0) + return res + +def standard_round(x0): + res = torch.where(x0 >= 0, torch.floor(x0 + 0.5), torch.ceil(x0 - 0.5)) + return res + +@triton.jit +def triton_round(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + mask = idx_block < N + x = tl.load(in_ptr0 + idx_block, mask=mask) + res = tl.where(x >= 0, tl.floor(x.to(tl.float32) + 0.5), tl.ceil(x.to(tl.float32) - 0.5)) + tl.store(out_ptr0 + idx_block, res, mask=mask) + + +types = [ + (torch.float32, 'float32'), + (torch.float16, 'float16'), + (torch.bfloat16, 'bfloat16'), +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + +ops = [ + ('round', triton_round, standard_round), +] + + +@pytest.mark.parametrize('opName, tritonOp, standOp', ops) +@pytest.mark.parametrize('dtype, sigtype', types) +@pytest.mark.parametrize('N, NUMEL', shapes) +def test_elementwise_common(opName, tritonOp, standOp, dtype, sigtype, N, NUMEL): + torch.manual_seed(0) + torch.txda.set_device(0) + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == 'int64': + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standOp(x0) + x0 = x0.cpu() + + output = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + tritonOp[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_eq.py b/test/wafer/ops/test_eq.py new file mode 100644 index 00000000..95a640c3 --- /dev/null +++ b/test/wafer/ops/test_eq.py @@ -0,0 +1,86 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_binary(x0, y0): + res = x0 == y0 + return res + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x == y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16, 'bfloat16'), + (torch.int8, "int8"), + (torch.int16, "int16"), + (torch.int32, "int32"), + (torch.int64, "int64"), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + y0 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + ans = standard_binary(x0, y0) + out = torch.zeros((N,), dtype=torch.bool).cpu() + x0_txda = x0.to("txda") + y0_txda = y0.to("txda") + out_txda = out.to("txda") + triton_elementwise_binary[1, 1, 1](x0_txda, y0_txda, out_txda, N, NUMEL) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_eq_2.py b/test/wafer/ops/test_eq_2.py new file mode 100644 index 00000000..a246a2cf --- /dev/null +++ b/test/wafer/ops/test_eq_2.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_eq(x0, x1): + return x0 == x1 + + +@triton.jit +def triton_eq(in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index, None) + tmp1 = tl.load(in_ptr1 + x_index, None) + tmp2 = tmp0 == tmp1 + tl.store(out_ptr0 + x_index, tmp2, None) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 2, 32768, 1024], + ]) +def test_eq(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_eq(x0, x1).to(eval('torch.' + dtype)) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval('torch.' + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_eq[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_exp.py b/test/wafer/ops/test_exp.py new file mode 100644 index 00000000..fefce7f5 --- /dev/null +++ b/test/wafer/ops/test_exp.py @@ -0,0 +1,63 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + +def torch_pointwise(x0): + res = torch.exp(x0.to(torch.float64)).to(x0.dtype) + return res + + +@triton.jit +def triton_exp(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp2 = tl.exp(tmp0) + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 2, 32768, 1024], + ] + ) + +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0) + y_cal = torch.zeros(shape, dtype = eval('torch.' + dtype)).cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_exp[ncore, 1, 1](x0_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_exp2.py b/test/wafer/ops/test_exp2.py new file mode 100644 index 00000000..5f54ee15 --- /dev/null +++ b/test/wafer/ops/test_exp2.py @@ -0,0 +1,64 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_exp2(x0): + res = torch.pow(2, x0, out=None) + return res + + +@triton.jit +def triton_exp2(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index, None) + tmp1 = tl.exp2(tmp0) + tl.store(out_ptr0 + x_index, tmp1, None) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 2, 32768, 1024], + ]) +def test_exp2(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_exp2(x0) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval('torch.' + dtype)).cpu() + x0_txda = x0.to("txda") + triton_res_txda = triton_res.to("txda") + triton_exp2[ncore, 1, 1](x0_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_exp_.py b/test/wafer/ops/test_exp_.py new file mode 100644 index 00000000..d128471b --- /dev/null +++ b/test/wafer/ops/test_exp_.py @@ -0,0 +1,106 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_unary(x0, dtype): + res = torch.exp(x0) + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + ret = tl.exp(x) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, "float32"), + # Expected dtype ['fp32', 'fp64'] + # (torch.float16, 'float16'), + # (torch.bfloat16, 'bfloat16'), + # (torch.int8, 'int8'), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_expand_dims.py b/test/wafer/ops/test_expand_dims.py new file mode 100644 index 00000000..c2b6a1d6 --- /dev/null +++ b/test/wafer/ops/test_expand_dims.py @@ -0,0 +1,76 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest + +@triton.jit +def fn_npu_(output_ptr, x_ptr,XB : tl.constexpr,YB : tl.constexpr,ZB : tl.constexpr): + xidx=tl.arange(0,XB) + yidx=tl.arange(0,YB) + zidx=tl.arange(0,ZB) + + idx=xidx[:,None,None]*YB*ZB+yidx[None,:,None]*ZB+zidx[None,None,:] + + X = tl.load(x_ptr+idx) + + ret = tl.expand_dims(X,2) + + oidx=xidx[:,None,None,None]*YB*ZB+yidx[None,:,None,None]*ZB+tl.arange(0,1)[None,None,:,None]+zidx[None,None,None,:] + + tl.store(output_ptr+oidx,ret) + +paras = [ + ('*fp32',eval('torch.float32'),2,256,16), + ('*fp32',eval('torch.float32'),8,8,4), + ('*fp16',eval('torch.float16'),2,256,16), + ('*fp16',eval('torch.float16'),8,8,4), + ('*i8',eval('torch.int8'),2,256,16), + ('*i8',eval('torch.int8'),8,8,4), +] + +@pytest.mark.parametrize('para_type,data_type,XB,YB,ZB', paras) +def test_npu(para_type,data_type,XB,YB,ZB): + + x = torch.randint(low=-128,high=128,size=(XB,YB,ZB),dtype=data_type).cpu() + a = x.unsqueeze(2) + + print(f"shape = {x.shape}") + print(x.dtype) + print(a[0,0:16,0,0]) + + output = torch.randint(1, (XB,YB,1,ZB), dtype=data_type).cpu() + + print(f"output.dtype={output.dtype}") + + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_[1,1,1](output_txda,x_txda, XB=XB, YB=YB, ZB=ZB, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output[0,0:16,0,0]) + + torch.testing.assert_close(output,a) \ No newline at end of file diff --git a/test/wafer/ops/test_extract_slice.py b/test/wafer/ops/test_extract_slice.py new file mode 100644 index 00000000..a1fbd0ef --- /dev/null +++ b/test/wafer/ops/test_extract_slice.py @@ -0,0 +1,44 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +from triton.experimental.tle.language import dsa as dl + + +@triton.jit +def triton_kernel(x_ptr, y_ptr, output_ptr, n_elements, BLOCK_SIZE: tl.constexpr): + pid = tl.program_id(axis=0) + block_start = pid * BLOCK_SIZE + offsets = block_start + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_elements + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + output = x + y + out_sub = dl.extract_slice(output, [block_start], [32], [1]) + out_idx = block_start + tl.arange(0, 32) + out_msk = out_idx < n_elements + tl.store(output_ptr + out_idx, out_sub, mask=out_msk) + + +def triton_func(x: torch.Tensor, y: torch.Tensor): + output = torch.empty_like(x) + n_elements = output.numel() + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + triton_kernel[grid](x_txda, y_txda, output_txda, n_elements, BLOCK_SIZE=1024) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + return output + + +def test_extract_slice(): + size = 1024 + x = torch.rand(size, device="cpu") + y = torch.rand(size, device="cpu") + torch_ref = x + y + triton_cal = triton_func(x, y) + print("max diff", (triton_cal[:32] - torch_ref[:32]).abs().max()) + torch.testing.assert_close(triton_cal[:32], torch_ref[:32]) diff --git a/test/wafer/ops/test_fdiv.py b/test/wafer/ops/test_fdiv.py new file mode 100644 index 00000000..6e52a5fa --- /dev/null +++ b/test/wafer/ops/test_fdiv.py @@ -0,0 +1,69 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_fdiv(x0, x1): + res = x0 / x1 + return res + + +@triton.jit +def triton_fdiv(in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + tmp0 = tl.load(in_ptr0 + x_index) + tmp1 = tl.load(in_ptr1 + x_index) + tmp2 = tl.fdiv(tmp0, tmp1) + tl.store(out_ptr0 + x_index, tmp2) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 2, 32768, 1024], + ['float16', (2, 4096, 8), 2, 32768, 1024], + ]) +def test_fdiv(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + y_tmp = test_common.generate_tensor(shape, dtype) + y0 = y_tmp.masked_fill(y_tmp == 0, 1) + y0 = y0.cpu() + + # torch结果 + y_ref = torch_fdiv(x0, y0).to(eval('torch.' + dtype)) + # triton结果 + y_cal = torch.zeros(shape, dtype=eval('torch.' + dtype)).cpu() + x0_txda = x0.to("txda") + y0_txda = y0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_fdiv[ncore, 1, 1](x0_txda, y0_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_floor.py b/test/wafer/ops/test_floor.py new file mode 100644 index 00000000..d0e0447c --- /dev/null +++ b/test/wafer/ops/test_floor.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_floor(x0, x1): + res = x0 + torch.floor(x1) + return res + + +@triton.jit +def triton_floor(in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = x_index < xnumel + tmp0 = tl.load(in_ptr0 + x_index, xmask) + tmp1 = tl.load(in_ptr1 + x_index, xmask) + tmp2 = tmp0 + tl.floor(tmp1) + tl.store(out_ptr0 + x_index, tmp2, xmask) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 2, 32768, 1024], + ]) +def test_floor(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + y_ref = torch_floor(x0, x1) + # triton结果 + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_floor[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_floordiv.py b/test/wafer/ops/test_floordiv.py new file mode 100644 index 00000000..0462757f --- /dev/null +++ b/test/wafer/ops/test_floordiv.py @@ -0,0 +1,75 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_func(x0, x1): + res = x0 // x1 + return res + + +@triton.jit +def triton_kernel(out_ptr0, in_ptr0, in_ptr1, N: tl.constexpr): + idx = tl.arange(0, N) + x = tl.load(in_ptr0 + idx) + y = tl.load(in_ptr1 + idx) + ret = x // y + tl.store(out_ptr0 + idx, ret) + + +def triton_func(x0, x1, N): + out = torch.empty_like(x0) + out_txda = out.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_kernel[1, 1, 1](out_txda, x0_txda, x1_txda, N) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + return out + + +types = [ + "int32", +] + +shapes = [ + 4, + 16, + 256, + 1024, +] + +@pytest.mark.parametrize("sigtype", types) +@pytest.mark.parametrize("N", shapes) +def test_floordiv(sigtype, N): + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + x1 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + x1 = x1.masked_fill(x1 == 0, 1) + + torch_ref = torch_func(x0, x1) + triton_cal = triton_func(x0, x1, N) + test_common.validate_cmp(sigtype, triton_cal, torch_ref) + diff --git a/test/wafer/ops/test_full.py b/test/wafer/ops/test_full.py new file mode 100644 index 00000000..7215f462 --- /dev/null +++ b/test/wafer/ops/test_full.py @@ -0,0 +1,97 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 +import pytest + +@triton.jit +def fn_npu_f32(output_ptr,XB : tl.constexpr,YB : tl.constexpr,ZB : tl.constexpr): + xidx=tl.arange(0,XB) + yidx=tl.arange(0,YB) + zidx=tl.arange(0,ZB) + + ret = tl.full((XB,YB,ZB),value = 100,dtype = tl.float32) + + oidx=xidx[:,None,None]*YB*ZB+yidx[None,:,None]*ZB+zidx[None,None,:] + + tl.store(output_ptr+oidx,ret) + +@triton.jit +def fn_npu_f16(output_ptr,XB : tl.constexpr,YB : tl.constexpr,ZB : tl.constexpr): + xidx=tl.arange(0,XB) + yidx=tl.arange(0,YB) + zidx=tl.arange(0,ZB) + + ret = tl.full((XB,YB,ZB),value = 100,dtype = tl.float16) + + oidx=xidx[:,None,None]*YB*ZB+yidx[None,:,None]*ZB+zidx[None,None,:] + + tl.store(output_ptr+oidx,ret) + +@triton.jit +def fn_npu_i8(output_ptr,XB : tl.constexpr,YB : tl.constexpr,ZB : tl.constexpr): + xidx=tl.arange(0,XB) + yidx=tl.arange(0,YB) + zidx=tl.arange(0,ZB) + + ret = tl.full((XB,YB,ZB),value = 100,dtype = tl.int8) + + oidx=xidx[:,None,None]*YB*ZB+yidx[None,:,None]*ZB+zidx[None,None,:] + + tl.store(output_ptr+oidx,ret) + +testlist = [ + (fn_npu_f32,'float32',torch.float32,2,256,16), + (fn_npu_f32,'float32',torch.float32,8,8,4), + + (fn_npu_f16,'float16',torch.float16,2,256,16), + (fn_npu_f16,'float16',torch.float16,8,8,4), + + (fn_npu_i8,'int8',torch.int8,2,256,16), + (fn_npu_i8,'int8',torch.int8,8,8,4), +] + +@pytest.mark.parametrize('testfunc, sigtype, dtype, XB, YB, ZB',testlist) +def test_npu(testfunc, sigtype, dtype, XB, YB, ZB): + + x = torch.full((XB,YB,ZB),100,dtype=dtype).cpu() + + print(f"shape = {x.shape}") + print(x.dtype) + print(x[0,0:16,0]) + + output = torch.randint(1, (XB,YB,ZB), dtype=dtype).cpu() + + print(f"output.dtype={output.dtype}") + + output_txda = output.to("txda") + testfunc[1,1,1](output_txda,XB,YB,ZB,debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output[0,0:16,0]) + + test_common.validate_cmp(sigtype,output,x) diff --git a/test/wafer/ops/test_ge.py b/test/wafer/ops/test_ge.py new file mode 100644 index 00000000..183ffea9 --- /dev/null +++ b/test/wafer/ops/test_ge.py @@ -0,0 +1,86 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_binary(x0, y0): + res = x0 >= y0 + return res + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x >= y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16, 'bfloat16'), + (torch.int8, "int8"), + (torch.int16, "int16"), + (torch.int32, "int32"), + (torch.int64, "int64"), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + y0 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + ans = standard_binary(x0, y0) + out = torch.zeros((N,), dtype=torch.bool).cpu() + x0_txda = x0.to("txda") + y0_txda = y0.to("txda") + out_txda = out.to("txda") + triton_elementwise_binary[1, 1, 1](x0_txda, y0_txda, out_txda, N, NUMEL) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_ge_2.py b/test/wafer/ops/test_ge_2.py new file mode 100644 index 00000000..49db5843 --- /dev/null +++ b/test/wafer/ops/test_ge_2.py @@ -0,0 +1,69 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_ge(x0, x1, dtype): + res = torch.where(torch.ge(x0, x1), torch.ones_like(x0), torch.zeros_like(x0)).to(eval('torch.' + dtype)) + return res + + +@triton.jit +def triton_ge(in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 >= tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize('param_list', + [ + ['float16', (2, 4096, 8), 2, 32768, 1024], + ['float32', (2, 4096, 8), 2, 32768, 1024], + ['int8', (2, 4096, 8), 2, 32768, 1024], + ]) +def test_ge(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_ge(x0, x1, dtype) + # triton结果 + triton_res = torch.empty_like(x0) + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_ge[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_gelu.py b/test/wafer/ops/test_gelu.py new file mode 100644 index 00000000..2d8df5df --- /dev/null +++ b/test/wafer/ops/test_gelu.py @@ -0,0 +1,102 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + +def standard_unary(x0, dtype): + res = x0 * 0.5 * (1.0 + torch.erf(x0 / torch.sqrt(torch.tensor(2.0)))) + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + ret = x * 0.5 * (1.0 + tl.erf(x / tl.sqrt(2.0))) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary(in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, 'float32'), + # (torch.float16, 'float16'), + # (torch.bfloat16, 'bfloat16'), + # (torch.int8, 'int8'), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize('dtype,sigtype', types) +@pytest.mark.parametrize('N,NUMEL', shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_gt.py b/test/wafer/ops/test_gt.py new file mode 100644 index 00000000..ce1c8617 --- /dev/null +++ b/test/wafer/ops/test_gt.py @@ -0,0 +1,67 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_gt(x0, x1, dtype): + res = torch.where(torch.gt(x0, x1), torch.ones_like(x0), torch.zeros_like(x0)).to(eval('torch.' + dtype)) + return res + + +@triton.jit +def triton_gt(in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 > tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 4096, 8), 2, 32768, 1024], + ]) +def test_gt(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_gt(x0, x1, dtype) + # triton结果 + triton_res = torch.empty_like(x0) + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_gt[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_hd_permute.py b/test/wafer/ops/test_hd_permute.py new file mode 100644 index 00000000..4c21d2a9 --- /dev/null +++ b/test/wafer/ops/test_hd_permute.py @@ -0,0 +1,65 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 + +X_SIZE = tl.constexpr(4) +Y_SIZE = tl.constexpr(64) +Z_SIZE = tl.constexpr(32) +NUMEL = tl.constexpr(X_SIZE.value * Y_SIZE.value * Z_SIZE.value) + + +def torch_permute(x): + return ( + x.reshape((X_SIZE, Y_SIZE, Z_SIZE)) + .permute(1, 0, 2) + .reshape((X_SIZE * Y_SIZE * Z_SIZE)) + ) + + +@triton.jit +def triton_permute(output_ptr, input_ptr): + x_index = tl.arange(0, X_SIZE * Y_SIZE * Z_SIZE) + input_local = tl.load(input_ptr + x_index) + output_local = ( + input_local.reshape((X_SIZE, Y_SIZE, Z_SIZE)) + .permute(1, 0, 2) + .reshape((X_SIZE * Y_SIZE * Z_SIZE)) + ) + tl.store(output_ptr + x_index, output_local) + + +def test_hd_permute(): + # 生成数据 + x = torch.randn(NUMEL).cpu() + # torch结果 + torch_res = torch_permute(x) + # triton结果 + triton_res = torch.randn(torch_res.shape, dtype=torch_res.dtype).cpu() + triton_res_txda = triton_res.to("txda") + x_txda = x.to("txda") + triton_permute[1, 1, 1](triton_res_txda, x_txda) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + torch.testing.assert_close(triton_res, torch_res, rtol=1e-3, atol=1e-3) diff --git a/test/wafer/ops/test_if_tensor.py b/test/wafer/ops/test_if_tensor.py new file mode 100644 index 00000000..897f7a3c --- /dev/null +++ b/test/wafer/ops/test_if_tensor.py @@ -0,0 +1,61 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def if_tensor_kernel( + kv_start_idx, # tensor + output_ptr, +): + pid = tl.program_id(0) + if kv_start_idx is not None: + value = tl.load(kv_start_idx + pid) + tl.store(output_ptr + pid, value) + + +# 测试函数 +def test_kernel(): + n = 8 + device = "cpu" + + kv_start_idx = torch.arange(n, dtype=torch.float32, device=device) + output1 = torch.zeros(n, dtype=torch.float32, device=device) + kv_start_idx_txda = kv_start_idx.to("txda") + output1_txda = output1.to("txda") + if_tensor_kernel[(n,)]( + kv_start_idx_txda, + output1_txda, + ) + with torch.no_grad(): + output1.copy_(output1_txda.cpu()) + + expected = torch.arange(n, dtype=torch.float32, device=device) + assert torch.allclose(output1, expected), f"Output {output1} != Expected {expected}" + print(f"RESULT: output1 = {output1}") + print("✅ Test passed!") + + +if __name__ == "__main__": + test_kernel() diff --git a/test/wafer/ops/test_insert_slice.py b/test/wafer/ops/test_insert_slice.py new file mode 100644 index 00000000..812eb521 --- /dev/null +++ b/test/wafer/ops/test_insert_slice.py @@ -0,0 +1,60 @@ +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +from triton.experimental.tle.language import dsa as dl + + +@triton.jit +def triton_kernel( + x_ptr, + y_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, + SLICE_OFFSET: tl.constexpr, + SLICE_SIZE: tl.constexpr, +): + pid = tl.program_id(axis=0) + block_start = pid * BLOCK_SIZE + offsets = block_start + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_elements + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + x_sub = dl.extract_slice(x, [block_start + SLICE_OFFSET], [SLICE_SIZE], [1]) + y_sub = dl.extract_slice(y, [block_start + SLICE_OFFSET], [SLICE_SIZE], [1]) + output_sub = x_sub + y_sub + output = tl.load(output_ptr + offsets, mask=mask) + output = dl.insert_slice( + output, output_sub, [block_start + SLICE_OFFSET], [SLICE_SIZE], [1] + ) + tl.store(output_ptr + offsets, output, mask=mask) + + +def triton_func(x: torch.Tensor, y: torch.Tensor, slice_offset: int, slice_size: int): + output = torch.empty_like(x) + n_elements = output.numel() + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + triton_kernel[grid]( + x_txda, y_txda, output_txda, n_elements, BLOCK_SIZE=1024, SLICE_OFFSET=0, SLICE_SIZE=32 + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + return output + + +def test_insert_slice(): + size = 1024 + slice_offset = 0 + slice_size = 32 + x = torch.rand(size, device="cpu") + y = torch.rand(size, device="cpu") + torch_ref = x + y + triton_cal = triton_func(x, y, slice_offset, slice_size) + torch.testing.assert_close( + triton_cal[slice_offset : slice_offset + slice_size], + torch_ref[slice_offset : slice_offset + slice_size], + ) diff --git a/test/wafer/ops/test_interleave.py b/test/wafer/ops/test_interleave.py new file mode 100644 index 00000000..1c966ffd --- /dev/null +++ b/test/wafer/ops/test_interleave.py @@ -0,0 +1,89 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_( + output_ptr, x_ptr, y_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr +): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + Y = tl.load(y_ptr + idx) + + ret = tl.interleave(X, Y) + + oidx = ( + xidx[:, None, None] * YB * ZB * 2 + + yidx[None, :, None] * ZB * 2 + + tl.arange(0, 2 * ZB)[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@pytest.mark.parametrize( + "para_type,data_type,XB,YB,ZB", + [ + ["float32", torch.float32, 2, 64, 16], + ["float32", torch.float32, 8, 8, 4], + ["float16", torch.float16, 2, 64, 16], + ["float16", torch.float16, 8, 8, 4], + ["int8", torch.int8, 2, 64, 32], + ["int8", torch.int8, 8, 8, 4], + ], +) +def test_interleave(para_type, data_type, XB, YB, ZB): + + x = torch.full((XB, YB, ZB), 100, dtype=data_type).cpu() + y = torch.full((XB, YB, ZB), 30, dtype=data_type).cpu() + + print(f"shape = {x.shape}") + print(x.dtype) + + output = torch.randint(1, (XB, YB, ZB * 2), dtype=data_type).cpu() + output1 = output + print(f"output.dtype={output.dtype}") + + ans = torch.stack((x, y), dim=-1).reshape(XB, YB, ZB * 2) + print(ans) + print(ans.shape) + + output_txda = output.to("txda") + x_txda = x.to("txda") + y_txda = y.to("txda") + fn_npu_[1, 1, 1](output_txda, x_txda, y_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + test_common.validate_cmp(para_type, ans, output) diff --git a/test/wafer/ops/test_invert.py b/test/wafer/ops/test_invert.py new file mode 100644 index 00000000..e5d82a88 --- /dev/null +++ b/test/wafer/ops/test_invert.py @@ -0,0 +1,72 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_invert(x0): + res = ~(x0) + return res + + +@triton.jit +def triton_invert( + in_ptr0, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = xindex < xnumel + x0 = xindex + tmp0 = tl.load(in_ptr0 + (x0), xmask) + tmp2 = ~tmp0 + tl.store(out_ptr0 + (xindex), tmp2, xmask) + + +@pytest.mark.parametrize( + "param_list", + [ + ["int8", (2, 4096, 8), 2, 32768, 1024], + ["int16", (2, 4096, 8), 2, 32768, 1024], + ["int32", (2, 4096, 8), 2, 32768, 1024], + ["int64", (2, 4096, 8), 2, 32768, 1024], + ["bool", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_invert(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_invert(x0) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + triton_res_txda = triton_res.to("txda") + triton_invert[ncore, 1, 1](x0_txda, triton_res_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_join.py b/test/wafer/ops/test_join.py new file mode 100644 index 00000000..4b3e517f --- /dev/null +++ b/test/wafer/ops/test_join.py @@ -0,0 +1,299 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_( + output_ptr, x_ptr, y_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr +): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + + idx = xidx[:, None] * YB + yidx[None, :] + + X = tl.load(x_ptr + idx) + Y = tl.load(y_ptr + idx) + + ret = tl.join(X, Y) + + oidx = ( + xidx[:, None, None] * YB * 2 + + yidx[None, :, None] * 2 + + tl.arange(0, 2)[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@pytest.mark.parametrize( + "para_type,data_type,XB,YB,ZB", + [ + ["float32", torch.float32, 4, 64, 4], + ["float32", torch.float32, 8, 8, 4], + ["float16", torch.float16, 4, 64, 4], + ["float16", torch.float16, 8, 8, 4], + ["int8", torch.int8, 4, 128, 4], + ["int8", torch.int8, 8, 8, 4], + ], +) +def test_join(para_type, data_type, XB, YB, ZB): + x = torch.full((XB, YB), 100, dtype=data_type).cpu() + y = torch.full((XB, YB), 30, dtype=data_type).cpu() + + ans = torch.stack((x, y), dim=-1) + print(ans) + + output = torch.randint(1, (XB, YB, 2), dtype=data_type).cpu() + output_txda = output.to("txda") + x_txda = x.to("txda") + y_txda = y.to("txda") + fn_npu_[1, 1, 1](output_txda, x_txda, y_txda, XB, YB, ZB, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + print(output) + test_common.validate_cmp(para_type, ans, output) + + +@triton.jit +def fn_npu_concat_axis_( + output_ptr, + x_ptr, + y_ptr, + XB: tl.constexpr, + YB: tl.constexpr, + ZB: tl.constexpr, + axis: tl.constexpr, # 0 / 1 / 2 +): + """ + 新增 kernel:将两个 shape=(XB, YB, ZB) 的张量沿 axis 拼接到 output。 + 只新增,不修改原有 fn_npu_。 + 说明:这个 kernel 假定输入张量是连续的、按行主序扁平化。 + """ + + # 构造三维索引 (i,j,k) -> linear index i*(YB*ZB) + j*ZB + k + xidx = tl.arange(0, XB) # (XB,) + yidx = tl.arange(0, YB) # (YB,) + zidx = tl.arange(0, ZB) # (ZB,) + + # idx shape: (XB, YB, ZB) + idx = ( + xidx[:, None, None] * (YB * ZB) + yidx[None, :, None] * ZB + zidx[None, None, :] + ) + + # 从输入加载完整块 + X = tl.load(x_ptr + idx) + Y = tl.load(y_ptr + idx) + + # 根据 axis 计算输出偏移并存储 + if axis == 0: + # out shape: (XB*2, YB, ZB) + oidx_x = idx # X 放在前半段 + oidx_y = ( + (xidx + XB)[:, None, None] * (YB * ZB) + + yidx[None, :, None] * ZB + + zidx[None, None, :] + ) + tl.store(output_ptr + oidx_x, X) + tl.store(output_ptr + oidx_y, Y) + + elif axis == 1: + # out shape: (XB, YB*2, ZB) + # 线性化为 i*(YB*2*ZB) + j*(ZB) + k + base_x = xidx[:, None, None] * (YB * 2 * ZB) + oidx_x = base_x + yidx[None, :, None] * ZB + zidx[None, None, :] + oidx_y = base_x + (yidx + YB)[None, :, None] * ZB + zidx[None, None, :] + tl.store(output_ptr + oidx_x, X) + tl.store(output_ptr + oidx_y, Y) + + elif axis == 2: + # out shape: (XB, YB, ZB*2) + # 线性化为 i*(YB*(ZB*2)) + j*(ZB*2) + k + base_x = xidx[:, None, None] * (YB * (ZB * 2)) + yidx[None, :, None] * (ZB * 2) + oidx_x = base_x + zidx[None, None, :] + oidx_y = base_x + (zidx + ZB)[None, None, :] + tl.store(output_ptr + oidx_x, X) + tl.store(output_ptr + oidx_y, Y) + + else: + # 不支持其它 axis,早退出 + return + + +@triton.jit +def fn_npu_concat_axis_tiled_( + output_ptr, + x_ptr, + y_ptr, + XB: tl.constexpr, + YB: tl.constexpr, + ZB: tl.constexpr, + axis: tl.constexpr, # 0/1/2 + TI: tl.constexpr, # tile size for i (XB dim) + TJ: tl.constexpr, # tile size for j (YB dim) + TK: tl.constexpr, # tile size for k (ZB dim) +): + """ + 分块 kernel:每个 program 处理一个 tile,支持沿 axis 拼接。 + - 输入 x,y 形状为 (XB, YB, ZB) + - 输出为拼接后的形状 (根据 axis) + - TI/TJ/TK 为每个方向的 tile 大小(constexpr) + """ + + # block id per dimension + bid_i = tl.program_id(0) + bid_j = tl.program_id(1) + bid_k = tl.program_id(2) + + # tile origin indices + i0 = bid_i * TI + j0 = bid_j * TJ + k0 = bid_k * TK + + # ranges within tile (实际长度考虑边界) + i_range = tl.arange(0, TI) + j_range = tl.arange(0, TJ) + k_range = tl.arange(0, TK) + + # compute actual masks for boundaries + ii = i0 + i_range # shape (TI,) + jj = j0 + j_range # (TJ,) + kk = k0 + k_range # (TK,) + + # masks whether indices are in bounds + mask_i = ii < XB + mask_j = jj < YB + mask_k = kk < ZB + + # create 3D grid of indices using broadcasting + # shapes: ii[:,None,None], jj[None,:,None], kk[None,None,:] + idx = ( + ii[:, None, None] * (YB * ZB) + jj[None, :, None] * ZB + kk[None, None, :] + ) # shape (TI, TJ, TK) but some entries out-of-bounds + + # combined mask for load/store (True where within X/Y/Z bounds) + mask = mask_i[:, None, None] & mask_j[None, :, None] & mask_k[None, None, :] + + # load X and Y with mask (out-of-bounds read returns undefined, so mask=False avoids) + X = tl.load(x_ptr + idx, mask=mask, other=0) + Y = tl.load(y_ptr + idx, mask=mask, other=0) + + # Now compute output base depending on axis + if axis == 0: + # out shape: (XB*2, YB, ZB) + # linear index for output: i*(YB*ZB) + j*ZB + k, but for Y we shift i by XB + oidx_x = idx # write X to i + oidx_y = ( + (ii[:, None, None] + XB) * (YB * ZB) + + jj[None, :, None] * ZB + + kk[None, None, :] + ) + # store with mask (only store valid positions) + tl.store(output_ptr + oidx_x, X, mask=mask) + tl.store(output_ptr + oidx_y, Y, mask=mask) + + elif axis == 1: + # out shape: (XB, YB*2, ZB) + # linear index: i*(YB*2*ZB) + j*(ZB) + k + base = ii[:, None, None] * (YB * 2 * ZB) + oidx_x = base + jj[None, :, None] * ZB + kk[None, None, :] + oidx_y = base + (jj[None, :, None] + YB) * ZB + kk[None, None, :] + tl.store(output_ptr + oidx_x, X, mask=mask) + tl.store(output_ptr + oidx_y, Y, mask=mask) + + elif axis == 2: + # out shape: (XB, YB, ZB*2) + # linear index: i*(YB*(ZB*2)) + j*(ZB*2) + k + base = ii[:, None, None] * (YB * (ZB * 2)) + jj[None, :, None] * (ZB * 2) + oidx_x = base + kk[None, None, :] + oidx_y = base + (kk[None, None, :] + ZB) + tl.store(output_ptr + oidx_x, X, mask=mask) + tl.store(output_ptr + oidx_y, Y, mask=mask) + + else: + return + + +def _ceil_div(a, b): + return (a + b - 1) // b + + +@pytest.mark.parametrize( + "para_type,data_type,XB,YB,ZB,axis", + [ + ["float32", torch.float32, 4, 4, 4, 0], + ["float32", torch.float32, 512, 256, 512, 0], + ["float32", torch.float32, 512, 256, 512, 1], + ["float32", torch.float32, 512, 256, 512, 2], + ["float16", torch.float16, 2, 8, 4, 1], + ["int8", torch.int8, 4, 8, 2, 2], + ], +) +def test_join_axis_added_tiled(para_type, data_type, XB, YB, ZB, axis): + """ + 增强测试:对小输入使用 grid=(1,1,1)(调用简单 kernel),对大输入使用 tiling kernel 并计算 grid。 + - 该测试仅新增,不改动原来的 fn_npu_ 或其他测试。 + """ + + x = torch.full((XB, YB, ZB), 100, dtype=data_type).cpu() + y = torch.full((XB, YB, ZB), 30, dtype=data_type).cpu() + ans = torch.cat((x, y), dim=axis) + + out_shape = list(x.shape) + out_shape[axis] = out_shape[axis] * 2 + output = torch.randint(1, tuple(out_shape), dtype=data_type).cpu() + + LARGE_THRESHOLD = 16 + + if max(XB, YB, ZB) <= LARGE_THRESHOLD: + output_txda = output.to("txda") + x_txda = x.to("txda") + y_txda = y.to("txda") + fn_npu_concat_axis_[1, 1, 1](output_txda, x_txda, y_txda, XB, YB, ZB, axis, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + else: + TI, TJ, TK = 16, 16, 16 + + # 计算每维需要多少 tile + gx = _ceil_div(XB, TI) + gy = _ceil_div(YB, TJ) + gz = _ceil_div(ZB, TK) + + # launch tiled kernel with grid (gx, gy, gz) + output_txda = output.to("txda") + x_txda = x.to("txda") + y_txda = y.to("txda") + fn_npu_concat_axis_tiled_[gx, gy, gz]( + output_txda, x_txda, y_txda, XB, YB, ZB, axis, TI, TJ, TK, debug=True + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + test_common.validate_cmp(para_type, ans, output) diff --git a/test/wafer/ops/test_lanzcos.py b/test/wafer/ops/test_lanzcos.py new file mode 100644 index 00000000..7245384d --- /dev/null +++ b/test/wafer/ops/test_lanzcos.py @@ -0,0 +1,284 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import numpy as np +import math +import pytest + + +@triton.jit +def lanczos_resize_kernel( + img_src_ptr, + img_dst_ptr, + img_coeffs_ptr, + src_rows, + src_cols, + dst_rows, + dst_cols, + R_H, + R_W, + C, + stride_in_h, + stride_in_w, + stride_in_c, + stride_out_h, + stride_out_w, + stride_out_c, + BLOCK_SIZE: tl.constexpr, +): + block_id_c = tl.program_id(0) + block_id_h = tl.program_id(1) + block_id_w = tl.program_id(2) + dest_h_offs = block_id_h * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + dest_w_offs = block_id_w * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + dest_offs = ( + block_id_c[None, None] * stride_out_c + + dest_h_offs[:, None] * stride_out_h + + dest_w_offs[None, :] * stride_out_w + ) + + RR_H = 1.0 / R_H + RR_W = 1.0 / R_W + + fy = (dest_h_offs + 0.5) * RR_H - 0.5 + sy = tl.floor(fy) + fx = (dest_w_offs + 0.5) * RR_W - 0.5 + sx = tl.floor(fx) + + idxY = tl.floor((fy - sy) * 24.999999).to(tl.int32) + idxX = tl.floor((fx - sx) * 24.999999).to(tl.int32) + tableIndex = idxY[:, None] * 25 + idxX[None, :] + res = tl.zeros((BLOCK_SIZE, BLOCK_SIZE), tl.float32) + + for ii in range(4): + for jj in range(4): + src_offsets = ( + block_id_c[None, None] * stride_in_c + + (tl.clamp((sy + ii - 1), 0, src_rows - 1)).to(tl.int32)[:, None] + * stride_in_h + + (tl.clamp((sx + jj - 1), 0, src_cols - 1)).to(tl.int32)[None, :] + * stride_in_w + ) + src_val = tl.load(img_src_ptr + src_offsets) + coeffs_offs = tableIndex[:, :] * 16 + (ii * 4 + jj)[None, None] + coeffs = tl.load(img_coeffs_ptr + coeffs_offs) + res = res + src_val * coeffs + dst_mask = (dest_h_offs[:, None] < dst_rows) & (dest_w_offs[None, :] < dst_cols) + res = tl.clamp(res, 0.0, 1.0) + tl.store(img_dst_ptr + dest_offs, res, mask=dst_mask) + + +def lanczos_resize_triton(img_src, img_dst, c_lanczosCoeffs, dst_rows, dst_cols): + N, C, src_rows, src_cols = img_src.shape + R_H = float(dst_rows) / src_rows + R_W = float(dst_cols) / src_cols + + stride_in_n, stride_in_c, stride_in_h, stride_in_w = img_src.stride() + stride_out_n, stride_out_c, stride_out_h, stride_out_w = img_dst.stride() + BLOCK_SIZE = 16 + grid = lambda meta: ( + C, + triton.cdiv(dst_rows, meta["BLOCK_SIZE"]), + triton.cdiv(dst_cols, meta["BLOCK_SIZE"]), + ) + img_src_txda = img_src.to("txda") + img_dst_txda = img_dst.to("txda") + c_lanczosCoeffs_txda = c_lanczosCoeffs.to("txda") + lanczos_resize_kernel[grid]( + img_src_txda, + img_dst_txda, + c_lanczosCoeffs_txda, + src_rows, + src_cols, + dst_rows, + dst_cols, + R_H, + R_W, + C, + stride_in_h, + stride_in_w, + stride_in_c, + stride_out_h, + stride_out_w, + stride_out_c, + BLOCK_SIZE=BLOCK_SIZE, + ) + with torch.no_grad(): + img_dst.copy_(img_dst_txda.cpu()) + return img_dst + + +def lanczos_resize_cpu(img_src, img_dst, img_coeffs, dst_rows, dst_cols): + N, C, src_rows, src_cols = img_src.shape + R_H = float(dst_rows) / src_rows + R_W = float(dst_cols) / src_cols + for i in range(dst_rows): + for j in range(dst_cols): + RR_H = 1.0 / R_H + RR_W = 1.0 / R_W + fy = (i + 0.5) * RR_H - 0.5 + sy = math.floor(fy) + fx = (j + 0.5) * RR_W - 0.5 + sx = math.floor(fx) + idxY = math.floor((fy - np.floor(fy)) * 24.999999) + idxX = math.floor((fx - np.floor(fx)) * 24.999999) + tableIndex = idxY * 25 + idxX + res = (0.0, 0.0, 0.0, 0.0) + for ii in range(4): + for jj in range(4): + idx_y = np.clip(sy + ii - 1, 0, src_rows - 1) + idx_x = np.clip(sx + jj - 1, 0, src_cols - 1) + src_val = img_src[0, :, idx_y, idx_x] + coeffs_offs = tableIndex * 16 + (ii * 4 + jj) + coeffs = img_coeffs[coeffs_offs] + res = res + src_val * coeffs + + img_dst[0, :, i, j] = np.clip(res, 0.0, 1.0) + + +@pytest.mark.parametrize( + "shapes", + [ + [360, 640, 140, 280], + ], +) +def test_lanzcos(shapes): + c_lanczosCoeffs = torch.randn(10000, dtype=torch.float32, device="cpu") / 4.0 + src_rows, src_cols, dst_rows, dst_cols = shapes + img_src = torch.randn(1, 4, src_rows, src_cols, dtype=torch.float32, device="cpu") + img_dst = torch.zeros( + (1, img_src.shape[1], dst_rows, dst_cols), + dtype=img_src.dtype, + device=img_src.device, + ) + resized_image = lanczos_resize_triton( + img_src, img_dst, c_lanczosCoeffs, dst_rows, dst_cols + ) + img_src_cpu = img_src.cpu().numpy() + img_dst_cpu = torch.zeros( + (1, img_src_cpu.shape[1], dst_rows, dst_cols), dtype=img_src.dtype, device="cpu" + ).numpy() + lanczos_resize_cpu( + img_src_cpu, img_dst_cpu, c_lanczosCoeffs.cpu().numpy(), dst_rows, dst_cols + ) + torch.testing.assert_close( + resized_image.cpu(), torch.from_numpy(img_dst_cpu), atol=1.0 / 255, rtol=0 + ) + + +def benchmark_test( + fn_ref, fn_triton, ref_args=(), triton_args=(), name="gen_fn", times=10, repeat=10 +): + import time + + print( + f"--------------------benchmark_{name} for {times * repeat} times--------------------" + ) + stream = torch.txda.current_stream() + # warm_up + stream.synchronize() + for _ in range(10): + fn_triton(*triton_args) + stream.synchronize() + + start = time.perf_counter() + for _ in range(times * repeat): + fn_triton(*triton_args) + stream.synchronize() + end = time.perf_counter() + + time_compiled = (end - start) / (times * repeat) + time_compiled *= 1000000 + + # warm_up + stream.synchronize() + for _ in range(10): + std = fn_ref(*ref_args) + stream.synchronize() + + start = time.perf_counter() + for _ in range(times * repeat): + std = fn_ref(*ref_args) + stream.synchronize() + end = time.perf_counter() + time_eager = (end - start) / (times * repeat) + time_eager *= 1000000 + + accelerated = (time_eager - time_compiled) / time_compiled * 100 + print( + f"Accelerated: {accelerated:.4f}% eager takes {time_eager:.3f} us, triton takes {time_compiled:.3f} us" + ) + + return accelerated, time_eager, time_compiled + + +if __name__ == "__main__": + c_lanczosCoeffs = torch.randn(10000, dtype=torch.float32, device="cpu") / 4.0 + + src_rows, src_cols = 360, 640 + dst_rows, dst_cols = 140, 280 + img_src = torch.randn(1, 4, src_rows, src_cols, dtype=torch.float32, device="cpu") + + print("==========run npu===============") + img_dst = torch.zeros( + (1, img_src.shape[1], dst_rows, dst_cols), + dtype=img_src.dtype, + device=img_src.device, + ) + resized_image = lanczos_resize_triton( + img_src, img_dst, c_lanczosCoeffs, dst_rows, dst_cols + ) + resized_cpu = resized_image.cpu().numpy() + print("==========run cpu===============") + img_src_cpu = img_src.cpu().numpy() + img_dst_cpu = torch.zeros( + (1, img_src_cpu.shape[1], dst_rows, dst_cols), dtype=img_src.dtype, device="cpu" + ).numpy() + lanczos_resize_cpu( + img_src_cpu, img_dst_cpu, c_lanczosCoeffs.cpu().numpy(), dst_rows, dst_cols + ) + + print("==========compare result===============") + diff = np.abs(resized_cpu - img_dst_cpu) + max_diff_value = np.max(diff) + print("max diff float = ", max_diff_value) + print("max diff * 255 int = ", int(max_diff_value * 255)) + torch.testing.assert_close( + resized_image.cpu(), torch.from_numpy(img_dst_cpu), atol=1.0 / 255, rtol=0 + ) + + print("==========profiling===============") + accelerate, eager_time, triton_time = benchmark_test( + lanczos_resize_cpu, + lanczos_resize_triton, + ref_args=( + img_src_cpu, + img_dst_cpu, + c_lanczosCoeffs.cpu().numpy(), + dst_rows, + dst_cols, + ), + triton_args=(img_src, img_dst, c_lanczosCoeffs, dst_rows, dst_cols), + name="lanzcos", + ) diff --git a/test/wafer/ops/test_launcher_empty_signature.py b/test/wafer/ops/test_launcher_empty_signature.py new file mode 100644 index 00000000..f18cc5bd --- /dev/null +++ b/test/wafer/ops/test_launcher_empty_signature.py @@ -0,0 +1,37 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import os +import pytest + +import triton +import triton.language as tl + + +@triton.jit +def _empty_kernel(): + return + + +@pytest.mark.interpreter +def test_launcher_empty_signature(): + grid = (1,) + _empty_kernel[grid]() + assert True diff --git a/test/wafer/ops/test_layernorm.py b/test/wafer/ops/test_layernorm.py new file mode 100644 index 00000000..707bccb6 --- /dev/null +++ b/test/wafer/ops/test_layernorm.py @@ -0,0 +1,146 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import torch +import triton +import triton.language as tl +import torch_txda # noqa: F401 + + +@triton.jit +def _layer_norm_fwd_fused( + X, # pointer to the input + Y, # pointer to the output + W, # pointer to the weights + B, # pointer to the biases + Mean, # pointer to the mean + Rstd, # pointer to the 1/std + stride, # how much to increase the pointer when moving by 1 row + N, # number of columns in X + eps, # epsilon to avoid division by zero + BLOCK_SIZE: tl.constexpr, +): + # Map the program id to the row of X and Y it should compute. + row = tl.program_id(0) + Y += row * stride + X += row * stride + # Compute mean + mean = 0 + _mean = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + a = tl.load(X + cols, mask=cols < N, other=0.0).to(tl.float32) + _mean += a + mean = tl.sum(_mean, axis=0) / N + # Compute variance + _var = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + x = tl.load(X + cols, mask=cols < N, other=0.0).to(tl.float32) + x = tl.where(cols < N, x - mean, 0.0) + _var += x * x + var = tl.sum(_var, axis=0) / N + rstd = 1 / tl.sqrt(var + eps) + # Write mean / rstd + tl.store(Mean + row, mean) + tl.store(Rstd + row, rstd) + # Normalize and apply linear transformation + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + mask = cols < N + w = tl.load(W + cols, mask=mask) + b = tl.load(B + cols, mask=mask) + x = tl.load(X + cols, mask=mask, other=0.0).to(tl.float32) + x_hat = (x - mean) * rstd + y = x_hat * w + b + # Write output + tl.store(Y + cols, y, mask=mask) + + +@torch.inference_mode() +def layer_norm(x, normalized_shape, weight, bias, eps=1e-5): + # allocate output + y = torch.empty_like(x) + # reshape input data into 2D tensor + x_arg = x.reshape(-1, x.shape[-1]) + M, N = x_arg.shape + mean = torch.empty((M,), dtype=torch.float32, device=x.device) + rstd = torch.empty((M,), dtype=torch.float32, device=x.device) + # Less than 64KB per feature: enqueue fused kernel + MAX_FUSED_SIZE = 65536 // x.element_size() + BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(N)) + # heuristics for number of warps + num_warps = min(max(BLOCK_SIZE // 256, 1), 8) + # enqueue kernel + x_arg_txda = x_arg.to("txda") + y_txda = y.to("txda") + weight_txda = weight.to("txda") + bias_txda = bias.to("txda") + mean_txda = mean.to("txda") + rstd_txda = rstd.to("txda") + kernel = _layer_norm_fwd_fused[(M,)]( # + x_arg_txda, + y_txda, + weight_txda, + bias_txda, + mean_txda, + rstd_txda, # + x_arg_txda.stride(0), + N, + eps, # + BLOCK_SIZE=BLOCK_SIZE, + num_warps=num_warps, + num_ctas=1, + ) + with torch.no_grad(): + y.copy_(y_txda.cpu()) + mean.copy_(mean_txda.cpu()) + rstd.copy_(rstd_txda.cpu()) + # print(kernel.asm['ttir']) + return y + + +def _layer_norm(M, N, dtype, eps=1e-5, device="cpu"): + # create data + x_shape = (M, N) + w_shape = (x_shape[-1],) + weight = torch.rand(w_shape, dtype=dtype, device=device, requires_grad=True) + bias = torch.rand(w_shape, dtype=dtype, device=device, requires_grad=True) + x = -2.3 + 0.5 * torch.randn(x_shape, dtype=dtype, device=device) + dy = 0.1 * torch.randn_like(x) + x.requires_grad_(True) + # forward pass + y_tri = layer_norm(x, w_shape, weight, bias, eps) + y_ref = torch.nn.functional.layer_norm(x, w_shape, weight, bias, eps).to(dtype) + # compare + assert torch.allclose(y_tri, y_ref, atol=1e-2, rtol=0) + print(f"layernorm {M},{N} {dtype} passed") + + +def test_layernorm(): + _layer_norm(128, 128, torch.float16) + _layer_norm(128, 128, torch.bfloat16) + _layer_norm(128, 128, torch.float32) + + # _layer_norm(128, 3, torch.bfloat16) + # _layer_norm(128, 16, torch.bfloat16) + # _layer_norm(128, 37, torch.bfloat16) + # _layer_norm(128, 781, torch.bfloat16) diff --git a/test/wafer/ops/test_ldst.py b/test/wafer/ops/test_ldst.py new file mode 100644 index 00000000..cfc96efc --- /dev/null +++ b/test/wafer/ops/test_ldst.py @@ -0,0 +1,650 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import triton.language.math as tl_math +import pytest +import test_common +import random + + +def test_ldst_indirect_00(): + + @triton.jit + def triton_ldst_indirect_00_kernel( + out_ptr0, in_ptr0, in_ptr1, OFFSET0: tl.constexpr, XS: tl.constexpr + ): + pid = tl.program_id(0) + offset1 = tl.load(in_ptr0 + OFFSET0) + idx_in1 = offset1 + pid * XS + tl.arange(0, XS) + tmp0 = tl.load(in_ptr1 + idx_in1) + tmp1 = tl.exp(tmp0) + idx_out0 = pid * XS + tl.arange(0, XS) + tl.store(out_ptr0 + idx_out0, tmp1) + + def triton_ldst_indirect_00_func(x0, x1, s, xs): + n = x1.numel() + ns = n - s + assert ns == xs, "test only single core" + y0 = torch.empty((ns,), dtype=x1.dtype, device=x1.device) + y0_txda = y0.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_ldst_indirect_00_kernel[ns // xs, 1, 1](y0_txda, x0_txda, x1_txda, OFFSET0=s, XS=xs) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_00_func(x0, x1, s): + offset = x0[s] + return torch.exp(x1[offset:]) + + DEV = "cpu" + DTYPE = torch.float32 + offset = 0 + N0, N1 = 16, 16 + blocksize = 16 + assert N0 > offset, "offset must be < N0" + N1 = N1 + offset + x0 = torch.arange(0, N0, dtype=torch.int32, device=DEV) + x1 = torch.randn((N1,), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_00_func(x0, x1, offset) + triton_cal = triton_ldst_indirect_00_func(x0, x1, offset, blocksize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_01(): + + @triton.jit + def triton_ldst_indirect_01_kernel( + out_ptr0, in_ptr0, in_ptr1, OFFSET0: tl.constexpr, XS: tl.constexpr + ): + pid = tl.program_id(0) + offset1 = tl.load(in_ptr0 + OFFSET0) + idx_in1 = offset1 + pid * XS + tl.arange(0, XS) + tmp0 = tl.load(in_ptr1 + idx_in1) + tmp1 = tl_math.exp(tmp0) + idx_out0 = pid * XS + tl.arange(0, XS) + tl.store(out_ptr0 + idx_out0, tmp1) + + def triton_ldst_indirect_01_func(x0, x1, s, xs): + n = x1.numel() + ns = n - s + assert ns == xs, "test only single core" + y0 = torch.empty((ns,), dtype=x1.dtype, device=x1.device) + y0_txda = y0.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_ldst_indirect_01_kernel[ns // xs, 1, 1](y0_txda, x0_txda, x1_txda, OFFSET0=s, XS=xs) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_01_func(x0, x1, s): + offset = x0[s] + return torch.exp(x1[offset:]) + + DEV = "cpu" + DTYPE = torch.float32 + offset = 0 + N0, N1 = 16, 16 + blocksize = 16 + assert N0 > offset, "offset must be < N0" + N1 = N1 + offset + x0 = torch.arange(0, N0, device=DEV) # int64 + x1 = torch.randn((N1,), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_01_func(x0, x1, offset) + triton_cal = triton_ldst_indirect_01_func(x0, x1, offset, blocksize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_02(): + + @triton.jit + def triton_ldst_indirect_02_kernel(out_ptr0, in_ptr0, in_ptr1, XS: tl.constexpr): + pid = tl.program_id(0) + for i in tl.range(0, XS): + tmp0 = tl.load(in_ptr0 + i) + tmp1 = tl.load(in_ptr1 + tmp0) + tmp2 = tl_math.exp(tmp1) + tl.store(out_ptr0 + i, tmp2) + + def triton_ldst_indirect_02_func(x0, x1, xs): + n0 = x0.numel() + assert n0 == xs, "test only single core" + y0 = torch.empty((n0,), dtype=x1.dtype, device=x1.device) + y0_txda = y0.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_ldst_indirect_02_kernel[n0 // xs, 1, 1](y0_txda, x0_txda, x1_txda, XS=xs) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_02_func(x0, x1): + return torch.exp(x1[x0]) + + DEV = "cpu" + DTYPE = torch.float32 + offset = 8 + N0, N1 = 16, 32 + blocksize = 16 + assert N1 >= N0 + offset, "N1 must be >= N0+offset" + assert N0 == blocksize, "N0 must be == blocksize" + x0 = offset + torch.arange(0, N0, device=DEV) # int64 + x1 = torch.randn((N1,), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_02_func(x0, x1) + triton_cal = triton_ldst_indirect_02_func(x0, x1, blocksize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_03(): + + @triton.jit + def triton_ldst_indirect_03_kernel(out_ptr0, in_ptr0, in_ptr1, XS: tl.constexpr): + pid = tl.program_id(0) + in_idx0 = pid * XS + tl.arange(0, XS) + tmp0 = tl.load(in_ptr0 + in_idx0) + tmp1 = tl.load(in_ptr1 + tmp0) + tmp2 = tl_math.exp(tmp1) + out0_idx = pid * XS + tl.arange(0, XS) + tl.store(out_ptr0 + out0_idx, tmp2) + + def triton_ldst_indirect_03_func(x0, x1, xs): + n0 = x0.numel() + assert n0 == xs, "test only single core" + y0 = torch.empty((n0,), dtype=x1.dtype, device=x1.device) + y0_txda = y0.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_ldst_indirect_03_kernel[n0 // xs, 1, 1](y0_txda, x0_txda, x1_txda, XS=xs) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_03_func(x0, x1): + return torch.exp(x1[x0]) + + DEV = "cpu" + DTYPE = torch.float32 + offset = 8 + N0, N1 = 16, 32 + blocksize = 16 + assert N1 >= N0 + offset, "N1 must be >= N0+offset" + assert N0 == blocksize, "N0 must be == blocksize" + x0 = offset + torch.arange(0, N0, device=DEV) # int64 + x1 = torch.randn((N1,), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_03_func(x0, x1) + triton_cal = triton_ldst_indirect_03_func(x0, x1, blocksize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_04(): + + @triton.jit + def triton_ldst_indirect_04_kernel(out_ptr0, in_ptr0, in_ptr1, XS: tl.constexpr): + pid = tl.program_id(0) + in_idx0 = pid * XS + tl.arange(0, XS) + tmp0 = tl.load(in_ptr0 + in_idx0) + tmp0min = tl.min(tmp0, axis=0) + tmp0max = tl.max(tmp0, axis=0) + tmp0 = tmp0 * 2.0 + tmp0 = tl.clamp(tmp0, tmp0min, tmp0max) + tmp0 = tmp0.to(tl.int32) + tmp1 = tl.load(in_ptr1 + tmp0) + tmp2 = tl_math.exp(tmp1) + out0_idx = pid * XS + tl.arange(0, XS) + tl.store(out_ptr0 + out0_idx, tmp2) + + def triton_ldst_indirect_04_func(x0, x1, xs): + n0 = x0.numel() + assert n0 == xs, "test only single core" + y0 = torch.empty((n0,), dtype=x1.dtype, device=x1.device) + y0_txda = y0.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_ldst_indirect_04_kernel[n0 // xs, 1, 1](y0_txda, x0_txda, x1_txda, XS=xs) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_04_func(x0, x1): + x0min = torch.min(x0) + x0max = torch.max(x0) + idx = torch.clamp(x0 * 2, x0min, x0max) + return torch.exp(x1[idx.to(torch.int32)]) + + DEV = "cpu" + DTYPE = torch.float32 + offset = 8 + N0, N1 = 16, 32 + blocksize = 16 + assert N1 >= N0 + offset, "N1 must be >= N0+offset" + assert N0 == blocksize, "N0 must be == blocksize" + x0 = offset + torch.arange(0, N0, dtype=torch.float32, device=DEV) + x1 = torch.randn((N1,), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_04_func(x0, x1) + triton_cal = triton_ldst_indirect_04_func(x0, x1, blocksize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_05(): + + @triton.jit + def triton_ldst_indirect_05_kernel( + out_ptr0, in_ptr1, in_ptr2, stride_in_r, XS: tl.constexpr, RS: tl.constexpr + ): + pid = tl.program_id(0) + in_idx0 = pid * XS + tl.arange(0, XS) + in_idx1 = tl.arange(0, RS) + tmp0 = tl.arange(0, XS) + tmp1 = tl.load(in_ptr1 + in_idx1) + in_idx2 = tmp0[:, None] * stride_in_r + tmp1[None, :] + tmp2 = tl.load(in_ptr2 + in_idx2) + tmp2 = tl_math.exp(tmp2) + out0_idx = in_idx0[:, None] * RS + in_idx1[None, :] + tl.store(out_ptr0 + out0_idx, tmp2) + + def triton_ldst_indirect_05_func(xc, x2, xs, rs): + nr = x2.size()[0] + nc = xc.numel() + stride_in_r = x2.stride()[0] + assert nr == xs, "test only single core" + y0 = torch.empty((nr, nc), dtype=x2.dtype, device=x2.device) + y0_txda = y0.to("txda") + xc_txda = xc.to("txda") + x2_txda = x2.to("txda") + triton_ldst_indirect_05_kernel[nr // xs, 1, 1]( + y0_txda, xc_txda, x2_txda, stride_in_r, XS=xs, RS=rs + ) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_05_func(xr, xc, x2): + flatten_idx = (xr[:, None] * x2.stride()[0] + xc[None, :]).flatten() + extracted = x2.flatten()[flatten_idx].reshape([xr.numel(), xc.numel()]) + return torch.exp(extracted) + + DEV = "cpu" + DTYPE = torch.float32 + offset = 8 + N0, N1 = 16, 32 + blocksize = 8 + lowdimsize = N0 + assert N1 >= N0 + offset, "N1 must be >= N0+offset" + assert N0 == lowdimsize, "N0 must be == lowdimsize" + xc = offset + torch.arange(0, N0, device=DEV) + xr = torch.arange(0, blocksize, device=DEV) + x2 = torch.randn((blocksize, N1), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_05_func(xr, xc, x2) + triton_cal = triton_ldst_indirect_05_func(xc, x2, blocksize, lowdimsize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_06(): + + @triton.jit + def triton_ldst_indirect_06_kernel( + out_ptr0, + in_ptr0, + in_ptr1, + in_ptr2, + stride_in_r, + XS: tl.constexpr, + RS: tl.constexpr, + ): + pid = tl.program_id(0) + in_idx0 = pid * XS + tl.arange(0, XS) + in_idx1 = tl.arange(0, RS) + tmp0 = tl.load(in_ptr0 + in_idx0) + tmp1 = tl.load(in_ptr1 + in_idx1) + in_idx2 = tmp0[:, None] * stride_in_r + tmp1[None, :] + tmp2 = tl.load(in_ptr2 + in_idx2) + tmp2 = tl_math.exp(tmp2) + out0_idx = in_idx0[:, None] * RS + in_idx1[None, :] + tl.store(out_ptr0 + out0_idx, tmp2) + + def triton_ldst_indirect_06_func(xr, xc, x2, xs, rs): + nr = x2.size()[0] + nc = xc.numel() + stride_in_r = x2.stride()[0] + assert nr == xs, "test only single core" + y0 = torch.empty((nr, nc), dtype=x2.dtype, device=x2.device) + y0_txda = y0.to("txda") + xr_txda = xr.to("txda") + xc_txda = xc.to("txda") + x2_txda = x2.to("txda") + triton_ldst_indirect_06_kernel[nr // xs, 1, 1]( + y0_txda, xr_txda, xc_txda, x2_txda, stride_in_r, XS=xs, RS=rs + ) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_06_func(xr, xc, x2): + flatten_idx = (xr[:, None] * x2.stride()[0] + xc[None, :]).flatten() + extracted = x2.flatten()[flatten_idx].reshape([xr.numel(), xc.numel()]) + return torch.exp(extracted) + + DEV = "cpu" + DTYPE = torch.float32 + offset = 8 + N0, N1 = 16, 32 + blocksize = 4 + lowdimsize = N0 + assert N1 >= N0 + offset, "N1 must be >= N0+offset" + assert N0 == lowdimsize, "N0 must be == lowdimsize" + xc = offset + torch.arange(0, N0, device=DEV) + xr = torch.arange(0, blocksize, device=DEV) + x2 = torch.randn((blocksize, N1), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_06_func(xr, xc, x2) + triton_cal = triton_ldst_indirect_06_func(xr, xc, x2, blocksize, lowdimsize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_07(): + + @triton.jit + def triton_ldst_indirect_07_kernel( + out_ptr0, + in_ptr0, + in_ptr1, + in_ptr2, + stride_in_r, + XS: tl.constexpr, + RS: tl.constexpr, + ): + pid = tl.program_id(0) + in_idx0 = pid * XS + tl.arange(0, XS) + in_idx1 = tl.arange(0, RS) + tmp0 = tl.load(in_ptr0 + in_idx0) + tmp1 = tl.load(in_ptr1 + in_idx1) + in_idx2 = tmp0[:, None] * stride_in_r + tmp1[None, :] + tmp2 = tl.load(in_ptr2 + in_idx2) + out0_idx = in_idx0[:, None] * RS + in_idx1[None, :] + tl.store(out_ptr0 + out0_idx, tmp2) + + def triton_ldst_indirect_07_func(xr, xc, x2, xs, rs): + nr = x2.size()[0] + nc = xc.numel() + stride_in_r = x2.stride()[0] + assert nr == xs, "test only single core" + y0 = torch.empty((nr, nc), dtype=x2.dtype, device=x2.device) + y0_txda = y0.to("txda") + xr_txda = xr.to("txda") + xc_txda = xc.to("txda") + x2_txda = x2.to("txda") + triton_ldst_indirect_07_kernel[nr // xs, 1, 1]( + y0_txda, xr_txda, xc_txda, x2_txda, stride_in_r, XS=xs, RS=rs + ) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_07_func(xr, xc, x2): + flatten_idx = (xr[:, None] * x2.stride()[0] + xc[None, :]).flatten() + extracted = x2.flatten()[flatten_idx].reshape([xr.numel(), xc.numel()]) + return extracted + + DEV = "cpu" + DTYPE = torch.float32 + offset = 8 + N0, N1 = 16, 32 + blocksize = 4 + lowdimsize = N0 + assert N1 >= N0 + offset, "N1 must be >= N0+offset" + assert N0 == lowdimsize, "N0 must be == lowdimsize" + xc = offset + torch.arange(0, N0, device=DEV) + xr = torch.arange(0, blocksize, device=DEV) + x2 = torch.randn((blocksize, N1), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_07_func(xr, xc, x2) + triton_cal = triton_ldst_indirect_07_func(xr, xc, x2, blocksize, lowdimsize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_08(): + + @triton.jit + def triton_ldst_indirect_08_kernel( + out_ptr0, + in_ptr_xc, + in_ptr_x2, + stride_in_r, + OUT_COLS: tl.constexpr, + XS: tl.constexpr, + RS: tl.constexpr, + ): + pid = tl.program_id(0) + row_idx_full = pid * XS + tl.arange(0, XS) + col_pos = tl.arange(0, RS) + xc_vals = tl.load(in_ptr_xc + col_pos) + row_arange = tl.arange(0, XS) + gather_flat = row_arange[:, None] * stride_in_r + xc_vals[None, :] + vals = tl.load(in_ptr_x2 + gather_flat) + vals = tl_math.exp(vals) + out_flat = row_idx_full[:, None] * OUT_COLS + xc_vals[None, :] + tl.store(out_ptr0 + out_flat, vals) + + def triton_ldst_indirect_08_func(xc, x2, xs, rs): + nr = x2.size(0) + out_cols = x2.size(1) + stride_in_r = x2.stride(0) + assert nr == xs, "test only single core" + y0 = torch.zeros((nr, out_cols), dtype=x2.dtype, device=x2.device) + y0_txda = y0.to("txda") + xc_txda = xc.to("txda") + x2_txda = x2.to("txda") + triton_ldst_indirect_08_kernel[nr // xs, 1, 1]( + y0_txda, xc_txda, x2_txda, stride_in_r, OUT_COLS=out_cols, XS=xs, RS=rs + ) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + xc.copy_(xc_txda.cpu()) + return y0 + + def torch_ldst_indirect_08_func(xr, xc, x2): + out = torch.zeros((xr.numel(), x2.size(1)), dtype=x2.dtype, device=x2.device) + gathered = torch.exp(x2[xr[:, None], xc[None, :]]) + out.scatter_(1, xc.expand(xr.numel(), -1), gathered) + return out + + DEV = "cpu" + DTYPE = torch.float32 + offset = 8 + N0, N1 = 16, 32 + blocksize = 8 + lowdimsize = N0 + assert N1 >= N0 + offset, "N1 must be >= N0+offset" + assert N0 == lowdimsize, "N0 must be == lowdimsize" + xc = offset + torch.arange(0, N0, device=DEV) + xr = torch.arange(0, blocksize, device=DEV) + x2 = torch.randn((blocksize, N1), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_08_func(xr, xc, x2) + triton_cal = triton_ldst_indirect_08_func(xc, x2, blocksize, lowdimsize) + torch.testing.assert_close(triton_cal, torch_ref) + + +def test_ldst_indirect_09(): + + @triton.jit + def triton_ldst_indirect_09_kernel( + out_ptr0, + in_ptr1, + in_ptr2, + stride_in_r, + offset: tl.constexpr, + XS: tl.constexpr, + RS: tl.constexpr, + ): + pid = tl.program_id(0) + in_idx0 = tl.arange(0, XS) + in_idx1 = tl.arange(0, RS) + tmp0 = pid * XS + tl.load(in_ptr1 + in_idx0) + tmp1 = tl.arange(0, RS) + offset + in_idx2 = tmp0[:, None] * stride_in_r + tmp1[None, :] + tmp2 = tl.load(in_ptr2 + in_idx2) + tmp2 = tl_math.exp(tmp2) + out0_idx = pid * XS * RS + in_idx0[:, None] * RS + in_idx1[None, :] + tl.store(out_ptr0 + out0_idx, tmp2) + + def triton_ldst_indirect_09_func(xr, x2, offset, xs, rs): + nr = xr.numel() + nc = rs + stride_in_r = x2.stride()[0] + y0 = torch.empty((nr, nc), dtype=x2.dtype, device=x2.device) + y0_txda = y0.to("txda") + xr_txda = xr.to("txda") + x2_txda = x2.to("txda") + triton_ldst_indirect_09_kernel[nr // xs, 1, 1]( + y0_txda, xr_txda, x2_txda, stride_in_r, offset=offset, XS=xs, RS=rs + ) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + def torch_ldst_indirect_09_func(xr, xc, x2): + flatten_idx = (xr[:, None] * x2.stride()[0] + xc[None, :]).flatten() + extracted = x2.flatten()[flatten_idx].reshape([xr.numel(), xc.numel()]) + return torch.exp(extracted) + + DEV = "cpu" + DTYPE = torch.float32 + offset = 8 + N0, N1 = 16, 32 + blocksize = 8 + lowdimsize = N0 + assert N1 >= N0 + offset, "N1 must be >= N0+offset" + assert N0 == lowdimsize, "N0 must be == lowdimsize" + xc = offset + torch.arange(0, N0, device=DEV) + xr = torch.arange(0, blocksize, device=DEV) + x2 = torch.randn((blocksize, N1), dtype=DTYPE, device=DEV) + torch_ref = torch_ldst_indirect_09_func(xr, xc, x2) + triton_cal = triton_ldst_indirect_09_func(xr, x2, offset, blocksize, lowdimsize) + torch.testing.assert_close(triton_cal, torch_ref) + + +@triton.jit +def unstructured_mask_2d_kernel( + in_ptr, out_ptr, mask_m_ptr, mask_n_ptr, m, n, M: tl.constexpr, N: tl.constexpr +): + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + mask_m = tl.load(mask_m_ptr + offs_m, mask=offs_m < m, other=0) != 0 + mask_n = tl.load(mask_n_ptr + offs_n, mask=offs_n < n, other=0) != 0 + + in_ptrs = in_ptr + offs_m[:, None] * N + offs_n[None, :] + # dim 0 with unstructured mask. + v = tl.load(in_ptrs, mask=mask_m[:, None] and offs_n[None, :] < n, other=-2) + out_ptrs = out_ptr + offs_m[:, None] * N + offs_n[None, :] + # dim 1 with unstructured mask. + tl.store(out_ptrs, v, mask=offs_m[:, None] < m and mask_n[None, :]) + + +# helper to get torch dtype from string +def torch_dtype(dtype_str): + return eval(f"torch.{dtype_str}") + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (8, 16)], + ], +) +def test_unstructured_mask_2d(param_list): + dtype_str, shape = param_list + dtype = torch_dtype(dtype_str) + M, N = shape + + # make deterministic + random.seed(0) + torch.manual_seed(0) + + # input: use distinct values per element for easy checking + # use arange and cast to dtype + total = M * N + if dtype.is_floating_point: + in_tensor = ( + torch.arange(total, dtype=torch.float32).reshape(M, N).to(dtype).cpu() + ) + else: + in_tensor = torch.arange(total, dtype=torch.int64).reshape(M, N).to(dtype).cpu() + + # masks: random 0/1 tensors (1D) + mask_m = torch.randint(0, 2, (M,), dtype=torch.int32).cpu() # rows + mask_n = torch.randint(0, 2, (N,), dtype=torch.int32).cpu() # cols + + # out: initialize with a sentinel so we can tell which positions are untouched + if dtype.is_floating_point: + sentinel = torch.tensor(-999.0, dtype=torch.float32).to(dtype) + else: + sentinel = torch.tensor(-999, dtype=torch.int64).to(dtype) + + out_init = torch.full((M, N), sentinel.item(), dtype=dtype).cpu() + out = out_init.clone() + + # call kernel: single program covering full matrix; M,N passed as constexpr + # signature: (in_ptr, out_ptr, mask_m_ptr, mask_n_ptr, m, n, M:tl.constexpr, N:tl.constexpr) + # set m=M, n=N to simplify masks (see analysis) + in_tensor_txda = in_tensor.to("txda") + out_txda = out.to("txda") + mask_m_txda = mask_m.to("txda") + mask_n_txda = mask_n.to("txda") + unstructured_mask_2d_kernel[1, 1](in_tensor_txda, out_txda, mask_m_txda, mask_n_txda, M, N, M=M, N=N) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + # construct reference output according to kernel logic described in analysis: + # when mask_n[j] == 0 -> out should remain the initial sentinel (kernel does not store) + # when mask_n[j] == 1: + # if mask_m[i] == 1 -> out[i,j] == in[i,j] + # else -> out[i,j] == -2 + expected = out_init.clone() + for i in range(M): + row_mask = bool(mask_m[i].item()) + for j in range(N): + col_mask = bool(mask_n[j].item()) + if not col_mask: + # kernel does not store here; keep initial sentinel + expected[i, j] = out_init[i, j] + else: + if row_mask: + expected[i, j] = in_tensor[i, j] + else: + # -2 with the same dtype + if dtype.is_floating_point: + expected[i, j] = torch.tensor(-2.0, dtype=torch.float32).to( + dtype + ) + else: + expected[i, j] = torch.tensor(-2, dtype=expected.dtype) + + # validate using project's common validator + test_common.validate_cmp(dtype_str, out, expected) + + +if __name__ == "__main__": + test_ldst_indirect_08() + print("success: test_ldst_indirect_05") diff --git a/test/wafer/ops/test_le.py b/test/wafer/ops/test_le.py new file mode 100644 index 00000000..a5cc2c91 --- /dev/null +++ b/test/wafer/ops/test_le.py @@ -0,0 +1,86 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_binary(x0, y0): + res = x0 <= y0 + return res + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x <= y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16, 'bfloat16'), + (torch.int8, "int8"), + (torch.int16, "int16"), + (torch.int32, "int32"), + (torch.int64, "int64"), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + y0 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + ans = standard_binary(x0, y0) + out = torch.zeros((N,), dtype=torch.bool).cpu() + x0_txda = x0.to("txda") + y0_txda = y0.to("txda") + out_txda = out.to("txda") + triton_elementwise_binary[1, 1, 1](x0_txda, y0_txda, out_txda, N, NUMEL) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_load.py b/test/wafer/ops/test_load.py new file mode 100644 index 00000000..374b18dd --- /dev/null +++ b/test/wafer/ops/test_load.py @@ -0,0 +1,158 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + +# eg: pytest -v test.py::test_add +############################# + + +@triton.jit +def triton_load_store( + in_ptr0, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = xindex < xnumel + x0 = xindex + tmp0 = tl.load(in_ptr0 + (x0), xmask) + tmp2 = tmp0 + tl.store(out_ptr0 + (xindex), tmp2, xmask) + + +# require: all data (4d and 5d) can be placed into but without ub overflow +@triton.jit +def triton_load_store_multi_d( + in_ptr0, + out_ptr0, + BLOCK_0: tl.constexpr, + BLOCK_1: tl.constexpr, + BLOCK_2: tl.constexpr, + BLOCK_3: tl.constexpr, + BLOCK_4: tl.constexpr, + SHAPE_0: tl.constexpr, + SHAPE_1: tl.constexpr, + SHAPE_2: tl.constexpr, + SHAPE_3: tl.constexpr, + SHAPE_4: tl.constexpr, + STRIDE_0: tl.constexpr, + STRIDE_1: tl.constexpr, + STRIDE_2: tl.constexpr, + STRIDE_3: tl.constexpr, + STRIDE_4: tl.constexpr, +): + offsets = tl.program_id(0) + + offsets = offsets + tl.arange(0, BLOCK_0) * STRIDE_0 + masks = tl.arange(0, BLOCK_0) < SHAPE_0 + if (BLOCK_1 * BLOCK_2 * BLOCK_3 * BLOCK_4) > 1: + offsets = offsets[:, None] + tl.arange(0, BLOCK_1)[None, :] * STRIDE_1 + masks = masks[:, None] & (tl.arange(0, BLOCK_1)[None, :] < SHAPE_1) + if (BLOCK_2 * BLOCK_3 * BLOCK_4) > 1: + offsets = offsets[:, :, None] + tl.arange(0, BLOCK_2)[None, None, :] * STRIDE_2 + masks = masks[:, :, None] & (tl.arange(0, BLOCK_2)[None, None, :] < SHAPE_2) + if (BLOCK_3 * BLOCK_4) > 1: + offsets = ( + offsets[:, :, :, None] + + tl.arange(0, BLOCK_3)[None, None, None, :] * STRIDE_3 + ) + masks = masks[:, :, :, None] & ( + tl.arange(0, BLOCK_3)[None, None, None, :] < SHAPE_3 + ) + if BLOCK_4 > 1: + offsets = ( + offsets[:, :, :, :, None] + + tl.arange(0, BLOCK_4)[None, None, None, None, :] * STRIDE_4 + ) + masks = masks[:, :, :, :, None] & ( + tl.arange(0, BLOCK_4)[None, None, None, None, :] < SHAPE_4 + ) + + tmp_in = tl.load(in_ptr0 + offsets, masks) + tmp_out = tmp_in + tl.store(out_ptr0 + offsets, tmp_out, masks) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ["float32", (8, 8, 4), 2, 128, 64], + ["float16", (8, 8, 4), 2, 128, 64], + ["int8", (8, 8, 4), 2, 128, 64], + ["int8", (8, 7, 4), 2, 128, 64], + ], +) +def test_load_store(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = x0 + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_load_store[(ncore,)](x0_txda, y_cal_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (8, 4, 16, 16)], + ["float16", (8, 4, 16, 16)], + ["int8", (8, 4, 16, 16)], + ["float32", (8, 8, 4, 4)], + ["float16", (8, 8, 4, 4)], + ["int8", (8, 8, 4, 4)], + ["float32", (4, 8, 2, 16, 16)], + ["float16", (4, 8, 2, 16, 16)], + ["int8", (8, 8, 8, 16, 16)], + ["float32", (16, 8, 8, 4, 4)], + ["float16", (16, 8, 8, 4, 4)], + ["int8", (16, 8, 8, 4, 4)], + ], +) +def test_load_store_multi_d(param_list): + dtype, shape = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + y_expect = x0 + y_actual = test_common.generate_tensor(shape, dtype).cpu() + + blocks = list(x0.size()) + shapes = list(x0.stride()) + while len(blocks) < 5: + blocks.append(1) + shapes.append(1) + x0_txda = x0.to("txda") + y_actual_txda = y_actual.to("txda") + triton_load_store_multi_d[(1,)](x0_txda, y_actual_txda, *blocks, *blocks, *shapes) + with torch.no_grad(): + y_actual.copy_(y_actual_txda.cpu()) + test_common.validate_cmp(dtype, y_actual, y_expect) diff --git a/test/wafer/ops/test_load_store.py b/test/wafer/ops/test_load_store.py new file mode 100644 index 00000000..46bd90f8 --- /dev/null +++ b/test/wafer/ops/test_load_store.py @@ -0,0 +1,242 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +@triton.jit +def triton_load_store( + in_ptr0, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = x_index < xnumel + tmp0 = tl.load(in_ptr0 + x_index, xmask) + tmp2 = tmp0 + tl.store(out_ptr0 + x_index, tmp2, xmask) + + +# require: all data (4d and 5d) can be placed into but without ub overflow +@triton.jit +def triton_load_store_multi_d( + in_ptr0, + out_ptr0, + BLOCK_0: tl.constexpr, + BLOCK_1: tl.constexpr, + BLOCK_2: tl.constexpr, + BLOCK_3: tl.constexpr, + BLOCK_4: tl.constexpr, + SHAPE_0: tl.constexpr, + SHAPE_1: tl.constexpr, + SHAPE_2: tl.constexpr, + SHAPE_3: tl.constexpr, + SHAPE_4: tl.constexpr, + STRIDE_0: tl.constexpr, + STRIDE_1: tl.constexpr, + STRIDE_2: tl.constexpr, + STRIDE_3: tl.constexpr, + STRIDE_4: tl.constexpr, +): + offsets = tl.program_id(0) + + offsets = offsets + tl.arange(0, BLOCK_0) * STRIDE_0 + masks = tl.arange(0, BLOCK_0) < SHAPE_0 + if (BLOCK_1 * BLOCK_2 * BLOCK_3 * BLOCK_4) > 1: + offsets = offsets[:, None] + tl.arange(0, BLOCK_1)[None, :] * STRIDE_1 + masks = masks[:, None] & (tl.arange(0, BLOCK_1)[None, :] < SHAPE_1) + if (BLOCK_2 * BLOCK_3 * BLOCK_4) > 1: + offsets = offsets[:, :, None] + tl.arange(0, BLOCK_2)[None, None, :] * STRIDE_2 + masks = masks[:, :, None] & (tl.arange(0, BLOCK_2)[None, None, :] < SHAPE_2) + if (BLOCK_3 * BLOCK_4) > 1: + offsets = ( + offsets[:, :, :, None] + + tl.arange(0, BLOCK_3)[None, None, None, :] * STRIDE_3 + ) + masks = masks[:, :, :, None] & ( + tl.arange(0, BLOCK_3)[None, None, None, :] < SHAPE_3 + ) + if BLOCK_4 > 1: + offsets = ( + offsets[:, :, :, :, None] + + tl.arange(0, BLOCK_4)[None, None, None, None, :] * STRIDE_4 + ) + masks = masks[:, :, :, :, None] & ( + tl.arange(0, BLOCK_4)[None, None, None, None, :] < SHAPE_4 + ) + + tmp_in = tl.load(in_ptr0 + offsets, masks) + tmp_out = tmp_in + tl.store(out_ptr0 + offsets, tmp_out, masks) + + +@triton.jit +def triton_load_store_sle_mask( + in_ptr0, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = x_index <= xnumel + tmp0 = tl.load(in_ptr0 + x_index, xmask) + tmp2 = tmp0 + tl.store(out_ptr0 + x_index, tmp2, xmask) + + +@triton.jit +def triton_load_store_sge_mask( + in_ptr0, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = xnumel + 2 >= x_index + 1 + tmp0 = tl.load(in_ptr0 + x_index, xmask) + tmp2 = tmp0 + tl.store(out_ptr0 + x_index, tmp2, xmask) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ["float32", (8, 8, 4), 2, 128, 64], + ["float16", (8, 8, 4), 2, 128, 64], + ["int8", (8, 8, 4), 2, 128, 64], + ], +) +def test_load_store(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + y_ref = x0 + # triton结果 + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_load_store[(ncore,)](x0_txda, y_cal_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, y_cal, y_ref) + + +# 因为ascend的arrange op 支持任意的正整数,但是triton只能是2的指数倍,不然会报 ValueError: arange's range must be a power of 2 ,因此调整数字 +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (8, 4, 16, 16)], + ["float16", (8, 4, 16, 16)], + ["int8", (8, 4, 16, 16)], + ["float32", (8, 8, 4, 4)], + ["float16", (8, 8, 4, 4)], + ["int8", (8, 8, 4, 4)], + ["float32", (4, 8, 2, 16, 16)], + ["float16", (4, 8, 2, 16, 16)], + ["int8", (8, 8, 8, 16, 16)], + ["float32", (16, 8, 8, 4, 4)], + ["float16", (16, 8, 8, 4, 4)], + ["int8", (16, 8, 8, 4, 4)], + ], +) +def test_load_store_multi_d(param_list): + # 生成数据 + dtype, shape = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + y_expect = x0 + y_actual = test_common.generate_tensor(shape, dtype).cpu() + # triton结果 + blocks = list(x0.size()) + shapes = list(x0.stride()) + while len(blocks) < 5: + blocks.append(1) + shapes.append(1) + x0_txda = x0.to("txda") + y_actual_txda = y_actual.to("txda") + triton_load_store_multi_d[(1,)](x0_txda, y_actual_txda, *blocks, *blocks, *shapes) + with torch.no_grad(): + y_actual.copy_(y_actual_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, y_actual, y_expect) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ["float32", (8, 8, 4), 2, 128, 64], + ["float16", (8, 8, 4), 2, 128, 64], + ["int8", (8, 8, 4), 2, 128, 64], + ], +) +def test_load_store_sle_mask(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + y_ref = x0 + # triton结果 + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_load_store_sle_mask[(ncore,)](x0_txda, y_cal_txda, x0_txda.numel() - 1, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, y_cal, y_ref) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ["float32", (8, 8, 4), 2, 128, 64], + ["float16", (8, 8, 4), 2, 128, 64], + ["int8", (8, 8, 4), 2, 128, 64], + ], +) +def test_load_store_sge_mask(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + y_ref = x0 + # triton结果 + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_load_store_sge_mask[(ncore,)](x0_txda, y_cal_txda, x0_txda.numel() - 1, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_log.py b/test/wafer/ops/test_log.py new file mode 100644 index 00000000..276c470f --- /dev/null +++ b/test/wafer/ops/test_log.py @@ -0,0 +1,106 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_unary(x0, dtype): + res = torch.log(x0) + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + ret = tl.math.log(x) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, "float32"), + # Expected dtype ['fp32', 'fp64'] + # (torch.float16, "float16"), + # (torch.bfloat16, 'bfloat16'), + # (torch.int8, 'int8'), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_log2.py b/test/wafer/ops/test_log2.py new file mode 100644 index 00000000..095c3b71 --- /dev/null +++ b/test/wafer/ops/test_log2.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_log2(x0): + res = torch.log2(x0) + return res + + +@triton.jit +def triton_log2(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_inedx = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_inedx, None) + tmp2 = tl.log2(tmp0) + tl.store(out_ptr0 + x_inedx, tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_log2(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_log2(x0) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + triton_res_txda = triton_res.to("txda") + triton_log2[ncore, 1, 1](x0_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_log_2.py b/test/wafer/ops/test_log_2.py new file mode 100644 index 00000000..c8929469 --- /dev/null +++ b/test/wafer/ops/test_log_2.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_exp2(x0): + res = torch.pow(2, x0, out=None) + return res + + +@triton.jit +def triton_exp2(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index, None) + tmp1 = tl.exp2(tmp0) + tl.store(out_ptr0 + x_index, tmp1, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_exp2(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_exp2(x0) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + triton_res_txda = triton_res.to("txda") + triton_exp2[ncore, 1, 1](x0_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_logical_and.py b/test/wafer/ops/test_logical_and.py new file mode 100644 index 00000000..47c1b6a4 --- /dev/null +++ b/test/wafer/ops/test_logical_and.py @@ -0,0 +1,71 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_logical_and(x0, x1): + res = torch.logical_and(x0, x1) + return res + + +@triton.jit +def triton_logical_and( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index) + tmp1 = tl.load(in_ptr1 + x_index) + tmp2 = tmp0.logical_and(tmp1) + tl.store(out_ptr0 + x_index, tmp2) + + +@pytest.mark.parametrize( + "param_list", + [ + ["bool", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_and(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch.logical_and(x0, x1) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_logical_and[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_logical_or.py b/test/wafer/ops/test_logical_or.py new file mode 100644 index 00000000..604b0f34 --- /dev/null +++ b/test/wafer/ops/test_logical_or.py @@ -0,0 +1,71 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_logical_or(x0, x1): + res = torch.logical_or(x0, x1) + return res + + +@triton.jit +def triton_logical_or( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index) + tmp1 = tl.load(in_ptr1 + x_index) + tmp2 = tmp0.logical_or(tmp1) + tl.store(out_ptr0 + x_index, tmp2) + + +@pytest.mark.parametrize( + "param_list", + [ + ["bool", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_logical_or(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_logical_or(x0, x1) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_logical_or[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_lshift.py b/test/wafer/ops/test_lshift.py new file mode 100644 index 00000000..ba0b2877 --- /dev/null +++ b/test/wafer/ops/test_lshift.py @@ -0,0 +1,105 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import time +import test_common +import torch +import torch_txda # noqa: F401 + + +def standard_unary(x0, dtype): + res = x0 << 2 + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + tmp = tl.cast(2, tl.int8) + ret = x << tmp + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + # (torch.float32, 'float32'), + # (torch.float16, 'float16'), + # (torch.bfloat16, 'bfloat16'), + (torch.int8, "int8"), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + # print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + # print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_lt.py b/test/wafer/ops/test_lt.py new file mode 100644 index 00000000..7a9008a4 --- /dev/null +++ b/test/wafer/ops/test_lt.py @@ -0,0 +1,70 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_lt(x0, x1): + return x0 < x1 + + +@triton.jit +def triton_lt( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index, None) + tmp1 = tl.load(in_ptr1 + x_index, None) + tmp2 = tmp0 < tmp1 + tl.store(out_ptr0 + x_index, tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (32,), 1, 32, 32], + ], +) +def test_lt(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_lt(x0, x1).to(eval("torch." + dtype)) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_lt[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_max_dim0.py b/test/wafer/ops/test_max_dim0.py new file mode 100644 index 00000000..c0a951b9 --- /dev/null +++ b/test/wafer/ops/test_max_dim0.py @@ -0,0 +1,139 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def standard_max(x0, dim, dtype): + (res, maxindex) = torch.max(x0, dim) + return res + + +@triton.jit +def triton_max_dim0( + in_ptr0, + out_ptr0, + M: tl.constexpr, + N: tl.constexpr, + MNUMEL: tl.constexpr, + NNUMEL: tl.constexpr, +): + mblk_idx = tl.arange(0, MNUMEL) + nblk_idx = tl.arange(0, NNUMEL) + + mmask = mblk_idx < M + nmask = nblk_idx < N + + mask = (mmask[:, None]) & (nmask[None, :]) + + idx = mblk_idx[:, None] * N + nblk_idx[None, :] + + x = tl.load(in_ptr0 + idx, mask=mask, other=-float("inf")) + + ret = tl.max(x, 0) + + tl.store(out_ptr0 + nblk_idx, ret, mask=nmask) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16,'bfloat16'), TODO: waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), TODO: waiting for supporting or testing + # (torch.int32,'int32'), TODO: waiting for supporting or testing + # (torch.int64,'int64'), TODO: waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (57, 3, 64, 16), + (57, -32, 64, 32), + (57, 37, 64, 64), + (57, -256, 64, 256), + (57, 263, 64, 512), + (64, 3, 64, 16), + (64, -32, 64, 32), + (64, 37, 64, 64), + (64, -256, 64, 256), + (64, 263, 64, 512), + (3, 3, 8, 8), + (-32, 3, 32, 8), + (37, 3, 64, 8), + (-256, 3, 256, 8), + (263, 3, 512, 8), + (3, 1, 8, 8), + (-32, 1, 32, 8), + (37, 1, 64, 8), + (-256, 1, 256, 8), + (263, 1, 512, 8), +] + +map_for_64_t = {37: (31, 32), 263: (107, 128)} +map_for_32_t = {263: (137, 256)} + + +# @pytest.mark.parametrize('dtype, sigtype',[(torch.float32,'float32'),]) +@pytest.mark.parametrize("M, N, MNUMEL, NNUMEL", [(64, -32, 64, 32)]) +@pytest.mark.parametrize("dtype, sigtype", types) +# @pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_max_dim0(dtype, sigtype, M, N, MNUMEL, NNUMEL): + + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == "float32" or sigtype == "bfloat16" or sigtype == "int32": + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"max : ({M}, {N}) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(M, N), dtype=sigtype) + + ans = standard_max(x0, 0, dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_max_dim0[1, 1, 1](x0_txda, output_txda, M=M, N=N, MNUMEL=MNUMEL, NNUMEL=NNUMEL) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_max_dim1.py b/test/wafer/ops/test_max_dim1.py new file mode 100644 index 00000000..11bdfc98 --- /dev/null +++ b/test/wafer/ops/test_max_dim1.py @@ -0,0 +1,149 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def standard_max(x0, dim, dtype): + (res, maxindex) = torch.max(x0, dim) + return res + + +@triton.jit +def triton_max_dim1( + in_ptr0, + out_ptr0, + M: tl.constexpr, + N: tl.constexpr, + MNUMEL: tl.constexpr, + NNUMEL: tl.constexpr, +): + mblk_idx = tl.arange(0, MNUMEL) + nblk_idx = tl.arange(0, NNUMEL) + + mmask = mblk_idx < M + nmask = nblk_idx < N + + mask = (mmask[:, None]) & (nmask[None, :]) + + idx = mblk_idx[:, None] * N + nblk_idx[None, :] + + if in_ptr0.dtype == tl.int8: + padding = -128 + else: + padding = -float("inf") + + x = tl.load(in_ptr0 + idx, mask=mask, other=padding) + + ret = tl.max(x, 1) + + tl.store(out_ptr0 + mblk_idx, ret, mask=mmask) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16,'bfloat16'), waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), waiting for supporting or testing + # (torch.int32,'int32'), waiting for supporting or testing + # (torch.int64,'int64'), waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (57, 3, 64, 16), + (57, -32, 64, 32), + (57, 37, 64, 64), + (57, -256, 64, 256), + (57, 263, 64, 512), + (64, 3, 64, 16), + (64, -32, 64, 32), + (64, 37, 64, 64), + (64, -256, 64, 256), + (64, 263, 64, 512), + (3, 3, 8, 8), + (-32, 3, 32, 8), + (37, 3, 64, 8), + (-256, 3, 256, 8), + (263, 3, 512, 8), + (3, 1, 8, 8), + (-32, 1, 32, 8), + (37, 1, 64, 8), + (-256, 1, 256, 8), + (263, 1, 512, 8), +] + +map_for_64_t = {37: (31, 32), 263: (107, 128)} +map_for_32_t = {263: (137, 256)} + + +@pytest.mark.parametrize( + "M, N, MNUMEL, NNUMEL", + [ + (64, -32, 64, 32), + ], +) +# @pytest.mark.parametrize('M, N',[(263,3),(-256,3)]) +@pytest.mark.parametrize("dtype, sigtype", types) +# @pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_max_dim1(dtype, sigtype, M, N, MNUMEL, NNUMEL): + + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == "float32" or sigtype == "bfloat16" or sigtype == "int32": + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"max : ({M}, {N}) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(M, N), dtype=sigtype) + + ans = standard_max(x0, 1, dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((M,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_max_dim1[1, 1, 1](x0_txda, output_txda, M=M, N=N, MNUMEL=MNUMEL, NNUMEL=NNUMEL) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_max_vector.py b/test/wafer/ops/test_max_vector.py new file mode 100644 index 00000000..1d6de543 --- /dev/null +++ b/test/wafer/ops/test_max_vector.py @@ -0,0 +1,94 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_(x0, dtype): + res, index = torch.max(x0, 0, keepdim=True) + return res + + +@triton.jit +def triton_max_vector(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + + if in_ptr0.dtype == tl.int8: + padding = -128 + else: + padding = -float("inf") + + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N, other=padding) + ret = tl.max(x, 0) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < 1) + + +types = [ + (torch.float32, "float32"), + # (torch.float16,'float16'), TODO : fix reduceConverter bug + # (torch.bfloat16,'bfloat16'), waiting for supporting or testing + # (torch.int8,'int8'), TODO : fix compiler bug + # (torch.int16,'int16'), waiting for supporting or testing + # (torch.int32,'int32'), waiting for supporting or testing + # (torch.int64,'int64'), waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +# @pytest.mark.skip(reason="randomly failed") +@pytest.mark.parametrize("dtype, sigtype", types) +@pytest.mark.parametrize("N, NUMEL", shapes) +def test_reduce_dim0_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_(x0, dtype) + x0 = x0.cpu() + + output = torch.zeros((1,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_max_vector[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_maximum.py b/test/wafer/ops/test_maximum.py new file mode 100644 index 00000000..8da1f2de --- /dev/null +++ b/test/wafer/ops/test_maximum.py @@ -0,0 +1,72 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_maximum(x0, x1): + res = torch.maximum(x0, x1) + return res + + +@triton.jit +def triton_maximum( + in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = x_index < xnumel + tmp0 = tl.load(in_ptr0 + x_index, xmask) + tmp1 = tl.load(in_ptr1 + x_index, xmask) + tmp2 = tl.maximum(tmp0, tmp1) + tl.store(out_ptr0 + x_index, tmp2, xmask) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_maximum(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_maximum(x0, x1) + # triton结果 + triton_res = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_maximum[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_mean_dim0.py b/test/wafer/ops/test_mean_dim0.py new file mode 100644 index 00000000..0ec611b8 --- /dev/null +++ b/test/wafer/ops/test_mean_dim0.py @@ -0,0 +1,162 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def standard_mean(x0, dim, dtype): + res = torch.mean(x0, dim, dtype=dtype) + return res + + +@triton.jit +def triton_mean_dim0( + in_ptr0, + out_ptr0, + M: tl.constexpr, + N: tl.constexpr, + MNUMEL: tl.constexpr, + NNUMEL: tl.constexpr, +): + mblk_idx = tl.arange(0, MNUMEL) + nblk_idx = tl.arange(0, NNUMEL) + + mmask = mblk_idx < M + nmask = nblk_idx < N + + mask = (mmask[:, None]) & (nmask[None, :]) + + idx = mblk_idx[:, None] * N + nblk_idx[None, :] + + x = tl.load(in_ptr0 + idx, mask=mask, other=0) + + if x.dtype == tl.bfloat16: + ret = (tl.sum(x.to(tl.float32), 0) / M).to(tl.bfloat16) + elif x.dtype == tl.float16: + ret = tl.sum(x, 0) / M + else: + ret = tl.sum(x.to(tl.float32), 0) / M + + tl.store(out_ptr0 + nblk_idx, ret, mask=nmask) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16,'bfloat16'), TODO: waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), TODO: waiting for supporting or testing + # (torch.int32,'int32'), TODO: waiting for supporting or testing + # (torch.int64,'int64'), TODO: waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (57, 3, 64, 16), + (57, -32, 64, 32), + (57, 37, 64, 64), + (57, -256, 64, 256), + (57, 263, 64, 512), + (64, 3, 64, 16), + (64, -32, 64, 32), + (64, 37, 64, 64), + (64, -256, 64, 256), + (64, 263, 64, 512), + (3, 3, 8, 8), + (-32, 3, 32, 8), + (37, 3, 64, 8), + (-256, 3, 256, 8), + (263, 3, 512, 8), + (3, 1, 8, 8), + (-32, 1, 32, 8), + (37, 1, 64, 8), + (-256, 1, 256, 8), + (263, 1, 512, 8), +] + +map_for_64_t = {37: (31, 32)} +map_for_32_t = {263: (137, 256)} + + +# @pytest.mark.parametrize('dtype, sigtype',[(torch.float32,'float32'),]) +@pytest.mark.parametrize( + "M, N, MNUMEL, NNUMEL", + [ + (57, 3, 64, 16), + (64, -32, 64, 32), + (37, 3, 64, 8), + (263, 1, 512, 8), + (-256, 3, 256, 8), + ], +) +@pytest.mark.parametrize("dtype, sigtype", types) +# @pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_mean_dim0(dtype, sigtype, M, N, MNUMEL, NNUMEL): + + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + res_dtype = dtype + res_sigtype = sigtype + should_cast_to_fp32 = ["int8", "int16", "int32", "int64", "float32", "bfloat16"] + + if sigtype in should_cast_to_fp32: + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + if sigtype != "bfloat16": + res_dtype = torch.float32 + res_sigtype = "float32" + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(M, N), dtype=sigtype) + + ans = standard_mean(x0, 0, res_dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((N,), dtype=res_dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_mean_dim0[1, 1, 1]( + x0_txda, output_txda, M=M, N=N, MNUMEL=MNUMEL, NNUMEL=NNUMEL, debug=True + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(res_sigtype, output, ans) diff --git a/test/wafer/ops/test_mean_dim1.py b/test/wafer/ops/test_mean_dim1.py new file mode 100644 index 00000000..5df56c0b --- /dev/null +++ b/test/wafer/ops/test_mean_dim1.py @@ -0,0 +1,162 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def standard_mean(x0, dim, dtype): + res = torch.mean(x0, dim, dtype=dtype) + return res + + +@triton.jit +def triton_mean_dim1( + in_ptr0, + out_ptr0, + M: tl.constexpr, + N: tl.constexpr, + MNUMEL: tl.constexpr, + NNUMEL: tl.constexpr, +): + mblk_idx = tl.arange(0, MNUMEL) + nblk_idx = tl.arange(0, NNUMEL) + + mmask = mblk_idx < M + nmask = nblk_idx < N + + mask = (mmask[:, None]) & (nmask[None, :]) + + idx = mblk_idx[:, None] * N + nblk_idx[None, :] + + x = tl.load(in_ptr0 + idx, mask=mask, other=0) + + if x.dtype == tl.bfloat16: + ret = (tl.sum(x.to(tl.float32), 1) / N).to(tl.bfloat16) + elif x.dtype == tl.float16: + ret = tl.sum(x, 1) / N + else: + ret = tl.sum(x.to(tl.float32), 1) / N + + tl.store(out_ptr0 + mblk_idx, ret, mask=mmask) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16,'bfloat16'), TODO: waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), TODO: waiting for supporting or testing + # (torch.int32,'int32'), TODO: waiting for supporting or testing + # (torch.int64,'int64'), TODO: waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (57, 3, 64, 16), + (57, -32, 64, 32), + (57, 37, 64, 64), + (57, -256, 64, 256), + (57, 263, 64, 512), + (64, 3, 64, 16), + (64, -32, 64, 32), + (64, 37, 64, 64), + (64, -256, 64, 256), + (64, 263, 64, 512), + (3, 3, 8, 8), + (-32, 3, 32, 8), + (37, 3, 64, 8), + (-256, 3, 256, 8), + (263, 3, 512, 8), + (3, 1, 8, 8), + (-32, 1, 32, 8), + (37, 1, 64, 8), + (-256, 1, 256, 8), + (263, 1, 512, 8), +] + +map_for_64_t = {37: (31, 32)} +map_for_32_t = {263: (137, 256)} + + +# @pytest.mark.parametrize('dtype, sigtype',[(torch.float32,'float32'),]) +@pytest.mark.parametrize( + "M, N, MNUMEL, NNUMEL", + [ + (57, 3, 64, 16), + (64, -32, 64, 32), + (37, 3, 64, 8), + (263, 1, 512, 8), + (-256, 3, 256, 8), + ], +) +@pytest.mark.parametrize("dtype, sigtype", types) +# @pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_mean_dim1(dtype, sigtype, M, N, MNUMEL, NNUMEL): + + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + res_dtype = dtype + res_sigtype = sigtype + should_cast_to_fp32 = ["int8", "int16", "int32", "int64", "float32", "bfloat16"] + + if sigtype in should_cast_to_fp32: + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + if sigtype != "bfloat16": + res_dtype = torch.float32 + res_sigtype = "float32" + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(M, N), dtype=sigtype) + + ans = standard_mean(x0, 1, res_dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((M,), dtype=res_dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_mean_dim1[1, 1, 1]( + x0_txda, output_txda, M=M, N=N, MNUMEL=MNUMEL, NNUMEL=NNUMEL, debug=True + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(res_sigtype, output, ans) diff --git a/test/wafer/ops/test_mean_vector.py b/test/wafer/ops/test_mean_vector.py new file mode 100644 index 00000000..b641f578 --- /dev/null +++ b/test/wafer/ops/test_mean_vector.py @@ -0,0 +1,110 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def standard_mean(x0, dim, dtype): + res = torch.mean(x0, dim, keepdim=True, dtype=dtype) + return res + + +@triton.jit +def triton_mean_dim0(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N, other=0) + + if x.dtype == tl.bfloat16: + ret = (tl.sum(x.to(tl.float32), 0) / N).to(tl.bfloat16) + elif x.dtype == tl.float16: + ret = tl.sum(x, 0) / N + else: + ret = tl.sum(x.to(tl.float32), 0) / N + + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < 1) + + +types = [ + (torch.float32, "float32"), + # (torch.float16,'float16'), TODO: should fix reduceConverter's bug + # (torch.bfloat16,'bfloat16'), TODO: waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), TODO: waiting for supporting or testing + # (torch.int32,'int32'), TODO: waiting for supporting or testing + # (torch.int64,'int64'), TODO: waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: (31, 32)} + + +@pytest.mark.parametrize("dtype, sigtype", types) +@pytest.mark.parametrize("N, NUMEL", shapes) +def test_mean_dim0(dtype, sigtype, N, NUMEL): + + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N][0] if N in map_for_64_t else N + NUMEL = map_for_64_t[N][1] if N in map_for_64_t else NUMEL + + res_dtype = dtype + res_sigtype = sigtype + should_cast_to_fp32 = ["int8", "int16", "int32", "int64", "float32"] + + if sigtype in should_cast_to_fp32: + res_dtype = torch.float32 + res_sigtype = "float32" + + print(f"sum : ({N},) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_mean(x0, 0, res_dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((1,), dtype=res_dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_mean_dim0[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(res_sigtype, output, ans) diff --git a/test/wafer/ops/test_min_dim0.py b/test/wafer/ops/test_min_dim0.py new file mode 100644 index 00000000..b8808ed0 --- /dev/null +++ b/test/wafer/ops/test_min_dim0.py @@ -0,0 +1,171 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +@pytest.mark.skip(reason="to be supported by bishengir-compile") +def test_min_dim0_3d(): + + def torch_func(x, dim): + res = torch.min(x, dim) + return res + + @triton.jit + def triton_kernel( + out_ptr0, in_ptr0, N0: tl.constexpr, N1: tl.constexpr, N2: tl.constexpr + ): + idx0 = tl.arange(0, N0) + idx1 = tl.arange(0, N1) + idx2 = tl.arange(0, N2) + in_idx = ( + idx2[None, None, :] + + idx1[None, :, None] * N2 + + idx0[:, None, None] * N2 * N1 + ) + tmp0 = tl.load(in_ptr0 + in_idx) + tmp1 = tl.min(tmp0, 0) + out_idx = idx2[None, :] + idx1[:, None] * N2 + tl.store(out_ptr0 + out_idx, tmp1) + + def triton_func(x0, dim): + N0, N1, N2 = x0.size() + y0 = test_common.generate_tensor(shape=(N1, N2), dtype="float32").cpu() + y0_txda = y0.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y0_txda, x0_txda, N0, N1, N2) + with torch.no_grad(): + y0.copy_(y0_txda.cpu()) + return y0 + + dim = 0 + N0, N1, N2 = 1, 22, 13 + x0 = test_common.generate_tensor(shape=(N0, N1, N2), dtype="float32").cpu() + torch_ref = torch_func(x0, dim) + triton_cal = triton_func(x0, dim) + test_common.validate_cmp("float32", triton_cal, torch_ref) + + +def standard_min(x0, dim, dtype): + res, index = torch.min(x0, dim) + return res + + +@triton.jit +def triton_min_dim0( + in_ptr0, + out_ptr0, + M: tl.constexpr, + N: tl.constexpr, + MNUMEL: tl.constexpr, + NNUMEL: tl.constexpr, +): + mblk_idx = tl.arange(0, MNUMEL) + nblk_idx = tl.arange(0, NNUMEL) + + mmask = mblk_idx < M + nmask = nblk_idx < N + + mask = (mmask[:, None]) & (nmask[None, :]) + + idx = mblk_idx[:, None] * N + nblk_idx[None, :] + if in_ptr0.dtype == tl.int8: + padding = 127 + else: + padding = float("inf") + x = tl.load(in_ptr0 + idx, mask=mask, other=padding) + + ret = tl.min(x, 0) + + tl.store(out_ptr0 + nblk_idx, ret, mask=nmask) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16,'bfloat16'), TODO: waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), TODO: waiting for supporting or testing + # (torch.int32,'int32'), TODO: waiting for supporting or testing + # (torch.int64,'int64'), TODO: waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +# shapes=[ +# (57,3,64,16), (57,-32,64,32), (57,37,64,64), (57,-256,64,256), (57,263,64,512), +# (64,3,64,16), (64,-32,64,32), (64,37,64,64), (64,-256,64,256), (64,263,64,512), +# (3,3,8,8), (-32,3,32,8), (37,3,64,8), (-256,3,256,8), (263,3,512,8), +# (3,1,8,8), (-32,1,32,8), (37,1,64,8), (-256,1,256,8), (263,1,512,8), +# ] +shapes = [ + (64, -32, 64, 32), +] + +map_for_64_t = {37: (31, 32), 263: (107, 128)} +map_for_32_t = {263: (137, 256)} + + +# @pytest.mark.parametrize('dtype, sigtype',[(torch.float32,'float32'),]) +@pytest.mark.parametrize("M, N, MNUMEL, NNUMEL", [(64, -32, 64, 32)]) +@pytest.mark.parametrize("dtype, sigtype", types) +# @pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_min_dim0(dtype, sigtype, M, N, MNUMEL, NNUMEL): + + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == "float32" or sigtype == "bfloat16" or sigtype == "int32": + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"min : ({M}, {N}) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(M, N), dtype=sigtype) + + ans = standard_min(x0, 0, dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_min_dim0[1, 1, 1](x0_txda, output_txda, M=M, N=N, MNUMEL=MNUMEL, NNUMEL=NNUMEL) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_min_dim1.py b/test/wafer/ops/test_min_dim1.py new file mode 100644 index 00000000..455dd1f5 --- /dev/null +++ b/test/wafer/ops/test_min_dim1.py @@ -0,0 +1,147 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def standard_min(x0, dim, dtype): + (res, minindex) = torch.min(x0, dim) + return res + + +@triton.jit +def triton_min_dim1( + in_ptr0, + out_ptr0, + M: tl.constexpr, + N: tl.constexpr, + MNUMEL: tl.constexpr, + NNUMEL: tl.constexpr, +): + mblk_idx = tl.arange(0, MNUMEL) + nblk_idx = tl.arange(0, NNUMEL) + + mmask = mblk_idx < M + nmask = nblk_idx < N + + mask = (mmask[:, None]) & (nmask[None, :]) + + idx = mblk_idx[:, None] * N + nblk_idx[None, :] + if in_ptr0.dtype == tl.int8: + padding = 127 + else: + padding = float("inf") + x = tl.load(in_ptr0 + idx, mask=mask, other=padding) + + ret = tl.min(x, 1) + + tl.store(out_ptr0 + mblk_idx, ret, mask=mmask) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16,'bfloat16'), waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), waiting for supporting or testing + # (torch.int32,'int32'), waiting for supporting or testing + # (torch.int64,'int64'), waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (57, 3, 64, 16), + (57, -32, 64, 32), + (57, 37, 64, 64), + (57, -256, 64, 256), + (57, 263, 64, 512), + (64, 3, 64, 16), + (64, -32, 64, 32), + (64, 37, 64, 64), + (64, -256, 64, 256), + (64, 263, 64, 512), + (3, 3, 8, 8), + (-32, 3, 32, 8), + (37, 3, 64, 8), + (-256, 3, 256, 8), + (263, 3, 512, 8), + (3, 1, 8, 8), + (-32, 1, 32, 8), + (37, 1, 64, 8), + (-256, 1, 256, 8), + (263, 1, 512, 8), +] + +map_for_64_t = {37: (31, 32), 263: (107, 128)} +map_for_32_t = {263: (137, 256)} + + +@pytest.mark.parametrize( + "M, N, MNUMEL, NNUMEL", + [ + (64, -32, 64, 32), + ], +) +# @pytest.mark.parametrize('M, N',[(263,3),(-256,3)]) +@pytest.mark.parametrize("dtype, sigtype", types) +# @pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_min_dim1(dtype, sigtype, M, N, MNUMEL, NNUMEL): + + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == "float32" or sigtype == "bfloat16" or sigtype == "int32": + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"min : ({M}, {N}) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(M, N), dtype=sigtype) + + ans = standard_min(x0, 1, dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((M,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_min_dim1[1, 1, 1](x0_txda, output_txda, M=M, N=N, MNUMEL=MNUMEL, NNUMEL=NNUMEL) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_min_vector.py b/test/wafer/ops/test_min_vector.py new file mode 100644 index 00000000..0e0f77c9 --- /dev/null +++ b/test/wafer/ops/test_min_vector.py @@ -0,0 +1,93 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest +import triton +import triton.language as tl +import test_common +import torch +import torch_txda # noqa: F401 + + +def standard_(x0, dtype): + res, index = torch.min(x0, 0, keepdim=True) + return res + + +@triton.jit +def triton_min_vector(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + if in_ptr0.dtype == tl.int8: + padding = 127 + else: + padding = float("inf") + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N, other=padding) + + ret = tl.min(x, 0) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < 1) + + +types = [ + (torch.float32, "float32"), + # (torch.float16,'float16'), TODO : fix reduceConverter bug + # (torch.bfloat16,'bfloat16'), waiting for supporting or testing + # (torch.int8,'int8'), TODO : fix compiler bug + # (torch.int16,'int16'), waiting for supporting or testing + # (torch.int32,'int32'), waiting for supporting or testing + # (torch.int64,'int64'), waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +# @pytest.mark.skip(reason="randomly failed") +@pytest.mark.parametrize("dtype, sigtype", types) +@pytest.mark.parametrize("N, NUMEL", shapes) +def test_reduce_dim0_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_(x0, dtype) + x0 = x0.cpu() + + output = torch.zeros((1,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_min_vector[1, 1, 1](x0_txda, output_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_minimum.py b/test/wafer/ops/test_minimum.py new file mode 100644 index 00000000..95066879 --- /dev/null +++ b/test/wafer/ops/test_minimum.py @@ -0,0 +1,72 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_minimum(x0, x1): + res = torch.minimum(x0, x1) + return res + + +@triton.jit +def triton_minimum( + in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + x_index = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = x_index < xnumel + tmp0 = tl.load(in_ptr0 + x_index, xmask) + tmp1 = tl.load(in_ptr1 + x_index, xmask) + tmp2 = tl.minimum(tmp0, tmp1) + tl.store(out_ptr0 + x_index, tmp2, xmask) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_minimum(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + y_ref = torch_minimum(x0, x1) + # triton结果 + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_minimum[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_mod.py b/test/wafer/ops/test_mod.py new file mode 100644 index 00000000..18b8549a --- /dev/null +++ b/test/wafer/ops/test_mod.py @@ -0,0 +1,94 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0, x1, dtype): + output_dtype = x0.dtype + if dtype == "float16": + x0 = x0.to(torch.float32) + x1 = x1.to(torch.float32) + elif dtype == "float32": + x0 = x0.to(torch.float64) + x1 = x1.to(torch.float64) + res = torch.div(x0, x1, rounding_mode="trunc") + res = x0 - x1 * res + # Compute the reference in higher precision, then match the kernel output dtype. + return res.to(output_dtype) + + +@triton.jit +def triton_mod( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 % tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + if dtype == "int8": + x0 = torch.randint( + low=1, high=127, size=shape, dtype=eval("torch." + dtype) + ).cpu() + x1 = torch.randint( + low=1, high=127, size=shape, dtype=eval("torch." + dtype) + ).cpu() + else: + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1, dtype) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_mod[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + # test_common.validate_cmp(dtype, y_cal, y_ref.cpu()) + if dtype == "int8": + torch.equal(y_cal, y_ref) + else: + res = torch.isclose(y_cal, y_ref, rtol=1e-3, atol=1e-3, equal_nan=True) + if not res.all(): + max_diff = torch.max((y_ref - y_cal)).item() + raise ValueError(f"Tensors are not close, diff is {max_diff}") diff --git a/test/wafer/ops/test_mul.py b/test/wafer/ops/test_mul.py new file mode 100644 index 00000000..9b32fb5d --- /dev/null +++ b/test/wafer/ops/test_mul.py @@ -0,0 +1,71 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0, x1): + res = x0 * x1 + return res + + +@triton.jit +def triton_mul( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 * tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_mul[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_nearest.py b/test/wafer/ops/test_nearest.py new file mode 100644 index 00000000..e20d0fc8 --- /dev/null +++ b/test/wafer/ops/test_nearest.py @@ -0,0 +1,171 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import math +import numpy as np +import pytest + + +@triton.jit +def nearest_resize_kernel( + img_src_ptr, + img_dst_ptr, + src_rows, + src_cols, + dst_rows, + dst_cols, + RR_H, + RR_W, + C, + stride_in_h, + stride_in_w, + stride_in_c, + stride_out_h, + stride_out_w, + stride_out_c, + BLOCK_SIZE: tl.constexpr, +): + # RR_H和RR_W分别为高和宽的缩放比例 + block_id_c = tl.program_id(0) + block_id_h = tl.program_id(1) + block_id_w = tl.program_id(2) + dest_h_offs = block_id_h * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + dest_w_offs = block_id_w * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + dest_offs = ( + block_id_c[None, None] * stride_out_c + + dest_h_offs[:, None] * stride_out_h + + dest_w_offs[None, :] * stride_out_w + ) + # 根据output image的坐标值(dest_h_offs, dest_w_offs)计算input image的坐标值(sy, sx) + fy = dest_h_offs * RR_H + sy = tl.floor(fy) + fx = dest_w_offs * RR_W + sx = tl.floor(fx) + + src_offsets = ( + block_id_c[None, None] * stride_in_c + + tl.clamp(sy, 0, src_rows - 1)[:, None].to(tl.int32) * stride_in_h + + tl.clamp(sx, 0, src_cols - 1)[None, :].to(tl.int32) * stride_in_w + ) + src_val = tl.load(img_src_ptr + src_offsets) + dst_mask = (dest_h_offs[:, None] < dst_rows) & (dest_w_offs[None, :] < dst_cols) + tl.store(img_dst_ptr + dest_offs, src_val, mask=dst_mask) + + +def triton_kernel(img_src, img_dst): + N, C, src_rows, src_cols = img_src.shape + _, _, dst_rows, dst_cols = img_dst.shape + R_H = float(dst_rows) / src_rows + R_W = float(dst_cols) / src_cols + RR_H = 1.0 / R_H + RR_W = 1.0 / R_W + stride_in_n, stride_in_c, stride_in_h, stride_in_w = img_src.stride() + stride_out_n, stride_out_c, stride_out_h, stride_out_w = img_dst.stride() + bs = 16 + grid = lambda meta: ( + C, + triton.cdiv(dst_rows, meta["BLOCK_SIZE"]), + triton.cdiv(dst_cols, meta["BLOCK_SIZE"]), + ) + img_src_txda = img_src.to("txda") + img_dst_txda = img_dst.to("txda") + nearest_resize_kernel[grid]( + img_src_txda, + img_dst_txda, + src_rows, + src_cols, + dst_rows, + dst_cols, + RR_H, + RR_W, + C, + stride_in_h, + stride_in_w, + stride_in_c, + stride_out_h, + stride_out_w, + stride_out_c, + bs, + ) + with torch.no_grad(): + img_dst.copy_(img_dst_txda.cpu()) + return img_dst + + +def nearest_resize_cpu(img_src, img_dst): + N, C, src_rows, src_cols = img_src.shape + _, _, dst_rows, dst_cols = img_dst.shape + # RR_H和RR_W分别为高和宽的缩放比例 + RR_H = src_rows / float(dst_rows) + RR_W = src_cols / float(dst_cols) + # 根据output image的坐标值(i,j)计算input image的坐标值(sy, sx) + for i in range(dst_rows): + for j in range(dst_cols): + fy = i * RR_H + sy = math.floor(fy) + fx = j * RR_W + sx = math.floor(fx) + src_val = img_src[ + 0, :, np.clip(sy, 0, src_rows - 1), np.clip(sx, 0, src_cols - 1) + ] + img_dst[0, :, i, j] = src_val + return img_dst + + +@pytest.mark.parametrize( + "shapes", + [ + [360, 640, 140, 280], + ], +) +def test_nearest(shapes): + src_rows, src_cols, dst_rows, dst_cols = shapes + img_src = torch.rand(1, 4, src_rows, src_cols, dtype=torch.float32, device="cpu") + img_dst = torch.zeros( + (1, img_src.shape[1], dst_rows, dst_cols), + dtype=img_src.dtype, + device=img_src.device, + ) + torch_ref = nearest_resize_cpu(img_src.cpu(), img_dst.cpu()) + triton_cal = triton_kernel(img_src, img_dst) + torch.testing.assert_close(torch_ref.cpu(), triton_cal) + + +if __name__ == "__main__": + src_rows, src_cols = 360, 640 + dst_rows, dst_cols = 140, 280 + img_src = torch.rand(1, 4, src_rows, src_cols, dtype=torch.float32, device="cpu") + img_dst = torch.zeros( + (1, img_src.shape[1], dst_rows, dst_cols), + dtype=img_src.dtype, + device=img_src.device, + ) + + assert ( + img_src.shape[0] == 1 + ), "currently supports only shape[0] == 1 which does not change the functionality of thie case" + torch_ref = nearest_resize_cpu(img_src.cpu(), img_dst.cpu()) + triton_cal = triton_kernel(img_src, img_dst) + torch.testing.assert_close(torch_ref.cpu(), triton_cal) + print("success") diff --git a/test/wafer/ops/test_neg.py b/test/wafer/ops/test_neg.py new file mode 100644 index 00000000..72ac0f3f --- /dev/null +++ b/test/wafer/ops/test_neg.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest +import triton +import triton.language as tl +import time +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_neg(x0): + res = -x0 + return res + + +@triton.jit +def triton_neg(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1 = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = -tmp0 + tl.store(out_ptr0 + (x0), tmp1, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float16", (8, 8), 8, 8, 8], + ["float32", (8, 8), 8, 8, 8], + ["int8", (2, 4096, 8), 32, 2048, 64], + ], +) +def test_neg(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_neg(x0) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + y_cal_txda = y_cal.to("txda") + triton_neg[ncore, 1, 1](x0_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_npu_indexing.py b/test/wafer/ops/test_npu_indexing.py new file mode 100644 index 00000000..47cc8ff2 --- /dev/null +++ b/test/wafer/ops/test_npu_indexing.py @@ -0,0 +1,197 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import time + + +def foo(a, b, c): + Z, Y, X, R = (1, 1, 64, 64) + y = a + b + y = y.sum(-1) + y = y.unsqueeze(3) + y = y.broadcast_to(Z, Y, X, R) + b + y = c + y.permute(0, 1, 3, 2) + return y + + +@triton.jit +def triton_foo( + in_ptr0, + in_ptr1, + in_ptr2, + out_ptr0, + BLOCK1: tl.constexpr, + BLOCK1_SUB: tl.constexpr, + BLOCK2: tl.constexpr, + Z: tl.constexpr, + Y: tl.constexpr, + X: tl.constexpr, + R: tl.constexpr, + Z_STRIDE: tl.constexpr, + Y_STRIDE: tl.constexpr, + X_STRIDE: tl.constexpr, + R_STRIDE: tl.constexpr, + Z_STRIDE1: tl.constexpr, + Y_STRIDE1: tl.constexpr, + X_STRIDE1: tl.constexpr, + R_STRIDE1: tl.constexpr, +): + offset: tl.constexpr = tl.program_id(0) * BLOCK1 + base1 = tl.arange(0, BLOCK1_SUB) + base2 = tl.arange(0, BLOCK2) + nsub: tl.constexpr = BLOCK1 // BLOCK1_SUB + # loops1 : tl.constexpr = nsub * Y * Z + loops1: tl.constexpr = nsub + loops2: tl.constexpr = R // BLOCK2 + + for z in range(Z): + for y in range(Y): + for loop1 in range(loops1): + # y = (loop1 // nsub) % Y + # z = loop1 // nsub // Y + # off1 = (loop1 % nsub) + off1 = loop1 + x = offset + (off1 * BLOCK1_SUB) + base1[:, None] + x1 = offset + (off1 * BLOCK1_SUB) + base1[None, :] + _tmp4 = tl.full([BLOCK1_SUB, BLOCK2], 0, tl.float32) + for loop2 in range(loops2): + r = loop2 * BLOCK2 + base2[None, :] + tmp0 = tl.load( + in_ptr0 + + ( + R_STRIDE * r + + (X_STRIDE * x) + + (Y_STRIDE * y) + + (Z_STRIDE * z) + ), + None, + ) + tmp1 = tl.load( + in_ptr1 + + ( + R_STRIDE * r + + (X_STRIDE * x) + + (Y_STRIDE * y) + + (Z_STRIDE * z) + ), + None, + ) + tmp2 = tmp0 + tmp1 + _tmp4 = _tmp4 + tmp2 + tmp4 = tl.sum(_tmp4, 1)[:, None] + tmp5 = tmp4.reshape(BLOCK1_SUB, 1).broadcast_to(BLOCK1_SUB, BLOCK2) + + for loop2 in range(loops2): + r = loop2 * BLOCK2 + base2[None, :] + tmp6 = tl.load( + in_ptr1 + + ( + R_STRIDE * r + + (X_STRIDE * x) + + (Y_STRIDE * y) + + (Z_STRIDE * z) + ), + None, + ) + tmp7 = tmp6 + tmp5 + r1 = loop2 * BLOCK2 + base2[:, None] + tmp8 = tl.load( + in_ptr2 + + ( + R_STRIDE1 * x1 + + (X_STRIDE1 * r1) + + (Y_STRIDE1 * y) + + (Z_STRIDE1 * z) + ), + None, + ) + + tmp9 = tmp8.reshape(BLOCK2, BLOCK1_SUB) + tmp7.reshape( + BLOCK1_SUB, BLOCK2 + ).permute(1, 0) + tl.store( + out_ptr0 + + ( + R_STRIDE1 * x1 + + (X_STRIDE1 * r1) + + (Y_STRIDE1 * y) + + (Z_STRIDE1 * z) + ), + tmp9, + None, + ) + + +def foo_triton_wrapper(a, b, c): + NBLOCKS = 1 + BLOCK1 = a.shape[2] // NBLOCKS + BLOCK1_SUB = 64 + BLOCK2 = 64 + + value = torch.empty_strided( + (c.shape[0], c.shape[1], c.shape[2], c.shape[3]), + (c.stride()[0], c.stride()[1], c.stride()[2], c.stride()[3]), + dtype=torch.float32, + ).cpu() + + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + value_txda = value.to("txda") + triton_foo[NBLOCKS, 1, 1]( + a_txda, + b_txda, + c_txda, + value_txda, + BLOCK1, + BLOCK1_SUB, + BLOCK2, + a_txda.shape[0], + a_txda.shape[1], + a_txda.shape[2], + a_txda.shape[3], + a_txda.stride()[0], + a_txda.stride()[1], + a_txda.stride()[2], + a_txda.stride()[3], + c_txda.stride()[0], + c_txda.stride()[1], + c_txda.stride()[2], + c_txda.stride()[3], + ) + with torch.no_grad(): + value.copy_(value_txda.cpu()) + return value + + +def test_npu_indexing(): + Z, Y, X, R = (1, 1, 64, 64) + a = torch.randn((Z, Y, X, R), dtype=torch.float32).cpu() + b = torch.randn((Z, Y, X, R), dtype=torch.float32).cpu() + c = torch.randn((Z, Y, R, X), dtype=torch.float32).cpu() + r = foo_triton_wrapper(a, b, c) + r1 = foo(a, b, c) + print(r[0, 0, 0:8, 0:8]) + print(r1[0, 0, 0:8, 0:8]) + torch.testing.assert_close(r, r1) diff --git a/test/wafer/ops/test_npu_indexing2.py b/test/wafer/ops/test_npu_indexing2.py new file mode 100644 index 00000000..c60be277 --- /dev/null +++ b/test/wafer/ops/test_npu_indexing2.py @@ -0,0 +1,121 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import time + + +def foo(a, b, c): + y = a + b + c + y = y.sum(dim=1) + return y + + +@triton.jit +def triton_codegen2( + in_ptr0, + in_ptr1, + in_ptr2, + out_ptr0, + XBLOCK: tl.constexpr, + XBLOCK_SUB: tl.constexpr, + RBLOCK: tl.constexpr, +): + ynumel = 2 + rnumel = 256 + xnumel = 128 + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + base2 = tl.arange(0, RBLOCK) + loops2: tl.constexpr = rnumel // RBLOCK + for y in range(ynumel): + y0 = y + for loop1 in range(loops1): + x = offset + (loop1 * XBLOCK_SUB) + base1 + x1 = offset + (loop1 * XBLOCK_SUB) + base1[None, :] + _tmp6 = tl.full([XBLOCK_SUB, RBLOCK], 0, tl.float32) + for loop2 in range(loops2): + r2 = loop2 * RBLOCK + base2[:, None] + tmp0 = tl.load( + in_ptr0 + (x1 + (128 * r2) + (128 * 256 * y0)), + None, + eviction_policy="evict_last", + ) + tmp1 = tl.load( + in_ptr1 + (x1 + (128 * r2) + (128 * 256 * y0)), + None, + eviction_policy="evict_last", + ) + tmp3 = tl.load( + in_ptr2 + (x1 + (128 * r2) + (128 * 256 * y0)), + None, + eviction_policy="evict_last", + ) + tmp2 = tmp0 + tmp1 + tmp4 = tmp2 + tmp3 + tmp5 = tl.reshape(tmp4, [RBLOCK, XBLOCK_SUB]) + tmp7 = _tmp6 + tmp5 + _tmp6 = tmp7 + tmp6 = tl.sum(_tmp6, 0).reshape(XBLOCK_SUB) + + tl.store(out_ptr0 + (x + (128 * y0)), tmp6, None) + + +def foo_triton_wrapper(a, b, c): + NBLOCKS = 2 + BLOCK1 = a.shape[2] // NBLOCKS + BLOCK1_SUB = 32 + BLOCK2 = 32 + + value = torch.empty_strided( + (c.shape[0], c.shape[2]), (c.shape[2], 1), dtype=torch.float16 + ).cpu() + + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + value_txda = value.to("txda") + triton_codegen2[NBLOCKS, 1, 1](a_txda, b_txda, c_txda, value_txda, BLOCK1, BLOCK1_SUB, BLOCK2) + with torch.no_grad(): + value.copy_(value_txda.cpu()) + + return value + + +def test_npu_indexing2(): + + Y, X, R = (2, 256, 128) + a = torch.randn((Y, X, R), dtype=torch.float16).cpu() + b = torch.randn((Y, X, R), dtype=torch.float16).cpu() + c = torch.randn((Y, X, R), dtype=torch.float16).cpu() + r = foo_triton_wrapper(a, b, c) + r1 = foo(a, b, c) + print( + r[ + 0:8, + 0:8, + ] + ) + print(r1[0:8, 0:8]) + torch.testing.assert_close(r, r1, rtol=1e-3, atol=1e-3) diff --git a/test/wafer/ops/test_or.py b/test/wafer/ops/test_or.py new file mode 100644 index 00000000..433b2b29 --- /dev/null +++ b/test/wafer/ops/test_or.py @@ -0,0 +1,71 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_or(x0, x1): + res = x0 | x1 + return res + + +@triton.jit +def triton_or( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x_index = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + x_index, None) + tmp1 = tl.load(in_ptr1 + x_index, None) + tmp2 = tmp0 | tmp1 + tl.store(out_ptr0 + x_index, tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["int32", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_or(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + torch_res = torch_or(x0, x1) + # triton结果 + triton_res = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_res_txda = triton_res.to("txda") + triton_or[ncore, 1, 1](x0_txda, x1_txda, triton_res_txda, xblock, xblock_sub) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_permute.py b/test/wafer/ops/test_permute.py new file mode 100644 index 00000000..d8bb98e1 --- /dev/null +++ b/test/wafer/ops/test_permute.py @@ -0,0 +1,172 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import time + + +@triton.jit +def triton_foo( + in_ptr0, + in_ptr1, + in_ptr2, + out_ptr0, + BLOCK1: tl.constexpr, + BLOCK1_SUB: tl.constexpr, + BLOCK2: tl.constexpr, + X: tl.constexpr, + Y: tl.constexpr, + Z: tl.constexpr, + R: tl.constexpr, + Z_STRIDE: tl.constexpr, + Y_STRIDE: tl.constexpr, + X_STRIDE: tl.constexpr, + R_STRIDE: tl.constexpr, + X_STRIDE1: tl.constexpr, + Y_STRIDE1: tl.constexpr, + Z_STRIDE1: tl.constexpr, + R_STRIDE1: tl.constexpr, +): + offset: tl.constexpr = tl.program_id(0) * BLOCK1 + base1 = tl.arange(0, BLOCK1_SUB) + base2 = tl.arange(0, BLOCK2) + nsub: tl.constexpr = BLOCK1 // BLOCK1_SUB + # loops1 : tl.constexpr = nsub * Y * Z + loops1: tl.constexpr = nsub + loops2: tl.constexpr = R // BLOCK2 + + for z in range(Z): + for y in range(Y): + for loop1 in range(loops1): + off1 = loop1 + x = offset + (off1 * BLOCK1_SUB) + base1[:, None] + x1 = offset + (off1 * BLOCK1_SUB) + base1[None, :] + + for loop2 in range(loops2): + r = loop2 * BLOCK2 + base2[None, :] + r1 = loop2 * BLOCK2 + base2[:, None] + tmp0 = tl.load( + in_ptr0 + + ( + (R_STRIDE * r) + + (X_STRIDE * x) + + (Y_STRIDE * y) + + (Z_STRIDE * z) + ), + None, + ) + tmp1 = tl.load( + in_ptr1 + + ( + (R_STRIDE * r) + + (X_STRIDE * x) + + (Y_STRIDE * y) + + (Z_STRIDE * z) + ), + None, + ) + tmp2 = tmp0 + tmp1 + + tmp8 = tl.load( + in_ptr2 + + ( + R_STRIDE1 * r + + X_STRIDE1 * x + + (Y_STRIDE1 * y) + + (Z_STRIDE1 * z) + ), + None, + ) + tmp9 = tmp8 + tmp2 + tl.store( + out_ptr0 + + ( + R_STRIDE1 * r + + X_STRIDE1 * x + + (Y_STRIDE1 * y) + + (Z_STRIDE1 * z) + ), + tmp9, + None, + ) + + +def foo_triton_wrapper(a, b, c): + NBLOCKS = 32 if c.shape[0] >= 256 else 1 + BLOCK1 = c.shape[0] // NBLOCKS + BLOCK1_SUB = BLOCK1 if BLOCK1 < 64 else 64 + BLOCK2 = c.shape[3] if c.shape[3] < 64 else 64 + + value = torch.empty_strided( + (c.shape[0], c.shape[1], c.shape[2], c.shape[3]), + (c.stride()[0], c.stride()[1], c.stride()[2], c.stride()[3]), + dtype=torch.float32, + ).cpu() + + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + value_txda = value.to("txda") + triton_foo[NBLOCKS, 1, 1]( + a_txda, + b_txda, + c_txda, + value_txda, + BLOCK1, + BLOCK1_SUB, + BLOCK2, + c_txda.shape[0], + c_txda.shape[1], + c_txda.shape[2], + c_txda.shape[3], + a_txda.stride()[0], + a_txda.stride()[1], + a_txda.stride()[2], + a_txda.stride()[3], + c_txda.stride()[0], + c_txda.stride()[1], + c_txda.stride()[2], + c_txda.stride()[3], + ) + with torch.no_grad(): + value.copy_(value_txda.cpu()) + return value + + +def foo(a, b, c): + y = a + b + y = c + y.permute(2, 1, 0, 3) + return y + + +def test_permute_handwritten(): + + Z, Y, X, R = (1, 12, 4096, 8) + a = torch.randn((Z, Y, X, R), dtype=torch.float32).cpu() + b = torch.randn((Z, Y, X, R), dtype=torch.float32).cpu() + c = torch.randn((X, Y, Z, R), dtype=torch.float32).cpu() + r = foo_triton_wrapper(a, b, c) + r1 = foo(a, b, c) + print(r[0, 0, 0:8, 0:8]) + print(r1[0, 0, 0:8, 0:8]) + torch.testing.assert_close(r, r1, rtol=1e-3, atol=1e-3) diff --git a/test/wafer/ops/test_permute_full.py b/test/wafer/ops/test_permute_full.py new file mode 100644 index 00000000..6eca87a8 --- /dev/null +++ b/test/wafer/ops/test_permute_full.py @@ -0,0 +1,203 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_021(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + # XB,YB,1 + X = tl.load(x_ptr + idx) + + ret = tl.permute(X, (0, 2, 1)) + + oidx = ( + xidx[:, None, None] * YB * ZB + zidx[None, :, None] * YB + yidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@triton.jit +def fn_npu_102(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + # XB,YB,1 + X = tl.load(x_ptr + idx) + + ret = tl.permute(X, (1, 0, 2)) + + oidx = ( + yidx[:, None, None] * XB * ZB + xidx[None, :, None] * ZB + zidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@triton.jit +def fn_npu_210(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + # XB,YB,1 + X = tl.load(x_ptr + idx) + + ret = tl.permute(X, (2, 1, 0)) + + oidx = ( + zidx[:, None, None] * YB * XB + yidx[None, :, None] * XB + xidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@triton.jit +def fn_npu_201(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + # XB,YB,1 + X = tl.load(x_ptr + idx) + + ret = tl.permute(X, (2, 0, 1)) + + oidx = ( + zidx[:, None, None] * YB * XB + xidx[None, :, None] * YB + yidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@triton.jit +def fn_npu_120(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + # XB,YB,1 + X = tl.load(x_ptr + idx) + + ret = tl.permute(X, (1, 2, 0)) + + oidx = ( + yidx[:, None, None] * ZB * XB + zidx[None, :, None] * XB + xidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@pytest.mark.parametrize( + "para_type,data_type,XB,YB,ZB", + [ + # ['float32',eval('torch.float32'),2,4,3], + ["float32", eval("torch.float32"), 2, 4, 8], + # ['float32',eval('torch.float32'),2,4,37], + ["float32", eval("torch.float32"), 2, 4, 64], + # ['float32',eval('torch.float32'),2,4,781], + # ['float16',eval('torch.float16'),2,4,3], + ["float16", eval("torch.float16"), 2, 4, 8], + # ['float16',eval('torch.float16'),2,4,37], + ["float16", eval("torch.float16"), 2, 4, 64], + # ['float16',eval('torch.float16'),2,4,781], + # ['int8',eval('torch.int8'),2,4,3], + ["int8", eval("torch.int8"), 2, 4, 8], + # ['int8',eval('torch.int8'),2,4,37], + ["int8", eval("torch.int8"), 2, 4, 64], + # ['int8',eval('torch.int8'),2,4,781], + ], +) +def test_permute(para_type, data_type, XB, YB, ZB): + + x = torch.randint(low=0, high=2, size=(XB, YB, ZB), dtype=data_type).cpu() + + output = torch.randint(1, (XB, ZB, YB), dtype=data_type).cpu() + torch_021 = torch.permute(x, (0, 2, 1)) + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_021[1, 1, 1](output_txda, x_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch_021) + + print(" test permute 021 passed") + + output = torch.randint(1, (YB, XB, ZB), dtype=data_type).cpu() + torch_102 = torch.permute(x, (1, 0, 2)) + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_102[1, 1, 1](output_txda, x_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch_102) + + print(" test permute 102 passed") + + output = torch.randint(1, (ZB, XB, YB), dtype=data_type).cpu() + torch_201 = torch.permute(x, (2, 0, 1)) + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_201[1, 1, 1](output_txda, x_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch_201) + + print(" test permute 201 passed") + + output = torch.randint(1, (ZB, YB, XB), dtype=data_type).cpu() + torch_210 = torch.permute(x, (2, 1, 0)) + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_210[1, 1, 1](output_txda, x_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch_210) + + print(" test permute 210 passed") + + output = torch.randint(1, (YB, ZB, XB), dtype=data_type).cpu() + torch_120 = torch.permute(x, (1, 2, 0)) + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_120[1, 1, 1](output_txda, x_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch_120) + + print(" test permute 120 passed") diff --git a/test/wafer/ops/test_permute_reshape.py b/test/wafer/ops/test_permute_reshape.py new file mode 100644 index 00000000..4ab75c62 --- /dev/null +++ b/test/wafer/ops/test_permute_reshape.py @@ -0,0 +1,102 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import time + + +@triton.jit +def triton_foo( + in_ptr0, + in_ptr1, + in_ptr2, + out_ptr0, + BLOCK1: tl.constexpr, + BLOCK1_SUB: tl.constexpr, + BLOCK2: tl.constexpr, + S: tl.constexpr, + N: tl.constexpr, + D: tl.constexpr, +): + offset: tl.constexpr = tl.program_id(0) * BLOCK1 + base1 = tl.arange(0, BLOCK1_SUB) + base2 = tl.arange(0, BLOCK2) + loops1: tl.constexpr = BLOCK1 // BLOCK1_SUB + loops2: tl.constexpr = D // BLOCK2 + + for loop1 in range(loops1): + off1 = loop1 + s = offset + (off1 * BLOCK1_SUB) + base1[:, None] + for n in range(N): + for loop2 in range(loops2): + d = loop2 * BLOCK2 + base2[None, :] + tmp0 = tl.load(in_ptr0 + ((32768 * n) + (8 * s) + d), None) + tmp1 = tl.load(in_ptr1 + ((32768 * n) + (8 * s) + d), None) + tmp2 = tmp0 + tmp1 + + tmp3 = tl.load(in_ptr2 + ((8 * n) + d + (96 * s)), None) + tmp9 = tmp3 + tmp2 + tl.store(out_ptr0 + ((8 * n) + d + (96 * s)), tmp9, None) + + +def foo_triton_wrapper(a, b, c): + NBLOCKS = 32 if a.shape[2] >= 256 else 1 + BLOCK1 = a.shape[2] // NBLOCKS + BLOCK1_SUB = BLOCK1 if BLOCK1 < 64 else 64 + BLOCK2 = a.shape[3] if a.shape[3] < 64 else 64 + + value = torch.empty_strided( + (c.shape[0], c.shape[1], c.shape[2]), + (c.stride()[0], c.stride()[1], c.stride()[2]), + dtype=torch.float32, + ).cpu() + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + value_txda = value.to("txda") + triton_foo[NBLOCKS, 1, 1]( + a_txda, b_txda, c_txda, value_txda, BLOCK1, BLOCK1_SUB, BLOCK2, a_txda.shape[2], a_txda.shape[1], a_txda.shape[3] + ) + with torch.no_grad(): + value.copy_(value_txda.cpu()) + + return value + + +def foo(a, b, c): + B, N, S, D = (1, 12, 4096, 8) + y = a + b + y = c + y.permute(2, 0, 1, 3).reshape(S, B, N * D) + return y + + +def test_permute_reshape(): + B, N, S, D = (1, 12, 4096, 8) + a = torch.randn((B, N, S, D), dtype=torch.float32).cpu() + b = torch.randn((B, N, S, D), dtype=torch.float32).cpu() + c = torch.randn((S, B, N * D), dtype=torch.float32).cpu() + r = foo_triton_wrapper(a, b, c) + r1 = foo(a, b, c) + print(r[0:8, 0, 0:8]) + print(r1[0:8, 0, 0:8]) + torch.testing.assert_close(r, r1, rtol=1e-3, atol=1e-3) diff --git a/test/wafer/ops/test_precise_div.py b/test/wafer/ops/test_precise_div.py new file mode 100644 index 00000000..c9ab478c --- /dev/null +++ b/test/wafer/ops/test_precise_div.py @@ -0,0 +1,71 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import test_common + + +def torch_divRn(x0, x1): + return x0 / x1 + + +@triton.jit +def triton_divRn( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = XBLOCK // XBLOCK_SUB + for loop1 in range(loops1): + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tl.div_rn(tmp0, tmp1) + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 32, 2048, 64], + ], +) +def test_divRn(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype) + x2 = x1.masked_fill(x1 == 0, 1) + x2 = x2.cpu() + y_ref = torch_divRn(x0, x2) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x2_txda = x2.to("txda") + y_cal_txda = y_cal.to("txda") + triton_divRn[ncore, 1, 1](x0_txda, x2_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_precise_sqrt.py b/test/wafer/ops/test_precise_sqrt.py new file mode 100644 index 00000000..48ec2289 --- /dev/null +++ b/test/wafer/ops/test_precise_sqrt.py @@ -0,0 +1,60 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + +torch.set_printoptions(precision=10) + + +@triton.jit +def sqrtrn_kernel(x_ptr, y_ptr, output_ptr, n_elements, BLOCK_SIZE: tl.constexpr): + id = tl.program_id(axis=0) + start = id * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_elements + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + + output = x + tl.sqrt_rn(y) + tl.store(output_ptr + offsets, output, mask=mask) + + +def sqrtrn(x: torch.Tensor, y: torch.Tensor): + output = torch.empty_like(y) + grid = lambda meta: (triton.cdiv(output.numel(), meta["BLOCK_SIZE"]),) + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + sqrtrn_kernel[grid](x_txda, y_txda, output_txda, output_txda.numel(), BLOCK_SIZE=512) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + return output + + +def test_sqrtrn_fp32(): + size = 10240 + x = torch.abs(torch.randn(size, device="cpu", dtype=torch.float32)) + y = torch.abs(torch.randn(size, device="cpu", dtype=torch.float32)) + ref = x + torch.sqrt(y) + cal = sqrtrn(x, y) + torch.testing.assert_close(cal, ref, rtol=1e-06, atol=1e-06, equal_nan=True) diff --git a/test/wafer/ops/test_ravel.py b/test/wafer/ops/test_ravel.py new file mode 100644 index 00000000..3717a739 --- /dev/null +++ b/test/wafer/ops/test_ravel.py @@ -0,0 +1,75 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.ravel(X) + + oidx = tl.arange(0, XB * YB * ZB) + tl.store(output_ptr + oidx, ret) + + +testlist = [ + ("float32", torch.float32, 2, 256, 16), + ("float32", torch.float32, 8, 8, 4), + ("float16", torch.float16, 2, 256, 16), + ("float16", torch.float16, 8, 8, 4), + ("int8", torch.int8, 2, 256, 16), + ("int8", torch.int8, 8, 8, 4), +] + + +@pytest.mark.parametrize("sigtype, dtype, XB, YB, ZB", testlist) +def test_ravel(sigtype, dtype, XB, YB, ZB): + + x = torch.randint(low=-128, high=128, size=(XB, YB, ZB), dtype=dtype).cpu() + ans = torch.ravel(x) + + print(ans[0:16]) + + output = torch.randint(1, (XB * YB * ZB,), dtype=dtype).cpu() + + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_[1, 1, 1](output_txda, x_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + print(output[0:16]) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_reduce_count_vector.py b/test/wafer/ops/test_reduce_count_vector.py new file mode 100644 index 00000000..c0e8c934 --- /dev/null +++ b/test/wafer/ops/test_reduce_count_vector.py @@ -0,0 +1,181 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_count(x0, cmp_val, dim): + res = (x0 == cmp_val).sum(dim=dim) + return res + + +def standard_gt(x0, cmp_val, dim): + res = (x0 > cmp_val).sum(dim=dim) + return res + + +def standard_lt(x0, cmp_val, dim): + res = (x0 < cmp_val).sum(dim=dim) + return res + + +@triton.jit +def triton_count( + in_ptr0, out_ptr0, cmp_val, dim: tl.constexpr, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, N) + x = tl.load(in_ptr0 + idx_block) + + tmp3 = x == cmp_val + # tmp3 bool -> tl.float32 + tmp4 = tmp3.to(tl.float32) + res = tl.sum(tmp4, dim) + + # Reducing the input vector produces one count; the output is a scalar. + tl.store(out_ptr0, res) + + +@triton.jit +def triton_gt( + in_ptr0, out_ptr0, cmp_val, dim: tl.constexpr, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, N) + x = tl.load(in_ptr0 + idx_block) + + tmp3 = x > cmp_val + # tmp3 bool -> tl.float32 + tmp4 = tmp3.to(tl.float32) + res = tl.sum(tmp4, dim) + + # Reducing the input vector produces one count; the output is a scalar. + tl.store(out_ptr0, res) + + +@triton.jit +def triton_lt( + in_ptr0, out_ptr0, cmp_val, dim: tl.constexpr, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, N) + x = tl.load(in_ptr0 + idx_block) + + tmp3 = x < cmp_val + # tmp3 bool -> tl.float32 + tmp4 = tmp3.to(tl.float32) + res = tl.sum(tmp4, dim) + + # Reducing the input vector produces one count; the output is a scalar. + tl.store(out_ptr0, res) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + (torch.bfloat16, "bfloat16"), + (torch.int8, "int8"), + (torch.int16, "int16"), + (torch.int32, "int32"), + (torch.int64, "int64"), +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (32, 32), +] + +map_for_64_t = {37: 31} + +CPM_VAL_INT = 8 +CPM_VAL_FLOAT = 0.5 + +# TO BE FIXED with mask +ops = [ + ("counti", triton_count, standard_count, CPM_VAL_INT), + ("countf", triton_gt, standard_gt, CPM_VAL_FLOAT), + ("countf", triton_lt, standard_lt, CPM_VAL_FLOAT), +] + + +def judge_continue(opName, sigtype): + if opName == "counti" and "int" in sigtype: + return False + if opName == "countf" and "float" in sigtype: + return False + return True + + +@pytest.mark.parametrize("opName, tritonOp, standOp, cmp_val", ops) +@pytest.mark.parametrize("dtype, sigtype", types) +@pytest.mark.parametrize("N, NUMEL", shapes) +def test_reduce_count_vector( + opName, tritonOp, standOp, cmp_val, dtype, sigtype, N, NUMEL +): + if judge_continue(opName, sigtype): + return + torch.manual_seed(0) + torch.txda.set_device(0) + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + ans = standOp(x0, cmp_val, 0) + x0 = x0.cpu() + + output = torch.tensor(0, dtype=torch.float32).cpu() + x0_txda = x0.to("txda") + # The reduction must write exactly one scalar, including at an offset. + storage = torch.full((129,), -83.0, dtype=torch.float32).to("txda") + output_txda = storage[64:65].view(()) + tritonOp[1, 1, 1](x0_txda, output_txda, cmp_val, dim=0, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + guards = storage.cpu() + assert torch.equal(guards[:64], torch.full((64,), -83.0)) + assert torch.equal(guards[65:], torch.full((64,), -83.0)) + output = output.cpu().to(torch.int32) + # print(f'x0:{x0}\ntriton:{output}\ntorch:{ans}') + assert torch.equal(output, ans) + + +if __name__ == "__main__": + dtype = torch.float32 + sigtype = "float32" + allshape = [(3, 32)] + for shape in allshape: + test_reduce_count_vector( + "countf", + triton_lt, + standard_lt, + CPM_VAL_FLOAT, + dtype, + sigtype, + shape[0], + shape[1], + ) diff --git a/test/wafer/ops/test_reduce_mean.py b/test/wafer/ops/test_reduce_mean.py new file mode 100644 index 00000000..e1256ad6 --- /dev/null +++ b/test/wafer/ops/test_reduce_mean.py @@ -0,0 +1,112 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common +import numpy as np + + +def numpy_mean_pr(x0, x1): + res = np.mean(x0, axis=-1) + x1 + return res + + +@triton.jit +def triton_mean_pr( + out_ptr0, + in_ptr0, + in_ptr1, + xnumel, + rnumel, + XBLOCK: tl.constexpr, + XBLOCK_SUB: tl.constexpr, + RBLOCK: tl.constexpr, +): + xoffset = tl.program_id(0) * XBLOCK + rindex = tl.arange(0, RBLOCK)[None, :] + rmask = rindex < rnumel + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB) + xmask = xindex[:, None] < xnumel + x0 = xindex + r1 = rindex + tmp0 = tl.load(in_ptr0 + (r1 + (RBLOCK * x0[:, None])), xmask & rmask) + tmp4 = tl.load(in_ptr1 + (x0), xindex < xnumel) + tmp1 = tl.reshape(tmp0, [XBLOCK_SUB, RBLOCK]) + tmp3 = tl.sum(tmp1, 1) / RBLOCK + tmp5 = tmp3 + tmp4 + tl.store(out_ptr0 + (xindex), tmp5, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (8, 8, 4), 8, 2], + ["float32", (8, 8, 64), 8, 2], + ["float32", (8, 8, 1024), 8, 2], + ["float16", (8, 8, 4), 8, 2], + ["float16", (8, 8, 64), 8, 2], + ["float16", (8, 8, 1024), 8, 2], + ["int8", (8, 8, 4), 8, 2], + ["int8", (8, 8, 64), 8, 2], + ["int8", (8, 8, 1024), 8, 2], + ], +) +def test_mean_pr(param_list): + dtype, shape, ncore, xblock_sub = param_list + import math + + numel = math.prod(shape) + xblock = numel // shape[-1] // ncore + rblock = shape[-1] + assert ncore * xblock * shape[-1] == numel + xn1 = np.random.randn(shape[0], shape[1], shape[2]).astype(eval("np." + dtype)) + xn2 = np.random.randn(shape[0], shape[1]).astype(eval("np." + dtype)) + x0 = torch.tensor(xn1).cpu() + x1 = torch.tensor(xn2).cpu() + y_ref = numpy_mean_pr(xn1, xn2) + if dtype == "int8": + y_cal = test_common.generate_tensor(shape[:-1], "float32").cpu() + else: + y_cal = test_common.generate_tensor(shape[:-1], dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_mean_pr[ncore, 1, 1]( + y_cal_txda, x0_txda, x1_txda, x1_txda.numel(), rblock, xblock, xblock_sub, rblock + ) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + if dtype == "int8": + assert torch.allclose( + torch.tensor(y_ref.astype(np.float32)).cpu(), + y_cal, + rtol=1e-03, + atol=1e-03, + equal_nan=True, + ) + else: + assert torch.allclose( + torch.tensor(y_ref).cpu(), y_cal, rtol=1e-03, atol=1e-03, equal_nan=True + ) diff --git a/test/wafer/ops/test_reduce_sum.py b/test/wafer/ops/test_reduce_sum.py new file mode 100644 index 00000000..6c562bf2 --- /dev/null +++ b/test/wafer/ops/test_reduce_sum.py @@ -0,0 +1,109 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import numpy as np + +import os +import sys + +parent_dir = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +sys.path.append(parent_dir) +import test_common + + +# PR: Pointiwise-Reduction pattern, reduction in last axis +def numpy_sum_pr(x0, x1): + res = np.sum(x0, axis=-1) + x1 + return res + + +@triton.jit +def triton_sum_pr( + out_ptr0, + in_ptr0, + in_ptr1, + xnumel, + rnumel, + XBLOCK: tl.constexpr, + XBLOCK_SUB: tl.constexpr, + RBLOCK: tl.constexpr, +): + xoffset = tl.program_id(0) * XBLOCK + rindex = tl.arange(0, RBLOCK)[None, :] + rmask = rindex < rnumel + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB) + xmask = xindex[:, None] < xnumel + x0 = xindex + r1 = rindex + tmp0 = tl.load(in_ptr0 + (r1 + (RBLOCK * x0[:, None])), xmask & rmask) + tmp4 = tl.load(in_ptr1 + (x0), xindex < xnumel) + tmp1 = tl.reshape(tmp0, [XBLOCK_SUB, RBLOCK]) + tmp3 = tl.sum(tmp1, 1) + tmp5 = tmp3 + tmp4 + tl.store(out_ptr0 + (xindex), tmp5, None) + + +# fp16 use numpy +@pytest.mark.parametrize( + "param_list", + [ + # ['float32', (8, 8, 4), 8, 2], + ["float32", (8, 8, 64), 8, 2], + ["float32", (8, 8, 512), 8, 2], + ["float16", (8, 8, 4), 8, 2], + ["float16", (8, 8, 64), 8, 2], + ["float16", (8, 8, 512), 8, 2], + ["int8", (8, 8, 4), 8, 2], + ["int8", (8, 8, 64), 8, 2], + ["int8", (8, 8, 512), 8, 2], + ], +) +def test_sum_pr(param_list): + dtype, shape, ncore, xblock_sub = param_list + import math + + numel = math.prod(shape) + xblock = numel // shape[-1] // ncore + rblock = shape[-1] + assert ncore * xblock * shape[-1] == numel + xn1 = np.random.randn(shape[0], shape[1], shape[2]).astype(eval("np." + dtype)) + xn2 = np.random.randn(shape[0], shape[1]).astype(eval("np." + dtype)) + x0 = torch.tensor(xn1).cpu() + x1 = torch.tensor(xn2).cpu() + if dtype == "int8": + y_ref = numpy_sum_pr(xn1, xn2).astype(np.int8) + else: + y_ref = numpy_sum_pr(xn1, xn2) + y_cal = test_common.generate_tensor(shape[:-1], dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + triton_sum_pr[ncore, 1, 1]( + y_cal_txda, x0_txda, x1_txda, x1_txda.numel(), rblock, xblock, xblock_sub, rblock + ) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, torch.tensor(y_ref).cpu()) diff --git a/test/wafer/ops/test_reshape.py b/test/wafer/ops/test_reshape.py new file mode 100644 index 00000000..c16bf369 --- /dev/null +++ b/test/wafer/ops/test_reshape.py @@ -0,0 +1,75 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.reshape(X, (ZB, XB * YB)) + + oidx = tl.arange(0, ZB)[:, None] * XB * YB + tl.arange(0, XB * YB)[None, :] + + tl.store(output_ptr + oidx, ret) + + +testlist = [ + ("float32", torch.float32, 2, 256, 16), + ("float32", torch.float32, 8, 8, 4), + ("float16", torch.float16, 2, 256, 16), + ("float16", torch.float16, 8, 8, 4), + ("int8", torch.int8, 2, 256, 16), + ("int8", torch.int8, 8, 8, 4), +] + + +@pytest.mark.parametrize("sigtype, dtype, XB, YB, ZB", testlist) +def test_ravel(sigtype, dtype, XB, YB, ZB): + + x = torch.randint(low=-128, high=128, size=(XB, YB, ZB), dtype=dtype).cpu() + ans = torch.reshape(x, (ZB, XB * YB)) + + print(ans[0, 0:16]) + + output = torch.randint(1, (ZB, XB * YB), dtype=dtype).cpu() + + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_[1, 1, 1](output_txda, x_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output[0, 0:16]) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_rms_norm.py b/test/wafer/ops/test_rms_norm.py new file mode 100644 index 00000000..09937393 --- /dev/null +++ b/test/wafer/ops/test_rms_norm.py @@ -0,0 +1,154 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import torch +import triton +import triton.language as tl +import torch_txda # noqa: F401 + + +@triton.jit +def _rms_norm_fwd_fused( + X, # pointer to the input + Y, # pointer to the output + W, # pointer to the weights + stride, # how much to increase the pointer when moving by 1 row + N, # number of columns in X + eps, # epsilon to avoid division by zero + BLOCK_SIZE: tl.constexpr, +): + # Map the program id to the row of X and Y it should compute. + row = tl.program_id(0) + Y += row * stride + X += row * stride + # Compute variance + _var = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + x = tl.load(X + cols, mask=cols < N, other=0.0).to(tl.float32) + _var += x * x + var = tl.sum(_var, axis=0) / N + rstd = 1 / tl.sqrt(var + eps) + # Normalize and apply linear transformation + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + mask = cols < N + w = tl.load(W + cols, mask=mask).to(tl.float32) + x = tl.load(X + cols, mask=mask, other=0.0).to(tl.float32) + x_hat = x * rstd + y = x_hat * w + # Write output + tl.store(Y + cols, y.to(tl.float16), mask=mask) + + +# have to change the block_size +@torch.inference_mode() +def rms_norm(x, weight, eps, out=None): + # allocate output, tl.store save y in tl.float16 + y = torch.empty_like(x, dtype=torch.float16) if out is None else out + # reshape input data into 2D tensor + x_arg = x.view(-1, x.shape[-1]) + M, N = x_arg.shape + # Less than 64KB per feature: enqueue fused kernel + MAX_FUSED_SIZE = 65536 // x.element_size() + BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(N)) + if N > BLOCK_SIZE: + raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.") + # heuristics for number of warps + num_warps = min(max(BLOCK_SIZE // 256, 1), 8) + BLOCK_SIZE = 128 * 2 * 2 * 2 * 2 * 2 * 2 + num_warps = 8 + # enqueue kernel + x_arg_txda = x_arg.to("txda") + y_txda = y.to("txda") + weight_txda = weight.to("txda") + kernel = _rms_norm_fwd_fused[(M,)]( + x_arg_txda, + y_txda, + weight_txda, + x_arg_txda.stride(0), + N, + eps, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=num_warps, + ) + with torch.no_grad(): + y.copy_(y_txda.cpu()) + return y, kernel + + +def _rms_norm(shape, datatype): + x = torch.randn(shape[0], shape[1], dtype=datatype, device="cpu") + weight = torch.randn(shape[1], dtype=datatype, device="cpu") + y, kernel = rms_norm(x, weight, eps=1e-5) + eps1 = 1e-5 + if datatype == torch.bfloat16 or datatype == torch.float16: + x = x.to(torch.float32) + rms = torch.sqrt(x.pow(2).mean(-1, keepdim=True) + eps1) # 计算均方根 + x_norm = x / rms # 标准化 + y_ref = weight * x_norm + y_ref = y_ref.to(torch.float16) + torch.testing.assert_close(y_ref, y, rtol=1e-3, atol=1e-3) + + +def test_cases(): + _rms_norm((16, 256), torch.float16) + _rms_norm((16, 256), torch.float32) + _rms_norm((16, 256), torch.bfloat16) + _rms_norm((128, 3), torch.bfloat16) + _rms_norm((128, 16), torch.bfloat16) + _rms_norm((128, 37), torch.bfloat16) + _rms_norm((128, 64), torch.bfloat16) + _rms_norm((16, 256), torch.float16) + _rms_norm((16, 256), torch.float32) + _rms_norm((16, 256), torch.bfloat16) + + _rms_norm((64, 64), torch.float16) + _rms_norm((64, 64), torch.float32) + _rms_norm((64, 64), torch.bfloat16) + + _rms_norm((1, 128), torch.float16) + _rms_norm((1, 128), torch.float32) + _rms_norm((1, 128), torch.bfloat16) + + _rms_norm((33, 128), torch.float16) + _rms_norm((33, 128), torch.float32) + _rms_norm((33, 128), torch.bfloat16) + + _rms_norm((128, 3), torch.float16) + _rms_norm((128, 3), torch.float32) + _rms_norm((128, 3), torch.bfloat16) + + _rms_norm((128, 16), torch.float16) + _rms_norm((128, 16), torch.float32) + _rms_norm((128, 16), torch.bfloat16) + + _rms_norm((128, 37), torch.float16) + _rms_norm((128, 37), torch.float32) + _rms_norm((128, 37), torch.bfloat16) + + _rms_norm((128, 64), torch.float16) + _rms_norm((128, 64), torch.float32) + _rms_norm((128, 64), torch.bfloat16) + + _rms_norm((128, 181), torch.float16) + _rms_norm((128, 181), torch.float32) + _rms_norm((128, 181), torch.bfloat16) diff --git a/test/wafer/ops/test_rotary_embedding.py b/test/wafer/ops/test_rotary_embedding.py new file mode 100644 index 00000000..b545f334 --- /dev/null +++ b/test/wafer/ops/test_rotary_embedding.py @@ -0,0 +1,191 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +"""Rotary embedding kernel implemented by Triton. + +GPT-NeoX style +""" + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def rotary_embedding_kernel( + state, # [num_tokens, head_num, head_dim] + cos, # [num_tokens, 1, head_dim // 2] + sin, # [num_tokens, 1, head_dim // 2] + stride_state_n, + stride_state_h, + stride_state_d, + stride_cos_n, + stride_cos_d, + # stride_sin_n, + # stride_sin_d, + num_tokens, + num_heads, + BLOCK_N: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_D: tl.constexpr, +): + token_index = tl.program_id(0) + token_range = token_index * BLOCK_N + tl.arange(0, BLOCK_N) + head_index = tl.program_id(1) + head_range = head_index * BLOCK_H + tl.arange(0, BLOCK_H) + + dim_range_x = tl.arange(0, BLOCK_D // 2) + dim_range_y = tl.arange(BLOCK_D // 2, BLOCK_D) + + state_x_offset = ( + token_range[:, None, None] * stride_state_n + + head_range[None, :, None] * stride_state_h + + dim_range_x[None, None, :] * stride_state_d + ) + state_y_offset = ( + token_range[:, None, None] * stride_state_n + + head_range[None, :, None] * stride_state_h + + dim_range_y[None, None, :] * stride_state_d + ) + + cos_sim_offset = ( + token_range[:, None, None] * stride_cos_n + + dim_range_x[None, None, :] * stride_cos_d + ) + + state_x = tl.load( + state + state_x_offset, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + other=0.0, + ) + state_y = tl.load( + state + state_y_offset, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + other=0.0, + ) + + cos_loaded = tl.load( + cos + cos_sim_offset, + mask=token_range[:, None, None] < num_tokens, + other=0.0, + ) + sin_loaded = tl.load( + sin + cos_sim_offset, + mask=token_range[:, None, None] < num_tokens, + other=0.0, + ) + + out_x = state_x * cos_loaded - state_y * sin_loaded + out_y = state_x * sin_loaded + state_y * cos_loaded + + tl.store( + state + state_x_offset, + out_x, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + ) + tl.store( + state + state_y_offset, + out_y, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + ) + + +@torch.inference_mode() +def rotary_embedding(state, cos, sin): + num_tokens = state.shape[0] + num_heads = state.shape[1] + head_dim = state.shape[2] + + # BLOCK_N = 32 + BLOCK_N = 16 + BLOCK_H = 4 + grid = ( + triton.cdiv(num_tokens, BLOCK_N), + triton.cdiv(num_heads, BLOCK_H), + ) + if head_dim >= 128: + num_warps = 8 + else: + num_warps = 4 + + state_txda = state.to("txda") + cos_txda = cos.to("txda") + sin_txda = sin.to("txda") + kernel = rotary_embedding_kernel[grid]( + state_txda, + cos_txda, + sin_txda, + state_txda.stride(0), + state_txda.stride(1), + state_txda.stride(2), + cos_txda.stride(0), + cos_txda.stride(2), + # sin.stride(0), + # sin.stride(2), + num_tokens, + num_heads, + BLOCK_N=BLOCK_N, + BLOCK_H=BLOCK_H, + BLOCK_D=head_dim, + num_warps=num_warps, + num_stages=1, + ) + with torch.no_grad(): + state.copy_(state_txda.cpu()) + return + + +def torch_rotary_embedding(state, cos, sin): + _, _, dim = state.shape + state_x = state[:, :, 0 : dim // 2] + state_y = state[:, :, dim // 2 : dim] + out_x = state_x * cos - state_y * sin + out_y = state_x * sin + state_y * cos + return torch.cat((out_x, out_y), dim=-1) + + +def rotary_emb(tokens, heads, headdim, dtype): + tokens_num = tokens + num_heads = heads + head_dim = headdim + max_positions = 1024 + + # torch.float16 has floating point problem in Triton 2.0.0 + # But it works fine in Triton 2.1.0 + state = torch.randn((tokens_num, num_heads, head_dim), dtype=dtype, device="cpu") + cos_shape = (tokens_num, 1, head_dim // 2) + cos = -1.2 + 0.5 * torch.randn(cos_shape, dtype=dtype, device="cpu") + sin = -2.0 + 0.5 * torch.randn(cos_shape, dtype=dtype, device="cpu") + # forward pass + torch_result = torch_rotary_embedding(state, cos, sin) + rotary_embedding(state, cos, sin) + triton_result = state # state is modified in-place + torch.testing.assert_close(torch_result, triton_result, rtol=1e-3, atol=1e-3) + + +def test_cases(): + rotary_emb(256, 96, 128, torch.float16) + rotary_emb(256, 96, 128, torch.float32) diff --git a/test/wafer/ops/test_rotatry_gpt.py b/test/wafer/ops/test_rotatry_gpt.py new file mode 100644 index 00000000..09e7292e --- /dev/null +++ b/test/wafer/ops/test_rotatry_gpt.py @@ -0,0 +1,202 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +""" +Rotary embedding kernel implemented by Triton. +GPT-J style +""" + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def rotary_embedding_kernel( + state, # [num_tokens, head_num, head_dim] + cos, # [num_tokens, 1, head_dim // 2] + sin, # [num_tokens, 1, head_dim // 2] + stride_state_n, + stride_state_h, + stride_state_d, + stride_cos_n, + stride_cos_d, + # stride_sin_n, + # stride_sin_d, + num_tokens, + num_heads, + BLOCK_N: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_D: tl.constexpr, +): + token_index = tl.program_id(0) + token_range = token_index * BLOCK_N + tl.arange(0, BLOCK_N) + head_index = tl.program_id(1) + head_range = head_index * BLOCK_H + tl.arange(0, BLOCK_H) + + dim_range = tl.arange(0, BLOCK_D // 2) + dim_range_x = dim_range * 2 + dim_range_y = dim_range * 2 + 1 + + # tl.device_print("dim x", dim_range_x) + # tl.device_print("dim y", dim_range_y) + + state_x_offset = ( + token_range[:, None, None] * stride_state_n + + head_range[None, :, None] * stride_state_h + + dim_range_x[None, None, :] * stride_state_d + ) + + state_y_offset = ( + token_range[:, None, None] * stride_state_n + + head_range[None, :, None] * stride_state_h + + dim_range_y[None, None, :] * stride_state_d + ) + + cos_sim_offset = ( + token_range[:, None, None] * stride_cos_n + + dim_range[None, None, :] * stride_cos_d + ) + + state_x = tl.load( + state + state_x_offset, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + other=0.0, + ) + state_y = tl.load( + state + state_y_offset, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + other=0.0, + ) + + cos_loaded = tl.load( + cos + cos_sim_offset, + mask=token_range[:, None, None] < num_tokens, + other=0.0, + ) + sin_loaded = tl.load( + sin + cos_sim_offset, + mask=token_range[:, None, None] < num_tokens, + other=0.0, + ) + + out_x = state_x * cos_loaded - state_y * sin_loaded + out_y = state_x * sin_loaded + state_y * cos_loaded + + tl.store( + state + state_x_offset, + out_x, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + ) + tl.store( + state + state_y_offset, + out_y, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + ) + + +@torch.inference_mode() +def rotary_embedding(state, cos, sin): + num_tokens = state.shape[0] + num_heads = state.shape[1] + head_dim = state.shape[2] + + BLOCK_N = 8 + BLOCK_H = 4 + grid = ( + triton.cdiv(num_tokens, BLOCK_N), + triton.cdiv(num_heads, BLOCK_H), + ) + if head_dim >= 128: + num_warps = 8 + else: + num_warps = 4 + + state_txda = state.to("txda") + cos_txda = cos.to("txda") + sin_txda = sin.to("txda") + kernel = rotary_embedding_kernel[grid]( + state_txda, + cos_txda, + sin_txda, + state_txda.stride(0), + state_txda.stride(1), + state_txda.stride(2), + cos_txda.stride(0), + cos_txda.stride(2), + # sin.stride(0), + # sin.stride(2), + num_tokens, + num_heads, + BLOCK_N=BLOCK_N, + BLOCK_H=BLOCK_H, + BLOCK_D=head_dim, + num_warps=num_warps, + num_stages=1, + ) + with torch.no_grad(): + state.copy_(state_txda.cpu()) + # print(kernel.asm['ttir']) + return + + +def torch_rotary_embedding(state, cos, sin): + _, _, dim = state.shape + state_x = state[:, :, 0:dim:2] + state_y = state[:, :, 1:dim:2] + out_x = state_x * cos - state_y * sin + out_y = state_x * sin + state_y * cos + out = torch.empty_like(state).cpu() + out[:, :, 0:dim:2] = out_x + out[:, :, 1:dim:2] = out_y + return out + + +def _rotary_emb(dtype): + tokens_num = 128 + num_heads = 96 + head_dim = 64 + max_positions = 1024 + + # torch.float16 has floating point problem in Triton 2.0.0 + # But it works fine in Triton 2.1.0 + state = torch.randn((tokens_num, num_heads, head_dim), dtype=dtype, device="cpu") + cos_shape = (tokens_num, 1, head_dim // 2) + cos = -1.2 + 0.5 * torch.randn(cos_shape, dtype=dtype, device="cpu") + sin = -2.0 + 0.5 * torch.randn(cos_shape, dtype=dtype, device="cpu") + # forward pass + torch_result = torch_rotary_embedding(state, cos, sin) + rotary_embedding(state, cos, sin) + triton_result = state # state is modified in-place + # print(torch_result[1][0]) + # print(triton_result[1][0]) + # Note: This test is not accurate enough. + assert torch.allclose(torch_result, triton_result, atol=1e-2, rtol=1e-7) + + +def test_rotary_emb(): + _rotary_emb(torch.float16) + _rotary_emb(torch.float32) diff --git a/test/wafer/ops/test_rotaty_embedding_gpt.py b/test/wafer/ops/test_rotaty_embedding_gpt.py new file mode 100644 index 00000000..8a7cfa84 --- /dev/null +++ b/test/wafer/ops/test_rotaty_embedding_gpt.py @@ -0,0 +1,190 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +"""Rotary embedding kernel implemented by Triton. + +GPT-NeoX style +""" + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def rotary_embedding_kernel( + state, # [num_tokens, head_num, head_dim] + cos, # [num_tokens, 1, head_dim // 2] + sin, # [num_tokens, 1, head_dim // 2] + stride_state_n, + stride_state_h, + stride_state_d, + stride_cos_n, + stride_cos_d, + # stride_sin_n, + # stride_sin_d, + num_tokens, + num_heads, + BLOCK_N: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_D: tl.constexpr, +): + token_index = tl.program_id(0) + token_range = token_index * BLOCK_N + tl.arange(0, BLOCK_N) + head_index = tl.program_id(1) + head_range = head_index * BLOCK_H + tl.arange(0, BLOCK_H) + + dim_range_x = tl.arange(0, BLOCK_D // 2) + dim_range_y = tl.arange(BLOCK_D // 2, BLOCK_D) + + state_x_offset = ( + token_range[:, None, None] * stride_state_n + + head_range[None, :, None] * stride_state_h + + dim_range_x[None, None, :] * stride_state_d + ) + state_y_offset = ( + token_range[:, None, None] * stride_state_n + + head_range[None, :, None] * stride_state_h + + dim_range_y[None, None, :] * stride_state_d + ) + + cos_sim_offset = ( + token_range[:, None, None] * stride_cos_n + + dim_range_x[None, None, :] * stride_cos_d + ) + + state_x = tl.load( + state + state_x_offset, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + other=0.0, + ) + state_y = tl.load( + state + state_y_offset, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + other=0.0, + ) + + cos_loaded = tl.load( + cos + cos_sim_offset, + mask=token_range[:, None, None] < num_tokens, + other=0.0, + ) + sin_loaded = tl.load( + sin + cos_sim_offset, + mask=token_range[:, None, None] < num_tokens, + other=0.0, + ) + + out_x = state_x * cos_loaded - state_y * sin_loaded + out_y = state_x * sin_loaded + state_y * cos_loaded + + tl.store( + state + state_x_offset, + out_x, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + ) + tl.store( + state + state_y_offset, + out_y, + mask=(token_range[:, None, None] < num_tokens) + & (head_range[None, :, None] < num_heads), + ) + + +@torch.inference_mode() +def rotary_embedding(state, cos, sin): + num_tokens = state.shape[0] + num_heads = state.shape[1] + head_dim = state.shape[2] + + BLOCK_N = 16 + BLOCK_H = 4 + grid = ( + triton.cdiv(num_tokens, BLOCK_N), + triton.cdiv(num_heads, BLOCK_H), + ) + if head_dim >= 128: + num_warps = 8 + else: + num_warps = 4 + + state_txda = state.to("txda") + cos_txda = cos.to("txda") + sin_txda = sin.to("txda") + kernel = rotary_embedding_kernel[grid]( + state_txda, + cos_txda, + sin_txda, + state_txda.stride(0), + state_txda.stride(1), + state_txda.stride(2), + cos_txda.stride(0), + cos_txda.stride(2), + # sin.stride(0), + # sin.stride(2), + num_tokens, + num_heads, + BLOCK_N=BLOCK_N, + BLOCK_H=BLOCK_H, + BLOCK_D=head_dim, + num_warps=num_warps, + num_stages=1, + ) + with torch.no_grad(): + state.copy_(state_txda.cpu()) + return + + +def torch_rotary_embedding(state, cos, sin): + _, _, dim = state.shape + state_x = state[:, :, 0 : dim // 2] + state_y = state[:, :, dim // 2 : dim] + out_x = state_x * cos - state_y * sin + out_y = state_x * sin + state_y * cos + return torch.cat((out_x, out_y), dim=-1) + + +def rotary_emb(tokens, heads, headdim, dtype): + tokens_num = tokens + num_heads = heads + head_dim = headdim + # max_positions = 1024 + + # torch.float16 has floating point problem in Triton 2.0.0 + # But it works fine in Triton 2.1.0 + state = torch.randn((tokens_num, num_heads, head_dim), dtype=dtype, device="cpu") + cos_shape = (tokens_num, 1, head_dim // 2) + cos = -1.2 + 0.5 * torch.randn(cos_shape, dtype=dtype, device="cpu") + sin = -2.0 + 0.5 * torch.randn(cos_shape, dtype=dtype, device="cpu") + # forward pass + torch_result = torch_rotary_embedding(state, cos, sin) + rotary_embedding(state, cos, sin) + triton_result = state # state is modified in-place + torch.testing.assert_close(torch_result, triton_result, rtol=1e-3, atol=1e-3) + + +def test_cases(): + rotary_emb(256, 96, 128, torch.float16) + rotary_emb(256, 96, 128, torch.float32) diff --git a/test/wafer/ops/test_rshift.py b/test/wafer/ops/test_rshift.py new file mode 100644 index 00000000..fed4f486 --- /dev/null +++ b/test/wafer/ops/test_rshift.py @@ -0,0 +1,105 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import time +import test_common +import torch +import torch_txda # noqa: F401 + + +def standard_unary(x0, dtype): + res = x0 >> 2 + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + tmp = tl.cast(2, tl.int8) + ret = x >> tmp + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + # (torch.float32, 'float32'), + # (torch.float16, 'float16'), + # (torch.bfloat16, 'bfloat16'), + (torch.int8, "int8"), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + # print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + # print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_rsqrt.py b/test/wafer/ops/test_rsqrt.py new file mode 100644 index 00000000..43c77329 --- /dev/null +++ b/test/wafer/ops/test_rsqrt.py @@ -0,0 +1,80 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import numpy as np +import pytest +import test_common + + +def numpy_rsqrt(x0, x1): + res = x0 + 1.0 / (np.sqrt(x1)) + return res + + +@triton.jit +def triton_rsqrt( + in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = xindex < xnumel + x0 = xindex + tmp0 = tl.load(in_ptr0 + (x0), xmask) + tmp1 = tl.load(in_ptr1 + (x0), xmask) + tmp2 = tmp0 + tl.rsqrt(tmp1) + tl.store(out_ptr0 + (xindex), tmp2, xmask) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_rsqrt(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = np.abs(np.random.randn(shape[0], shape[1], shape[2])).astype( + eval("np." + dtype) + ) + x1 = np.abs(np.random.randn(shape[0], shape[1], shape[2])).astype( + eval("np." + dtype) + ) + x0_npu = torch.tensor(x0).cpu() + x1_npu = torch.tensor(x1).cpu() + # numpy结果 + numpy_res = numpy_rsqrt(x0, x1) + # triton结果 + triton_res = test_common.generate_tensor(shape, dtype).cpu() + x0_npu_txda = x0_npu.to("txda") + x1_npu_txda = x1_npu.to("txda") + triton_res_txda = triton_res.to("txda") + triton_rsqrt[ncore, 1, 1]( + x0_npu_txda, x1_npu_txda, triton_res_txda, x0_npu_txda.numel(), xblock, xblock_sub + ) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, triton_res, torch.tensor(numpy_res).cpu()) diff --git a/test/wafer/ops/test_scalar_calc.py b/test/wafer/ops/test_scalar_calc.py new file mode 100644 index 00000000..6f487fa5 --- /dev/null +++ b/test/wafer/ops/test_scalar_calc.py @@ -0,0 +1,774 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +from triton.language.extra import libdevice +import pytest +import test_common + + +### add +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_add_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tmp0 + 2.0 + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = y + 2.0 + return torch.tensor(y) + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### sub +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_sub_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tmp0 - 2.0 + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = y - 2.0 + return torch.tensor(y) + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### mul +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_mul_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tmp0 * 2.0 + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = y * 2.0 + return torch.tensor(y) + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### div +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_div_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tmp0 / 2.0 + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = y / 2.0 + return torch.tensor(y) + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### remf +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_remf_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tmp0 % 2.0 + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = y - 2.0 * torch.div(y, 2.0, rounding_mode="trunc") + return torch.tensor(y) + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### negf +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_negf_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = -tmp0 + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = -y + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### cmpf +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_cmpf_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = (tmp0 > 0.5).to(tmp0.dtype) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = (y > 0.5).to(y.dtype) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### ceil +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_ceil_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.math.ceil(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.ceil(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### floor +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_floor_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.math.floor(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.floor(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### maximum(propagate_nan == tl.PropagateNan.ALL) +# setting propagate_nan=tl.PropagateNan.ALL to generate arith::MaximumFOp +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_maximum_nanall_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + tl.static_assert(N > 1) + tmp0 = tl.load(in_ptr0 + 0) + tmp1 = tl.load(in_ptr0 + 1) + tmp1 = tl.maximum(tmp0, tmp1, propagate_nan=tl.PropagateNan.ALL) + tl.store(out_ptr0 + 0, tmp1) + + def torch_func(x0): + y0 = x0[0] + y1 = x0[1] + y = torch.maximum(y0, y1) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### maximum(propagate_nan == tl.PropagateNan.NONE) +# setting propagate_nan=tl.PropagateNan.NONE to generate arith::MaxNumFOp +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_maximum_nannone_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + tl.static_assert(N > 1) + tmp0 = tl.load(in_ptr0 + 0) + tmp1 = tl.load(in_ptr0 + 1) + tmp1 = tl.maximum(tmp0, tmp1, propagate_nan=tl.PropagateNan.ALL) + tl.store(out_ptr0 + 0, tmp1) + + def torch_func(x0): + y0 = x0[0] + y1 = x0[1] + y = torch.fmax(y0, y1) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### minimum(propagate_nan == tl.PropagateNan.ALL) +# setting propagate_nan=tl.PropagateNan.ALL to generate arith::MinimumFOp +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_minimum_nanall_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + tl.static_assert(N > 1) + tmp0 = tl.load(in_ptr0 + 0) + tmp1 = tl.load(in_ptr0 + 1) + tmp1 = tl.minimum(tmp0, tmp1, propagate_nan=tl.PropagateNan.ALL) + tl.store(out_ptr0 + 0, tmp1) + + def torch_func(x0): + y0 = x0[0] + y1 = x0[1] + y = torch.minimum(y0, y1) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### minimum(propagate_nan == tl.PropagateNan.NONE) +# setting propagate_nan=tl.PropagateNan.NONE to generate arith::MinNumFOp +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_minimum_nannone_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + tl.static_assert(N > 1) + tmp0 = tl.load(in_ptr0 + 0) + tmp1 = tl.load(in_ptr0 + 1) + tmp1 = tl.minimum(tmp0, tmp1, propagate_nan=tl.PropagateNan.NONE) + tl.store(out_ptr0 + 0, tmp1) + + def torch_func(x0): + y0 = x0[0] + y1 = x0[1] + y = torch.fmin(y0, y1) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### extf +@pytest.mark.parametrize("param_list", [["float16", "float32", 16]]) +def test_scalar_extf_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tmp0.to(tl.float32) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = y.to(torch.float32) + return y + + src_dtype, dst_dtype, N = param_list + x0 = test_common.generate_tensor((N,), src_dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dst_dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dst_dtype, y_cal[0], y_ref) + + +### truncf +@pytest.mark.parametrize("param_list", [["float32", "float16", 16]]) +def test_scalar_truncf_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tmp0.to(tl.float16) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = y.to(torch.float16) + return y + + src_dtype, dst_dtype, N = param_list + x0 = test_common.generate_tensor((N,), src_dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dst_dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dst_dtype, y_cal[0], y_ref) + + +### exp +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_exp_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.math.exp(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.exp(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### exp2 +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_exp_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.math.exp2(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.exp2(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### log +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_log_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp0 = tl.abs(tmp0) + tmp1 = tl.log(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.abs(y) + y = torch.log(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### log2 +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_log2_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp0 = tl.abs(tmp0) + tmp1 = tl.log2(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.abs(y) + y = torch.log2(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### sin +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_sin_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.sin(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.sin(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### cos +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_cos_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.cos(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.cos(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### abs +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_abs_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.abs(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.abs(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### erf +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_erf_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.erf(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.erf(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### sqrt +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_sqrt_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp0 = tl.abs(tmp0) + tmp1 = tl.math.sqrt(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.abs(y) + y = torch.sqrt(y) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### rsqrt +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_rsqrt_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp0 = tl.abs(tmp0) + tmp1 = tl.math.rsqrt(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.abs(y) + y = torch.rsqrt(y) + return y.clone().detach() + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### tanh +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_tanh_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + idx = 0 + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = libdevice.tanh(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def torch_func(x0): + y = x0[0] + y = torch.tanh(y) + return y.clone().detach() + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) + + +### sum +@pytest.mark.parametrize("param_list", [["float32", 16]]) +def test_scalar_sum_calc(param_list): + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, N: tl.constexpr): + tmp0 = tl.load(in_ptr0 + tl.arange(0, N)) + tmp1 = tl.sum(tmp0, 0) + tl.store(out_ptr0 + 0, tmp1) + + def torch_func(x0): + y = torch.sum(x0, 0) + return y + + dtype, N = param_list + x0 = test_common.generate_tensor((N,), dtype).cpu() + y_ref = torch_func(x0) + y_cal = test_common.generate_tensor((1,), dtype).cpu() + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](y_cal_txda, x0_txda, N=N) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal[0], y_ref) diff --git a/test/wafer/ops/test_sigmoid.py b/test/wafer/ops/test_sigmoid.py new file mode 100644 index 00000000..8c6b7e17 --- /dev/null +++ b/test/wafer/ops/test_sigmoid.py @@ -0,0 +1,71 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_sigmoid(x0, x1): + res = x0 + torch.sigmoid(x1) + return res + + +@triton.jit +def triton_sigmoid( + in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = xindex < xnumel + x0 = xindex + tmp0 = tl.load(in_ptr0 + (x0), xmask) + tmp1 = tl.load(in_ptr1 + (x0), xmask) + tmp2 = tmp0 + tl.sigmoid(tmp1) + tl.store(out_ptr0 + (xindex), tmp2, xmask) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_sigmoid(param_list): + # 生成数据 + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + # torch结果 + y_ref = torch_sigmoid(x0, x1) + # triton结果 + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_sigmoid[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + # 比较结果 + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_silu.py b/test/wafer/ops/test_silu.py new file mode 100644 index 00000000..07c1f6cf --- /dev/null +++ b/test/wafer/ops/test_silu.py @@ -0,0 +1,106 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_unary(x0, dtype): + res = x0 * (1 / (1 + torch.exp(-x0))) + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + ret = x * (1 / (1 + tl.math.exp(-x))) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, "float32"), + # 只支持'fp32', 'fp64' ValueError: Expected dtype ['fp32', 'fp64'] but got fp16 + # (torch.float16, "float16"), + # (torch.bfloat16, 'bfloat16'), + # (torch.int8, 'int8'), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_silu_and_mul.py b/test/wafer/ops/test_silu_and_mul.py new file mode 100644 index 00000000..b842e421 --- /dev/null +++ b/test/wafer/ops/test_silu_and_mul.py @@ -0,0 +1,211 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + +fast_expf = tl.math.exp + + +@triton.jit +def _silu_and_mul_kernel( + gateup_ptr, + out_ptr, + N: tl.constexpr, + stride_gum: tl.constexpr, + stride_gun: tl.constexpr, + stride_om: tl.constexpr, + stride_on: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + """silu and mul kernel.""" + m_id = tl.program_id(0) + + up_ptr = gateup_ptr + N * stride_gun + + offs_n = tl.arange(0, BLOCK_SIZE_N) + gate_ptrs = gateup_ptr + m_id * stride_gum + offs_n * stride_gun + up_ptrs = up_ptr + m_id * stride_gum + offs_n * stride_gun + out_ptrs = out_ptr + m_id * stride_om + offs_n * stride_on + + for _ in range(0, N, BLOCK_SIZE_N): + gate = tl.load(gate_ptrs).to(tl.float32) + up = tl.load(up_ptrs).to(tl.float32) + + gate = gate / (1 + fast_expf(-gate)) + out = gate * up + + tl.store(out_ptrs, out) + + gate_ptrs += BLOCK_SIZE_N * stride_gun + up_ptrs += BLOCK_SIZE_N * stride_gun + out_ptrs += BLOCK_SIZE_N * stride_on + + +@triton.jit +def _silu_and_mul_no_align_kernel( + gateup_ptr, + out_ptr, + N: tl.constexpr, + stride_gum: tl.constexpr, + stride_gun: tl.constexpr, + stride_om: tl.constexpr, + stride_on: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + """silu and mul kernel.""" + m_id = tl.program_id(0) + + up_ptr = gateup_ptr + N * stride_gun + + offs_n = tl.arange(0, BLOCK_SIZE_N) + gate_ptrs = gateup_ptr + m_id * stride_gum + offs_n * stride_gun + up_ptrs = up_ptr + m_id * stride_gum + offs_n * stride_gun + out_ptrs = out_ptr + m_id * stride_om + offs_n * stride_on + + for n in range(0, N, BLOCK_SIZE_N): + mask = n + offs_n < N + gate = tl.load(gate_ptrs, mask=mask, other=0.0).to(tl.float32) + up = tl.load(up_ptrs, mask=mask, other=0.0).to(tl.float32) + + gate = gate / (1 + fast_expf(-gate)) + out = gate * up + + tl.store(out_ptrs, out, mask=mask) + + gate_ptrs += BLOCK_SIZE_N * stride_gun + up_ptrs += BLOCK_SIZE_N * stride_gun + out_ptrs += BLOCK_SIZE_N * stride_on + + +def silu_and_mul(gate_up: torch.Tensor, out: torch.Tensor = None): + """silu and mul.""" + assert gate_up.dim() == 2 + + M = gate_up.size(0) + N = gate_up.size(-1) // 2 + if out is None: + out_shape = (M, N) + out = gate_up.new_empty(out_shape) + + BLOCK_SIZE_N = triton.next_power_of_2(N) + BLOCK_SIZE_N = min(BLOCK_SIZE_N, 1024) + num_warps = 4 + num_stages = 2 + grid = (M,) + if N % BLOCK_SIZE_N == 0: + gate_up_txda = gate_up.to("txda") + out_txda = out.to("txda") + _silu_and_mul_kernel[grid]( + gate_up_txda, + out_txda, + N, + stride_gum=gate_up_txda.stride(0), + stride_gun=gate_up_txda.stride(1), + stride_om=out_txda.stride(0), + stride_on=out_txda.stride(1), + BLOCK_SIZE_N=BLOCK_SIZE_N, + num_warps=num_warps, + num_stages=num_stages, + ) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + else: + gate_up_txda = gate_up.to("txda") + out_txda = out.to("txda") + _silu_and_mul_no_align_kernel[grid]( + gate_up_txda, + out_txda, + N, + stride_gum=gate_up_txda.stride(0), + stride_gun=gate_up_txda.stride(1), + stride_om=out_txda.stride(0), + stride_on=out_txda.stride(1), + BLOCK_SIZE_N=BLOCK_SIZE_N, + num_warps=num_warps, + num_stages=num_stages, + ) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + return out + + +class TestSiluAndMul: + + @pytest.fixture + def seqlen(self): + yield 11 + + @pytest.fixture + def feat_size(self, request): + yield request.param + + @pytest.fixture + def x(self, seqlen, feat_size): + yield torch.rand(seqlen, feat_size, dtype=torch.float16, device="cpu") + + @pytest.fixture + def gt(self, x): + gate, up = x.chunk(2, -1) + gate = torch.nn.functional.silu(gate) + yield gate * up + + @pytest.mark.parametrize("feat_size", [256, 768], indirect=True) + def test_silu_and_mul(self, x, gt): + out = silu_and_mul(x) + torch.testing.assert_close(out, gt) + + +def _gt(x): + gate, up = x.chunk(2, -1) + gate = torch.nn.functional.silu(gate) + return gate * up + + +def _test_silu_and_mul(x): + return silu_and_mul(x) + + +def test(): + seqlen = 11 + feat_size = 128 + x = torch.rand(seqlen, feat_size, dtype=torch.float16, device="cpu") + + gt = _gt(x) + tt = _test_silu_and_mul(x) + + print("max diff", (gt - tt).abs().max()) + + # configs = [] + # configs.append( + # triton.testing.Benchmark( + # x_names=['op'], + # x_vals=['fwd'], + # line_arg='provider', + # line_vals=['triton', 'pytorch'], + # line_names=['Triton', 'PyTorch'], + # ylabel='ms', + # plot_name='', + # args={}, + # )) + + # @triton.testing.perf_report(configs) + # def bench_fn(op, provider, device='cpu'): + # warmup = 100 + # rep = 200 + + # if 'triton' in provider: + # # fn = lambda: test_paged_attention(conti_q, blocked_kv, block_offsets, start_loc, seq_lens, history_lens, feat_dim_v) + # fn = lambda: silu_and_mul(x) + # if 'pytorch' in provider: + # fn = lambda: _gt(x) + + # ms = triton.testing.do_bench(fn, warmup=warmup, rep=rep) + # return ms + + # bench_fn.run(show_plots=True, print_data=True) + + +if __name__ == "__main__": + test() diff --git a/test/wafer/ops/test_sin.py b/test/wafer/ops/test_sin.py new file mode 100644 index 00000000..b8c48d1c --- /dev/null +++ b/test/wafer/ops/test_sin.py @@ -0,0 +1,105 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_unary(x0, dtype): + res = torch.sin(x0) + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + ret = tl.sin(x) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, "float32"), + # (torch.float16, 'float16'), + # (torch.bfloat16, 'bfloat16'), + # (torch.int8, 'int8'), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_softmax.py b/test/wafer/ops/test_softmax.py new file mode 100644 index 00000000..88e27f55 --- /dev/null +++ b/test/wafer/ops/test_softmax.py @@ -0,0 +1,142 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import test_common +import pytest + + +def naive_softmax(x): + # read MN elements ; write M elements + x_max = x.max(dim=1)[0] + # read MN + M elements ; write MN elements + z = x - x_max[:, None] + # read MN elements ; write MN elements + numerator = torch.exp(z) + # read MN elements ; write M elements + denominator = numerator.sum(dim=1) + # read MN + M elements ; write MN elements + ret = numerator / denominator[:, None] + # in total: read 5MN + 2M elements ; wrote 3MN + 2M elements + return ret + + +@triton.jit +def softmax_kernel( + output_ptr, + input_ptr, + input_row_stride, + output_row_stride, + n_rows, + n_cols, + BLOCK_SIZE: tl.constexpr, +): + # starting row of the program + row_start = tl.program_id(0) + row_step = tl.num_programs(0) + for row_idx in tl.range(row_start, n_rows, row_step): + # The stride represents how much we need to increase the pointer to advance 1 row + row_start_ptr = input_ptr + row_idx * input_row_stride + # The block size is the next power of two greater than n_cols, so we can fit each + # row in a single block + col_offsets = tl.arange(0, BLOCK_SIZE) + input_ptrs = row_start_ptr + col_offsets + # Load the row into SRAM, using a mask since BLOCK_SIZE may be > than n_cols + mask = col_offsets < n_cols + row = tl.load(input_ptrs, mask=mask, other=-float("inf")) + # Subtract maximum for numerical stability + row_minus_max = row - tl.max(row, axis=0) + numerator = tl.exp(row_minus_max) + denominator = tl.sum(numerator, axis=0) + softmax_output = numerator / denominator + # Write back output to DRAM + output_row_start_ptr = output_ptr + row_idx * output_row_stride + output_ptrs = output_row_start_ptr + col_offsets + tl.store(output_ptrs, softmax_output, mask=mask) + + +kernels = {} + + +def softmax(x, stream): + n_rows, n_cols = x.shape + + BLOCK_SIZE = triton.next_power_of_2(n_cols) + + y = torch.empty_like(x) + + kernel, num_programs = kernels.get(BLOCK_SIZE, (None, 0)) + if kernel is None: + num_programs = 32 + kernel = softmax_kernel + kernels[BLOCK_SIZE] = (kernel, num_programs) + + num_programs = min(num_programs, n_rows) + + # Create a number of persistent programs. + y_txda = y.to("txda") + x_txda = x.to("txda") + kernel[(num_programs, 1, 1)]( + y_txda, x_txda, x_txda.stride(0), y_txda.stride(0), n_rows, n_cols, BLOCK_SIZE + ) + with torch.no_grad(): + y.copy_(y_txda.cpu()) + return y + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + (torch.bfloat16, "bfloat16"), +] + +shapes = [ + (1823, 781), + (1823, 2), + (1823, 4), + (1823, -32), + (1823, -100), + (1823, -256), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype, sigtype", types) +@pytest.mark.parametrize("M, N", shapes) +def test_softmax(dtype, sigtype, M, N): + torch.txda.set_device(0) + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M] if M in map_for_64_t else M + N = map_for_64_t[N] if N in map_for_64_t else N + + device = torch.txda.current_device() + stream = torch.txda.current_stream(device).txda_stream + torch.manual_seed(0) + x = torch.randn(M, N, dtype=dtype, device="cpu") + y_triton = softmax(x, stream) + y_torch = torch.softmax(x, axis=1) + test_common.validate_cmp(sigtype, y_triton, y_torch) diff --git a/test/wafer/ops/test_split.py b/test/wafer/ops/test_split.py new file mode 100644 index 00000000..7a04a5de --- /dev/null +++ b/test/wafer/ops/test_split.py @@ -0,0 +1,86 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl + +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_( + output_ptr, x_ptr, output_ptr1, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr +): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + + xx, yy = tl.split(X) + + oidx = xidx[:, None] * YB + yidx[None, :] + + tl.store(output_ptr + oidx, xx) + tl.store(output_ptr1 + oidx, yy) + + +@pytest.mark.parametrize( + "para_type,data_type,XB,YB,ZB", + [ + ["float32", torch.float32, 16, 256, 2], + ["float32", torch.float32, 8, 8, 2], + ["float16", torch.float16, 16, 256, 2], + ["float16", torch.float16, 8, 8, 2], + ["int8", torch.int8, 8, 128, 2], + ["int8", torch.int8, 8, 8, 2], + ], +) +def test_split(para_type, data_type, XB, YB, ZB): + + x = torch.randint(low=-128, high=128, size=(XB, YB, ZB), dtype=data_type).cpu() + + a, b = torch.split(x, 1, dim=-1) + a = a.reshape(XB, YB) + b = b.reshape(XB, YB) + print(a) + print(b) + + output = torch.randint(1, (XB, YB), dtype=data_type).cpu() + output1 = torch.randint(1, (XB, YB), dtype=data_type).cpu() + output_txda = output.to("txda") + x_txda = x.to("txda") + output1_txda = output1.to("txda") + fn_npu_[1, 1, 1](output_txda, x_txda, output1_txda, XB, YB, ZB, debug=True) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output1.copy_(output1_txda.cpu()) + + print(output) + print(output1) + + test_common.validate_cmp(para_type, a, output) + test_common.validate_cmp(para_type, b, output1) diff --git a/test/wafer/ops/test_sqrt.py b/test/wafer/ops/test_sqrt.py new file mode 100644 index 00000000..40381f51 --- /dev/null +++ b/test/wafer/ops/test_sqrt.py @@ -0,0 +1,106 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest + +import triton +import triton.language as tl +import test_common + +import torch +import torch_txda # noqa: F401 + + +def standard_unary(x0, dtype): + res = torch.sqrt(x0) + return res + + +def standard_binary(x0, y0, dtype): + res = x0 + y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + ret = tl.sqrt(x) + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary( + in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr +): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x + y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + (torch.float32, "float32"), + # Expected dtype ['fp32', 'fp64'] but got fp16 + # (torch.float16, "float16"), + # (torch.bfloat16, 'bfloat16'), + # (torch.int8, 'int8'), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize("dtype,sigtype", types) +@pytest.mark.parametrize("N,NUMEL", shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype) + + ans = standard_unary(x0, dtype) + x0 = x0.cpu() + print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + out_txda = out.to("txda") + triton_elementwise_unary[1, 1, 1](x0_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_store_scalar.py b/test/wafer/ops/test_store_scalar.py new file mode 100644 index 00000000..06b15bd2 --- /dev/null +++ b/test/wafer/ops/test_store_scalar.py @@ -0,0 +1,50 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +# load with mask, store with scalar +@triton.jit +def sum_kernel_1(inp, mid, M, BLOCK_SIZE: tl.constexpr): + pid = tl.program_id(0) + # 0 / 1 / 2 * 4 + (0,1,2,3) + offset = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + inp_ptrs = inp + offset + mask = offset < M + inp_val = tl.load(inp_ptrs, mask=mask).to(tl.float32) + sum_val = tl.sum(inp_val) + mid_ptr = mid + pid + tl.store(mid_ptr, sum_val) + + +def test_case(): + inp = torch.ones(16, device="cpu", dtype=torch.float32) + mid = torch.empty(4, device="cpu", dtype=torch.float32) + inp_txda = inp.to("txda") + mid_txda = mid.to("txda") + sum_kernel_1[(4, 1, 1)](inp_txda, mid_txda, 16, 4) + with torch.no_grad(): + mid.copy_(mid_txda.cpu()) + ref = torch.tensor([4.0, 4.0, 4.0, 4.0], device="cpu", dtype=torch.float32) + assert torch.allclose(mid, ref, rtol=1e-03, atol=1e-03, equal_nan=True) diff --git a/test/wafer/ops/test_sub.py b/test/wafer/ops/test_sub.py new file mode 100644 index 00000000..9fe16695 --- /dev/null +++ b/test/wafer/ops/test_sub.py @@ -0,0 +1,72 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import numpy as np +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0, x1): + res = x0 - x1 + return res + + +@triton.jit +def triton_sub( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 - tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int32", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_sub[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_sum.py b/test/wafer/ops/test_sum.py new file mode 100644 index 00000000..856f2957 --- /dev/null +++ b/test/wafer/ops/test_sum.py @@ -0,0 +1,160 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import pytest +import test_common +import time + + +@triton.jit +def sum_loop_high( + in_ptr0, + in_ptr1, + in_ptr2, + out_ptr0, + rnumel, + xnumel, + XBLOCK: tl.constexpr, + XBLOCK_SUB: tl.constexpr, + RBLOCK: tl.constexpr, +): + R = rnumel + X = xnumel + xoffset = tl.program_id(0) * XBLOCK + xbase = tl.arange(0, XBLOCK_SUB) + rbase = tl.arange(0, RBLOCK) + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + xbase + x0 = xindex[None, :] + _tmp6 = tl.full([RBLOCK, XBLOCK_SUB], 0, tl.float32) + for roffset in range(0, rnumel, RBLOCK): + rindex = roffset + rbase + rmask = None + r1 = rindex[:, None] + tmp0 = tl.load(in_ptr0 + (X * r1 + (x0)), rmask) + tmp1 = tl.load(in_ptr1 + (X * r1 + (x0)), rmask) + tmp3 = tl.load(in_ptr2 + (X * r1 + (x0)), rmask) + tmp2 = tmp0 + tmp1 + tmp4 = tmp2 + tmp3 + _tmp6 = _tmp6 + tmp4 + tmp6 = tl.sum(_tmp6, 0) + tl.store(out_ptr0 + (xindex), tmp6, None) + + +@triton.jit +def sum_loop_low( + in_ptr0, + in_ptr1, + in_ptr2, + out_ptr0, + xnumel, + ynumel, + XBLOCK: tl.constexpr, + RBLOCK: tl.constexpr, +): + X = xnumel + Y = ynumel + xoffset = tl.program_id(0) * XBLOCK + xindex = xoffset + tl.arange(0, XBLOCK) + + x0 = xindex[:, None] + rbase = tl.arange(0, RBLOCK) + _tmp6 = tl.full([XBLOCK, RBLOCK], 0, tl.float32) + for roffset in range(0, ynumel, RBLOCK): + rindex = roffset + rbase + rmask = None + r1 = rindex[None, :] + tmp0 = tl.load(in_ptr0 + (r1 + (Y * x0)), rmask) + tmp1 = tl.load(in_ptr1 + (r1 + (Y * x0)), rmask) + tmp3 = tl.load(in_ptr2 + (r1 + (Y * x0)), rmask) + tmp2 = tmp0 + tmp1 + tmp4 = tmp2 + tmp3 + _tmp6 = _tmp6 + tmp4 + tmp6 = tl.sum(_tmp6, 1) + + tl.store(out_ptr0 + (xindex), tmp6, None) + + +def foo(a, b, c): + y = a + b + c + y = y.sum(0) + return y + + +def bar(a, b, c): + y = a + b + c + y = y.sum(1) + return y + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (64, 8192), 1, 2, 256, 16], + ], +) +def test_case_1(param_list): + dtype, shape, ncore, XB, YB, ZB = param_list + a = test_common.generate_tensor(shape, dtype).cpu() + b = test_common.generate_tensor(shape, dtype).cpu() + c = test_common.generate_tensor(shape, dtype).cpu() + value = torch.empty_strided((a.shape[0],), (1,)).cpu() + + std_low_ret = bar(a, b, c) + print(f"std_low_ret = {std_low_ret[0:8]}") + XBLOCK = 64 + RBLOCK = 32 + NBLOCKS = a.shape[0] // XBLOCK + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + value_txda = value.to("txda") + sum_loop_low[NBLOCKS, 1, 1](a_txda, b_txda, c_txda, value_txda, a_txda.shape[0], a_txda.shape[1], XBLOCK, RBLOCK) + with torch.no_grad(): + value.copy_(value_txda.cpu()) + triton_low_ret = value + print(f"triton_low_ret = {triton_low_ret[0:8]}") + torch.testing.assert_close(std_low_ret, triton_low_ret, rtol=1e-3, atol=1e-3) + + std_ret2 = foo(a, b, c) + print(f"std_ret2 = {std_ret2[0:8]}") + NBLOCKS = 32 + XBLOCK = a.shape[1] // NBLOCKS + XBLOCK_SUB = min(64, max(XBLOCK // 2, 32)) + RBLOCK = 64 + + value2 = torch.empty_strided((a.shape[1],), (1,)).cpu() + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + value2_txda = value2.to("txda") + sum_loop_high[NBLOCKS, 1, 1]( + a_txda, b_txda, c_txda, value2_txda, a_txda.shape[0], a_txda.shape[1], XBLOCK, XBLOCK_SUB, RBLOCK + ) + with torch.no_grad(): + value2.copy_(value2_txda.cpu()) + triton_ret2 = value2 + print(f"triton_ret2 = {triton_ret2[0:8]}") + torch.testing.assert_close(std_ret2, triton_ret2, rtol=1e-3, atol=1e-3) diff --git a/test/wafer/ops/test_sum_dim0.py b/test/wafer/ops/test_sum_dim0.py new file mode 100644 index 00000000..5ccef822 --- /dev/null +++ b/test/wafer/ops/test_sum_dim0.py @@ -0,0 +1,150 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def standard_sum(x0, dim, dtype): + res = torch.sum(x0, dim, dtype=dtype) + return res + + +@triton.jit +def triton_sum_dim0( + in_ptr0, + out_ptr0, + M: tl.constexpr, + N: tl.constexpr, + MNUMEL: tl.constexpr, + NNUMEL: tl.constexpr, +): + mblk_idx = tl.arange(0, MNUMEL) + nblk_idx = tl.arange(0, NNUMEL) + + mmask = mblk_idx < M + nmask = nblk_idx < N + + mask = (mmask[:, None]) & (nmask[None, :]) + + idx = mblk_idx[:, None] * N + nblk_idx[None, :] + + x = tl.load(in_ptr0 + idx, mask=mask, other=0) + + ret = tl.sum(x, 0) + + tl.store(out_ptr0 + nblk_idx, ret, mask=nmask) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16,'bfloat16'), TODO: waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), TODO: waiting for supporting or testing + # (torch.int32,'int32'), TODO: waiting for supporting or testing + # (torch.int64,'int64'), TODO: waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (57, 3, 64, 16), + (57, -32, 64, 32), + (57, 37, 64, 64), + (57, -256, 64, 256), + (57, 263, 64, 512), + (64, 3, 64, 16), + (64, -32, 64, 32), + (64, 37, 64, 64), + (64, -256, 64, 256), + (64, 263, 64, 512), + (3, 3, 8, 8), + (-32, 3, 32, 8), + (37, 3, 64, 8), + (-256, 3, 256, 8), + (263, 3, 512, 8), + (3, 1, 8, 8), + (-32, 1, 32, 8), + (37, 1, 64, 8), + (-256, 1, 256, 8), + (263, 1, 512, 8), +] + +map_for_64_t = {37: (31, 32), 263: (107, 128)} +map_for_32_t = {263: (137, 256)} + + +# @pytest.mark.parametrize('dtype, sigtype',[(torch.float32,'float32'),]) +@pytest.mark.parametrize( + "M, N, MNUMEL, NNUMEL", + [ + (57, 3, 64, 16), + (64, -32, 64, 32), + (37, 3, 64, 8), + (263, 1, 512, 8), + (-256, 3, 256, 8), + ], +) +@pytest.mark.parametrize("dtype, sigtype", types) +# @pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_sum_dim0(dtype, sigtype, M, N, MNUMEL, NNUMEL): + + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == "float32" or sigtype == "bfloat16" or sigtype == "int32": + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(M, N), dtype=sigtype) + + ans = standard_sum(x0, 0, dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_sum_dim0[1, 1, 1]( + x0_txda, output_txda, M=M, N=N, MNUMEL=MNUMEL, NNUMEL=NNUMEL, debug=True + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_sum_dim1.py b/test/wafer/ops/test_sum_dim1.py new file mode 100644 index 00000000..9f7f4dc5 --- /dev/null +++ b/test/wafer/ops/test_sum_dim1.py @@ -0,0 +1,150 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl +import time + +import torch +import torch_txda # noqa: F401 +import test_common + + +def standard_sum(x0, dim, dtype): + res = torch.sum(x0, dim, dtype=dtype) + return res + + +@triton.jit +def triton_sum_dim1( + in_ptr0, + out_ptr0, + M: tl.constexpr, + N: tl.constexpr, + MNUMEL: tl.constexpr, + NNUMEL: tl.constexpr, +): + mblk_idx = tl.arange(0, MNUMEL) + nblk_idx = tl.arange(0, NNUMEL) + + mmask = mblk_idx < M + nmask = nblk_idx < N + + mask = (mmask[:, None]) & (nmask[None, :]) + + idx = mblk_idx[:, None] * N + nblk_idx[None, :] + + x = tl.load(in_ptr0 + idx, mask=mask, other=0) + + ret = tl.sum(x, 1) + + tl.store(out_ptr0 + mblk_idx, ret, mask=mmask) + + +types = [ + (torch.float32, "float32"), + (torch.float16, "float16"), + # (torch.bfloat16,'bfloat16'), waiting for supporting or testing + (torch.int8, "int8"), + # (torch.int16,'int16'), waiting for supporting or testing + # (torch.int32,'int32'), waiting for supporting or testing + # (torch.int64,'int64'), waiting for supporting or testing +] + +# if shape axis = 32/256 , then actual shape = axis/element_size() +shapes = [ + (57, 3, 64, 16), + (57, -32, 64, 32), + (57, 37, 64, 64), + (57, -256, 64, 256), + (57, 263, 64, 512), + (64, 3, 64, 16), + (64, -32, 64, 32), + (64, 37, 64, 64), + (64, -256, 64, 256), + (64, 263, 64, 512), + (3, 3, 8, 8), + (-32, 3, 32, 8), + (37, 3, 64, 8), + (-256, 3, 256, 8), + (263, 3, 512, 8), + (3, 1, 8, 8), + (-32, 1, 32, 8), + (37, 1, 64, 8), + (-256, 1, 256, 8), + (263, 1, 512, 8), +] + +map_for_64_t = {37: (31, 32), 263: (107, 128)} +map_for_32_t = {263: (137, 256)} + + +@pytest.mark.parametrize( + "M, N, MNUMEL, NNUMEL", + [ + (57, 3, 64, 16), + (64, -32, 64, 32), + (37, 3, 64, 8), + (263, 1, 512, 8), + (-256, 3, 256, 8), + ], +) +# @pytest.mark.parametrize('M, N',[(263,3),(-256,3)]) +@pytest.mark.parametrize("dtype, sigtype", types) +# @pytest.mark.parametrize('M, N, MNUMEL, NNUMEL',shapes) +def test_sum_dim1(dtype, sigtype, M, N, MNUMEL, NNUMEL): + + M = (-M) // torch.tensor(0, dtype=dtype).element_size() if M < 0 else M + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + M = map_for_64_t[M][0] if M in map_for_64_t else M + MNUMEL = map_for_64_t[M][1] if M in map_for_64_t else MNUMEL + N = map_for_64_t[N][0] if N in map_for_64_t else N + NNUMEL = map_for_64_t[N][1] if N in map_for_64_t else NNUMEL + + elif sigtype == "float32" or sigtype == "bfloat16" or sigtype == "int32": + M = map_for_32_t[M][0] if M in map_for_32_t else M + MNUMEL = map_for_32_t[M][1] if M in map_for_32_t else MNUMEL + N = map_for_32_t[N][0] if N in map_for_32_t else N + NNUMEL = map_for_32_t[N][1] if N in map_for_32_t else NNUMEL + + print(f"sum : ({M}, {N}) {dtype} {sigtype}") + x0 = test_common.generate_tensor(shape=(M, N), dtype=sigtype) + + ans = standard_sum(x0, 1, dtype) + + x0 = x0.cpu() + print(ans) + + output = torch.zeros((M,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + output_txda = output.to("txda") + triton_sum_dim1[1, 1, 1]( + x0_txda, output_txda, M=M, N=N, MNUMEL=MNUMEL, NNUMEL=NNUMEL, debug=True + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + + test_common.validate_cmp(sigtype, output, ans) diff --git a/test/wafer/ops/test_sum_vector.py b/test/wafer/ops/test_sum_vector.py new file mode 100644 index 00000000..738a0363 --- /dev/null +++ b/test/wafer/ops/test_sum_vector.py @@ -0,0 +1,92 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + +# from dlblas.utils.libentry import libentry +import pytest +from test_common import generate_tensor, validate_cmp, _32bit_dtypes, _16bit_dtypes + + +def torch_func(x0): + return torch.sum(x0) + + +@pytest.mark.parametrize("dtype", _32bit_dtypes) +@pytest.mark.parametrize("shape", [(1,), (4,), (8,), (32,), (64,), (1024,)]) +def test_sum(dtype, shape): + + # @libentry() + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, XBLOCK: tl.constexpr): + idx = tl.arange(0, XBLOCK) + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.sum(tmp0) + tl.store(out_ptr0 + idx, tmp1) + + def triton_func(x0): + out = x0[0] + out_txda = out.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](out_txda, x0_txda, x0_txda.numel()) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + return out + + x0 = generate_tensor(shape=shape, dtype=dtype).cpu() + torch_ref = torch_func(x0) + triton_cal = triton_func(x0) + validate_cmp(dtype, torch_ref, triton_cal) + + +@triton.jit +def _reduce_combine(a, b): + return a + b + + +@pytest.mark.parametrize("dtype", _32bit_dtypes) +@pytest.mark.parametrize("shape", [(1,), (4,), (8,), (32,), (64,), (1024,)]) +def test_reduce_sum(dtype, shape): + + # @libentry() + @triton.jit + def triton_kernel(out_ptr0, in_ptr0, XBLOCK: tl.constexpr): + idx = tl.arange(0, XBLOCK) + tmp0 = tl.load(in_ptr0 + idx) + tmp1 = tl.reduce(tmp0, 0, _reduce_combine) + tl.store(out_ptr0 + idx, tmp1) + + def triton_func(x0): + out = x0[0] + out_txda = out.to("txda") + x0_txda = x0.to("txda") + triton_kernel[1, 1, 1](out_txda, x0_txda, x0_txda.numel()) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + return out + + x0 = generate_tensor(shape=shape, dtype=dtype).cpu() + torch_ref = torch_func(x0) + triton_cal = triton_func(x0) + validate_cmp(dtype, torch_ref, triton_cal) diff --git a/test/wafer/ops/test_swap.py b/test/wafer/ops/test_swap.py new file mode 100644 index 00000000..23318b84 --- /dev/null +++ b/test/wafer/ops/test_swap.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import triton +import triton.language as tl +import torch_txda # noqa: F401 +import pytest + + +@triton.jit +def swap_kernel( + x_ptr, # *Pointer* to first inout vector. + y_ptr, # *Pointer* to second inout vector. + BLOCK_SIZE: tl.constexpr, # Number of elements each program should process. + # NOTE: `constexpr` so it can be used as a shape value. +): + pid = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0. + block_start = pid * BLOCK_SIZE + offsets = block_start + tl.arange(0, BLOCK_SIZE) + x = tl.load(x_ptr + offsets) + y = tl.load(y_ptr + offsets) + tl.store(x_ptr + offsets, y) + tl.store(y_ptr + offsets, x) + + +def swap(x: torch.Tensor, y: torch.Tensor, size): + n_elements = x.numel() + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) + x_txda = x.to("txda") + y_txda = y.to("txda") + swap_kernel[grid](x_txda, y_txda, BLOCK_SIZE=size) + with torch.no_grad(): + x.copy_(x_txda.cpu()) + y.copy_(y_txda.cpu()) + + +@pytest.mark.parametrize( + "shape", [(1,), (2,), (4,), (8,), (16,), (128,), (512,), (1024,)] +) +def test(shape): + x = torch.rand(shape).cpu() + y = torch.rand(shape).cpu() + assert not torch.equal(x, y) + x_ = x.clone() + y_ = y.clone() + swap(x, y, shape[0]) + torch.testing.assert_close(x, y_, rtol=1e-04, atol=1e-04, equal_nan=True) + torch.testing.assert_close(y, x_, rtol=1e-04, atol=1e-04, equal_nan=True) diff --git a/test/wafer/ops/test_swiglu.py b/test/wafer/ops/test_swiglu.py new file mode 100644 index 00000000..68075acb --- /dev/null +++ b/test/wafer/ops/test_swiglu.py @@ -0,0 +1,109 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +# from .utils import calculate_settings +def standard_binary(e, g): + ee = e.to(torch.float32) + f = ee * torch.sigmoid(ee) + h = (f * g).to(g.dtype) + return h + + +@triton.jit +def _fg_kernel( + e, + g, + h, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + block_idx = tl.program_id(0) + offsets = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_elements + + e_row = tl.load(e + offsets, mask=mask, other=0).to(tl.float32) + g_row = tl.load(g + offsets, mask=mask, other=0) # .to(tl.float32) + + # f = e * sigmoid(e) + f_row = e_row * tl.sigmoid(e_row) # e_row / (1 + tl.exp(-e_row)) + # f_row = f_row.to(g_row.dtype) # bf16 should always cast to fp32 when calculating + # h = f * g + h_row = (f_row * g_row).to(g_row.dtype) + # Store h + tl.store(h + offsets, h_row, mask=mask) + + +pass + + +def swiglu_fg_kernel(e, g): + batch, seq_len, hd = e.shape + n_elements = e.numel() + h = torch.empty((batch, seq_len, hd), dtype=e.dtype, device="cpu") + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) + e_txda = e.to("txda") + g_txda = g.to("txda") + h_txda = h.to("txda") + kk = _fg_kernel[grid]( + e_txda, + g_txda, + h_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + h.copy_(h_txda.cpu()) + print(kk.asm["ttir"]) + return h + + +pass + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 128, 128)], + ["float16", (2, 128, 128)], + ["bfloat16", (2, 128, 128)], + ], +) +def test_case(param_list): + dtype, size = param_list + torch.manual_seed(0) + x = torch.rand(size, device="cpu", dtype=eval("torch." + dtype)) + y = torch.rand(size, device="cpu", dtype=eval("torch." + dtype)) + std_ret = standard_binary(x, y) + print(f"std_ret= {std_ret}") + ret = swiglu_fg_kernel(x, y) + print(f"ret= {ret}") + test_common.validate_cmp(dtype, std_ret, ret) + + +pass diff --git a/test/wafer/ops/test_swizzle2d.py b/test/wafer/ops/test_swizzle2d.py new file mode 100644 index 00000000..89feb4af --- /dev/null +++ b/test/wafer/ops/test_swizzle2d.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + i = tl.arange(0, 4) + j = tl.arange(0, 4) + xx, yy = tl.swizzle2d(i, j, size_i=4, size_j=4, size_g=2) + + tl.store(output_ptr + tl.arange(0, 4), xx) + tl.store(x_ptr + tl.arange(0, 4), yy) + + +@pytest.mark.parametrize( + "param_list", + [ + ["int32", (2, 256, 16), 1, 2, 256, 16], + ], +) +def test_case(param_list): + dtype, shape, ncore, XB, YB, ZB = param_list + x = test_common.generate_tensor((4,), dtype).cpu() + a = torch.tensor( + [[0, 2, 1, 3], [4, 6, 5, 7], [8, 10, 9, 11], [12, 14, 13, 15]], + dtype=eval("torch." + dtype), + ).cpu() + output = torch.randint(1, (4,), dtype=eval("torch." + dtype)).cpu() + output_txda = output.to("txda") + x_txda = x.to("txda") + fn_npu_[ncore, 1, 1](output_txda, x_txda, XB, YB, ZB) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + x.copy_(x_txda.cpu()) + print(f"output={output}") + triton_ret = output[:, None] * 4 + x[None, :] + print(f"triton_ret={triton_ret}") + torch.testing.assert_close(triton_ret, a) diff --git a/test/wafer/ops/test_template.py b/test/wafer/ops/test_template.py new file mode 100644 index 00000000..8d389833 --- /dev/null +++ b/test/wafer/ops/test_template.py @@ -0,0 +1,71 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import pytest + +import triton +import triton.language as tl + +import time +import torch +import torch_txda # noqa: F401 +import test_common + +NBLOCKS = 1 +X_SIZE = tl.constexpr(4) +Y_SIZE = tl.constexpr(64) +Z_SIZE = tl.constexpr(32) +NUMEL = tl.constexpr(X_SIZE.value * Y_SIZE.value * Z_SIZE.value) + + +def fn(input): + output = ( + input.reshape((X_SIZE, Y_SIZE, Z_SIZE)) + .permute((1, 0, 2)) + .reshape((X_SIZE * Y_SIZE * Z_SIZE)) + ) + return output + + +@triton.jit +def fn_kernel(output_ptr, input_ptr): + col_offsets = tl.arange(0, X_SIZE * Y_SIZE * Z_SIZE) + input_local = tl.load(input_ptr + col_offsets) + input_local = ( + input_local.reshape((X_SIZE, Y_SIZE, Z_SIZE)) + .permute((1, 0, 2)) + .reshape((X_SIZE * Y_SIZE * Z_SIZE)) + ) + tl.store(output_ptr + col_offsets, input_local) + + +def test_cases(): + input = torch.randn(NUMEL, dtype=torch.float16).cpu() + output = torch.randn(NUMEL, dtype=torch.float16).cpu() + output2 = torch.randn(NUMEL, dtype=torch.float16).cpu() + output_txda = output.to("txda") + input_txda = input.to("txda") + fn_kernel[1, 1, 1](output_txda, input_txda) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output2 = fn(input) + test_common.validate_cmp("float16", output, output2) + print("data validation passed") diff --git a/test/wafer/ops/test_tensor_get_item.py b/test/wafer/ops/test_tensor_get_item.py new file mode 100644 index 00000000..509ece81 --- /dev/null +++ b/test/wafer/ops/test_tensor_get_item.py @@ -0,0 +1,57 @@ +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +# Enable Wafer bounded tensor slicing backed by TLE DSA operations. +import triton.language.extra.wafer.slicing # noqa: F401 + + +@triton.jit +def triton_kernel( + x_ptr, + y_ptr, + output_ptr, + POS: tl.constexpr, + N: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + pid = tl.program_id(axis=0) + start = pid * N + offsets = tl.arange(0, BLOCK_SIZE_N) + mask = offsets < N + x = tl.load(x_ptr + start + offsets, mask=mask) + y = tl.load(y_ptr + start + offsets, mask=mask) + out_left = x[:POS] + y[:POS] + out_right = x[POS:] - y[POS:] + out_left_offsets = tl.arange(0, POS) + tl.store(output_ptr + start + out_left_offsets, out_left) + out_right_offsets = POS + out_left_offsets + tl.store( + output_ptr + start + out_right_offsets, out_right, mask=out_right_offsets < N + ) + + +def triton_func(x: torch.Tensor, y: torch.Tensor, pos: int): + output = torch.empty_like(x) + M = x.size(0) + N = x.size(1) + BLOCK_SIZE_N = triton.next_power_of_2(N) + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + triton_kernel[(M,)](x_txda, y_txda, output_txda, POS=pos, N=N, BLOCK_SIZE_N=BLOCK_SIZE_N) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + return output + + +def test_tensor_get_item(): + size = (32, 32) + mid_pos = size[1] // 2 + x = torch.rand(size, device="cpu") + y = torch.rand(size, device="cpu") + torch_add_ref = x + y + torch_sub_ref = x - y + triton_cal = triton_func(x, y, mid_pos) + torch.testing.assert_close(triton_cal[:, :mid_pos], torch_add_ref[:, :mid_pos]) + torch.testing.assert_close(triton_cal[:, mid_pos:], torch_sub_ref[:, mid_pos:]) diff --git a/test/wafer/ops/test_trans_3d.py b/test/wafer/ops/test_trans_3d.py new file mode 100644 index 00000000..ba303e3f --- /dev/null +++ b/test/wafer/ops/test_trans_3d.py @@ -0,0 +1,94 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import logging +import math + +import pytest +import test_common +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + +@triton.jit +def fn_npu_102(output_ptr, x_ptr, YB: tl.constexpr, ZB: tl.constexpr, KB: tl.constexpr): + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + kidx = tl.arange(0, KB) + idx = yidx[:, None, None] * ZB * KB + zidx[None, :, None] * KB + kidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.trans(X, 1, 0, 2) + + oidx = ( + zidx[:, None, None] * YB * KB + yidx[None, :, None] * KB + kidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@triton.jit +def fn_npu_021(output_ptr, x_ptr, YB: tl.constexpr, ZB: tl.constexpr, KB: tl.constexpr): + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + kidx = tl.arange(0, KB) + idx = yidx[:, None, None] * ZB * KB + zidx[None, :, None] * KB + kidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.trans(X, 0, 2, 1) + + oidx = ( + yidx[:, None, None] * ZB * KB + kidx[None, :, None] * ZB + zidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@pytest.mark.parametrize("shape", [(16, 8, 32)]) +@pytest.mark.parametrize("dtype", ["float32"]) +def test_permute_3d(shape, dtype): + logging.debug(f"dtype:{dtype} shape:{shape}") + + data_type = eval("torch." + dtype) + x = torch.randint(low=0, high=2, size=shape, dtype=data_type).cpu() + + triton_res = torch.empty((shape[1], shape[0], shape[2]), dtype=data_type).cpu() + torch_res = torch.permute(x, (1, 0, 2)) + triton_res_txda = triton_res.to("txda") + x_txda = x.to("txda") + fn_npu_102[1, 1, 1](triton_res_txda, x_txda, shape[0], shape[1], shape[2]) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + test_common.validate_cmp(dtype, triton_res, torch_res) + + triton_res = torch.empty((shape[0], shape[2], shape[1]), dtype=data_type).cpu() + torch_res = torch.permute(x, (0, 2, 1)) + triton_res_txda = triton_res.to("txda") + x_txda = x.to("txda") + fn_npu_021[1, 1, 1](triton_res_txda, x_txda, shape[0], shape[1], shape[2]) + with torch.no_grad(): + triton_res.copy_(triton_res_txda.cpu()) + test_common.validate_cmp(dtype, triton_res, torch_res) diff --git a/test/wafer/ops/test_triton_eq.py b/test/wafer/ops/test_triton_eq.py new file mode 100644 index 00000000..40b547d2 --- /dev/null +++ b/test/wafer/ops/test_triton_eq.py @@ -0,0 +1,70 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0, x1): + res = x0 == x1 + return res + + +@triton.jit +def triton_test( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 == tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype=eval("torch." + "bool")).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_test[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_triton_le.py b/test/wafer/ops/test_triton_le.py new file mode 100644 index 00000000..dcbb0678 --- /dev/null +++ b/test/wafer/ops/test_triton_le.py @@ -0,0 +1,70 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0, x1): + res = x0 <= x1 + return res + + +@triton.jit +def triton_le( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 <= tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype=eval("torch." + "bool")).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_le[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_triton_lt.py b/test/wafer/ops/test_triton_lt.py new file mode 100644 index 00000000..a7cee14d --- /dev/null +++ b/test/wafer/ops/test_triton_lt.py @@ -0,0 +1,70 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0, x1): + res = x0 < x1 + return res + + +@triton.jit +def triton_lt( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 < tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype=eval("torch." + "bool")).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_lt[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_triton_neq.py b/test/wafer/ops/test_triton_neq.py new file mode 100644 index 00000000..2b8180a5 --- /dev/null +++ b/test/wafer/ops/test_triton_neq.py @@ -0,0 +1,70 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_pointwise(x0, x1): + res = x0 != x1 + return res + + +@triton.jit +def triton_neq( + in_ptr0, in_ptr1, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + for loop1 in range(loops1): + x0_prime = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1 + tmp0 = tl.load(in_ptr0 + (x0), None) + tmp1 = tl.load(in_ptr1 + (x0), None) + tmp2 = tmp0 != tmp1 + tl.store(out_ptr0 + (x0), tmp2, None) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 4096, 8), 2, 32768, 1024], + ["float16", (2, 4096, 8), 2, 32768, 1024], + ["int8", (2, 4096, 8), 2, 32768, 1024], + ], +) +def test_case(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_pointwise(x0, x1) + y_cal = torch.zeros(shape, dtype=eval("torch." + "bool")).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_neq[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_umulhi.py b/test/wafer/ops/test_umulhi.py new file mode 100644 index 00000000..e49e4e25 --- /dev/null +++ b/test/wafer/ops/test_umulhi.py @@ -0,0 +1,67 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import numpy as np +import triton +import triton.language as tl +from numpy.random import RandomState + + +# inp the two 32 bit signed integers. +@triton.jit +def umulhi_kernel(X, Y, Z, N: tl.constexpr): + offs = tl.arange(0, N) + x = tl.load(X + offs) + y = tl.load(Y + offs) + z = tl.umulhi(x, y) + tl.store(Z + tl.arange(0, N), z) + + +# accuracy reference +def umulhi32(a, b): + a_64 = a.astype(np.int64) + b_64 = b.astype(np.int64) + product_64 = a_64 * b_64 + # get the high part + result_high_32 = product_64 >> 32 + return result_high_32.astype(np.int32) + + +def test_umulhi(): + N = 128 + x = torch.randint(low=0, high=2000, size=(N,), dtype=torch.int32) + y = torch.randint(low=0, high=2000, size=(N,), dtype=torch.int32) + xx = x.cpu() + yy = y.cpu() + z_tri = torch.zeros(size=(N,), dtype=torch.int32).cpu() + xx_txda = xx.to("txda") + yy_txda = yy.to("txda") + z_tri_txda = z_tri.to("txda") + umulhi_kernel[(1,)](xx_txda, yy_txda, z_tri_txda, N=N) + with torch.no_grad(): + z_tri.copy_(z_tri_txda.cpu()) + + xxx = x.numpy() + yyy = y.numpy() + z_ref = umulhi32(xxx, yyy) + z_ref1 = torch.from_numpy(z_ref).cpu() + assert torch.equal(z_tri, z_ref1) diff --git a/test/wafer/ops/test_unlign_sum.py b/test/wafer/ops/test_unlign_sum.py new file mode 100644 index 00000000..de99b18d --- /dev/null +++ b/test/wafer/ops/test_unlign_sum.py @@ -0,0 +1,82 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import triton +import triton.language as tl + +# [128,65] -> [128,128] +# [128,128] -> [128,65] mask 作用 +# [128,65] -> [128,64] [128,1] mask 作用 + + +@triton.jit +def triton_unlign( + in_ptr0, + out_ptr0, + x0_numel, + r1_numel, + XBLOCK: tl.constexpr, + XBLOCK_SUB: tl.constexpr, + RBLOCK: tl.constexpr, +): + offset = tl.program_id(0) * XBLOCK + base1 = tl.arange(0, XBLOCK_SUB) + loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB + base2 = tl.arange(0, RBLOCK) + loops2: tl.constexpr = (r1_numel + RBLOCK - 1) // RBLOCK + for loop1 in range(loops1): + x = offset + (loop1 * XBLOCK_SUB) + base1 + x0 = offset + (loop1 * XBLOCK_SUB) + base1[:, None] + xmask = x0 < x0_numel + _tmp2 = tl.full([XBLOCK_SUB, RBLOCK], 0, tl.float32) + for loop2 in range(loops2): + r1_prime = loop2 * RBLOCK + base2[:, None] + r1 = loop2 * RBLOCK + base2[None, :] + rmask = r1 < r1_numel + tmp0 = tl.load( + in_ptr0 + (r1 + (65 * x0)), + rmask & xmask, + eviction_policy="evict_first", + other=0.0, + ) + tmp1 = tl.reshape(tmp0, [XBLOCK_SUB, RBLOCK]) + tmp3 = _tmp2 + tmp1 + _tmp2 = tmp3 + tmp2 = tl.sum(_tmp2, 1).reshape(XBLOCK_SUB, 1) + tl.store(out_ptr0 + (x0), tmp2, xmask) + + +def test_cases(): + size = (128, 65) + b = weights = torch.randn((size), dtype=torch.float32).cpu() + c = torch.sum(b, dim=1) + ret = ( + torch.randn((size[0]), device="cpu", dtype=torch.float32).cpu().reshape(size[0]) + ) + b_txda = b.to("txda") + ret_txda = ret.to("txda") + triton_unlign[1, 1, 1](b_txda, ret_txda, size[0], size[1], size[0], size[0], 32) + with torch.no_grad(): + ret.copy_(ret_txda.cpu()) + assert torch.allclose(c, ret, rtol=1e-03, atol=1e-03, equal_nan=True) diff --git a/test/wafer/ops/test_unused_func_arg.py b/test/wafer/ops/test_unused_func_arg.py new file mode 100644 index 00000000..991a5832 --- /dev/null +++ b/test/wafer/ops/test_unused_func_arg.py @@ -0,0 +1,95 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl +import test_common +import torch +import torch_txda # noqa: F401 +import pytest +import math + + +def expand_to_next_power_of_two(a): + if a <= 0: + raise ValueError("must >0") + if (math.log2(a)).is_integer(): + return a + return 2 ** math.ceil(math.log2(a)) + + +@triton.jit +def triton_unused_func_arg_kernel( + output_ptr, + x_ptr, + X: tl.constexpr, + Y: tl.constexpr, + Z: tl.constexpr, + XNUMEL: tl.constexpr, + YNUMEL: tl.constexpr, + ZNUMEL: tl.constexpr, +): + xidx = tl.arange(0, XNUMEL) + yidx = tl.arange(0, YNUMEL) + zidx = tl.arange(0, ZNUMEL) + Xmask = xidx < X + Ymask = yidx < Y + Zmask = zidx < Z + oidx = xidx[:, None, None] * Y * Z + yidx[None, :, None] * Z + zidx[None, None, :] + mask = (Xmask[:, None, None]) & (Ymask[None, :, None]) & (Zmask[None, None, :]) + abc = tl.load(x_ptr + oidx, mask=mask) + ret = tl.zeros_like(abc) + tl.store(output_ptr + oidx, ret, mask=mask) + + +testlist = [ + (triton_unused_func_arg_kernel, "int8", torch.int8, 2, 255, 9), + (triton_unused_func_arg_kernel, "int16", torch.int16, 3, 5, 3), + (triton_unused_func_arg_kernel, "int32", torch.int32, 2, 255, 9), + (triton_unused_func_arg_kernel, "int64", torch.int64, 2, 5, 3), + (triton_unused_func_arg_kernel, "float16", torch.float16, 55, 5, 16), + (triton_unused_func_arg_kernel, "float16", torch.float16, 4, 5, 17), + (triton_unused_func_arg_kernel, "float16", torch.float16, 6, 5, 15), + (triton_unused_func_arg_kernel, "float16", torch.float16, 2, 1928, 3), + (triton_unused_func_arg_kernel, "float32", torch.float32, 2, 255, 9), + (triton_unused_func_arg_kernel, "bfloat16", torch.bfloat16, 3, 5, 3), + (triton_unused_func_arg_kernel, "bool", torch.bool, 3, 5, 3), +] + + +@pytest.mark.parametrize("testfunc, sigtype, dtype, X, Y, Z", testlist) +def test_npu(testfunc, sigtype, dtype, X, Y, Z): + XNUMEL = expand_to_next_power_of_two(X) + YNUMEL = expand_to_next_power_of_two(Y) + ZNUMEL = expand_to_next_power_of_two(Z) + x = torch.full((X, Y, Z), 10, dtype=dtype).cpu() + y = torch.full((X, Y, Z), 0, dtype=dtype).cpu() + output = torch.full((X, Y, Z), 5, dtype=dtype).cpu() + output_txda = output.to("txda") + x_txda = x.to("txda") + testfunc[1, 1, 1](output_txda, x_txda, X, Y, Z, XNUMEL, YNUMEL, ZNUMEL) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + test_common.validate_cmp(sigtype, output, y) + + +if __name__ == "__main__": + test_npu(triton_unused_func_arg_kernel, "bool", torch.bool, 3, 5, 3) diff --git a/test/wafer/ops/test_view.py b/test/wafer/ops/test_view.py new file mode 100644 index 00000000..06ff496a --- /dev/null +++ b/test/wafer/ops/test_view.py @@ -0,0 +1,70 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.reshape(X, (ZB, XB * YB)) + + oidx = tl.arange(0, ZB)[:, None] * XB * YB + tl.arange(0, XB * YB)[None, :] + + tl.store(output_ptr + oidx, ret) + + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 256, 16), 1, 2, 256, 16], + ['float32', (8, 8, 4), 1, 8, 8, 4], + ['float16', (2, 256, 16), 1, 2, 256, 16], + ['float16', (8, 8, 4), 1, 8, 8, 4], + ['int8', (2, 256, 16), 1, 2, 256, 16], + ['int8', (8, 8, 4), 1, 8, 8, 4], + ] + ) +def test_case(param_list): + dtype, shape, ncore, XB, YB, ZB = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = x0.view(ZB, XB * YB).cpu() + print(f"y_ref = {y_ref[0, 0:4]}") + y_cal = torch.empty((ZB, XB * YB), dtype=eval('torch.' + dtype)).cpu() + + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + fn_npu_[ncore, 1, 1](y_cal_txda, x0_txda, XB, YB, ZB) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + print(f"y_cal = {y_cal[0, 0:4]}") + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_where_lt.py b/test/wafer/ops/test_where_lt.py new file mode 100644 index 00000000..de41bf79 --- /dev/null +++ b/test/wafer/ops/test_where_lt.py @@ -0,0 +1,66 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +def torch_where_lt_case1(x0, x1): + res = torch.where(x0 < x1, x0, 1) + return res + +@triton.jit +def triton_where_lt_case1(in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = xindex < xnumel + x0 = xindex + tmp0 = tl.load(in_ptr0 + (x0), xmask) + tmp1 = tl.load(in_ptr1 + (x0), xmask) + tmp2 = tmp0 < tmp1 + tmp3 = tl.where(tmp2, tmp0, 1) + tl.store(out_ptr0 + (xindex), tmp3, xmask) + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 1024, 8), 2, 8192, 1024], + ['float16', (2, 1024, 8), 2, 8192, 1024], + ['int8', (2, 1024, 8), 2, 8192, 1024], + ] + ) + +def test_where_lt_case1(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_where_lt_case1(x0, x1) + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_where_lt_case1[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) \ No newline at end of file diff --git a/test/wafer/ops/test_where_mask.py b/test/wafer/ops/test_where_mask.py new file mode 100644 index 00000000..0d3404a4 --- /dev/null +++ b/test/wafer/ops/test_where_mask.py @@ -0,0 +1,64 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + +def torch_where_lt_case2(x0, x1): + res = torch.where(x0 < x1, x0, x1) + return res + +@triton.jit +def triton_where_lt_case2(in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = xindex < xnumel + x0 = xindex + tmp0 = tl.load(in_ptr0 + (x0), xmask) + tmp1 = tl.load(in_ptr1 + (x0), xmask) + tmp2 = tmp0 < tmp1 + tmp3 = tl.where(tmp2, tmp0, tmp1) + tl.store(out_ptr0 + (xindex), tmp3, xmask) + +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 1024, 8), 2, 8192, 1024], + ['float16', (2, 1024, 8), 2, 8192, 1024], + ['int8', (2, 1024, 8), 2, 8192, 1024], + ] + ) +def test_where_lt_case2(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_where_lt_case2(x0, x1) + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_where_lt_case2[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_where_var.py b/test/wafer/ops/test_where_var.py new file mode 100644 index 00000000..04dc1738 --- /dev/null +++ b/test/wafer/ops/test_where_var.py @@ -0,0 +1,65 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + +def torch_where_lt_case2(x0, x1): + res = torch.where(x0 < x1, x0, x1) + return res + +@triton.jit +def triton_where_lt_case2(in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr): + xoffset = tl.program_id(0) * XBLOCK + for xoffset_sub in range(0, XBLOCK, XBLOCK_SUB): + xindex = xoffset + xoffset_sub + tl.arange(0, XBLOCK_SUB)[:] + xmask = xindex < xnumel + x0 = xindex + tmp0 = tl.load(in_ptr0 + (x0), xmask) + tmp1 = tl.load(in_ptr1 + (x0), xmask) + tmp2 = tmp0 < tmp1 + tmp3 = tl.where(tmp2, tmp0, tmp1) + tl.store(out_ptr0 + (xindex), tmp3, xmask) + +# @pytest.mark.usefixtures("pytest_runonce") +@pytest.mark.parametrize('param_list', + [ + ['float32', (2, 1024, 8), 2, 8192, 1024], + ['float16', (2, 1024, 8), 2, 8192, 1024], + ['int8', (2, 1024, 8), 2, 8192, 1024], + ] + ) +def test_where_lt_case2(param_list): + dtype, shape, ncore, xblock, xblock_sub = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + x1 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch_where_lt_case2(x0, x1) + y_cal = test_common.generate_tensor(shape, dtype).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + y_cal_txda = y_cal.to("txda") + triton_where_lt_case2[ncore, 1, 1](x0_txda, x1_txda, y_cal_txda, x0_txda.numel(), xblock, xblock_sub) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + test_common.validate_cmp(dtype, y_cal, y_ref) \ No newline at end of file diff --git a/test/wafer/ops/test_xor.py b/test/wafer/ops/test_xor.py new file mode 100644 index 00000000..cd9603ba --- /dev/null +++ b/test/wafer/ops/test_xor.py @@ -0,0 +1,102 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + +import pytest +import triton +import triton.language as tl +import time +import test_common +import torch +import torch_txda # noqa: F401 + + +def standard_unary(x0, dtype): + res = x0 + return res + + +def standard_binary(x0, y0, dtype): + res = x0 ^ y0 + return res + + +@triton.jit +def triton_elementwise_unary(in_ptr0, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + ret = x + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +@triton.jit +def triton_elementwise_binary(in_ptr0, in_ptr1, out_ptr0, N: tl.constexpr, NUMEL: tl.constexpr): + idx_block = tl.arange(0, NUMEL) + x = tl.load(in_ptr0 + idx_block, mask=idx_block < N) + y = tl.load(in_ptr1 + idx_block, mask=idx_block < N) + ret = x ^ y + tl.store(out_ptr0 + idx_block, ret, mask=idx_block < N) + + +types = [ + # (torch.float32, 'float32'), + # (torch.float16, 'float16'), + # (torch.bfloat16, 'bfloat16'), + (torch.int8, 'int8'), + # (torch.int16, 'int16'), + # (torch.int32, 'int32'), + # (torch.int64, 'int64'), +] + +shapes = [ + (3, 32), + (-32, 32), + (37, 64), + (-256, 256), + (781, 1024), +] + +map_for_64_t = {37: 31} + + +@pytest.mark.parametrize('dtype,sigtype', types) +@pytest.mark.parametrize('N,NUMEL', shapes) +def test_elementwsie_common(dtype, sigtype, N, NUMEL): + N = (-N) // torch.tensor(0, dtype=dtype).element_size() if N < 0 else N + + if sigtype == "int64": + N = map_for_64_t[N] if N in map_for_64_t else N + + print(f"elementwise : ({N},) {dtype} {sigtype}") + + x0 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + x1 = test_common.generate_tensor(shape=(N,), dtype=sigtype).cpu() + ans = standard_binary(x0, x1, dtype) + print(ans) + + out = torch.zeros((N,), dtype=dtype).cpu() + x0_txda = x0.to("txda") + x1_txda = x1.to("txda") + out_txda = out.to("txda") + triton_elementwise_binary[1, 1, 1](x0_txda, x1_txda, out_txda, N=N, NUMEL=NUMEL, debug=True) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + print(out) + + test_common.validate_cmp(sigtype, out, ans) diff --git a/test/wafer/ops/test_xor_sum.py b/test/wafer/ops/test_xor_sum.py new file mode 100644 index 00000000..9b92d5ac --- /dev/null +++ b/test/wafer/ops/test_xor_sum.py @@ -0,0 +1,89 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_( + in_ptr0, out_ptr0, xnumel, ynumel, XBLOCK: tl.constexpr, RBLOCK: tl.constexpr +): + X = xnumel + Y = ynumel + xoffset = tl.program_id(0) * XBLOCK + xindex = xoffset + tl.arange(0, XBLOCK) + + x0 = xindex[:, None] + rbase = tl.arange(0, RBLOCK) + _tmp6 = tl.full([XBLOCK, RBLOCK], 0, tl.int32) + for roffset in range(0, ynumel, RBLOCK): + rindex = roffset + rbase + rmask = None + r1 = rindex[None, :] + tmp0 = tl.load(in_ptr0 + (r1 + (Y * x0)), rmask) + _tmp6 = _tmp6 ^ tmp0 + tmp6 = tl.xor_sum(_tmp6, 1) + + tl.store(out_ptr0 + (xindex), tmp6, None) + + +def bar(tensor): + N, M = tensor.shape + result = torch.zeros(N, dtype=tensor.dtype, device=tensor.device) + for i in range(N): + row_xor_sum = 0 + for j in range(M): + row_xor_sum ^= tensor[i, j].item() + result[i] = row_xor_sum + return result + + +@pytest.mark.parametrize( + "param_list", + [ + ["int8", (64, 32), 64, 32], + ["int32", (64, 32), 64, 32], + ], +) +def test_case(param_list): + dtype, shape, xblock, rblock = param_list + a = test_common.generate_tensor(shape, dtype).cpu() + + std_ret = bar(a) + print(f"std_ret={std_ret}") + + value = torch.empty_strided((a.shape[0],), (1,), dtype=eval("torch." + dtype)).cpu() + XBLOCK = xblock + RBLOCK = rblock + NBLOCK = a.shape[0] // XBLOCK + a_txda = a.to("txda") + value_txda = value.to("txda") + fn_npu_[NBLOCK, 1, 1](a_txda, value_txda, a_txda.shape[0], a_txda.shape[1], XBLOCK, RBLOCK) + with torch.no_grad(): + value.copy_(value_txda.cpu()) + print(f"triton_ret={value}") + + torch.testing.assert_close(value, std_ret) diff --git a/test/wafer/ops/test_zeros.py b/test/wafer/ops/test_zeros.py new file mode 100644 index 00000000..faaa6314 --- /dev/null +++ b/test/wafer/ops/test_zeros.py @@ -0,0 +1,125 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_f32(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.zeros((XB, YB, ZB), dtype=tl.float32) + + oidx = ( + xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@triton.jit +def fn_npu_f16(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.zeros((XB, YB, ZB), dtype=tl.float16) + + oidx = ( + xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@triton.jit +def fn_npu_i8(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.zeros((XB, YB, ZB), dtype=tl.int8) + + oidx = ( + xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 256, 16), 1, 2, 256, 16], + ["float32", (8, 8, 4), 1, 8, 8, 4], + ["float16", (2, 256, 16), 1, 2, 256, 16], + ["float16", (8, 8, 4), 1, 8, 8, 4], + ["int8", (2, 256, 16), 1, 2, 256, 16], + ["int8", (8, 8, 4), 1, 8, 8, 4], + ], +) +def test_case(param_list): + dtype, shape, ncore, XB, YB, ZB = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + + y_ref = torch.full((XB, YB, ZB), 0, dtype=eval("torch." + dtype)).cpu() + print(f"y_ref = {y_ref[0, 0, 0:4]}") + + y_cal = torch.randint(1, (XB, YB, ZB), dtype=eval("torch." + dtype)).cpu() + if dtype == "float32": + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + fn_npu_f32[ncore, 1, 1](y_cal_txda, x0_txda, XB, YB, ZB) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + elif dtype == "float16": + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + fn_npu_f16[ncore, 1, 1](y_cal_txda, x0_txda, XB, YB, ZB) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + else: + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + fn_npu_i8[ncore, 1, 1](y_cal_txda, x0_txda, XB, YB, ZB) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + print(f"y_cal = {y_cal[0, 0, 0:4]}") + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/ops/test_zeroslike.py b/test/wafer/ops/test_zeroslike.py new file mode 100644 index 00000000..88b092dd --- /dev/null +++ b/test/wafer/ops/test_zeroslike.py @@ -0,0 +1,73 @@ +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. + + +import triton +import triton.language as tl +import torch +import torch_txda # noqa: F401 +import pytest +import test_common + + +@triton.jit +def fn_npu_(output_ptr, x_ptr, XB: tl.constexpr, YB: tl.constexpr, ZB: tl.constexpr): + xidx = tl.arange(0, XB) + yidx = tl.arange(0, YB) + zidx = tl.arange(0, ZB) + + idx = xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + + X = tl.load(x_ptr + idx) + + ret = tl.zeros_like(X) + + oidx = ( + xidx[:, None, None] * YB * ZB + yidx[None, :, None] * ZB + zidx[None, None, :] + ) + + tl.store(output_ptr + oidx, ret) + + +@pytest.mark.parametrize( + "param_list", + [ + ["float32", (2, 256, 16), 1, 2, 256, 16], + ["float32", (8, 8, 4), 1, 8, 8, 4], + ["float16", (2, 256, 16), 1, 2, 256, 16], + ["float16", (8, 8, 4), 1, 8, 8, 4], + ["int8", (2, 256, 16), 1, 2, 256, 16], + ["int8", (8, 8, 4), 1, 8, 8, 4], + ], +) +def test_case(param_list): + dtype, shape, ncore, XB, YB, ZB = param_list + x0 = test_common.generate_tensor(shape, dtype).cpu() + y_ref = torch.zeros_like(x0, dtype=eval("torch." + dtype)).cpu() + print(f"y_ref = {y_ref[0, 0, 0:4]}") + y_cal = torch.zeros(shape, dtype=eval("torch." + dtype)).cpu() + + y_cal_txda = y_cal.to("txda") + x0_txda = x0.to("txda") + fn_npu_[ncore, 1, 1](y_cal_txda, x0_txda, XB, YB, ZB) + with torch.no_grad(): + y_cal.copy_(y_cal_txda.cpu()) + print(f"y_cal = {y_cal[0, 0, 0:4]}") + test_common.validate_cmp(dtype, y_cal, y_ref) diff --git a/test/wafer/runtime/conftest.py b/test/wafer/runtime/conftest.py new file mode 100644 index 00000000..1b235d15 --- /dev/null +++ b/test/wafer/runtime/conftest.py @@ -0,0 +1,7 @@ +"""Native math tests use ordinary TXDA tensors and the production launcher.""" +import pytest + + +@pytest.fixture(autouse=True) +def require_device(wafer_device): + return wafer_device diff --git a/test/wafer/runtime/test_autotune.py b/test/wafer/runtime/test_autotune.py new file mode 100644 index 00000000..155e03ff --- /dev/null +++ b/test/wafer/runtime/test_autotune.py @@ -0,0 +1,33 @@ +"""Autotuner timing and cached invocation with real torch_txda tensor inputs.""" +import math +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +from triton.testing import do_bench + + +def test_native_autotune(): + measured = [] + + def benchmark(fn, quantiles): + times = do_bench(fn, warmup=1, rep=3, quantiles=quantiles) + assert all(math.isfinite(t) and t > 0 for t in times) + measured.append(times) + return times + + @triton.autotune(configs=[triton.Config({'BLOCK': 64}), triton.Config({'BLOCK': 128})], + key=['N'], do_bench=benchmark) + @triton.jit + def add_one(X, Y, N: tl.constexpr, BLOCK: tl.constexpr): + i = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(Y + i, tl.load(X + i, i < N, 0) + 1, i < N) + + host = torch.arange(257, dtype=torch.float32) + x, y = (host).to("txda"), (torch.zeros_like(host)).to("txda") + add_one[lambda meta: (triton.cdiv(257, meta['BLOCK']),)](x, y, N=257) + assert len(measured) == 2 + assert add_one.best_config.kwargs['BLOCK'] in (64, 128) + add_one[lambda meta: (triton.cdiv(257, meta['BLOCK']),)](x, y, N=257) + assert len(measured) == 2 # Cache hit must not repeat benchmarking. + torch.testing.assert_close(y.cpu(), host + 1, rtol=0, atol=0) diff --git a/test/wafer/runtime/test_elementwise.py b/test/wafer/runtime/test_elementwise.py new file mode 100644 index 00000000..073ff3e2 --- /dev/null +++ b/test/wafer/runtime/test_elementwise.py @@ -0,0 +1,29 @@ +"""Add/mul with CPU references, matching the Ascend comparison tolerances.""" +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def binary_kernel(X, Y, Z, N: tl.constexpr, MUL: tl.constexpr, B: tl.constexpr): + i = tl.program_id(0) * B + tl.arange(0, B) + x, y = tl.load(X + i, i < N, 0), tl.load(Y + i, i < N, 0) + z = x * y if MUL else x + y + tl.store(Z + i, z, i < N) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16, torch.int32]) +@pytest.mark.parametrize("size", [1, 31, 256, 257]) +@pytest.mark.parametrize("multiply", [False, True]) +def test_binary(dtype, size, multiply): + a = (torch.arange(size) % 17 - 8).to(dtype) + b = (torch.arange(size) % 7 - 3).to(dtype) + x, y = (a).to("txda"), (b).to("txda") + z = (torch.zeros_like(a)).to("txda") + binary_kernel[(triton.cdiv(size, 256),)](x, y, z, size, multiply, 256) + expected = a * b if multiply else a + b + # These bounded integer-valued inputs are exactly representable in all dtypes. + torch.testing.assert_close(z.cpu(), expected, rtol=0, atol=0) + diff --git a/test/wafer/runtime/test_matmul.py b/test/wafer/runtime/test_matmul.py new file mode 100644 index 00000000..a1f6a6e3 --- /dev/null +++ b/test/wafer/runtime/test_matmul.py @@ -0,0 +1,33 @@ +"""Native TXDA allocation, masked GEMM, FP32 accumulation and CPU comparison.""" +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def matmul_kernel(A, B, C, M: tl.constexpr, N: tl.constexpr, K: tl.constexpr): + m = tl.program_id(0) * 16 + tl.arange(0, 16) + n = tl.program_id(1) * 16 + tl.arange(0, 16) + k = tl.arange(0, 16) + acc = tl.full((16, 16), 0, tl.float32) + for start in range(tl.cdiv(K, 16)): + kk = start * 16 + k + a = tl.load(A + m[:, None] * K + kk[None, :], (m[:, None] < M) & (kk[None, :] < K), 0) + b = tl.load(B + kk[:, None] * N + n[None, :], (kk[:, None] < K) & (n[None, :] < N), 0) + acc = tl.dot(a, b, acc) + tl.store(C + m[:, None] * N + n[None, :], acc, (m[:, None] < M) & (n[None, :] < N)) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("shape", [(16, 16, 16), (32, 64, 32), (17, 19, 33)]) +def test_native_matmul(dtype, shape): + m, n, k = shape + a = (torch.arange(m * k).reshape(m, k) % 9 - 4).to(dtype) / 8 + b = (torch.arange(k * n).reshape(k, n) % 7 - 3).to(dtype) / 8 + x, y = (a).to("txda"), (b).to("txda") + z = (torch.zeros((m, n), dtype=dtype)).to("txda") + matmul_kernel[(triton.cdiv(m, 16), triton.cdiv(n, 16))](x, y, z, m, n, k) + torch.testing.assert_close(z.cpu(), (a.float() @ b.float()).to(dtype), rtol=1e-3, atol=1e-3) + diff --git a/test/wafer/runtime/test_reduction.py b/test/wafer/runtime/test_reduction.py new file mode 100644 index 00000000..89d19eaf --- /dev/null +++ b/test/wafer/runtime/test_reduction.py @@ -0,0 +1,34 @@ +"""Reduction and softmax, with no torch_txda arithmetic in the oracle.""" +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def row_kernel(X, Y, N: tl.constexpr, SOFTMAX: tl.constexpr, B: tl.constexpr): + row = tl.program_id(0) + i = tl.arange(0, B) + x = tl.load(X + row * N + i, i < N, other=0).to(tl.float32) + if SOFTMAX: + x = tl.where(i < N, x, float('-inf')) + exp = tl.exp(x - tl.max(x, 0)) + value = exp / tl.sum(exp, 0) + tl.store(Y + row * N + i, value, i < N) + else: + tl.store(Y + row, tl.sum(x, 0)) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16]) +@pytest.mark.parametrize("columns", [32, 127, 256]) +@pytest.mark.parametrize("softmax", [False, True]) +def test_row_reduce(dtype, columns, softmax): + host = (torch.arange(4 * columns).reshape(4, columns) % 17 - 8).to(dtype) / 8 + expected = host.float().softmax(1).to(dtype) if softmax else host.float().sum(1) + x = (host).to("txda") + y = (torch.zeros_like(expected)).to("txda") + row_kernel[(4,)](x, y, columns, softmax, triton.next_power_of_2(columns)) + tolerance = 1e-4 if dtype == torch.float32 else 1e-3 + torch.testing.assert_close(y.cpu(), expected, rtol=tolerance, atol=tolerance) + diff --git a/test/wafer/runtime/test_tensor_runtime.py b/test/wafer/runtime/test_tensor_runtime.py new file mode 100644 index 00000000..1a4cbce4 --- /dev/null +++ b/test/wafer/runtime/test_tensor_runtime.py @@ -0,0 +1,52 @@ +"""Native counterpart to Ascend load/store cases; exercises the tensor ABI.""" +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def copy_kernel(X, Y, N: tl.constexpr, B: tl.constexpr): + i = tl.program_id(0) * B + tl.arange(0, B) + tl.store(Y + i, tl.load(X + i, i < N, 0), i < N) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16, + torch.int8, torch.int16, torch.int32, torch.int64]) +@pytest.mark.parametrize("size", [1, 31, 256, 257]) +def test_native_copy(dtype, size): + host = (torch.arange(size) % 61 - 30).to(dtype) + x = host.to("txda") + # Keep the offset/boundary check explicit after removing device_tensor. + storage = torch.full((size + 128,), -83, dtype=dtype).to("txda") + y = storage[64:64 + size] + copy_kernel[(triton.cdiv(size, 256),)](x, y, size, 256) + torch.testing.assert_close(y.cpu(), host, rtol=0, atol=0) + actual = storage.cpu() + torch.testing.assert_close(actual[:64], torch.full((64,), -83, dtype=dtype), rtol=0, atol=0) + torch.testing.assert_close(actual[64 + size:], torch.full((64,), -83, dtype=dtype), rtol=0, atol=0) + + +def test_explicit_stream_and_cache(): + # Framework events and the compiler launcher must use the same stream. + stream = torch.txda.Stream() + host = torch.arange(257, dtype=torch.float32) + storage = torch.cat((host, torch.zeros(15), torch.zeros_like(host))).to("txda") + x, y = storage[:257], storage[272:] + assert x.untyped_storage().data_ptr() == y.untyped_storage().data_ptr() + with torch.txda.stream(stream): + driver = triton.runtime.driver.active + assert driver.get_current_stream(torch.txda.current_device()) == stream.txda_stream + start, end = torch.txda.Event(enable_timing=True), torch.txda.Event(enable_timing=True) + start.record() + first = copy_kernel[(2,)](x, y, 257, 256) + # The intermediate stays on TXDA between launches; both views retain + # their shared allocation and the destination's nonzero offset. + second = copy_kernel[(2,)](y, x, 257, 256) + end.record() + end.synchronize() + assert first is second + assert start.elapsed_time(end) >= 0 + torch.testing.assert_close(y.cpu(), host, rtol=0, atol=0) + torch.testing.assert_close(x.cpu(), host, rtol=0, atol=0) diff --git a/test/wafer/suites/accepted.txt b/test/wafer/suites/accepted.txt new file mode 100644 index 00000000..abe97d86 --- /dev/null +++ b/test/wafer/suites/accepted.txt @@ -0,0 +1,2866 @@ +test/wafer/runtime/test_autotune.py::test_native_autotune +test/wafer/runtime/test_elementwise.py::test_binary[False-1-dtype0] +test/wafer/runtime/test_elementwise.py::test_binary[False-1-dtype1] +test/wafer/runtime/test_elementwise.py::test_binary[False-1-dtype2] +test/wafer/runtime/test_elementwise.py::test_binary[False-1-dtype3] +test/wafer/runtime/test_elementwise.py::test_binary[False-256-dtype0] +test/wafer/runtime/test_elementwise.py::test_binary[False-256-dtype1] +test/wafer/runtime/test_elementwise.py::test_binary[False-256-dtype2] +test/wafer/runtime/test_elementwise.py::test_binary[False-256-dtype3] +test/wafer/runtime/test_elementwise.py::test_binary[False-257-dtype0] +test/wafer/runtime/test_elementwise.py::test_binary[False-257-dtype1] +test/wafer/runtime/test_elementwise.py::test_binary[False-257-dtype2] +test/wafer/runtime/test_elementwise.py::test_binary[False-257-dtype3] +test/wafer/runtime/test_elementwise.py::test_binary[False-31-dtype0] +test/wafer/runtime/test_elementwise.py::test_binary[False-31-dtype1] +test/wafer/runtime/test_elementwise.py::test_binary[False-31-dtype2] +test/wafer/runtime/test_elementwise.py::test_binary[False-31-dtype3] +test/wafer/runtime/test_elementwise.py::test_binary[True-1-dtype0] +test/wafer/runtime/test_elementwise.py::test_binary[True-1-dtype1] +test/wafer/runtime/test_elementwise.py::test_binary[True-1-dtype2] +test/wafer/runtime/test_elementwise.py::test_binary[True-1-dtype3] +test/wafer/runtime/test_elementwise.py::test_binary[True-256-dtype0] +test/wafer/runtime/test_elementwise.py::test_binary[True-256-dtype1] +test/wafer/runtime/test_elementwise.py::test_binary[True-256-dtype2] +test/wafer/runtime/test_elementwise.py::test_binary[True-256-dtype3] +test/wafer/runtime/test_elementwise.py::test_binary[True-257-dtype0] +test/wafer/runtime/test_elementwise.py::test_binary[True-257-dtype1] +test/wafer/runtime/test_elementwise.py::test_binary[True-257-dtype2] +test/wafer/runtime/test_elementwise.py::test_binary[True-257-dtype3] +test/wafer/runtime/test_elementwise.py::test_binary[True-31-dtype0] +test/wafer/runtime/test_elementwise.py::test_binary[True-31-dtype1] +test/wafer/runtime/test_elementwise.py::test_binary[True-31-dtype2] +test/wafer/runtime/test_elementwise.py::test_binary[True-31-dtype3] +test/wafer/runtime/test_matmul.py::test_native_matmul[shape0-dtype0] +test/wafer/runtime/test_matmul.py::test_native_matmul[shape0-dtype1] +test/wafer/runtime/test_matmul.py::test_native_matmul[shape1-dtype0] +test/wafer/runtime/test_matmul.py::test_native_matmul[shape1-dtype1] +test/wafer/runtime/test_matmul.py::test_native_matmul[shape2-dtype0] +test/wafer/runtime/test_matmul.py::test_native_matmul[shape2-dtype1] +test/wafer/runtime/test_reduction.py::test_row_reduce[False-127-dtype0] +test/wafer/runtime/test_reduction.py::test_row_reduce[False-127-dtype1] +test/wafer/runtime/test_reduction.py::test_row_reduce[False-256-dtype0] +test/wafer/runtime/test_reduction.py::test_row_reduce[False-256-dtype1] +test/wafer/runtime/test_reduction.py::test_row_reduce[False-32-dtype0] +test/wafer/runtime/test_reduction.py::test_row_reduce[False-32-dtype1] +test/wafer/runtime/test_reduction.py::test_row_reduce[True-127-dtype0] +test/wafer/runtime/test_reduction.py::test_row_reduce[True-127-dtype1] +test/wafer/runtime/test_reduction.py::test_row_reduce[True-256-dtype0] +test/wafer/runtime/test_reduction.py::test_row_reduce[True-256-dtype1] +test/wafer/runtime/test_reduction.py::test_row_reduce[True-32-dtype0] +test/wafer/runtime/test_reduction.py::test_row_reduce[True-32-dtype1] +test/wafer/runtime/test_tensor_runtime.py::test_explicit_stream_and_cache +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[1-dtype0] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[1-dtype1] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[1-dtype2] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[1-dtype3] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[1-dtype4] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[1-dtype5] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[1-dtype6] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[256-dtype0] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[256-dtype1] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[256-dtype2] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[256-dtype3] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[256-dtype4] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[256-dtype5] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[256-dtype6] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[257-dtype0] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[257-dtype1] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[257-dtype2] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[257-dtype3] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[257-dtype4] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[257-dtype5] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[257-dtype6] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[31-dtype0] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[31-dtype1] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[31-dtype2] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[31-dtype3] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[31-dtype4] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[31-dtype5] +test/wafer/runtime/test_tensor_runtime.py::test_native_copy[31-dtype6] +test/wafer/native_math/test_log1p.py::test_log1p[param_list0] +test/wafer/native_math/test_log1p.py::test_log1p_special_values[dtype0] +test/wafer/native_math/test_log1p.py::test_log1p_special_values[dtype1] +test/wafer/native_math/test_log1p.py::test_log1p_special_values[dtype2] +test/wafer/native_math/test_multi_return.py::test_correctness_functional[1.0-dtype0-1e-08-0.05-2-512-4096] +test/wafer/native_math/test_multi_return.py::test_correctness_functional[1.0-dtype1-1e-08-1e-06-2-512-4096] +test/wafer/native_math/test_relu.py::test_relu[param_list0] +test/wafer/native_math/test_relu.py::test_relu[param_list1] +test/wafer/native_math/test_relu.py::test_relu_special_values[dtype0] +test/wafer/native_math/test_relu.py::test_relu_special_values[dtype1] +test/wafer/native_math/test_relu.py::test_relu_special_values[dtype2] +test/wafer/native_math/test_unary.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/native_math/test_unary.py::test_elementwsie_common[-256-256-dtype1-float16] +test/wafer/native_math/test_unary.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/native_math/test_unary.py::test_elementwsie_common[-32-32-dtype1-float16] +test/wafer/native_math/test_unary.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/native_math/test_unary.py::test_elementwsie_common[3-32-dtype1-float16] +test/wafer/native_math/test_unary.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/native_math/test_unary.py::test_elementwsie_common[37-64-dtype1-float16] +test/wafer/native_math/test_unary.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/native_math/test_unary.py::test_elementwsie_common[781-1024-dtype1-float16] +test/wafer/native_math/test_unary.py::test_isnan[256-bfloat16] +test/wafer/native_math/test_unary.py::test_isnan[256-float16] +test/wafer/native_math/test_unary.py::test_isnan[256-float32] +test/wafer/native_math/test_unary.py::test_scalar_tanh_calc[param_list0] +test/wafer/native_math/test_unary.py::test_unary_special_values[atan-dtype0] +test/wafer/native_math/test_unary.py::test_unary_special_values[atan-dtype1] +test/wafer/native_math/test_unary.py::test_unary_special_values[atan-dtype2] +test/wafer/native_math/test_unary.py::test_unary_special_values[isnan-dtype0] +test/wafer/native_math/test_unary.py::test_unary_special_values[isnan-dtype1] +test/wafer/native_math/test_unary.py::test_unary_special_values[isnan-dtype2] +test/wafer/native_math/test_unary.py::test_unary_special_values[tanh-dtype0] +test/wafer/native_math/test_unary.py::test_unary_special_values[tanh-dtype1] +test/wafer/native_math/test_unary.py::test_unary_special_values[tanh-dtype2] +test/wafer/ops/test_2d_permute.py::test_cases[16-256] +test/wafer/ops/test_2d_permute.py::test_cases[16-32] +test/wafer/ops/test_2d_permute.py::test_cases[16-64] +test/wafer/ops/test_2d_permute.py::test_cases[32-256] +test/wafer/ops/test_2d_permute.py::test_cases[32-32] +test/wafer/ops/test_2d_permute.py::test_cases[32-64] +test/wafer/ops/test_3Dgrid.py::test_3dgrid[size0] +test/wafer/ops/test_abs.py::test_case[param_list0] +test/wafer/ops/test_abs.py::test_case[param_list1] +test/wafer/ops/test_abs_2.py::test_abs[param_list0] +test/wafer/ops/test_abs_2.py::test_abs[param_list1] +test/wafer/ops/test_add.py::test_all_blocks_parallel[param_list0] +test/wafer/ops/test_add.py::test_all_blocks_parallel[param_list1] +test/wafer/ops/test_add.py::test_all_blocks_parallel[param_list2] +test/wafer/ops/test_add.py::test_case[param_list0] +test/wafer/ops/test_add.py::test_case[param_list1] +test/wafer/ops/test_add_multi_return.py::test_all_blocks_parallel[param_list0] +test/wafer/ops/test_add_multi_return.py::test_all_blocks_parallel[param_list1] +test/wafer/ops/test_add_multi_return.py::test_all_blocks_parallel[param_list2] +test/wafer/ops/test_add_multi_return.py::test_case[param_list0] +test/wafer/ops/test_add_multi_return.py::test_case[param_list1] +test/wafer/ops/test_advance.py::test_advance_supplement[shape0-float32] +test/wafer/ops/test_advance.py::test_advance_supplement[shape0-int16] +test/wafer/ops/test_advance.py::test_advance_supplement[shape0-int32] +test/wafer/ops/test_advance.py::test_advance_supplement[shape1-float32] +test/wafer/ops/test_advance.py::test_advance_supplement[shape1-int16] +test/wafer/ops/test_advance.py::test_advance_supplement[shape1-int32] +test/wafer/ops/test_advance.py::test_advance_supplement[shape2-float32] +test/wafer/ops/test_advance.py::test_advance_supplement[shape2-int16] +test/wafer/ops/test_advance.py::test_advance_supplement[shape2-int32] +test/wafer/ops/test_advance.py::test_advance_supplement[shape3-float32] +test/wafer/ops/test_advance.py::test_advance_supplement[shape3-int16] +test/wafer/ops/test_advance.py::test_advance_supplement[shape3-int32] +test/wafer/ops/test_advance.py::test_advance_with_boundary_check[shape0-float32] +test/wafer/ops/test_advance.py::test_advance_with_boundary_check[shape0-int16] +test/wafer/ops/test_advance.py::test_advance_with_boundary_check[shape0-int32] +test/wafer/ops/test_advance.py::test_advance_with_boundary_check[shape1-float32] +test/wafer/ops/test_advance.py::test_advance_with_boundary_check[shape1-int16] +test/wafer/ops/test_advance.py::test_advance_with_boundary_check[shape1-int32] +test/wafer/ops/test_advance.py::test_npu[*fp16-data_type2-2-256-16] +test/wafer/ops/test_advance.py::test_npu[*fp16-data_type3-8-8-4] +test/wafer/ops/test_advance.py::test_npu[*fp32-data_type0-2-256-16] +test/wafer/ops/test_advance.py::test_npu[*fp32-data_type1-8-8-4] +test/wafer/ops/test_advance.py::test_npu[*i8-data_type4-2-256-16] +test/wafer/ops/test_advance.py::test_npu[*i8-data_type5-8-8-4] +test/wafer/ops/test_and.py::test_and[param_list0] +test/wafer/ops/test_arange.py::test_case[param_list1] +test/wafer/ops/test_arange.py::test_case_access[param_list1] +test/wafer/ops/test_arange.py::test_arange_invalid_range[invalid_param_list0] +test/wafer/ops/test_arange.py::test_arange_invalid_revinput[invalid_param_list0] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-0-shape0-float32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-0-shape0-int32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-0-shape1-float32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-0-shape1-int32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-0-shape2-float32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-0-shape2-int32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-1-shape1-float32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-1-shape1-int32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-1-shape2-float32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-1-shape2-int32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-2-shape2-float32] +test/wafer/ops/test_associative_scan.py::test_scan[False-maximum-2-shape2-int32] +test/wafer/ops/test_associative_scan_multi_input.py::test_multi_input_prefix_sum[shape0-0] +test/wafer/ops/test_associative_scan_multi_input.py::test_multi_input_prefix_sum[shape1-0] +test/wafer/ops/test_associative_scan_multi_input.py::test_multi_input_prefix_sum[shape2-1] +test/wafer/ops/test_block_ptr.py::test_npu[*fp16-data_type2-2-256-16] +test/wafer/ops/test_block_ptr.py::test_npu[*fp16-data_type3-8-8-4] +test/wafer/ops/test_block_ptr.py::test_npu[*fp32-data_type0-2-256-16] +test/wafer/ops/test_block_ptr.py::test_npu[*fp32-data_type1-8-8-4] +test/wafer/ops/test_block_ptr.py::test_npu[*i8-data_type4-2-256-16] +test/wafer/ops/test_block_ptr.py::test_npu[*i8-data_type5-8-8-4] +test/wafer/ops/test_broadcast_op.py::test_broadcast +test/wafer/ops/test_cat_dim.py::test_cat_dim0 +test/wafer/ops/test_cat_dim.py::test_cat_dim1 +test/wafer/ops/test_cdiv.py::test_cdiv[param_list0] +test/wafer/ops/test_ceil.py::test_ceil[param_list0] +test/wafer/ops/test_clamp.py::test_clamp[param_list0] +test/wafer/ops/test_clamp.py::test_clamp[param_list1] +test/wafer/ops/test_conv.py::test_conv_transpose2d_param[2-32-32-3-32-32-5-1-2-float16] +test/wafer/ops/test_conv.py::test_conv_transpose2d_param[2-4-4-3-8-8-2-1-1-float16] +test/wafer/ops/test_cos.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_cos.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_cos.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_cos.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_cos.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_cos_2.py::test_cos[param_list0] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[-32-1-32-8-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[-32-3-32-8-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[3-1-8-8-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[3-3-8-8-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[37-1-64-8-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[37-3-64-8-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[57--32-64-32-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[57-3-64-16-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[57-37-64-64-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[64--32-64-32-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[64-3-64-16-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_eq_dim0_common[64-37-64-64-dtype0-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[-32-1-32-8-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[-32-1-32-8-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[-32-1-32-8-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[-32-3-32-8-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[-32-3-32-8-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[-32-3-32-8-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[3-1-8-8-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[3-1-8-8-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[3-1-8-8-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[3-3-8-8-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[3-3-8-8-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[3-3-8-8-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[37-1-64-8-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[37-1-64-8-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[37-1-64-8-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[37-3-64-8-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[37-3-64-8-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[37-3-64-8-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57--32-64-32-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57--32-64-32-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57--32-64-32-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57-3-64-16-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57-3-64-16-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57-3-64-16-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57-37-64-64-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57-37-64-64-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[57-37-64-64-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64--32-64-32-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64--32-64-32-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64--32-64-32-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64-3-64-16-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64-3-64-16-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64-3-64-16-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64-37-64-64-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64-37-64-64-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_gt_dim0_common[64-37-64-64-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64--32-64-32-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64--32-64-32-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64--32-64-32-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64-3-64-16-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64-3-64-16-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64-3-64-16-dtype2-int8] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64-37-64-64-dtype0-float32] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64-37-64-64-dtype1-float16] +test/wafer/ops/test_count_dim0.py::test_count_lt_dim0_common[64-37-64-64-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[-32-1-32-8-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[-32-3-32-8-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[3-1-8-8-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[3-3-8-8-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[37-1-64-8-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[37-3-64-8-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[57--32-64-32-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[57-3-64-16-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[64--32-64-32-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_eq_dim0_common[64-3-64-16-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[-32-1-32-8-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[-32-1-32-8-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[-32-1-32-8-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[-32-3-32-8-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[-32-3-32-8-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[-32-3-32-8-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[3-1-8-8-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[3-1-8-8-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[3-1-8-8-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[3-3-8-8-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[3-3-8-8-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[3-3-8-8-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[37-1-64-8-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[37-1-64-8-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[37-1-64-8-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[37-3-64-8-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[37-3-64-8-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[37-3-64-8-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[57--32-64-32-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[57--32-64-32-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[57--32-64-32-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[57-3-64-16-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[57-3-64-16-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[57-3-64-16-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[64--32-64-32-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[64--32-64-32-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[64--32-64-32-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[64-3-64-16-dtype0-float32] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[64-3-64-16-dtype1-float16] +test/wafer/ops/test_count_dim1.py::test_count_gt_dim0_common[64-3-64-16-dtype2-int8] +test/wafer/ops/test_count_dim1.py::test_count_lt_dim0_common[57--32-64-32-dtype0-int8] +test/wafer/ops/test_count_dim1.py::test_count_lt_dim0_common[64--32-64-32-dtype0-int8] +test/wafer/ops/test_cumprod.py::test_cumprod[False-0-shape0-bfloat16] +test/wafer/ops/test_cumprod.py::test_cumprod[False-0-shape0-float16] +test/wafer/ops/test_cumprod.py::test_cumprod[False-0-shape0-float32] +test/wafer/ops/test_cumprod.py::test_cumprod[False-0-shape0-int16] +test/wafer/ops/test_cumprod.py::test_cumprod[False-0-shape0-int32] +test/wafer/ops/test_cumprod.py::test_cumprod[False-0-shape0-int64] +test/wafer/ops/test_cumprod.py::test_cumprod[False-1-shape0-bfloat16] +test/wafer/ops/test_cumprod.py::test_cumprod[False-1-shape0-float16] +test/wafer/ops/test_cumprod.py::test_cumprod[False-1-shape0-float32] +test/wafer/ops/test_cumprod.py::test_cumprod[False-1-shape0-int16] +test/wafer/ops/test_cumprod.py::test_cumprod[False-1-shape0-int32] +test/wafer/ops/test_cumprod.py::test_cumprod[False-1-shape0-int64] +test/wafer/ops/test_cumsum.py::test_cumsum[False-0-shape0-bfloat16] +test/wafer/ops/test_cumsum.py::test_cumsum[False-0-shape0-float32] +test/wafer/ops/test_cumsum.py::test_cumsum[False-0-shape0-int16] +test/wafer/ops/test_cumsum.py::test_cumsum[False-0-shape0-int32] +test/wafer/ops/test_cumsum.py::test_cumsum[False-0-shape0-int64] +test/wafer/ops/test_cumsum.py::test_cumsum[False-1-shape0-bfloat16] +test/wafer/ops/test_cumsum.py::test_cumsum[False-1-shape0-float32] +test/wafer/ops/test_cumsum.py::test_cumsum[False-1-shape0-int16] +test/wafer/ops/test_cumsum.py::test_cumsum[False-1-shape0-int32] +test/wafer/ops/test_cumsum.py::test_cumsum[False-1-shape0-int64] +test/wafer/ops/test_debug_barrier.py::test_case[param_list0] +test/wafer/ops/test_device_print.py::test_device_print_fp32[float32] +test/wafer/ops/test_device_print.py::test_device_print_int32[int32] +test/wafer/ops/test_device_print.py::test_device_print_int16[int16] +test/wafer/ops/test_device_print.py::test_device_print_int8[int8] +test/wafer/ops/test_device_print.py::test_device_print_fp16[float16] +test/wafer/ops/test_div.py::test_case[param_list0] +test/wafer/ops/test_div.py::test_case[param_list1] +test/wafer/ops/test_div.py::test_case[param_list2] +test/wafer/ops/test_elementwise_ceil.py::test_elementwise_common[-256-256-dtype0-float32-ceil-triton_ceil-standard_ceil] +test/wafer/ops/test_elementwise_ceil.py::test_elementwise_common[-32-32-dtype0-float32-ceil-triton_ceil-standard_ceil] +test/wafer/ops/test_elementwise_ceil.py::test_elementwise_common[3-32-dtype0-float32-ceil-triton_ceil-standard_ceil] +test/wafer/ops/test_elementwise_ceil.py::test_elementwise_common[37-64-dtype0-float32-ceil-triton_ceil-standard_ceil] +test/wafer/ops/test_elementwise_ceil.py::test_elementwise_common[781-1024-dtype0-float32-ceil-triton_ceil-standard_ceil] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[-256-256-dtype0-float32-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[-256-256-dtype1-float16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[-256-256-dtype2-bfloat16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[-32-32-dtype0-float32-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[-32-32-dtype1-float16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[-32-32-dtype2-bfloat16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[3-32-dtype0-float32-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[3-32-dtype1-float16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[3-32-dtype2-bfloat16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[37-64-dtype0-float32-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[37-64-dtype1-float16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[37-64-dtype2-bfloat16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[781-1024-dtype0-float32-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[781-1024-dtype1-float16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_clip.py::test_elementwise_common[781-1024-dtype2-bfloat16-clamp-triton_clamp-standard_clamp] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype0-float32-f2i16-triton_f2i16-standard_f2i16-int16] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype0-float32-f2i32-triton_f2i32-standard_f2i32-int32] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype0-float32-f2i64-triton_f2i64-standard_f2i64-int64] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype0-float32-f2i8-triton_f2i8-standard_f2i8-int8] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype1-float16-f2i16-triton_f2i16-standard_f2i16-int16] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype1-float16-f2i32-triton_f2i32-standard_f2i32-int32] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype1-float16-f2i64-triton_f2i64-standard_f2i64-int64] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype1-float16-f2i8-triton_f2i8-standard_f2i8-int8] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype2-bfloat16-f2i16-triton_f2i16-standard_f2i16-int16] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype2-bfloat16-f2i32-triton_f2i32-standard_f2i32-int32] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype2-bfloat16-f2i64-triton_f2i64-standard_f2i64-int64] +test/wafer/ops/test_elementwise_f2i.py::test_elementwise_common[3-32-dtype2-bfloat16-f2i8-triton_f2i8-standard_f2i8-int8] +test/wafer/ops/test_elementwise_floor.py::test_elementwise_common[-256-256-dtype0-float32-floor-triton_floor-standard_floor] +test/wafer/ops/test_elementwise_floor.py::test_elementwise_common[-32-32-dtype0-float32-floor-triton_floor-standard_floor] +test/wafer/ops/test_elementwise_floor.py::test_elementwise_common[3-32-dtype0-float32-floor-triton_floor-standard_floor] +test/wafer/ops/test_elementwise_floor.py::test_elementwise_common[37-64-dtype0-float32-floor-triton_floor-standard_floor] +test/wafer/ops/test_elementwise_floor.py::test_elementwise_common[781-1024-dtype0-float32-floor-triton_floor-standard_floor] +test/wafer/ops/test_elementwise_i2f.py::test_elementwise_common[3-32-dtype0-int32-i2f32-triton_i2f_float32-standard_i2f_float32-float32] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[-256-256-dtype0-float32-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[-256-256-dtype1-float16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[-256-256-dtype2-bfloat16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[-32-32-dtype0-float32-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[-32-32-dtype1-float16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[-32-32-dtype2-bfloat16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[3-32-dtype0-float32-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[3-32-dtype1-float16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[3-32-dtype2-bfloat16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[37-64-dtype0-float32-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[37-64-dtype1-float16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[37-64-dtype2-bfloat16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[781-1024-dtype0-float32-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[781-1024-dtype1-float16-round-triton_round-standard_round] +test/wafer/ops/test_elementwise_round.py::test_elementwise_common[781-1024-dtype2-bfloat16-round-triton_round-standard_round] +test/wafer/ops/test_eq.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_eq.py::test_elementwsie_common[-256-256-dtype1-float16] +test/wafer/ops/test_eq.py::test_elementwsie_common[-256-256-dtype2-int8] +test/wafer/ops/test_eq.py::test_elementwsie_common[-256-256-dtype3-int16] +test/wafer/ops/test_eq.py::test_elementwsie_common[-256-256-dtype4-int32] +test/wafer/ops/test_eq.py::test_elementwsie_common[-256-256-dtype5-int64] +test/wafer/ops/test_eq.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_eq.py::test_elementwsie_common[-32-32-dtype1-float16] +test/wafer/ops/test_eq.py::test_elementwsie_common[-32-32-dtype2-int8] +test/wafer/ops/test_eq.py::test_elementwsie_common[-32-32-dtype3-int16] +test/wafer/ops/test_eq.py::test_elementwsie_common[-32-32-dtype4-int32] +test/wafer/ops/test_eq.py::test_elementwsie_common[-32-32-dtype5-int64] +test/wafer/ops/test_eq.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_eq.py::test_elementwsie_common[3-32-dtype1-float16] +test/wafer/ops/test_eq.py::test_elementwsie_common[3-32-dtype2-int8] +test/wafer/ops/test_eq.py::test_elementwsie_common[3-32-dtype3-int16] +test/wafer/ops/test_eq.py::test_elementwsie_common[3-32-dtype4-int32] +test/wafer/ops/test_eq.py::test_elementwsie_common[3-32-dtype5-int64] +test/wafer/ops/test_eq.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_eq.py::test_elementwsie_common[37-64-dtype1-float16] +test/wafer/ops/test_eq.py::test_elementwsie_common[37-64-dtype2-int8] +test/wafer/ops/test_eq.py::test_elementwsie_common[37-64-dtype3-int16] +test/wafer/ops/test_eq.py::test_elementwsie_common[37-64-dtype4-int32] +test/wafer/ops/test_eq.py::test_elementwsie_common[37-64-dtype5-int64] +test/wafer/ops/test_eq.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_eq.py::test_elementwsie_common[781-1024-dtype1-float16] +test/wafer/ops/test_eq.py::test_elementwsie_common[781-1024-dtype2-int8] +test/wafer/ops/test_eq.py::test_elementwsie_common[781-1024-dtype3-int16] +test/wafer/ops/test_eq.py::test_elementwsie_common[781-1024-dtype4-int32] +test/wafer/ops/test_eq.py::test_elementwsie_common[781-1024-dtype5-int64] +test/wafer/ops/test_eq_2.py::test_eq[param_list0] +test/wafer/ops/test_exp.py::test_case[param_list0] +test/wafer/ops/test_exp2.py::test_exp2[param_list0] +test/wafer/ops/test_exp_.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_exp_.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_exp_.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_exp_.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_exp_.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_expand_dims.py::test_npu[*fp16-data_type2-2-256-16] +test/wafer/ops/test_expand_dims.py::test_npu[*fp16-data_type3-8-8-4] +test/wafer/ops/test_expand_dims.py::test_npu[*fp32-data_type0-2-256-16] +test/wafer/ops/test_expand_dims.py::test_npu[*fp32-data_type1-8-8-4] +test/wafer/ops/test_expand_dims.py::test_npu[*i8-data_type4-2-256-16] +test/wafer/ops/test_expand_dims.py::test_npu[*i8-data_type5-8-8-4] +test/wafer/ops/test_extract_slice.py::test_extract_slice +test/wafer/ops/test_fdiv.py::test_fdiv[param_list0] +test/wafer/ops/test_fdiv.py::test_fdiv[param_list1] +test/wafer/ops/test_floor.py::test_floor[param_list0] +test/wafer/ops/test_floordiv.py::test_floordiv[1024-int32] +test/wafer/ops/test_floordiv.py::test_floordiv[16-int32] +test/wafer/ops/test_floordiv.py::test_floordiv[256-int32] +test/wafer/ops/test_floordiv.py::test_floordiv[4-int32] +test/wafer/ops/test_full.py::test_npu[fn_npu_f16-float16-dtype2-2-256-16] +test/wafer/ops/test_full.py::test_npu[fn_npu_f16-float16-dtype3-8-8-4] +test/wafer/ops/test_full.py::test_npu[fn_npu_f32-float32-dtype0-2-256-16] +test/wafer/ops/test_full.py::test_npu[fn_npu_f32-float32-dtype1-8-8-4] +test/wafer/ops/test_full.py::test_npu[fn_npu_i8-int8-dtype4-2-256-16] +test/wafer/ops/test_full.py::test_npu[fn_npu_i8-int8-dtype5-8-8-4] +test/wafer/ops/test_ge.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_ge.py::test_elementwsie_common[-256-256-dtype1-float16] +test/wafer/ops/test_ge.py::test_elementwsie_common[-256-256-dtype2-int8] +test/wafer/ops/test_ge.py::test_elementwsie_common[-256-256-dtype3-int16] +test/wafer/ops/test_ge.py::test_elementwsie_common[-256-256-dtype4-int32] +test/wafer/ops/test_ge.py::test_elementwsie_common[-256-256-dtype5-int64] +test/wafer/ops/test_ge.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_ge.py::test_elementwsie_common[-32-32-dtype1-float16] +test/wafer/ops/test_ge.py::test_elementwsie_common[-32-32-dtype2-int8] +test/wafer/ops/test_ge.py::test_elementwsie_common[-32-32-dtype3-int16] +test/wafer/ops/test_ge.py::test_elementwsie_common[-32-32-dtype4-int32] +test/wafer/ops/test_ge.py::test_elementwsie_common[-32-32-dtype5-int64] +test/wafer/ops/test_ge.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_ge.py::test_elementwsie_common[3-32-dtype1-float16] +test/wafer/ops/test_ge.py::test_elementwsie_common[3-32-dtype2-int8] +test/wafer/ops/test_ge.py::test_elementwsie_common[3-32-dtype3-int16] +test/wafer/ops/test_ge.py::test_elementwsie_common[3-32-dtype4-int32] +test/wafer/ops/test_ge.py::test_elementwsie_common[3-32-dtype5-int64] +test/wafer/ops/test_ge.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_ge.py::test_elementwsie_common[37-64-dtype1-float16] +test/wafer/ops/test_ge.py::test_elementwsie_common[37-64-dtype2-int8] +test/wafer/ops/test_ge.py::test_elementwsie_common[37-64-dtype3-int16] +test/wafer/ops/test_ge.py::test_elementwsie_common[37-64-dtype4-int32] +test/wafer/ops/test_ge.py::test_elementwsie_common[37-64-dtype5-int64] +test/wafer/ops/test_ge.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_ge.py::test_elementwsie_common[781-1024-dtype1-float16] +test/wafer/ops/test_ge.py::test_elementwsie_common[781-1024-dtype2-int8] +test/wafer/ops/test_ge.py::test_elementwsie_common[781-1024-dtype3-int16] +test/wafer/ops/test_ge.py::test_elementwsie_common[781-1024-dtype4-int32] +test/wafer/ops/test_ge.py::test_elementwsie_common[781-1024-dtype5-int64] +test/wafer/ops/test_ge_2.py::test_ge[param_list0] +test/wafer/ops/test_ge_2.py::test_ge[param_list1] +test/wafer/ops/test_ge_2.py::test_ge[param_list2] +test/wafer/ops/test_gelu.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_gelu.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_gelu.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_gelu.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_gelu.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_gt.py::test_gt[param_list0] +test/wafer/ops/test_hd_permute.py::test_hd_permute +test/wafer/ops/test_if_tensor.py::test_kernel +test/wafer/ops/test_insert_slice.py::test_insert_slice +test/wafer/ops/test_interleave.py::test_interleave[float16-data_type2-2-64-16] +test/wafer/ops/test_interleave.py::test_interleave[float16-data_type3-8-8-4] +test/wafer/ops/test_interleave.py::test_interleave[float32-data_type0-2-64-16] +test/wafer/ops/test_interleave.py::test_interleave[float32-data_type1-8-8-4] +test/wafer/ops/test_interleave.py::test_interleave[int8-data_type4-2-64-32] +test/wafer/ops/test_interleave.py::test_interleave[int8-data_type5-8-8-4] +test/wafer/ops/test_invert.py::test_invert[param_list0] +test/wafer/ops/test_invert.py::test_invert[param_list1] +test/wafer/ops/test_invert.py::test_invert[param_list2] +test/wafer/ops/test_invert.py::test_invert[param_list3] +test/wafer/ops/test_invert.py::test_invert[param_list4] +test/wafer/ops/test_join.py::test_join[float16-data_type2-4-64-4] +test/wafer/ops/test_join.py::test_join[float16-data_type3-8-8-4] +test/wafer/ops/test_join.py::test_join[float32-data_type0-4-64-4] +test/wafer/ops/test_join.py::test_join[float32-data_type1-8-8-4] +test/wafer/ops/test_join.py::test_join[int8-data_type4-4-128-4] +test/wafer/ops/test_join.py::test_join[int8-data_type5-8-8-4] +test/wafer/ops/test_join.py::test_join_axis_added_tiled[float16-data_type4-2-8-4-1] +test/wafer/ops/test_join.py::test_join_axis_added_tiled[float32-data_type0-4-4-4-0] +test/wafer/ops/test_join.py::test_join_axis_added_tiled[float32-data_type1-512-256-512-0] +test/wafer/ops/test_join.py::test_join_axis_added_tiled[float32-data_type2-512-256-512-1] +test/wafer/ops/test_join.py::test_join_axis_added_tiled[float32-data_type3-512-256-512-2] +test/wafer/ops/test_join.py::test_join_axis_added_tiled[int8-data_type5-4-8-2-2] +test/wafer/ops/test_lanzcos.py::test_lanzcos[shapes0] +test/wafer/ops/test_launcher_empty_signature.py::test_launcher_empty_signature +test/wafer/ops/test_layernorm.py::test_layernorm +test/wafer/ops/test_ldst.py::test_ldst_indirect_00 +test/wafer/ops/test_ldst.py::test_ldst_indirect_01 +test/wafer/ops/test_ldst.py::test_ldst_indirect_02 +test/wafer/ops/test_ldst.py::test_ldst_indirect_03 +test/wafer/ops/test_ldst.py::test_ldst_indirect_04 +test/wafer/ops/test_ldst.py::test_ldst_indirect_05 +test/wafer/ops/test_ldst.py::test_ldst_indirect_06 +test/wafer/ops/test_ldst.py::test_ldst_indirect_07 +test/wafer/ops/test_ldst.py::test_ldst_indirect_08 +test/wafer/ops/test_ldst.py::test_ldst_indirect_09 +test/wafer/ops/test_ldst.py::test_unstructured_mask_2d[param_list0] +test/wafer/ops/test_le.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_le.py::test_elementwsie_common[-256-256-dtype1-float16] +test/wafer/ops/test_le.py::test_elementwsie_common[-256-256-dtype2-int8] +test/wafer/ops/test_le.py::test_elementwsie_common[-256-256-dtype3-int16] +test/wafer/ops/test_le.py::test_elementwsie_common[-256-256-dtype4-int32] +test/wafer/ops/test_le.py::test_elementwsie_common[-256-256-dtype5-int64] +test/wafer/ops/test_le.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_le.py::test_elementwsie_common[-32-32-dtype1-float16] +test/wafer/ops/test_le.py::test_elementwsie_common[-32-32-dtype2-int8] +test/wafer/ops/test_le.py::test_elementwsie_common[-32-32-dtype3-int16] +test/wafer/ops/test_le.py::test_elementwsie_common[-32-32-dtype4-int32] +test/wafer/ops/test_le.py::test_elementwsie_common[-32-32-dtype5-int64] +test/wafer/ops/test_le.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_le.py::test_elementwsie_common[3-32-dtype1-float16] +test/wafer/ops/test_le.py::test_elementwsie_common[3-32-dtype2-int8] +test/wafer/ops/test_le.py::test_elementwsie_common[3-32-dtype3-int16] +test/wafer/ops/test_le.py::test_elementwsie_common[3-32-dtype4-int32] +test/wafer/ops/test_le.py::test_elementwsie_common[3-32-dtype5-int64] +test/wafer/ops/test_le.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_le.py::test_elementwsie_common[37-64-dtype1-float16] +test/wafer/ops/test_le.py::test_elementwsie_common[37-64-dtype2-int8] +test/wafer/ops/test_le.py::test_elementwsie_common[37-64-dtype3-int16] +test/wafer/ops/test_le.py::test_elementwsie_common[37-64-dtype4-int32] +test/wafer/ops/test_le.py::test_elementwsie_common[37-64-dtype5-int64] +test/wafer/ops/test_le.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_le.py::test_elementwsie_common[781-1024-dtype1-float16] +test/wafer/ops/test_le.py::test_elementwsie_common[781-1024-dtype2-int8] +test/wafer/ops/test_le.py::test_elementwsie_common[781-1024-dtype3-int16] +test/wafer/ops/test_le.py::test_elementwsie_common[781-1024-dtype4-int32] +test/wafer/ops/test_le.py::test_elementwsie_common[781-1024-dtype5-int64] +test/wafer/ops/test_load.py::test_load_store[param_list0] +test/wafer/ops/test_load.py::test_load_store[param_list1] +test/wafer/ops/test_load.py::test_load_store[param_list2] +test/wafer/ops/test_load.py::test_load_store[param_list3] +test/wafer/ops/test_load.py::test_load_store[param_list4] +test/wafer/ops/test_load.py::test_load_store[param_list5] +test/wafer/ops/test_load.py::test_load_store[param_list6] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list0] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list10] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list11] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list1] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list2] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list3] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list4] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list5] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list6] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list7] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list8] +test/wafer/ops/test_load.py::test_load_store_multi_d[param_list9] +test/wafer/ops/test_load_store.py::test_load_store[param_list0] +test/wafer/ops/test_load_store.py::test_load_store[param_list1] +test/wafer/ops/test_load_store.py::test_load_store[param_list2] +test/wafer/ops/test_load_store.py::test_load_store[param_list3] +test/wafer/ops/test_load_store.py::test_load_store[param_list4] +test/wafer/ops/test_load_store.py::test_load_store[param_list5] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list0] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list10] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list11] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list1] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list2] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list3] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list4] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list5] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list6] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list7] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list8] +test/wafer/ops/test_load_store.py::test_load_store_multi_d[param_list9] +test/wafer/ops/test_load_store.py::test_load_store_sge_mask[param_list0] +test/wafer/ops/test_load_store.py::test_load_store_sge_mask[param_list1] +test/wafer/ops/test_load_store.py::test_load_store_sge_mask[param_list2] +test/wafer/ops/test_load_store.py::test_load_store_sge_mask[param_list3] +test/wafer/ops/test_load_store.py::test_load_store_sge_mask[param_list4] +test/wafer/ops/test_load_store.py::test_load_store_sge_mask[param_list5] +test/wafer/ops/test_load_store.py::test_load_store_sle_mask[param_list0] +test/wafer/ops/test_load_store.py::test_load_store_sle_mask[param_list1] +test/wafer/ops/test_load_store.py::test_load_store_sle_mask[param_list2] +test/wafer/ops/test_load_store.py::test_load_store_sle_mask[param_list3] +test/wafer/ops/test_load_store.py::test_load_store_sle_mask[param_list4] +test/wafer/ops/test_load_store.py::test_load_store_sle_mask[param_list5] +test/wafer/ops/test_log.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_log.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_log.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_log.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_log.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_log2.py::test_log2[param_list0] +test/wafer/ops/test_log_2.py::test_exp2[param_list0] +test/wafer/ops/test_logical_and.py::test_and[param_list0] +test/wafer/ops/test_logical_or.py::test_logical_or[param_list0] +test/wafer/ops/test_lshift.py::test_elementwsie_common[-256-256-dtype0-int8] +test/wafer/ops/test_lshift.py::test_elementwsie_common[-32-32-dtype0-int8] +test/wafer/ops/test_lshift.py::test_elementwsie_common[3-32-dtype0-int8] +test/wafer/ops/test_lshift.py::test_elementwsie_common[37-64-dtype0-int8] +test/wafer/ops/test_lshift.py::test_elementwsie_common[781-1024-dtype0-int8] +test/wafer/ops/test_lt.py::test_lt[param_list0] +test/wafer/ops/test_max_dim0.py::test_max_dim0[dtype0-float32-64--32-64-32] +test/wafer/ops/test_max_dim0.py::test_max_dim0[dtype1-float16-64--32-64-32] +test/wafer/ops/test_max_dim0.py::test_max_dim0[dtype2-int8-64--32-64-32] +test/wafer/ops/test_max_dim1.py::test_max_dim1[dtype0-float32-64--32-64-32] +test/wafer/ops/test_max_dim1.py::test_max_dim1[dtype1-float16-64--32-64-32] +test/wafer/ops/test_max_dim1.py::test_max_dim1[dtype2-int8-64--32-64-32] +test/wafer/ops/test_max_vector.py::test_reduce_dim0_common[-256-256-dtype0-float32] +test/wafer/ops/test_max_vector.py::test_reduce_dim0_common[-32-32-dtype0-float32] +test/wafer/ops/test_max_vector.py::test_reduce_dim0_common[3-32-dtype0-float32] +test/wafer/ops/test_max_vector.py::test_reduce_dim0_common[37-64-dtype0-float32] +test/wafer/ops/test_max_vector.py::test_reduce_dim0_common[781-1024-dtype0-float32] +test/wafer/ops/test_maximum.py::test_maximum[param_list0] +test/wafer/ops/test_maximum.py::test_maximum[param_list1] +test/wafer/ops/test_maximum.py::test_maximum[param_list2] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype0-float32--256-3-256-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype0-float32-263-1-512-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype0-float32-37-3-64-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype0-float32-57-3-64-16] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype0-float32-64--32-64-32] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype1-float16--256-3-256-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype1-float16-263-1-512-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype1-float16-37-3-64-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype1-float16-57-3-64-16] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype1-float16-64--32-64-32] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype2-int8--256-3-256-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype2-int8-263-1-512-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype2-int8-37-3-64-8] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype2-int8-57-3-64-16] +test/wafer/ops/test_mean_dim0.py::test_mean_dim0[dtype2-int8-64--32-64-32] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype0-float32--256-3-256-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype0-float32-263-1-512-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype0-float32-37-3-64-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype0-float32-57-3-64-16] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype0-float32-64--32-64-32] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype1-float16--256-3-256-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype1-float16-263-1-512-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype1-float16-37-3-64-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype1-float16-57-3-64-16] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype1-float16-64--32-64-32] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype2-int8--256-3-256-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype2-int8-263-1-512-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype2-int8-37-3-64-8] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype2-int8-57-3-64-16] +test/wafer/ops/test_mean_dim1.py::test_mean_dim1[dtype2-int8-64--32-64-32] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[-256-256-dtype0-float32] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[-256-256-dtype1-int8] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[-32-32-dtype0-float32] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[-32-32-dtype1-int8] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[3-32-dtype0-float32] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[3-32-dtype1-int8] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[37-64-dtype0-float32] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[37-64-dtype1-int8] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[781-1024-dtype0-float32] +test/wafer/ops/test_mean_vector.py::test_mean_dim0[781-1024-dtype1-int8] +test/wafer/ops/test_min_dim0.py::test_min_dim0[dtype0-float32-64--32-64-32] +test/wafer/ops/test_min_dim0.py::test_min_dim0[dtype1-float16-64--32-64-32] +test/wafer/ops/test_min_dim0.py::test_min_dim0[dtype2-int8-64--32-64-32] +test/wafer/ops/test_min_dim1.py::test_min_dim1[dtype0-float32-64--32-64-32] +test/wafer/ops/test_min_dim1.py::test_min_dim1[dtype1-float16-64--32-64-32] +test/wafer/ops/test_min_dim1.py::test_min_dim1[dtype2-int8-64--32-64-32] +test/wafer/ops/test_min_vector.py::test_reduce_dim0_common[-256-256-dtype0-float32] +test/wafer/ops/test_min_vector.py::test_reduce_dim0_common[-32-32-dtype0-float32] +test/wafer/ops/test_min_vector.py::test_reduce_dim0_common[3-32-dtype0-float32] +test/wafer/ops/test_min_vector.py::test_reduce_dim0_common[37-64-dtype0-float32] +test/wafer/ops/test_min_vector.py::test_reduce_dim0_common[781-1024-dtype0-float32] +test/wafer/ops/test_minimum.py::test_minimum[param_list0] +test/wafer/ops/test_minimum.py::test_minimum[param_list1] +test/wafer/ops/test_minimum.py::test_minimum[param_list2] +test/wafer/ops/test_mod.py::test_case[param_list0] +test/wafer/ops/test_mod.py::test_case[param_list1] +test/wafer/ops/test_mul.py::test_case[param_list0] +test/wafer/ops/test_mul.py::test_case[param_list1] +test/wafer/ops/test_nearest.py::test_nearest[shapes0] +test/wafer/ops/test_neg.py::test_neg[param_list0] +test/wafer/ops/test_neg.py::test_neg[param_list1] +test/wafer/ops/test_neg.py::test_neg[param_list2] +test/wafer/ops/test_npu_indexing.py::test_npu_indexing +test/wafer/ops/test_npu_indexing2.py::test_npu_indexing2 +test/wafer/ops/test_or.py::test_or[param_list0] +test/wafer/ops/test_permute.py::test_permute_handwritten +test/wafer/ops/test_permute_full.py::test_permute[float16-data_type2-2-4-8] +test/wafer/ops/test_permute_full.py::test_permute[float16-data_type3-2-4-64] +test/wafer/ops/test_permute_full.py::test_permute[float32-data_type0-2-4-8] +test/wafer/ops/test_permute_full.py::test_permute[float32-data_type1-2-4-64] +test/wafer/ops/test_permute_full.py::test_permute[int8-data_type4-2-4-8] +test/wafer/ops/test_permute_full.py::test_permute[int8-data_type5-2-4-64] +test/wafer/ops/test_permute_reshape.py::test_permute_reshape +test/wafer/ops/test_precise_div.py::test_divRn[param_list0] +test/wafer/ops/test_precise_sqrt.py::test_sqrtrn_fp32 +test/wafer/ops/test_ravel.py::test_ravel[float16-dtype2-2-256-16] +test/wafer/ops/test_ravel.py::test_ravel[float16-dtype3-8-8-4] +test/wafer/ops/test_ravel.py::test_ravel[float32-dtype0-2-256-16] +test/wafer/ops/test_ravel.py::test_ravel[float32-dtype1-8-8-4] +test/wafer/ops/test_ravel.py::test_ravel[int8-dtype4-2-256-16] +test/wafer/ops/test_ravel.py::test_ravel[int8-dtype5-8-8-4] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype0-float32-countf-triton_gt-standard_gt-0.5] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype0-float32-countf-triton_lt-standard_lt-0.5] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype1-float16-countf-triton_gt-standard_gt-0.5] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype1-float16-countf-triton_lt-standard_lt-0.5] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype2-bfloat16-countf-triton_gt-standard_gt-0.5] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype2-bfloat16-countf-triton_lt-standard_lt-0.5] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype3-int8-counti-triton_count-standard_count-8] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype4-int16-counti-triton_count-standard_count-8] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype5-int32-counti-triton_count-standard_count-8] +test/wafer/ops/test_reduce_count_vector.py::test_reduce_count_vector[32-32-dtype6-int64-counti-triton_count-standard_count-8] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list0] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list1] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list2] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list3] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list4] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list5] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list6] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list7] +test/wafer/ops/test_reduce_mean.py::test_mean_pr[param_list8] +test/wafer/ops/test_reduce_sum.py::test_sum_pr[param_list0] +test/wafer/ops/test_reduce_sum.py::test_sum_pr[param_list1] +test/wafer/ops/test_reduce_sum.py::test_sum_pr[param_list2] +test/wafer/ops/test_reduce_sum.py::test_sum_pr[param_list3] +test/wafer/ops/test_reduce_sum.py::test_sum_pr[param_list4] +test/wafer/ops/test_reduce_sum.py::test_sum_pr[param_list5] +test/wafer/ops/test_reduce_sum.py::test_sum_pr[param_list6] +test/wafer/ops/test_reduce_sum.py::test_sum_pr[param_list7] +test/wafer/ops/test_reshape.py::test_ravel[float16-dtype2-2-256-16] +test/wafer/ops/test_reshape.py::test_ravel[float16-dtype3-8-8-4] +test/wafer/ops/test_reshape.py::test_ravel[float32-dtype0-2-256-16] +test/wafer/ops/test_reshape.py::test_ravel[float32-dtype1-8-8-4] +test/wafer/ops/test_reshape.py::test_ravel[int8-dtype4-2-256-16] +test/wafer/ops/test_reshape.py::test_ravel[int8-dtype5-8-8-4] +test/wafer/ops/test_rms_norm.py::test_cases +test/wafer/ops/test_rotary_embedding.py::test_cases +test/wafer/ops/test_rotatry_gpt.py::test_rotary_emb +test/wafer/ops/test_rotaty_embedding_gpt.py::test_cases +test/wafer/ops/test_rshift.py::test_elementwsie_common[-256-256-dtype0-int8] +test/wafer/ops/test_rshift.py::test_elementwsie_common[-32-32-dtype0-int8] +test/wafer/ops/test_rshift.py::test_elementwsie_common[3-32-dtype0-int8] +test/wafer/ops/test_rshift.py::test_elementwsie_common[37-64-dtype0-int8] +test/wafer/ops/test_rshift.py::test_elementwsie_common[781-1024-dtype0-int8] +test/wafer/ops/test_rsqrt.py::test_rsqrt[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_abs_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_add_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_ceil_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_cmpf_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_cos_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_div_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_erf_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_exp_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_extf_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_floor_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_log2_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_log_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_maximum_nanall_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_maximum_nannone_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_minimum_nanall_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_minimum_nannone_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_mul_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_negf_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_remf_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_rsqrt_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_sin_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_sqrt_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_sub_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_sum_calc[param_list0] +test/wafer/ops/test_scalar_calc.py::test_scalar_truncf_calc[param_list0] +test/wafer/ops/test_sigmoid.py::test_sigmoid[param_list0] +test/wafer/ops/test_silu.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_silu.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_silu.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_silu.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_silu.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_silu_and_mul.py::TestSiluAndMul::test_silu_and_mul[256] +test/wafer/ops/test_silu_and_mul.py::TestSiluAndMul::test_silu_and_mul[768] +test/wafer/ops/test_silu_and_mul.py::test +test/wafer/ops/test_sin.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_sin.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_sin.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_sin.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_sin.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_softmax.py::test_softmax[1823--100-dtype0-float32] +test/wafer/ops/test_softmax.py::test_softmax[1823--100-dtype1-float16] +test/wafer/ops/test_softmax.py::test_softmax[1823--100-dtype2-bfloat16] +test/wafer/ops/test_softmax.py::test_softmax[1823--256-dtype0-float32] +test/wafer/ops/test_softmax.py::test_softmax[1823--256-dtype1-float16] +test/wafer/ops/test_softmax.py::test_softmax[1823--256-dtype2-bfloat16] +test/wafer/ops/test_softmax.py::test_softmax[1823--32-dtype0-float32] +test/wafer/ops/test_softmax.py::test_softmax[1823--32-dtype1-float16] +test/wafer/ops/test_softmax.py::test_softmax[1823--32-dtype2-bfloat16] +test/wafer/ops/test_softmax.py::test_softmax[1823-2-dtype0-float32] +test/wafer/ops/test_softmax.py::test_softmax[1823-2-dtype1-float16] +test/wafer/ops/test_softmax.py::test_softmax[1823-2-dtype2-bfloat16] +test/wafer/ops/test_softmax.py::test_softmax[1823-4-dtype0-float32] +test/wafer/ops/test_softmax.py::test_softmax[1823-4-dtype1-float16] +test/wafer/ops/test_softmax.py::test_softmax[1823-4-dtype2-bfloat16] +test/wafer/ops/test_softmax.py::test_softmax[1823-781-dtype0-float32] +test/wafer/ops/test_softmax.py::test_softmax[1823-781-dtype1-float16] +test/wafer/ops/test_softmax.py::test_softmax[1823-781-dtype2-bfloat16] +test/wafer/ops/test_split.py::test_split[float16-data_type2-16-256-2] +test/wafer/ops/test_split.py::test_split[float16-data_type3-8-8-2] +test/wafer/ops/test_split.py::test_split[float32-data_type0-16-256-2] +test/wafer/ops/test_split.py::test_split[float32-data_type1-8-8-2] +test/wafer/ops/test_split.py::test_split[int8-data_type4-8-128-2] +test/wafer/ops/test_split.py::test_split[int8-data_type5-8-8-2] +test/wafer/ops/test_sqrt.py::test_elementwsie_common[-256-256-dtype0-float32] +test/wafer/ops/test_sqrt.py::test_elementwsie_common[-32-32-dtype0-float32] +test/wafer/ops/test_sqrt.py::test_elementwsie_common[3-32-dtype0-float32] +test/wafer/ops/test_sqrt.py::test_elementwsie_common[37-64-dtype0-float32] +test/wafer/ops/test_sqrt.py::test_elementwsie_common[781-1024-dtype0-float32] +test/wafer/ops/test_store_scalar.py::test_case +test/wafer/ops/test_sub.py::test_case[param_list0] +test/wafer/ops/test_sub.py::test_case[param_list1] +test/wafer/ops/test_sub.py::test_case[param_list2] +test/wafer/ops/test_sub.py::test_case[param_list3] +test/wafer/ops/test_sum.py::test_case_1[param_list0] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype0-float32--256-3-256-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype0-float32-263-1-512-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype0-float32-37-3-64-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype0-float32-57-3-64-16] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype0-float32-64--32-64-32] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype1-float16--256-3-256-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype1-float16-263-1-512-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype1-float16-37-3-64-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype1-float16-57-3-64-16] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype1-float16-64--32-64-32] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype2-int8--256-3-256-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype2-int8-263-1-512-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype2-int8-37-3-64-8] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype2-int8-57-3-64-16] +test/wafer/ops/test_sum_dim0.py::test_sum_dim0[dtype2-int8-64--32-64-32] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype0-float32--256-3-256-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype0-float32-263-1-512-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype0-float32-37-3-64-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype0-float32-57-3-64-16] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype0-float32-64--32-64-32] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype1-float16--256-3-256-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype1-float16-263-1-512-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype1-float16-37-3-64-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype1-float16-57-3-64-16] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype1-float16-64--32-64-32] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype2-int8--256-3-256-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype2-int8-263-1-512-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype2-int8-37-3-64-8] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype2-int8-57-3-64-16] +test/wafer/ops/test_sum_dim1.py::test_sum_dim1[dtype2-int8-64--32-64-32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape0-float32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape0-int32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape1-float32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape1-int32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape2-float32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape2-int32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape3-float32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape3-int32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape4-float32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape4-int32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape5-float32] +test/wafer/ops/test_sum_vector.py::test_reduce_sum[shape5-int32] +test/wafer/ops/test_sum_vector.py::test_sum[shape0-float32] +test/wafer/ops/test_sum_vector.py::test_sum[shape0-int32] +test/wafer/ops/test_sum_vector.py::test_sum[shape1-float32] +test/wafer/ops/test_sum_vector.py::test_sum[shape1-int32] +test/wafer/ops/test_sum_vector.py::test_sum[shape2-float32] +test/wafer/ops/test_sum_vector.py::test_sum[shape2-int32] +test/wafer/ops/test_sum_vector.py::test_sum[shape3-float32] +test/wafer/ops/test_sum_vector.py::test_sum[shape3-int32] +test/wafer/ops/test_sum_vector.py::test_sum[shape4-float32] +test/wafer/ops/test_sum_vector.py::test_sum[shape4-int32] +test/wafer/ops/test_sum_vector.py::test_sum[shape5-float32] +test/wafer/ops/test_sum_vector.py::test_sum[shape5-int32] +test/wafer/ops/test_swap.py::test[shape0] +test/wafer/ops/test_swap.py::test[shape1] +test/wafer/ops/test_swap.py::test[shape2] +test/wafer/ops/test_swap.py::test[shape3] +test/wafer/ops/test_swap.py::test[shape4] +test/wafer/ops/test_swap.py::test[shape5] +test/wafer/ops/test_swap.py::test[shape6] +test/wafer/ops/test_swap.py::test[shape7] +test/wafer/ops/test_swiglu.py::test_case[param_list0] +test/wafer/ops/test_swiglu.py::test_case[param_list1] +test/wafer/ops/test_swizzle2d.py::test_case[param_list0] +test/wafer/ops/test_template.py::test_cases +test/wafer/ops/test_tensor_get_item.py::test_tensor_get_item +test/wafer/ops/test_trans_3d.py::test_permute_3d[float32-shape0] +test/wafer/ops/test_triton_eq.py::test_case[param_list0] +test/wafer/ops/test_triton_eq.py::test_case[param_list1] +test/wafer/ops/test_triton_eq.py::test_case[param_list2] +test/wafer/ops/test_triton_le.py::test_case[param_list0] +test/wafer/ops/test_triton_le.py::test_case[param_list1] +test/wafer/ops/test_triton_le.py::test_case[param_list2] +test/wafer/ops/test_triton_lt.py::test_case[param_list0] +test/wafer/ops/test_triton_lt.py::test_case[param_list1] +test/wafer/ops/test_triton_lt.py::test_case[param_list2] +test/wafer/ops/test_triton_neq.py::test_case[param_list0] +test/wafer/ops/test_triton_neq.py::test_case[param_list1] +test/wafer/ops/test_triton_neq.py::test_case[param_list2] +test/wafer/ops/test_umulhi.py::test_umulhi +test/wafer/ops/test_unlign_sum.py::test_cases +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-bfloat16-dtype9-3-5-3] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-bool-dtype10-3-5-3] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-float16-dtype4-55-5-16] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-float16-dtype5-4-5-17] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-float16-dtype6-6-5-15] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-float16-dtype7-2-1928-3] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-float32-dtype8-2-255-9] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-int16-dtype1-3-5-3] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-int32-dtype2-2-255-9] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-int64-dtype3-2-5-3] +test/wafer/ops/test_unused_func_arg.py::test_npu[triton_unused_func_arg_kernel-int8-dtype0-2-255-9] +test/wafer/ops/test_view.py::test_case[param_list0] +test/wafer/ops/test_view.py::test_case[param_list1] +test/wafer/ops/test_view.py::test_case[param_list2] +test/wafer/ops/test_view.py::test_case[param_list3] +test/wafer/ops/test_view.py::test_case[param_list4] +test/wafer/ops/test_view.py::test_case[param_list5] +test/wafer/ops/test_where_lt.py::test_where_lt_case1[param_list0] +test/wafer/ops/test_where_lt.py::test_where_lt_case1[param_list1] +test/wafer/ops/test_where_lt.py::test_where_lt_case1[param_list2] +test/wafer/ops/test_where_mask.py::test_where_lt_case2[param_list0] +test/wafer/ops/test_where_mask.py::test_where_lt_case2[param_list1] +test/wafer/ops/test_where_mask.py::test_where_lt_case2[param_list2] +test/wafer/ops/test_where_var.py::test_where_lt_case2[param_list0] +test/wafer/ops/test_where_var.py::test_where_lt_case2[param_list1] +test/wafer/ops/test_where_var.py::test_where_lt_case2[param_list2] +test/wafer/ops/test_xor.py::test_elementwsie_common[-256-256-dtype0-int8] +test/wafer/ops/test_xor.py::test_elementwsie_common[-32-32-dtype0-int8] +test/wafer/ops/test_xor.py::test_elementwsie_common[3-32-dtype0-int8] +test/wafer/ops/test_xor.py::test_elementwsie_common[37-64-dtype0-int8] +test/wafer/ops/test_xor.py::test_elementwsie_common[781-1024-dtype0-int8] +test/wafer/ops/test_xor_sum.py::test_case[param_list0] +test/wafer/ops/test_xor_sum.py::test_case[param_list1] +test/wafer/ops/test_zeros.py::test_case[param_list0] +test/wafer/ops/test_zeros.py::test_case[param_list1] +test/wafer/ops/test_zeros.py::test_case[param_list2] +test/wafer/ops/test_zeros.py::test_case[param_list3] +test/wafer/ops/test_zeros.py::test_case[param_list4] +test/wafer/ops/test_zeros.py::test_case[param_list5] +test/wafer/ops/test_zeroslike.py::test_case[param_list0] +test/wafer/ops/test_zeroslike.py::test_case[param_list1] +test/wafer/ops/test_zeroslike.py::test_case[param_list2] +test/wafer/ops/test_zeroslike.py::test_case[param_list3] +test/wafer/ops/test_zeroslike.py::test_case[param_list4] +test/wafer/ops/test_zeroslike.py::test_case[param_list5] +third_party/wafer/examples/flagtree/test_tle_cumsum.py::test_cumsum_1d_full_block +third_party/wafer/examples/flagtree/test_tle_cumsum.py::test_cumsum_1d_masked +third_party/wafer/examples/flagtree/test_tle_cumsum.py::test_cumsum_2d_unsupported +third_party/wafer/examples/flagtree/test_tle_cumsum.py::test_cumsum_reverse_unsupported +third_party/wafer/examples/flagtree/test_tle_dsa_arith.py::TestTLEDsaArith::test_arith[shape0-block0-add-] +third_party/wafer/examples/flagtree/test_tle_dsa_arith.py::TestTLEDsaArith::test_arith[shape0-block0-div-] +third_party/wafer/examples/flagtree/test_tle_dsa_arith.py::TestTLEDsaArith::test_arith[shape0-block0-max-] +third_party/wafer/examples/flagtree/test_tle_dsa_arith.py::TestTLEDsaArith::test_arith[shape0-block0-min-] +third_party/wafer/examples/flagtree/test_tle_dsa_arith.py::TestTLEDsaArith::test_arith[shape0-block0-mul-] +third_party/wafer/examples/flagtree/test_tle_dsa_arith.py::TestTLEDsaArith::test_arith[shape0-block0-sub-] +third_party/wafer/examples/flagtree/test_tle_dsa_bridge.py::TestTLEDsaBridge::test_bridge[shape0-block0] +third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py::TestTLEPipelineEndToEnd::test_elementwise_add_basic +third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py::TestTLEPipelineEndToEnd::test_elementwise_add_different_dtypes +third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py::TestTLEPipelineEndToEnd::test_elementwise_add_different_sizes +third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py::TestTLEPipelineEndToEnd::test_elementwise_add_edge_cases +third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py::TestTLEPipelineEndToEnd::test_tle_module_import +third_party/wafer/examples/flagtree/test_tle_dsa_rand.py::test_rand_invalid_n_out +third_party/wafer/examples/flagtree/test_tle_dsa_rand.py::test_rand_uniform_stats +third_party/wafer/examples/flagtree/test_tle_dsa_rand.py::test_randgen_invalid_n_out +third_party/wafer/examples/flagtree/test_tle_dsa_rand.py::test_randgen_shape_and_determinism +third_party/wafer/examples/flagtree/test_tle_dsa_rand.py::test_randn_normal_stats +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestSlice::test_extract_dyn +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestSlice::test_extract_static +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestSlice::test_insert_default +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestSlice::test_insert_strided +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestSlice::test_member +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestTile::test_extract_dyn_multi +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestTile::test_extract_scalar +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestTile::test_insert_dyn_scalar +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestTile::test_insert_multi +third_party/wafer/examples/flagtree/test_tle_dsa_slice.py::TestTile::test_insert_oop +third_party/wafer/examples/test_assert.py::test_assert_scalar[True] +third_party/wafer/examples/test_assert.py::test_assert_tensor[cond_list0] +third_party/wafer/examples/test_assert.py::test_assert_tensor[cond_list3] +third_party/wafer/examples/test_bare_matmul.py::test_bare_matmul[128-dtype1] +third_party/wafer/examples/test_bare_matmul.py::test_bare_matmul[256-dtype2] +third_party/wafer/examples/test_bare_matmul.py::test_bare_matmul[64-dtype0] +third_party/wafer/examples/test_bare_matmul_acc.py::test_bare_matmul_acc[128-dtype1] +third_party/wafer/examples/test_bare_matmul_acc.py::test_bare_matmul_acc[256-dtype2] +third_party/wafer/examples/test_bare_matmul_acc.py::test_bare_matmul_acc[64-dtype0] +third_party/wafer/examples/test_blockptr_complex_offset.py::test +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-128-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-32-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-128-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-32-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[32-64-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-128-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-32-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-128-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-False-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-False-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-False-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e2m1-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e2m1-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e2m1-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e4m3-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e4m3-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e4m3-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e4m3-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e5m2-bf16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e5m2-e4m3-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e5m2-e5m2-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[64-64-64-True-True-True-e5m2-fp16-4-16-1] +third_party/wafer/examples/test_dot_scaled.py::test_scaled_dot[128-64-64-True-False-False-e2m1-e4m3-4-16-1] +third_party/wafer/examples/test_embedding.py::test_embedding[1152-2048-dtype0] +third_party/wafer/examples/test_flip.py::test_flip[float32-8-64] +third_party/wafer/examples/test_fma.py::test_fma[1024-dtype0] +third_party/wafer/examples/test_fp8_conversion.py::test_e5m2_to_fp16_all_encodings +third_party/wafer/examples/test_gather.py::test_gather[src_shape0-indices_shape0-0] +third_party/wafer/examples/test_histogram.py::test_histogram[1024-8] +third_party/wafer/examples/test_libdevice.py::test_libdevice_erf +third_party/wafer/examples/test_libdevice.py::test_libdevice_rename +third_party/wafer/examples/test_libdevice.py::test_special[ceil-ceil-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[ceil-ceil-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[finitef-isfinite-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[finitef-isfinite-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[floor-floor-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[floor-floor-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[fmod-fmod-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[fmod-fmod-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[isinf-isinf-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[isinf-isinf-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[isnan-isnan-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[isnan-isnan-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[pow-pow-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[pow-pow-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[rint-round-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[rint-round-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[tanh-tanh-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[tanh-tanh-4-float32] +third_party/wafer/examples/test_libdevice.py::test_special[trunc-trunc-128-float32] +third_party/wafer/examples/test_libdevice.py::test_special[trunc-trunc-4-float32] +third_party/wafer/examples/test_load_store_mod.py::test[128] +third_party/wafer/examples/test_load_store_mod.py::test[16] +third_party/wafer/examples/test_load_store_mod.py::test[32] +third_party/wafer/examples/test_load_store_mod.py::test[64] +third_party/wafer/examples/test_load_store_mod.py::test[8] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-128-dtype60] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-128-dtype61] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-128-dtype62] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-48-dtype54] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-48-dtype55] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-48-dtype56] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-64-dtype57] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-64-dtype58] +third_party/wafer/examples/test_matmul.py::test_matmul[128-128-64-dtype59] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-128-dtype69] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-128-dtype70] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-128-dtype71] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-48-dtype63] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-48-dtype64] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-48-dtype65] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-64-dtype66] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-64-dtype67] +third_party/wafer/examples/test_matmul.py::test_matmul[128-156-64-dtype68] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-128-dtype78] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-128-dtype79] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-128-dtype80] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-48-dtype72] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-48-dtype73] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-48-dtype74] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-64-dtype75] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-64-dtype76] +third_party/wafer/examples/test_matmul.py::test_matmul[128-512-64-dtype77] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-128-dtype6] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-128-dtype7] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-128-dtype8] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-48-dtype0] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-48-dtype1] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-48-dtype2] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-64-dtype3] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-64-dtype4] +third_party/wafer/examples/test_matmul.py::test_matmul[48-128-64-dtype5] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-128-dtype15] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-128-dtype16] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-128-dtype17] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-48-dtype10] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-48-dtype11] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-48-dtype9] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-64-dtype12] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-64-dtype13] +third_party/wafer/examples/test_matmul.py::test_matmul[48-156-64-dtype14] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-128-dtype24] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-128-dtype25] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-128-dtype26] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-48-dtype18] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-48-dtype19] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-48-dtype20] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-64-dtype21] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-64-dtype22] +third_party/wafer/examples/test_matmul.py::test_matmul[48-512-64-dtype23] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-128-dtype33] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-128-dtype34] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-128-dtype35] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-48-dtype27] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-48-dtype28] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-48-dtype29] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-64-dtype30] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-64-dtype31] +third_party/wafer/examples/test_matmul.py::test_matmul[64-128-64-dtype32] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-128-dtype42] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-128-dtype43] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-128-dtype44] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-48-dtype36] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-48-dtype37] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-48-dtype38] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-64-dtype39] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-64-dtype40] +third_party/wafer/examples/test_matmul.py::test_matmul[64-156-64-dtype41] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-128-dtype51] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-128-dtype52] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-128-dtype53] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-48-dtype45] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-48-dtype46] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-48-dtype47] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-64-dtype48] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-64-dtype49] +third_party/wafer/examples/test_matmul.py::test_matmul[64-512-64-dtype50] +third_party/wafer/examples/test_modulo.py::test_1d +third_party/wafer/examples/test_modulo.py::test_2d +third_party/wafer/examples/test_modulo.py::test_side_by_side_masked_loop +third_party/wafer/examples/test_modulo.py::test_stacked_masked_loop +third_party/wafer/examples/test_modulo.py::test_torch_inductor_pattern +third_party/wafer/examples/test_modulo.py::test_wrap_stacked +third_party/wafer/examples/test_nested_loops.py::test_nested2_complex_body +third_party/wafer/examples/test_nested_loops.py::test_nested2_use_loop_results +third_party/wafer/examples/test_nested_loops.py::test_nested2_use_same_level_loop_result +third_party/wafer/examples/test_nested_loops.py::test_nested3 +third_party/wafer/examples/test_pipeline.py::test_pipeline_gemm[128] +third_party/wafer/examples/test_pipeline.py::test_pipeline_gemm[16] +third_party/wafer/examples/test_pipeline.py::test_pipeline_gemm[32] +third_party/wafer/examples/test_pipeline.py::test_pipeline_gemm[33] +third_party/wafer/examples/test_pipeline.py::test_pipeline_gemm[64] +third_party/wafer/examples/test_pipeline.py::test_pipeline_gemm[96] +third_party/wafer/examples/test_precision_modes.py::test_integer_modes[0-dtype0] +third_party/wafer/examples/test_precision_modes.py::test_integer_modes[1-dtype1] +third_party/wafer/examples/test_precision_modes.py::test_integer_modes[2-dtype2] +third_party/wafer/examples/test_print.py::test_print +third_party/wafer/examples/test_scalar_store.py::test +third_party/wafer/examples/test_sign_extend.py::test_sign_extend +third_party/wafer/examples/test_sort.py::test_sort[float32-False-8-64] +third_party/wafer/examples/test_tensor_index_iterargs.py::test_integer_tensor +third_party/wafer/examples/test_tensor_index_iterargs.py::test_tensor_indices_nested +third_party/wafer/examples/test_tensor_index_iterargs.py::test_tensor_indices_nested_with_mask +third_party/wafer/examples/tle/test_tle_dsa_noc_gemm_4096.py::test_noc_gemm[random-256-256-64] +third_party/wafer/examples/tle/test_tle_dsa_noc_gemm_4096.py::test_noc_gemm[random-4096-4096-1024] +third_party/wafer/examples/tle/test_tle_dsa_noc_gemm_4096.py::test_noc_gemm[structured-256-256-64] +third_party/wafer/examples/tle/test_tle_dsa_noc_gemm_4096.py::test_noc_gemm[structured-4096-4096-1024] diff --git a/test/wafer/test_benchmark_interface.py b/test/wafer/test_benchmark_interface.py new file mode 100644 index 00000000..5f5259fe --- /dev/null +++ b/test/wafer/test_benchmark_interface.py @@ -0,0 +1,31 @@ +import importlib +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + + +def source_driver(wafer_modules): + module = importlib.import_module("_wafer_under_test.driver") + driver = object.__new__(module.DICPDriver) + driver.target = "wafer" + return driver + + +def test_wafer_benchmark_uses_runtime_and_preserves_other_cache_policies(monkeypatch, wafer_modules): + driver = source_driver(wafer_modules) + runtime = SimpleNamespace(Event=Mock(), synchronize=Mock()) + monkeypatch.setattr(wafer_modules[2], "get_runtime", lambda: runtime) + assert driver.get_device_interface() is runtime + driver.clear_cache(driver.get_empty_cache_for_benchmark()) + # DICP shares clear_cache across backends; tensor-based policies still flush. + tensor_cache = SimpleNamespace(zero_=Mock()) + driver.clear_cache(tensor_cache) + tensor_cache.zero_.assert_called_once_with() + + +def test_raw_sdk_runtime_cannot_silently_provide_benchmark_timing(monkeypatch, wafer_modules): + driver = source_driver(wafer_modules) + monkeypatch.setattr(wafer_modules[2], "get_runtime", lambda: SimpleNamespace(current_device=lambda: 0)) + with pytest.raises(RuntimeError, match="torch_txda Event"): + driver.get_device_interface() diff --git a/test/wafer/test_cache_and_runtime.py b/test/wafer/test_cache_and_runtime.py new file mode 100644 index 00000000..54683fc3 --- /dev/null +++ b/test/wafer/test_cache_and_runtime.py @@ -0,0 +1,223 @@ +import os +import subprocess +import sys +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +from triton.backends.compiler import GPUTarget + + +TARGET = GPUTarget("wafer", "wafer", 32) + + +def test_modes_are_snapshotted_and_separate_cache_keys( + monkeypatch, wafer_modules, fake_toolchain +): + _, compiler, _ = wafer_modules + monkeypatch.setenv("USE_SIM_MODE", "0") + monkeypatch.setenv("WAFER_ENABLE_RUNTIME", "0") + offline = compiler.WaferBackend(TARGET) + monkeypatch.setenv("WAFER_ENABLE_RUNTIME", "1") + hardware = compiler.WaferBackend(TARGET) + monkeypatch.setenv("USE_SIM_MODE", "1") + simulator = compiler.WaferBackend(TARGET) + monkeypatch.setenv("WAFER_ENABLE_RUNTIME", "0") + for backend, final in ((offline, "o"), (hardware, "so"), (simulator, "so")): + stages = {} + backend.add_stages(stages, backend.parse_options({})) + assert backend.binary_ext == final + assert list(stages)[-1] == final + assert len({backend.hash() for backend in (offline, hardware, simulator)}) == 3 + + +def test_archive_update_invalidates_both_compiler_and_link_cache( + monkeypatch, wafer_modules, fake_toolchain, tmp_path +): + _, compiler, _ = wafer_modules + _, libraries, _ = fake_toolchain + linker, _ = compiler._runtime_link_inputs() + monkeypatch.setenv("USE_SIM_MODE", "0") + monkeypatch.setenv("WAFER_ENABLE_RUNTIME", "1") + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path / "cache")) + monkeypatch.setenv("WAFER_DEVICE_LOG_ABI", "wafer") + commands = [] + + def link(command): + from pathlib import Path + + commands.append(command) + Path(command[-1]).write_bytes(b"linked " + libraries[3].read_bytes()) + + monkeypatch.setattr(compiler, "_run_tool", link) + before = compiler.WaferBackend(TARGET).hash() + metadata = {} + first = compiler.object_to_binary(b"same object", metadata) + path_before = metadata["kernel_path"] + assert compiler.object_to_binary(b"same object", {}) == first + assert sum(str(command[0]) == str(linker) for command in commands) == 1 + # Even a same-size edit with the old mtime restored must invalidate. + old_stat = libraries[3].stat() + libraries[3].write_bytes(b"X" * old_stat.st_size) + os.utime(libraries[3], ns=(old_stat.st_atime_ns, old_stat.st_mtime_ns)) + second = compiler.object_to_binary(b"same object", metadata) + assert compiler.WaferBackend(TARGET).hash() != before + assert second != first and metadata["kernel_path"] != path_before + assert sum(str(command[0]) == str(linker) for command in commands) == 2 + + +def test_launcher_sdk_and_compiler_changes_invalidate_cache( + wafer_modules, fake_toolchain +): + _, _, runtime = wafer_modules + tools, _, sdk = fake_toolchain + keys = [runtime._launcher_cache_key("same source")] + for path in ( + sdk / "include/tx_runtime.h", + sdk / "lib/libhpgr.so", + tools["clang++"], + ): + path.write_bytes(path.read_bytes() + b" changed") + keys.append(runtime._launcher_cache_key("same source")) + assert len(set(keys)) == 4 + + +def test_device_log_abi_invalidates_cache(monkeypatch, wafer_modules, fake_toolchain): + _, compiler, _ = wafer_modules + monkeypatch.setenv("USE_SIM_MODE", "0") + monkeypatch.setenv("WAFER_ENABLE_RUNTIME", "1") + monkeypatch.setenv("WAFER_DEVICE_LOG_ABI", "wafer") + before = compiler.WaferBackend(TARGET).hash() + monkeypatch.setenv("WAFER_DEVICE_LOG_ABI", "rcs") + assert compiler.WaferBackend(TARGET).hash() != before + + +def test_repeated_compiled_kernel_launch_initializes_once(monkeypatch, wafer_modules): + from triton.compiler import compiler + + _, _, runtime = wafer_modules + launch = Mock() + launcher_cls = Mock(return_value=launch) + driver = SimpleNamespace( + utils=runtime.WaferUtils(), + launcher_cls=launcher_cls, + get_current_device=lambda: 0, + get_current_stream=lambda device: None, + get_current_target=lambda: TARGET, + ) + monkeypatch.setattr(compiler, "driver", SimpleNamespace(active=driver)) + monkeypatch.setattr(compiler, "max_shared_mem", lambda device: 4096) + kernel = compiler.CompiledKernel.__new__(compiler.CompiledKernel) + kernel.module, kernel.function, kernel._run = None, None, None + kernel.src = object() + kernel.name, kernel.kernel = "test_kernel", b"ELF lifetime token" + kernel.metadata = SimpleNamespace(shared=0, num_warps=1) + kernel.packed_metadata = () + kernel.metadata_group, kernel.hash = {}, "unit-test" + kernel[(1, 1, 1)](123) + kernel[(1, 1, 1)](456) + kernel.launch_metadata((1, 1, 1), None) + assert kernel.module is kernel.kernel + assert launcher_cls.call_count == 1 + assert launch.call_count == 2 + + +def test_jit_and_compiled_kernel_argument_contracts(monkeypatch, wafer_modules): + _, _, runtime = wafer_modules + launch = Mock() + monkeypatch.setattr( + runtime, "compile_launcher", lambda source: SimpleNamespace(launch=launch) + ) + src = SimpleNamespace( + fn=SimpleNamespace(arg_names=["pointer", "alpha", "BLOCK"]), + signature={"BLOCK": "constexpr", "alpha": "fp32", "pointer": "*fp32"}, + constants={(2,): 256}, + ) + metadata = object() + launcher = runtime.WaferLauncher(src, metadata) + prefix = (1, 1, 1, None, 0, (), None, None, None) + launcher(*prefix, 123, 1.25, 256) + launcher(*prefix, 123, 1.25) + expected = (*prefix[:5], metadata, *prefix[6:], 123, 1.25) + assert launch.call_args_list[0].args == expected + assert launch.call_args_list[1].args == expected + with pytest.raises(TypeError, match="expected 2 runtime arguments"): + launcher(*prefix, 123) + + +def test_linker_library_lookup(monkeypatch, wafer_modules, tmp_path): + _, compiler, _ = wafer_modules + library = tmp_path / "libc.a" + library.write_bytes(b"archive") + query = Mock(return_value=str(library) + "\n") + monkeypatch.setattr(compiler.subprocess, "check_output", query) + assert compiler._find_linker_library("gcc", "libc.a") == library + for output in ("libc.a", str(tmp_path / "missing.a")): + query.return_value = output + with pytest.raises(RuntimeError, match="could not locate"): + compiler._find_linker_library("gcc", "libc.a") + query.side_effect = subprocess.CalledProcessError(1, "gcc") + with pytest.raises(subprocess.CalledProcessError): + compiler._find_linker_library("gcc", "libc.a") + + +def test_half_scalars_are_not_silently_packed_as_fp32(wafer_modules): + _, _, runtime = wafer_modules + for scalar in ("fp16", "bf16"): + with pytest.raises(NotImplementedError, match="pass an fp32 scalar"): + runtime.make_launcher({0: scalar}) + assert "get_pointer" in runtime.make_launcher({0: "*" + scalar}) + + +def test_runtime_selection_does_not_hide_native_abi_errors(monkeypatch, wafer_modules): + import builtins + + _, _, runtime = wafer_modules + mock_wafer_runtime = object() + fallback = object() + monkeypatch.setitem(sys.modules, "torch", SimpleNamespace(txda=mock_wafer_runtime)) + monkeypatch.setitem(sys.modules, "torch_txda", SimpleNamespace()) + monkeypatch.setattr(runtime, "_KuiperRuntime", lambda: fallback) + assert runtime.get_runtime() is mock_wafer_runtime + monkeypatch.setitem(sys.modules, "torch_txda", None) + assert runtime.get_runtime() is fallback + original_import = builtins.__import__ + + def broken_import(name, *args, **kwargs): + if name == "torch_txda": + raise OSError("undefined symbol: required_runtime_api") + return original_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", broken_import) + with pytest.raises(OSError, match="required_runtime_api"): + runtime.get_runtime() + + +def test_rcs_adaptation_uses_cached_private_copies( + monkeypatch, wafer_modules, fake_toolchain, tmp_path +): + _, compiler, _ = wafer_modules + linker, libraries = compiler._runtime_link_inputs() + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path / "cache")) + originals = [path.read_bytes() for path in libraries] + commands = [] + + def adapt(command): + from pathlib import Path + + commands.append(command) + Path(command[-1]).write_bytes(Path(command[-2]).read_bytes() + b" adapted") + + monkeypatch.setattr(compiler, "_run_tool", adapt) + fingerprint = compiler._link_fingerprint(linker, libraries, "rcs") + first = compiler._adapt_logging_libraries(libraries, fingerprint) + second = compiler._adapt_logging_libraries(libraries, fingerprint) + assert first == second and len(commands) == 4 + assert [path.read_bytes() for path in libraries] == originals + assert all(a != b for a, b in zip(first[:4], libraries[:4])) + assert first[4:] == libraries[4:] + assert all( + "--redefine-sym=tx8_kernel_printf=rcs_kernel_printf" in command + for command in commands + ) diff --git a/test/wafer/test_cluster_launch.py b/test/wafer/test_cluster_launch.py new file mode 100644 index 00000000..4b900c78 --- /dev/null +++ b/test/wafer/test_cluster_launch.py @@ -0,0 +1,52 @@ +"""Check generated native launch dispatch and arguments without device work.""" +import os +from types import SimpleNamespace + +import pytest + + +@pytest.mark.parametrize("mode,marker", [("simt", "0x71"), ("cluster", "0x72")]) +def test_native_launch_mode(mode, marker, wafer_modules, tmp_path, monkeypatch): + if not os.getenv("KUIPER_ROOT"): + pytest.skip("Native launcher test requires Kuiper headers") + _, _, runtime = wafer_modules + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path / "cache")) + stubs = r''' +static txError_t test_simt(const char *name, uint64_t elf, uint64_t len, + dim3 grid, dim3 block, void *args, uint32_t arglen, uint32_t shared, txStream_t stream) { + if (strcmp(name, "test") || len != 4 || !elf || grid.x != 16 || grid.y != 1 || grid.z != 1 || + block.x != 1 || block.y != 1 || block.z != 1 || shared != 0 || stream != nullptr || + arglen != 7 * sizeof(uint64_t) || ((uint64_t *)args)[0] != 42 || ((uint64_t *)args)[1] != 16) + return (txError_t)0x73; + return (txError_t)0x71; +} +static txError_t test_cluster(const char *name, uint64_t elf, uint64_t len, dim3 cluster, + dim3 grid, dim3 block, void *args, uint32_t arglen, uint32_t shared, txStream_t stream) { + if (cluster.x != 1 || cluster.y != 1 || cluster.z != 1) return (txError_t)0x74; + txError_t result = test_simt(name, elf, len, grid, block, args, arglen, shared, stream); + return result == (txError_t)0x71 ? (txError_t)0x72 : result; +} +#define txLaunchKernelGGL test_simt +#define txLaunchClusterKernelGGL test_cluster +''' + source = runtime.make_launcher({0: "i32"}, mode) + source = source.replace('#include "tx_runtime.h"', '#include "tx_runtime.h"\n' + stubs) + launch = runtime.compile_launcher(source).launch + binary = tmp_path / "dummy.so" + binary.write_bytes(b"test") + metadata = SimpleNamespace(kernel_path=str(binary), name="test") + with pytest.raises(RuntimeError, match=marker): + launch(16, 1, 1, None, 0, metadata, None, None, None, 42) + if mode == "cluster": + for grid in [(17, 1, 1), (8, 2, 1), (8, 1, 2)]: + with pytest.raises(ValueError, match="cluster grid"): + launch(*grid, None, 0, metadata, None, None, None, 42) + + +def test_cluster_option_validation_and_cache(wafer_modules): + _, compiler, _ = wafer_modules + assert compiler.WaferOptions().hash() != compiler.WaferOptions(launch_mode="cluster").hash() + with pytest.raises(ValueError, match="launch_mode"): + compiler.WaferOptions(launch_mode="invalid") + with pytest.raises(ValueError, match="one cluster"): + compiler.WaferOptions(launch_mode="cluster", cluster_dims=(2, 1, 1)) diff --git a/test/wafer/test_commonir_abi.py b/test/wafer/test_commonir_abi.py new file mode 100644 index 00000000..248e7f4c --- /dev/null +++ b/test/wafer/test_commonir_abi.py @@ -0,0 +1,51 @@ +"""Keep the shared CommonIR loader independent of Wafer's five-value ABI.""" + +import importlib.util +from pathlib import Path +import sys +from types import ModuleType, SimpleNamespace +from unittest.mock import Mock + +import pytest + + +@pytest.mark.parametrize( + "target,mix_mode", [("mlu", False), ("maca", False), ("ascend", True)] +) +def test_commonir_legacy_loader_contract(monkeypatch, target, mix_mode): + calls = [] + handle = object() + + def load_binary(name, binary, shared, device, *, mix_mode): + calls.append((name, binary, shared, device, mix_mode)) + return handle, 123, 4, 0 + + driver = SimpleNamespace( + get_current_device=lambda: 0, + launcher_cls=Mock(return_value=object()), + utils=SimpleNamespace(load_binary=load_binary), + ) + # Import the real loader while isolating the unavailable vendor drivers. + package = ModuleType("_wafer_commonir_abi_test") + package.__path__ = [] + dependency = ModuleType(package.__name__ + ".backend") + dependency.commonir_backend = SimpleNamespace(get_driver=lambda: driver) + monkeypatch.setitem(sys.modules, package.__name__, package) + monkeypatch.setitem(sys.modules, dependency.__name__, dependency) + path = Path(__file__).resolve().parents[2] / "backend/commonir/compiler.py" + spec = importlib.util.spec_from_file_location(package.__name__ + ".compiler", path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + + kernel = module.CompiledKernel.__new__(module.CompiledKernel) + kernel.module = None + kernel.name = target + "_kernel" + kernel.kernel = b"device binary" + kernel.src = object() + kernel.metadata = SimpleNamespace(shared=32, mix_mode=mix_mode) + kernel._init_handles() + kernel.launch_metadata((1, 1, 1), None) + assert calls == [(kernel.name, kernel.kernel, 32, 0, mix_mode)] + assert kernel.module is handle + assert kernel.function == 123 + driver.launcher_cls.assert_called_once_with(kernel.src, kernel.metadata) diff --git a/test/wafer/test_crt_diagnostics.py b/test/wafer/test_crt_diagnostics.py new file mode 100644 index 00000000..041476bb --- /dev/null +++ b/test/wafer/test_crt_diagnostics.py @@ -0,0 +1,94 @@ +"""Check the diagnostic ABI without triggering a fatal assertion on the card.""" +import os +from pathlib import Path +import shutil +import subprocess + +import pytest + + +def test_assert_reports_and_terminates_with_ndebug(tmp_path): + compiler = shutil.which("cc") + if not compiler: + pytest.skip("C compiler required") + root = Path(__file__).resolve().parents[2] + (tmp_path / "wafer.h").write_text(""" + #define INTRNISIC_RUN_SWITCH + #define KCORE_LOG_ERROR 3 + void tsm_ep_log(const char *, const char *, unsigned, unsigned, const char *, ...); + void __assert_func(const char *, int, const char *, const char *) __attribute__((noreturn)); + """) + source = tmp_path / "check.c" + source.write_text(r''' + #include + #include + #include + #include + static int reported; + void tsm_ep_log(const char *file, const char *func, unsigned line, + unsigned level, const char *format, ...) { + va_list args; + char text[256]; + va_start(args, format); + vsnprintf(text, sizeof(text), format, args); + va_end(args); + if (level != 3 || strcmp(text, + "kernel.py(line 17, col 9)::tile (2, 3, 4): bad value\n")) exit(21); + reported = 1; + } + void __assert_func(const char *file, int line, const char *func, const char *expr) { + if (!reported || strcmp(file, "kernel.py") || line != 17 || + strcmp(func, "__Assert") || strcmp(expr, "bad value")) exit(22); + exit(0); + } + void __Assert(const char *, ...); + int main(void) { + __Assert("bad value", "kernel.py", 17, 9, 2, 3, 4); + return 23; + } + ''') + exe = tmp_path / "check" + subprocess.run([ + compiler, "-O2", "-DNDEBUG", "-Werror=implicit-function-declaration", + "-I", str(tmp_path), str(root / "third_party/wafer/crt/lib/Wafer/assert.c"), + str(source), "-o", str(exe), + ], check=True) + subprocess.run([str(exe)], check=True, timeout=10) + + +@pytest.mark.parametrize("with_sdk_init", [False, True]) +def test_assert_links_to_firmware_without_posix_imports(wafer_modules, monkeypatch, tmp_path, with_sdk_init): + if not os.getenv("WAFER_DEPS_ROOT"): + pytest.skip("Wafer SDK required for RISC-V link test") + _, compiler, _ = wafer_modules + linker, libraries = compiler._runtime_link_inputs() + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path / "cache")) + source = tmp_path / "probe.c" + # module_init pulls the real intrinsic/common archives and their libc + # dependencies. An assertion-only object cannot expose missing libgloss. + source.write_text('''#include + extern void module_init(void *); + void probe(void *args) { + ''' + ('module_init(args);' if with_sdk_init else '') + ''' + __assert_func("probe", 1, "probe", "false"); + } + ''') + obj = tmp_path / "probe.o" + subprocess.run([str(linker), "-fPIC", "-O2", "-march=rv64imfdc", "-mabi=lp64d", + "-c", str(source), "-o", str(obj)], check=True) + original_libc = next(p for p in libraries if p.name == "libc.a") + before = compiler.file_fingerprint(original_libc) + metadata = {} + compiler.object_to_binary(obj.read_bytes(), metadata, simulator=False, log_abi="rcs") + nm = compiler._find_llvm_tool("llvm-nm") + undefined = subprocess.check_output([str(nm), "--undefined-only", "--format=posix", + metadata["kernel_path"]], text=True) + imports = {line.split()[0] for line in undefined.splitlines()} + assert "__assert_func" in imports + allowed = {"__assert_func", "__get_pid", "get_log_level", "monitor_write_log", + "rcs_ep_log", "rcs_kernel_printf", "rcs_kernel_vprintf", + "rt_free", "rt_malloc", "rt_thread_mdelay"} + assert imports <= allowed, imports - allowed + if not with_sdk_init: + assert imports == {"__assert_func"} + assert compiler.file_fingerprint(original_libc) == before diff --git a/test/wafer/test_crt_fp8.py b/test/wafer/test_crt_fp8.py new file mode 100644 index 00000000..164922af --- /dev/null +++ b/test/wafer/test_crt_fp8.py @@ -0,0 +1,47 @@ +"""Run the production C decoder with host-only SPM address translation.""" +import ctypes +import math +from pathlib import Path +import shutil +import struct +import subprocess + +import pytest + + +def test_e5m2_to_fp16_all_encodings(tmp_path): + compiler = shutil.which("cc") + if compiler is None: + pytest.skip("A host C compiler is required for the CRT decoder test") + source = Path(__file__).resolve().parents[2] / "third_party/wafer/crt/lib/Wafer/mxfp_fp16.c" + text = source.read_text() + start = text.index("void __FP8E5M2_FP16(") + function = text[start:text.index("\n/**", start)] + stub = "#include \nstatic uint64_t get_spm_memory_mapping_wrapper(uint64_t p) { return p; }\n" + host_source = tmp_path / "decoder.c" + host_source.write_text(stub + function) + library = tmp_path / "decoder.so" + subprocess.run([compiler, "-O2", "-shared", "-fPIC", str(host_source), "-o", str(library)], check=True) + decode = ctypes.CDLL(str(library)).__FP8E5M2_FP16 + decode.argtypes = [ctypes.POINTER(ctypes.c_uint8), ctypes.POINTER(ctypes.c_uint16), ctypes.c_uint32] + decode.restype = None + src = (ctypes.c_uint8 * 256)(*range(256)) + # Sentinels check that elem_count bounds the destination writes. + dst = (ctypes.c_uint16 * 258)(*([0x1234] * 258)) + decode(src, dst, 256) + assert list(dst)[256:] == [0x1234, 0x1234] + for code, bits in enumerate(list(dst)[:256]): + sign, exponent, fraction = code >> 7, (code >> 2) & 31, code & 3 + actual = struct.unpack("> 15 == sign, hex(code) + if exponent == 31: + assert math.isnan(actual) if fraction else math.isinf(actual) + assert bits & 1023 == fraction * 256, hex(code) + else: + magnitude = math.ldexp(fraction / 4, -14) if exponent == 0 else math.ldexp(1 + fraction / 4, exponent - 15) + expected = -magnitude if sign else magnitude + expected_bits = struct.unpack(" +#include +#define INTRNISIC_RUN_SWITCH +#define SYNCHRONOUS_INTRINSIC_SWITCH +using Data_Format = int; +constexpr int I_CGRA = 0; +struct TsmPeripheralInstr { int tag; int a[1]; int b[1]; }; +struct TsmPeripheral { + void RandGen(TsmPeripheralInstr*, uint64_t a, uint64_t b, uint64_t c, + uint64_t d, uint64_t e, uint32_t bytes, Data_Format fmt) { + assert(a == 0x1000 && b == 0x2000 && c == 0x3000 && d == 0x4000 && e == 0x5000); + assert(bytes == 256 && fmt == 11); + } +}; +struct Intrinsic { TsmPeripheral* peripheral_pointer; }; +Intrinsic* g_intrinsic() { static TsmPeripheral p; static Intrinsic i{&p}; return &i; } +void TsmExecute(TsmPeripheralInstr*) {} +''' + path = tmp_path / "rand.cpp" + path.write_text(source + implementation + ''' +int main() { __RandGen((uint64_t*)0x1000, (uint64_t*)0x2000, (uint64_t*)0x3000, + (uint64_t*)0x4000, (uint64_t*)0x5000, 256, 11); } +''') + exe = tmp_path / "rand" + subprocess.run([compiler, "-std=c++17", str(path), "-o", str(exe)], check=True) + subprocess.run([str(exe)], check=True, timeout=10) diff --git a/test/wafer/test_dense_constants.py b/test/wafer/test_dense_constants.py new file mode 100644 index 00000000..d56cbd49 --- /dev/null +++ b/test/wafer/test_dense_constants.py @@ -0,0 +1,41 @@ +"""Shape vectors and nonuniform tensors must lower and bufferize on Wafer.""" +import os +import subprocess + +import pytest + + +@pytest.mark.parametrize("dtype,shape,values", [ + ("i64", "2", "[16, 512]"), + ("i32", "2x3", "[[1, 2, 3], [-4, 5, 6]]"), + ("f32", "2x2", "[[1.25, -2.5], [3.0, 0.0]]"), + ("index", "2", "[16, 512]"), + ("index", "2", "16"), +]) +def test_non_splat_constant_bufferization(dtype, shape, values, tmp_path): + from triton.backends.dicp_triton.wafer import _find_wafer_opt + + ty = f"tensor<{shape}x{dtype}>" + indices = [f"%i{axis}" for axis in range(len(shape.split("x")))] + args = ", ".join(f"{index}: index" for index in indices) + source = tmp_path / "constant.mlir" + source.write_text(f"""module {{ + func.func @constant({args}) -> {dtype} {{ + %value = arith.constant dense<{values}> : {ty} + %element = tensor.extract %value[{', '.join(indices)}] : {ty} + return %element : {dtype} + }} + }}""") + result = subprocess.run([ + os.getenv("WAFER_TEST_OPT") or str(_find_wafer_opt()), str(source), + "--linalg-to-mk=precision-mode=2", + "--one-shot-bufferize=bufferize-function-boundaries", + "--convert-bufferization-to-memref", + ], capture_output=True, text=True, timeout=30) + assert result.returncode == 0, result.stderr + assert "arith.constant dense<" not in result.stdout + if values.startswith("["): + assert f"memref<{shape}x{dtype}>" in result.stdout + else: + # A dynamic lookup into a splat may fold to the scalar itself. + assert f"arith.constant {values} : {dtype}" in result.stdout diff --git a/test/wafer/test_device_print_abi.py b/test/wafer/test_device_print_abi.py new file mode 100644 index 00000000..f3f8ec53 --- /dev/null +++ b/test/wafer/test_device_print_abi.py @@ -0,0 +1,35 @@ +"""Compile device printing without loading a kernel on the card.""" +import re + +import pytest +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget +from triton.compiler import ASTSource + + +@triton.jit +def print_vector(X, HEX: tl.constexpr): + value = tl.load(X + tl.arange(0, 8)) + tl.device_print("abi", value, hex=HEX) + + +@triton.jit +def print_scalar(X, HEX: tl.constexpr): + tl.device_print("abi", tl.load(X), hex=HEX) + + +@pytest.mark.parametrize("kernel", [print_scalar, print_vector]) +@pytest.mark.parametrize("dtype,hex_mode,expected", [ + ("fp16", False, r"fpext half .* to double"), + ("bf16", False, r"fpext bfloat .* to double"), + ("i16", False, r"sext i16 .* to i32"), + ("u16", False, r"zext i16 .* to i32"), + ("fp16", True, r"zext i16 .* to i32"), +]) +def test_printf_uses_c_vararg_promotions(kernel, dtype, hex_mode, expected): + source = ASTSource(kernel, signature={"X": "*" + dtype, "HEX": "constexpr"}, + constexprs={"HEX": hex_mode}) + result = triton.compile(source, target=GPUTarget("wafer", "wafer", 32)) + assert re.search(expected, result.asm["llir"]), result.asm["llir"] + diff --git a/test/wafer/test_elf_audit.py b/test/wafer/test_elf_audit.py new file mode 100644 index 00000000..639fa734 --- /dev/null +++ b/test/wafer/test_elf_audit.py @@ -0,0 +1,71 @@ +"""NoC references must not authorize missing exports or unrelated imports.""" +import importlib.util +from pathlib import Path +import struct + +import pytest + + +@pytest.fixture +def audit(tmp_path, monkeypatch): + source = Path(__file__).resolve().parents[2] / "scripts/wafer/audit_wafer_elf.py" + spec = importlib.util.spec_from_file_location("wafer_elf_audit", source) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + header = bytearray(64) + header[:6] = b"\x7fELF\x02\x01" + struct.pack_into(" 0 diff --git a/test/wafer/test_frontend_isolation.py b/test/wafer/test_frontend_isolation.py new file mode 100644 index 00000000..44c03571 --- /dev/null +++ b/test/wafer/test_frontend_isolation.py @@ -0,0 +1,180 @@ +"""Run against the isolated wheel; no vendor runtime or device is required.""" + +import os +from pathlib import Path +import subprocess +import sys + +import pytest +import triton +import triton.language as tl +from triton._C.libtriton import ir, wafer +from triton.backends.compiler import GPUTarget +from triton.backends.dicp_triton.wafer import WaferBackend, _find_wafer_opt +from triton.compiler import ASTSource + + +pytestmark = pytest.mark.skipif( + getattr(wafer, "build_role", None) != "frontend", + reason="requires the Wafer-only frontend wheel", +) + + +@triton.jit +def loop_kernel(out): + i = tl.arange(0, 16) + x = i.to(tl.float32) + for _ in range(2): + x = x + 1.0 + tl.store(out + i, x) + + +@triton.jit +def non_power_of_two(out): + i = tl.arange(0, 3) + tl.store(out + i, i.to(tl.float32)) + + +@triton.jit +def scalar_copy_kernel(src, out): + tl.store(out, tl.load(src)) + + +@triton.jit +def member_slice_kernel(out): + x = tl.full((16,), 1, tl.float32) + sub = x.extract_slice(offsets=(0,), sizes=(8,), strides=(1,)) + y = x.insert_slice(sub + 1, offsets=(8,)) + tl.store(out + tl.arange(0, 16), y) + + +@triton.jit +def bounded_slice_kernel(out): + x = tl.arange(0, 16).to(tl.float32) + left = x[:8] + right = x[8:] + tl.store(out + tl.arange(0, 8)[:, None], (left + right)[:, None]) + + +def make_module(fn, signature=None, constexprs=None): + backend = WaferBackend(GPUTarget("wafer", "wafer", 32)) + options = backend.parse_options({}) + context = ir.context() + ir.load_dialects(context) + backend.load_dialects(context) + codegen = backend.get_codegen_implementation(options) + return ASTSource(fn, signature=signature or {"out": "*fp32"}, constexprs=constexprs).make_ir( + backend.target, options, codegen, {}, context + ) + + +def test_package_has_no_original_dicp_or_cann(): + from triton._C import libtriton + + assert not hasattr(libtriton, "dicp_triton") + root = Path(triton.__file__).parent + assert not (root / "language/extra/deeplink").exists() + assert not (root / "backends/dicp_triton/dicp_opt").exists() + text = str(make_module(scalar_copy_kernel, {"src": "*fp32", "out": "*fp32"})) + assert "dicp.disable_addptr_fold" not in text + assert "tt.load" in text and "tt.store" in text + + +@pytest.mark.parametrize("target", ["wafer", "wafer-cache-before-tle", "wafer-cache-after-tle", "wafer-slices"]) +def test_codegen_in_separate_processes(target): + result = subprocess.run( + [sys.executable, str(Path(__file__).resolve()), target], + capture_output=True, text=True, timeout=60, + ) + assert result.returncode == 0, result.stdout + result.stderr + + +def test_cache_tracks_vendor_sources_without_importing_them(tmp_path, monkeypatch): + from triton.runtime import cache + import sysconfig + + root = tmp_path / "triton" + # A package initializer must never run just to compute a cache key. Both + # package and nested source edits must nevertheless invalidate that key. + source = root / "language/extra/vendor_probe/ops.py" + files = { + "runtime/cache.py": "# cache input\n", + "_C/libtriton." + sysconfig.get_config_var("EXT_SUFFIX").split(".")[-1]: "binary input", + "language/extra/__init__.py": "raise RuntimeError('unexpected import')\n", + "language/extra/vendor_probe/__init__.py": "raise RuntimeError('unexpected import')\n", + "language/extra/vendor_probe/ops.py": "VALUE = 1\n", + } + for name, text in files.items(): + path = root / name + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(text) + monkeypatch.setattr(cache, "__file__", str(root / "runtime/cache.py")) + initial = cache.triton_key.__wrapped__() + source.write_text("VALUE = 2\n") + changed_source = cache.triton_key.__wrapped__() + source.with_name("__init__.py").write_text("raise RuntimeError('still must not run')\n") + assert len({initial, changed_source, cache.triton_key.__wrapped__()}) == 3 + + +def test_frontend_text_can_enter_wafer_lowering(tmp_path): + source = tmp_path / "frontend.mlir" + source.write_text(str(make_module(loop_kernel))) + result = subprocess.run( + [os.getenv("WAFER_TEST_OPT") or str(_find_wafer_opt()), str(source), + "--triton-to-core-dialects", "-o", str(tmp_path / "core.mlir")], + capture_output=True, text=True, timeout=60, + ) + assert result.returncode == 0, result.stderr + assert "linalg" in (tmp_path / "core.mlir").read_text() + + +def test_tle_text_can_enter_separate_tools(tmp_path): + from test_tle_frontend import local_kernel + + source = tmp_path / "dsa.mlir" + source.write_text(str(make_module(local_kernel))) + assert "dsa.alloc" in source.read_text() + result = subprocess.run( + [os.getenv("WAFER_TEST_OPT") or str(_find_wafer_opt()), str(source), + "--triton-to-core-dialects", "--tle-to-mk", "--dsa-memory-to-core", + "-o", str(tmp_path / "core.mlir")], capture_output=True, text=True, timeout=60, + ) + assert result.returncode == 0, result.stderr + assert "dsa.alloc" not in (tmp_path / "core.mlir").read_text() + + +if __name__ == "__main__": + assert wafer.build_role == "frontend" + if sys.argv[1] == "wafer-slices": + import triton.language.extra.wafer.slicing + + text = str(make_module(bounded_slice_kernel)) + assert text.count('"dsa.extract_slice"') == 2 + assert "tt.expand_dims" in text + assert not any(".deeplink.cann" in name for name in sys.modules) + sys.exit(0) + if sys.argv[1].startswith("wafer-cache-"): + from triton.runtime.cache import triton_key + + if sys.argv[1] == "wafer-cache-before-tle": + triton_key() + import triton.experimental.tle.language # registers Wafer tensor members + members = (tl.tensor.extract_slice, tl.tensor.insert_slice, tl.tensor.__getitem__) + original_tanh = getattr(tl.math, "tanh", None) + triton_key() + assert not any(".deeplink.cann" in name for name in sys.modules) + assert members == (tl.tensor.extract_slice, tl.tensor.insert_slice, tl.tensor.__getitem__) + assert getattr(tl.math, "tanh", None) is original_tanh + text = str(make_module(member_slice_kernel)) + assert '"dsa.extract_slice"' in text and '"dsa.insert_slice"' in text + sys.exit(0) + from triton.runtime.cache import triton_key + triton_key() + original_tanh = getattr(tl.math, "tanh", None) + module = str(make_module(loop_kernel)) + assert "scf.for" in module and "tt.store" in module + assert "dicp.disable_addptr_fold" not in module + assert not any(".deeplink.cann" in name for name in sys.modules) + assert getattr(tl.math, "tanh", None) is original_tanh + with pytest.raises(triton.compiler.errors.CompilationError, match="power of 2"): + make_module(non_power_of_two) diff --git a/test/wafer/test_loader_isolation.py b/test/wafer/test_loader_isolation.py new file mode 100644 index 00000000..e6dcd64b --- /dev/null +++ b/test/wafer/test_loader_isolation.py @@ -0,0 +1,58 @@ +"""Exercise upstream Triton's loader call without loading a device program.""" + +from types import SimpleNamespace + +import pytest +from triton.compiler import compiler + + +def make_kernel(monkeypatch, load): + launcher = object() + driver = SimpleNamespace( + utils=SimpleNamespace(load_binary=load), + get_current_device=lambda: 3, + get_current_target=lambda: SimpleNamespace(warp_size=32), + launcher_cls=lambda src, metadata: launcher, + ) + monkeypatch.setattr(compiler, "driver", SimpleNamespace(active=driver)) + monkeypatch.setattr(compiler, "max_shared_mem", lambda device: 1024) + kernel = object.__new__(compiler.CompiledKernel) + kernel.module = None + kernel.function = None + kernel.src = object() + kernel.name = "wafer_entry" + kernel.kernel = b"ELF" + kernel.metadata = SimpleNamespace(shared=64, num_warps=1) + kernel.metadata_group = {} + kernel.hash = "test-hash" + return kernel, launcher + + +def test_upstream_triton_uses_native_five_result_loader(monkeypatch): + calls = [] + native = (object(), object(), 7, 2, 1024) + + def load(*args): + calls.append(args) + return native + + kernel, launcher = make_kernel(monkeypatch, load) + kernel._init_handles() + assert calls == [("wafer_entry", b"ELF", 64, 3)] + assert (kernel.module, kernel.function, kernel.n_regs, kernel.n_spills, kernel.n_max_threads) == native + assert kernel._run is launcher + kernel._init_handles() + assert len(calls) == 1 + + +def test_loader_error_is_not_retried_with_another_signature(monkeypatch): + calls = [] + + def load(*args): + calls.append(args) + raise TypeError("invalid binary in vendor loader") + + kernel, _ = make_kernel(monkeypatch, load) + with pytest.raises(TypeError, match="invalid binary"): + kernel._init_handles() + assert len(calls) == 1 diff --git a/test/wafer/test_log1p_lowering.py b/test/wafer/test_log1p_lowering.py new file mode 100644 index 00000000..65dc026e --- /dev/null +++ b/test/wafer/test_log1p_lowering.py @@ -0,0 +1,21 @@ +"""Wafer must retain stable scalar libm log1p calls through its actual pipeline.""" +import pytest + + +@pytest.mark.parametrize("dtype,symbol", [("f32", "log1pf"), ("f64", "log1p")]) +def test_stable_log1p_lowering(dtype, symbol, wafer_modules, monkeypatch): + from triton.backends.dicp_triton.wafer import _find_wafer_opt + + _, compiler, _ = wafer_modules + # Exercise the source pipeline with the installed, versioned compiler tool. + monkeypatch.setattr(compiler, "_find_wafer_opt", _find_wafer_opt) + ir = f"""module {{ + func.func @stable_log1p(%x: {dtype}) -> {dtype} {{ + %y = math.log1p %x : {dtype} + return %y : {dtype} + }} + }}""" + llvm = compiler.wafer_ir_to_llir(ir, {}) + assert f"@{symbol}(" in llvm + assert "@llvm.log." not in llvm + assert "fadd" not in llvm diff --git a/test/wafer/test_module_runtime.py b/test/wafer/test_module_runtime.py new file mode 100644 index 00000000..6aa1c7a1 --- /dev/null +++ b/test/wafer/test_module_runtime.py @@ -0,0 +1,68 @@ +"""SDK module ownership, native dispatch and GIL behavior without a device.""" +import ctypes +import os +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + + +def test_module_owner_cleanup(wafer_modules, monkeypatch): + _, _, runtime = wafer_modules + state = {"device": 3} + def load(pointer, binary, size): + assert state["device"] == 0 and size == 3 + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_void_p)).contents.value = 123 + return 0 + def get_function(pointer, module, name): + assert module.value == 123 and name == b"kernel" + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_void_p)).contents.value = 456 + return 0 + def unload(module): + assert state["device"] == 0 and module.value == 123 + return 0 + library = SimpleNamespace(txModuleLoad=Mock(side_effect=load), + txModuleGetFunction=Mock(side_effect=get_function), + txModuleUnload=Mock(side_effect=unload)) + sdk = SimpleNamespace(library=library, current_device=lambda: state["device"], + set_device=lambda device: state.update(device=device)) + monkeypatch.setattr(runtime, "_KuiperRuntime", lambda: sdk) + owner = runtime._LoadedModule("kernel", b"ELF", 0) + assert owner.function == 456 and state["device"] == 3 + owner.close() + owner.close() + assert library.txModuleUnload.call_count == 1 and state["device"] == 3 + library.txModuleGetFunction.side_effect = lambda *args: 0x42 + with pytest.raises(RuntimeError, match="txModuleGetFunction.*42"): + runtime._LoadedModule("kernel", b"ELF", 0) + assert library.txModuleUnload.call_count == 2 and state["device"] == 3 + + +@pytest.mark.parametrize("mode", ["simt", "cluster"]) +def test_native_module_dispatch_releases_gil(mode, wafer_modules, tmp_path, monkeypatch): + if not os.getenv("KUIPER_ROOT"): + pytest.skip("Kuiper headers required") + _, _, runtime = wafer_modules + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path / "cache")) + stubs = r''' +static txError_t test_module(txFunction_t function, dim3 grid, dim3 block, + void *args, uint32_t length, uint32_t shared, txStream_t stream) { + if (PyGILState_Check() || function != (txFunction_t)456 || grid.x != 4 || + length != 7 * sizeof(uint64_t) || ((uint64_t*)args)[0] != 42) + return (txError_t)0x76; + return (txError_t)0x75; +} +static txError_t test_cluster_module(txFunction_t function, dim3 cluster, + dim3 grid, dim3 block, void *args, uint32_t length, uint32_t shared, txStream_t stream) { + if (cluster.x != 1 || cluster.y != 1 || cluster.z != 1) return (txError_t)0x77; + return test_module(function, grid, block, args, length, shared, stream); +} +#define txLaunchKernel test_module +#define txLaunchClusterKernel test_cluster_module +''' + source = runtime.make_launcher({0: "i32"}, mode).replace( + '#include "tx_runtime.h"', '#include "tx_runtime.h"\n' + stubs) + launch = runtime.compile_launcher(source).launch + metadata = SimpleNamespace(kernel_path=str(tmp_path / "not-read.so"), name="test") + with pytest.raises(RuntimeError, match="0x75"): + launch(4, 1, 1, None, 456, metadata, None, None, None, 42) diff --git a/test/wafer/test_mxfp_reference.py b/test/wafer/test_mxfp_reference.py new file mode 100644 index 00000000..f460dd24 --- /dev/null +++ b/test/wafer/test_mxfp_reference.py @@ -0,0 +1,40 @@ +"""Check the independent scaled-dot oracle against fixed format encodings.""" + +import importlib.util +from pathlib import Path + +import pytest + +torch = pytest.importorskip("torch") +path = Path(__file__).resolve().parents[2] / "third_party/wafer/examples/_wafer_reference.py" +spec = importlib.util.spec_from_file_location("wafer_mxfp_reference", path) +reference = importlib.util.module_from_spec(spec) +spec.loader.exec_module(reference) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +def test_mxfp_known_encodings_and_scales(dtype): + # Both nibble order and all signed FP4 values, including signed zero. + packed = torch.tensor([[0x10, 0x32, 0x54, 0x76, 0x98, 0xBA, 0xDC, 0xFE] * 2], dtype=torch.uint8) + expected = torch.tensor([[0, .5, 1, 1.5, 2, 3, 4, 6, -0., -.5, -1, -1.5, -2, -3, -4, -6] * 2], dtype=dtype) + one = torch.tensor([[127]], dtype=torch.uint8) + actual = reference.upcast_mxfp_cpu(packed, one, "e2m1", dtype) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + assert torch.equal(torch.signbit(actual), torch.signbit(expected)) + formats = [ + ("e4m3", [0, 1, 0x38, 0x40, 0x7E, 0x7F, 0x80, 0xB8], + [0, 2**-9, 1, 2, 448, float("nan"), -0., -1]), + ("e5m2", [0, 1, 0x3C, 0x40, 0x7B, 0x7C, 0x7F, 0xBC], + [0, 2**-16, 1, 2, 57344, float("inf"), float("nan"), -1]), + ] + for name, codes, values in formats: + packed = torch.tensor([codes * 4], dtype=torch.uint8) + actual = reference.upcast_mxfp_cpu(packed, one, name, dtype) + torch.testing.assert_close(actual, torch.tensor([values * 4], dtype=dtype), + rtol=0, atol=0, equal_nan=True) + for scale, value in [(0, 0), (126, .5), (127, 1), (128, 2), (255, float("nan"))]: + actual = reference.upcast_mxfp_cpu( + torch.full((1, 16), 0x22, dtype=torch.uint8), + torch.tensor([[scale]], dtype=torch.uint8), "e2m1", dtype) + torch.testing.assert_close(actual, torch.full((1, 32), value, dtype=dtype), + rtol=0, atol=0, equal_nan=True) diff --git a/test/wafer/test_naming_compat.py b/test/wafer/test_naming_compat.py new file mode 100644 index 00000000..3d26eb63 --- /dev/null +++ b/test/wafer/test_naming_compat.py @@ -0,0 +1,75 @@ +"""Compatibility behavior at the Wafer naming migration boundaries.""" + +import os +from pathlib import Path +import shutil +import subprocess + +import pytest + + +def test_legacy_imports_keep_one_discoverable_backend(wafer_modules): + from triton.backends import _find_concrete_subclasses + from triton.backends.compiler import BaseBackend + + _, compiler, runtime = wafer_modules + assert compiler.TXDABackend is compiler.WaferBackend + assert compiler.TXDAOptions is compiler.WaferOptions + assert runtime.TXDALauncher is runtime.WaferLauncher + assert runtime.TXDAUtils is runtime.WaferUtils + assert _find_concrete_subclasses(compiler, BaseBackend) is compiler.WaferBackend + + +def test_runtime_dependencies_require_wafer_configuration( + tmp_path, monkeypatch, wafer_modules +): + _, compiler, _ = wafer_modules + root = tmp_path / "sdk" + (root / "lib").mkdir(parents=True) + for name in ("libcommon_util.a", "libinstr_tx81.a", "liblibc_stub.a"): + (root / "lib" / name).write_bytes(b"sdk") + toolchain = tmp_path / "toolchain" + (toolchain / "bin").mkdir(parents=True) + (toolchain / "bin/riscv64-unknown-elf-gcc").touch() + archive_dir = tmp_path / "crt" + archive_dir.mkdir() + (archive_dir / "libvr.a").touch() + for name in ("libm.a", "libc.a", "libgcc.a", "libgloss.a"): + (archive_dir / name).touch() + monkeypatch.setenv("XUANTIE_NAME", str(toolchain)) + monkeypatch.setenv("WAFER_RUNTIME_LIB_DIR", str(archive_dir)) + monkeypatch.delenv("WAFER_DEPS_ROOT", raising=False) + with pytest.raises(RuntimeError, match="WAFER_DEPS_ROOT is not set"): + compiler._runtime_link_inputs() + monkeypatch.setenv("WAFER_DEPS_ROOT", str(root)) + monkeypatch.setattr(compiler, "_find_linker_library", lambda linker, name: archive_dir / name) + _, libraries = compiler._runtime_link_inputs() + assert all(path.parent == root / "lib" for path in libraries[:3]) + + +@pytest.mark.parametrize("mode", ["environment", "canonical", "both"]) +@pytest.mark.parametrize("setting", ["WAFER_DEPS_ROOT", "WAFER_BSP_INCLUDE_DIR"]) +def test_cmake_configuration_precedence(tmp_path, mode, setting): + cmake = shutil.which("cmake") + if cmake is None: + pytest.skip("CMake is required to verify configuration compatibility") + repo = Path(__file__).resolve().parents[2] + script = tmp_path / "config.cmake" + script.write_text( + f'include("{repo / "third_party/wafer/cmake/WaferConfig.cmake"}")\n' + f'if(NOT {setting} STREQUAL EXPECT)\n' + ' message(FATAL_ERROR "configuration precedence mismatch")\n' + 'endif()\n' + ) + arguments = [] + if mode != "environment": + arguments.append(f"-D{setting}=/canonical") + expected = "/environment" if mode == "environment" else "/canonical" + environment = os.environ.copy() + environment.pop(setting, None) + if mode != "canonical": + environment[setting] = "/environment" + subprocess.run( + [cmake, *arguments, f"-DEXPECT={expected}", "-P", str(script)], + check=True, capture_output=True, text=True, env=environment, + ) diff --git a/test/wafer/test_native_errors.py b/test/wafer/test_native_errors.py new file mode 100644 index 00000000..0f7e5d28 --- /dev/null +++ b/test/wafer/test_native_errors.py @@ -0,0 +1,49 @@ +"""Exercise generated C++ validation without submitting work to a device.""" + +import os +from types import SimpleNamespace + +import pytest + + +@pytest.fixture +def native_launcher(monkeypatch, tmp_path, wafer_modules): + if not os.environ.get("KUIPER_ROOT"): + pytest.skip("Native launcher validation needs Kuiper headers and libhpgr") + _, _, runtime = wafer_modules + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path / "cache")) + return runtime.compile_launcher( + runtime.make_launcher({0: "*fp32", 1: "i32"}) + ).launch + + +def test_native_invalid_arguments(native_launcher, tmp_path): + empty = tmp_path / "empty.so" + empty.touch() + + def call(metadata, pointer=123, stream=None, grid=(1, 1, 1), enter=None): + return native_launcher( + *grid, stream, 0, metadata, None, enter, None, pointer, 256 + ) + + with pytest.raises(AttributeError, match="kernel_path"): + call(SimpleNamespace()) + with pytest.raises(TypeError): + call(SimpleNamespace(kernel_path=123, name="test")) + with pytest.raises(OSError, match="empty"): + call(SimpleNamespace(kernel_path=str(empty), name="test")) + with pytest.raises(FileNotFoundError): + call(SimpleNamespace(kernel_path=str(tmp_path / "missing.so"), name="test")) + with pytest.raises(TypeError): + call(SimpleNamespace(), stream=object()) + with pytest.raises(AttributeError, match="data_ptr"): + call(SimpleNamespace(), pointer=object()) + with pytest.raises(ValueError, match="nonnegative"): + call(SimpleNamespace(), grid=(-1, 1, 1)) + assert call(SimpleNamespace(), grid=(0, 1, 1)) is None + + def failing_hook(metadata): + raise ValueError("hook failed") + + with pytest.raises(ValueError, match="hook failed"): + call(SimpleNamespace(), enter=failing_hook) diff --git a/test/wafer/test_noc_sync.py b/test/wafer/test_noc_sync.py new file mode 100644 index 00000000..62405371 --- /dev/null +++ b/test/wafer/test_noc_sync.py @@ -0,0 +1,89 @@ +"""Stress the production request/ack protocol on a host SPM model.""" +from pathlib import Path +import shutil +import subprocess + +import pytest + + +def test_noc_initializer_only_clears_protocol_words(tmp_path): + compiler = shutil.which("c++") + if compiler is None: + pytest.skip("A C++ compiler is required for the initializer test") + path = Path(__file__).resolve().parents[2] / "third_party/wafer/crt/lib/Wafer/noc_init.c" + source = path.read_text() + program = r''' +#include +#include +#define SINGLE_SPM_SYNC_ADDR 0x2f0320 +static uint32_t spm[22]; +static uintptr_t get_spm_memory_mapping(uintptr_t address) { + assert(address == SINGLE_SPM_SYNC_ADDR); + return (uintptr_t)&spm[1]; +} +''' + program += source[source.index("void __NoCRingInit"):] + program += r''' +int main() { + for (auto &word : spm) word = 0xdeadbeef; + __NoCRingInit(); + __NoCRingInit(); + for (int i = 0; i < 22; ++i) + assert(spm[i] == ((i == 1 || i == 2) ? 0 : 0xdeadbeef)); +} +''' + src, exe = tmp_path / "init.cpp", tmp_path / "init" + src.write_text(program) + subprocess.run([compiler, "-std=c++17", "-O2", str(src), "-o", str(exe)], check=True) + subprocess.run([str(exe)], check=True, timeout=10) + + +@pytest.mark.parametrize("tiles", [3, 16]) +def test_noc_sync_repeated_rounds_and_launches(tmp_path, tiles): + compiler = shutil.which("c++") + if compiler is None: + pytest.skip("A C++ compiler is required for the ring protocol test") + path = Path(__file__).resolve().parents[2] / "third_party/wafer/crt/lib/Wafer/send.c" + source = path.read_text() + protocol = source[source.index("static void noc_memory_fence"):source.index("// Send to the next tile")] + # MMIO accesses become atomic host-memory accesses, avoiding C++ data races. + protocol = protocol.replace("volatile uint32_t", "std::atomic") + program = r''' +#include +#include +#include +#include +#include +#include +#define SINGLE_SPM_SYNC_ADDR 0 +static std::atomic spm[16][2]{}; +static thread_local int tile; +static uintptr_t get_spm_memory_mapping(uintptr_t) { return (uintptr_t)spm[tile]; } +static uintptr_t get_tile_spm_addr_base(int i, int, int) { return (uintptr_t)spm[i]; } +''' + program += protocol + program += f''' +int main() {{ + const int count = {tiles}; + for (int launch = 0; launch < 2; ++launch) {{ + std::vector workers; + for (int i = 0; i < count; ++i) workers.emplace_back([i, count]() {{ + tile = i; + for (int round = 0; round < 1000; ++round) {{ + if (i == 1 && round % 7 == 0) + std::this_thread::sleep_for(std::chrono::microseconds(1)); + noc_ring_sync((i + count - 1) % count, (i + 1) % count); + }} + }}); + for (auto &worker : workers) worker.join(); + for (int i = 0; i < count; ++i) {{ + assert(spm[i][0] == 0); + assert(spm[i][1] == 0); + }} + }} +}} +''' + src, exe = tmp_path / "sync.cpp", tmp_path / "sync" + src.write_text(program) + subprocess.run([compiler, "-std=c++17", "-O2", "-pthread", str(src), "-o", str(exe)], check=True) + subprocess.run([str(exe)], check=True, timeout=20) diff --git a/test/wafer/test_patch_profiles.py b/test/wafer/test_patch_profiles.py new file mode 100644 index 00000000..6fb9c579 --- /dev/null +++ b/test/wafer/test_patch_profiles.py @@ -0,0 +1,160 @@ +"""Check profile composition against the actual pinned Triton Git objects.""" + +import importlib.util +import json +from pathlib import Path +import subprocess +import sys + +import pytest + +ROOT = Path(__file__).resolve().parents[2] +spec = importlib.util.spec_from_file_location("apply_triton_profile", ROOT / "scripts/wafer/apply_triton_profile.py") +profiles = importlib.util.module_from_spec(spec) +spec.loader.exec_module(profiles) + + +@pytest.fixture +def triton_source(tmp_path): + source = tmp_path / "triton" + subprocess.run(["git", "clone", "--shared", "--no-checkout", str(ROOT / "third_party/triton"), str(source)], + check=True, capture_output=True) + commit = "c3c476f357f1e9768ea4e45aa5c17528449ab9ef" + subprocess.run(["git", "-C", str(source), "checkout", "--detach", commit], check=True, capture_output=True) + return source + + +@pytest.mark.parametrize("profile", ["ascend", "wafer-tools", "wafer-frontend"]) +def test_profiles_apply_idempotently_and_preserve_edits(triton_source, profile): + source = triton_source + first = profiles.apply_profile(source, profile) + diff = profiles.git(source, "diff", "--binary") + assert profiles.apply_profile(source, profile) == first + assert profiles.apply_profile(source, profile, check=True) == first + assert profiles.git(source, "diff", "--binary") == diff + edited = source / "python/src/ir.cc" + edited.write_text(edited.read_text() + "\n// independent user change\n") + with pytest.raises(RuntimeError, match="Refusing to replace modified") as error: + profiles.apply_profile(source, profile) + assert "--force" in str(error.value) + assert "including staged edits" in str(error.value) + assert "No automatic backup" in str(error.value) + assert edited.read_text().endswith("// independent user change\n") + + +@pytest.mark.parametrize("previous,target", [ + ("wafer-frontend", "ascend"), ("ascend", "wafer-tools"), +]) +def test_force_replaces_profile_and_tracked_edits_only(triton_source, tmp_path, previous, target): + source = triton_source + profiles.apply_profile(source, previous) + original_readme = profiles.git(source, "show", "HEAD:README.md") + (source / "README.md").write_text("staged edit outside the profile\n") + (source / "staged-only.txt").write_text("new tracked file\n") + profiles.git(source, "add", "README.md", "staged-only.txt") + (source / "README.md").write_text("unstaged edit outside the profile\n") + (source / "python/src/ir.cc").write_text("independent edit inside a profile\n") + (source / "force-untracked.txt").write_text("keep untracked\n") + exclude = source / ".git/info/exclude" + exclude.write_text(exclude.read_text() + "\nforce-ignored.txt\n") + (source / "force-ignored.txt").write_text("keep ignored\n") + record = tmp_path / "profile.json" + outside = tmp_path / "outside-source.txt" + outside.write_text("keep the caller's other files\n") + + result = subprocess.run([ + sys.executable, str(ROOT / "scripts/wafer/apply_triton_profile.py"), + "--source", str(source), "--profile", target, "--force", "--record", str(record), + ], capture_output=True, text=True) + + assert result.returncode == 0, result.stderr + assert "WARNING:" in result.stderr and "No automatic backup" in result.stderr + assert str(source) in result.stderr + assert json.loads(record.read_text()) == profiles.apply_profile(source, target, check=True) + assert (source / "README.md").read_bytes() == original_readme + assert not (source / "staged-only.txt").exists() + assert profiles.git(source, "diff", "--cached", "--binary") == b"" + assert (source / "force-untracked.txt").read_text() == "keep untracked\n" + assert (source / "force-ignored.txt").read_text() == "keep ignored\n" + assert outside.read_text() == "keep the caller's other files\n" + + +@pytest.mark.parametrize("ignored", [False, True]) +def test_force_refuses_untracked_obstructions_without_resetting(triton_source, ignored): + source = triton_source + profiles.git(source, "rm", "--cached", "README.md") + (source / "README.md").write_text("keep the untracked replacement\n") + if ignored: + exclude = source / ".git/info/exclude" + exclude.write_text(exclude.read_text() + "\n/README.md\n") + (source / "python/src/ir.cc").write_text("keep the tracked edit too\n") + working = profiles.git(source, "diff", "--binary") + staged = profiles.git(source, "diff", "--cached", "--binary") + + with pytest.raises(RuntimeError, match="untracked or ignored files"): + profiles.apply_profile(source, "ascend", force=True) + + assert profiles.git(source, "diff", "--binary") == working + assert profiles.git(source, "diff", "--cached", "--binary") == staged + assert (source / "README.md").read_text() == "keep the untracked replacement\n" + + +def test_force_does_not_bypass_pinned_commit(triton_source): + source = triton_source + # The source cache can be shallow, so create a different revision locally. + profiles.git(source, "-c", "user.name=Profile Test", "-c", "user.email=profile-test@example.invalid", + "-c", "commit.gpgsign=false", "commit", "--allow-empty", "-m", "Different test revision") + head = profiles.git(source, "rev-parse", "HEAD") + (source / "README.md").write_text("keep edits on the wrong revision\n") + with pytest.raises(RuntimeError, match="requires Triton"): + profiles.apply_profile(source, "ascend", force=True) + assert profiles.git(source, "rev-parse", "HEAD") == head + assert (source / "README.md").read_text() == "keep edits on the wrong revision\n" + + +def test_force_cannot_be_combined_with_check(triton_source): + source = triton_source + (source / "README.md").write_text("check must not discard this\n") + with pytest.raises(RuntimeError, match="--check and --force"): + profiles.apply_profile(source, "ascend", check=True, force=True) + result = subprocess.run([ + sys.executable, str(ROOT / "scripts/wafer/apply_triton_profile.py"), + "--source", str(source), "--profile", "ascend", "--check", "--force", + ], capture_output=True, text=True) + assert result.returncode == 2 + assert "not allowed with argument" in result.stderr + assert (source / "README.md").read_text() == "check must not discard this\n" + + +def test_force_validates_patches_before_discarding_edits(triton_source, tmp_path, monkeypatch): + source = triton_source + config_root = tmp_path / "invalid-profile" + patch_dir = config_root / "patch/triton" + patch_dir.mkdir(parents=True) + (patch_dir / "invalid.patch").write_text("not a valid patch\n") + catalog = config_root / "profiles.json" + catalog.write_text(json.dumps({ + "triton_commit": profiles.git(source, "rev-parse", "HEAD").decode().strip(), + "profiles": {"ascend": ["patch/triton/invalid.patch"]}, + })) + monkeypatch.setattr(profiles, "ROOT", config_root) + monkeypatch.setattr(profiles, "CATALOG", catalog) + (source / "README.md").write_text("keep staged edit\n") + profiles.git(source, "add", "README.md") + (source / "README.md").write_text("keep unstaged edit\n") + working = profiles.git(source, "diff", "--binary") + staged = profiles.git(source, "diff", "--cached", "--binary") + + with pytest.raises(subprocess.CalledProcessError): + profiles.apply_profile(source, "ascend", force=True) + + assert profiles.git(source, "diff", "--binary") == working + assert profiles.git(source, "diff", "--cached", "--binary") == staged + + +def test_force_rejects_a_source_inside_the_repository(triton_source): + source = triton_source + (source / "README.md").write_text("keep changes outside the selected subdirectory\n") + with pytest.raises(RuntimeError, match="repository root"): + profiles.apply_profile(source / "python", "ascend", force=True) + assert (source / "README.md").read_text() == "keep changes outside the selected subdirectory\n" diff --git a/test/wafer/test_precision_modes.py b/test/wafer/test_precision_modes.py new file mode 100644 index 00000000..bd94d972 --- /dev/null +++ b/test/wafer/test_precision_modes.py @@ -0,0 +1,64 @@ +"""Precision policy must affect both compilation options and cache identity.""" +import pytest +from triton.backends.compiler import GPUTarget + + +def test_precision_mode_snapshot_and_override(monkeypatch, wafer_modules, fake_toolchain): + _, compiler, _ = wafer_modules + monkeypatch.setenv("WAFER_ENABLE_RUNTIME", "0") + monkeypatch.delenv("PRECISION_MODE", raising=False) + monkeypatch.setenv("PRECISION_PRIORITY", "1") + legacy = compiler.WaferBackend(GPUTarget("wafer", "wafer", 32)) + assert legacy.parse_options({}).precision_mode == 2 + before = legacy.hash() + monkeypatch.setenv("PRECISION_MODE", "1") + assert legacy.hash() == before + modern = compiler.WaferBackend(GPUTarget("wafer", "wafer", 32)) + assert modern.parse_options({}).precision_mode == 1 + assert modern.hash() != before + assert modern.parse_options({"precision_mode": 0}).precision_mode == 0 + assert modern.parse_options({"precision_mode": 2}).hash() != modern.parse_options({}).hash() + monkeypatch.setenv("PRECISION_MODE", "invalid") + with pytest.raises(ValueError, match="PRECISION_MODE"): + compiler.WaferBackend(GPUTarget("wafer", "wafer", 32)) + + +def test_strided_output_copyback(tmp_path): + import os + import re + import subprocess + from triton.backends.dicp_triton.wafer import _find_wafer_opt + source = tmp_path / "stride.mlir" + source.write_text('''module { + func.func @slice(%input: memref<4xf32>, %output: memref<8xf32>) { + %one = arith.constant 1.0 : f32 + %view = memref.subview %output[0] [4] [2] : memref<8xf32> to memref<4xf32, strided<[2]>> + "mk.addvs"(%input, %one, %view) : (memref<4xf32>, f32, memref<4xf32, strided<[2]>>) -> () + return + } + }''') + result = subprocess.run([os.getenv("WAFER_TEST_OPT") or _find_wafer_opt(), str(source), + "--materialize-strided-linalg-inputs"], + capture_output=True, text=True, timeout=30) + assert result.returncode == 0, result.stderr + copies = re.findall(r"memref.copy (%[\w]+), (%[\w]+)", result.stdout) + assert len(copies) == 2, result.stdout + assert copies[0] == copies[1][::-1], result.stdout + assert copies[0][0] != copies[0][1], result.stdout + + +def test_pipeline_option_cache_isolation(monkeypatch, wafer_modules, fake_toolchain): + _, compiler, _ = wafer_modules + monkeypatch.setenv("WAFER_ENABLE_RUNTIME", "0") + monkeypatch.setenv("TRITON_PIPELINE", "0") + plain = compiler.WaferBackend(GPUTarget("wafer", "wafer", 32)) + before = plain.hash() + monkeypatch.setenv("TRITON_PIPELINE", "1") + pipelined = compiler.WaferBackend(GPUTarget("wafer", "wafer", 32)) + assert not plain.parse_options({}).enable_pipeline + assert plain.hash() == before + assert pipelined.parse_options({}).enable_pipeline + assert pipelined.hash() != before + override = pipelined.parse_options({"enable_pipeline": False}) + assert not override.enable_pipeline + assert override.hash() != pipelined.parse_options({}).hash() diff --git a/test/wafer/test_scalar_copy.py b/test/wafer/test_scalar_copy.py new file mode 100644 index 00000000..acef00d0 --- /dev/null +++ b/test/wafer/test_scalar_copy.py @@ -0,0 +1,37 @@ +"""Scalar copies used by debug reductions must survive the real SPM lowering.""" +import os +import re +import subprocess + +import pytest + + +@pytest.mark.parametrize("dtype", ["i1", "i8", "i32", "f32"]) +def test_scalar_spm_copy(dtype, tmp_path): + from triton.backends.dicp_triton.wafer import _find_wafer_opt + + source = tmp_path / "scalar-copy.mlir" + source.write_text(f"""module {{ + func.func @copy(%value: {dtype}) -> {dtype} {{ + %src = memref.alloc() : memref<{dtype}> + %dst = memref.alloc() : memref<{dtype}> + memref.store %value, %src[] : memref<{dtype}> + memref.copy %src, %dst : memref<{dtype}> to memref<{dtype}> + %result = memref.load %dst[] : memref<{dtype}> + return %result : {dtype} + }} + }}""") + result = subprocess.run([ + os.getenv("WAFER_TEST_OPT") or str(_find_wafer_opt()), str(source), + "--spmd-allocate-shared-memory", "--expand-strided-metadata", + "--lower-affine", "--mk-to-wafer", + ], capture_output=True, text=True, timeout=30) + assert result.returncode == 0, result.stderr + assert "memref.copy" not in result.stdout + assert "linalg.transpose" not in result.stdout + # Both the original access and the copy must retain scalar memory ops; + # WaferToLLVM needs the isSpm marker to apply the device SPM address map. + loads = re.findall(r"memref.load[^\n]+", result.stdout) + stores = re.findall(r"memref.store[^\n]+", result.stdout) + assert len(loads) == len(stores) == 2, result.stdout + assert all("isSpm = 1" in op for op in loads + stores), result.stdout diff --git a/test/wafer/test_tle_frontend.py b/test/wafer/test_tle_frontend.py new file mode 100644 index 00000000..947b1faf --- /dev/null +++ b/test/wafer/test_tle_frontend.py @@ -0,0 +1,96 @@ +"""Compile the TLE frontend against the installed Wafer/Triton IR bindings.""" +import pytest +import triton +import triton.language as tl +import triton.experimental.tle.language as tle +from triton._C.libtriton import ir +from triton.backends.compiler import GPUTarget +from triton.compiler import ASTSource +from triton.compiler.compiler import make_backend + + +@triton.jit +def local_kernel(out): + buf = tle.dsa.alloc((16,), tl.float32) + i = tl.arange(0, 16) + ptr = tle.dsa.local_ptr(buf, [i]) + tl.store(ptr, i.to(tl.float32)) + tl.store(out + i, tl.load(ptr)) + + +@triton.jit +def buffer_helper(buf, out): + i = tl.arange(0, 16) + tl.store(out + i, tl.load(tle.dsa.local_ptr(buf, [i]))) + + +@triton.jit +def helper_kernel(out): + buf = tle.dsa.alloc((16,), tl.float32) + buffer_helper(buf, out) + + +@triton.jit +def remote_kernel(out): + buf = tle.dsa.alloc((16,), tl.float32) + remote = tle.remote(buf, tl.program_id(0)) + ptr = tle.dsa.local_ptr(remote, [tl.arange(0, 16)]) + tl.store(ptr, tl.full((16,), 1, tl.float32)) + + +@triton.jit +def copy_kernel(out): + src = tle.dsa.alloc((16,), tl.float32) + dst = tle.dsa.alloc((16,), tl.float32) + tle.dsa.copy(src, dst, (16,)) + tle.dsa.copy(dst, out, (16,)) + + +@triton.jit +def barrier_kernel(out): + tle.distributed_barrier() + + +@triton.jit +def invalid_alloc_kernel(out): + tle.dsa.alloc((-1,), tl.float32) + + +@triton.jit +def pipeline_kernel(out): + for i in tle.dsa.pipeline(0, 16, num_stages=2): + tl.store(out + i, i.to(tl.float32)) + + +def make_ttir(fn): + target = GPUTarget("wafer", "wafer", 32) + backend = make_backend(target) + options = backend.parse_options({}) + context = ir.context() + ir.load_dialects(context) + backend.load_dialects(context) + src = ASTSource(fn, signature={"out": "*fp32"}) + return str(src.make_ir(target, options, backend.get_codegen_implementation(options), + backend.get_module_map(), context)) + + +@pytest.mark.parametrize("kernel,operations", [ + (local_kernel, ["dsa.alloc", "dsa.local_pointers", "tt.load", "tt.store"]), + (helper_kernel, ["dsa.alloc", "tt.call", "dsa.local_pointers"]), + (remote_kernel, ["dsa.remote_pointers", "tt.get_program_id"]), + (copy_kernel, ["dsa.copy"]), + (pipeline_kernel, ["scf.for", "tt.num_stages = 2"]), +]) +def test_tle_frontend(kernel, operations): + module = make_ttir(kernel) + for operation in operations: + assert operation in module + + +@pytest.mark.parametrize("kernel,message", [ + (barrier_kernel, "cross-tile CRT implementation"), + (invalid_alloc_kernel, "must be a positive integer"), +]) +def test_tle_rejects_unsupported_semantics(kernel, message): + with pytest.raises(triton.compiler.errors.CompilationError, match=message): + make_ttir(kernel) diff --git a/test/wafer/verify_wafer_acceptance.py b/test/wafer/verify_wafer_acceptance.py new file mode 100644 index 00000000..6cfcb037 --- /dev/null +++ b/test/wafer/verify_wafer_acceptance.py @@ -0,0 +1,143 @@ +#!/usr/bin/env python3 + +import argparse +import os +import tempfile +from pathlib import Path + + +def parse_args(): + workspace_root = Path(__file__).resolve().parents[4] + parser = argparse.ArgumentParser( + description="Compile a Triton GEMM through Wafer to LLVM dialect IR." + ) + parser.add_argument( + "--reference", + type=Path, + default=workspace_root / "gemm_ll_0.mlir", + help="LLVM dialect IR used for structural comparison.", + ) + parser.add_argument( + "--output-dir", + type=Path, + help="Artifact directory outside the Git worktree (default: temporary directory).", + ) + return parser.parse_args() + + +def require_markers(text, markers, label): + missing = [marker for marker in markers if marker not in text] + if missing: + raise RuntimeError(f"{label} is missing structural markers: {missing}") + + +def main(): + args = parse_args() + output_dir = args.output_dir or Path( + tempfile.mkdtemp(prefix="wafer-gemm-acceptance-", dir="/tmp") + ) + output_dir = output_dir.resolve() + repo_root = Path(__file__).resolve().parents[2] + if output_dir == repo_root or repo_root in output_dir.parents: + raise RuntimeError("Acceptance artifacts must be stored outside the Git worktree") + output_dir.mkdir(parents=True, exist_ok=True) + + if not args.reference.is_file(): + raise FileNotFoundError(f"Reference IR not found: {args.reference}") + + os.environ["DICP_BACKEND"] = "wafer" + os.environ["USE_SIM_MODE"] = "1" + os.environ["TRITON_ALWAYS_COMPILE"] = "1" + os.environ["TRITON_DUMP_PATH"] = str(output_dir) + + import triton + import triton.language as tl + from triton._C import libtriton + from triton.backends import backends + from triton.backends.compiler import GPUTarget + from triton.compiler import ASTSource + + if "dicp_triton" not in backends: + raise RuntimeError(f"dicp_triton backend was not discovered: {list(backends)}") + + target = GPUTarget("wafer", "wafer", 32) + backend = backends["dicp_triton"].compiler(target) + backend.load_dialects(libtriton.ir.context()) + + import triton.language.extra.wafer # noqa: F401 + + @triton.jit + def wafer_gemm( + a_base, + b_base, + out, + M, + N, + K, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, + ): + rows = tl.arange(0, BLOCK_M) + columns = tl.arange(0, BLOCK_N) + reduction = tl.arange(0, BLOCK_K) + a_offsets = rows[:, None] * K + reduction[None, :] + b_offsets = reduction[:, None] * N + columns[None, :] + out_offsets = rows[:, None] * N + columns[None, :] + a_ptrs = a_base + a_offsets + b_ptrs = b_base + b_offsets + out_ptrs = out + out_offsets + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + accumulator = tl.dot(a, b) + tl.store(out_ptrs, accumulator) + + source = ASTSource( + fn=wafer_gemm, + signature={ + "a_base": "*bf16", + "b_base": "*bf16", + "out": "*bf16", + "M": "i32", + "N": "i32", + "K": "i32", + "BLOCK_M": "constexpr", + "BLOCK_N": "constexpr", + "BLOCK_K": "constexpr", + }, + constexprs={"BLOCK_M": 128, "BLOCK_N": 256, "BLOCK_K": 64}, + ) + kernel = triton.compile(source, target=target) + required_stages = {"ttir", "coreir", "wafer_ir", "llir"} + missing_stages = sorted(required_stages.difference(kernel.asm)) + if missing_stages: + raise RuntimeError(f"Wafer compilation is missing stages: {missing_stages}") + + llvm_mlir_path = output_dir / "llvm.mlir" + if not llvm_mlir_path.is_file(): + raise RuntimeError(f"Wafer did not dump LLVM dialect IR: {llvm_mlir_path}") + + generated = llvm_mlir_path.read_text(encoding="utf-8") + reference = args.reference.read_text(encoding="utf-8") + common_markers = ( + "llvm.func", + "@__Gemm", + 'section = "ExportedDYNSYMTab"', + "triton_tsm.spm_use", + ) + require_markers(generated, common_markers, "generated LLVM dialect IR") + require_markers(reference, common_markers, "reference LLVM dialect IR") + require_markers(generated, ("@wafer_gemm",), "generated LLVM dialect IR") + + print("Wafer GEMM compiler acceptance passed:") + print(f" triton: {triton.__file__}") + print(f" libtriton: {libtriton.__file__}") + print(f" backend: {target}") + print(f" stages: {list(kernel.asm)}") + print(f" llvm dialect IR: {llvm_mlir_path}") + print(f" reference: {args.reference.resolve()}") + print(f" structural markers: {', '.join(common_markers)}") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/test/wafer/verify_wafer_examples.py b/test/wafer/verify_wafer_examples.py new file mode 100644 index 00000000..f58e149d --- /dev/null +++ b/test/wafer/verify_wafer_examples.py @@ -0,0 +1,196 @@ +#!/usr/bin/env python3 + +import argparse +import functools +import os +import subprocess +import tempfile +from pathlib import Path + + +def parse_args(): + repo_root = Path(__file__).resolve().parents[2] + workspace_root = repo_root.parents[1] + parser = argparse.ArgumentParser( + description="Compile representative Triton kernels with baseline and new Wafer compilers." + ) + parser.add_argument( + "--baseline-opt", + type=Path, + default=workspace_root / "DLCompiler/third_party/wafer/build_manual/install/bin/wafer-opt", + ) + parser.add_argument( + "--new-opt", + type=Path, + default=repo_root / "third_party/wafer/build_manual/third_party/wafer/bin/wafer-opt", + ) + parser.add_argument("--output-dir", type=Path) + return parser.parse_args() + + +@functools.lru_cache(maxsize=None) +def uses_wafer_pass_names(wafer_opt): + help_text = subprocess.check_output([str(wafer_opt), "--help"], text=True) + return "--mk-to-wafer" in help_text + + +def run_stage(wafer_opt, source, output, arguments): + if not uses_wafer_pass_names(wafer_opt): + raise RuntimeError(f"Compiler does not support Wafer pass names: {wafer_opt}") + subprocess.run( + [str(wafer_opt), str(source), *arguments, "-o", str(output)], + check=True, + ) + + +def lower_case(wafer_opt, ttir_path, output_dir): + coreir_path = output_dir / "coreir.mlir" + wafer_ir_path = output_dir / "wafer_ir.mlir" + llvm_path = output_dir / "llvm.mlir" + run_stage( + wafer_opt, + ttir_path, + coreir_path, + [ + "--triton-to-core-dialects", + "--tle-to-mk", + "--dsa-memory-to-core", + "--linalg-tiling", + "--core-dialects-to-mk", + "--linalg-fusion", + "--legalize-tensor-form-loops", + "--one-shot-bufferize", + "--convert-bufferization-to-memref", + "--cse", + "--canonicalize", + ], + ) + run_stage( + wafer_opt, + coreir_path, + wafer_ir_path, + ["--spmd-allocate-shared-memory", "--expand-strided-metadata", "--lower-affine", "--mk-to-wafer", "--cse"], + ) + run_stage( + wafer_opt, + wafer_ir_path, + llvm_path, + [ + "--wafer-memref-to-llvm", + "--addr-to-llvm", + "--convert-scf-to-cf", + "--convert-math-to-llvm", + "--convert-math-to-libm", + "--convert-cf-to-llvm", + "--convert-func-to-llvm", + "--expand-strided-metadata", + "--finalize-memref-to-llvm", + "--kernel-arg-buffer", + "--wafer-to-llvm", + "--convert-arith-to-llvm", + "--reconcile-unrealized-casts", + "--canonicalize", + "--export-kernel-symbols", + ], + ) + if "llvm.func" not in llvm_path.read_text(encoding="utf-8"): + raise RuntimeError(f"LLVM dialect output has no llvm.func: {llvm_path}") + + +def main(): + args = parse_args() + repo_root = Path(__file__).resolve().parents[2] + output_dir = (args.output_dir or Path(tempfile.mkdtemp(prefix="wafer-examples-", dir="/tmp"))).resolve() + if output_dir == repo_root or repo_root in output_dir.parents: + raise RuntimeError("Regression artifacts must be stored outside the Git worktree") + for compiler in (args.baseline_opt, args.new_opt): + if not compiler.is_file(): + raise FileNotFoundError(compiler) + + os.environ["DICP_BACKEND"] = "wafer" + os.environ["USE_SIM_MODE"] = "1" + + import triton + import triton.language as tl + from triton._C import libtriton + from triton.backends import backends + from triton.backends.compiler import GPUTarget + from triton.compiler import ASTSource + + @triton.jit + def addptr(in_ptr, out_ptr): + for offset in range(0, 10, 2): + first = in_ptr + 1 + offset + second = first + 1 + tl.store(out_ptr + 1 + offset, tl.load(first)) + tl.store(out_ptr + 2 + offset, tl.load(second)) + + @triton.jit + def block_copy(in_ptr, out_ptr): + source = tl.make_block_ptr(in_ptr + 8, (2, 2), (2, 1), (0, 0), (2, 2), (1, 0)) + destination = tl.make_block_ptr(out_ptr, (2, 2), (2, 1), (0, 0), (2, 2), (1, 0)) + tl.store(destination, tl.load(source, boundary_check=(0,)), boundary_check=(0,)) + + @triton.jit + def reduce_2d(in_ptr, out_ptr, stride, elements, BLOCK_SIZE: tl.constexpr): + row = tl.program_id(0) + block = tl.make_block_ptr(in_ptr, (elements * tl.num_programs(0),), (1,), (stride * row,), (BLOCK_SIZE,), (0,)) + tl.store(out_ptr + row, tl.sum(tl.load(block, boundary_check=(0,)), axis=0)) + + @triton.jit + def scan_1d(out_ptr, in_ptr, elements, M: tl.constexpr, N: tl.constexpr): + offsets = tl.arange(0, M) + values = tl.load(in_ptr + offsets, mask=offsets < elements, other=0.0) + result = tl.cumsum(values).reshape((1, M)).broadcast_to((N, M)) + rows = tl.arange(0, N) + columns = tl.arange(0, M) + tl.store(out_ptr + M * rows[:, None] + columns[None, :], result, mask=columns[None, :] < elements) + + cases = { + "addptr": ASTSource(addptr, {"in_ptr": "*fp32", "out_ptr": "*fp32"}, {}), + "blockptr": ASTSource(block_copy, {"in_ptr": "*fp32", "out_ptr": "*fp32"}, {}), + "reduce": ASTSource( + reduce_2d, + {"in_ptr": "*fp32", "out_ptr": "*fp32", "stride": "i32", "elements": "i32", "BLOCK_SIZE": "constexpr"}, + {"BLOCK_SIZE": 32}, + ), + "scan": ASTSource( + scan_1d, + {"out_ptr": "*fp32", "in_ptr": "*fp32", "elements": "i32", "M": "constexpr", "N": "constexpr"}, + {"M": 32, "N": 2}, + ), + } + + target = GPUTarget("wafer", "wafer", 32) + backend = backends["dicp_triton"].compiler(target) + options = backend.parse_options({}) + context = libtriton.ir.context() + libtriton.ir.load_dialects(context) + backend.load_dialects(context) + output_dir.mkdir(parents=True, exist_ok=True) + + for name, source in cases.items(): + case_dir = output_dir / name + case_dir.mkdir() + module = source.make_ir( + target, + options, + backend.get_codegen_implementation(options), + backend.get_module_map(), + context, + ) + module = backend.make_ttir(module, {}, options) + ttir_path = case_dir / "input.ttir.mlir" + ttir_path.write_text(str(module), encoding="utf-8") + + for label, compiler in (("baseline", args.baseline_opt), ("new", args.new_opt)): + result_dir = case_dir / label + result_dir.mkdir() + lower_case(compiler, ttir_path, result_dir) + print(f"PASS {name}: {label}") + + print(f"Artifacts: {output_dir}") + + +if __name__ == "__main__": + main() diff --git a/test/wafer/verify_wafer_installation.py b/test/wafer/verify_wafer_installation.py new file mode 100644 index 00000000..bc43505e --- /dev/null +++ b/test/wafer/verify_wafer_installation.py @@ -0,0 +1,81 @@ +#!/usr/bin/env python3 +"""Check an installed Wafer wheel and mode/cache separation without device work. + +Run outside the repository root after sourcing init_wafer_env.sh. SDK and LLVM +are still required; the CRT archive must come from the installed wheel. +""" + +import argparse +import importlib +import importlib.metadata +import json +import os +from pathlib import Path + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", required=True, type=Path) + args = parser.parse_args() + os.environ.pop("WAFER_RUNTIME_LIB_DIR", None) + os.environ["USE_SIM_MODE"] = "0" + + import triton + from triton.backends import _find_concrete_subclasses + from triton.backends.compiler import BaseBackend, GPUTarget + from triton.backends.dicp_triton import wafer, wafer_runtime + from triton.compiler import ASTSource + from verify_wafer_runtime import wafer_vector + from audit_wafer_elf import audit_kernel + + repo = Path(__file__).resolve().parents[2] + package = Path(triton.__file__).resolve().parent + assert repo not in package.parents, f"Source checkout shadows installed wheel: {package}" + assert wafer.TXDABackend is wafer.WaferBackend + assert wafer.TXDAOptions is wafer.WaferOptions + assert wafer_runtime.TXDALauncher is wafer_runtime.WaferLauncher + assert wafer_runtime.TXDAUtils is wafer_runtime.WaferUtils + assert _find_concrete_subclasses(wafer, BaseBackend) is wafer.WaferBackend + canonical = importlib.import_module("triton.language.extra.wafer.libdevice") + legacy = importlib.import_module("triton.language.extra.txda.libdevice") + exports = [name for name in dir(canonical) if not name.startswith("_")] + assert exports + for name in exports: + assert getattr(legacy, name) is getattr(canonical, name), name + + _, libraries = wafer._runtime_link_inputs() + archive = libraries[3].resolve() + assert archive == Path(wafer.__file__).resolve().parent / "lib/libvr.a", archive + source = ASTSource( + fn=wafer_vector, + signature=dict(lhs="*fp32", rhs="*fp32", output="*fp32", alpha="fp32", + size="i32", BLOCK="constexpr"), + constexprs=dict(BLOCK=256), + ) + compiled = [] + for runtime in (False, True, False, True): + os.environ["WAFER_ENABLE_RUNTIME"] = str(int(runtime)) + kernel = triton.compile(source, target=GPUTarget("wafer", "wafer", 32)) + extension = "so" if runtime else "o" + assert list(kernel.asm) == ["source", "ttir", "coreir", "wafer_ir", "llir", extension] + if runtime: + audit_kernel(kernel.metadata.kernel_path, kernel.metadata.device_log_abi) + compiled.append({"runtime": runtime, "hash": kernel.hash, "extension": extension}) + assert compiled[0] == compiled[2] + assert compiled[1] == compiled[3] + assert compiled[0]["hash"] != compiled[1]["hash"] + report = { + "installed_package": str(package), + "version": importlib.metadata.version("triton"), + "packaged_crt": str(archive), + "legacy_language_exports_checked": len(exports), + "mode_sequence": compiled, + "device_work_submitted": False, + } + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(report, indent=2) + "\n") + print(json.dumps(report, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/test/wafer/verify_wafer_runtime.py b/test/wafer/verify_wafer_runtime.py new file mode 100755 index 00000000..e64ab24d --- /dev/null +++ b/test/wafer/verify_wafer_runtime.py @@ -0,0 +1,431 @@ +#!/usr/bin/env python3 +"""Numerical Wafer acceptance through the installed CompiledKernel launcher. + +Run outside the source root after sourcing the workspace environment. Requires +USE_SIM_MODE=0, WAFER_ENABLE_RUNTIME=1, WAFER_RUNTIME_LIB_DIR, and a matching +WAFER_DEVICE_LOG_ABI (rcs on Kuiper 1.4). --torch uses actual TXDA tensors. +""" + +import argparse +import ctypes +from contextlib import ExitStack, contextmanager +import os +from pathlib import Path +import sys +import time + +import numpy as np +import triton +from triton import knobs +import triton.language as tl +from triton.backends.compiler import GPUTarget +from triton.compiler import ASTSource +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "scripts/wafer")) +from audit_wafer_elf import audit_kernel + + +@triton.jit +def wafer_vector(lhs, rhs, output, alpha, size, BLOCK: tl.constexpr): + offset = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offset < size + value = tl.load(lhs + offset, mask=mask) + tl.load(rhs + offset, mask=mask) + alpha + tl.store(output + offset, value, mask=mask) + + +@triton.jit +def wafer_reduction(values, output, size, BLOCK: tl.constexpr): + offset = tl.arange(0, BLOCK) + values = tl.load(values + offset, mask=offset < size, other=0.0) + tl.store(output, tl.sum(values, 0)) + + +@triton.jit +def wafer_matmul( + lhs, rhs, output, M, N, K, BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr +): + rows = tl.arange(0, BM) + cols = tl.arange(0, BN) + inner = tl.arange(0, BK) + a = tl.load(lhs + rows[:, None] * K + inner[None, :]) + b = tl.load(rhs + inner[:, None] * N + cols[None, :]) + result = tl.dot(a, b) + tl.store(output + rows[:, None] * N + cols[None, :], result) + + +def compile_kernel(fn, signature, constants): + source = ASTSource(fn=fn, signature=signature, constexprs=constants) + started = time.monotonic() + kernel = triton.compile(source, target=GPUTarget("wafer", "wafer", 32)) + assert list(kernel.asm) == ["source", "ttir", "coreir", "wafer_ir", "llir", "so"] + audit_kernel(kernel.metadata.kernel_path, kernel.metadata.device_log_abi) + print( + f"compiled {kernel.name} in {time.monotonic() - started:.3f}s: {kernel.metadata.kernel_path}", + flush=True, + ) + return kernel + + +class RawRuntime: + def __init__(self): + root = Path(os.environ["KUIPER_ROOT"]) + self.library = ctypes.CDLL(str(root / "lib/libhpgr.so")) + for name, argtypes in { + "txSetDevice": [ctypes.c_uint32], + "txMalloc": [ctypes.POINTER(ctypes.c_void_p), ctypes.c_uint64], + "txFree": [ctypes.c_void_p], + "txMemcpy": [ + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_uint64, + ctypes.c_int, + ], + }.items(): + fn = getattr(self.library, name) + fn.argtypes, fn.restype = argtypes, ctypes.c_int + self.check(self.library.txSetDevice(0), "txSetDevice") + + @staticmethod + def check(status, operation): + if status: + raise RuntimeError(f"{operation} failed with status 0x{status:x}") + + @contextmanager + def buffer(self, host): + pointer = ctypes.c_void_p() + self.check( + self.library.txMalloc(ctypes.byref(pointer), host.nbytes), "txMalloc" + ) + buffer = RawBuffer(self, pointer, host) + failed = False + try: + self.check( + self.library.txMemcpy( + pointer, ctypes.c_void_p(host.ctypes.data), host.nbytes, 1 + ), + "H2D", + ) + yield buffer + except BaseException: + failed = True + raise + finally: + status = self.library.txFree(pointer) + if status and failed: + print(f"cleanup txFree also failed: 0x{status:x}", flush=True) + else: + self.check(status, "txFree") + + +class RawBuffer: + def __init__(self, runtime, pointer, host): + self.runtime, self.pointer, self.host = runtime, pointer, host + + def data_ptr(self): + return self.pointer.value + + def reset(self): + self.runtime.check( + self.runtime.library.txMemcpy( + self.pointer, + ctypes.c_void_p(self.host.ctypes.data), + self.host.nbytes, + 1, + ), + "reset output H2D", + ) + + def cpu(self): + result = np.empty_like(self.host) + self.runtime.check( + self.runtime.library.txMemcpy( + ctypes.c_void_p(result.ctypes.data), self.pointer, result.nbytes, 2 + ), + "D2H", + ) + return result + + +def invoke(kernel, args, grid, iterations, reset, verify): + calls = [] + saved = (knobs.runtime.launch_enter_hook, knobs.runtime.launch_exit_hook) + knobs.runtime.launch_enter_hook = lambda metadata: calls.append( + ("enter", metadata.get()["name"]) + ) + knobs.runtime.launch_exit_hook = lambda metadata: calls.append( + ("exit", metadata.get()["name"]) + ) + try: + reset() + kernel[grid](*args) + verify() + launcher, module = kernel._run, kernel.module + memory_before = free_device_memory() + elapsed = 0.0 + for _ in range(iterations - 1): + reset() + started = time.monotonic() + kernel[grid](*args) + elapsed += time.monotonic() - started + verify() + print( + f"repeat launch mean={elapsed / (iterations - 1) * 1000:.3f}ms; device free memory change={free_device_memory() - memory_before} bytes", + flush=True, + ) + assert ( + kernel._run is launcher and kernel.module is module and module is not None + ) + assert calls == [ + (event, kernel.name) + for _ in range(iterations) + for event in ("enter", "exit") + ] + finally: + knobs.runtime.launch_enter_hook, knobs.runtime.launch_exit_hook = saved + + +def free_device_memory(): + library = ctypes.CDLL(str(Path(os.environ["KUIPER_ROOT"]) / "lib/libhpgr.so")) + query = library.txMemGetInfo + query.argtypes = [ctypes.POINTER(ctypes.c_uint64), ctypes.POINTER(ctypes.c_uint64)] + query.restype = ctypes.c_int + free, total = ctypes.c_uint64(), ctypes.c_uint64() + RawRuntime.check(query(ctypes.byref(free), ctypes.byref(total)), "txMemGetInfo") + return free.value + + +def run_case(case, iterations, torch_mode, compile_only=False): + if case in ("vector", "grid"): + size = 256 if case == "vector" else 700 + hosts = [ + np.arange(size, dtype=np.float32) * 0.5, + np.arange(size, dtype=np.float32)[::-1].copy() * 0.25, + np.zeros(size, dtype=np.float32), + ] + expected = hosts[0] + hosts[1] + np.float32(1.25) + kernel = compile_kernel( + wafer_vector, + dict( + lhs="*fp32", + rhs="*fp32", + output="*fp32", + alpha="fp32", + size="i32", + BLOCK="constexpr", + ), + dict(BLOCK=256), + ) + scalars, grid = [1.25, size], (triton.cdiv(size, 256), 1, 1) + elif case == "reduction": + hosts = [np.arange(1, 257, dtype=np.float32), np.zeros(1, dtype=np.float32)] + expected = np.array([32896.0], dtype=np.float32) + kernel = compile_kernel( + wafer_reduction, + dict(values="*fp32", output="*fp32", size="i32", BLOCK="constexpr"), + dict(BLOCK=256), + ) + scalars, grid = [256], (1, 1, 1) + else: + hosts = [ + np.full((128, 64), 0x3F80, dtype=np.uint16), + np.full((64, 256), 0x3F80, dtype=np.uint16), + np.zeros((128, 256), dtype=np.uint16), + ] + expected = np.full((128, 256), 0x4280, dtype=np.uint16) + kernel = compile_kernel( + wafer_matmul, + dict( + lhs="*bf16", + rhs="*bf16", + output="*bf16", + M="i32", + N="i32", + K="i32", + BM="constexpr", + BN="constexpr", + BK="constexpr", + ), + dict(BM=128, BN=256, BK=64), + ) + scalars, grid = [128, 256, 64], (1, 1, 1) + + if compile_only: + print(f"PASS compile/link/ELF audit: {case}; no device launch", flush=True) + return + + with ExitStack() as stack: + if torch_mode: + import torch + import torch_txda # noqa: F401 + from triton.backends.dicp_triton.wafer_runtime import get_runtime + + assert get_runtime() is torch.txda + buffers = [ + ( + torch.from_numpy(host).view(torch.bfloat16).to("txda") + if host.dtype == np.uint16 + else torch.from_numpy(host).to("txda") + ) + for host in hosts + ] + stream = torch.txda.Stream() + stack.enter_context(torch.txda.stream(stream)) + from triton.runtime import driver + + assert driver.active.get_current_stream(0) == stream.txda_stream + args = buffers + scalars + reset_host = torch.from_numpy(hosts[-1]) + if hosts[-1].dtype == np.uint16: + reset_host = reset_host.view(torch.bfloat16) + + def reset(): + buffers[-1].copy_(reset_host) + + else: + runtime = RawRuntime() + buffers = [stack.enter_context(runtime.buffer(host)) for host in hosts] + # Exercise both integer pointers and data_ptr() objects in one launch. + args = [buffers[0].data_ptr(), *buffers[1:], *scalars] + reset = buffers[-1].reset + + def verify(): + if torch_mode: + stream.synchronize() + output = buffers[-1].cpu() + result = ( + output.view(torch.uint16).numpy() + if hosts[-1].dtype == np.uint16 + else output.numpy() + ) + else: + result = buffers[-1].cpu() + np.testing.assert_array_equal(result, expected) + + invoke(kernel, args, grid, iterations, reset, verify) + print( + f"PASS {case}: {expected.size} elements, grid={grid}, iterations={iterations}, " + f"{'TXDA tensor/non-default stream' if torch_mode else 'raw pointers/data_ptr'}, hooks and every iteration verified", + flush=True, + ) + + +def run_jit(iterations, compile_only=False): + import torch + from triton.runtime import driver + + if compile_only: + active = driver.active + saved = active.get_current_device, active.get_current_stream + active.get_current_device = lambda: 0 + active.get_current_stream = lambda device: None + try: + for size in (700, 1): + # warmup specializes tensor dtypes and arguments without loading + # a device binary. CPU tensors are sufficient for this check. + values = torch.empty(size, dtype=torch.float32) + kernel = wafer_vector.warmup( + values, + values, + values, + 1.25, + size, + BLOCK=256, + grid=(triton.cdiv(size, 256),), + ) + audit_kernel( + kernel.metadata.kernel_path, kernel.metadata.device_log_abi + ) + assert ((4,) in kernel.src.constants) == (size == 1) + print( + f"PASS JIT compile/link/ELF audit: size={size}, constexpr/specialization; no device launch", + flush=True, + ) + finally: + active.get_current_device, active.get_current_stream = saved + return + + import torch_txda # noqa: F401 + + assert driver.active.get_active_torch_device() == torch.device("txda", 0) + stream = torch.txda.Stream() + for size in (700, 1): + host = torch.arange(size, dtype=torch.float32) + lhs, rhs = host.to("txda"), (host * 0.5).to("txda") + output = torch.empty_like(lhs) + prepared = wafer_vector.warmup( + lhs, rhs, output, 1.25, size, BLOCK=256, grid=(triton.cdiv(size, 256),) + ) + audit_kernel(prepared.metadata.kernel_path, prepared.metadata.device_log_abi) + expected = host * 1.5 + 1.25 + reset_host = torch.zeros_like(host) + with torch.txda.stream(stream): + output.copy_(reset_host) + kernel = wafer_vector[(triton.cdiv(size, 256),)]( + lhs, rhs, output, 1.25, size, BLOCK=256 + ) + torch.testing.assert_close(output.cpu(), expected, rtol=0, atol=0) + launcher = kernel._run + memory_before = free_device_memory() + elapsed = 0.0 + for _ in range(iterations - 1): + output.copy_(reset_host) + started = time.monotonic() + cached = wafer_vector[(triton.cdiv(size, 256),)]( + lhs, rhs, output, 1.25, size, BLOCK=256 + ) + elapsed += time.monotonic() - started + assert cached is kernel and cached._run is launcher + torch.testing.assert_close(output.cpu(), expected, rtol=0, atol=0) + stream.synchronize() + print( + f"repeat JIT launch mean={elapsed / (iterations - 1) * 1000:.3f}ms; device free memory change={free_device_memory() - memory_before} bytes", + flush=True, + ) + assert kernel.metadata.device_log_abi == os.environ["WAFER_DEVICE_LOG_ABI"] + print( + f"PASS JIT: size={size}, constexpr=256, iterations={iterations}, specialization/cache reuse and every iteration verified", + flush=True, + ) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--case", + choices=("vector", "grid", "reduction", "gemm", "jit", "all"), + default="all", + ) + parser.add_argument("--iterations", type=int, default=2) + parser.add_argument("--torch", action="store_true") + parser.add_argument( + "--compile-only", + action="store_true", + help="Compile, link and audit without allocating or launching on the device", + ) + args = parser.parse_args() + if args.case == "jit" and not args.torch: + parser.error("--case jit requires --torch") + if args.iterations < 2: + parser.error("--iterations must be at least 2 to verify initialization reuse") + if os.getenv("USE_SIM_MODE") != "0" or os.getenv("WAFER_ENABLE_RUNTIME") != "1": + parser.error( + "Set USE_SIM_MODE=0 and WAFER_ENABLE_RUNTIME=1 before starting Python" + ) + started = time.monotonic() + cases = ( + ("vector", "reduction", "gemm", "grid") if args.case == "all" else (args.case,) + ) + if args.case == "all" and args.torch: + cases += ("jit",) + for case in cases: + if case == "jit": + run_jit(args.iterations, args.compile_only) + else: + run_case(case, args.iterations, args.torch, args.compile_only) + print( + f"{'Compile-only checks' if args.compile_only else 'Hardware acceptance'} passed in {time.monotonic() - started:.2f}s", + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/test/wafer/verify_wafer_runtime.sh b/test/wafer/verify_wafer_runtime.sh new file mode 100755 index 00000000..aa10e2f4 --- /dev/null +++ b/test/wafer/verify_wafer_runtime.sh @@ -0,0 +1,18 @@ +#!/usr/bin/env bash + +WAFER_WORKSPACE_ROOT=${WAFER_WORKSPACE_ROOT:-} + +set -euo pipefail + +SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd) +source "$SCRIPT_DIR/../../scripts/wafer/init_wafer_env.sh" +export USE_SIM_MODE=0 WAFER_ENABLE_RUNTIME=1 +export TRITON_CACHE_DIR=${TRITON_CACHE_DIR:-$WAFER_WORKSPACE_ROOT/build/wafer-acceptance/cache} +export TRITON_DUMP_PATH=${TRITON_DUMP_PATH:-$WAFER_WORKSPACE_ROOT/build/wafer-acceptance/dump} +RESULT_DIR=${WAFER_RESULT_DIR:-$WAFER_WORKSPACE_ROOT/build/wafer-acceptance} +mkdir -p "$RESULT_DIR" "$TRITON_CACHE_DIR" "$TRITON_DUMP_PATH" +cd "$WAFER_WORKSPACE_ROOT" +LOG_PATH="$RESULT_DIR/acceptance-$(date +%Y%m%d-%H%M%S-%N).log" +echo "Acceptance log: $LOG_PATH" +timeout --signal=TERM --kill-after=5s "${WAFER_TEST_TIMEOUT:-120}s" \ + "$PYTHON" "$SCRIPT_DIR/verify_wafer_runtime.py" "$@" 2>&1 | tee "$LOG_PATH" diff --git a/test/wafer/verify_wafer_torch_stack.py b/test/wafer/verify_wafer_torch_stack.py new file mode 100755 index 00000000..a13023d1 --- /dev/null +++ b/test/wafer/verify_wafer_torch_stack.py @@ -0,0 +1,51 @@ +#!/usr/bin/env python3 +"""Validate the paired torch/torch_txda/txops release with native TXDNN calls.""" + +import argparse +import importlib.metadata + +import torch +import torch_txda # noqa: F401 +import txops + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--matmul", + action="store_true", + help="Also reproduce the vendor BF16 matmul case (use an external timeout)", + ) + args = parser.parse_args() + for name in ("torch", "torch-txda", "txops", "triton"): + print(f"{name}: {importlib.metadata.version(name)}", flush=True) + assert torch.txda.is_available() and torch.txda.device_count() > 0 + host = torch.arange(256, dtype=torch.float32).reshape(1, 256) + torch.testing.assert_close(host.to("txda").cpu(), host, rtol=0, atol=0) + handle = txops.txdnn.create() + layout = txops.txdnn.txLayout.NCX + # Use TXDNN's descriptor API and compare its returned host tensor. + lhs = txops.txdnn.tensor_like(handle, host, layout) + rhs = txops.txdnn.tensor_like(handle, torch.ones_like(host) * 2, layout) + result = txops.txdnn.add(handle, lhs, rhs).tensor() + torch.testing.assert_close(result, host + 2, rtol=0, atol=0) + print("PASS native txdnn.add: 256 FP32 elements", flush=True) + try: + txops.txdnn.add(handle, object(), rhs) + except TypeError: + print("PASS native txdnn invalid argument propagation", flush=True) + else: + raise AssertionError("TXDNN accepted an invalid tensor descriptor") + if not args.matmul: + return + a = torch.ones((128, 64), dtype=torch.bfloat16) + b = torch.ones((64, 256), dtype=torch.bfloat16) + da, db = (txops.txdnn.tensor_like(handle, value, layout) for value in (a, b)) + print("Launching vendor txdnn.matmul", flush=True) + result = txops.txdnn.matmul(handle, da, db).tensor() + torch.testing.assert_close(result, a @ b, rtol=0, atol=0) + print("PASS native txdnn.matmul: BF16 128x64 @ 64x256", flush=True) + + +if __name__ == "__main__": + main() diff --git a/third_party/wafer/CMakeLists.txt b/third_party/wafer/CMakeLists.txt new file mode 100755 index 00000000..7a0c0bee --- /dev/null +++ b/third_party/wafer/CMakeLists.txt @@ -0,0 +1,161 @@ +set(WAFER_BUILD_ROLE "tools" CACHE STRING "Wafer component: frontend or tools") +set_property(CACHE WAFER_BUILD_ROLE PROPERTY STRINGS frontend tools) + +# The Wafer Python extension contains only the TTIR/DSA frontend. FLIR's +# dialects and lowering live in the separate wafer-opt build. Neither build +# links the original DLCompiler C++ plugin. +if(WAFER_BUILD_ROLE STREQUAL "frontend") + if(NOT TRITON_BUILD_PYTHON_MODULE) + message(FATAL_ERROR "Wafer frontend requires TRITON_BUILD_PYTHON_MODULE") + endif() + set(WAFER_TLE_BUILD_CONVERSIONS OFF) + add_subdirectory(third_party/tle) + add_triton_plugin(TritonWafer + ${CMAKE_CURRENT_SOURCE_DIR}/python/triton_wafer_frontend.cc + LINK_LIBS TleDsaIR MLIRArithDialect MLIRLinalgDialect MLIRTensorDialect + MLIRVectorDialect MLIRFuncDialect MLIRFuncAllExtensions + MLIRArithTransforms MLIRBufferizationTransforms + MLIRLinalgTransforms MLIRSCFTransforms MLIRTensorTransforms) + target_link_libraries(TritonWafer PRIVATE Python3::Module pybind11::headers) + return() +elseif(NOT WAFER_BUILD_ROLE STREQUAL "tools") + message(FATAL_ERROR "WAFER_BUILD_ROLE must be frontend or tools") +endif() + +if(TARGET TritonToLinalg OR TARGET TritonStructuredIR) + message(FATAL_ERROR + "Wafer tools require a separate build tree without the DICP plugin. " + "Build the original DICP plugin and Wafer in separate build trees.") +endif() + +include("${CMAKE_CURRENT_LIST_DIR}/cmake/WaferConfig.cmake") + +if(NOT DEFINED WAFER_DEPS_ROOT) + if(DEFINED ENV{WAFER_DEPS_ROOT}) + set(WAFER_DEPS_ROOT $ENV{WAFER_DEPS_ROOT}) + else() + message(STATUS "WAFER_DEPS_ROOT not set, CRT/Profiler will be disabled") + set(WAFER_ENABLE_CRT OFF) + endif() +endif() + +# Check if WAFER_DEPS_ROOT exists to enable CRT/Profiler +if(DEFINED WAFER_DEPS_ROOT AND EXISTS "${WAFER_DEPS_ROOT}") + set(WAFER_ENABLE_CRT ON) +else() + set(WAFER_ENABLE_CRT OFF) + if(DEFINED WAFER_DEPS_ROOT) + message(WARNING "WAFER_DEPS_ROOT set but path does not exist: ${WAFER_DEPS_ROOT}") + endif() +endif() +# Enable ccache if available +find_program(CCACHE_PROGRAM ccache) +if(CCACHE_PROGRAM) + set(CMAKE_C_COMPILER_LAUNCHER ${CCACHE_PROGRAM}) + set(CMAKE_CXX_COMPILER_LAUNCHER ${CCACHE_PROGRAM}) +endif() + +if(NOT DEFINED USE_HOST_PROFILE) + if(DEFINED ENV{USE_HOST_PROFILE}) + set(USE_HOST_PROFILE $ENV{USE_HOST_PROFILE}) + endif() +endif() + +if(NOT DEFINED ENABLE_PROFILING) + if(DEFINED ENV{ENABLE_PROFILING}) + set(ENABLE_PROFILING $ENV{ENABLE_PROFILING}) + endif() +endif() + +if(NOT DEFINED NO_INTRNISIC_RUN) + if(DEFINED ENV{NO_INTRNISIC_RUN}) + set(NO_INTRNISIC_RUN $ENV{NO_INTRNISIC_RUN}) + endif() +endif() + +if(NOT DEFINED ENABLE_SYNCHRONOUS_INTRINSIC) + if(DEFINED ENV{ENABLE_SYNCHRONOUS_INTRINSIC}) + set(ENABLE_SYNCHRONOUS_INTRINSIC $ENV{ENABLE_SYNCHRONOUS_INTRINSIC}) + endif() +endif() + +set(CMAKE_EXPORT_COMPILE_COMMANDS ON) + +# Set FLAGTREE_BACKEND for FLIR to skip test/tools directories +set(FLAGTREE_BACKEND "wafer") + +set(XUANTIE_NAME Xuantie-900-gcc-elf-newlib-x86_64-V2.10.2) +set(INSTALL_WAFER_DIR ${CMAKE_INSTALL_PREFIX}/triton/backends/wafer/) +install(CODE "file(MAKE_DIRECTORY \"${INSTALL_WAFER_DIR}\")") + +include_directories(${CMAKE_CURRENT_SOURCE_DIR}/include) +include_directories(${CMAKE_CURRENT_BINARY_DIR}/include) +include_directories(${CMAKE_CURRENT_SOURCE_DIR}/crt/include) + +if(NOT DEFINED WAFER_SDK_INCLUDE_DIR OR + NOT EXISTS "${WAFER_SDK_INCLUDE_DIR}/instr_def.h") + message(FATAL_ERROR + "WAFER_SDK_INCLUDE_DIR must contain instr_def.h for Wafer lowering") +endif() +include_directories(${WAFER_SDK_INCLUDE_DIR}) + +# WAFER_DEPS_ROOT include only if CRT enabled +if(WAFER_ENABLE_CRT) + include_directories(${WAFER_DEPS_ROOT}/include) +endif() + +# FLIR (FlagTree Intermediate Representation) as local submodule +# FLIR provides TritonTilingExt and TritonStructured dialects +set(FLIR_SOURCE_DIR ${CMAKE_CURRENT_SOURCE_DIR}/third_party/flir) +set(FLIR_BINARY_DIR ${CMAKE_CURRENT_BINARY_DIR}/third_party/flir) +include_directories(${FLIR_SOURCE_DIR}/include) +include_directories(${FLIR_BINARY_DIR}/include) + +# TLE (Wafer Language Extension) - DSA dialect needed by TLEToMK pass +# Bundled in third_party/tle so wafer is independent of FlagTree +set(TLE_SOURCE_DIR ${CMAKE_CURRENT_SOURCE_DIR}/third_party/tle) +set(TLE_BINARY_DIR ${CMAKE_CURRENT_BINARY_DIR}/third_party/tle) +# Make "tle/include/..." includes work (for lib/Conversion/TLEToMK/TLEToMK.cpp) +include_directories(${CMAKE_CURRENT_SOURCE_DIR}/third_party) +include_directories(${CMAKE_CURRENT_BINARY_DIR}/third_party) +# Make "third_party/tle/include/..." includes work (for bin/RegisterTritonDialects.h) +include_directories(${CMAKE_CURRENT_SOURCE_DIR}) +include_directories(${CMAKE_CURRENT_BINARY_DIR}) + +# When built via FlagTree (cmake source), FlagTree builds TLE itself. +# Pass -DWAFER_USE_EXTERNAL_TLE=ON to let the parent cmake provide TLE. +option(WAFER_USE_EXTERNAL_TLE "Use TLE provided by the parent cmake (e.g. FlagTree)" OFF) +if(NOT WAFER_USE_EXTERNAL_TLE) + add_subdirectory(third_party/tle) +endif() +add_subdirectory(lib) +add_subdirectory(include) +add_subdirectory(third_party/flir) +add_subdirectory(bin) + +# Conditionally add CRT and Profiler +if(WAFER_ENABLE_CRT) + add_subdirectory(crt) + add_subdirectory(profiler) +else() + message(STATUS "Building wafer without CRT/Profiler") +endif() + +if(TRITON_BUILD_PYTHON_MODULE) + # FIXME: Unify the libraries for Wafer into fewer ones + # FLIR libraries (from third_party/flir): + # - TritonTilingExtIR, TritonStructuredIR: Dialect libraries + # - TritonToLinalg, TritonToStructured: Conversion passes + add_triton_plugin(TritonWafer ${CMAKE_CURRENT_SOURCE_DIR}/python/triton_wafer.cc + LINK_LIBS TritonSharedAnalysis TritonSharedAnalysisStructured MagicKernelIR + WaferIR TritonTilingExtIR TritonStructuredIR TritonToCoreDialects + TritonToLinalg WaferTritonToStructured StructuredToMemref LinalgToMagicKernel + TritonArithToLinalg CoreDialectsToMK WaferToLLVM AllocateSharedMemory ExportKernelSymbols + WaferMemrefToLLVM MKToWafer + MLIRFuncAllExtensions) + target_link_libraries(TritonWafer PRIVATE Python3::Module pybind11::headers) +endif() +#if(TRITON_BUILD_UT) +# add_subdirectory(unittest) +#endif() +#add_subdirectory(test) diff --git a/third_party/wafer/README.md b/third_party/wafer/README.md new file mode 100755 index 00000000..f9d687ed --- /dev/null +++ b/third_party/wafer/README.md @@ -0,0 +1,80 @@ +# Local Build and Runtime Configuration + +Repository-level build, environment, and packaging tools live in +`scripts/wafer/`. Standalone `verify_wafer_*` acceptance programs live in +`test/wafer/` alongside pytest tests, but are invoked explicitly rather than +collected by pytest. Device acceptance programs require the local hardware +environment. The packaging entry `setup_on_wafer.py` remains at repository root. + +Prepare the pinned LLVM toolchain, Triton source checkout, Wafer SDK, and +matching Torch/Kuiper runtime locally. Dependency acquisition is outside this +project's installation interface. Missing dependencies must be supplied by the +user; the Wafer setup scripts do not download replacements. + +Set `LLVM_SYSPATH` and `WAFER_DEPS_ROOT`, then run the repository-root +`scripts/wafer/setup_wafer_env.sh`, `scripts/wafer/compile_wafer.sh`, and `scripts/wafer/install_wafer.sh` entries with the +prepared Python environment. `scripts/wafer/migrate_wafer_env.sh --compiler-only` checks local +compiler prerequisites and writes the activation file; it does not install +dependencies. Initialize the pinned `third_party/triton` checkout locally before +building. Python build requirements must already be installed. + +The project target is `GPUTarget("wafer", "wafer", 32)`. Device log configuration +accepts `WAFER_DEVICE_LOG_ABI=wafer` for the SDK's original logging ABI or `rcs` +for the RCS firmware ABI. SDK binary symbols and SDK directory layouts retain +their vendor-defined spelling. Torch device allocation continues to use the +matching `torch_txda` extension and its registered `txda` device. + +Set `WAFER_BSP_INCLUDE_DIR` when the local SDK provides BSP headers in a +different directory. An explicit CMake value takes precedence over the +environment. Otherwise the vendor SDK's default BSP layout is used. + +# Wafer Triton Plugin + +## Current Status + +This directory contains the Wafer external Triton plugin imported for the +DLCompiler Triton 3.5 migration. The import includes the bundled FLIR and TLE +sources and excludes prior build trees, generated IR, binaries, caches, and +machine-specific paths. + +The plugin is adapted to LLVM/MLIR 22 and connected to the combined DLCompiler +and Wafer plugin build. A real Triton BF16 GEMM containing `tl.dot` is verified +through TTIR, CoreIR, TXIR, LLVM dialect IR, LLVM IR, and object generation. + +The integration includes Kuiper linking, a device launcher, and native +`torch_txda` acceptance checks. Device execution requires a compatible local +runtime, firmware, and available device. Compiler-only checks do not establish +device correctness. `test/wafer/verify_wafer_runtime.py --compile-only --case all` +checks compilation, linking, and ELF structure without launching a kernel; +device acceptance must be run separately in the prepared runtime environment. + +## Source Layout + +- `backend/`: Python backend implementation imported for later adaptation +- `bin/`: Wafer command-line tools and dialect registration +- `include/` and `lib/`: dialects, analyses, and conversion passes +- `crt/` and `profiler/`: optional Wafer runtime components +- `third_party/flir/`: bundled FLIR sources +- `third_party/tle/`: bundled TLE sources + +## Optional Wafer Dependencies + +`WAFER_DEPS_ROOT` is optional at plugin configuration time. When it is unset or +points to a missing directory, the top-level CMake configuration skips both +`crt/` and `profiler/`. + +The main compiler and dialect sources remain available without Wafer runtime +dependencies. A later build that enables CRT or the profiler must provide a +valid `WAFER_DEPS_ROOT` and a compatible LLVM toolchain containing Clang. + +## LLVM Environment + +Prepare the Triton-pinned LLVM/MLIR 22 environment from the repository root: + +```bash +./scripts/wafer/setup_llvm22_env.sh +source llvm22_env.sh +``` + +Device code generation and linking require a matching LLVM RISC-V toolchain +and Wafer runtime libraries in addition to the compiler-only package. diff --git a/third_party/wafer/backend/__init__.py b/third_party/wafer/backend/__init__.py new file mode 100755 index 00000000..89ac7c87 --- /dev/null +++ b/third_party/wafer/backend/__init__.py @@ -0,0 +1,8 @@ +"""External Wafer plugin metadata; execution uses the DICP Wafer backend.""" + +from .compiler import WaferExternalBackend +from .driver import WaferExternalDriver +from .logger_config import setup_logger +from . import wafer_tools + +__all__ = ["WaferExternalBackend", "WaferExternalDriver", "setup_logger", "wafer_tools"] diff --git a/third_party/wafer/backend/compiler.py b/third_party/wafer/backend/compiler.py new file mode 100755 index 00000000..7b03c166 --- /dev/null +++ b/third_party/wafer/backend/compiler.py @@ -0,0 +1,27 @@ +from triton.backends.compiler import BaseBackend + + +class WaferExternalBackend(BaseBackend): + binary_ext = "o" + + @classmethod + def supports_target(cls, target): + return False + + def hash(self): + return "wafer-external" + + def parse_options(self, options): + return options + + def add_stages(self, stages, options): + raise RuntimeError( + "The wafer runtime backend is installed by setup_on_wafer.py as " + "triton.backends.dicp_triton." + ) + + def load_dialects(self, context): + pass + + def get_module_map(self): + return {} diff --git a/third_party/wafer/backend/driver.py b/third_party/wafer/backend/driver.py new file mode 100755 index 00000000..98d965c5 --- /dev/null +++ b/third_party/wafer/backend/driver.py @@ -0,0 +1,20 @@ +from triton.backends.compiler import GPUTarget +from triton.backends.driver import DriverBase + + +class WaferExternalDriver(DriverBase): + @classmethod + def is_active(cls): + return False + + def get_current_target(self): + return GPUTarget("wafer", 0, 32) + + def get_active_torch_device(self): + return "cpu" + + def get_benchmarker(self): + raise RuntimeError( + "The wafer runtime backend is installed by setup_on_wafer.py as " + "triton.backends.dicp_triton." + ) diff --git a/third_party/wafer/backend/include/logger.h b/third_party/wafer/backend/include/logger.h new file mode 100755 index 00000000..ab6e05b8 --- /dev/null +++ b/third_party/wafer/backend/include/logger.h @@ -0,0 +1,106 @@ +#pragma once +#include +#include +#include +#include +#include + +namespace simple_logger { + +enum LogLevel { DEBUG = 0, INFO = 1, WARN = 2, ERROR = 3 }; + +class Logger { +public: + explicit Logger(LogLevel level = INFO) : current_level_(level) {} + + void setLogLevel(LogLevel level) { current_level_ = level; } + + void log(LogLevel level, const char *format, ...) { + if (level < current_level_) { + return; + } + + va_list args; + va_start(args, format); + int msg_len = std::vsnprintf(nullptr, 0, format, args); + va_end(args); + + if (msg_len <= 0) { + const char *err = "<>"; + msg_len = static_cast(std::strlen(err)); + } + + auto now = std::chrono::system_clock::now(); + auto now_ms = std::chrono::time_point_cast(now); + auto ms = now_ms.time_since_epoch() % 1000; + + std::time_t now_time_t = std::chrono::system_clock::to_time_t(now_ms); + std::tm tm{}; + localtime_r(&now_time_t, &tm); + + char time_buf[64]; + std::strftime(time_buf, sizeof(time_buf), "%Y%m%d %H:%M:%S", &tm); + size_t time_len = std::strlen(time_buf); + std::snprintf(time_buf + time_len, sizeof(time_buf) - time_len, ".%03d", + static_cast(ms.count())); + time_len += 4; + + const char *level_prefix; + switch (level) { + case DEBUG: + level_prefix = "[DEBUG] "; + break; + case INFO: + level_prefix = "[INFO ] "; + break; + case WARN: + level_prefix = "[WARN ] "; + break; + case ERROR: + level_prefix = "[ERROR] "; + break; + default: + level_prefix = "[?????] "; + } + size_t level_len = std::strlen(level_prefix); + + size_t total_add_len = + 1 + time_len + 2 + level_len + + static_cast( + msg_len > 0 ? msg_len + : static_cast(std::strlen("<>"))) + + 1; + + std::string buffer; + buffer.resize(total_add_len); + char *out = &buffer[0]; + + char *p = out; + *p++ = '['; + std::memcpy(p, time_buf, time_len); + p += time_len; + *p++ = ']'; + + std::memcpy(p, level_prefix, level_len); + p += level_len; + + if (msg_len > 0) { + va_start(args, format); + std::vsnprintf(p, static_cast(msg_len) + 1, format, args); + va_end(args); + p += msg_len; + } else { + const char *err = "<>"; + std::memcpy(p, err, std::strlen(err)); + p += std::strlen(err); + } + + *p++ = '\n'; + printf("%s", buffer.c_str()); + } + +private: + LogLevel current_level_; +}; + +} // namespace simple_logger diff --git a/third_party/wafer/backend/logger_config.py b/third_party/wafer/backend/logger_config.py new file mode 100755 index 00000000..f347b208 --- /dev/null +++ b/third_party/wafer/backend/logger_config.py @@ -0,0 +1,136 @@ +# logger_config.py +import logging +import os +import sys + +# Standard mapping: custom number (0~4) -> logging constant +CUSTOM_NUMBER_TO_LOGGING = { + 0: logging.DEBUG, 1: logging.INFO, 2: logging.WARNING, 3: logging.ERROR, 4: logging.CRITICAL +} + +# Reverse for validation +LOGGING_TO_CUSTOM_NUMBER = {v: k for k, v in CUSTOM_NUMBER_TO_LOGGING.items()} + +# Standard level names for validation +STANDARD_LEVEL_NAMES = { + name.upper(): getattr(logging, name) + for name in ['DEBUG', 'INFO', 'WARNING', 'ERROR', 'CRITICAL'] +} + + +def get_log_level_from_env(env_var='WAFER_LOG_LEVEL', default='info'): + """ + Read log level from environment variable. + Supports: + - String: 'DEBUG', 'info', 'WARNING', etc. (case-insensitive) + - Number: '0', '1', '2', '3', '4' + Returns a standard logging level integer (e.g., logging.INFO = 20). + """ + fallback = os.getenv("TX_LOG_LEVEL", default) if env_var == "WAFER_LOG_LEVEL" else default + raw_value = os.getenv(env_var, fallback).strip() + if not raw_value: + raw_value = default + + # Try to interpret as integer (custom 0~4) + try: + num = int(raw_value) + if num in CUSTOM_NUMBER_TO_LOGGING: + return CUSTOM_NUMBER_TO_LOGGING[num] + else: + print(f"Warning: Invalid numeric log level '{num}'. Must be 0~4. Using default '{default}'.") + except ValueError: + # Not an integer, treat as string + level_str = raw_value.upper() + if level_str in STANDARD_LEVEL_NAMES: + return STANDARD_LEVEL_NAMES[level_str] + else: + print(f"Warning: Invalid log level string '{raw_value}'. " + f"Expected one of {list(STANDARD_LEVEL_NAMES.keys())} or 0~4. Using default '{default}'.") + + # Fallback to default + fallback_level_str = str(default).upper() + if fallback_level_str in STANDARD_LEVEL_NAMES: + return STANDARD_LEVEL_NAMES[fallback_level_str] + elif default.isdigit(): + fallback_num = int(default) + if fallback_num in CUSTOM_NUMBER_TO_LOGGING: + return CUSTOM_NUMBER_TO_LOGGING[fallback_num] + # Final fallback + return logging.INFO + + +# Custom log level mapping: standard level names to custom integers (0~4) +CUSTOM_LEVEL_MAP = {'DEBUG': 0, 'INFO': 1, 'WARNING': 2, 'ERROR': 3, 'CRITICAL': 4} + + +def log_level_name_to_custom_number(level_name: str) -> int: + """ + Convert a standard log level name (e.g., 'INFO') to a custom integer (0~4). + Case-insensitive. + """ + level_name = level_name.upper() + if level_name not in CUSTOM_LEVEL_MAP: + raise ValueError(f"Unsupported log level: {level_name}") + return CUSTOM_LEVEL_MAP[level_name] + + +def logger_to_custom_level_number(logger) -> int: + """ + Get the effective log level of the given logger and convert it to a custom integer (0~4). + """ + effective_level = logger.getEffectiveLevel() + level_name = logging.getLevelName(effective_level) + + # Handle non-standard or unrecognized levels + if not isinstance(level_name, str) or level_name.startswith("Level "): + raise ValueError(f"Unrecognized log level value: {effective_level}") + + return log_level_name_to_custom_number(level_name) + + +def log_at_current_level(logger, message): + current_level = logger.getEffectiveLevel() + if current_level <= logging.DEBUG: + logger.debug(message) + elif current_level <= logging.INFO: + logger.info(message) + elif current_level <= logging.WARNING: + logger.warning(message) + elif current_level <= logging.ERROR: + logger.error(message) + else: + logger.critical(message) + + +def setup_logger(name='wafer'): + """ + Set up and return a unified logger instance. + Log level is controlled by the LOG_LEVEL environment variable (default: INFO). + Logs are output to both console and file. + """ + log_file = f"{name}.log" + logger = logging.getLogger(name) + + if not logger.handlers: + log_level = get_log_level_from_env() + + # Set logger to lowest level; actual filtering is done by handlers + logger.setLevel(log_level) + + formatter = logging.Formatter(fmt='[%(asctime)s.%(msecs)03d][%(levelname)s]%(name)s:%(message)s', + datefmt='%Y%m%d %H:%M:%S') + + # File handler + file_handler = logging.FileHandler(log_file, encoding='utf-8') + file_handler.setLevel(log_level) + file_handler.setFormatter(formatter) + + # Console handler + console_handler = logging.StreamHandler(sys.stdout) + console_handler.setLevel(log_level) + console_handler.setFormatter(formatter) + + logger.addHandler(file_handler) + logger.addHandler(console_handler) + + return logger diff --git a/third_party/wafer/backend/name.conf b/third_party/wafer/backend/name.conf new file mode 100755 index 00000000..89a6ae8d --- /dev/null +++ b/third_party/wafer/backend/name.conf @@ -0,0 +1 @@ +wafer diff --git a/third_party/wafer/backend/txda_tools.py b/third_party/wafer/backend/txda_tools.py new file mode 100644 index 00000000..885ddb59 --- /dev/null +++ b/third_party/wafer/backend/txda_tools.py @@ -0,0 +1,2 @@ +"""Compatibility import for the renamed Wafer helper module.""" +from .wafer_tools import * # noqa: F401,F403 diff --git a/third_party/wafer/backend/wafer_tools.py b/third_party/wafer/backend/wafer_tools.py new file mode 100755 index 00000000..5025340f --- /dev/null +++ b/third_party/wafer/backend/wafer_tools.py @@ -0,0 +1,195 @@ +import os +import shutil +import subprocess +import hashlib + +from .logger_config import setup_logger + +logger = setup_logger("wafer_launch") + +_dump_dir_cache = None +dump_cmd_count = 0 + + +def _get_dump_env_path(): + path = os.getenv("TRITON_DUMP_PATH", "") + if not path: + return "" + os.makedirs(path, exist_ok=True) + return path + + +def get_dump_dir(): + global _dump_dir_cache + # 如果已缓存,直接返回 + if _dump_dir_cache is not None: + return _dump_dir_cache + + base_dir = _get_dump_env_path() + # 查找第一个不存在的dumpN目录 + index = 1 + while True: + dir_name = f"dump{index}" + full_path = os.path.join(base_dir, dir_name) + if not os.path.exists(full_path): + os.makedirs(full_path) + _dump_dir_cache = full_path # 缓存结果 + logger.debug(f"make dump dir:{full_path}") + break + index += 1 + return _dump_dir_cache + + +def runLoweringCmd(destFile: str, args: list): + isAlwaysCompile = os.getenv("TRITON_ALWAYS_COMPILE", "0").lower() in ("1", "true", "yes") + if isAlwaysCompile or not os.path.exists(destFile): + if os.getenv("MLIR_ENABLE_DUMP", "0") == "1": + subprocess.check_call(args, stderr=subprocess.STDOUT) + else: + subprocess.check_call(args, stdout=subprocess.DEVNULL) + else: + logger.debug(f"Skip lowering {destFile}") + + +def is_use_profile(): + return os.getenv("ENABLE_PROFILING", "0").lower() in ("1", "true", "yes") + + +def is_enable_kernel_file_cache(): + return os.getenv("ENABLE_KERNEL_FILE_CACHE", "1").lower() in ("1", "true", "yes") + + +def get_kernel_cache_size(): + if is_enable_kernel_file_cache(): + kernel_size_str = os.getenv("KERNEL_FILE_SIZE", "1024") + try: + kernel_size = int(kernel_size_str) + return str(kernel_size) + except ValueError: + raise ValueError(f"Illegal input KERNEL_FILE_SIZE '{kernel_size}', need Integer number.") + else: + raise ValueError("Must set ENABLE_KERNEL_FILE_CACHE=1 first") + + +def dump_ir_if_needed(files): + path = get_dump_dir() + for f in files: + shutil.copy(f, os.path.join(path, os.path.basename(f))) + + +def dump_file_if_needed(src_file, dest_file_name): + path = get_dump_dir() + shutil.copy(src_file, os.path.join(path, dest_file_name)) + + +def dump_cmd_if_needed(cmd: list, flag: str): + path = get_dump_dir() + global dump_cmd_count + if dump_cmd_count == 0: + open_type = 'w' + else: + open_type = 'a' + file_path = os.path.join(path, "cmds.txt") + str_cmd = ' '.join(map(str, cmd)) + dump_cmd = f"{flag}:{str_cmd}\n\n" + + # 将字符串写入指定文件 + with open(file_path, open_type, encoding='utf-8') as f: + f.write(dump_cmd) + dump_cmd_count += 1 + + +def get_llvm_bin_path(bin_name: str) -> str: + path = os.getenv("LLVM_BINARY_DIR", "") + if path == "": + raise Exception("LLVM_BINARY_DIR is not set.") + return os.path.join(path, bin_name) + + +def get_tsm_opt_path() -> str: + """ + Find wafer-opt binary in the following order: + 1. Relative to backend directory (after setup_on_wafer.py install) + 2. WAFER_BUILD_DIR environment variable + 3. Default build_manual location (for development) + 4. System PATH + """ + # 1. Check relative to backend (installed package) + backend_dir = os.path.dirname(__file__) + installed_path = os.path.join(backend_dir, "bin", "wafer-opt") + if os.path.exists(installed_path): + return installed_path + + # 2. Check WAFER_BUILD_DIR environment variable + build_dir = os.getenv("WAFER_BUILD_DIR", "") + if build_dir: + env_path = os.path.join(build_dir, "bin", "wafer-opt") + if os.path.exists(env_path): + return env_path + + # 3. Check default build_manual location (development) + # Navigate from third_party/wafer/backend/ to build_manual/ + dev_path = os.path.join(backend_dir, "..", "..", "build_manual", "install", "bin", "wafer-opt") + dev_path = os.path.abspath(dev_path) + if os.path.exists(dev_path): + return dev_path + + # 4. Check PATH + for path_dir in os.getenv("PATH", "").split(os.pathsep): + path_binary = os.path.join(path_dir, "wafer-opt") + if os.path.exists(path_binary): + return path_binary + + raise RuntimeError( + "wafer-opt not found!\n" + "Searched:\n" + f" - {installed_path}\n" + f" - {dev_path}\n" + "Please either:\n" + " 1. Run setup_on_wafer.py install (copies binaries to package)\n" + " 2. Set WAFER_BUILD_DIR to your build output directory\n" + " 3. Run scripts/wafer/compile_wafer.sh to build the tools" + ) + + +def get_wafer_deps_path(sub_name: str) -> str: + path = os.getenv("WAFER_DEPS_ROOT", "") + if path == "": + raise Exception("WAFER_DEPS_ROOT is not set.") + return os.path.join(path, sub_name) + + +def get_kuiper_path(sub_name: str) -> str: + """Get kuiper path from environment variable or use default""" + kuiper_path = os.getenv("KUIPER_PATH", "/usr/local/kuiper") + return os.path.join(kuiper_path, sub_name) + + +def get_wafer_profiler_path() -> str: + path = os.path.join(os.path.dirname(get_tsm_opt_path()), "wafer-profiler") + return path + + +def is_dump_args_profile(): + env_value = os.getenv("DUMP_KERNEL_ARGS", "0").strip() + try: + num_value = int(env_value) + except ValueError: + num_value = 0 + return num_value + + +def is_debug(): + debug_value = os.getenv("DEBUG", "OFF").strip().upper() + return debug_value in ["ON", "TRUE", "1", "YES"] + + +def calculate_str_md5(string: str): + str_hash = hashlib.md5(string).hexdigest() + return str_hash + + +def calculate_file_md5(file_path): + with open(file_path, 'rb') as f: + file_bytes = f.read() + return calculate_str_md5(file_bytes) diff --git a/third_party/wafer/benchmark/benchmark.py b/third_party/wafer/benchmark/benchmark.py new file mode 100755 index 00000000..d4e29076 --- /dev/null +++ b/third_party/wafer/benchmark/benchmark.py @@ -0,0 +1,1668 @@ +#!/usr/bin/env python3 +""" +使用triton的do_bench来测量性能,支持各种GPU设备 + +使用方法: + export PYTHONPATH="${FLAGGEMS_ROOT}/src" + bash third_party/wafer/scripts/run_wafer.sh pytest test_abs_cuda_time.py +""" + +import pytest +import torch +import flag_gems + +# 导入triton的do_bench,triton会自动处理不同设备的兼容性 +try: + import triton + _do_bench = triton.testing.do_bench +except ImportError: + raise ImportError("triton不可用,请安装triton以支持性能测试") + +# 测试配置常量 +BASE_SHAPES = [(256, 256), (4096, 4096), (16384, 16384)] +# BASE_SHAPES = [(256, 256)] +DTYPE = torch.float16 +WARMUP = 1 +REPETITION = 3 + +# 索引和形状限制常量(避免OOM) +MAX_INDEX_SIZE = 64 +MAX_BATCH_SIZE = 8 +MAX_SEQ_LEN = 256 +MAX_KRON_1D_SIZE = 32 +MAX_KRON_2D_SIZE = 16 +MAX_KRON_ND_SIZE = 8 +MAX_ISIN_TEST_ELEMENTS = 100 + +# 所有算子列表(从 op_list.md 提取,按字母顺序排序) +OP_LIST = [ + 'abs', + 'add', + 'all', + 'amax', + 'angle', + 'any', + 'argmax', + 'argmin', + 'arange', + 'bitwise_and', + 'bitwise_not', + 'bitwise_or', + 'cat', + 'concat_and_cache_mla', + 'contiguous', + 'cos', + 'count_nonzero', + 'cross_entropy_loss', + 'cumsum', + 'diag_embed', + 'diagonal', + 'div', + 'dot', + 'dropout', + 'elu', + 'embedding', + 'eq', + 'erf', + 'exp', + 'eye', + 'fill', + 'flash_attention_forward', + 'flash_mla', + 'flip', + 'floor_divide', + 'full', + 'full_like', + 'fused_add_rms_norm', + 'gather', + 'gelu', + 'gelu_and_mul', + 'ge', + 'glu', + 'gt', + 'hstack', + 'index', + 'index_put', + 'index_select', + 'isin', + 'isfinite', + 'isinf', + 'isclose', + 'isnan', + 'layer_norm', + 'le', + 'linspace', + 'log', + 'log_softmax', + 'logical_and', + 'logical_not', + 'logical_or', + 'logical_xor', + 'lt', + 'masked_fill', + 'masked_select', + 'max', + 'maximum', + 'mean', + 'min', + 'minimum', + 'mm', + 'mse_loss', + 'mul', + 'multinomial', + 'mv', + 'matmul', + 'nan_to_num', + 'ne', + 'neg', + 'nll_loss', + 'nonzero', + 'normal', + 'ones', + 'ones_like', + 'outer', + 'pad', + 'pow', + 'prod', + 'rand', + 'rand_like', + 'randn', + 'randn_like', + 'reciprocal', + 'relu', + 'remainder', + 'repeat_interleave', + 'reshape_and_cache', + 'reshape_and_cache_flash', + 'resolve_conj', + 'resolve_neg', + 'rms_norm', + 'rsqrt', + 'rsub', + 'scaled_dot_product_attention', + 'scatter', + 'select', + 'sigmoid', + 'silu', + 'silu_and_mul', + 'slice_scatter', + 'softmax', + 'sort', + 'stack', + 'sub', + 'sum', + 'tanh', + 'threshold', + 'tile', + 'to', + 'topk', + 'triu', + 'unique', + 'vector_norm', + 'vdot', + 'vstack', + 'where', + 'zeros', + 'zeros_like', + #'batch_norm','bitwise_xor', 'conv1d', 'conv2d', 'conv_depthwise2d','cummax', 'cummin', 'diag', 'group_norm', + #, 'index_add','kron', 'lerp', 'log_sigmoid', 'polar', 'quantile', 'randperm', 'var_mean', +] + +# 全局标志,确保只启用一次 +_flag_gems_enabled = False + + +def _get_do_bench(): + """获取do_bench函数,triton会自动处理不同设备的兼容性""" + return _do_bench + + +def _format_shape_str(op_name, inputs, config, shape): + """格式化shape字符串用于显示和报告""" + if config['type'] == 'matrix': + if op_name == 'addmm': + return f"bias{inputs[0].shape},mat1{inputs[1].shape},mat2{inputs[2].shape}" + elif op_name == 'bmm': + return f"mat1{inputs[0].shape},mat2{inputs[1].shape}" + elif op_name == 'mv': + return f"mat{inputs[0].shape},vec{inputs[1].shape}" + else: + return f"mat1{inputs[0].shape},mat2{inputs[1].shape}" + elif config['type'] == 'ternary': + if op_name == 'lerp': + weight_str = str(inputs[2]) if isinstance(inputs[2], (int, float)) else str(inputs[2].shape) + return f"inp{inputs[0].shape},end{inputs[1].shape},weight{weight_str}" + else: + return f"inp1{inputs[0].shape},inp2{inputs[1].shape},inp3{inputs[2].shape}" + elif config['type'] == 'binary': + return f"{inputs[0].shape}" + elif config['type'] == 'multi_input': + return f"inputs{len(inputs[0])}x{inputs[0][0].shape if inputs[0] else '()'}" + elif config['type'] == 'special_constructor': + if op_name == 'arange': + end = shape[0] if isinstance(shape, tuple) and len(shape) > 0 else ( + shape if isinstance(shape, int) else 256) + return f"arange(0, {end})" + elif op_name == 'linspace': + steps = shape[0] if isinstance(shape, tuple) and len(shape) > 0 else ( + shape if isinstance(shape, int) else 256) + return f"linspace(0, 1, {steps})" + else: + return str(shape) + else: + return str(inputs.shape if hasattr(inputs, 'shape') else shape) + + +def _store_result(op_name, shape_str, config, avg_time_ms, avg_time_us, elapsed_time_ms, error=None): + """存储测试结果到pytest._op_perf_results""" + if not hasattr(pytest, '_op_perf_results'): + pytest._op_perf_results = [] + + result = { + 'op': op_name, + 'shape': shape_str, + 'config': str(config.get('extra_args', {})) if config.get('extra_args') else '', + 'dtype': str(DTYPE).split('.')[-1], + 'avg_time_ms': avg_time_ms, + 'avg_time_us': avg_time_us, + 'elapsed_time_ms': elapsed_time_ms, + 'type': config['type'], + } + + if error: + result['error'] = str(error) + + pytest._op_perf_results.append(result) + + +def _ensure_flag_gems_enabled(): + """确保FlagGems已启用,避免重复注册""" + global _flag_gems_enabled + if not _flag_gems_enabled: + try: + flag_gems.enable() + _flag_gems_enabled = True + except RuntimeError as e: + # 如果已经启用,忽略重复注册的错误 + if "already a kernel registered" not in str(e): + raise + _flag_gems_enabled = True + + +def parse_op_list(): + """解析并返回所有算子列表""" + return OP_LIST + + +def get_op_config(op_name): + """ + 为每个算子返回测试配置(shape适配和调用方式) + + Args: + op_name: 算子名称 + + Returns: + dict: 包含type、shapes、extra_args、dtype等配置的字典 + """ + op_name_lower = op_name.lower() + + # 真正需要跳过的算子(不适合性能测试) + skip_ops = { + 'fused_add_rms_norm', # 复合算子 + 'concat_and_cache_mla', # 特殊缓存算子 + 'reshape_and_cache', # 特殊缓存算子 + 'reshape_and_cache_flash', # 特殊缓存算子 + 'flash_attention_forward', # 需要很多参数 + 'flash_mla', # 需要很多参数 + 'scaled_dot_product_attention', # 需要很多参数 + 'nonzero', # 返回indices,不适合性能测试 + 'unique', # 返回indices,不适合性能测试 + 'to', # dtype转换,不适合性能测试 + 'contiguous', # 内存操作,不适合性能测试 + 'resolve_neg', # 特殊算子 + 'resolve_conj', # 特殊算子 + 'gelu_and_mul', # 复合算子 + 'silu_and_mul', # 复合算子 + } + + # 需要固定配置的特殊算子(参考FlagGems/tests中的测试用例) + special_fixed_config_ops = { + 'batch_norm', # 需要 weight, bias, running_mean, running_var 等 + 'layer_norm', # 需要 normalized_shape, weight, bias 等 + 'group_norm', # 需要 num_groups, weight, bias 等 + 'rms_norm', # 可能需要特殊参数 + 'cross_entropy_loss', # 需要 target + 'nll_loss', # 需要 target + 'mse_loss', # 需要 target + 'embedding', # 需要 indices + 'gather', # 需要 indices + 'index', # 需要 indices + 'index_add', # 需要 indices + 'index_put', # 需要 indices + 'index_select', # 需要 index + 'scatter', # 需要 index + 'slice_scatter', # 需要 index + 'multinomial', # 需要 num_samples + 'topk', # 需要 k + 'quantile', # 需要 q + 'where', # 需要 condition + 'masked_fill', # 需要 mask + 'masked_select', # 需要 mask + 'select', # 需要 dim 和 index + 'diagonal', # 需要 offset, dim1, dim2 + 'diag', # 需要 offset + 'diag_embed', # 需要 offset, dim1, dim2 + 'pad', # 需要 pad 参数 + 'tile', # 需要 dims 参数 + 'repeat_interleave', # 需要 repeats 参数 + 'kron', # 需要两个输入 + 'outer', # 需要两个输入 + 'vdot', # 需要两个输入 + 'dot', # 需要两个输入 + 'polar', # 需要两个输入 + 'atan2', # 需要两个输入 + 'hypot', # 需要两个输入 + 'fmod', # 需要两个输入 + 'isin', # 需要test_elements参数 + 'conv1d', # 需要 weight, bias, stride, padding 等 + 'conv2d', # 需要 weight, bias, stride, padding 等 + 'conv_depthwise2d', # 需要 weight, bias, stride, padding 等 + 'fill', # 需要 value 参数 + } + + if op_name_lower in skip_ops: + return {'type': 'skip', 'reason': '不适合性能测试'} + + # 初始化config字典 + config = { + 'type': 'unary', # default + 'shapes': BASE_SHAPES, + 'extra_args': {}, + 'call_func': None, + } + + if op_name_lower in special_fixed_config_ops: + config['type'] = 'special_fixed_config' + # 为每个特殊算子设置固定配置 + if op_name_lower == 'batch_norm': + # batch_norm: (N, C, H, W) -> (N, C, H, W) + config['shapes'] = [(16, 32, 32, 32), (32, 64, 64, 64), (64, 128, 128, 128)] + config['extra_args'] = {'eps': 1e-5, 'momentum': 0.1, 'training': True} + elif op_name_lower == 'layer_norm': + # layer_norm: (N, C, H, W) -> (N, C, H, W) + config['shapes'] = [(16, 32, 32, 32), (32, 64, 64, 64), (64, 128, 128, 128)] + config['extra_args'] = {'eps': 1e-5} + elif op_name_lower == 'group_norm': + # group_norm: (N, C, H, W) -> (N, C, H, W) + config['shapes'] = [(16, 32, 32, 32), (32, 64, 64, 64), (64, 128, 128, 128)] + config['extra_args'] = {'num_groups': 8, 'eps': 1e-5} + elif op_name_lower == 'rms_norm': + # rms_norm: (N, C) -> (N, C) + config['shapes'] = [(256, 512), (4096, 4096), (16384, 16384)] + config['extra_args'] = {'eps': 1e-5} + elif op_name_lower == 'conv1d': + # conv1d: (N, C, L) -> (N, C_out, L_out) + config['shapes'] = [(16, 32, 128), (32, 64, 256), (64, 128, 512)] + config['extra_args'] = {'stride': 1, 'padding': 1, 'dilation': 1} + elif op_name_lower == 'conv2d': + # conv2d: (N, C, H, W) -> (N, C_out, H_out, W_out) + config['shapes'] = [(16, 32, 32, 32), (32, 64, 64, 64), (64, 128, 128, 128)] + config['extra_args'] = {'stride': 1, 'padding': 1, 'dilation': 1, 'groups': 1} + elif op_name_lower == 'conv_depthwise2d': + # conv_depthwise2d: (N, C, H, W) -> (N, C, H_out, W_out) + config['shapes'] = [(16, 32, 32, 32), (32, 64, 64, 64), (64, 128, 128, 128)] + config['extra_args'] = {'stride': 1, 'padding': 1, 'dilation': 1} + elif op_name_lower == 'embedding': + # embedding: (num_embeddings, embedding_dim), indices -> (indices_shape, embedding_dim) + config['shapes'] = [(256, 512), (4096, 4096), (16384, 16384)] + config['extra_args'] = {'num_embeddings': 10000, 'embedding_dim': 512} + elif op_name_lower == 'gather': + # gather: input, dim, index -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'dim': 0} + elif op_name_lower == 'index_select': + # index_select: input, dim, index -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'dim': 0} + elif op_name_lower == 'scatter': + # scatter: input, dim, index, src -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'dim': 0} + elif op_name_lower == 'topk': + # topk: input, k -> (values, indices) + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'k': 10, 'dim': -1} + elif op_name_lower == 'multinomial': + # multinomial: input, num_samples -> indices + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'num_samples': 10} + elif op_name_lower == 'quantile': + # quantile: input, q -> output + # quantile不支持float16,需要使用float32或float64 + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'q': 0.5, 'dim': -1} + config['dtype'] = torch.float32 + elif op_name_lower == 'where': + # where: condition, x, y -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {} + elif op_name_lower == 'masked_fill': + # masked_fill: input, mask, value -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'value': 0.0} + elif op_name_lower == 'masked_select': + # masked_select: input, mask -> output (1D) + config['shapes'] = BASE_SHAPES + config['extra_args'] = {} + elif op_name_lower == 'select': + # select: input, dim, index -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'dim': 0, 'index': 0} + elif op_name_lower == 'diagonal': + # diagonal: input, offset, dim1, dim2 -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'offset': 0, 'dim1': 0, 'dim2': 1} + elif op_name_lower == 'diag': + # diag: input, diagonal -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'diagonal': 0} + elif op_name_lower == 'diag_embed': + # diag_embed: input, offset, dim1, dim2 -> output + config['shapes'] = [(256, ), (4096, ), (16384, )] + config['extra_args'] = {'offset': 0, 'dim1': -2, 'dim2': -1} + elif op_name_lower == 'pad': + # pad: input, pad -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'pad': (1, 1, 1, 1), 'mode': 'constant', 'value': 0.0} + elif op_name_lower == 'tile': + # tile: input, dims -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'dims': (2, 2)} + elif op_name_lower == 'repeat_interleave': + # repeat_interleave: input, repeats -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'repeats': 2, 'dim': -1} + elif op_name_lower == 'kron': + # kron: input1, input2 -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {} + elif op_name_lower == 'outer': + # outer: input1, input2 -> output + config['shapes'] = [(256, ), (4096, ), (16384, )] + config['extra_args'] = {} + elif op_name_lower == 'vdot': + # vdot: input1, input2 -> scalar + config['shapes'] = [(256, ), (4096, ), (16384, )] + config['extra_args'] = {} + elif op_name_lower == 'dot': + # dot: input1, input2 -> scalar or 1D + config['shapes'] = [(256, ), (4096, ), (16384, )] + config['extra_args'] = {} + elif op_name_lower == 'polar': + # polar: abs, angle -> complex + # polar不支持float16,需要使用float32 + config['shapes'] = BASE_SHAPES + config['extra_args'] = {} + config['dtype'] = torch.float32 + elif op_name_lower == 'atan2': + # atan2: input1, input2 -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {} + elif op_name_lower == 'hypot': + # hypot: input1, input2 -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {} + elif op_name_lower == 'fmod': + # fmod: input1, input2 -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {} + elif op_name_lower == 'slice_scatter': + # slice_scatter: input, dim, src, start, end, step -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'dim': 0, 'start': 0, 'end': 128, 'step': 1} + elif op_name_lower == 'isin': + # isin: elements, test_elements -> output + config['shapes'] = [(256, ), (4096, ), (16384, )] + config['extra_args'] = {} + elif op_name_lower == 'fill': + # fill: input, value -> output + config['shapes'] = BASE_SHAPES + config['extra_args'] = {'value': 1.0} + elif op_name_lower == 'cross_entropy_loss': + # cross_entropy_loss: input, target -> loss + config['shapes'] = [(256, 10), (4096, 100), (16384, 1000)] + config['extra_args'] = {} + elif op_name_lower == 'nll_loss': + # nll_loss: input, target -> loss + config['shapes'] = [(256, 10), (4096, 100), (16384, 1000)] + config['extra_args'] = {} + elif op_name_lower == 'mse_loss': + # mse_loss: input, target -> loss + config['shapes'] = BASE_SHAPES + config['extra_args'] = {} + return config + + # 规约算子:保持原shape,默认不带dim(测试全局规约性能) + reduction_ops = { + 'sum', 'mean', 'max', 'min', 'prod', 'all', 'any', 'amax', 'amin', 'argmax', 'argmin', 'std', 'var', 'var_mean', + 'count_nonzero', 'norm', 'vector_norm' + } + + # 累积算子:需要dim参数 + cumulative_ops = {'cummax', 'cummin', 'cumsum'} + + # 位运算算子:需要整数类型 + bitwise_ops = {'bitwise_and', 'bitwise_or', 'bitwise_not', 'bitwise_xor'} + + # 需要特殊数据类型的二元算子 + binary_ops_special_dtype = { + 'floor_divide': torch.float32, # floor_divide不支持float16,需要使用float32 + 'polar': torch.float32, # polar不支持float16,需要使用float32或float64 + } + + # 需要特殊数据类型的构造函数算子 + constructor_ops_special_dtype = { + 'randperm': torch.int64, # randperm只支持整数类型(int16/int32/int64),使用int64 + } + + # 二元算子:需要两个相同shape的输入 + binary_ops = { + 'add', 'sub', 'mul', 'div', 'pow', 'maximum', 'minimum', 'eq', 'ne', 'lt', 'le', 'gt', 'ge', 'remainder', + 'logical_and', 'logical_or', 'logical_xor', 'fmod', 'atan2', 'hypot', 'rsub', # rsub是二元算子(reverse subtract) + 'isclose', # isclose是二元算子,需要两个输入和rtol/atol参数 + } + + # 三元算子:需要三个输入 + ternary_ops = { + 'lerp', # lerp(input, end, weight) 需要三个输入,weight可以是标量或张量 + } + + # 一元逻辑算子 + unary_logical_ops = { + 'logical_not', # logical_not是一元算子 + } + + # 矩阵运算:需要特定shape + # mm: (M, K) x (K, N) -> (M, N) + # bmm: (B, M, K) x (B, K, N) -> (B, M, N) + # addmm: bias + (M, K) x (K, N) -> (M, N) + # mv: (M, N) x (N,) -> (M,) + matrix_ops = { + 'mm': lambda s: (s[0], s[0]), # 返回 (M, N),内部会创建 (M, K) 和 (K, N) + 'bmm': lambda s: (8, s[0], s[0]), # 返回 batch size + 'addmm': lambda s: (s[0], s[0]), # 返回 (M, N) + 'mv': lambda s: (s[0], s[0]), # 返回 (M, N) + 'matmul': lambda s: (s[0], s[0]), # 返回 (M, N) + } + + # 需要特殊参数的算子 + special_ops = { + 'softmax': {'dim': -1}, 'log_softmax': {'dim': -1}, 'dropout': {'p': 0.5, 'train': + True}, # dropout使用train参数,不是training + 'elu': {'alpha': 1.0}, # elu需要alpha参数 + 'flip': {'dims': (0, )}, # flip需要dims参数,默认在dim=0上翻转 + 'gelu': {}, 'silu': {}, 'relu': {}, 'sigmoid': {}, 'tanh': {}, 'exp': {}, 'log': {}, 'sqrt': {}, 'rsqrt': {}, + 'abs': {}, 'neg': {}, 'reciprocal': {}, 'cos': {}, 'sin': {}, 'erf': {}, 'angle': {}, # 一元算子 + 'glu': {}, # 一元算子 + 'log_sigmoid': {}, # 一元算子 + 'isfinite': {}, # 一元算子 + 'isinf': {}, # 一元算子 + 'isnan': {}, # 一元算子 + 'nan_to_num': {}, # 一元算子 + 'logical_not': {}, # 一元逻辑算子 + 'threshold': {'threshold': 0.5, 'value': 0.0}, # threshold需要threshold和value参数 + 'triu': {'diagonal': 0}, # triu需要diagonal参数 + 'sort': {'dim': -1}, # sort需要dim参数 + 'isclose': {'rtol': 1e-5, 'atol': 1e-8}, # isclose需要rtol和atol参数 + } + + # 需要多个输入的算子 + multi_input_ops = { + 'cat': lambda s: ([torch.randn(s, dtype=DTYPE, device=flag_gems.device) for _ in range(3)], {'dim': 0}), + 'stack': lambda s: ([torch.randn(s, dtype=DTYPE, device=flag_gems.device) for _ in range(3)], {'dim': 0}), + 'hstack': lambda s: ([torch.randn(s, dtype=DTYPE, device=flag_gems.device) for _ in range(3)], {}), + 'vstack': lambda s: ([torch.randn(s, dtype=DTYPE, device=flag_gems.device) for _ in range(3)], {}), + } + + # 构造函数算子(不需要输入) + constructor_ops = { + 'ones', 'zeros', 'eye', 'rand', 'randn', 'full', 'empty', 'ones_like', 'zeros_like', 'rand_like', 'randn_like', + 'full_like', 'empty_like', 'normal', # normal需要mean和std参数,但可以设置默认值 + 'randperm', # randperm需要n参数 + } + + # 需要特殊参数的构造函数算子 + special_constructor_ops = {'arange', 'linspace'} + + # 根据算子类型设置配置 + if op_name_lower in reduction_ops: + config['type'] = 'reduction' + # 默认不带dim参数,测试全局规约性能(与原始测试用例 test_accuracy_sum_without_dim 一致) + # vector_norm需要ord参数,默认使用2 + if op_name_lower == 'vector_norm': + config['extra_args'] = {'ord': 2} + else: + config['extra_args'] = {} + elif op_name_lower in cumulative_ops: + config['type'] = 'cumulative' + # 累积算子需要dim参数,默认使用最后一个维度 + config['extra_args'] = {'dim': -1} + elif op_name_lower in bitwise_ops: + config['type'] = 'bitwise' + # 位运算需要整数类型,使用int16 + config['dtype'] = torch.int16 + elif op_name_lower in binary_ops_special_dtype: + config['type'] = 'binary' + # 需要特殊数据类型的二元算子 + config['dtype'] = binary_ops_special_dtype[op_name_lower] + elif op_name_lower in ternary_ops: + config['type'] = 'ternary' + # 三元算子需要三个输入 + elif op_name_lower in binary_ops: + config['type'] = 'binary' + # 如果isclose在special_ops中,需要合并extra_args + if op_name_lower == 'isclose': + config['extra_args'] = special_ops.get('isclose', {}) + elif op_name_lower in unary_logical_ops: + config['type'] = 'unary' + # 一元逻辑算子使用special_ops中的配置 + if op_name_lower in special_ops: + config['extra_args'] = special_ops[op_name_lower] + elif op_name_lower in matrix_ops: + config['type'] = 'matrix' + config['shapes'] = [matrix_ops[op_name_lower](s) for s in BASE_SHAPES] + elif op_name_lower in special_ops: + config['type'] = 'unary' + config['extra_args'] = special_ops[op_name_lower] + elif op_name_lower in multi_input_ops: + config['type'] = 'multi_input' + config['call_func'] = multi_input_ops[op_name_lower] + elif op_name_lower in special_constructor_ops: + config['type'] = 'special_constructor' + # arange和linspace需要特殊参数,不使用shape + config['shapes'] = BASE_SHAPES + elif op_name_lower in constructor_ops: + config['type'] = 'constructor' + # 构造函数使用shape作为输出shape + config['shapes'] = BASE_SHAPES + # 如果构造函数需要特殊数据类型,设置它 + if op_name_lower in constructor_ops_special_dtype: + config['dtype'] = constructor_ops_special_dtype[op_name_lower] + + return config + + +def create_test_inputs(op_name, shape, config): + """ + 根据算子类型和配置创建测试输入 + + Args: + op_name: 算子名称 + shape: 输入shape + config: 算子配置字典 + + Returns: + tensor或tuple: 测试输入,根据算子类型可能是单个tensor或tuple + """ + # device = flag_gems.device + device = "cpu" + # 获取数据类型,如果配置中指定了则使用配置的,否则使用默认的DTYPE + dtype = config.get('dtype', DTYPE) + + if config['type'] == 'bitwise': + # 位运算需要整数类型,使用randint生成 + if op_name == 'bitwise_not': + # 一元位运算 + return torch.randint(low=-0x7FFF, high=0x7FFF, size=shape, dtype=dtype, device=device).to(flag_gems.device) + else: + # 二元位运算 + return (torch.randint(low=-0x7FFF, high=0x7FFF, size=shape, dtype=dtype, + device=device).to(flag_gems.device), + torch.randint(low=-0x7FFF, high=0x7FFF, size=shape, dtype=dtype, + device=device).to(flag_gems.device)) + elif config['type'] == 'ternary': + # 三元算子需要三个输入 + if op_name == 'lerp': + # lerp(input, end, weight) - weight可以是标量或张量,这里使用标量 + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), 0.5 # weight作为标量 + ) + else: + # 其他三元算子 + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), + torch.randn(shape, dtype=dtype, + device=device).to(flag_gems.device), torch.randn(shape, dtype=dtype, + device=device).to(flag_gems.device)) + elif config['type'] == 'binary': + return (torch.randn(shape, dtype=dtype, + device=device).to(flag_gems.device), torch.randn(shape, dtype=dtype, + device=device).to(flag_gems.device)) + elif config['type'] == 'matrix': + if op_name == 'mm': + # mm: (M, K) x (K, N) -> (M, N) + M, N = shape + K = N # 使用N作为K + return (torch.randn((M, K), dtype=DTYPE, + device=device).to(flag_gems.device), torch.randn((K, N), dtype=DTYPE, + device=device).to(flag_gems.device)) + elif op_name == 'bmm': + # bmm: (B, M, K) x (B, K, N) -> (B, M, N) + B, M, N = shape + K = N + return (torch.randn((B, M, K), dtype=DTYPE, + device=device).to(flag_gems.device), torch.randn((B, K, N), dtype=DTYPE, + device=device).to(flag_gems.device)) + elif op_name == 'addmm': + # addmm: bias + (M, K) x (K, N) -> (M, N) + M, N = shape + K = N + return (torch.randn((M, ), dtype=DTYPE, device=device).to(flag_gems.device), # bias + torch.randn((M, K), dtype=DTYPE, device=device).to(flag_gems.device), # mat1 + torch.randn((K, N), dtype=DTYPE, device=device).to(flag_gems.device) # mat2 + ) + elif op_name == 'mv': + # mv: (M, N) x (N,) -> (M,) + M, N = shape + return (torch.randn((M, N), dtype=DTYPE, + device=device).to(flag_gems.device), torch.randn((N, ), dtype=DTYPE, + device=device).to(flag_gems.device)) + elif op_name == 'matmul': + # matmul: 同mm + M, N = shape + K = N + return (torch.randn((M, K), dtype=DTYPE, + device=device).to(flag_gems.device), torch.randn((K, N), dtype=DTYPE, + device=device).to(flag_gems.device)) + elif config['type'] == 'multi_input': + if config['call_func']: + inputs, kwargs = config['call_func'](shape) + return (inputs, kwargs) + elif config['type'] == 'constructor': + # 构造函数:shape作为输出shape参数 + return shape + elif config['type'] == 'special_constructor': + # arange和linspace需要特殊参数,返回None表示需要特殊处理 + return None + elif config['type'] == 'special_fixed_config': + # 特殊固定配置算子,根据算子类型创建相应的输入 + return _create_special_fixed_config_inputs(op_name, shape, config, dtype, device=device) + + # 默认:一元算子或规约算子 + dtype = config.get('dtype', DTYPE) + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) + + +def _create_special_fixed_config_inputs(op_name, shape, config, dtype, device): + """为特殊固定配置算子创建测试输入""" + if op_name == 'batch_norm': + # batch_norm: input (N, C, H, W), weight (C,), bias (C,), running_mean (C,), running_var (C,) + N, C, H, W = shape + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.randn((C, ), dtype=dtype, device=device).to(flag_gems.device), # weight + torch.randn((C, ), dtype=dtype, device=device).to(flag_gems.device), # bias + torch.randn((C, ), dtype=dtype, device=device).to(flag_gems.device), # running_mean + torch.randn((C, ), dtype=dtype, device=device).to(flag_gems.device), # running_var + ) + elif op_name == 'layer_norm': + # layer_norm: input (N, C, H, W), normalized_shape (C, H, W), weight (C*H*W,), bias (C*H*W,) + N, C, H, W = shape + normalized_shape = (C, H, W) + norm_size = C * H * W + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + normalized_shape, torch.randn((norm_size, ), dtype=dtype, device=device).to(flag_gems.device), # weight + torch.randn((norm_size, ), dtype=dtype, device=device).to(flag_gems.device), # bias + ) + elif op_name == 'group_norm': + # group_norm: input (N, C, H, W), num_groups, weight (C,), bias (C,) + N, C, H, W = shape + num_groups = config['extra_args'].get('num_groups', 8) + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + num_groups, torch.randn((C, ), dtype=dtype, device=device).to(flag_gems.device), # weight + torch.randn((C, ), dtype=dtype, device=device).to(flag_gems.device), # bias + ) + elif op_name == 'rms_norm': + # rms_norm: input (N, C), weight (C,) + N, C = shape + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.randn((C, ), dtype=dtype, device=device).to(flag_gems.device), # weight + ) + elif op_name == 'conv1d': + # conv1d: input (N, C, L), weight (C_out, C, K), bias (C_out,) + N, C, L = shape + C_out = C # 输出通道数等于输入通道数 + K = 3 # 卷积核大小 + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.randn((C_out, C, K), dtype=dtype, device=device).to(flag_gems.device), # weight + torch.randn((C_out, ), dtype=dtype, device=device).to(flag_gems.device), # bias (可选) + ) + elif op_name == 'conv2d': + # conv2d: input (N, C, H, W), weight (C_out, C, K, K), bias (C_out,) + N, C, H, W = shape + C_out = C # 输出通道数等于输入通道数 + K = 3 # 卷积核大小 + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.randn((C_out, C, K, K), dtype=dtype, device=device).to(flag_gems.device), # weight + torch.randn((C_out, ), dtype=dtype, device=device).to(flag_gems.device), # bias (可选) + ) + elif op_name == 'conv_depthwise2d': + # conv_depthwise2d: input (N, C, H, W), weight (C, 1, K, K), bias (C,) + N, C, H, W = shape + K = 3 # 卷积核大小 + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.randn((C, 1, K, K), dtype=dtype, device=device).to(flag_gems.device), # weight + torch.randn((C, ), dtype=dtype, device=device).to(flag_gems.device), # bias (可选) + ) + elif op_name == 'embedding': + # embedding: weight (num_embeddings, embedding_dim), indices (Batch, M) + # 根据测试用例,indices应该是较小的2D shape,如(Batch, M),而不是直接使用shape + num_embeddings = config['extra_args'].get('num_embeddings', 4096) + embedding_dim = config['extra_args'].get('embedding_dim', 512) + # 限制indices的shape,避免内存溢出 + # 使用较小的batch和sequence length,参考测试用例中的(Batch, M)格式 + batch_size = min(shape[0] if len(shape) > 0 else 4, MAX_BATCH_SIZE) + seq_len = min(shape[-1] if len(shape) > 1 else 128, MAX_SEQ_LEN) + indices_shape = (batch_size, seq_len) + return (torch.randn((num_embeddings, embedding_dim), dtype=dtype, device=device).to(flag_gems.device), # weight + torch.randint(0, num_embeddings, size=indices_shape, dtype=torch.long, + device=device).to(flag_gems.device), # indices + ) + elif op_name == 'gather': + # gather: input, dim, index -> output + dim = config['extra_args'].get('dim', 0) + index_shape = list(shape) + index_shape[dim] = min(shape[dim], MAX_INDEX_SIZE) + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + dim, torch.randint(0, shape[dim], size=tuple(index_shape), dtype=torch.long, + device=device).to(flag_gems.device), # index + ) + elif op_name == 'index_select': + # index_select: input, dim, index -> output + dim = config['extra_args'].get('dim', 0) + index_size = min(shape[dim], MAX_INDEX_SIZE) + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + dim, torch.randint(0, shape[dim], size=(index_size, ), dtype=torch.long, + device=device).to(flag_gems.device), # index + ) + elif op_name == 'scatter': + # scatter: input, dim, index, src -> output + dim = config['extra_args'].get('dim', 0) + index_shape = list(shape) + index_shape[dim] = min(shape[dim], 64) + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + dim, torch.randint(0, shape[dim], size=tuple(index_shape), dtype=torch.long, + device=device).to(flag_gems.device), # index + torch.randn(tuple(index_shape), dtype=dtype, device=device).to(flag_gems.device), # src + ) + elif op_name == 'slice_scatter': + # slice_scatter: input, dim, src, start, end, step -> output + # slice_scatter(input, dim=dim, src=src, start=start, end=end, step=step) + dim = config['extra_args'].get('dim', 0) + start = config['extra_args'].get('start', 0) + end = config['extra_args'].get('end', shape[dim]) + step = config['extra_args'].get('step', 1) + # 计算 src 的形状 + size = shape[dim] + start = start % size + end = end % (size + 1) + if end < start: + end, start = start, end + elif end == start: + end = size + src_size = (end - start + step - 1) // step + src_shape = list(shape) + src_shape[dim] = src_size + return ( + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + dim, + torch.randn(tuple(src_shape), dtype=dtype, device=device).to(flag_gems.device), # src + start, + end, + step, + ) + elif op_name == 'index': + # index: input, indices -> output + # indices 是一个列表,包含多个索引张量,每个对应输入的一个维度 + # 为了简化,我们为每个维度创建一个索引张量 + indices = [] + for i, dim_size in enumerate(shape): + # 为每个维度创建一个较小的索引张量 + index_size = min(dim_size, MAX_INDEX_SIZE) + indices.append( + torch.randint(0, dim_size, size=(index_size, ), dtype=torch.long, device=device).to(flag_gems.device)) + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + indices, # indices list + ) + elif op_name == 'index_add': + # index_add: input, dim, index, source -> output + # index_add(input, dim, index, source, alpha=1) + dim = config['extra_args'].get('dim', 0) + index_max = shape[dim] + index_len = min(index_max, MAX_INDEX_SIZE) + index = torch.randperm(index_len, device=device).to(flag_gems.device) # 1D索引张量 + src_shape = list(shape) + src_shape[dim] = index_len + source = torch.randn(tuple(src_shape), dtype=dtype, device=device).to(flag_gems.device) + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + dim, index, # index tensor + source, # source tensor + ) + elif op_name == 'index_put': + # index_put: input, indices, values -> output + # index_put(input, indices, values, accumulate=False) + # indices 是一个列表,包含多个索引张量,每个对应输入的一个维度 + indices = [] + for i, dim_size in enumerate(shape): + # 为每个维度创建一个较小的索引张量 + index_size = min(dim_size, MAX_INDEX_SIZE) + indices.append( + torch.randint(0, dim_size, size=(index_size, ), dtype=torch.long, device=device).to(flag_gems.device)) + # values 的形状需要与 indices 广播后的形状匹配 + # 简化处理:使用第一个索引张量的形状作为 values 的形状 + if len(indices) > 0: + values_shape = indices[0].shape + else: + values_shape = (MAX_INDEX_SIZE, ) + values = torch.randn(values_shape, dtype=dtype, device=device) + return (torch.randn(shape, dtype=dtype, device=device), # input + indices, # indices list + values, # values tensor + ) + elif op_name == 'topk': + # topk: input, k, dim -> (values, indices) + k = config['extra_args'].get('k', 10) + dim = config['extra_args'].get('dim', -1) + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) # input + elif op_name == 'multinomial': + # multinomial: input (probs), num_samples -> indices + # input需要是概率分布,每行和为1 + probs = torch.rand(shape, dtype=dtype, device=device).to(flag_gems.device) + probs = probs / probs.sum(dim=-1, keepdim=True) # 归一化为概率分布 + return probs + elif op_name == 'quantile': + # quantile: input, q, dim -> output + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) # input + elif op_name == 'where': + # where: condition, x, y -> output + condition = torch.rand(shape, dtype=dtype, device=device).to(flag_gems.device) > 0.5 + return (condition, torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # x + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # y + ) + elif op_name == 'masked_fill': + # masked_fill: input, mask, value -> output + value = config['extra_args'].get('value', 0.0) + return ( + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.rand(shape, dtype=dtype, device=device).to(flag_gems.device) > 0.5, # mask + value, + ) + elif op_name == 'masked_select': + # masked_select: input, mask -> output (1D) + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.rand(shape, dtype=dtype, device=device).to(flag_gems.device) > 0.5, # mask + ) + elif op_name == 'select': + # select: input, dim, index -> output + dim = config['extra_args'].get('dim', 0) + index = config['extra_args'].get('index', 0) + return ( + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + dim, + index, + ) + elif op_name == 'diagonal': + # diagonal: input, offset, dim1, dim2 -> output + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) # input + elif op_name == 'diag': + # diag: input, diagonal -> output + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) # input + elif op_name == 'diag_embed': + # diag_embed: input (1D), offset, dim1, dim2 -> output + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) # input + elif op_name == 'pad': + # pad: input, pad -> output + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) # input + elif op_name == 'tile': + # tile: input, dims -> output + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) # input + elif op_name == 'repeat_interleave': + # repeat_interleave: input, repeats, dim -> output + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) # input + elif op_name in ['kron', 'outer', 'vdot', 'dot', 'polar', 'atan2', 'hypot', 'fmod']: + # 二元算子:需要两个输入 + if op_name == 'kron': + # kron 的输出大小是输入大小的乘积,需要限制输入大小避免内存溢出 + # 参考测试用例中的 KRON_SHAPES,使用较小的形状 + if len(shape) == 1: + # 1D: 限制大小避免OOM + kron_shape1 = (min(shape[0], MAX_KRON_1D_SIZE), ) + kron_shape2 = (min(shape[0], MAX_KRON_1D_SIZE), ) + elif len(shape) == 2: + # 2D: 限制大小避免OOM + kron_shape1 = (min(shape[0], MAX_KRON_2D_SIZE), min(shape[1], MAX_KRON_2D_SIZE)) + kron_shape2 = (min(shape[0], MAX_KRON_2D_SIZE), min(shape[1], MAX_KRON_2D_SIZE)) + else: + # 高维: 限制每个维度大小避免OOM + kron_shape1 = tuple(min(s, MAX_KRON_ND_SIZE) for s in shape) + kron_shape2 = tuple(min(s, MAX_KRON_ND_SIZE) for s in shape) + return ( + torch.randn(kron_shape1, dtype=dtype, device=device).to(flag_gems.device), + torch.randn(kron_shape2, dtype=dtype, device=device).to(flag_gems.device), + ) + elif op_name in ['outer', 'vdot', 'dot']: + # 这些算子需要1D输入 + return ( + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), + ) + else: + return ( + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), + ) + elif op_name == 'isin': + # isin: elements, test_elements -> output + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # elements + torch.randn((min(MAX_ISIN_TEST_ELEMENTS, shape[0]), ), dtype=dtype, + device=device).to(flag_gems.device), # test_elements + ) + elif op_name == 'fill': + # fill: input, value -> output + value = config['extra_args'].get('value', 1.0) + return ( + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + value, + ) + elif op_name == 'cross_entropy_loss': + # cross_entropy_loss: input (N, C), target (N,) + N, C = shape + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.randint(0, C, size=(N, ), dtype=torch.long, device=device).to(flag_gems.device), # target + ) + elif op_name == 'nll_loss': + # nll_loss: input (N, C), target (N,) + N, C = shape + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.randint(0, C, size=(N, ), dtype=torch.long, device=device).to(flag_gems.device), # target + ) + elif op_name == 'mse_loss': + # mse_loss: input, target -> loss + return (torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # input + torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device), # target + ) + else: + # 默认:一元算子 + return torch.randn(shape, dtype=dtype, device=device).to(flag_gems.device) + + +def call_op(op_name, inputs, config, shape=None): + """ + 调用算子 + + 注意:由于性能测试需要多次调用,使用全局启用的flag_gems.enable() + 而不是每次调用use_gems(),避免重复注册错误。 + + Args: + op_name: 算子名称 + inputs: 输入tensor或tuple + config: 算子配置字典 + shape: 可选的shape参数(用于构造函数) + + Returns: + tensor: 算子输出 + """ + # 某些算子需要使用torch.nn.functional + nn_functional_ops = { + 'dropout': torch.nn.functional.dropout, 'elu': torch.nn.functional.elu, 'relu': torch.nn.functional.relu, + 'gelu': torch.nn.functional.gelu, 'silu': torch.nn.functional.silu, 'sigmoid': torch.nn.functional.sigmoid, + 'tanh': torch.nn.functional.tanh, 'log_sigmoid': + torch.nn.functional.logsigmoid, # log_sigmoid在torch.nn.functional中是logsigmoid + } + + if op_name in nn_functional_ops: + op_func = nn_functional_ops[op_name] + else: + # 特殊处理:vector_norm在torch.linalg中 + if op_name == 'vector_norm': + try: + op_func = torch.linalg.vector_norm + except AttributeError: + # 如果torch.linalg不存在,尝试torch.vector_norm + op_func = getattr(torch, 'vector_norm', None) + else: + op_func = getattr(torch, op_name, None) + if op_func is None: + # 尝试下划线版本 + op_func = getattr(torch, op_name + '_', None) + if op_func is None: + # 尝试torch.nn.functional + op_func = getattr(torch.nn.functional, op_name, None) + if op_func is None: + # 尝试torch.linalg(对于其他linalg算子) + if hasattr(torch, 'linalg'): + op_func = getattr(torch.linalg, op_name, None) + + # 对于special_fixed_config类型的算子,op_func可能为None(在_call_special_fixed_config_op中直接调用) + if op_func is None and config.get('type') != 'special_fixed_config': + raise ValueError(f"未找到算子: {op_name}") + + extra_args = config.get('extra_args', {}) + + # 直接调用,因为flag_gems已经全局启用 + return _call_op_impl(op_func, op_name, inputs, config, extra_args, shape) + + +def _call_special_fixed_config_op(op_func, op_name, inputs, config, extra_args, shape=None): + """调用特殊固定配置算子""" + if op_name == 'batch_norm': + # batch_norm(input, weight, bias, running_mean, running_var, ...) + return torch.nn.functional.batch_norm(inputs[0], inputs[1], inputs[2], inputs[3], inputs[4], **extra_args) + elif op_name == 'layer_norm': + # layer_norm(input, normalized_shape, weight, bias, ...) + return torch.layer_norm(inputs[0], inputs[1], weight=inputs[2], bias=inputs[3], **extra_args) + elif op_name == 'group_norm': + # group_norm(input, num_groups, weight, bias, ...) + # num_groups 已经通过位置参数传递,不应该再通过 extra_args 传递 + filtered_args = {k: v for k, v in extra_args.items() if k != 'num_groups'} + return torch.nn.functional.group_norm(inputs[0], inputs[1], weight=inputs[2], bias=inputs[3], **filtered_args) + elif op_name == 'rms_norm': + # rms_norm(input, weight, ...) + return torch.nn.functional.layer_norm(inputs[0], (inputs[0].shape[-1], ), weight=inputs[1], **extra_args) + elif op_name == 'conv1d': + # conv1d(input, weight, bias=None, ...) + return torch.nn.functional.conv1d(inputs[0], inputs[1], bias=inputs[2], **extra_args) + elif op_name == 'conv2d': + # conv2d(input, weight, bias=None, ...) + return torch.nn.functional.conv2d(inputs[0], inputs[1], bias=inputs[2], **extra_args) + elif op_name == 'conv_depthwise2d': + # conv_depthwise2d(input, weight, bias=None, ...) + return torch.nn.functional.conv2d(inputs[0], inputs[1], bias=inputs[2], groups=inputs[0].shape[1], **extra_args) + elif op_name == 'embedding': + # embedding(indices, weight, ...) + # num_embeddings 和 embedding_dim 只是用于创建 weight 的参数,不应该传递给 embedding 函数 + filtered_args = {k: v for k, v in extra_args.items() if k not in ['num_embeddings', 'embedding_dim']} + return torch.nn.functional.embedding(inputs[1], inputs[0], **filtered_args) + elif op_name == 'gather': + # gather(input, dim, index, ...) + # dim 已经通过位置参数传递,不应该再通过 extra_args 传递 + filtered_args = {k: v for k, v in extra_args.items() if k != 'dim'} + return torch.gather(inputs[0], inputs[1], inputs[2], **filtered_args) + elif op_name == 'index_select': + # index_select(input, dim, index, ...) + # dim 已经通过位置参数传递,不应该再通过 extra_args 传递 + filtered_args = {k: v for k, v in extra_args.items() if k != 'dim'} + return torch.index_select(inputs[0], inputs[1], inputs[2], **filtered_args) + elif op_name == 'scatter': + # scatter(input, dim, index, src, ...) + # dim 已经通过位置参数传递,不应该再通过 extra_args 传递 + filtered_args = {k: v for k, v in extra_args.items() if k != 'dim'} + return torch.scatter(inputs[0], inputs[1], inputs[2], inputs[3], **filtered_args) + elif op_name == 'slice_scatter': + # slice_scatter(input, dim=dim, src=src, start=start, end=end, step=step) + # dim, start, end, step 已经通过位置参数传递,不应该再通过 extra_args 传递 + filtered_args = {k: v for k, v in extra_args.items() if k not in ['dim', 'start', 'end', 'step']} + return torch.slice_scatter(inputs[0], dim=inputs[1], src=inputs[2], start=inputs[3], end=inputs[4], + step=inputs[5], **filtered_args) + elif op_name == 'index': + # index(input, indices) -> output + # indices 是一个列表,包含多个索引张量 + return torch.ops.aten.index(inputs[0], inputs[1]) + elif op_name == 'index_add': + # index_add(input, dim, index, source, alpha=1) -> output + # dim 已经通过位置参数传递,不应该再通过 extra_args 传递 + filtered_args = {k: v for k, v in extra_args.items() if k != 'dim'} + return torch.index_add(inputs[0], inputs[1], inputs[2], inputs[3], **filtered_args) + elif op_name == 'index_put': + # index_put(input, indices, values, accumulate=False) -> output + # indices 是一个列表,包含多个索引张量 + return torch.index_put(inputs[0], inputs[1], inputs[2], **extra_args) + elif op_name == 'topk': + # topk(input, k, dim, ...) -> (values, indices) + k = extra_args.get('k', 10) + dim = extra_args.get('dim', -1) + result = torch.topk(inputs, k, dim=dim) + return result.values # 只返回values用于性能测试 + elif op_name == 'multinomial': + # multinomial(input, num_samples, ...) + num_samples = extra_args.get('num_samples', 10) + return torch.multinomial(inputs, num_samples, **{k: v for k, v in extra_args.items() if k != 'num_samples'}) + elif op_name == 'quantile': + # quantile(input, q, dim, ...) + q = extra_args.get('q', 0.5) + dim = extra_args.get('dim', -1) + return torch.quantile(inputs, q, dim=dim, **{k: v for k, v in extra_args.items() if k not in ['q', 'dim']}) + elif op_name == 'where': + # where(condition, x, y, ...) + return torch.where(inputs[0], inputs[1], inputs[2], **extra_args) + elif op_name == 'masked_fill': + # masked_fill(input, mask, value, ...) + # value 已经通过位置参数传递,不应该再通过 extra_args 传递 + filtered_args = {k: v for k, v in extra_args.items() if k != 'value'} + return torch.masked_fill(inputs[0], inputs[1], inputs[2], **filtered_args) + elif op_name == 'masked_select': + # masked_select(input, mask, ...) + return torch.masked_select(inputs[0], inputs[1], **extra_args) + elif op_name == 'select': + # select(input, dim, index, ...) + # dim 和 index 已经通过位置参数传递,不应该再通过 extra_args 传递 + filtered_args = {k: v for k, v in extra_args.items() if k not in ['dim', 'index']} + return torch.select(inputs[0], inputs[1], inputs[2], **filtered_args) + elif op_name == 'diagonal': + # diagonal(input, offset, dim1, dim2, ...) + offset = extra_args.get('offset', 0) + dim1 = extra_args.get('dim1', 0) + dim2 = extra_args.get('dim2', 1) + return torch.diagonal(inputs, offset=offset, dim1=dim1, dim2=dim2) + elif op_name == 'diag': + # diag(input, diagonal, ...) + diagonal = extra_args.get('diagonal', 0) + return torch.diag(inputs, diagonal=diagonal) + elif op_name == 'diag_embed': + # diag_embed(input, offset, dim1, dim2, ...) + offset = extra_args.get('offset', 0) + dim1 = extra_args.get('dim1', -2) + dim2 = extra_args.get('dim2', -1) + return torch.diag_embed(inputs, offset=offset, dim1=dim1, dim2=dim2) + elif op_name == 'pad': + # pad(input, pad, mode, value, ...) + pad = extra_args.get('pad', (1, 1, 1, 1)) + mode = extra_args.get('mode', 'constant') + value = extra_args.get('value', 0.0) + return torch.nn.functional.pad(inputs, pad, mode=mode, value=value) + elif op_name == 'tile': + # tile(input, dims, ...) + dims = extra_args.get('dims', (2, 2)) + return torch.tile(inputs, dims) + elif op_name == 'repeat_interleave': + # repeat_interleave(input, repeats, dim, ...) + repeats = extra_args.get('repeats', 2) + dim = extra_args.get('dim', -1) + return torch.repeat_interleave(inputs, repeats, dim=dim) + elif op_name == 'kron': + # kron(input1, input2, ...) + return torch.kron(inputs[0], inputs[1], **extra_args) + elif op_name == 'outer': + # outer(input1, input2, ...) + return torch.outer(inputs[0], inputs[1], **extra_args) + elif op_name == 'vdot': + # vdot(input1, input2, ...) + return torch.vdot(inputs[0], inputs[1], **extra_args) + elif op_name == 'dot': + # dot(input1, input2, ...) + return torch.dot(inputs[0], inputs[1], **extra_args) + elif op_name == 'polar': + # polar(abs, angle, ...) + return torch.polar(inputs[0], inputs[1], **extra_args) + elif op_name == 'atan2': + # atan2(input1, input2, ...) + return torch.atan2(inputs[0], inputs[1], **extra_args) + elif op_name == 'hypot': + # hypot(input1, input2, ...) + return torch.hypot(inputs[0], inputs[1], **extra_args) + elif op_name == 'fmod': + # fmod(input1, input2, ...) + return torch.fmod(inputs[0], inputs[1], **extra_args) + elif op_name == 'isin': + # isin(elements, test_elements, ...) + return torch.isin(inputs[0], inputs[1], **extra_args) + elif op_name == 'fill': + # fill(input, value, ...) - 注意:fill是inplace操作 + result = inputs[0].clone() + result.fill_(inputs[1]) + return result + elif op_name == 'cross_entropy_loss': + # cross_entropy_loss(input, target, ...) + return torch.nn.functional.cross_entropy(inputs[0], inputs[1], **extra_args) + elif op_name == 'nll_loss': + # nll_loss(input, target, ...) + return torch.nn.functional.nll_loss(inputs[0], inputs[1], **extra_args) + elif op_name == 'mse_loss': + # mse_loss(input, target, ...) + return torch.nn.functional.mse_loss(inputs[0], inputs[1], **extra_args) + else: + raise ValueError(f"未知的特殊固定配置算子: {op_name}") + + +def _call_op_impl(op_func, op_name, inputs, config, extra_args, shape=None): + """实际的算子调用实现""" + if config['type'] == 'bitwise': + # 位运算:一元或二元 + if op_name == 'bitwise_not': + return op_func(inputs, **extra_args) + else: + return op_func(inputs[0], inputs[1], **extra_args) + elif config['type'] == 'ternary': + # 三元算子 + if op_name == 'lerp': + # lerp(input, end, weight) - weight可以是标量或张量 + return op_func(inputs[0], inputs[1], weight=inputs[2], **extra_args) + else: + return op_func(inputs[0], inputs[1], inputs[2], **extra_args) + elif config['type'] == 'binary': + return op_func(inputs[0], inputs[1], **extra_args) + elif config['type'] == 'matrix': + if op_name == 'addmm': + # addmm(bias, mat1, mat2, ...) + return op_func(inputs[0], inputs[1], inputs[2], **extra_args) + elif op_name in ['bmm', 'mm', 'matmul']: + return op_func(inputs[0], inputs[1], **extra_args) + elif op_name == 'mv': + return op_func(inputs[0], inputs[1], **extra_args) + elif config['type'] == 'multi_input': + input_list, kwargs = inputs + merged_kwargs = {**extra_args, **kwargs} + return op_func(input_list, **merged_kwargs) + elif config['type'] == 'constructor': + # 构造函数使用shape作为参数 + if op_name == 'full_like': + # full_like(input, fill_value) 需要参考tensor和fill_value + ref_tensor = torch.randn(shape, dtype=DTYPE, device=flag_gems.device) + fill_value = 1.0 + return op_func(ref_tensor, fill_value, dtype=DTYPE, device=flag_gems.device, **extra_args) + elif 'like' in op_name: + # ones_like等需要参考tensor(inputs就是shape,需要创建参考tensor) + ref_tensor = torch.randn(shape, dtype=DTYPE, device=flag_gems.device) + return op_func(ref_tensor, **extra_args) + elif op_name == 'eye': + # eye需要矩阵大小,使用shape的第一个维度 + n = shape[0] if isinstance(shape, tuple) and len(shape) > 0 else shape + m = shape[1] if isinstance(shape, tuple) and len(shape) > 1 else n + return op_func(n, m, dtype=DTYPE, device=flag_gems.device, **extra_args) + elif op_name == 'normal': + # normal(mean, std, size) 需要mean和std参数 + mean = 0.0 + std = 1.0 + return op_func(mean, std, shape, dtype=DTYPE, device=flag_gems.device, **extra_args) + elif op_name == 'randperm': + # randperm(n) 需要n参数,使用shape的第一个维度 + # randperm只支持整数类型(int16/int32/int64),使用配置中的dtype或默认int64 + dtype = config.get('dtype', torch.int64) + n = shape[0] if isinstance(shape, tuple) and len(shape) > 0 else (shape if isinstance(shape, int) else 256) + return op_func(n, dtype=dtype, device=flag_gems.device, **extra_args) + elif op_name == 'full': + # full(size, fill_value) 需要fill_value参数 + fill_value = 1.0 + return op_func(shape, fill_value, dtype=DTYPE, device=flag_gems.device, **extra_args) + else: + # ones, zeros, rand, randn, empty等直接使用shape + return op_func(shape, dtype=DTYPE, device=flag_gems.device, **extra_args) + elif config['type'] == 'special_constructor': + # arange和linspace需要特殊参数 + if op_name == 'arange': + # arange(start, end, step, ...) + # 使用shape的第一个维度作为end值 + end = shape[0] if isinstance(shape, tuple) and len(shape) > 0 else ( + shape if isinstance(shape, int) else 256) + return op_func(0, end, dtype=DTYPE, device=flag_gems.device, **extra_args) + elif op_name == 'linspace': + # linspace(start, end, steps, ...) + # 使用shape的第一个维度作为steps值 + steps = shape[0] if isinstance(shape, tuple) and len(shape) > 0 else ( + shape if isinstance(shape, int) else 256) + return op_func(0.0, 1.0, steps, dtype=DTYPE, device=flag_gems.device, **extra_args) + else: + raise ValueError(f"未知的特殊构造函数算子: {op_name}") + elif config['type'] == 'cumulative': + # 累积算子:cummax和cummin返回命名元组(values, indices),cumsum返回tensor + result = op_func(inputs, **extra_args) + if op_name in ['cummax', 'cummin']: + # 返回values部分用于性能测试 + return result.values + return result + elif config['type'] == 'special_fixed_config': + # 特殊固定配置算子 + return _call_special_fixed_config_op(op_func, op_name, inputs, config, extra_args, shape) + elif op_name == 'dropout': + # dropout需要位置参数:dropout(input, p, train) + # 使用torch.nn.functional.dropout + p = extra_args.get('p', 0.5) + train = extra_args.get('train', True) + return torch.nn.functional.dropout(inputs, p, train) + elif op_name == 'flip': + # flip需要位置参数:flip(input, dims) + dims = extra_args.get('dims', (0, )) + return op_func(inputs, dims) + elif op_name == 'threshold': + # threshold需要位置参数:threshold(input, threshold, value) + threshold = extra_args.get('threshold', 0.5) + value = extra_args.get('value', 0.0) + return op_func(inputs, threshold, value) + elif op_name == 'triu': + # triu需要位置参数:triu(input, diagonal) + diagonal = extra_args.get('diagonal', 0) + return op_func(inputs, diagonal) + elif op_name == 'sort': + # sort需要位置参数:sort(input, dim) + dim = extra_args.get('dim', -1) + result = op_func(inputs, dim=dim) + # sort返回(values, indices),只返回values用于性能测试 + return result.values if hasattr(result, 'values') else result[0] + elif op_name == 'vector_norm': + # vector_norm需要ord参数,dim和keepdim可选 + ord = extra_args.get('ord', 2) + dim = extra_args.get('dim', None) + keepdim = extra_args.get('keepdim', False) + if dim is not None: + return op_func(inputs, ord=ord, dim=dim, keepdim=keepdim) + else: + return op_func(inputs, ord=ord, keepdim=keepdim) + elif op_name == 'isclose': + # isclose需要两个输入和rtol/atol参数 + if config['type'] == 'binary': + rtol = extra_args.get('rtol', 1e-5) + atol = extra_args.get('atol', 1e-8) + return op_func(inputs[0], inputs[1], rtol=rtol, atol=atol) + else: + # 如果只有一个输入,创建第二个输入 + inp2 = torch.randn_like(inputs) + rtol = extra_args.get('rtol', 1e-5) + atol = extra_args.get('atol', 1e-8) + return op_func(inputs, inp2, rtol=rtol, atol=atol) + else: + # 一元算子或规约算子 + return op_func(inputs, **extra_args) + + +def test_op_performance(op_name, shape, config): + """ + 测试单个算子的性能 + + 使用triton的do_bench来测量性能,支持各种GPU设备。 + 注意:使用全局启用的flag_gems.enable(),避免每次调用use_gems()导致的重复注册。 + + Args: + op_name: 算子名称 + shape: 测试shape + config: 算子配置字典 + + Returns: + float: 平均耗时(毫秒) + """ + do_bench = _get_do_bench() + _ensure_flag_gems_enabled() + + try: + # 跳过需要特殊参数的算子 + if config.get('type') == 'skip': + pytest.skip(f"算子 {op_name}: {config.get('reason', '需要特殊参数')}") + + # 创建测试输入 + inputs = create_test_inputs(op_name, shape, config) + + # 定义要测试的函数 + if inputs is None and config['type'] == 'special_constructor': + + def test_fn(): + return call_op(op_name, None, config, shape=shape) + + # special_constructor的shape格式化 + if op_name == 'arange': + end = shape[0] if isinstance(shape, tuple) and len(shape) > 0 else ( + shape if isinstance(shape, int) else 256) + shape_str = f"arange(0, {end})" + elif op_name == 'linspace': + steps = shape[0] if isinstance(shape, tuple) and len(shape) > 0 else ( + shape if isinstance(shape, int) else 256) + shape_str = f"linspace(0, 1, {steps})" + else: + shape_str = str(shape) + else: + + def test_fn(): + return call_op(op_name, inputs, config, shape=shape) + + shape_str = _format_shape_str(op_name, inputs, config, shape) + + # 使用do_bench测量性能(返回中位数,单位:毫秒) + avg_time_ms = do_bench(test_fn, warmup=WARMUP, rep=REPETITION, return_mode="median") + avg_time_us = avg_time_ms * 1000 + elapsed_time_ms = avg_time_ms * REPETITION + + # 存储结果 + _store_result(op_name, shape_str, config, avg_time_ms, avg_time_us, elapsed_time_ms) + + # 运行时打印 + print(f" shape={shape_str:<35} avg_time={avg_time_us:>10.2f} us ({avg_time_ms:>8.4f} ms)", flush=True) + + return avg_time_ms + + except Exception as e: + # 记录失败的测试 + shape_str = str(shape) + _store_result(op_name, shape_str, config, None, None, None, error=str(e)) + print(f" shape={shape_str:<35} ERROR: {str(e)}", flush=True) + pytest.skip(f"算子 {op_name} 测试失败: {e}") + + +# 为了pytest兼容性,创建一个通用的测试函数 +@pytest.mark.parametrize("op_name", []) +def test_op_performance_pytest(op_name): + """pytest版本的测试函数(通过parametrize动态生成)""" + ops = parse_op_list() + if op_name not in ops: + pytest.skip(f"未知算子: {op_name}") + + config = get_op_config(op_name) + # 使用第一个shape作为默认测试 + shape = config['shapes'][0] if config['shapes'] else BASE_SHAPES[0] + test_op_performance(op_name, shape, config) + + +def print_summary_report(): + """ + 打印性能测试总结报告 + + 按算子分组显示所有测试结果,包括shape、配置、耗时等信息。 + """ + if not hasattr(pytest, '_op_perf_results') or not pytest._op_perf_results: + return + + results = pytest._op_perf_results + + # 按算子分组 + op_groups = {} + for r in results: + op = r['op'] + if op not in op_groups: + op_groups[op] = [] + op_groups[op].append(r) + + print("\n" + "=" * 120) + print(" " * 40 + "算子性能测试报告") + print("=" * 120) + + # 打印每个算子的结果 + for op_name in sorted(op_groups.keys()): + op_results = op_groups[op_name] + print(f"\n算子: {op_name.upper()}") + print("=" * 120) + print(f"{'Shape':<35} {'Config':<20} {'Avg Time (us)':<18} {'Avg Time (ms)':<18} {'Status':<15}") + print("=" * 120) + + for r in op_results: + shape_str = r['shape'][:34] if len(r['shape']) > 34 else r['shape'] + config_str = r['config'][:19] if len(r['config']) > 19 else r['config'] + + if r.get('error'): + print(f"{shape_str:<35} {config_str:<20} {'N/A':<18} {'N/A':<18} {'ERROR':<15}") + print(f" Error: {r['error']}") + elif r['avg_time_us'] is not None: + print( + f"{shape_str:<35} {config_str:<20} {r['avg_time_us']:>15.2f} {r['avg_time_ms']:>15.4f} {'OK':<15}" + ) + else: + print(f"{shape_str:<35} {config_str:<20} {'N/A':<18} {'N/A':<18} {'SKIPPED':<15}") + + print("\n" + "=" * 120) + print(f"测试配置:") + print(f" * Warmup: {WARMUP}") + print(f" * Repetition: {REPETITION}") + print(f" * Dtype: {DTYPE}") + _ensure_flag_gems_enabled() + print(f" * Device: {flag_gems.device}") + print(f" * Benchmark Method: triton.do_bench") + print("=" * 120 + "\n") + + +def print_csv_report(filename='performance_report.csv'): + """ + 将性能测试报告输出为CSV格式 + + Args: + filename: CSV文件名,默认为'performance_report.csv' + + Note: + 使用&作为分隔符,避免shape字段中的逗号导致格式混乱。 + """ + import csv + import os + + if not hasattr(pytest, '_op_perf_results') or not pytest._op_perf_results: + print("没有性能测试结果可输出") + return + + results = pytest._op_perf_results + + # CSV文件路径 + csv_path = os.path.join(os.path.dirname(__file__), filename) + + # 写入CSV文件 + with open(csv_path, 'w', newline='', encoding='utf-8') as csvfile: + fieldnames = [ + 'Operator', 'Shape', 'Config', 'Dtype', 'Type', 'Avg Time (us)', 'Avg Time (ms)', 'Elapsed Time (ms)', + 'Status', 'Error' + ] + # 使用&作为分隔符,避免shape字段中的逗号导致格式混乱 + writer = csv.DictWriter(csvfile, fieldnames=fieldnames, delimiter='&', quoting=csv.QUOTE_MINIMAL) + + # 写入表头 + writer.writeheader() + + # 写入数据 + for r in results: + status = 'OK' + if r.get('error'): + status = 'ERROR' + elif r['avg_time_us'] is None: + status = 'SKIPPED' + + # 将所有字段转换为字符串,确保CSV格式正确 + # csv模块会自动为包含逗号的字段添加引号 + writer.writerow({ + 'Operator': + str(r['op']), 'Shape': + str(r['shape']), # 包含逗号的字段会被自动用引号括起来 + 'Config': + str(r['config']), # 包含逗号的字段会被自动用引号括起来 + 'Dtype': + str(r['dtype']), 'Type': + str(r['type']), 'Avg Time (us)': + f"{r['avg_time_us']:.2f}" if r['avg_time_us'] is not None else 'N/A', 'Avg Time (ms)': + f"{r['avg_time_ms']:.4f}" if r['avg_time_ms'] is not None else 'N/A', 'Elapsed Time (ms)': + f"{r['elapsed_time_ms']:.4f}" if r.get('elapsed_time_ms') is not None else 'N/A', 'Status': + str(status), 'Error': + str(r.get('error', '')) + }) + + print(f"\n性能测试报告已保存为CSV格式: {csv_path}") + + +@pytest.hookimpl(trylast=True) +def pytest_sessionfinish(session, exitstatus): + """pytest session结束时打印总结报告并输出CSV""" + print_summary_report() + print_csv_report() + + +if __name__ == "__main__": + # 启用FlagGems(会自动检测可用的GPU设备) + _ensure_flag_gems_enabled() + + # 初始化结果列表 + pytest._op_perf_results = [] + + # 获取所有算子 + ops = parse_op_list() + print(f"找到 {len(ops)} 个算子,开始性能测试...\n") + + # 运行所有测试 + for op_name in ops: + config = get_op_config(op_name) + print(f"测试算子: {op_name} (类型: {config['type']})...") + + # 跳过需要特殊参数的算子 + if config.get('type') == 'skip': + print(f" 跳过: {config.get('reason', '需要特殊参数')}") + continue + + for shape in config.get('shapes', []): + try: + test_op_performance(op_name, shape, config) + except Exception as e: + print(f" 警告: {op_name} shape={shape} 测试失败: {e}") + + # 打印性能测试报告 + print("\n" + "=" * 100) + print_summary_report() + + # 输出CSV格式报告 + print_csv_report() diff --git a/third_party/wafer/bin/CMakeLists.txt b/third_party/wafer/bin/CMakeLists.txt new file mode 100755 index 00000000..36058cf7 --- /dev/null +++ b/third_party/wafer/bin/CMakeLists.txt @@ -0,0 +1,110 @@ +get_property(dialect_libs GLOBAL PROPERTY MLIR_DIALECT_LIBS) +get_property(conversion_libs GLOBAL PROPERTY MLIR_CONVERSION_LIBS) +get_property(extension_libs GLOBAL PROPERTY MLIR_EXTENSION_LIBS) +get_property(triton_libs GLOBAL PROPERTY TRITON_LIBS) + +# TODO: Workaround include path +# include_directories(${CMAKE_CURRENT_SOURCE_DIR}/../third_party/wafer/include) +# include_directories(${CMAKE_CURRENT_BINARY_DIR}/../third_party/wafer/include) + +add_llvm_executable(wafer-opt wafer-opt.cpp PARTIAL_SOURCES_INTENDED) + +# TODO: what's this? +llvm_update_compile_flags(wafer-opt) +target_link_libraries(wafer-opt PRIVATE + ${dialect_libs} + ${conversion_libs} + ${extension_libs} + ${triton_libs} + TritonSharedAnalysis + TleDsaIR + TleDsaToCore + # MLIR core + MLIROptLib + MLIRPass + MLIRRegisterAllExtensions + MLIRRegisterAllPasses + MLIRTransforms +) + +mlir_check_all_link_libraries(wafer-opt) + +# add_llvm_executable(wafer-reduce wafer-reduce.cpp PARTIAL_SOURCES_INTENDED) +# mlir_check_all_link_libraries(wafer-reduce) + +# llvm_update_compile_flags(wafer-reduce) +# target_link_libraries(wafer-reduce PRIVATE +# ${dialect_libs} +# ${conversion_libs} +# ${extension_libs} +# ${triton_libs} +# # tests +# TritonTestAnalysis +# TritonTestDialectTritonGPU +# TritonAMDGPUTestAnalysis +# # MLIR core +# MLIRReduceLib +# MLIRPass +# MLIRTransforms +# MLIRMathTestPasses +# ) + +# mlir_check_all_link_libraries(wafer-reduce) + +# add_llvm_executable(wafer-lsp wafer-lsp.cpp PARTIAL_SOURCES_INTENDED) + +# llvm_update_compile_flags(wafer-lsp) +# target_link_libraries(wafer-lsp PRIVATE +# ${dialect_libs} +# ${conversion_libs} +# ${extension_libs} +# ${triton_libs} +# # tests +# TritonTestAnalysis +# TritonTestDialectTritonGPU +# TritonAMDGPUTestAnalysis +# # MLIR core +# MLIRLspServerLib +# MLIRPass +# MLIRTransforms +# MLIRMathTestPasses +# ) + +# mlir_check_all_link_libraries(wafer-lsp) + + +# add_llvm_executable(wafer-llvm-opt +# wafer-llvm-opt.cpp + +# PARTIAL_SOURCES_INTENDED +# DEPENDS +# intrinsics_gen +# SUPPORT_PLUGINS +# ) +# target_link_libraries(wafer-llvm-opt PRIVATE +# TritonLLVMIR + +# LLVMAnalysis +# LLVMCore +# LLVMSupport +# LLVMOption +# LLVMCodeGen +# ) +# export_executable_symbols_for_plugins(wafer-llvm-opt) + + +# add_llvm_executable(wafer-tensor-layout wafer-tensor-layout.cpp PARTIAL_SOURCES_INTENDED) +# target_link_libraries(wafer-tensor-layout PRIVATE +# ${triton_libs} +# ${conversion_libs} +# ${extension_libs} +# ${dialect_libs} +# TritonTestAnalysis +# TritonTestDialectTritonGPU +# TritonAMDGPUTestAnalysis +# MLIRMathTestPasses +# ) + +install(TARGETS wafer-opt + RUNTIME DESTINATION ${INSTALL_WAFER_DIR}/bin +) diff --git a/third_party/wafer/bin/RegisterTritonDialects.h b/third_party/wafer/bin/RegisterTritonDialects.h new file mode 100755 index 00000000..8c3d8f3f --- /dev/null +++ b/third_party/wafer/bin/RegisterTritonDialects.h @@ -0,0 +1,132 @@ +#pragma once +#include "Address/Dialect/IR/AddressDialect.h" +#include "Address/Transforms/Passes.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/TritonGPU/IR/Dialect.h" +#include "triton/Dialect/TritonNvidiaGPU/IR/Dialect.h" + +#include "triton/Dialect/Triton/Transforms/Passes.h" +#include "triton/Dialect/TritonGPU/Transforms/Passes.h" +#include "triton/Dialect/TritonNvidiaGPU/Transforms/Passes.h" + +#include "triton/Conversion/TritonGPUToLLVM/Passes.h" +#include "triton/Conversion/TritonToTritonGPU/Passes.h" +#include "triton/Target/LLVMIR/Passes.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Transforms/AllInterfaces.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/ValueBoundsOpInterfaceImpl.h" +#include "mlir/Dialect/SCF/Transforms/BufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/Dialect/Tensor/IR/TensorInferTypeOpInterfaceImpl.h" +#include "mlir/Dialect/Tensor/Transforms/BufferizableOpInterfaceImpl.h" + +#include "magic-kernel/Conversion/TLEToMK/Passes.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "third_party/tle/include/tle-dsa/Conversion/DsaToCore/DsaToCore.h" +#include "third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.h" +#include "triton-shared/Conversion/ConvertTritonPtr/Passes.h" +#include "triton-shared/Conversion/ReconcilePtrCasts/Passes.h" +#include "triton-shared/Conversion/StructuredToMemref/Passes.h" +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h" +#include "triton-shared/Conversion/TritonPtrToMemref/Passes.h" +#include "triton-shared/Conversion/TritonToCoreDialects/Passes.h" +#include "triton-shared/Conversion/TritonToLinalg/Passes.h" +#include "triton-shared/Conversion/TritonToStructured/Passes.h" +#include "triton-shared/Conversion/TritonToUnstructured/Passes.h" +#include "triton-shared/Conversion/UnstructuredToMemref/Passes.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" +#include "wafer/Conversion/LinalgFusion/Passes.h" +#include "wafer/Conversion/LinalgTiling/Passes.h" +#include "wafer/Dialect/IR/WaferDialect.h" + +#include "magic-kernel/Conversion/CoreDialectsToMK/Passes.h" +#include "magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.h" +#include "magic-kernel/Conversion/LinalgToMK/Passes.h" +#include "magic-kernel/Transforms/Passes.h" +#include "magic-kernel/Conversion/MKPipeline/Passes.h" +#include "wafer/Transforms/Passes.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "wafer/Conversion/AllocateSharedMemory/Passes.h" +#include "wafer/Conversion/ExportKernelSymbols/Passes.h" +#include "wafer/Conversion/MKToWafer/Passes.h" +#include "wafer/Conversion/WaferMemrefToLLVM/Passes.h" +#include "wafer/Conversion/WaferToLLVM/KernelArgBufferPass.h" +#include "wafer/Conversion/WaferToLLVM/Passes.h" + +#include "magic-kernel/Transforms/BufferizableOpInterfaceImpl.h" + +#include "mlir/InitAllDialects.h" +#include "mlir/InitAllExtensions.h" +#include "mlir/InitAllPasses.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Pass/PassRegistry.h" + +inline void registerTritonDialects(mlir::DialectRegistry ®istry) { + mlir::registerAllPasses(); + mlir::triton::registerTritonPasses(); + mlir::registerLinalgPasses(); + mlir::dsa::registerDsaMemoryToCorePass(); + mlir::triton::registerTLEToMKPass(); + mlir::triton::nvidia_gpu::registerTritonNvidiaGPUPasses(); + mlir::triton::registerTritonToLinalgPass(); + mlir::triton::registerTritonToStructuredPass(); + mlir::triton::registerTritonToUnstructuredPass(); + mlir::triton::registerTritonArithToLinalgPasses(); + mlir::triton::registerConvertTritonToTritonGPUPass(); + mlir::triton::registerStructuredToMemrefPasses(); + mlir::triton::registerUnstructuredToMemref(); + mlir::triton::registerTritonPtrToMemref(); + mlir::triton::registerTritonToCoreDialectsPass(); + mlir::triton::registerReconcilePtrCasts(); + mlir::triton::gpu::registerAllocateSharedMemoryPass(); + mlir::triton::gpu::registerTritonGPUAllocateWarpGroups(); + mlir::triton::gpu::registerTritonGPUGlobalScratchAllocationPass(); + mlir::registerLLVMDIScope(); + + // Core dialects to MK layer conversion passes + mlir::triton::registerWaferMemrefToLLVMPass(); + mlir::triton::registerLinalgToMKPass(); + mlir::triton::registerMKTransformsPasses(); + mlir::triton::registerMKPipelinePasses(); + mlir::triton::registerWaferTransformsPasses(); + mlir::triton::registerCoreDialectsToMKPass(); + mlir::triton::registerLegalizeTensorFormLoopsPass(); + mlir::addr::registerAddrToLLVMPass(); + mlir::triton::registerLinalgTilingPass(); + mlir::triton::registerLinalgFusionPass(); + + // Wafer specific conversion passes + mlir::triton::registerMKToWaferPass(); + mlir::triton::alloc::registerAllocateSharedMemoryPass(); + mlir::triton::registerWaferToLLVMPass(); + mlir::triton::registerExportKernelSymbols(); + mlir::triton::registerKernelArgBufferPass(); + + // Register LLVM 22's standard external models and Wafer's custom model. + mlir::registerAllExtensions(registry); + mlir::linalg::registerAllDialectInterfaceImplementations(registry); + mlir::scf::registerBufferizableOpInterfaceExternalModels(registry); + mlir::scf::registerValueBoundsOpInterfaceExternalModels(registry); + mlir::tensor::registerBufferizableOpInterfaceExternalModels(registry); + mlir::tensor::registerInferTypeOpInterfaceExternalModels(registry); + mlir::mk::registerBufferizableOpInterfaceExternalModels(registry); + + registry.insert< + mlir::triton::TritonDialect, mlir::cf::ControlFlowDialect, + mlir::triton::nvidia_gpu::TritonNvidiaGPUDialect, + mlir::triton::gpu::TritonGPUDialect, mlir::math::MathDialect, + mlir::arith::ArithDialect, mlir::scf::SCFDialect, mlir::gpu::GPUDialect, + mlir::LLVM::LLVMDialect, + mlir::ttx::TritonTilingExtDialect, mlir::tts::TritonStructuredDialect, + mlir::linalg::LinalgDialect, mlir::func::FuncDialect, + mlir::tensor::TensorDialect, mlir::memref::MemRefDialect, + mlir::affine::AffineDialect, mlir::bufferization::BufferizationDialect, + mlir::mk::MagicKernelDialect, mlir::wafer::WaferDialect, + mlir::addr::AddressDialect, mlir::dsa::DsaDialect>(); +} diff --git a/third_party/wafer/bin/wafer-llvm-opt.cpp b/third_party/wafer/bin/wafer-llvm-opt.cpp new file mode 100755 index 00000000..1ec804cb --- /dev/null +++ b/third_party/wafer/bin/wafer-llvm-opt.cpp @@ -0,0 +1,121 @@ +/// Trimmed down clone of llvm opt to be able to test triton custom llvm ir +/// passes. +#include "lib/Target/LLVMIR/LLVMPasses.h" +#include "llvm/CodeGen/CommandFlags.h" +#include "llvm/IR/Constants.h" +#include "llvm/IR/DataLayout.h" +#include "llvm/IR/LLVMContext.h" +#include "llvm/IR/Module.h" +#include "llvm/IR/Verifier.h" +#include "llvm/IRReader/IRReader.h" +#include "llvm/Passes/PassBuilder.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/Error.h" +#include "llvm/Support/FileSystem.h" +#include "llvm/Support/InitLLVM.h" +#include "llvm/Support/SourceMgr.h" +#include "llvm/Support/SystemUtils.h" +#include "llvm/Support/ToolOutputFile.h" +#include "llvm/TargetParser/Triple.h" +#include + +using namespace llvm; + +static cl::opt InputFilename(cl::Positional, + cl::desc(""), + cl::init("-"), + cl::value_desc("filename")); + +static cl::opt OutputFilename("o", + cl::desc("Override output filename"), + cl::value_desc("filename")); + +static cl::opt ClDataLayout("data-layout", + cl::desc("data layout string to use"), + cl::value_desc("layout-string"), + cl::init("")); +static cl::opt + TargetTriple("mtriple", cl::desc("Override target triple for module")); + +static cl::opt + BreakStructPhiNodes("break-struct-phi-nodes", + llvm::cl::desc("run pass to break phi struct"), + cl::init(false)); + +namespace { +static std::function makeOptimizingPipeline() { + return [](Module *m) -> Error { + PipelineTuningOptions tuningOptions; + PassBuilder pb(nullptr, tuningOptions); + + LoopAnalysisManager lam; + FunctionAnalysisManager fam; + CGSCCAnalysisManager cgam; + ModuleAnalysisManager mam; + pb.registerModuleAnalyses(mam); + pb.registerCGSCCAnalyses(cgam); + pb.registerFunctionAnalyses(fam); + pb.registerLoopAnalyses(lam); + pb.crossRegisterProxies(lam, fam, cgam, mam); + + ModulePassManager mpm; + llvm::FunctionPassManager fpm; + if (BreakStructPhiNodes) + fpm.addPass(BreakStructPhiNodesPass()); + mpm.addPass(createModuleToFunctionPassAdaptor(std::move(fpm))); + mpm.run(*m, mam); + return Error::success(); + }; +} +} // namespace + +int main(int argc, char **argv) { + InitLLVM X(argc, argv); + cl::ParseCommandLineOptions( + argc, argv, "llvm .bc -> .bc modular optimizer and analysis printer\n"); + + LLVMContext Context; + SMDiagnostic Err; + + // Load the input module... + auto SetDataLayout = [](StringRef, StringRef) -> std::optional { + if (ClDataLayout.empty()) + return std::nullopt; + return ClDataLayout; + }; + std::unique_ptr M; + M = parseIRFile(InputFilename, Err, Context, ParserCallbacks(SetDataLayout)); + if (!M) { + Err.print(argv[0], errs()); + return 1; + } + // If we are supposed to override the target triple or data layout, do so now. + if (!TargetTriple.empty()) + M->setTargetTriple(Triple::normalize(TargetTriple)); + auto optPipeline = makeOptimizingPipeline(); + if (auto err = optPipeline(M.get())) { + llvm::errs() << "Failed to optimize LLVM IR " << err << "\n"; + } + + if (verifyModule(*M, &errs())) { + errs() << argv[0] << ": " << InputFilename + << ": error: input module is broken!\n"; + return 1; + } + + // Write to standard output. + std::unique_ptr Out; + // Default to standard output. + if (OutputFilename.empty()) + OutputFilename = "-"; + std::error_code EC; + sys::fs::OpenFlags Flags = sys::fs::OF_TextWithCRLF; + Out.reset(new ToolOutputFile(OutputFilename, EC, Flags)); + if (EC) { + errs() << EC.message() << '\n'; + return 1; + } + Out->os() << *M << "\n"; + Out->keep(); + return 0; +} diff --git a/third_party/wafer/bin/wafer-lsp.cpp b/third_party/wafer/bin/wafer-lsp.cpp new file mode 100755 index 00000000..f95036dc --- /dev/null +++ b/third_party/wafer/bin/wafer-lsp.cpp @@ -0,0 +1,10 @@ +#include "./RegisterTritonDialects.h" + +#include "mlir/Tools/mlir-lsp-server/MlirLspServerMain.h" + +int main(int argc, char **argv) { + mlir::DialectRegistry registry; + registerTritonDialects(registry); + + return mlir::failed(mlir::MlirLspServerMain(argc, argv, registry)); +} diff --git a/third_party/wafer/bin/wafer-opt.cpp b/third_party/wafer/bin/wafer-opt.cpp new file mode 100755 index 00000000..2d257077 --- /dev/null +++ b/third_party/wafer/bin/wafer-opt.cpp @@ -0,0 +1,11 @@ +#include "./RegisterTritonDialects.h" + +#include "mlir/Tools/mlir-opt/MlirOptMain.h" + +int main(int argc, char **argv) { + mlir::DialectRegistry registry; + registerTritonDialects(registry); + + return mlir::asMainReturnCode(mlir::MlirOptMain( + argc, argv, "Triton (GPU) optimizer driver\n", registry)); +} diff --git a/third_party/wafer/bin/wafer-reduce.cpp b/third_party/wafer/bin/wafer-reduce.cpp new file mode 100755 index 00000000..8235f8fc --- /dev/null +++ b/third_party/wafer/bin/wafer-reduce.cpp @@ -0,0 +1,11 @@ +#include "./RegisterTritonDialects.h" + +#include "mlir/Tools/mlir-reduce/MlirReduceMain.h" + +int main(int argc, char **argv) { + mlir::DialectRegistry registry; + registerTritonDialects(registry); + + mlir::MLIRContext context(registry); + return mlir::failed(mlir::mlirReduceMain(argc, argv, context)); +} diff --git a/third_party/wafer/bin/wafer-tensor-layout.cpp b/third_party/wafer/bin/wafer-tensor-layout.cpp new file mode 100755 index 00000000..cc121b3e --- /dev/null +++ b/third_party/wafer/bin/wafer-tensor-layout.cpp @@ -0,0 +1,232 @@ +#include "RegisterTritonDialects.h" + +#include "mlir/AsmParser/AsmParser.h" +#include "mlir/AsmParser/AsmParserState.h" +#include "mlir/IR/MLIRContext.h" + +#include "triton/Dialect/TritonGPU/IR/Dialect.h" +#include "triton/Dialect/TritonNvidiaGPU/IR/Dialect.h" + +#include "llvm/Support/CommandLine.h" +#include "llvm/Support/ErrorOr.h" +#include "llvm/Support/FileSystem.h" +#include "llvm/Support/MemoryBuffer.h" +#include "llvm/Support/SourceMgr.h" +#include "llvm/Support/raw_ostream.h" + +using namespace llvm; +using namespace mlir; + +// A CLI tool to print the layout of a tensor. +// +// clang-format off +// Example usage: +// +// triton-tensor-layout -l "#ttg.nvidia_mma<{versionMajor = 3, versionMinor = 0, warpsPerCTA = [8, 1], CTAsPerCGA = [1, 1], CTASplitNum = [1, 1], CTAOrder = [1, 0], instrShape = [16, 256, 32]}>" -t "tensor<128x256xf16>" +// +// triton-tensor-layout -i input.mlir -t "tensor<1x128x128xf16>" -o output.txt +// +// triton-tensor-layout -i input.mlir -t "tensor<1x128x128xf16>" -o output.txt -alias-names="blocked,mma" -use-hw-view +// +// An input file usually looks like: +// ''' +// #mma = #ttg.amd_mfma<{versionMajor = 2, versionMinor = 0, warpsPerCTA = [1, 1, 8], instrShape = [32, 32], isTransposed = false}> +// #blocked = #ttg.blocked<{sizePerThread = [1, 8, 1], threadsPerWarp = [1, 16, 4], warpsPerCTA = [1, 1, 8], order = [0, 1, 2]}> +// ''' +// clang-format on + +//===--------------------------------------------------------------------===// +// CLI options +//===--------------------------------------------------------------------===// + +cl::OptionCategory PrinterCategory("Available Print Options", + "Options for the tensor layout printing."); + +static cl::opt InputFile( + "i", cl::desc("File that contains the tensor data layout attributes"), + cl::init(""), cl::value_desc("filename"), cl::cat(PrinterCategory)); + +static cl::opt + OutputFile("o", cl::desc("Output file to write the layout into"), + cl::init(""), cl::value_desc("filename"), + cl::cat(PrinterCategory)); + +static cl::opt + DataLayoutStr("l", cl::desc("Tensor data layout attribute in string"), + cl::value_desc("layout-string"), cl::init(""), + cl::cat(PrinterCategory)); + +static cl::list + AliasName("alias-names", + cl::desc("A list of alias names (separated by comma) of the " + "layout attributes in the input file"), + cl::value_desc("name1,name2,name3,..."), cl::CommaSeparated, + cl::ZeroOrMore, cl::cat(PrinterCategory)); + +static cl::opt UseHWPointOfView( + "use-hw-view", + llvm::cl::desc( + "Print the layout in hardware point of view. This means the output is " + "from the warp's perspective. Otherwise, the output is from the " + "tensor's perspective (e.g., each element maps to xxx thread)."), + cl::init(false), cl::cat(PrinterCategory)); + +static cl::opt TensorStr( + "t", cl::desc("Tensor shape and element type (e.g., tensor<2x2xf32>)"), + cl::init(""), cl::value_desc("tensor-type"), cl::cat(PrinterCategory)); + +//===--------------------------------------------------------------------===// +// Helper functions +//===--------------------------------------------------------------------===// + +LogicalResult layoutPrint(RankedTensorType tensorType, raw_ostream &os) { + // DistributedEncodingTrait and SharedEncodingTrait implements the + // toLinearLayout interface. + mlir::Attribute layout = tensorType.getEncoding(); + if (isa(layout)) { + os << triton::gpu::getLayoutStr(tensorType, UseHWPointOfView); + return success(); + } + + llvm::errs() << "Unsupported tensor layout attribute: " + << tensorType.getEncoding() << "\n"; + return failure(); +} + +LogicalResult printLayoutFromFile(MLIRContext *context, StringRef filename, + ArrayRef names, + TensorType tensorTy, raw_string_ostream &ss) { + if (filename.empty()) + return success(); + + llvm::ErrorOr> fileOrErr = + llvm::MemoryBuffer::getFileOrSTDIN(filename); + if (std::error_code ec = fileOrErr.getError()) { + llvm::errs() << "Could not open input file: " << ec.message() << "\n"; + return failure(); + } + + llvm::SourceMgr sourceMgr; + sourceMgr.AddNewSourceBuffer(std::move(*fileOrErr), llvm::SMLoc()); + ParserConfig config(context); + auto asmState = AsmParserState(); + + Block parsedIR; + if (failed(parseAsmSourceFile(sourceMgr, &parsedIR, config, &asmState))) { + llvm::errs() << "Fail to parse the input file: " << filename << "\n"; + return failure(); + } + + auto printLambda = [&](StringRef name, mlir::Attribute attr) { + ss << "Print layout attribute: #" << name << " = " << attr << "\n"; + + auto rankedTensorTy = RankedTensorType::get( + tensorTy.getShape(), tensorTy.getElementType(), attr); + + return layoutPrint(rankedTensorTy, ss); + }; + + if (names.empty()) + // If no alias name is given, we print all layout attributes in the file. + for (const auto &def : asmState.getAttributeAliasDefs()) { + if (failed(printLambda(def.name, def.value))) + return failure(); + } + else { + // Print the layout attributes with the given alias names. + for (const auto &alias : names) { + auto def = asmState.getAttributeAliasDef(alias); + if (!def) { + llvm::errs() << "Can't find the layout attribute: " << alias << "\n"; + return failure(); + } + + if (failed(printLambda(alias, def->value))) + return failure(); + + ss << "\n"; + } + } + + return success(); +} + +LogicalResult printLayoutFromString(MLIRContext *context, + StringRef layoutAttrStr, + TensorType tensorTy, + raw_string_ostream &ss) { + if (layoutAttrStr.empty()) + return success(); + + mlir::Attribute layout = parseAttribute(layoutAttrStr, context); + if (!layout) { + llvm::errs() << "Invalid layout attribute: " << layoutAttrStr << "\n"; + return failure(); + } + + auto rankedTensorTy = RankedTensorType::get( + tensorTy.getShape(), tensorTy.getElementType(), layout); + + ss << "Print layout attribute: " << layout << "\n"; + + return layoutPrint(rankedTensorTy, ss); +} + +//===--------------------------------------------------------------------===// +// Main entry point +//===--------------------------------------------------------------------===// + +int main(int argc, char **argv) { + cl::HideUnrelatedOptions(PrinterCategory); + cl::ParseCommandLineOptions(argc, argv, "tensor layout printer\n"); + + DialectRegistry registry; + registerTritonDialects(registry); + + MLIRContext ctx(registry); + ctx.loadAllAvailableDialects(); + + if (TensorStr.empty()) { + llvm::errs() << "Must specify the tensor type argument\n"; + return 1; + } + + mlir::Type parsedTy = parseType(TensorStr, &ctx); + if (!parsedTy) { + llvm::errs() << "Fail to parse the tensor type argument: " << TensorStr + << "\n"; + return 1; + } + + TensorType tensorType = dyn_cast(parsedTy); + if (!tensorType) { + llvm::errs() << "Invalid tensor type argument: " << TensorStr << "\n"; + return 1; + } + + std::string storage; + raw_string_ostream ss(storage); + + if (failed(printLayoutFromFile(&ctx, InputFile, AliasName, tensorType, ss))) + return 1; + + if (failed(printLayoutFromString(&ctx, DataLayoutStr, tensorType, ss))) + return 1; + + if (OutputFile.empty()) { + llvm::outs() << ss.str(); + } else { + std::error_code ec; + llvm::raw_fd_ostream outFs(OutputFile, ec, llvm::sys::fs::OF_Text); + if (ec) { + llvm::errs() << "Error: " << ec.message() << " : unable to open " + << OutputFile << " for output\n"; + return 1; + } + outFs << ss.str(); + outFs.close(); + } + + return 0; +} diff --git a/third_party/wafer/cmake/WaferConfig.cmake b/third_party/wafer/cmake/WaferConfig.cmake new file mode 100644 index 00000000..49c33f9e --- /dev/null +++ b/third_party/wafer/cmake/WaferConfig.cmake @@ -0,0 +1,5 @@ +foreach(name WAFER_DEPS_ROOT WAFER_SDK_INCLUDE_DIR WAFER_RT_THREAD_SMP_ROOT WAFER_BSP_INCLUDE_DIR) + if(NOT DEFINED ${name} AND DEFINED ENV{${name}}) + set(${name} "$ENV{${name}}") + endif() +endforeach() diff --git a/third_party/wafer/crt/CMakeLists.txt b/third_party/wafer/crt/CMakeLists.txt new file mode 100755 index 00000000..034d5c0e --- /dev/null +++ b/third_party/wafer/crt/CMakeLists.txt @@ -0,0 +1,171 @@ +include("${CMAKE_CURRENT_LIST_DIR}/../cmake/WaferConfig.cmake") + +cmake_minimum_required(VERSION 3.18) + +set(TARGET Wafer) +# Set TARGET from environment variable +if(NOT DEFINED TARGET) + if(DEFINED ENV{CRT_TARGET}) + set(TARGET $ENV{CRT_TARGET}) + else() + message(FATAL_ERROR "CRT_TARGET environment variable is not defined") + endif() +endif() + +if(NOT DEFINED WAFER_RT_THREAD_SMP_ROOT) + if(DEFINED ENV{WAFER_RT_THREAD_SMP_ROOT}) + set(WAFER_RT_THREAD_SMP_ROOT $ENV{WAFER_RT_THREAD_SMP_ROOT}) + else() + message(FATAL_ERROR "WAFER_RT_THREAD_SMP_ROOT environment variable is not defined") + endif() +endif() + +if(NOT DEFINED XUANTIE_NAME) + if(DEFINED ENV{XUANTIE_NAME}) + set(XUANTIE_NAME $ENV{XUANTIE_NAME}) + else() + message(FATAL_ERROR "XUANTIE_NAME environment variable is not defined") + endif() +endif() + +# Set LLVM_SYSPATH from environment variable +if(NOT DEFINED LLVM_SYSPATH) + if(DEFINED ENV{LLVM_SYSPATH}) + set(LLVM_SYSPATH $ENV{LLVM_SYSPATH}) + else() + message(FATAL_ERROR "LLVM_SYSPATH environment variable is not defined") + endif() +endif() + +if(NOT DEFINED WAFER_DEPS_ROOT) + if(DEFINED ENV{WAFER_DEPS_ROOT}) + set(WAFER_DEPS_ROOT $ENV{WAFER_DEPS_ROOT}) + else() + message(FATAL_ERROR "WAFER_DEPS_ROOT environment variable is not defined") + endif() +endif() + +# Build for simulator or hardware +if(NOT DEFINED USE_SIM_MODE) + if(DEFINED ENV{USE_SIM_MODE}) + set(USE_SIM_MODE $ENV{USE_SIM_MODE}) + else() + set(USE_SIM_MODE OFF) + message(STATUS "Building for hardware (USE_SIM_MODE not set)") + endif() +endif() + +# Project name and version +project(VendorRuntime LANGUAGES CXX C) + +# Define standard include directories +include_directories(${WAFER_DEPS_ROOT}/include) +include_directories(${CMAKE_CURRENT_SOURCE_DIR}) +include_directories(${CMAKE_CURRENT_SOURCE_DIR}/include/${TARGET}) +include_directories(${CMAKE_CURRENT_BINARY_DIR}) +include_directories(${WAFER_RT_THREAD_SMP_ROOT}/interface/op_fw_sim_if/peripheral/include/) + +# Set build type default +if(NOT CMAKE_BUILD_TYPE) + set(CMAKE_BUILD_TYPE Release CACHE STRING "Build type (default Release)" FORCE) +endif() + +# Library name: vr stands for Vendor Runtime +set(VENDOR_RUNTIME_LIB vr) + +# Collect all source files from the vendor directory +file(GLOB_RECURSE VENDOR_SOURCES lib/${TARGET}/*.c) + +set(CMAKE_SYSTEM_NAME Generic) +set(CMAKE_C_COMPILER ${LLVM_SYSPATH}/bin/clang) +set(CMAKE_CXX_COMPILER ${LLVM_SYSPATH}/bin/clang++) + +if (USE_SIM_MODE) + # Define simulator specific compile options + set(SIMULATOR_COMPILE_OPTIONS + -fPIC + -DUSE_SIM_MODE + ) + + add_library(${VENDOR_RUNTIME_LIB} SHARED ${VENDOR_SOURCES}) + + # Apply simulator specific settings to our target + target_compile_options(${VENDOR_RUNTIME_LIB} PRIVATE ${SIMULATOR_COMPILE_OPTIONS}) + + # Set properties for the library + set_target_properties(${VENDOR_RUNTIME_LIB} PROPERTIES + POSITION_INDEPENDENT_CODE ON + LIBRARY_OUTPUT_DIRECTORY ${CMAKE_CURRENT_BINARY_DIR}/lib + SUFFIX ".so" + ) +else() + # Define RISC-V target triple + set(RISCV_TRIPLE "riscv64-unknown-elf") + set(CMAKE_SYSTEM_PROCESSOR riscv) + + include_directories(${WAFER_DEPS_ROOT}/include) + if(IS_ABSOLUTE "${XUANTIE_NAME}") + set(WAFER_XUANTIE_ROOT "${XUANTIE_NAME}") + else() + set(WAFER_XUANTIE_ROOT "${WAFER_DEPS_ROOT}/${XUANTIE_NAME}") + endif() + include_directories(${WAFER_XUANTIE_ROOT}/riscv64-unknown-elf/include) + if(NOT DEFINED WAFER_BSP_INCLUDE_DIR) + set(WAFER_BSP_INCLUDE_DIR "${WAFER_DEPS_ROOT}/include/bsp/xuantie_riscv_tx81/board_riscv_tx81/include") + endif() + include_directories("${WAFER_BSP_INCLUDE_DIR}") + include_directories(${WAFER_DEPS_ROOT}/include/rtthread/include) + + # Define RISC-V specific compile options + set(RISCV_COMPILE_OPTIONS + --target=${RISCV_TRIPLE} + -march=rv64imfdc + -mabi=lp64d + -mcmodel=medany + -DCONFIG_TX8_KERNEL_PRINTF_SUPPORT=1 + ) + + if (ENABLE_PROFILING) + # 只需要定义,不需要链接,编译kernel.so的时候链接profile + set(RISCV_COMPILE_OPTIONS ${RISCV_COMPILE_OPTIONS} + -DENABLE_PROFILING + ) + endif() + + if (NO_INTRNISIC_RUN) + set(RISCV_COMPILE_OPTIONS ${RISCV_COMPILE_OPTIONS} + -DNO_INTRNISIC_RUN + ) + endif() + + if (ENABLE_SYNCHRONOUS_INTRINSIC) + set(RISCV_COMPILE_OPTIONS ${RISCV_COMPILE_OPTIONS} + -DENABLE_SYNCHRONOUS_INTRINSIC + ) + endif() + + if(CMAKE_BUILD_TYPE STREQUAL "Debug") + set(RISCV_COMPILE_OPTIONS ${RISCV_COMPILE_OPTIONS} + -DRT_USING_DEBUG + ) + endif() + + # Add the library target + add_library(${VENDOR_RUNTIME_LIB} STATIC ${VENDOR_SOURCES}) + + # Apply RISC-V specific settings to our target + target_compile_options(${VENDOR_RUNTIME_LIB} PRIVATE ${RISCV_COMPILE_OPTIONS}) + target_link_options(${VENDOR_RUNTIME_LIB} PRIVATE --target=${RISCV_TRIPLE}) + + # Set properties for the library + set_target_properties(${VENDOR_RUNTIME_LIB} PROPERTIES + POSITION_INDEPENDENT_CODE ON + ARCHIVE_OUTPUT_DIRECTORY ${CMAKE_CURRENT_BINARY_DIR}/lib + SUFFIX ".a" + ) + +endif() + +install(TARGETS ${VENDOR_RUNTIME_LIB} + ARCHIVE DESTINATION ${INSTALL_WAFER_DIR}/lib +) diff --git a/third_party/wafer/crt/README.md b/third_party/wafer/crt/README.md new file mode 100755 index 00000000..3750c3c7 --- /dev/null +++ b/third_party/wafer/crt/README.md @@ -0,0 +1,2 @@ +This folder contains the low level API implementation for various ML +controller or accelerator. diff --git a/third_party/wafer/crt/include/Wafer/op_gelu.h b/third_party/wafer/crt/include/Wafer/op_gelu.h new file mode 100755 index 00000000..c14683bf --- /dev/null +++ b/third_party/wafer/crt/include/Wafer/op_gelu.h @@ -0,0 +1,19 @@ +#ifndef CRT_TARGET_GELU_H +#define CRT_TARGET_GELU_H + +#include "wafer.h" + +hybrid_value get_ptr_value_by_idx(void *addr, size_t idx, Data_Format dtype); + +void set_ptr_value_by_idx(void *addr, hybrid_value value, uint64_t idx, + Data_Format dtype); +void get_erf_value(void *in_addr, void *out_addr, uint64_t count, + Data_Format dtype); +void get_tanh_value(uint64_t *in, uint64_t *imm, uint64_t *out, + uint32_t elem_count, uint16_t fmt); +void op_gelu_none(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt); +void op_gelu_tanh(uint64_t *src, uint64_t *imm, uint64_t *dst, + uint32_t elem_count, uint16_t fmt); + +#endif // OP_GELU_H diff --git a/third_party/wafer/crt/include/Wafer/op_reduce_mul_impl.h b/third_party/wafer/crt/include/Wafer/op_reduce_mul_impl.h new file mode 100755 index 00000000..9c3c68dd --- /dev/null +++ b/third_party/wafer/crt/include/Wafer/op_reduce_mul_impl.h @@ -0,0 +1,10 @@ +#ifndef CRT_TARGET_REDUCE_MUL_H +#define CRT_TARGET_REDUCE_MUL_H + +#define CONFIG_NO_PLATFORM_HOOK_H +#include "instr_adapter.h" + +void op_reduce_mul_impl(void *in, void *out, Data_Shape shape, + uint32_t reduce_dim, Data_Format fmt); + +#endif // CRT_TARGET_REDUCE_MUL_H diff --git a/third_party/wafer/crt/include/Wafer/wafer.h b/third_party/wafer/crt/include/Wafer/wafer.h new file mode 100755 index 00000000..6edc6e23 --- /dev/null +++ b/third_party/wafer/crt/include/Wafer/wafer.h @@ -0,0 +1,97 @@ +//===----------------------- wafer.h ---------------------------*- C -*-----===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +#ifndef CRT_TARGET_WAFER_H +#define CRT_TARGET_WAFER_H + +#define CONFIG_NO_PLATFORM_HOOK_H +#include "instr_adapter.h" +#include "instr_adapter_plat.h" +#include "instr_def.h" +#include "instr_operator.h" +#include "lib_log.h" +#include +#include +#include + +typedef enum { + UNKNOWN = 0, + SPM = 1, + DDR = 2, +} MemorySpace; + +// Neural engine activate mode +typedef enum { + None = 0, + ENRelu = 1, + ENLeakRelu = 2, +} ActFuncMode; +typedef union { + int32_t i; + uint32_t u; + float f; +} tmp_32suf; + +typedef union { + int16_t i16; + float fp32; + uint8_t data[4]; // 按字节访问 +} hybrid_value; +#ifdef __cplusplus +extern "C" { +#endif + +float set_value2float32(Data_Format fmt, int8_t *value); +hybrid_value set_float2value(Data_Format dtype, float value); + +uint32_t get_dtype_size_new(Data_Format fmt); + +uint32_t get_cx_align_base_new(uint32_t c, Data_Format fmt); + +uint64_t next_power_of_two_64(uint64_t x); + +bool is_contiguous(int *shape, int *strides, int elem_bytes); + +bool no_reverse_memory_access(int *stride, int rank); + +// Copy data byte by byte +void wafer_memcpy(char *srcPtr, char *dstPtr, int *src_shape, int *src_stride, + int *dst_shape, int *dst_stride, int rank, + uint32_t elem_bytes); + +void legalizeMemoryOpAttribute(int *src_shape, int *src_stride, int *dst_shape, + int *dst_stride, int rank, uint32_t *elem_bytes, + uint32_t *fmt); + +// Use in simulation mode, return the spm address mapping +int8_t *get_spm_memory_mapping(uint64_t offset); +// Hardware mode will use add the spmMappingOffset to get the real spm address +// Simulation mode will call get_spm_memory_mapping +int8_t *get_spm_memory_mapping_wrapper(uint64_t offset); + +#ifdef USE_SIM_MODE +#else +void atomic_barrier_in(); +void atomic_barrier_out(); +#endif + +#ifdef __cplusplus +} +#endif + +#ifdef NO_INTRNISIC_RUN +#define INTRNISIC_RUN_SWITCH return +#else +#define INTRNISIC_RUN_SWITCH +#endif + +#ifdef ENABLE_SYNCHRONOUS_INTRINSIC +#define SYNCHRONOUS_INTRINSIC_SWITCH TsmWaitfinish() +#else +#define SYNCHRONOUS_INTRINSIC_SWITCH +#endif + +#endif // CRT_TARGET_WAFER_H diff --git a/third_party/wafer/crt/lib/Wafer/abs.c b/third_party/wafer/crt/lib/Wafer/abs.c new file mode 100755 index 00000000..81c2f9bc --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/abs.c @@ -0,0 +1,33 @@ +//===------------------------- abs.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::AbsVVOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __AbsVV(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->AbsVV(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/argmax.c b/third_party/wafer/crt/lib/Wafer/argmax.c new file mode 100755 index 00000000..db4930c2 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/argmax.c @@ -0,0 +1,73 @@ +//===------------------------ argmax.c ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::ArgMax see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __ArgMax(uint64_t *src, uint64_t *dst0, uint64_t *dst1, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + volatile void *max_val = + (volatile void *)get_spm_memory_mapping((uint64_t)dst0); + volatile void *max_idx = + (volatile void *)get_spm_memory_mapping((uint64_t)dst1); + if (elem_count == 1) { + volatile void *src_data = + (volatile void *)get_spm_memory_mapping((uint64_t)src); + *(uint32_t *)max_idx = 0; + switch (fmt) { + case Fmt_FP16: + case Fmt_BF16: + *(uint16_t *)max_val = *(uint16_t *)src_data; + break; + case Fmt_FP32: + case Fmt_TF32: + *(uint32_t *)max_val = *(uint32_t *)src_data; + break; + default: + assert(0 && "ArgMax: Unsupport dtype"); + } + return; + } + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmPeripheralInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + ; + + cmd->ArgMax(&inst, (uint64_t)src, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + + TsmWaitfinish(); + + switch (fmt) { + case Fmt_FP16: + case Fmt_BF16: + *(uint16_t *)max_val = *(uint16_t *)&inst.param.wb_data0; + break; + case Fmt_FP32: + case Fmt_TF32: + *(uint32_t *)max_val = *(uint32_t *)&inst.param.wb_data0; + break; + default: + assert(0 && "ArgMax: Unsupport dtype"); + } + + *(uint32_t *)max_idx = *(uint32_t *)&inst.param.wb_data1; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/argmin.c b/third_party/wafer/crt/lib/Wafer/argmin.c new file mode 100755 index 00000000..0fe8f8a7 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/argmin.c @@ -0,0 +1,73 @@ +//===------------------------ argmin.c ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::ArgMin see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __ArgMin(uint64_t *src, uint64_t *dst0, uint64_t *dst1, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + volatile void *min_val = + (volatile void *)get_spm_memory_mapping((uint64_t)dst0); + volatile void *min_idx = + (volatile void *)get_spm_memory_mapping((uint64_t)dst1); + // Create command buffer. + if (elem_count == 1) { + volatile void *src_data = + (volatile void *)get_spm_memory_mapping((uint64_t)src); + *(uint32_t *)min_idx = 0; + switch (fmt) { + case Fmt_FP16: + case Fmt_BF16: + *(uint16_t *)min_val = *(uint16_t *)src_data; + break; + case Fmt_FP32: + case Fmt_TF32: + *(uint32_t *)min_val = *(uint32_t *)src_data; + break; + default: + assert(0 && "ArgMax: Unsupport dtype"); + } + return; + } + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmPeripheralInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + ; + + cmd->ArgMin(&inst, (uint64_t)src, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + + TsmWaitfinish(); + + switch (fmt) { + case Fmt_FP16: + case Fmt_BF16: + *(uint16_t *)min_val = *(uint16_t *)&inst.param.wb_data0; + break; + case Fmt_FP32: + case Fmt_TF32: + *(uint32_t *)min_val = *(uint32_t *)&inst.param.wb_data0; + break; + default: + assert(0 && "ArgMin: Unsupport dtype"); + } + + *(uint32_t *)min_idx = *(uint32_t *)&inst.param.wb_data1; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/arith.c b/third_party/wafer/crt/lib/Wafer/arith.c new file mode 100755 index 00000000..69318fc6 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/arith.c @@ -0,0 +1,240 @@ +//===------------------------ arith.c ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::ArithOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __AddVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + RND_MODE round, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->AddVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + round, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} + +void __SubVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + RND_MODE round, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->SubVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + round, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __MulVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + RND_MODE round, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->MulVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + round, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} + +void __DivVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + RND_MODE round, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->DivVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + round, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __AddVS(uint64_t *src0, uint32_t src1, uint64_t *dst, uint32_t elem_count, + RND_MODE round, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->AddVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, round, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __SubVS(uint64_t *src0, uint32_t src1, uint64_t *dst, uint32_t elem_count, + RND_MODE round, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->SubVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, round, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __MulVS(uint64_t *src0, uint32_t src1, uint64_t *dst, uint32_t elem_count, + RND_MODE round, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->MulVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, round, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __DivVS(uint64_t *src0, uint32_t src1, uint64_t *dst, uint32_t elem_count, + RND_MODE round, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->DivVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, round, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __MaxVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + RND_MODE reserved, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->MaxVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + reserved, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __MinVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + RND_MODE reserved, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->MinVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + reserved, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/assert.c b/third_party/wafer/crt/lib/Wafer/assert.c new file mode 100755 index 00000000..a1387504 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/assert.c @@ -0,0 +1,42 @@ +// ===------------------------ assert.c +// ------------------------------------===// + +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. + +// ===---------------------------------------------------------------------===// + +// Enable wafer kernel assert support + +#include "wafer.h" +#include +#include +#include + +void __Assert(const char *message, ...) { + INTRNISIC_RUN_SWITCH; + va_list args; + va_start(args, message); + + char *file = va_arg(args, char *); + int line = va_arg(args, int); + int col = va_arg(args, int); + int pidX = va_arg(args, int); + int pidY = va_arg(args, int); + int pidZ = va_arg(args, int); + va_end(args); + +#ifdef USE_SIM_MODE + printf("%s(line %d, col %d)::tile (%d, %d, %d): %s\n", file, line, col, pidX, + pidY, pidZ, message); + abort(); +#else + tsm_ep_log(__FILE__, __func__, __LINE__, KCORE_LOG_ERROR, + "%s(line %d, col %d)::tile (%d, %d, %d): %s\n", file, line, col, + pidX, pidY, pidZ, message); + // RT_ASSERT is an RT-Thread macro, not an exported firmware function. + // Call the SDK's non-returning newlib assertion entry directly: assert(0) + // would disappear from Release builds when NDEBUG is defined. + __assert_func(file, line, __func__, message); +#endif +} diff --git a/third_party/wafer/crt/lib/Wafer/atomic_barrier_in.c b/third_party/wafer/crt/lib/Wafer/atomic_barrier_in.c new file mode 100755 index 00000000..be65db94 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/atomic_barrier_in.c @@ -0,0 +1,21 @@ +//===------------------------ AtomicBarrierIn.c +//----------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::AtomicBarrierIn see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __AtomicBarrierIn() { +#ifdef USE_SIM_MODE +#else + atomic_barrier_in(); + SYNCHRONOUS_INTRINSIC_SWITCH; +#endif +} diff --git a/third_party/wafer/crt/lib/Wafer/atomic_barrier_out.c b/third_party/wafer/crt/lib/Wafer/atomic_barrier_out.c new file mode 100755 index 00000000..5b7e0706 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/atomic_barrier_out.c @@ -0,0 +1,20 @@ +//===------------------------ AtomicBarrierOut.c --------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::AtomicBarrierOut see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __AtomicBarrierOut() { +#ifdef USE_SIM_MODE +#else + atomic_barrier_out(); + SYNCHRONOUS_INTRINSIC_SWITCH; +#endif +} diff --git a/third_party/wafer/crt/lib/Wafer/barrier.c b/third_party/wafer/crt/lib/Wafer/barrier.c new file mode 100755 index 00000000..df93932e --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/barrier.c @@ -0,0 +1,17 @@ +//===------------------------ Barrier.c -----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Barrier see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Barrier() { + INTRNISIC_RUN_SWITCH; + TsmWaitfinish(); +} diff --git a/third_party/wafer/crt/lib/Wafer/bf16_fp16.c b/third_party/wafer/crt/lib/Wafer/bf16_fp16.c new file mode 100755 index 00000000..4e7b8c8c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/bf16_fp16.c @@ -0,0 +1,32 @@ +//===------------------------ bf16_fp16.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::BF16_FP16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __BF16_FP16(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BF16_FP16(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/bf16_fp32.c b/third_party/wafer/crt/lib/Wafer/bf16_fp32.c new file mode 100755 index 00000000..ba5e490a --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/bf16_fp32.c @@ -0,0 +1,31 @@ +//===------------------------ bf16_fp32.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::BF16_FP32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __BF16_FP32(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BF16_FP32(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/bf16_int16.c b/third_party/wafer/crt/lib/Wafer/bf16_int16.c new file mode 100755 index 00000000..5809400e --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/bf16_int16.c @@ -0,0 +1,33 @@ +//===------------------------ bf16_int16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::BF16_INT16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __BF16_INT16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BF16_INT16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/bf16_int32.c b/third_party/wafer/crt/lib/Wafer/bf16_int32.c new file mode 100755 index 00000000..38878dfd --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/bf16_int32.c @@ -0,0 +1,33 @@ +//===------------------------ bf16_int32.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::BF16_INT32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __BF16_INT32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BF16_INT32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/bf16_int8.c b/third_party/wafer/crt/lib/Wafer/bf16_int8.c new file mode 100755 index 00000000..979fb73d --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/bf16_int8.c @@ -0,0 +1,32 @@ +//===------------------------ bf16_int8.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::BF16_INT8 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __BF16_INT8(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BF16_INT8(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/bf16_tf32.c b/third_party/wafer/crt/lib/Wafer/bf16_tf32.c new file mode 100755 index 00000000..0619528f --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/bf16_tf32.c @@ -0,0 +1,32 @@ +//===------------------------ bf16_tf32.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::BF16_TF32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __BF16_TF32(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BF16_TF32(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/bilinear.c b/third_party/wafer/crt/lib/Wafer/bilinear.c new file mode 100755 index 00000000..efddeb7d --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/bilinear.c @@ -0,0 +1,40 @@ +//===------------------------ bilinear.c ----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Bilinear see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Bilinear(uint64_t *src, uint64_t *dst, uint16_t src_n, uint16_t src_h, + uint16_t src_w, uint16_t src_c, uint16_t dst_n, uint16_t dst_h, + uint16_t dst_w, uint16_t dst_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmPeripheralInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + ; + + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + Data_Shape shape2 = {dst_n, dst_h, dst_w, dst_c}; + cmd->Bilinear(&inst, (uint64_t)src, (uint64_t)dst, shape1, shape2, + (src_w - 1) / (dst_w - 1), (src_h - 1) / (dst_h - 1), + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/bit2fp.c b/third_party/wafer/crt/lib/Wafer/bit2fp.c new file mode 100755 index 00000000..3acb7518 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/bit2fp.c @@ -0,0 +1,37 @@ +//===------------------------ bit2fp.c ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Bit2Fp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Bit2Fp(uint64_t *src, uint64_t *target, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmPeripheralInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + // assert(elem_count % 8 == 0); + + cmd->Bit2Fp(&inst, (uint64_t)src, (uint64_t)target, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/channelnorm.c b/third_party/wafer/crt/lib/Wafer/channelnorm.c new file mode 100755 index 00000000..f903539a --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/channelnorm.c @@ -0,0 +1,156 @@ +//===------------------------ channelnorm.c -------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::channelnorm/dechannelnorm. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" +#include + +void __ChannelNorm(uint64_t *src, uint64_t *dst, uint16_t n, uint16_t h, + uint16_t w, uint16_t c, uint16_t c0, uint16_t bit_width) { + INTRNISIC_RUN_SWITCH; + int calign_base = bit_width == 8 ? 128 : 64; + int dtype_size = bit_width / 8; + int cx = c / calign_base; + + TsmDataMove *dm = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr dm_param = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + uint32_t inner_dim_size = c; + St_StrideIteration src_it = {0}, dst_it = {0}; + + // align cx + if (cx > 0) { + uint32_t elem_size = calign_base * dtype_size; + // byte number + src_it.stride0 = inner_dim_size * dtype_size; + src_it.iteration0 = h * w; + src_it.stride1 = elem_size; + src_it.iteration1 = cx; + src_it.stride2 = h * w * c * dtype_size; + src_it.iteration2 = n; + + dst_it.stride0 = elem_size; + dst_it.iteration0 = h * w; + dst_it.stride1 = dst_it.iteration0 * dst_it.stride0; + dst_it.iteration1 = cx; + dst_it.stride2 = n * h * w * elem_size; + dst_it.iteration2 = n; + + dm->GatherScatter(&dm_param, (uint64_t)src, (uint64_t)dst, elem_size, + &src_it, &dst_it); + TsmExecute(&dm_param); + TsmWaitfinish(); + } + + // align c0 + if (c0 > 0) { + uint32_t src_offset = cx * calign_base * dtype_size; + uint32_t dst_offset = cx * h * w * calign_base * dtype_size; + int32_t c0_valid = inner_dim_size - cx * calign_base; + int32_t elem_size = c0_valid * dtype_size; + + src_it.stride0 = inner_dim_size * dtype_size; + src_it.iteration0 = n * h * w; + src_it.stride1 = n * h * w * inner_dim_size * dtype_size; + src_it.iteration1 = 1; + src_it.stride2 = n * h * w * inner_dim_size * dtype_size; + src_it.iteration2 = 1; + + dst_it.stride0 = c0 * dtype_size; + dst_it.iteration0 = n * h * w; + dst_it.stride1 = n * h * w * c0 * dtype_size; + dst_it.iteration1 = 1; + dst_it.stride2 = n * h * w * c0 * dtype_size; + dst_it.iteration2 = 1; + + dm->GatherScatter(&dm_param, (uint64_t)src + src_offset, + (uint64_t)dst + dst_offset, elem_size, &src_it, &dst_it); + TsmExecute(&dm_param); + TsmWaitfinish(); + } +} + +void __DechannelNorm(uint64_t *src, uint64_t *dst, uint16_t n, uint16_t h, + uint16_t w, uint16_t c, uint16_t c0, uint16_t bit_width) { + INTRNISIC_RUN_SWITCH; + int calign_base = bit_width == 8 ? 128 : 64; + int dtype_size = bit_width / 8; + int cx = c / calign_base; + + TsmDataMove *dm = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr dm_param = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + uint32_t inner_dim_size = c; + St_StrideIteration src_it = {0}, dst_it = {0}; + + // align cx + if (cx > 0) { + uint32_t elem_size = calign_base * dtype_size; + // byte number + + src_it.stride0 = h * w * elem_size; + src_it.iteration0 = cx; + src_it.stride1 = elem_size; + src_it.iteration1 = h * w; + src_it.stride2 = cx * h * w * elem_size; + src_it.iteration2 = n; + + dst_it.stride0 = elem_size; + dst_it.iteration0 = cx; + dst_it.stride1 = inner_dim_size * dtype_size; + dst_it.iteration1 = h * w; + dst_it.stride2 = h * w * inner_dim_size * dtype_size; + dst_it.iteration2 = n; + + dm->GatherScatter(&dm_param, (uint64_t)src, (uint64_t)dst, elem_size, + &src_it, &dst_it); + TsmExecute(&dm_param); + TsmWaitfinish(); + } + + // align c0 + if (c0 > 0) { + uint32_t src_offset = cx * calign_base * dtype_size; + uint32_t dst_offset = cx * h * w * calign_base * dtype_size; + int32_t c0_valid = inner_dim_size - cx * calign_base; + int32_t elem_size = c0_valid * dtype_size; + + src_it.stride0 = c0 * dtype_size; + src_it.iteration0 = n * h * w; + src_it.stride1 = n * h * w * c0 * dtype_size; + src_it.iteration1 = 1; + src_it.stride2 = n * h * w * c0 * dtype_size; + src_it.iteration2 = 1; + + dst_it.stride0 = inner_dim_size * dtype_size; + dst_it.iteration0 = n * h * w; + dst_it.stride1 = n * h * w * inner_dim_size * dtype_size; + dst_it.iteration1 = 1; + dst_it.stride2 = n * h * w * inner_dim_size * dtype_size; + dst_it.iteration2 = 1; + + dm->GatherScatter(&dm_param, (uint64_t)src + src_offset, + (uint64_t)dst + dst_offset, elem_size, &src_it, &dst_it); + TsmExecute(&dm_param); + TsmWaitfinish(); + } +} diff --git a/third_party/wafer/crt/lib/Wafer/common.c b/third_party/wafer/crt/lib/Wafer/common.c new file mode 100755 index 00000000..2587ce7c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/common.c @@ -0,0 +1,19 @@ +//===----------------------- common.c -------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Implement common helper functions in this file. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +// WORKAROUND for undefined symbols in libkcorert.a +int main(int argc, char **argv) { return 0; } + +int get_app_version() { return 1; } + +int nvram_get_val() { return 1; } diff --git a/third_party/wafer/crt/lib/Wafer/concat.c b/third_party/wafer/crt/lib/Wafer/concat.c new file mode 100755 index 00000000..1b87ea36 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/concat.c @@ -0,0 +1,41 @@ +//===------------------------ concat.c ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Concat see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Concat(uint64_t *src1, uint16_t src1_n, uint16_t src1_h, uint16_t src1_w, + uint16_t src1_c, uint64_t *src2, uint16_t src2_n, uint16_t src2_h, + uint16_t src2_w, uint16_t src2_c, uint64_t *dst, uint16_t dst_n, + uint16_t dst_h, uint16_t dst_w, uint16_t dst_c, uint32_t dim, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src1_n, src1_h, src1_w, src1_c}; + Data_Shape shape2 = {src2_n, src2_h, src2_w, src2_c}; + Data_Shape shape3 = {dst_n, dst_h, dst_w, dst_c}; + cmd->Concat(&inst, (uint64_t)src1, shape1, (uint64_t)src2, shape2, + (uint64_t)dst, shape3, dim, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/conv.c b/third_party/wafer/crt/lib/Wafer/conv.c new file mode 100755 index 00000000..a5c284d9 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/conv.c @@ -0,0 +1,67 @@ +//===------------------------ conv.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TsmConv, see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +// The arguments list is aligned with TsmConv in WaferOps.td +void __Conv(int64_t opType, int64_t *srcAct, int64_t *srcActDims, + int64_t *weight, int64_t *weightDims, bool enBias, int64_t *bias, + bool enNegScale, int64_t *negScale, bool enPosScale, + int64_t *posScale, bool enSparse, int64_t *sparse, bool enPsum, + int64_t *psum, int64_t *pads, int64_t *unpads, int64_t *strides, + int64_t *dilations, bool enLeakyRelu, int64_t srcActFmt, + int64_t weightFmt, int64_t dstFmt, int64_t *dst, int64_t *dstDims) { + INTRNISIC_RUN_SWITCH; + // Create convolution command buffer. + TsmConv *conv = g_intrinsic()->conv_pointer; + TsmNeInstr inst = {I_NEUR, + { + 0, + }, + { + 0, + }}; + + // Convert to nhwc format + Data_Shape shape = {(uint16_t)srcActDims[0], (uint16_t)srcActDims[1], + (uint16_t)srcActDims[2], (uint16_t)srcActDims[3]}; + + Data_Shape wshape = {(uint16_t)weightDims[0], (uint16_t)weightDims[1], + (uint16_t)weightDims[2], (uint16_t)weightDims[3]}; + + Data_Shape dstShape = {(uint16_t)dstDims[0], (uint16_t)dstDims[1], + (uint16_t)dstDims[2], (uint16_t)dstDims[3]}; + + conv->AddInput(&inst, (int64_t)srcAct, shape, (Data_Format)srcActFmt); + conv->AddWeight(&inst, (uint64_t)weight, wshape, (Data_Format)weightFmt); + conv->AddBias(&inst, enBias, (uint64_t)bias); + conv->AddOutput(&inst, (uint64_t)dst, dstShape, (Data_Format)dstFmt); + conv->SetOpType(&inst, opType); + conv->SetNegativeAxisScale(&inst, enNegScale, (uint64_t)negScale); + conv->SetPositiveAxisScale(&inst, enPosScale, (uint64_t)posScale); + conv->SetSparse(&inst, enSparse, (uint64_t)sparse); + // FIXME: Should we have psum format instead? + conv->SetPsum(&inst, enPsum, (uint64_t)psum, (Data_Format)dstFmt); + conv->SetPads(&inst, pads[0], pads[1], pads[2], pads[3]); + conv->SetUnPads(&inst, unpads[0], unpads[1], unpads[2], unpads[3]); + conv->SetKernelStrides(&inst, strides[0], strides[1], strides[2], strides[3]); + conv->SetDilations(&inst, dilations[0], dilations[1]); + if (enLeakyRelu) + conv->EnableLeakyRelu(&inst); + else + conv->EnableRelu(&inst); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/cos.c b/third_party/wafer/crt/lib/Wafer/cos.c new file mode 100755 index 00000000..9d3f68bd --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/cos.c @@ -0,0 +1,33 @@ +//===------------------------ cos.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Cos see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Cos(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmTranscendental *cmd = g_intrinsic()->transcendental_pointer; + TsmTranscendentalInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Cos(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/count.c b/third_party/wafer/crt/lib/Wafer/count.c new file mode 100755 index 00000000..a7413486 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/count.c @@ -0,0 +1,34 @@ +//===------------------------ count.c -------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Count see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Count(uint64_t *src, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmPeripheralInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + ; + + cmd->Count(&inst, (uint64_t)src, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/empty.c b/third_party/wafer/crt/lib/Wafer/empty.c new file mode 100755 index 00000000..e3ad42e2 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/empty.c @@ -0,0 +1,7 @@ +#include + +void sqrt(int x) { assert(0 && "Unexpected execute sqrt function\n"); } +void floor(int x) { assert(0 && "Unexpected execute floor function\n"); } +void fmin(int x) { assert(0 && "Unexpected execute fmin function\n"); } +void fmax(int x) { assert(0 && "Unexpected execute fmax function\n"); } +void ceil(int x) { assert(0 && "Unexpected execute fmax function\n"); } diff --git a/third_party/wafer/crt/lib/Wafer/exp.c b/third_party/wafer/crt/lib/Wafer/exp.c new file mode 100755 index 00000000..b2a0493f --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/exp.c @@ -0,0 +1,33 @@ +//===------------------------ exp.c ---------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Exp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Exp(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmTranscendental *cmd = g_intrinsic()->transcendental_pointer; + TsmTranscendentalInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Exp(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/explp.c b/third_party/wafer/crt/lib/Wafer/explp.c new file mode 100755 index 00000000..e1c6e6d0 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/explp.c @@ -0,0 +1,33 @@ +//===------------------------ explp.c -------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Explp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Explp(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmTranscendental *cmd = g_intrinsic()->transcendental_pointer; + TsmTranscendentalInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Explp(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp16_bf16.c b/third_party/wafer/crt/lib/Wafer/fp16_bf16.c new file mode 100755 index 00000000..e38a313d --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp16_bf16.c @@ -0,0 +1,33 @@ +//===------------------------ fp16_bf16.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP16_BF16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP16_BF16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP16_BF16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp16_fp32.c b/third_party/wafer/crt/lib/Wafer/fp16_fp32.c new file mode 100755 index 00000000..aae15461 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp16_fp32.c @@ -0,0 +1,31 @@ +//===------------------------ fp16_fp32.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP16_FP32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP16_FP32(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP16_FP32(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp16_int16.c b/third_party/wafer/crt/lib/Wafer/fp16_int16.c new file mode 100755 index 00000000..c0d295e0 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp16_int16.c @@ -0,0 +1,33 @@ +//===------------------------ fp16_int16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP16_INT16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP16_INT16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP16_INT16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp16_int32.c b/third_party/wafer/crt/lib/Wafer/fp16_int32.c new file mode 100755 index 00000000..ae40c930 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp16_int32.c @@ -0,0 +1,33 @@ +//===------------------------ fp16_int32.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP16_INT32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP16_INT32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP16_INT32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp16_int8.c b/third_party/wafer/crt/lib/Wafer/fp16_int8.c new file mode 100755 index 00000000..73a3c486 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp16_int8.c @@ -0,0 +1,32 @@ +//===------------------------ fp16_int8.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP16_INT8 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP16_INT8(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP16_INT8(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp16_tf32.c b/third_party/wafer/crt/lib/Wafer/fp16_tf32.c new file mode 100755 index 00000000..68a91956 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp16_tf32.c @@ -0,0 +1,32 @@ +//===------------------------ fp16_tf32.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP16_TF32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP16_TF32(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP16_TF32(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp32_bf16.c b/third_party/wafer/crt/lib/Wafer/fp32_bf16.c new file mode 100755 index 00000000..662f2a13 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp32_bf16.c @@ -0,0 +1,33 @@ +//===------------------------ fp32_bf16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP32_BF16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP32_BF16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP32_BF16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp32_fp16.c b/third_party/wafer/crt/lib/Wafer/fp32_fp16.c new file mode 100755 index 00000000..47d64698 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp32_fp16.c @@ -0,0 +1,32 @@ +//===------------------------ fp32_fp16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP32_FP16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP32_FP16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP32_FP16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp32_int16.c b/third_party/wafer/crt/lib/Wafer/fp32_int16.c new file mode 100755 index 00000000..79ee8383 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp32_int16.c @@ -0,0 +1,33 @@ +//===------------------------ fp32_int16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP32_INT16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP32_INT16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP32_INT16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp32_int32.c b/third_party/wafer/crt/lib/Wafer/fp32_int32.c new file mode 100755 index 00000000..6efd31c0 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp32_int32.c @@ -0,0 +1,32 @@ +//===------------------------ fp32_int32.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP32_INT32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP32_INT32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP32_INT32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp32_int8.c b/third_party/wafer/crt/lib/Wafer/fp32_int8.c new file mode 100755 index 00000000..cc4e7cc2 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp32_int8.c @@ -0,0 +1,33 @@ +//===------------------------ fp32_int8.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP32_INT8 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP32_INT8(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP32_INT8(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/fp32_tf32.c b/third_party/wafer/crt/lib/Wafer/fp32_tf32.c new file mode 100755 index 00000000..418e0aa5 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/fp32_tf32.c @@ -0,0 +1,32 @@ +//===------------------------ fp32_tf32.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::FP32_TF32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __FP32_TF32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->FP32_TF32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/gatherscatter.c b/third_party/wafer/crt/lib/Wafer/gatherscatter.c new file mode 100755 index 00000000..a9221b07 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/gatherscatter.c @@ -0,0 +1,44 @@ +//===------------------------ gatherscatter.c -----------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::GatherScatter see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __GatherScatter(uint64_t *src, uint64_t *dst, uint32_t bytes, + uint32_t src_strideN, uint32_t src_strideH, + uint32_t src_strideW, uint32_t src_iterN, + uint32_t src_iterH, uint32_t src_iterW, + uint32_t dst_strideN, uint32_t dst_strideH, + uint32_t dst_strideW, uint32_t dst_iterN, + uint32_t dst_iterH, uint32_t dst_iterW) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + St_StrideIteration src_si = {src_strideW, src_iterW, src_strideH, + src_iterH, src_strideN, src_iterN}; + St_StrideIteration dst_si = {dst_strideW, dst_iterW, dst_strideH, + dst_iterH, dst_strideN, dst_iterN}; + + cmd->GatherScatter(&inst, (uint64_t)src, (uint64_t)dst, bytes, &src_si, + &dst_si); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/gelu_none.c b/third_party/wafer/crt/lib/Wafer/gelu_none.c new file mode 100755 index 00000000..4ef53e8c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/gelu_none.c @@ -0,0 +1,20 @@ +//===------------------------ gelu_none.c +//-----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::GeluNone see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "op_gelu.h" + +void __GeluNone(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + op_gelu_none(src, dst, elem_count, (Data_Format)fmt); + SYNCHRONOUS_INTRINSIC_SWITCH; +} diff --git a/third_party/wafer/crt/lib/Wafer/gelu_tanh.c b/third_party/wafer/crt/lib/Wafer/gelu_tanh.c new file mode 100755 index 00000000..bf5355ea --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/gelu_tanh.c @@ -0,0 +1,20 @@ +//===------------------------ gelu_tanh.c +//-----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::GeluTanh see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "op_gelu.h" + +void __GeluTanh(uint64_t *src, uint64_t *imm, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + op_gelu_tanh(src, imm, dst, elem_count, fmt); + SYNCHRONOUS_INTRINSIC_SWITCH; +} diff --git a/third_party/wafer/crt/lib/Wafer/gemm.c b/third_party/wafer/crt/lib/Wafer/gemm.c new file mode 100755 index 00000000..8dfb5303 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/gemm.c @@ -0,0 +1,58 @@ +//===------------------------ gemm.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TsmGemm, see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +// The arguments list is aligned with TsmConv in WaferOps.td +void __Gemm(int64_t *srcA, int64_t *srcB, int64_t *srcBias, int64_t *dst, + int32_t *dims, bool enPsum, int64_t *psum, bool enTransA, + bool enTransB, int64_t batchSizeA, int64_t batchSizeB, + int32_t reluMode, bool enBias, bool enNegScale, int64_t *negScale, + bool enPosScale, int64_t *posScale, int64_t srcFmt, + int64_t dstFmt) { + INTRNISIC_RUN_SWITCH; + // Create gemm command buffer. + TsmGemm *gemm = g_intrinsic()->gemm_pointer; + TsmNeInstr inst = {I_NEUR, + { + 0, + }, + { + 0, + }}; + + gemm->AddInput(&inst, (uint64_t)srcA, (uint64_t)srcB, (Data_Format)srcFmt); + gemm->ConfigMKN(&inst, (uint32_t)dims[0], (uint32_t)dims[1], + (uint32_t)dims[2]); + gemm->AddOutput(&inst, (uint64_t)dst, (Data_Format)dstFmt); + gemm->SetPsum(&inst, enPsum, (uint64_t)psum, (Data_Format)dstFmt); + gemm->SetTransflag(&inst, (uint8_t)enTransA, (uint8_t)enTransB); + // TODO: + // gemm->SetQuant(); + gemm->ConfigBatch(&inst, (uint32_t)batchSizeA, (uint32_t)batchSizeB); + gemm->AddBias(&inst, enBias, (uint64_t)srcBias); + gemm->SetNegativeAxisScale(&inst, enNegScale, (uint64_t)negScale); + gemm->SetPositiveAxisScale(&inst, enPosScale, (uint64_t)posScale); + switch (reluMode) { + case ENRelu: + gemm->EnableRelu(&inst); + break; + case ENLeakRelu: + gemm->EnableLeakyRelu(&inst); + break; + default: + break; + } + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; +} diff --git a/third_party/wafer/crt/lib/Wafer/img2col.c b/third_party/wafer/crt/lib/Wafer/img2col.c new file mode 100755 index 00000000..0a20df98 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/img2col.c @@ -0,0 +1,43 @@ +//===------------------------ img2col.c -----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Img2col see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Img2col(uint64_t *src, uint16_t src_n, uint16_t src_h, uint16_t src_w, + uint16_t src_c, uint64_t *dst, uint16_t dst_n, uint16_t dst_h, + uint16_t dst_w, uint16_t dst_c, uint64_t src_elem_num, + uint64_t dst_elem_num, uint16_t swr_n, uint16_t swr_h, + uint16_t swr_w, uint16_t swr_c, uint16_t pdr_n, uint16_t pdr_h, + uint16_t pdr_w, uint16_t pdr_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + Data_Shape shape2 = {dst_n, dst_h, dst_w, dst_c}; + Data_Shape shape3 = {swr_n, swr_h, swr_w, swr_c}; + Data_Shape shape4 = {pdr_n, pdr_h, pdr_w, pdr_c}; + cmd->Img2col(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + src_elem_num, dst_elem_num, shape3, shape4, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int16_bf16.c b/third_party/wafer/crt/lib/Wafer/int16_bf16.c new file mode 100755 index 00000000..bb9a841c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int16_bf16.c @@ -0,0 +1,33 @@ +//===------------------------ int16_bf16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT16_BF16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT16_BF16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT16_BF16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int16_fp16.c b/third_party/wafer/crt/lib/Wafer/int16_fp16.c new file mode 100755 index 00000000..202e992d --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int16_fp16.c @@ -0,0 +1,32 @@ +//===------------------------ int16_fp16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT16_FP16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT16_FP16(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT16_FP16(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int16_fp32.c b/third_party/wafer/crt/lib/Wafer/int16_fp32.c new file mode 100755 index 00000000..ad3837ac --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int16_fp32.c @@ -0,0 +1,32 @@ +//===------------------------ int16_fp32.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT16_FP32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT16_FP32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT16_FP32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int16_tf32.c b/third_party/wafer/crt/lib/Wafer/int16_tf32.c new file mode 100755 index 00000000..e2ecd8e8 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int16_tf32.c @@ -0,0 +1,32 @@ +//===------------------------ int16_tf32.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT16_TF32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT16_TF32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT16_TF32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int32_bf16.c b/third_party/wafer/crt/lib/Wafer/int32_bf16.c new file mode 100755 index 00000000..e76fc270 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int32_bf16.c @@ -0,0 +1,32 @@ +//===------------------------ int32_bf16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT32_BF16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT32_BF16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT32_BF16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int32_fp16.c b/third_party/wafer/crt/lib/Wafer/int32_fp16.c new file mode 100755 index 00000000..ac7cc1a7 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int32_fp16.c @@ -0,0 +1,33 @@ +//===------------------------ int32_fp16.cpp +//-------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT32_FP16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT32_FP16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT32_FP16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int32_fp32.c b/third_party/wafer/crt/lib/Wafer/int32_fp32.c new file mode 100755 index 00000000..30ac99ba --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int32_fp32.c @@ -0,0 +1,32 @@ +//===------------------------ int32_fp32.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT32_FP32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT32_FP32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT32_FP32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int32_tf32.c b/third_party/wafer/crt/lib/Wafer/int32_tf32.c new file mode 100755 index 00000000..07cbc5da --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int32_tf32.c @@ -0,0 +1,32 @@ +//===------------------------ int32_tf32.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT32_TF32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT32_TF32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT32_TF32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int8_bf16.c b/third_party/wafer/crt/lib/Wafer/int8_bf16.c new file mode 100755 index 00000000..1787f6c0 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int8_bf16.c @@ -0,0 +1,33 @@ +//===------------------------ int8_bf16.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT8_BF16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT8_BF16(uint64_t *src, uint64_t *dst, uint32_t zp, + uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT8_BF16(&inst, (uint64_t)src, zp, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int8_fp16.c b/third_party/wafer/crt/lib/Wafer/int8_fp16.c new file mode 100755 index 00000000..098a361f --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int8_fp16.c @@ -0,0 +1,33 @@ +//===------------------------ int8_fp16.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT8_FP16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT8_FP16(uint64_t *src, uint64_t *dst, uint32_t zp, + uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT8_FP16(&inst, (uint64_t)src, zp, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int8_fp32.c b/third_party/wafer/crt/lib/Wafer/int8_fp32.c new file mode 100755 index 00000000..e292bd6e --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int8_fp32.c @@ -0,0 +1,32 @@ +//===------------------------ int8_fp32.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT8_FP32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT8_FP32(uint64_t *src, uint64_t *dst, uint32_t zp, + uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT8_FP32(&inst, (uint64_t)src, zp, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/int8_tf32.c b/third_party/wafer/crt/lib/Wafer/int8_tf32.c new file mode 100755 index 00000000..8b8f7dee --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/int8_tf32.c @@ -0,0 +1,33 @@ +//===------------------------ int8_tf32.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::INT8_TF32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __INT8_TF32(uint64_t *src, uint64_t *dst, uint32_t zp, + uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->INT8_TF32(&inst, (uint64_t)src, zp, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/leakyrelu.c b/third_party/wafer/crt/lib/Wafer/leakyrelu.c new file mode 100755 index 00000000..c3d4f322 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/leakyrelu.c @@ -0,0 +1,34 @@ +//===------------------------ leakyrelu.c ---------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Leakyrelu see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Leakyrelu(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmActivation *cmd = g_intrinsic()->activation_pointer; + TsmActivationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Leakyrelu(&inst, (uint64_t)src, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/ln.c b/third_party/wafer/crt/lib/Wafer/ln.c new file mode 100755 index 00000000..4ea64268 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/ln.c @@ -0,0 +1,32 @@ +//===------------------------ ln.c ----------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Ln see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Ln(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmTranscendental *cmd = g_intrinsic()->transcendental_pointer; + TsmTranscendentalInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Ln(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/log2.c b/third_party/wafer/crt/lib/Wafer/log2.c new file mode 100755 index 00000000..52b662ef --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/log2.c @@ -0,0 +1,32 @@ +//===------------------------ log2.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Log2 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Log2(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmTranscendental *cmd = g_intrinsic()->transcendental_pointer; + TsmTranscendentalInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Log2(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/logic.c b/third_party/wafer/crt/lib/Wafer/logic.c new file mode 100755 index 00000000..b9c27491 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/logic.c @@ -0,0 +1,162 @@ +//===------------------------ logic.c -------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::LogicOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __AndVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmLogic *cmd = g_intrinsic()->logic_pointer; + TsmLogicInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->AndVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} + +void __OrVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmLogic *cmd = g_intrinsic()->logic_pointer; + TsmLogicInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->OrVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __XorVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmLogic *cmd = g_intrinsic()->logic_pointer; + TsmLogicInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->XorVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolNotV(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmLogic *cmd = g_intrinsic()->logic_pointer; + TsmLogicInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolNotV(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); +} + +void __BoolAndV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmLogic *cmd = g_intrinsic()->logic_pointer; + TsmLogicInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolAndV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); +} + +void __BoolOrV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmLogic *cmd = g_intrinsic()->logic_pointer; + TsmLogicInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolOrV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); +} + +void __BoolXorV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmLogic *cmd = g_intrinsic()->logic_pointer; + TsmLogicInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolXorV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); +} diff --git a/third_party/wafer/crt/lib/Wafer/lut16.c b/third_party/wafer/crt/lib/Wafer/lut16.c new file mode 100755 index 00000000..740357f3 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/lut16.c @@ -0,0 +1,35 @@ +//===------------------------ lut16.c -------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Lut16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Lut16(uint64_t *src, uint64_t *dst, uint64_t *lut16, + uint32_t src_elem_count, uint32_t lut_elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmPeripheralInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + ; + + cmd->Lut16(&inst, (uint64_t)src, (uint64_t)dst, (uint64_t)lut16, + src_elem_count, lut_elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/lut32.c b/third_party/wafer/crt/lib/Wafer/lut32.c new file mode 100755 index 00000000..2c739373 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/lut32.c @@ -0,0 +1,35 @@ +//===------------------------ lut32.c -------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Lut32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Lut32(uint64_t *src, uint64_t *dst, uint64_t *lut32, + uint32_t src_elem_count, uint32_t lut_elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmPeripheralInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + ; + + cmd->Lut32(&inst, (uint64_t)src, (uint64_t)dst, (uint64_t)lut32, + src_elem_count, lut_elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/mask_move.c b/third_party/wafer/crt/lib/Wafer/mask_move.c new file mode 100755 index 00000000..e5828871 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/mask_move.c @@ -0,0 +1,31 @@ +//===------------------------ mask_move.c ---------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::MaskMoveOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __MaskMove(uint64_t *src, uint64_t *target, uint32_t elem_count, + uint64_t *mask, int32_t fmt) { + INTRNISIC_RUN_SWITCH; + TsmMaskDataMove *move = g_intrinsic()->maskdatamove_pointer; + TsmMaskDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + move->MaskMove(&inst, (uint64_t)src, (uint64_t)mask, (uint64_t)target, + elem_count, (Data_Format)fmt); + + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; +} diff --git a/third_party/wafer/crt/lib/Wafer/memcpy.c b/third_party/wafer/crt/lib/Wafer/memcpy.c new file mode 100755 index 00000000..de347b4c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/memcpy.c @@ -0,0 +1,76 @@ +//===------------------------ memcpy.c ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::MemCopyOp, see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Memcpy(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + unsigned eleByte; + + switch (fmt) { + case Fmt_BOOL: { + // NOTE: Assume bool is 8 byte aligned + eleByte = 1; + elem_count = (elem_count + 7) / 8; + fmt = Fmt_INT8; + break; + } + case Fmt_INT8: { + eleByte = 1; + break; + } + case Fmt_INT16: + case Fmt_FP16: + case Fmt_BF16: { + eleByte = 2; + break; + } + case Fmt_INT32: + case Fmt_FP32: + case Fmt_TF32: { + eleByte = 4; + break; + } + case Fmt_INT64: { + eleByte = 8; + break; + } + default: + // Other formats are not supported. + assert(false && "Unsupported format\n"); + break; + } + + if (elem_count == 0) { + // If elem_count is 0, we don't need to do anything. + return; + } + + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + St_StrideIteration src_si = {1, 1, 1, 1, 1, 1}; + St_StrideIteration dst_si = {1, 1, 1, 1, 1, 1}; + + cmd->GatherScatter(&inst, (uint64_t)src, (uint64_t)dst, eleByte * elem_count, + &src_si, &dst_si); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); +} diff --git a/third_party/wafer/crt/lib/Wafer/memset.c b/third_party/wafer/crt/lib/Wafer/memset.c new file mode 100755 index 00000000..0c02283a --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/memset.c @@ -0,0 +1,49 @@ +//===------------------------ memset.c ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Memset see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Memset(char *dst, int value, int *dst_shape, int *dst_stride, int rank, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + // TODO: Use real stride and iteration, now accumulate all data to elem_count + int stride0 = 0; + int stride1 = 0; + int stride2 = 0; + + int iteration0 = 1; + int iteration1 = 1; + int iteration2 = 1; + + int elem_count = 1; + for (int i = 0; i < rank; i++) { + elem_count *= dst_shape[i]; + } + + St_StrideIteration si = {stride0, iteration0, stride1, + iteration1, stride1, iteration2}; + cmd->Memset(&inst, (uint64_t)dst, value, elem_count, &si, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/mirror.c b/third_party/wafer/crt/lib/Wafer/mirror.c new file mode 100755 index 00000000..786eb99e --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/mirror.c @@ -0,0 +1,37 @@ +//===------------------------ mirror.c ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Mirror see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Mirror(uint64_t *src, uint16_t src_n, uint16_t src_h, uint16_t src_w, + uint16_t src_c, uint64_t *dst, uint16_t dst_n, uint16_t dst_h, + uint16_t dst_w, uint16_t dst_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + Data_Shape shape2 = {dst_n, dst_h, dst_w, dst_c}; + cmd->Mirror(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/mxfp_bf16.c b/third_party/wafer/crt/lib/Wafer/mxfp_bf16.c new file mode 100755 index 00000000..364e2cc7 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/mxfp_bf16.c @@ -0,0 +1,291 @@ +//===------------------------ mxfp_bf16.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation +// tx::FP8E5M2ToBF16Op/tx::FP8E4M3ToBF16Op/tx::FP4E2M1ToBF16Op +// +//===----------------------------------------------------------------------===// +#include "wafer.h" +#include + +/** + * Converts an array of FP8 (E5M2) values to BF16 format + * + * @param src Input array of FP8 values (E5M2 format) + * @param dst Output array for BF16 values (must be pre-allocated) + * @param elem_count Number of elements to convert + * + * FP8 (E5M2) format: + * [S][EEEEE][MM] + * 1 sign bit, 5 exponent bits (bias=15), 2 mantissa bits + * + * BF16 output format: + * [S][EEEEEEEE][MMMMMMM] + * 1 sign bit, 8 exponent bits (bias=127), 7 mantissa bits + */ +void __FP8E5M2_BF16(uint8_t *src, uint16_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + src = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)src); + dst = (uint16_t *)get_spm_memory_mapping_wrapper((uint64_t)dst); + + for (uint32_t i = 0; i < elem_count; i++) { + // Extract FP8 components + uint8_t fp8 = src[i]; + uint8_t sign = fp8 & 0x80; // Isolate sign bit (10000000) + uint8_t exponent = (fp8 >> 2) & 0x1F; // Extract 5-bit exponent (01111100) + uint8_t mantissa = fp8 & 0x03; // Extract 2-bit mantissa (00000011) + + // Handle special cases + if (exponent == 0) { + // Handle FP8 E5M2 subnormal/zero (per OCP MX Spec 5.3.1: E=0 is + // subnormal/zero) + static const uint16_t subNormLut[] = {0x0000, 0x3780, 0x3800, 0x3840}; + + // Reconstruct BF16 format (per OCP MX Spec 5.3.1 and BF16 definition): + // [15] - Sign bit (from FP8 sign) + // [14:7] - 8-bit exponent (BF16 exponent) + // [6:0] - 7-bit mantissa (BF16 mantissa) + dst[i] = (uint16_t)(sign << 8) | subNormLut[mantissa]; + continue; + } + + if (exponent == 0x1F) { + // NaN/Infinity: Preserve sign and mantissa, set max exponent + dst[i] = (sign << 8) | (0x1F << 10) | (mantissa << 7); + continue; + } + + // Convert exponent from FP8 (bias=15) to BF16 (bias=127) + // Formula: E_bf16 = E_fp8 + (127 - 15) = E_fp8 + 112 + uint16_t bf16_exponent = (uint16_t)(exponent + 112) << 7; + + // Reconstruct BF16 format: + // [15] - Sign bit + // [14:7] - 8-bit exponent + // [6:0] - 7-bit mantissa (FP8's 2-bit mantissa becomes bits [6:5]) + dst[i] = (sign << 8) | // Sign bit at bit 15 + bf16_exponent | // Exponent at bits 14-7 + (mantissa << 5); // Mantissa at bits 6-5 (bits 4-0 zero) + } + SYNCHRONOUS_INTRINSIC_SWITCH; +} + +/** + * Converts an array of FP8 (E4M3) values to BF16 format + * + * @param src Input array of FP8 values (E4M3 format) + * @param dst Output array for BF16 values (must be pre-allocated) + * @param elem_count Number of elements to convert + * + * FP8 (E4M3) format: + * [S][EEEE][MMM] + * 1 sign bit, 4 exponent bits (bias=7), 3 mantissa bits + * + * Note: E4M3 has no infinities. Exponent=15 (0xF) represents NaNs + * + * BF16 output format: + * [S][EEEEEEEE][MMMMMMM] + * 1 sign bit, 8 exponent bits (bias=127), 7 mantissa bits + */ +void __FP8E4M3_BF16(uint8_t *src, uint16_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + src = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)src); + dst = (uint16_t *)get_spm_memory_mapping_wrapper((uint64_t)dst); + + for (uint32_t i = 0; i < elem_count; i++) { + // Extract FP8 components + uint8_t fp8 = src[i]; + uint8_t sign = fp8 & 0x80; // Isolate sign bit (10000000) + uint8_t exponent = (fp8 >> 3) & 0x0F; // Extract 4-bit exponent (00001111) + uint8_t mantissa = fp8 & 0x07; // Extract 3-bit mantissa (00000111) + + // Handle special cases + if (exponent == 0) { + // Denormal/subnormal: Flush to zero (preserving sign only) + // E4M3 denormals are not supported in this implementation + dst[i] = sign << 8; + continue; + } + + if (exponent == 0x0F) { + // NaN case (E4M3 has no infinities) + // Set BF16 exponent to all 1s (0xFF) and preserve mantissa + // Shift mantissa to top 3 bits of BF16 mantissa field + dst[i] = (sign << 8) | (0xFF << 7) | (mantissa << 4); + continue; + } + + // Convert exponent from FP8 (bias=7) to BF16 (bias=127) + // Formula: E_bf16 = E_fp8 + (127 - 7) = E_fp8 + 120 + uint8_t bf16_exponent = exponent + 120; + + // Reconstruct BF16 format: + // [15] - Sign bit + // [14:7] - 8-bit exponent + // [6:0] - 7-bit mantissa + // Shift FP8 mantissa to bits [6:4] of BF16 mantissa field + dst[i] = (sign << 8) | // Sign bit at position 15 + (bf16_exponent << 7) | // Exponent at bits 14-7 + (mantissa << 4); // Mantissa at bits 6-4 (bits 3-0 zero) + } + SYNCHRONOUS_INTRINSIC_SWITCH; +} + +/** + * Converts an array of FP8 (E4M3FN) values to BF16 format (supports subnormal + * numbers) + * + * @param src Input array of FP8 values (E4M3FN: Finite Normal + + * Subnormal) + * @param dst Output array for BF16 values (must be pre-allocated) + * @param elem_count Number of elements to convert + * + * Based on: OCP Microscaling Formats (MX) Specification v1.0 + * - FP8 E4M3 format: 1 sign bit (S) + 4 exponent bits (E, bias=7) + 3 mantissa + * bits (M) + * - Subnormal numbers (E=0): v = (-1)^S × 2^(1-bias) × (M/8) (Section 5.3.1) + * - NaN (E=0xF): No infinities defined; E=0xF represents NaNs (Table 2) + * - BF16 format: 1 sign bit + 8 exponent bits (bias=127) + 7 mantissa bits + * (IEEE 754 compatible) + */ +void __FP8E4M3FN_BF16(uint8_t *src, uint16_t *dst, uint32_t elem_count) { + // Memory mapping wrapper (retained as original) + src = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)src); + dst = (uint16_t *)get_spm_memory_mapping_wrapper((uint64_t)dst); + + for (uint32_t i = 0; i < elem_count; i++) { + uint8_t fp8 = src[i]; + + // 1. Extract FP8 E4M3FN core fields (per Section 5.3.1) + uint8_t sign_bit = (fp8 >> 7) & 0x01; // Sign bit: 0=positive, 1=negative + uint8_t exponent = (fp8 >> 3) & 0x0F; // 4-bit exponent (E: 0~0xF) + uint8_t mantissa = fp8 & 0x07; // 3-bit mantissa (M: 0~7) + + // 2. Handle special case: NaN (E=0xF, explicitly no infinities in E4M3 per + // Table 2) + if (fp8 == 0x7F || fp8 == 0xFF) { // Positive or negative NaN + // BF16 NaN rule: all-1s exponent (0xFF) + non-zero mantissa (preserve + // E4M3 mantissa) + uint16_t bf16_sign = (uint16_t)sign_bit << 15; // Sign bit at BF16 bit 15 + uint16_t bf16_exponent = 0xFF + << 7; // Exponent bits (BF16 bits 14~7) all 1s + uint16_t bf16_mantissa = + (uint16_t)mantissa + << 4; // E4M3 mantissa occupies top 3 bits of BF16 mantissa (bits 6~4) + dst[i] = bf16_sign | bf16_exponent | bf16_mantissa; + continue; + } + + // 3. Handle subnormal numbers (E=0, per Section 5.3.1 formula) + if (exponent == 0) { + + static const int denormsAndZeroLut[8] = {0x0000, 0x3b00, 0x3b80, 0x3bc0, + 0x3c00, 0x3c20, 0x3c40, 0x3c60}; + dst[i] = (sign_bit << 15) | denormsAndZeroLut[mantissa]; + + continue; + } + + // 4. Handle normal numbers (E=1~14, per Section 5.3.1 formula) + // 4.1 Exponent conversion: E_BF16 = (E_FP8 - bias_FP8) + bias_BF16 = E + + // (127-7) = E + 120 + uint8_t bf16_exponent = exponent + 120; + // (Note: E=1→121, E=14→134, all within BF16 normal exponent range [1,254], + // no overflow handling needed) + + // 4.2 Reconstruct BF16 format (sign + exponent + mantissa) + uint16_t bf16_sign = (uint16_t)sign_bit << 15; // BF16 bit 15: sign + uint16_t bf16_exp_bits = (uint16_t)bf16_exponent + << 7; // BF16 bits 14~7: exponent + uint16_t bf16_mant_bits = + (uint16_t)mantissa + << 4; // BF16 bits 6~4: E4M3 mantissa (lower 4 bits zero-padded) + dst[i] = bf16_sign | bf16_exp_bits | bf16_mant_bits; + } +} + +/** + * Converts packed FP4 (E2M1) values to BF16 format + * + * @param src Input array of packed FP4 values (2 values per byte) + * @param dst Output array for BF16 values (must be pre-allocated) + * @param elem_count Number of FP4 elements (not bytes) + * + * FP4 (E2M1) format (per element): + * [S][EE][M] + * 1 sign bit (bit3), 2 exponent bits (bit2-1), 1 mantissa bit (bit0) + * + * Storage format: + * Each byte contains two FP4 values: + * [S1 E1 E0 M1] [S0 E1 E0 M0] (high nibble first, corrected symmetry) + */ +void __FP4E2M1_BF16(uint8_t *src, uint16_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + src = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)src); + dst = (uint16_t *)get_spm_memory_mapping_wrapper((uint64_t)dst); + + // Constants matching Python logic + const uint16_t to_bias = 127; // BF16 exponent bias + const uint16_t to_m_bits = 7; // BF16 mantissa bits + const uint16_t to_point5 = 16128; // BF16 value for ±0.5 (subnormal mapping) + const uint16_t bias_offset = (to_bias - 1) + << to_m_bits; // (127-1) <<7 = 126<<7 + + for (uint32_t i = 0; i < (elem_count + 1) / 2; i++) { + uint8_t byte = src[i]; + + // Process high nibble (first element, bits 7-4) + if (2 * i + 1 < elem_count) { + uint8_t elem = (byte >> 4) & 0x0F; // Extract 4-bit FP4 element + uint16_t sign = elem & 0x08; // Sign bit (bit3 of element) + uint16_t exp = (elem >> 1) & 0x03; // Exponent bits (bit2-1 of element) + uint16_t mant = elem & 0x01; // Mantissa bit (bit0 of element) + + uint16_t result; + if (exp == 0) { + // Subnormal or zero (exp=0) + if (mant == 1) { + // Subnormal: map to ±0.5 (to_point5 + sign) + result = + (sign << 12) | to_point5; // Sign shifted to BF16 bit15 (3+12=15) + } else { + // Zero: only preserve sign + result = sign << 12; + } + } else { + // Normal number: base mantissa + bias offset + uint16_t abs = elem & 0x07; + uint16_t base = (abs << (to_m_bits - 1)) | + (sign << 12); // Mantissa shifted to BF16 bit6 (7-1=6) + result = base + bias_offset; + } + dst[2 * i + 1] = result; + } + + // Process low nibble (second element, bits 3-0) if needed + { + uint8_t elem = byte & 0x0F; // Extract 4-bit FP4 element + uint16_t sign = elem & 0x08; // Sign bit (bit3 of element) + uint16_t exp = (elem >> 1) & 0x03; // Exponent bits (bit2-1 of element) + uint16_t mant = elem & 0x01; // Mantissa bit (bit0 of element) + + uint16_t result; + if (exp == 0) { + if (mant == 1) { + result = (sign << 12) | to_point5; + } else { + result = sign << 12; + } + } else { + uint16_t abs = elem & 0x07; + uint16_t base = (abs << (to_m_bits - 1)) | (sign << 12); + result = base + bias_offset; + } + dst[2 * i] = result; + } + } + SYNCHRONOUS_INTRINSIC_SWITCH; +} diff --git a/third_party/wafer/crt/lib/Wafer/mxfp_fp16.c b/third_party/wafer/crt/lib/Wafer/mxfp_fp16.c new file mode 100755 index 00000000..7b6bcf6c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/mxfp_fp16.c @@ -0,0 +1,232 @@ +//===------------------------ mxfp_fp16.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation +// tx::FP8E5M2ToFP16Op/tx::FP8E4M3ToFP16Op/tx::FP4E2M1ToFP16Op +// +//===----------------------------------------------------------------------===// +#include "wafer.h" +#include + +/** + * Converts an array of FP8 (E5M2) values to FP16 format + * + * @param src Input array of FP8 values (E5M2 format) + * @param dst Output array for FP16 values (must be pre-allocated) + * @param elem_count Number of elements to convert + * + * FP8 (E5M2) format: + * [S][EEEEE][MM] + * 1 sign bit, 5 exponent bits (bias=15), 2 mantissa bits + * + * FP16 output format: + * [S][EEEE E][MMMM MMMMMM] + * 1 sign bit, 5 exponent bits (bias=15), 10 mantissa bits + */ +void __FP8E5M2_FP16(uint8_t *src, uint16_t *dst, uint32_t elem_count) { + src = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)src); + dst = (uint16_t *)get_spm_memory_mapping_wrapper((uint64_t)dst); + + // E5M2 and FP16 share the sign/exponent layout and exponent bias. + // Extending the significand preserves zero, subnormals, infinity and NaN + // payloads; exponent=0 must remain zero in the destination format. + for (uint32_t i = 0; i < elem_count; i++) { + dst[i] = (uint16_t)src[i] << 8; + } +} + +/** + * Converts an array of FP8 (E4M3) values to FP16 format + * + * @param src Input array of FP8 values (E4M3 format) + * @param dst Output array for FP16 values (must be pre-allocated) + * @param elem_count Number of elements to convert + * + * FP8 (E4M3) format: + * [S][EEEE][MMM] + * 1 sign bit, 4 exponent bits (bias=7), 3 mantissa bits + * + * FP16 output format: + * [S][EEEE E][MMMM MMMMMM] + * 1 sign bit, 5 exponent bits (bias=15), 10 mantissa bits + */ +void __FP8E4M3_FP16(uint8_t *src, uint16_t *dst, uint32_t elem_count) { + src = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)src); + dst = (uint16_t *)get_spm_memory_mapping_wrapper((uint64_t)dst); + + for (uint32_t i = 0; i < elem_count; i++) { + uint8_t fp8 = src[i]; + uint8_t sign = fp8 & 0x80; // Extract sign bit (10000000) + uint8_t exponent = (fp8 >> 3) & 0x0F; // Extract 4-bit exponent (00001111) + uint8_t mantissa = fp8 & 0x07; // Extract 3-bit mantissa (00000111) + + // Handle subnormal values (exponent = 0) + if (exponent == 0) { + // According to OCP spec: v = (-1)^S × 2^(1-7) × (0 + 2^(-3) × M) + // FP16 exponent = (1 - 7) + 15 = 9 (bias adjustment) + static const int denormsAndZeroLut[8] = {0x0000, 0x1800, 0x1C00, 0x1E00, + 0x2000, 0x2100, 0x2200, 0x2300}; + dst[i] = (sign << 8) | denormsAndZeroLut[mantissa]; + continue; + } + + // Handle NaN case (exponent = 0x0F for E4M3) + if (exponent == 0x0F) { + // Set FP16 max exponent (0x1F) and extend mantissa + dst[i] = (sign << 8) | (0x1F << 10) | (mantissa << 7); + continue; + } + + // Normal case conversion + // Convert exponent from FP8 (bias=7) to FP16 (bias=15): E_fp16 = E_fp8 + 8 + uint8_t fp16_exponent = exponent + 8; + + // Construct FP16: extend mantissa to 10 bits + dst[i] = (sign << 8) | // Sign bit (bit 15) + (fp16_exponent << 10) | // 5-bit exponent (bits 14-10) + (mantissa << 7); // 3-bit mantissa extended to 10 bits (bits 9-7) + } +} + +/** + * Converts an array of FP8 (E4M3FN) values to FP16 format (supports subnormal + * numbers) + * + * @param src Input array of FP8 values (E4M3FN: Finite Normal + + * Subnormal) + * @param dst Output array for FP16 values (must be pre-allocated) + * @param elem_count Number of elements to convert + * + * Based on: OCP Microscaling Formats (MX) Specification v1.0 + * - FP8 E4M3 format: 1 sign bit (S) + 4 exponent bits (E, bias=7) + 3 mantissa + * bits (M) + * - Subnormal numbers (E=0): v = (-1)^S × 2^(1-bias) × (M/8) (Section 5.3.1) + * - NaN (E=0xF): No infinities defined; E=0xF represents NaNs (Table 2) + * - FP16 format: 1 sign bit + 5 exponent bits (bias=15) + 10 mantissa bits + * (IEEE 754 compatible) + */ +void __FP8E4M3FN_FP16(uint8_t *src, uint16_t *dst, uint32_t elem_count) { + src = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)src); + dst = (uint16_t *)get_spm_memory_mapping_wrapper((uint64_t)dst); + + for (uint32_t i = 0; i < elem_count; i++) { + uint8_t fp8 = src[i]; + uint8_t sign = fp8 & 0x80; // Extract sign bit (10000000) + uint8_t exponent = (fp8 >> 3) & 0x0F; // Extract 4-bit exponent (00001111) + uint8_t mantissa = fp8 & 0x07; // Extract 3-bit mantissa (00000111) + + // Handle subnormal values (exponent = 0) + if (exponent == 0) { + // According to OCP spec: v = (-1)^S × 2^(1-7) × (0 + 2^(-3) × M) + // FP16 exponent = (1 - 7) + 15 = 9 (bias adjustment) + static const int denormsAndZeroLut[8] = {0x0000, 0x1800, 0x1C00, 0x1E00, + 0x2000, 0x2100, 0x2200, 0x2300}; + + dst[i] = (sign << 8) | denormsAndZeroLut[mantissa]; + continue; + } + + // Handle NaN case (exponent = 0x0F for E4M3) + if (fp8 == 0x7F || fp8 == 0xFF) { + // Set FP16 max exponent (0x1F) and extend mantissa + dst[i] = (sign << 8) | (0x1F << 10) | (mantissa << 7); + continue; + } + + // Normal case conversion + // Convert exponent from FP8 (bias=7) to FP16 (bias=15): E_fp16 = E_fp8 + 8 + uint8_t fp16_exponent = exponent + 8; + + // Construct FP16: extend mantissa to 10 bits + dst[i] = (sign << 8) | // Sign bit (bit 15) + (fp16_exponent << 10) | // 5-bit exponent (bits 14-10) + (mantissa << 7); // 3-bit mantissa extended to 10 bits (bits 9-7) + } +} + +/** + * Converts packed FP4 (E2M1) values to FP16 format + * + * @param src Input array of packed FP4 values (2 values per byte) + * @param dst Output array for FP16 values (must be pre-allocated) + * @param elem_count Number of FP4 elements (not bytes) + * + * FP4 (E2M1) format (per element): + * [S][EE][M] + * 1 sign bit (bit3), 2 exponent bits (bit2-1), 1 mantissa bit (bit0) + * + * Storage format: + * Each byte contains two FP4 values: + * [S1 E1 E0 M1] [S0 E1 E0 M0] (high nibble first, corrected symmetry) + * + * FP16 output format: + * [S][EEEE E][MMMM MMMMMM] + * 1 sign bit (bit15), 5 exponent bits (bit14-10), 10 mantissa bits (bit9-0) + */ +void __FP4E2M1_FP16(uint8_t *src, uint16_t *dst, uint32_t elem_count) { + src = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)src); + dst = (uint16_t *)get_spm_memory_mapping_wrapper((uint64_t)dst); + + // Constants matching Python logic + const uint16_t to_bias = 15; // FP16 exponent bias + const uint16_t to_m_bits = 10; // FP16 mantissa bits + const uint16_t to_point5 = 0x3800; // FP16 value for ±0.5 (subnormal mapping) + const uint16_t bias_offset = (to_bias - 1) + << to_m_bits; // (15-1) <<10 = 14<<10 + + for (uint32_t i = 0; i < (elem_count + 1) / 2; i++) { + uint8_t byte = src[i]; + // Process high nibble (first element, bits 7-4) + if (2 * i + 1 < elem_count) { + uint8_t elem = (byte >> 4) & 0x0F; // Extract 4-bit FP4 element + uint16_t sign = elem & 0x08; // Sign bit (bit3 of element) + uint16_t exp = (elem >> 1) & 0x03; // Exponent bits (bit2-1 of element) + uint16_t mant = elem & 0x01; // Mantissa bit (bit0 of element) + uint16_t result; + if (exp == 0) { + // Subnormal or zero (exp=0) + if (mant == 1) { + // Subnormal: map to ±0.5 (to_point5 + sign) + result = + (sign << 12) | to_point5; // Sign shifted to FP16 bit15 (3+12=15) + } else { + // Zero: only preserve sign + result = sign << 12; + } + } else { + // Normal number: base mantissa + bias offset + uint16_t abs = elem & 0x07; + uint16_t base = (abs << (to_m_bits - 1)) | + (sign << 12); // Mantissa shifted to FP16 bit9 (10-1=9) + result = base + bias_offset; + } + dst[2 * i + 1] = result; + } + + // Process low nibble (second element, bits 3-0) if needed + { + uint8_t elem = byte & 0x0F; // Extract 4-bit FP4 element + uint16_t sign = elem & 0x08; // Sign bit (bit3 of element) + uint16_t exp = (elem >> 1) & 0x03; // Exponent bits (bit2-1 of element) + uint16_t mant = elem & 0x01; // Mantissa bit (bit0 of element) + + uint16_t result; + if (exp == 0) { + if (mant == 1) { + result = (sign << 12) | to_point5; + } else { + result = sign << 12; + } + } else { + uint16_t abs = elem & 0x07; + uint16_t base = (abs << (to_m_bits - 1)) | + (sign << 12); // Mantissa shifted to FP16 bit9 (10-1=9) + result = base + bias_offset; + } + dst[2 * i] = result; + } + } +} diff --git a/third_party/wafer/crt/lib/Wafer/mxfp_scale_bf16.c b/third_party/wafer/crt/lib/Wafer/mxfp_scale_bf16.c new file mode 100755 index 00000000..7c2b7783 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/mxfp_scale_bf16.c @@ -0,0 +1,68 @@ +//===------------------------- mxfp_scale_bf16.c --------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::MXFPScaleBF16Op see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "instr_def.h" +#include "wafer.h" +#include + +/** + * Applies microscaling to BF16 data using E8M0 scale factors. + * + * Implements OCP Microscaling Format (MX) specification by scaling blocks + * of BF16 data with E8M0 scale factors. Special handling for NaN scales + * as per MX specification requirements. + * + * @param value Source BF16 data pointer (MXFP4-converted BF16 format) + * @param scale E8M0 scale factors pointer (1 per block) + * @param dst Destination BF16 data pointer + * @param elem_count Total elements in source (must be blocks*32) + */ +void __mxfpScaleBF16(uint16_t *value, uint8_t *scale, uint16_t *dst, + uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + const int scaling_block_size = 32; // MXFP4 block size per OCP spec + + // Obtain hardware-specific memory mapping for scale factors + scale = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)scale); + + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + // Main processing loop per scaling block + for (uint32_t block_idx = 0; block_idx < elem_count / scaling_block_size; + block_idx++) { + // Convert E8M0 scale to BF16 representation + uint16_t scale_bf16 = ((uint16_t)scale[block_idx]) << 7; + + // MX spec handling: E8M0 scale value 0xFF indicates NaN + if (scale[block_idx] == 0xFF) { + scale_bf16 = 0x7FC0; // BF16 quiet NaN encoding + } + + // Calculate block positions + uint16_t *block_src = value + block_idx * scaling_block_size; + uint16_t *block_dst = dst + block_idx * scaling_block_size; + + // Apply scaling via vector-scalar multiplication + cmd->MulVS(&inst, (uint64_t)block_src, (uint32_t)scale_bf16, + (uint64_t)block_dst, scaling_block_size, RND_NEAREST_EVEN, + Fmt_BF16); + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + } +} diff --git a/third_party/wafer/crt/lib/Wafer/mxfp_scale_fp16.c b/third_party/wafer/crt/lib/Wafer/mxfp_scale_fp16.c new file mode 100755 index 00000000..ed90307e --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/mxfp_scale_fp16.c @@ -0,0 +1,91 @@ +//===------------------------- mxfp_scale_fp16.c --------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::MXFPScaleFp16Op see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" +#include + +/** + * Applies microscaling to FP16 data using E8M0 scale factors. + * + * Implements OCP Microscaling Format (MX) specification by scaling blocks + * of FP16 data with E8M0 scale factors. Special handling for NaN scales + * as per MX specification requirements. + * + * @param value Source FP16 data pointer + * @param scale E8M0 scale factors pointer (1 per block) + * @param dst Destination FP16 data pointer + * @param elem_count Total elements in source (must be blocks*32) + */ +void __mxfpScaleFP16(uint16_t *value, uint8_t *scale, uint16_t *dst, + uint32_t elem_count) { + const int scaling_block_size = 32; // MXFP4 block size per OCP spec + + // Obtain hardware-specific memory mapping for scale factors + scale = (uint8_t *)get_spm_memory_mapping_wrapper((uint64_t)scale); + + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + for (uint32_t block_idx = 0; block_idx < elem_count / scaling_block_size; + block_idx++) { + uint16_t scale_fp16; + uint8_t scale_val = scale[block_idx]; + + // referrence to test_dot_scaled.py: + // ``` + // scale_fp32 = (scale.to(tl.uint32) << 23).to(tl.float32, bitcast=True) + // upcasted_scale = scale_fp32.to(tl.float16) + // ``` + // Pure bitwise E8M0 to FP16 conversion (matching Python impl) + if (scale_val == 0xFF) { + // MX spec: 0xFF = NaN; FP16 NaN (all exp 1s, non-zero mantissa) + scale_fp16 = 0x7FFF; + } else { + // Step 1: Extract float32 exponent (mimic Python's uint32 << 23) + // E8M0 8-bit value as float32 exponent (float32 bias 127) + uint32_t f32_exp = scale_val; + + // Step 2: Convert to FP16 exponent (FP16 bias 15) + // FP16 exp = (f32 actual exp) + 15 = (f32_exp - 127) + 15 = f32_exp - 112 + int16_t fp16_exp = (int16_t)f32_exp - 112; + + // Step 3: Handle exponent range (FP16 exp: 5-bit, 0~31) + if (fp16_exp > 31) { + // Overflow: positive infinity (all exp 1s, mantissa 0s) + scale_fp16 = 0x7C00; + } else if (fp16_exp < 0) { + // Underflow: zero + scale_fp16 = 0x0000; + } else { + // Normal range: construct FP16 (sign 0, 5-bit exp, 10-bit mantissa 0) + scale_fp16 = (uint16_t)(fp16_exp << 10); // Mantissa all 0s + } + } + + // Calculate block positions + uint16_t *block_src = value + block_idx * scaling_block_size; + uint16_t *block_dst = dst + block_idx * scaling_block_size; + + // Apply scaling via vector-scalar multiplication with FP16 format + cmd->MulVS(&inst, (uint64_t)block_src, (uint32_t)scale_fp16, + (uint64_t)block_dst, scaling_block_size, RND_NEAREST_EVEN, + Fmt_FP16); + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + } +} diff --git a/third_party/wafer/crt/lib/Wafer/nchw2nhwc.c b/third_party/wafer/crt/lib/Wafer/nchw2nhwc.c new file mode 100755 index 00000000..d94f3fc0 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/nchw2nhwc.c @@ -0,0 +1,37 @@ +//===------------------------ nchw2nhwc.c ---------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Nchw2nhwc see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Nchw2nhwc(uint64_t *src, uint64_t *dst, int32_t *src_shape, + int32_t *dst_shape, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_shape[0], src_shape[1], src_shape[2], src_shape[3]}; + Data_Shape shape2 = {dst_shape[0], dst_shape[1], dst_shape[2], dst_shape[3]}; + cmd->Nchw2nhwc(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/neg.c b/third_party/wafer/crt/lib/Wafer/neg.c new file mode 100755 index 00000000..b7e6eb85 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/neg.c @@ -0,0 +1,32 @@ +//===------------------------- neg.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::negVVOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __NegVV(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->NegVV(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + + TsmWaitfinish(); + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/nhwc2nchw.c b/third_party/wafer/crt/lib/Wafer/nhwc2nchw.c new file mode 100755 index 00000000..9229c835 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/nhwc2nchw.c @@ -0,0 +1,37 @@ +//===------------------------ nhwc2nchw.c ---------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Nhwc2nchw see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Nhwc2nchw(uint64_t *src, uint64_t *dst, int32_t *src_shape, + int32_t *dst_shape, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_shape[0], src_shape[1], src_shape[2], src_shape[3]}; + Data_Shape shape2 = {dst_shape[0], dst_shape[1], dst_shape[2], dst_shape[3]}; + cmd->Nhwc2nchw(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/noc_init.c b/third_party/wafer/crt/lib/Wafer/noc_init.c new file mode 100644 index 00000000..8ce19cad --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/noc_init.c @@ -0,0 +1,16 @@ +// Initialize the protocol-owned words before a 16-tile NoC launch. +#include +#include "tx81_spm.h" +#include "wafer.h" + +void __NoCRingInit(void) { + volatile uint32_t *state = + (volatile uint32_t *)get_spm_memory_mapping(SINGLE_SPM_SYNC_ADDR); + state[0] = 0; // request + state[1] = 0; // acknowledgement +#ifdef __riscv + __asm__ __volatile__("fence iorw, iorw" ::: "memory"); +#else + __sync_synchronize(); +#endif +} diff --git a/third_party/wafer/crt/lib/Wafer/op_gelu.c b/third_party/wafer/crt/lib/Wafer/op_gelu.c new file mode 100755 index 00000000..42e6498a --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/op_gelu.c @@ -0,0 +1,524 @@ +#include "op_gelu.h" + +const struct erff_data { + struct { + float erf, scale; + } tab[513]; +} __erff_data; + +const struct erff_data __erff_data = { + .tab = {{0x0.000000p+0, 0x1.20dd76p+0}, {0x1.20dbf4p-7, 0x1.20d8f2p+0}, + {0x1.20d770p-6, 0x1.20cb68p+0}, {0x1.b137e0p-6, 0x1.20b4d8p+0}, + {0x1.20c564p-5, 0x1.209546p+0}, {0x1.68e5d4p-5, 0x1.206cb4p+0}, + {0x1.b0fafep-5, 0x1.203b26p+0}, {0x1.f902a8p-5, 0x1.2000a0p+0}, + {0x1.207d48p-4, 0x1.1fbd28p+0}, {0x1.44703ep-4, 0x1.1f70c4p+0}, + {0x1.68591ap-4, 0x1.1f1b7ap+0}, {0x1.8c36bep-4, 0x1.1ebd56p+0}, + {0x1.b00812p-4, 0x1.1e565cp+0}, {0x1.d3cbf8p-4, 0x1.1de698p+0}, + {0x1.f7815ap-4, 0x1.1d6e14p+0}, {0x1.0d9390p-3, 0x1.1cecdcp+0}, + {0x1.1f5e1ap-3, 0x1.1c62fap+0}, {0x1.311fc2p-3, 0x1.1bd07cp+0}, + {0x1.42d7fcp-3, 0x1.1b3572p+0}, {0x1.548642p-3, 0x1.1a91e6p+0}, + {0x1.662a0cp-3, 0x1.19e5eap+0}, {0x1.77c2d2p-3, 0x1.19318cp+0}, + {0x1.895010p-3, 0x1.1874dep+0}, {0x1.9ad142p-3, 0x1.17aff0p+0}, + {0x1.ac45e4p-3, 0x1.16e2d8p+0}, {0x1.bdad72p-3, 0x1.160da4p+0}, + {0x1.cf076ep-3, 0x1.153068p+0}, {0x1.e05354p-3, 0x1.144b3cp+0}, + {0x1.f190aap-3, 0x1.135e30p+0}, {0x1.015f78p-2, 0x1.12695ep+0}, + {0x1.09eed6p-2, 0x1.116cd8p+0}, {0x1.127632p-2, 0x1.1068bap+0}, + {0x1.1af54ep-2, 0x1.0f5d16p+0}, {0x1.236bf0p-2, 0x1.0e4a08p+0}, + {0x1.2bd9dcp-2, 0x1.0d2fa6p+0}, {0x1.343ed6p-2, 0x1.0c0e0ap+0}, + {0x1.3c9aa8p-2, 0x1.0ae550p+0}, {0x1.44ed18p-2, 0x1.09b590p+0}, + {0x1.4d35f0p-2, 0x1.087ee4p+0}, {0x1.5574f4p-2, 0x1.07416cp+0}, + {0x1.5da9f4p-2, 0x1.05fd3ep+0}, {0x1.65d4b8p-2, 0x1.04b27cp+0}, + {0x1.6df50ap-2, 0x1.036140p+0}, {0x1.760abap-2, 0x1.0209a6p+0}, + {0x1.7e1594p-2, 0x1.00abd0p+0}, {0x1.861566p-2, 0x1.fe8fb0p-1}, + {0x1.8e0a02p-2, 0x1.fbbbbep-1}, {0x1.95f336p-2, 0x1.f8dc0ap-1}, + {0x1.9dd0d2p-2, 0x1.f5f0cep-1}, {0x1.a5a2acp-2, 0x1.f2fa4cp-1}, + {0x1.ad6896p-2, 0x1.eff8c4p-1}, {0x1.b52264p-2, 0x1.ecec78p-1}, + {0x1.bccfecp-2, 0x1.e9d5a8p-1}, {0x1.c47104p-2, 0x1.e6b498p-1}, + {0x1.cc0584p-2, 0x1.e38988p-1}, {0x1.d38d44p-2, 0x1.e054bep-1}, + {0x1.db081cp-2, 0x1.dd167cp-1}, {0x1.e275eap-2, 0x1.d9cf06p-1}, + {0x1.e9d68ap-2, 0x1.d67ea2p-1}, {0x1.f129d4p-2, 0x1.d32592p-1}, + {0x1.f86faap-2, 0x1.cfc41ep-1}, {0x1.ffa7eap-2, 0x1.cc5a8ap-1}, + {0x1.03693ap-1, 0x1.c8e91cp-1}, {0x1.06f794p-1, 0x1.c5701ap-1}, + {0x1.0a7ef6p-1, 0x1.c1efcap-1}, {0x1.0dff50p-1, 0x1.be6872p-1}, + {0x1.117894p-1, 0x1.bada5ap-1}, {0x1.14eab4p-1, 0x1.b745c6p-1}, + {0x1.1855a6p-1, 0x1.b3aafcp-1}, {0x1.1bb95cp-1, 0x1.b00a46p-1}, + {0x1.1f15ccp-1, 0x1.ac63e8p-1}, {0x1.226ae8p-1, 0x1.a8b828p-1}, + {0x1.25b8a8p-1, 0x1.a5074ep-1}, {0x1.28ff02p-1, 0x1.a1519ep-1}, + {0x1.2c3decp-1, 0x1.9d9762p-1}, {0x1.2f755cp-1, 0x1.99d8dap-1}, + {0x1.32a54cp-1, 0x1.961650p-1}, {0x1.35cdb4p-1, 0x1.925008p-1}, + {0x1.38ee8ap-1, 0x1.8e8646p-1}, {0x1.3c07cap-1, 0x1.8ab950p-1}, + {0x1.3f196ep-1, 0x1.86e96ap-1}, {0x1.42236ep-1, 0x1.8316d6p-1}, + {0x1.4525c8p-1, 0x1.7f41dcp-1}, {0x1.482074p-1, 0x1.7b6abcp-1}, + {0x1.4b1372p-1, 0x1.7791b8p-1}, {0x1.4dfebap-1, 0x1.73b714p-1}, + {0x1.50e24cp-1, 0x1.6fdb12p-1}, {0x1.53be26p-1, 0x1.6bfdf0p-1}, + {0x1.569244p-1, 0x1.681ff2p-1}, {0x1.595ea6p-1, 0x1.644156p-1}, + {0x1.5c2348p-1, 0x1.60625cp-1}, {0x1.5ee02ep-1, 0x1.5c8342p-1}, + {0x1.619556p-1, 0x1.58a446p-1}, {0x1.6442c0p-1, 0x1.54c5a6p-1}, + {0x1.66e86ep-1, 0x1.50e79ep-1}, {0x1.69865ep-1, 0x1.4d0a68p-1}, + {0x1.6c1c98p-1, 0x1.492e42p-1}, {0x1.6eab18p-1, 0x1.455366p-1}, + {0x1.7131e6p-1, 0x1.417a0cp-1}, {0x1.73b102p-1, 0x1.3da26ep-1}, + {0x1.762870p-1, 0x1.39ccc2p-1}, {0x1.789836p-1, 0x1.35f940p-1}, + {0x1.7b0058p-1, 0x1.32281ep-1}, {0x1.7d60d8p-1, 0x1.2e5992p-1}, + {0x1.7fb9c0p-1, 0x1.2a8dcep-1}, {0x1.820b12p-1, 0x1.26c508p-1}, + {0x1.8454d6p-1, 0x1.22ff72p-1}, {0x1.869712p-1, 0x1.1f3d3cp-1}, + {0x1.88d1cep-1, 0x1.1b7e98p-1}, {0x1.8b050ep-1, 0x1.17c3b6p-1}, + {0x1.8d30dep-1, 0x1.140cc4p-1}, {0x1.8f5544p-1, 0x1.1059eep-1}, + {0x1.91724ap-1, 0x1.0cab62p-1}, {0x1.9387f6p-1, 0x1.09014cp-1}, + {0x1.959652p-1, 0x1.055bd6p-1}, {0x1.979d68p-1, 0x1.01bb2cp-1}, + {0x1.999d42p-1, 0x1.fc3ee6p-2}, {0x1.9b95e8p-1, 0x1.f511aap-2}, + {0x1.9d8768p-1, 0x1.edeeeep-2}, {0x1.9f71cap-1, 0x1.e6d700p-2}, + {0x1.a1551ap-1, 0x1.dfca26p-2}, {0x1.a33162p-1, 0x1.d8c8aap-2}, + {0x1.a506b0p-1, 0x1.d1d2d0p-2}, {0x1.a6d50cp-1, 0x1.cae8dap-2}, + {0x1.a89c86p-1, 0x1.c40b08p-2}, {0x1.aa5d26p-1, 0x1.bd3998p-2}, + {0x1.ac16fcp-1, 0x1.b674c8p-2}, {0x1.adca14p-1, 0x1.afbcd4p-2}, + {0x1.af767ap-1, 0x1.a911f0p-2}, {0x1.b11c3cp-1, 0x1.a27456p-2}, + {0x1.b2bb68p-1, 0x1.9be438p-2}, {0x1.b4540ap-1, 0x1.9561c8p-2}, + {0x1.b5e630p-1, 0x1.8eed36p-2}, {0x1.b771e8p-1, 0x1.8886b2p-2}, + {0x1.b8f742p-1, 0x1.822e66p-2}, {0x1.ba764ap-1, 0x1.7be47ap-2}, + {0x1.bbef10p-1, 0x1.75a91ap-2}, {0x1.bd61a2p-1, 0x1.6f7c6ap-2}, + {0x1.bece0ep-1, 0x1.695e8cp-2}, {0x1.c03464p-1, 0x1.634fa6p-2}, + {0x1.c194b2p-1, 0x1.5d4fd4p-2}, {0x1.c2ef08p-1, 0x1.575f34p-2}, + {0x1.c44376p-1, 0x1.517de6p-2}, {0x1.c5920ap-1, 0x1.4bac00p-2}, + {0x1.c6dad2p-1, 0x1.45e99cp-2}, {0x1.c81de2p-1, 0x1.4036d0p-2}, + {0x1.c95b46p-1, 0x1.3a93b2p-2}, {0x1.ca930ep-1, 0x1.350052p-2}, + {0x1.cbc54cp-1, 0x1.2f7cc4p-2}, {0x1.ccf20cp-1, 0x1.2a0916p-2}, + {0x1.ce1962p-1, 0x1.24a554p-2}, {0x1.cf3b5cp-1, 0x1.1f518ap-2}, + {0x1.d0580cp-1, 0x1.1a0dc6p-2}, {0x1.d16f7ep-1, 0x1.14da0ap-2}, + {0x1.d281c4p-1, 0x1.0fb662p-2}, {0x1.d38ef0p-1, 0x1.0aa2d0p-2}, + {0x1.d49710p-1, 0x1.059f5ap-2}, {0x1.d59a34p-1, 0x1.00ac00p-2}, + {0x1.d6986cp-1, 0x1.f79184p-3}, {0x1.d791cap-1, 0x1.edeb40p-3}, + {0x1.d8865ep-1, 0x1.e46530p-3}, {0x1.d97636p-1, 0x1.daff4ap-3}, + {0x1.da6162p-1, 0x1.d1b982p-3}, {0x1.db47f4p-1, 0x1.c893cep-3}, + {0x1.dc29fcp-1, 0x1.bf8e1cp-3}, {0x1.dd0788p-1, 0x1.b6a856p-3}, + {0x1.dde0aap-1, 0x1.ade26cp-3}, {0x1.deb570p-1, 0x1.a53c42p-3}, + {0x1.df85eap-1, 0x1.9cb5bep-3}, {0x1.e0522ap-1, 0x1.944ec2p-3}, + {0x1.e11a3ep-1, 0x1.8c0732p-3}, {0x1.e1de36p-1, 0x1.83deeap-3}, + {0x1.e29e22p-1, 0x1.7bd5c8p-3}, {0x1.e35a12p-1, 0x1.73eba4p-3}, + {0x1.e41214p-1, 0x1.6c2056p-3}, {0x1.e4c638p-1, 0x1.6473b6p-3}, + {0x1.e5768cp-1, 0x1.5ce596p-3}, {0x1.e62322p-1, 0x1.5575c8p-3}, + {0x1.e6cc08p-1, 0x1.4e241ep-3}, {0x1.e7714ap-1, 0x1.46f066p-3}, + {0x1.e812fcp-1, 0x1.3fda6cp-3}, {0x1.e8b12ap-1, 0x1.38e1fap-3}, + {0x1.e94be4p-1, 0x1.3206dcp-3}, {0x1.e9e336p-1, 0x1.2b48dap-3}, + {0x1.ea7730p-1, 0x1.24a7b8p-3}, {0x1.eb07e2p-1, 0x1.1e233ep-3}, + {0x1.eb9558p-1, 0x1.17bb2cp-3}, {0x1.ec1fa2p-1, 0x1.116f48p-3}, + {0x1.eca6ccp-1, 0x1.0b3f52p-3}, {0x1.ed2ae6p-1, 0x1.052b0cp-3}, + {0x1.edabfcp-1, 0x1.fe6460p-4}, {0x1.ee2a1ep-1, 0x1.f2a902p-4}, + {0x1.eea556p-1, 0x1.e72372p-4}, {0x1.ef1db4p-1, 0x1.dbd32ap-4}, + {0x1.ef9344p-1, 0x1.d0b7a0p-4}, {0x1.f00614p-1, 0x1.c5d04ap-4}, + {0x1.f07630p-1, 0x1.bb1c98p-4}, {0x1.f0e3a6p-1, 0x1.b09bfcp-4}, + {0x1.f14e82p-1, 0x1.a64de6p-4}, {0x1.f1b6d0p-1, 0x1.9c31c6p-4}, + {0x1.f21ca0p-1, 0x1.92470ap-4}, {0x1.f27ff8p-1, 0x1.888d1ep-4}, + {0x1.f2e0eap-1, 0x1.7f036cp-4}, {0x1.f33f7ep-1, 0x1.75a960p-4}, + {0x1.f39bc2p-1, 0x1.6c7e64p-4}, {0x1.f3f5c2p-1, 0x1.6381e2p-4}, + {0x1.f44d88p-1, 0x1.5ab342p-4}, {0x1.f4a31ep-1, 0x1.5211ecp-4}, + {0x1.f4f694p-1, 0x1.499d48p-4}, {0x1.f547f2p-1, 0x1.4154bcp-4}, + {0x1.f59742p-1, 0x1.3937b2p-4}, {0x1.f5e490p-1, 0x1.31458ep-4}, + {0x1.f62fe8p-1, 0x1.297dbap-4}, {0x1.f67952p-1, 0x1.21df9ap-4}, + {0x1.f6c0dcp-1, 0x1.1a6a96p-4}, {0x1.f7068cp-1, 0x1.131e14p-4}, + {0x1.f74a6ep-1, 0x1.0bf97ep-4}, {0x1.f78c8cp-1, 0x1.04fc3ap-4}, + {0x1.f7cceep-1, 0x1.fc4b5ep-5}, {0x1.f80ba2p-1, 0x1.eeea8cp-5}, + {0x1.f848acp-1, 0x1.e1d4d0p-5}, {0x1.f8841ap-1, 0x1.d508fap-5}, + {0x1.f8bdf2p-1, 0x1.c885e0p-5}, {0x1.f8f63ep-1, 0x1.bc4a54p-5}, + {0x1.f92d08p-1, 0x1.b05530p-5}, {0x1.f96256p-1, 0x1.a4a54ap-5}, + {0x1.f99634p-1, 0x1.99397ap-5}, {0x1.f9c8a8p-1, 0x1.8e109cp-5}, + {0x1.f9f9bap-1, 0x1.83298ep-5}, {0x1.fa2974p-1, 0x1.78832cp-5}, + {0x1.fa57dep-1, 0x1.6e1c58p-5}, {0x1.fa84fep-1, 0x1.63f3f6p-5}, + {0x1.fab0dep-1, 0x1.5a08e8p-5}, {0x1.fadb84p-1, 0x1.505a18p-5}, + {0x1.fb04f6p-1, 0x1.46e66cp-5}, {0x1.fb2d40p-1, 0x1.3dacd2p-5}, + {0x1.fb5464p-1, 0x1.34ac36p-5}, {0x1.fb7a6cp-1, 0x1.2be38cp-5}, + {0x1.fb9f60p-1, 0x1.2351c2p-5}, {0x1.fbc344p-1, 0x1.1af5d2p-5}, + {0x1.fbe61ep-1, 0x1.12ceb4p-5}, {0x1.fc07fap-1, 0x1.0adb60p-5}, + {0x1.fc28d8p-1, 0x1.031ad6p-5}, {0x1.fc48c2p-1, 0x1.f7182ap-6}, + {0x1.fc67bcp-1, 0x1.e85c44p-6}, {0x1.fc85d0p-1, 0x1.da0006p-6}, + {0x1.fca2fep-1, 0x1.cc0180p-6}, {0x1.fcbf52p-1, 0x1.be5ecep-6}, + {0x1.fcdaccp-1, 0x1.b1160ap-6}, {0x1.fcf576p-1, 0x1.a4255ap-6}, + {0x1.fd0f54p-1, 0x1.978ae8p-6}, {0x1.fd286ap-1, 0x1.8b44e6p-6}, + {0x1.fd40bep-1, 0x1.7f5188p-6}, {0x1.fd5856p-1, 0x1.73af0cp-6}, + {0x1.fd6f34p-1, 0x1.685bb6p-6}, {0x1.fd8562p-1, 0x1.5d55ccp-6}, + {0x1.fd9ae2p-1, 0x1.529b9ep-6}, {0x1.fdafb8p-1, 0x1.482b84p-6}, + {0x1.fdc3e8p-1, 0x1.3e03d8p-6}, {0x1.fdd77ap-1, 0x1.3422fep-6}, + {0x1.fdea6ep-1, 0x1.2a875cp-6}, {0x1.fdfcccp-1, 0x1.212f62p-6}, + {0x1.fe0e96p-1, 0x1.181984p-6}, {0x1.fe1fd0p-1, 0x1.0f443ep-6}, + {0x1.fe3080p-1, 0x1.06ae14p-6}, {0x1.fe40a6p-1, 0x1.fcab14p-7}, + {0x1.fe504cp-1, 0x1.ec7262p-7}, {0x1.fe5f70p-1, 0x1.dcaf36p-7}, + {0x1.fe6e18p-1, 0x1.cd5ecap-7}, {0x1.fe7c46p-1, 0x1.be7e5ap-7}, + {0x1.fe8a00p-1, 0x1.b00b38p-7}, {0x1.fe9748p-1, 0x1.a202bep-7}, + {0x1.fea422p-1, 0x1.94624ep-7}, {0x1.feb090p-1, 0x1.87275ep-7}, + {0x1.febc96p-1, 0x1.7a4f6ap-7}, {0x1.fec836p-1, 0x1.6dd7fep-7}, + {0x1.fed374p-1, 0x1.61beaep-7}, {0x1.fede52p-1, 0x1.56011cp-7}, + {0x1.fee8d4p-1, 0x1.4a9cf6p-7}, {0x1.fef2fep-1, 0x1.3f8ff6p-7}, + {0x1.fefccep-1, 0x1.34d7dcp-7}, {0x1.ff064cp-1, 0x1.2a727ap-7}, + {0x1.ff0f76p-1, 0x1.205dacp-7}, {0x1.ff1852p-1, 0x1.169756p-7}, + {0x1.ff20e0p-1, 0x1.0d1d6ap-7}, {0x1.ff2924p-1, 0x1.03ede2p-7}, + {0x1.ff3120p-1, 0x1.f60d8ap-8}, {0x1.ff38d6p-1, 0x1.e4cc4ap-8}, + {0x1.ff4048p-1, 0x1.d4143ap-8}, {0x1.ff4778p-1, 0x1.c3e1a6p-8}, + {0x1.ff4e68p-1, 0x1.b430ecp-8}, {0x1.ff551ap-1, 0x1.a4fe84p-8}, + {0x1.ff5b90p-1, 0x1.9646f4p-8}, {0x1.ff61ccp-1, 0x1.8806d8p-8}, + {0x1.ff67d0p-1, 0x1.7a3adep-8}, {0x1.ff6d9ep-1, 0x1.6cdfccp-8}, + {0x1.ff7338p-1, 0x1.5ff276p-8}, {0x1.ff789ep-1, 0x1.536fc2p-8}, + {0x1.ff7dd4p-1, 0x1.4754acp-8}, {0x1.ff82dap-1, 0x1.3b9e40p-8}, + {0x1.ff87b2p-1, 0x1.30499cp-8}, {0x1.ff8c5cp-1, 0x1.2553eep-8}, + {0x1.ff90dcp-1, 0x1.1aba78p-8}, {0x1.ff9532p-1, 0x1.107a8cp-8}, + {0x1.ff9960p-1, 0x1.06918cp-8}, {0x1.ff9d68p-1, 0x1.f9f9d0p-9}, + {0x1.ffa14ap-1, 0x1.e77448p-9}, {0x1.ffa506p-1, 0x1.d58da6p-9}, + {0x1.ffa8a0p-1, 0x1.c4412cp-9}, {0x1.ffac18p-1, 0x1.b38a3ap-9}, + {0x1.ffaf6ep-1, 0x1.a36454p-9}, {0x1.ffb2a6p-1, 0x1.93cb12p-9}, + {0x1.ffb5bep-1, 0x1.84ba30p-9}, {0x1.ffb8b8p-1, 0x1.762d84p-9}, + {0x1.ffbb98p-1, 0x1.682100p-9}, {0x1.ffbe5ap-1, 0x1.5a90b0p-9}, + {0x1.ffc102p-1, 0x1.4d78bcp-9}, {0x1.ffc390p-1, 0x1.40d564p-9}, + {0x1.ffc606p-1, 0x1.34a306p-9}, {0x1.ffc862p-1, 0x1.28de12p-9}, + {0x1.ffcaa8p-1, 0x1.1d8318p-9}, {0x1.ffccd8p-1, 0x1.128ebap-9}, + {0x1.ffcef4p-1, 0x1.07fdb4p-9}, {0x1.ffd0fap-1, 0x1.fb99b8p-10}, + {0x1.ffd2eap-1, 0x1.e7f232p-10}, {0x1.ffd4cap-1, 0x1.d4fed8p-10}, + {0x1.ffd696p-1, 0x1.c2b9d0p-10}, {0x1.ffd84ep-1, 0x1.b11d70p-10}, + {0x1.ffd9f8p-1, 0x1.a02436p-10}, {0x1.ffdb90p-1, 0x1.8fc8c8p-10}, + {0x1.ffdd18p-1, 0x1.8005f0p-10}, {0x1.ffde90p-1, 0x1.70d6a4p-10}, + {0x1.ffdffap-1, 0x1.6235fcp-10}, {0x1.ffe154p-1, 0x1.541f34p-10}, + {0x1.ffe2a2p-1, 0x1.468daep-10}, {0x1.ffe3e2p-1, 0x1.397ceep-10}, + {0x1.ffe514p-1, 0x1.2ce898p-10}, {0x1.ffe63cp-1, 0x1.20cc76p-10}, + {0x1.ffe756p-1, 0x1.15246ep-10}, {0x1.ffe866p-1, 0x1.09ec86p-10}, + {0x1.ffe96ap-1, 0x1.fe41cep-11}, {0x1.ffea64p-1, 0x1.e97ba4p-11}, + {0x1.ffeb54p-1, 0x1.d57f52p-11}, {0x1.ffec3ap-1, 0x1.c245d4p-11}, + {0x1.ffed16p-1, 0x1.afc85ep-11}, {0x1.ffedeap-1, 0x1.9e0058p-11}, + {0x1.ffeeb4p-1, 0x1.8ce75ep-11}, {0x1.ffef76p-1, 0x1.7c7744p-11}, + {0x1.fff032p-1, 0x1.6caa0ep-11}, {0x1.fff0e4p-1, 0x1.5d79ecp-11}, + {0x1.fff18ep-1, 0x1.4ee142p-11}, {0x1.fff232p-1, 0x1.40daa4p-11}, + {0x1.fff2d0p-1, 0x1.3360ccp-11}, {0x1.fff366p-1, 0x1.266ea8p-11}, + {0x1.fff3f6p-1, 0x1.19ff46p-11}, {0x1.fff480p-1, 0x1.0e0de8p-11}, + {0x1.fff504p-1, 0x1.0295f0p-11}, {0x1.fff582p-1, 0x1.ef25d4p-12}, + {0x1.fff5fcp-1, 0x1.da0110p-12}, {0x1.fff670p-1, 0x1.c5b542p-12}, + {0x1.fff6dep-1, 0x1.b23a5ap-12}, {0x1.fff74ap-1, 0x1.9f8894p-12}, + {0x1.fff7aep-1, 0x1.8d986ap-12}, {0x1.fff810p-1, 0x1.7c629ap-12}, + {0x1.fff86cp-1, 0x1.6be022p-12}, {0x1.fff8c6p-1, 0x1.5c0a38p-12}, + {0x1.fff91cp-1, 0x1.4cda54p-12}, {0x1.fff96cp-1, 0x1.3e4a24p-12}, + {0x1.fff9bap-1, 0x1.305390p-12}, {0x1.fffa04p-1, 0x1.22f0b4p-12}, + {0x1.fffa4cp-1, 0x1.161be4p-12}, {0x1.fffa90p-1, 0x1.09cfa4p-12}, + {0x1.fffad0p-1, 0x1.fc0d56p-13}, {0x1.fffb0ep-1, 0x1.e577bcp-13}, + {0x1.fffb4ap-1, 0x1.cfd4a6p-13}, {0x1.fffb82p-1, 0x1.bb1a96p-13}, + {0x1.fffbb8p-1, 0x1.a74068p-13}, {0x1.fffbecp-1, 0x1.943d4ap-13}, + {0x1.fffc1ep-1, 0x1.8208bcp-13}, {0x1.fffc4ep-1, 0x1.709a8ep-13}, + {0x1.fffc7ap-1, 0x1.5feadap-13}, {0x1.fffca6p-1, 0x1.4ff208p-13}, + {0x1.fffccep-1, 0x1.40a8c2p-13}, {0x1.fffcf6p-1, 0x1.3207fcp-13}, + {0x1.fffd1ap-1, 0x1.2408eap-13}, {0x1.fffd3ep-1, 0x1.16a502p-13}, + {0x1.fffd60p-1, 0x1.09d5f8p-13}, {0x1.fffd80p-1, 0x1.fb2b7ap-14}, + {0x1.fffda0p-1, 0x1.e3bcf4p-14}, {0x1.fffdbep-1, 0x1.cd5528p-14}, + {0x1.fffddap-1, 0x1.b7e946p-14}, {0x1.fffdf4p-1, 0x1.a36eecp-14}, + {0x1.fffe0ep-1, 0x1.8fdc1cp-14}, {0x1.fffe26p-1, 0x1.7d2738p-14}, + {0x1.fffe3ep-1, 0x1.6b4702p-14}, {0x1.fffe54p-1, 0x1.5a329cp-14}, + {0x1.fffe68p-1, 0x1.49e178p-14}, {0x1.fffe7ep-1, 0x1.3a4b60p-14}, + {0x1.fffe90p-1, 0x1.2b6876p-14}, {0x1.fffea2p-1, 0x1.1d3120p-14}, + {0x1.fffeb4p-1, 0x1.0f9e1cp-14}, {0x1.fffec4p-1, 0x1.02a868p-14}, + {0x1.fffed4p-1, 0x1.ec929ap-15}, {0x1.fffee4p-1, 0x1.d4f4b4p-15}, + {0x1.fffef2p-1, 0x1.be6abcp-15}, {0x1.ffff00p-1, 0x1.a8e8ccp-15}, + {0x1.ffff0cp-1, 0x1.94637ep-15}, {0x1.ffff18p-1, 0x1.80cfdcp-15}, + {0x1.ffff24p-1, 0x1.6e2368p-15}, {0x1.ffff30p-1, 0x1.5c540cp-15}, + {0x1.ffff3ap-1, 0x1.4b581cp-15}, {0x1.ffff44p-1, 0x1.3b2652p-15}, + {0x1.ffff4ep-1, 0x1.2bb5ccp-15}, {0x1.ffff56p-1, 0x1.1cfe02p-15}, + {0x1.ffff60p-1, 0x1.0ef6c4p-15}, {0x1.ffff68p-1, 0x1.019842p-15}, + {0x1.ffff70p-1, 0x1.e9b5e8p-16}, {0x1.ffff78p-1, 0x1.d16f58p-16}, + {0x1.ffff7ep-1, 0x1.ba4f04p-16}, {0x1.ffff84p-1, 0x1.a447b8p-16}, + {0x1.ffff8cp-1, 0x1.8f4cccp-16}, {0x1.ffff92p-1, 0x1.7b5224p-16}, + {0x1.ffff98p-1, 0x1.684c22p-16}, {0x1.ffff9cp-1, 0x1.562facp-16}, + {0x1.ffffa2p-1, 0x1.44f21ep-16}, {0x1.ffffa6p-1, 0x1.34894ap-16}, + {0x1.ffffacp-1, 0x1.24eb72p-16}, {0x1.ffffb0p-1, 0x1.160f44p-16}, + {0x1.ffffb4p-1, 0x1.07ebd2p-16}, {0x1.ffffb8p-1, 0x1.f4f12ep-17}, + {0x1.ffffbcp-1, 0x1.db5ad0p-17}, {0x1.ffffc0p-1, 0x1.c304f0p-17}, + {0x1.ffffc4p-1, 0x1.abe09ep-17}, {0x1.ffffc6p-1, 0x1.95df98p-17}, + {0x1.ffffcap-1, 0x1.80f43ap-17}, {0x1.ffffccp-1, 0x1.6d1178p-17}, + {0x1.ffffd0p-1, 0x1.5a2ae0p-17}, {0x1.ffffd2p-1, 0x1.483488p-17}, + {0x1.ffffd4p-1, 0x1.372310p-17}, {0x1.ffffd6p-1, 0x1.26eb9ep-17}, + {0x1.ffffd8p-1, 0x1.1783cep-17}, {0x1.ffffdcp-1, 0x1.08e1bap-17}, + {0x1.ffffdep-1, 0x1.f5f7d8p-18}, {0x1.ffffdep-1, 0x1.db92b6p-18}, + {0x1.ffffe0p-1, 0x1.c282cep-18}, {0x1.ffffe2p-1, 0x1.aab7acp-18}, + {0x1.ffffe4p-1, 0x1.94219cp-18}, {0x1.ffffe6p-1, 0x1.7eb1a2p-18}, + {0x1.ffffe8p-1, 0x1.6a5972p-18}, {0x1.ffffe8p-1, 0x1.570b6ap-18}, + {0x1.ffffeap-1, 0x1.44ba86p-18}, {0x1.ffffeap-1, 0x1.335a62p-18}, + {0x1.ffffecp-1, 0x1.22df2ap-18}, {0x1.ffffeep-1, 0x1.133d96p-18}, + {0x1.ffffeep-1, 0x1.046aeap-18}, {0x1.fffff0p-1, 0x1.ecb9d0p-19}, + {0x1.fffff0p-1, 0x1.d21398p-19}, {0x1.fffff2p-1, 0x1.b8d094p-19}, + {0x1.fffff2p-1, 0x1.a0df10p-19}, {0x1.fffff2p-1, 0x1.8a2e26p-19}, + {0x1.fffff4p-1, 0x1.74adc8p-19}, {0x1.fffff4p-1, 0x1.604ea8p-19}, + {0x1.fffff4p-1, 0x1.4d0232p-19}, {0x1.fffff6p-1, 0x1.3aba86p-19}, + {0x1.fffff6p-1, 0x1.296a70p-19}, {0x1.fffff6p-1, 0x1.190562p-19}, + {0x1.fffff8p-1, 0x1.097f62p-19}, {0x1.fffff8p-1, 0x1.f59a20p-20}, + {0x1.fffff8p-1, 0x1.d9c736p-20}, {0x1.fffff8p-1, 0x1.bf716cp-20}, + {0x1.fffffap-1, 0x1.a6852cp-20}, {0x1.fffffap-1, 0x1.8eefd8p-20}, + {0x1.fffffap-1, 0x1.789fb8p-20}, {0x1.fffffap-1, 0x1.6383f8p-20}, + {0x1.fffffap-1, 0x1.4f8c96p-20}, {0x1.fffffap-1, 0x1.3caa62p-20}, + {0x1.fffffcp-1, 0x1.2acee2p-20}, {0x1.fffffcp-1, 0x1.19ec60p-20}, + {0x1.fffffcp-1, 0x1.09f5d0p-20}, {0x1.fffffcp-1, 0x1.f5bd96p-21}, + {0x1.fffffcp-1, 0x1.d9371ep-21}, {0x1.fffffcp-1, 0x1.be41dep-21}, + {0x1.fffffcp-1, 0x1.a4c89ep-21}, {0x1.fffffcp-1, 0x1.8cb738p-21}, + {0x1.fffffep-1, 0x1.75fa8ep-21}, {0x1.fffffep-1, 0x1.608078p-21}, + {0x1.fffffep-1, 0x1.4c37c0p-21}, {0x1.fffffep-1, 0x1.39100ep-21}, + {0x1.fffffep-1, 0x1.26f9e0p-21}, {0x1.fffffep-1, 0x1.15e682p-21}, + {0x1.fffffep-1, 0x1.05c804p-21}, {0x1.fffffep-1, 0x1.ed2254p-22}, + {0x1.fffffep-1, 0x1.d06ad6p-22}, {0x1.fffffep-1, 0x1.b551c8p-22}, + {0x1.fffffep-1, 0x1.9bc0a0p-22}, {0x1.fffffep-1, 0x1.83a200p-22}, + {0x1.fffffep-1, 0x1.6ce1aap-22}, {0x1.fffffep-1, 0x1.576c72p-22}, + {0x1.fffffep-1, 0x1.43302cp-22}, {0x1.fffffep-1, 0x1.301ba2p-22}, + {0x1.fffffep-1, 0x1.1e1e86p-22}, {0x1.fffffep-1, 0x1.0d2966p-22}, + {0x1.000000p+0, 0x1.fa5b50p-23}, {0x1.000000p+0, 0x1.dc3ae4p-23}, + {0x1.000000p+0, 0x1.bfd756p-23}, {0x1.000000p+0, 0x1.a517dap-23}, + {0x1.000000p+0, 0x1.8be4f8p-23}, {0x1.000000p+0, 0x1.74287ep-23}, + {0x1.000000p+0, 0x1.5dcd66p-23}, {0x1.000000p+0, 0x1.48bfd4p-23}, + {0x1.000000p+0, 0x1.34ecf8p-23}, {0x1.000000p+0, 0x1.224310p-23}, + {0x1.000000p+0, 0x1.10b148p-23}}, +}; + +hybrid_value get_ptr_value_by_idx_new(void *addr, size_t idx, + Data_Format dtype) { + hybrid_value value = {0}; + switch (dtype) { + case Fmt_FP16: + case Fmt_BF16: + value.i16 = ((int16_t *)addr)[idx]; + break; + case Fmt_FP32: + value.fp32 = ((float *)addr)[idx]; + break; + default: + break; + } + return value; +} + +void get_erf_value(void *in_addr, void *out_addr, uint64_t count, + Data_Format dtype) { + float TwoSqrtPI_f = 0.12837917f; + tmp_32suf Shift; + Shift.f = 65536.0f; + float OneThird_f = 0.33333334f; + float OneSqrtTwo_f = 0.7071067811865475f; + for (size_t idx = 0; idx < count; idx++) { + hybrid_value in_value = get_ptr_value_by_idx_new(in_addr, idx, dtype); + float in = set_value2float32(dtype, (int8_t *)&in_value); + tmp_32suf x; + x.f = in * OneSqrtTwo_f; + uint32_t ix = x.u; + uint32_t ia = ix & 0x7fffffff; + uint32_t sign = ix & ~0x7fffffff; + in = 0.5f * in; + if (ia < 0x20800000) { + hybrid_value value1 = + set_float2value(dtype, in * (1.0f + (TwoSqrtPI_f * x.f + x.f))); + set_ptr_value_by_idx(out_addr, value1, idx, dtype); + continue; + } + if (ia < 0x407b8000) { + tmp_32suf a, z, y; + a.u = ia; + z.f = a.f + Shift.f; + uint32_t i = z.u - Shift.u; + float r = z.f - Shift.f; + float erfr = __erff_data.tab[i].erf; + float scale = __erff_data.tab[i].scale; + float d = a.f - r; + float d2 = d * d; + y.f = -(OneThird_f * d + r); + y.f = (y.f * d2 + d) * scale + erfr; + y.u = y.u | sign; + hybrid_value value2 = set_float2value(dtype, in * (1.0f + y.f)); + set_ptr_value_by_idx(out_addr, value2, idx, dtype); + continue; + } + if (ia >= 0x7f800000) { + hybrid_value value3 = set_float2value( + dtype, in * (1.0f + ((1.0f - (float)(sign >> 30)) + 1.0f / x.f))); + set_ptr_value_by_idx(out_addr, value3, idx, dtype); + continue; + } + tmp_32suf tmp; + tmp.f = 1.0f; + tmp.u = sign | tmp.u; + hybrid_value value4 = set_float2value(dtype, in * (1.0f + tmp.f)); + set_ptr_value_by_idx(out_addr, value4, idx, dtype); + TsmWaitfinish(); + } +} + +void get_tanh_value(uint64_t *in, uint64_t *imm, uint64_t *out, + uint32_t elem_count, uint16_t fmt) { + union { + float f; + uint32_t i; + } converter; + converter.f = 0.044715f; + uint32_t a1 = converter.i; + converter.f = 0.79788456f; + uint32_t a2 = converter.i; + converter.f = 1.0f; + uint32_t a3 = converter.i; + converter.f = 0.5f; + uint32_t a4 = converter.i; + + uint16_t rnd_mode = 0; + + uint64_t imm_size = elem_count * sizeof(float); + uint64_t imm_a = (uint64_t)(imm); + uint64_t imm_b = (uint64_t)(imm) + imm_size; + uint64_t cycle_value = 0; + + uint64_t in_addr = (uint64_t)in; + uint64_t out_addr = (uint64_t)out; + + // TsmActivation *activation = + // (TsmActivation *)getTsmOpPointer()->activation_pointer; + // TsmActivationInstr activationParam = {I_CGRA, {0,}, {0,}}; + // TsmArith *arith = (TsmArith *)getTsmOpPointer()->arith_pointer; + // TsmArithInstr arithParam = {I_CGRA, {0,}, {0,}}; + // TsmConvert *convert = (TsmConvert *)getTsmOpPointer()->convert_pointer; + // CT_Param ct_params = {I_CGRA, {0}, {0}}; + + TsmActivation *activation = TsmNewActivation(); + TsmActivationInstr activationParam = {I_CGRA, + { + 0, + }, + { + 0, + }}; + TsmArith *arith = TsmNewArith(); + TsmArithInstr arithParam = {I_CGRA, + { + 0, + }, + { + 0, + }}; + TsmConvert *convert = TsmNewConvert(); + CT_Param ct_params = {I_CGRA, {0}, {0}}; + + if (fmt == Fmt_FP16) { + convert->FP16_FP32(&ct_params, in_addr, imm_a, elem_count); + cycle_value += TsmExecute(&ct_params); + + arith->MulVV(&arithParam, imm_a, imm_a, imm_b, elem_count, rnd_mode, + Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVV(&arithParam, imm_a, imm_b, imm_b, elem_count, rnd_mode, + Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVS(&arithParam, imm_b, a1, imm_b, elem_count, rnd_mode, Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->AddVV(&arithParam, imm_a, imm_b, imm_b, elem_count, rnd_mode, + Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + arith->MulVS(&arithParam, imm_b, a2, imm_b, elem_count, rnd_mode, Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + activation->Tanh(&activationParam, imm_b, imm_b, elem_count, Fmt_FP32); + cycle_value += TsmExecute(&activationParam); + + arith->AddVS(&arithParam, imm_b, a3, imm_b, elem_count, rnd_mode, Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVV(&arithParam, imm_a, imm_b, imm_b, elem_count, rnd_mode, + Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVS(&arithParam, imm_b, a4, imm_b, elem_count, rnd_mode, Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + convert->FP32_FP16(&ct_params, imm_b, out_addr, elem_count, + RND_NEAREST_EVEN); + cycle_value += TsmExecute(&ct_params); + } else if (fmt == Fmt_BF16) { + convert->BF16_FP32(&ct_params, in_addr, imm_a, elem_count); + cycle_value += TsmExecute(&ct_params); + + arith->MulVV(&arithParam, imm_a, imm_a, imm_b, elem_count, rnd_mode, + Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVV(&arithParam, imm_a, imm_b, imm_b, elem_count, rnd_mode, + Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVS(&arithParam, imm_b, a1, imm_b, elem_count, rnd_mode, Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->AddVV(&arithParam, imm_a, imm_b, imm_b, elem_count, rnd_mode, + Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVS(&arithParam, imm_b, a2, imm_b, elem_count, rnd_mode, Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + activation->Tanh(&activationParam, imm_b, imm_b, elem_count, Fmt_FP32); + cycle_value += TsmExecute(&activationParam); + arith->AddVS(&arithParam, imm_b, a3, imm_b, elem_count, rnd_mode, Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVV(&arithParam, imm_a, imm_b, imm_b, elem_count, rnd_mode, + Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + arith->MulVS(&arithParam, imm_b, a4, imm_b, elem_count, rnd_mode, Fmt_FP32); + cycle_value += TsmExecute(&arithParam); + + convert->FP32_BF16(&ct_params, imm_b, out_addr, elem_count, + RND_NEAREST_EVEN); + cycle_value += TsmExecute(&ct_params); + } else if (fmt == Fmt_FP32) { + arith->MulVV(&arithParam, in_addr, in_addr, out_addr, elem_count, rnd_mode, + fmt); + cycle_value += TsmExecute(&arithParam); + + arith->MulVV(&arithParam, in_addr, out_addr, out_addr, elem_count, rnd_mode, + fmt); + cycle_value += TsmExecute(&arithParam); + + arith->MulVS(&arithParam, out_addr, a1, out_addr, elem_count, rnd_mode, + fmt); + cycle_value += TsmExecute(&arithParam); + + arith->AddVV(&arithParam, in_addr, out_addr, out_addr, elem_count, rnd_mode, + fmt); + cycle_value += TsmExecute(&arithParam); + + arith->MulVS(&arithParam, out_addr, a2, out_addr, elem_count, rnd_mode, + fmt); + cycle_value += TsmExecute(&arithParam); + + activation->Tanh(&activationParam, out_addr, out_addr, elem_count, fmt); + cycle_value += TsmExecute(&activationParam); + + arith->AddVS(&arithParam, out_addr, a3, out_addr, elem_count, rnd_mode, + fmt); + cycle_value += TsmExecute(&arithParam); + + arith->MulVV(&arithParam, in_addr, out_addr, out_addr, elem_count, rnd_mode, + fmt); + cycle_value += TsmExecute(&arithParam); + + arith->MulVS(&arithParam, out_addr, a4, out_addr, elem_count, rnd_mode, + fmt); + cycle_value += TsmExecute(&arithParam); + } else { + } +} + +void op_gelu_none(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + TsmWaitfinish(); + uint8_t *in_ddr = (uint8_t *)get_spm_memory_mapping((uint64_t)(src)); + uint8_t *out_ddr = (uint8_t *)get_spm_memory_mapping((uint64_t)(dst)); + + get_erf_value(in_ddr, out_ddr, elem_count, fmt); + SYNCHRONOUS_INTRINSIC_SWITCH; + +#ifdef USING_RISCV + csi_dcache_clean_range((uint64_t *)out_ddr, elem_count * fmt); +#endif +} + +void op_gelu_tanh(uint64_t *src, uint64_t *imm, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + TsmWaitfinish(); + get_tanh_value(src, imm, dst, elem_count, fmt); + SYNCHRONOUS_INTRINSIC_SWITCH; +} diff --git a/third_party/wafer/crt/lib/Wafer/op_reduce_mul_impl.c b/third_party/wafer/crt/lib/Wafer/op_reduce_mul_impl.c new file mode 100755 index 00000000..21afac52 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/op_reduce_mul_impl.c @@ -0,0 +1,161 @@ +#include "op_reduce_mul_impl.h" +#include "wafer.h" +#include + +void op_reduce_mul_impl(void *in, void *out, Data_Shape shape, + uint32_t reduce_dim, Data_Format fmt) { + TsmArith *arith = TsmNewArith(); + TsmArithInstr arithIns = {I_CGRA, {0}, {0}}; + + uint32_t n = shape.n; + uint32_t h = shape.h; + uint32_t w = shape.w; + uint32_t c = shape.c; + + // reduce mul impl, only Cx is supported. + + // get_cx_align_base: 4, 8, 16, 32, 64 (> 32 except int_8) , 128 ( > 64 int_8) + uint32_t align_val = (c > 4) ? get_cx_align_base_new(c, fmt) : c; + // uint32_t one_align = get_cx_align_base(1, fmt); + uint32_t one_align = 1; + + // In triton, the channel dimension is always a power of two. + uint32_t cx_align = c; + if (cx_align > align_val && cx_align % align_val != 0) { + assert(cx_align % align_val == 0); + } else if (cx_align < align_val) { + assert(cx_align == next_power_of_two_64(cx_align) && + "cx_align should be power of two"); + } + uint64_t cx_align_mem_size = (uint64_t)cx_align * get_dtype_size_new(fmt); + uint64_t one_align_mem_size = (uint64_t)one_align * get_dtype_size_new(fmt); + + if (reduce_dim == 0) { + + TsmDataMoveInstr data_move_inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + // reduce C in NHWC. + int32_t nhw_cnt = n * h * w; + + if (cx_align == 1) { + TsmDataMove *data_move = TsmNewDataMove(); + St_StrideIteration src_si = {1, 1, 1, 1, 1, 1}; + St_StrideIteration dst_si = {1, 1, 1, 1, 1, 1}; + + data_move->GatherScatter(&data_move_inst, (uint64_t)in, (uint64_t)out, + nhw_cnt * cx_align_mem_size, &src_si, &dst_si); + + // Dispatch the command to accelerator + TsmExecute(&data_move_inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. + TsmDeleteDataMove(data_move); + } + + for (int32_t nhw_index = 0; nhw_index < nhw_cnt; nhw_index++) { + uint64_t src_in_addr = + (uint64_t)in + (uint64_t)nhw_index * cx_align_mem_size; + uint64_t dst_out_addr = + (uint64_t)out + (uint64_t)nhw_index * one_align_mem_size; + // init src + uint8_t bytes = get_dtype_size_new(fmt); + hybrid_value init_v = set_float2value(fmt, 1.0); + if (cx_align > c) { + + // op_reduce_micro_op_memset(src_in_addr + c * bytes, *(uint32_t + // *)&init_v, + // cx_align - c, fmt); + uint32_t loop = cx_align / align_val; + + for (int32_t loop_index = 1; loop_index < loop; loop_index++) { + uint64_t arith_src0_addr = src_in_addr; + uint64_t arith_src1_addr = + arith_src0_addr + (uint64_t)bytes * loop_index * align_val; + uint64_t arith_dst_addr = arith_src0_addr; + + // TsmArith *arith = (TsmArith *)getTsmOpPointer()->arith_pointer; + + arith->MulVV(&arithIns, arith_src0_addr, arith_src1_addr, + arith_dst_addr, align_val, RND_NEAREST_EVEN, fmt); + TsmExecute(&arithIns); + SYNCHRONOUS_INTRINSIC_SWITCH; + } + } + + uint32_t c_tail = (cx_align > align_val) ? align_val : cx_align; + for (uint32_t stride = c_tail / 2; stride > 0; stride = stride / 2) { + uint64_t arith_src0_addr = src_in_addr; + uint64_t arith_src1_addr = arith_src0_addr + (uint64_t)bytes * stride; + uint64_t arith_dst_addr = + (stride == 1) ? dst_out_addr : arith_src0_addr; + + // TsmArith *arith = (TsmArith *)getTsmOpPointer()->arith_pointer; + + arith->MulVV(&arithIns, arith_src0_addr, arith_src1_addr, + arith_dst_addr, stride, RND_NEAREST_EVEN, fmt); + TsmExecute(&arithIns); + SYNCHRONOUS_INTRINSIC_SWITCH; + } + } + } else if (reduce_dim == 1) { + TsmDataMove *data_move = TsmNewDataMove(); + + TsmDataMoveInstr data_move_inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + // reduce W in NHWC. + int32_t nh_cnt = n * h; + for (int32_t nh_index = 0; nh_index < nh_cnt; nh_index++) { + uint64_t src_in_addr = + (uint64_t)in + (uint64_t)nh_index * cx_align_mem_size; + uint64_t dst_out_addr = + (uint64_t)out + (uint64_t)nh_index * cx_align_mem_size; + // TsmOperatorPointer *opp = getTsmOpPointer(); + // TsmDataMove *data_move = (TsmDataMove *)opp->datamove_pointer; + + // Create command buffer. + + St_StrideIteration src_si = {1, 1, 1, 1, 1, 1}; + St_StrideIteration dst_si = {1, 1, 1, 1, 1, 1}; + + data_move->GatherScatter(&data_move_inst, (uint64_t)src_in_addr, + (uint64_t)dst_out_addr, cx_align_mem_size, + &src_si, &dst_si); + + // Dispatch the command to accelerator + TsmExecute(&data_move_inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + for (int32_t w_index = 1; w_index < w; w_index++) { + uint64_t arith_src0_addr = dst_out_addr; + uint64_t arith_src1_addr = + src_in_addr + (uint64_t)w_index * cx_align_mem_size; + uint64_t arith_dst_addr = dst_out_addr; + + // TsmArith *arith = (TsmArith *)getTsmOpPointer()->arith_pointer; + + arith->MulVV(&arithIns, arith_src0_addr, arith_src1_addr, + arith_dst_addr, cx_align, RND_NEAREST_EVEN, fmt); + TsmExecute(&arithIns); + SYNCHRONOUS_INTRINSIC_SWITCH; + } + } + // Destroy the command buffer. + TsmDeleteDataMove(data_move); + } else { + assert(0); + } + + TsmDeleteArith(arith); +} diff --git a/third_party/wafer/crt/lib/Wafer/pad.c b/third_party/wafer/crt/lib/Wafer/pad.c new file mode 100755 index 00000000..ea1be225 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/pad.c @@ -0,0 +1,40 @@ +//===------------------------ pad.c ---------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Pad see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Pad(uint64_t *src, uint16_t src_n, uint16_t src_h, uint16_t src_w, + uint16_t src_c, uint64_t *dst, uint16_t dst_n, uint16_t dst_h, + uint16_t dst_w, uint16_t dst_c, uint16_t pad_n, uint16_t pad_h, + uint16_t pad_w, uint16_t pad_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + Data_Shape shape2 = {dst_n, dst_h, dst_w, dst_c}; + Data_Shape shape3 = {pad_n, pad_h, pad_w, pad_c}; + cmd->Pad(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, shape3, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/pow.c b/third_party/wafer/crt/lib/Wafer/pow.c new file mode 100755 index 00000000..aa4ba03e --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/pow.c @@ -0,0 +1,467 @@ +#include "math.h" +#include "stdlib.h" +#include "wafer.h" +#include +// #include + +typedef union { + long i; + unsigned long u; + double f; + struct { + unsigned int lsw; + unsigned int msw; + } u32s; + struct { + unsigned int manl : 32; + unsigned int manh : 20; + unsigned int exp : 11; + unsigned int sign : 1; + } bits; +} Perf_64suf; + +typedef union { + int i; + unsigned int u; + float f; + struct { + unsigned int man : 23; + unsigned int exp : 8; + unsigned int sign : 1; + } bits; +} Perf_32suf; + +static unsigned long __L_tbl[] = { + 0x3ff0000000000000, 0x0000000000000000, 0x0000000000000000, + 0x8000000000000000, 0x3fefc00000000000, 0x0000000000000000, + 0x3f872c7ba2100000, 0xbd319b14945cf6ba, 0x3fef800000000000, + 0x0000000000000000, 0x3f9743ee86200000, 0xbd495539356f93dc, + 0x3fef400000000000, 0x0000000000000000, 0x3fa184b8e4c60000, + 0xbd52a0fadb83e4fe, 0x3fef000000000000, 0x0000000000000000, + 0x3fa77394c9da0000, 0xbd54e55443478fe0, 0x3feec00000000000, + 0x0000000000000000, 0x3fad6ebd1f200000, 0xbd2401fbaaa67e3c, + 0x3feea00000000000, 0x0000000000000000, 0x3fb0387efbcb0000, + 0xbd5e5897de9078d1, 0x3fee600000000000, 0x0000000000000000, + 0x3fb33f7cde150000, 0xbd48530f101bad19, 0x3fee200000000000, + 0x0000000000000000, 0x3fb64ce26c060000, 0x3d5c55ad0f90a497, + 0x3fede00000000000, 0x0000000000000000, 0x3fb960caf9ac0000, + 0xbd520d8486e9222a, 0x3feda00000000000, 0x0000000000000000, + 0x3fbc7b528b710000, 0xbd2c760bc9b188c4, 0x3fed800000000000, + 0x0000000000000000, 0x3fbe0b1ae8f30000, 0xbd054cda62d3926e, + 0x3fed400000000000, 0x0000000000000000, 0x3fc097e38ce60000, + 0x3d2924ae921f7eca, 0x3fed000000000000, 0x0000000000000000, + 0x3fc22dadc2ab0000, 0x3d5a4b69691d7994, 0x3fece00000000000, + 0x0000000000000000, 0x3fc2f9e32d5c0000, 0xbd117b2f1731efbe, + 0x3feca00000000000, 0x0000000000000000, 0x3fc494f863b90000, + 0xbd5065a2ac5ac35f, 0x3fec800000000000, 0x0000000000000000, + 0x3fc563dc29ff8000, 0x3d56590643906f2a, 0x3fec400000000000, + 0x0000000000000000, 0x3fc7046031c78000, 0x3d4f84be19cb9578, + 0x3fec000000000000, 0x0000000000000000, 0x3fc8a8980abf8000, + 0x3d5e9933354dbf17, 0x3febe00000000000, 0x0000000000000000, + 0x3fc97c1cb13c8000, 0xbd03f7a55cd2af4c, 0x3feba00000000000, + 0x0000000000000000, 0x3fcb2602497d8000, 0xbd565d3990cb67ba, + 0x3feb800000000000, 0x0000000000000000, 0x3fcbfc67a8000000, + 0xbd3667f21fa8423f, 0x3feb400000000000, 0x0000000000000000, + 0x3fcdac22d3e48000, 0xbd5f1680dd458fb2, 0x3feb200000000000, + 0x0000000000000000, 0x3fce857d3d360000, 0x3d4367bde40c5e6d, + 0x3feb000000000000, 0x0000000000000000, 0x3fcf5fd8a9060000, + 0x3d5f1a4847f7b278, 0x3feac00000000000, 0x0000000000000000, + 0x3fd08bce0d960000, 0xbd37204f55bbf90d, 0x3feaa00000000000, + 0x0000000000000000, 0x3fd0fa848044c000, 0xbd495def21f8497b, + 0x3fea600000000000, 0x0000000000000000, 0x3fd1d982c9d54000, + 0xbd58f7ca2cff7b90, 0x3fea400000000000, 0x0000000000000000, + 0x3fd249cd2b13c000, 0x3d4ad89a083e072a, 0x3fea200000000000, + 0x0000000000000000, 0x3fd2baa0c34c0000, 0xbd5e1410132ae5e4, + 0x3fe9e00000000000, 0x0000000000000000, 0x3fd39de8e1558000, + 0x3d5f6f7f2b4bd1c4, 0x3fe9c00000000000, 0x0000000000000000, + 0x3fd4106017c40000, 0xbd535d71963580ba, 0x3fe9a00000000000, + 0x0000000000000000, 0x3fd48365e695c000, 0x3d5796aa2981fdbc, + 0x3fe9800000000000, 0x0000000000000000, 0x3fd4f6fbb2cec000, + 0x3d3661e393a16b95, 0x3fe9400000000000, 0x0000000000000000, + 0x3fd5dfdcf1eec000, 0xbd51f1bbd2926f16, 0x3fe9200000000000, + 0x0000000000000000, 0x3fd6552b49988000, 0xbd5d8894dbdff331, + 0x3fe9000000000000, 0x0000000000000000, 0x3fd6cb0f6865c000, + 0x3d41d406db502403, 0x3fe8e00000000000, 0x0000000000000000, + 0x3fd7418acebc0000, 0xbd4ce2935fff809a, 0x3fe8a00000000000, + 0x0000000000000000, 0x3fd8304d90c10000, 0x3d5fd32a3ab0a4b5, + 0x3fe8800000000000, 0x0000000000000000, 0x3fd8a8980abfc000, + 0xbd266cccab240e90, 0x3fe8600000000000, 0x0000000000000000, + 0x3fd921800924c000, 0x3d5d3b7f711abd5c, 0x3fe8400000000000, + 0x0000000000000000, 0x3fd99b072a96c000, 0x3d3ac9bca36fd02e, + 0x3fe8200000000000, 0x0000000000000000, 0x3fda152f14298000, + 0x3d1b3d7b0e65d2ce, 0x3fe8000000000000, 0x0000000000000000, + 0x3fda8ff971810000, 0x3d44bc302ffa76fb, 0x3fe7e00000000000, + 0x0000000000000000, 0x3fdb0b67f4f48000, 0xbd57f00af09dc1c7, + 0x3fe7a00000000000, 0x0000000000000000, 0x3fdc043859e30000, + 0xbd22642415d47384, 0x3fe7800000000000, 0x0000000000000000, + 0x3fdc819dc2d44000, 0x3d5fe43895d8ac46, 0x3fe7600000000000, + 0x0000000000000000, 0x3fdcffae611ac000, 0x3d512b628e2d05d7, + 0x3fe7400000000000, 0x0000000000000000, 0x3fdd7e6c0abc4000, + 0xbd450e785694a8c6, 0x3fe7200000000000, 0x0000000000000000, + 0x3fddfdd89d588000, 0xbd51d4f639bb5cdf, 0x3fe7000000000000, + 0x0000000000000000, 0x3fde7df5fe538000, 0x3d45669df6a2b592, + 0x3fe6e00000000000, 0x0000000000000000, 0x3fdefec61b010000, + 0x3d5f855b4987c5d5, 0x3fe6c00000000000, 0x0000000000000000, + 0x3fdf804ae8d0c000, 0x3d4a0331af2e6fea, 0x3fe6a00000000000, + 0x0000000000000000, 0x3fe0014332be0000, 0x3cf9518ce032f41d, + 0x3fe6800000000000, 0x0000000000000000, 0x3fe042bd4b9a8000, + 0xbd3b3b3864c60011, 0x3fe6600000000000, 0x0000000000000000, + 0x3fe08494c66b8000, 0x3d5ddf82e1fe57c7, 0x3fe6400000000000, + 0x0000000000000000, 0x3fe0c6caaf0c6000, 0xbd54d20c519e12f4, + 0x3fe6200000000000, 0x0000000000000000, 0x3fe1096015dee000, + 0x3d43676289cd3dd4, 0x3fe6000000000000, 0x0000000000000000, + 0x3fe14c560fe68000, 0x3d55f101c141e670, 0x3fe5e00000000000, + 0x0000000000000000, 0x3fe18fadb6e2e000, 0xbd587cc95d0a2ee8, + 0x3fe5c00000000000, 0x0000000000000000, 0x3fe1d368296b6000, + 0xbd5b567e7ee54aef, 0x3fe5a00000000000, 0x0000000000000000, + 0x3fe217868b0c4000, 0xbd5030ab442ce320, 0x3fe5800000000000, + 0x0000000000000000, 0x3fe25c0a0463c000, 0xbd250520a377c7ec, + 0x3fe5800000000000, 0x0000000000000000, 0x3fe25c0a0463c000, + 0xbd250520a377c7ec, 0x3fe5600000000000, 0x0000000000000000, + 0x3fe2a0f3c3408000, 0xbd5f48e1a4725559, 0x3fe5400000000000, + 0x0000000000000000, 0x3fe2e644fac04000, 0x3d5faf6283bf2868, + 0x3fe5200000000000, 0x0000000000000000, 0x3fe32bfee370e000, + 0x3d5cd0cb4492f1bc, 0x3fe5000000000000, 0x0000000000000000, + 0x3fe37222bb708000, 0xbd5708b4b2b5056c, 0x3fe4e00000000000, + 0x0000000000000000, 0x3fe3b8b1c68fa000, 0x3d4bb4b69336b66e, + 0x3fe4c00000000000, 0x0000000000000000, 0x3fe3ffad4e750000, + 0xbd5c5432aeb609f5, 0x3fe4a00000000000, 0x0000000000000000, + 0x3fe44716a2c08000, 0x3d33106e404cabb7, 0x3fe4a00000000000, + 0x0000000000000000, 0x3fe44716a2c08000, 0x3d33106e404cabb7, + 0x3fe4800000000000, 0x0000000000000000, 0x3fe48eef19318000, + 0xbd49bcaf1aa4168a, 0x3fe4600000000000, 0x0000000000000000, + 0x3fe4d7380dcc4000, 0x3d31646b761c48de, 0x3fe4400000000000, + 0x0000000000000000, 0x3fe51ff2e3022000, 0xbd56879fa00b120a, + 0x3fe4200000000000, 0x0000000000000000, 0x3fe5692101d9c000, + 0xbd56b37dcf60e620, 0x3fe4200000000000, 0x0000000000000000, + 0x3fe5692101d9c000, 0xbd56b37dcf60e620, 0x3fe4000000000000, + 0x0000000000000000, 0x3fe5b2c3da198000, 0xbd5b8afe492bf6ff, + 0x3fe3e00000000000, 0x0000000000000000, 0x3fe5fcdce2728000, + 0xbd3125d6cbcd1095, 0x3fe3c00000000000, 0x0000000000000000, + 0x3fe6476d98ada000, 0xbd4bd9b32266d92c, 0x3fe3c00000000000, + 0x0000000000000000, 0x3fe6476d98ada000, 0xbd4bd9b32266d92c, + 0x3fe3a00000000000, 0x0000000000000000, 0x3fe6927781d94000, + 0xbd5aaf6f137a3d8c, 0x3fe3800000000000, 0x0000000000000000, + 0x3fe6ddfc2a790000, 0xbd3ce60916e52e91, 0x3fe3600000000000, + 0x0000000000000000, 0x3fe729fd26b70000, 0x3d4f1f5ae718f241, + 0x3fe3600000000000, 0x0000000000000000, 0x3fe729fd26b70000, + 0x3d4f1f5ae718f241, 0x3fe3400000000000, 0x0000000000000000, + 0x3fe7767c12968000, 0xbd46eb9612e0b4f3, 0x3fe3200000000000, + 0x0000000000000000, 0x3fe7c37a9227e000, 0x3d4fed21f9cb2cc5, + 0x3fe3000000000000, 0x0000000000000000, 0x3fe810fa51bf6000, + 0x3d47f5dc57266758, 0x3fe3000000000000, 0x0000000000000000, + 0x3fe810fa51bf6000, 0x3d47f5dc57266758, 0x3fe2e00000000000, + 0x0000000000000000, 0x3fe85efd062c6000, 0x3d45b338360c2ae2, + 0x3fe2c00000000000, 0x0000000000000000, 0x3fe8ad846cf36000, + 0x3d53481b85a54d7f, 0x3fe2c00000000000, 0x0000000000000000, + 0x3fe8ad846cf36000, 0x3d53481b85a54d7f, 0x3fe2a00000000000, + 0x0000000000000000, 0x3fe8fc924c89a000, 0x3d5908df8ec933b3, + 0x3fe2800000000000, 0x0000000000000000, 0x3fe94c287492c000, + 0x3d436c101ee13440, 0x3fe2800000000000, 0x0000000000000000, + 0x3fe94c287492c000, 0x3d436c101ee13440, 0x3fe2600000000000, + 0x0000000000000000, 0x3fe99c48be206000, 0x3d3e41fa0a62e6ae, + 0x3fe2400000000000, 0x0000000000000000, 0x3fe9ecf50bf44000, + 0xbd1d97ee9124773b, 0x3fe2400000000000, 0x0000000000000000, + 0x3fe9ecf50bf44000, 0xbd1d97ee9124773b, 0x3fe2200000000000, + 0x0000000000000000, 0x3fea3e2f4ac44000, 0xbd13f94e00e7d6bc, + 0x3fe2000000000000, 0x0000000000000000, 0x3fea8ff971810000, + 0x3d54bc302ffa76fb, 0x3fe2000000000000, 0x0000000000000000, + 0x3fea8ff971810000, 0x3d54bc302ffa76fb, 0x3fe1e00000000000, + 0x0000000000000000, 0x3feae255819f0000, 0x3d31659d8e2d7d38, + 0x3fe1c00000000000, 0x0000000000000000, 0x3feb35458761e000, + 0xbd570d0fa8f9603b, 0x3fe1c00000000000, 0x0000000000000000, + 0x3feb35458761e000, 0xbd570d0fa8f9603b, 0x3fe1a00000000000, + 0x0000000000000000, 0x3feb88cb9a2ac000, 0xbd55bdaf522a183c, + 0x3fe1a00000000000, 0x0000000000000000, 0x3feb88cb9a2ac000, + 0xbd55bdaf522a183c, 0x3fe1800000000000, 0x0000000000000000, + 0x3febdce9dcc96000, 0x3d2871a7610e40bd, 0x3fe1600000000000, + 0x0000000000000000, 0x3fec31a27dd00000, 0x3d569378d0928989, + 0x3fe1600000000000, 0x0000000000000000, 0x3fec31a27dd00000, + 0x3d569378d0928989, 0x3fe1400000000000, 0x0000000000000000, + 0x3fec86f7b7ea4000, 0x3d551167134e9647, 0x3fe1400000000000, + 0x0000000000000000, 0x3fec86f7b7ea4000, 0x3d551167134e9647, + 0x3fe1200000000000, 0x0000000000000000, 0x3fecdcebd2374000, + 0xbd49ad57391924a7, 0x3fe1200000000000, 0x0000000000000000, + 0x3fecdcebd2374000, 0xbd49ad57391924a7, 0x3fe1000000000000, + 0x0000000000000000, 0x3fed338120a6e000, 0xbd33167ccc538261, + 0x3fe0e00000000000, 0x0000000000000000, 0x3fed8aba045b0000, + 0x3d2c7a4ff65ddbc9, 0x3fe0e00000000000, 0x0000000000000000, + 0x3fed8aba045b0000, 0x3d2c7a4ff65ddbc9, 0x3fe0c00000000000, + 0x0000000000000000, 0x3fede298ec0ba000, 0x3d5819530c22d152, + 0x3fe0c00000000000, 0x0000000000000000, 0x3fede298ec0ba000, + 0x3d5819530c22d152, 0x3fe0a00000000000, 0x0000000000000000, + 0x3fee3b20546f6000, 0xbd556bde9f1f0d3d, 0x3fe0a00000000000, + 0x0000000000000000, 0x3fee3b20546f6000, 0xbd556bde9f1f0d3d, + 0x3fe0800000000000, 0x0000000000000000, 0x3fee9452c8a72000, + 0xbd5fb0e626c0de13, 0x3fe0800000000000, 0x0000000000000000, + 0x3fee9452c8a72000, 0xbd5fb0e626c0de13, 0x3fe0600000000000, + 0x0000000000000000, 0x3feeee32e2aec000, 0x3d597da24fd75f61, + 0x3fe0600000000000, 0x0000000000000000, 0x3feeee32e2aec000, + 0x3d597da24fd75f61, 0x3fe0400000000000, 0x0000000000000000, + 0x3fef48c34bd1e000, 0x3d52dd67591d81df, 0x3fe0400000000000, + 0x0000000000000000, 0x3fef48c34bd1e000, 0x3d52dd67591d81df, + 0x3fe0200000000000, 0x0000000000000000, 0x3fefa406bd244000, + 0x3d3ef5d00e390a00, 0x3fe0000000000000, 0x0000000000000000, + 0x3ff0000000000000, 0x0000000000000000, 0x3fe0000000000000, + 0x0000000000000000, 0x3ff0000000000000, 0x0000000000000000, +}; + +static unsigned long __Exp_tbl[256] = { + 0x0000000000000000, 0x0000000000000000, 0x0000163da9fb3335, + 0x3c9b3b4f1a88bf6e, 0x00002c9a3e778061, 0xbc7160139cd8dc5c, + 0x00004315e86e7f85, 0xbc905e7a108766d1, 0x000059b0d3158574, + 0x3c8cd2523567f613, 0x0000706b29ddf6de, 0xbc8bce8023f98efa, + 0x0000874518759bc8, 0x3c60f74e61e6c861, 0x00009e3ecac6f383, + 0x3c90a3e45b33d399, 0x0000b5586cf9890f, 0x3c979aa65d837b6d, + 0x0000cc922b7247f7, 0x3c8eb51a92fdeffc, 0x0000e3ec32d3d1a2, + 0x3c3ebe3d702f9cd1, 0x0000fb66affed31b, 0xbc6a033489906e0b, + 0x00011301d0125b51, 0xbc9556522a2fbd0e, 0x00012abdc06c31cc, + 0xbc5080ef8c4eea55, 0x0001429aaea92de0, 0xbc91c923b9d5f416, + 0x00015a98c8a58e51, 0x3c80d3e3e95c55af, 0x000172b83c7d517b, + 0xbc801b15eaa59348, 0x00018af9388c8dea, 0xbc8f1ff055de323d, + 0x0001a35beb6fcb75, 0x3c8b898c3f1353bf, 0x0001bbe084045cd4, + 0xbc96d99c7611eb26, 0x0001d4873168b9aa, 0x3c9aecf73e3a2f60, + 0x0001ed5022fcd91d, 0xbc8fe782cb86389d, 0x0002063b88628cd6, + 0x3c8a6f4144a6c38d, 0x00021f49917ddc96, 0x3c807a05b0e4047d, + 0x0002387a6e756238, 0x3c968efde3a8a894, 0x000251ce4fb2a63f, + 0x3c875e18f274487d, 0x00026b4565e27cdd, 0x3c80472b981fe7f2, + 0x000284dfe1f56381, 0xbc96b87b3f71085e, 0x00029e9df51fdee1, + 0x3c82f7e16d09ab31, 0x0002b87fd0dad990, 0xbc3d219b1a6fbffa, + 0x0002d285a6e4030b, 0x3c8b3782720c0ab4, 0x0002ecafa93e2f56, + 0x3c6e149289cecb8f, 0x000306fe0a31b715, 0x3c834d754db0abb6, + 0x00032170fc4cd831, 0x3c864201e2ac744c, 0x00033c08b26416ff, + 0x3c8fdd395dd3f84a, 0x000356c55f929ff1, 0xbc86a3803b8e5b04, + 0x000371a7373aa9cb, 0xbc924aedcc4b5068, 0x00038cae6d05d866, + 0xbc9907f81b512d8e, 0x0003a7db34e59ff7, 0xbc71d1e83e9436d2, + 0x0003c32dc313a8e5, 0xbc991919b3ce1b15, 0x0003dea64c123422, + 0x3c859f48a72a4c6d, 0x0003fa4504ac801c, 0xbc9312607a28698a, + 0x0004160a21f72e2a, 0xbc58a78f4817895b, 0x000431f5d950a897, + 0xbc7c2c9b67499a1b, 0x00044e086061892d, 0x3c4363ed60c2ac12, + 0x00046a41ed1d0057, 0x3c9666093b0664ef, 0x000486a2b5c13cd0, + 0x3c6ecce1daa10379, 0x0004a32af0d7d3de, 0x3c93ff8e3f0f1230, + 0x0004bfdad5362a27, 0x3c7690cebb7aafb0, 0x0004dcb299fddd0d, + 0x3c931dbdeb54e077, 0x0004f9b2769d2ca7, 0xbc8f94340071a38e, + 0x000516daa2cf6642, 0xbc87deccdc93a349, 0x0005342b569d4f82, + 0xbc78dec6bd0f385f, 0x000551a4ca5d920f, 0xbc861246ec7b5cf6, + 0x00056f4736b527da, 0x3c93350518fdd78e, 0x00058d12d497c7fd, + 0x3c7b98b72f8a9b05, 0x0005ab07dd485429, 0x3c9063e1e21c5409, + 0x0005c9268a5946b7, 0x3c34c7855019c6ea, 0x0005e76f15ad2148, + 0x3c9432e62b64c035, 0x000605e1b976dc09, 0xbc8ce44a6199769f, + 0x0006247eb03a5585, 0xbc8c33c53bef4da8, 0x0006434634ccc320, + 0xbc845378892be9ae, 0x0006623882552225, 0xbc93cedd78565858, + 0x00068155d44ca973, 0x3c5710aa807e1964, 0x0006a09e667f3bcd, + 0xbc93b3efbf5e2228, 0x0006c012750bdabf, 0xbc6a12ad8734b982, + 0x0006dfb23c651a2f, 0xbc6367efb86da9ee, 0x0006ff7df9519484, + 0xbc80dc3d54e08851, 0x00071f75e8ec5f74, 0xbc781f647e5a3ecf, + 0x00073f9a48a58174, 0xbc86ee4ac08b7db0, 0x00075feb564267c9, + 0xbc8619321e55e68a, 0x000780694fde5d3f, 0x3c909ccb5e09d4d3, + 0x0007a11473eb0187, 0xbc7b32dcb94da51d, 0x0007c1ed0130c132, + 0x3c94ecfd5467c06b, 0x0007e2f336cf4e62, 0x3c65ebe1abd66c55, + 0x00080427543e1a12, 0xbc88a1c52fb3cf42, 0x00082589994cce13, + 0xbc9369b6f13b3734, 0x0008471a4623c7ad, 0xbc805e843a19ff1e, + 0x000868d99b4492ed, 0xbc94d450d872576e, 0x00088ac7d98a6699, + 0x3c90ad675b0e8a00, 0x0008ace5422aa0db, 0x3c8db72fc1f0eab4, + 0x0008cf3216b5448c, 0xbc65b6609cc5e7ff, 0x0008f1ae99157736, + 0x3c7bf68359f35f44, 0x0009145b0b91ffc6, 0xbc93091fa71e3d83, + 0x00093737b0cdc5e5, 0xbc5da9b88b6c1e29, 0x00095a44cbc8520f, + 0xbc6c23f97c90b959, 0x00097d829fde4e50, 0xbc92434322f4f9aa, + 0x0009a0f170ca07ba, 0xbc85ca6cd7668e4b, 0x0009c49182a3f090, + 0x3c71affc2b91ce27, 0x0009e86319e32323, 0x3c6dd235e10a73bb, + 0x000a0c667b5de565, 0xbc87c50422622263, 0x000a309bec4a2d33, + 0x3c8b1c86e3e231d5, 0x000a5503b23e255d, 0xbc91bbd1d3bcbb15, + 0x000a799e1330b358, 0x3c90cc319cee31d2, 0x000a9e6b5579fdbf, + 0x3c8469846e735ab3, 0x000ac36bbfd3f37a, 0xbc82dfcd978e9db4, + 0x000ae89f995ad3ad, 0x3c8c1a7792cb3387, 0x000b0e07298db666, + 0xbc907b8f4ad1d9fa, 0x000b33a2b84f15fb, 0xbc55c3d956dcaeba, + 0x000b59728de5593a, 0xbc90a40e3da6f640, 0x000b7f76f2fb5e47, + 0xbc68d6f438ad9334, 0x000ba5b030a1064a, 0xbc91eee26b588a35, + 0x000bcc1e904bc1d2, 0x3c74ffd70a5fddcd, 0x000bf2c25bd71e09, + 0xbc91bdfbfa9298ac, 0x000c199bdd85529c, 0x3c736eae30af0cb3, + 0x000c40ab5fffd07a, 0x3c8ee3325c9ffd94, 0x000c67f12e57d14b, + 0x3c84e08fd10959ac, 0x000c8f6d9406e7b5, 0x3c63cdaf384e1a67, + 0x000cb720dcef9069, 0x3c676b2c6c921968, 0x000cdf0b555dc3fa, + 0xbc808a1883ccb5d2, 0x000d072d4a07897c, 0xbc8fad5d3ffffa6f, + 0x000d2f87080d89f2, 0xbc900dae3875a949, 0x000d5818dcfba487, + 0x3c74a385a63d07a7, 0x000d80e316c98398, 0xbc82919e2040220f, + 0x000da9e603db3285, 0x3c8e5a50d5c192ac, 0x000dd321f301b460, + 0x3c843a59ac016b4b, 0x000dfc97337b9b5f, 0xbc82d52107b43e1f, + 0x000e264614f5a129, 0xbc892ab93b470dc9, 0x000e502ee78b3ff6, + 0x3c74b604603a88d3, 0x000e7a51fbc74c83, 0x3c83c5ec519d7271, + 0x000ea4afa2a490da, 0xbc8ff7128fd391f0, 0x000ecf482d8e67f1, + 0xbc8dae98e223747d, 0x000efa1bee615a27, 0x3c8ec3bc41aa2008, + 0x000f252b376bba97, 0x3c842b94c3a9eb32, 0x000f50765b6e4540, + 0x3c8a64a931d185ee, 0x000f7bfdad9cbe14, 0xbc8e37bae43be3ed, + 0x000fa7c1819e90d8, 0x3c77893b4d91cd9d, 0x000fd3c22b8f71f1, + 0x3c5305c14160cc89, +}; + +#define ADD_BIT 0xc010100000000000 +#define EXPMASK 0xfff0000000000000 +#define EXP_SHIFTER 52776558134271.0 // 1.5000000000290754*2^45 +#define ONE 1.0 +#define c1l_c1h_8 2.035527417494173e-17 +#define A0_ha 0.16032146466245167 +#define A1_ha 0.2885390081778253 +#define A2_ha 0.20609929024234264 +#define A3_ha 0.4808983469629878 +#define B0_ha -0.18035669405180624 +#define B1_ha -0.3606737602222562 +#define B2_ha -0.2404491725744602 +#define B3_ha -0.7213475204444817 +#define Invln2_pow 1.4426950408889634 +#define C0_ha 0.0013333562801027398 +#define C1_ha 0.055504108664817364 +#define D0_ha 0.009618132633217303 +#define D1_ha 0.24022650695908054 +#define LN2_pow 0.6931471805599453 + +#define mant_mask_powf_ha (0x007fffff) +#define scaled_expon_powf_ha 0x3b800000 +#define shift_powf_ha 0x00008000 // int +#define mask_powf_ha 0xffff0000 // int +#define one_powf_ha 1.0 // single-precision +#define AddConst_powf_ha 211106232532992 +#define A0_powf_ha -0.3606803440419083 +#define A1_powf_ha 0.48090361399617193 +#define A2_powf_ha -0.72134752040442129 +#define A3_powf_ha 1.4426950408769454 +#define B0_powf_ha 0.055504596955955943 +#define B1_powf_ha 0.24022885514364503 +#define B2_powf_ha 0.69314718052020807 + +int32_t round_to_even2(float value) { + // 检查是否为正向无穷大 + if (value == INFINITY) { + return INT32_MAX; + } + // 检查是否为负无穷 + if (value == -INFINITY) { + return INT32_MIN; + } + // 获取整数部分 + float int_part = floor(fabs(value)); // 使用 fabs 来处理负数 + // 获取小数部分 + float frac_part = fabs(value) - int_part; + + if (frac_part < 0.5) { + // 如果小数部分小于 0.5,直接舍去 + return (int32_t)(int_part) * (value < 0 ? -1 : 1); + } else if (frac_part > 0.5) { + // 如果小数部分大于 0.5,直接进位 + return (int32_t)(int_part + 1) * (value < 0 ? -1 : 1); + } else { + // 如果小数部分恰好是 0.5,根据整数部分的奇偶性决定 + int32_t int_value = (int32_t)int_part; // 整数部分的整型值 + // printf("int_value = %d\n", int_value); + if (int_value % 2 == 0) { + // 如果整数部分是偶数, 舍去 + return (int32_t)(int_part) * (value < 0 ? -1 : 1); + } else { + // 如果整数部分是奇数, 进位 + return (int32_t)(int_part + 1) * (value < 0 ? -1 : 1); + } + } +} + +float powf(float src1, float src2) { + float nan = NAN; + if (src2 == 0) + return 1; + if (src1 == 1) + return 1; + if (src2 == 1) + return src1; + if (isnan(src1) | isnan(src2)) + return nan; + float y = (float)src2; + float round_y = round_to_even2(y); + float diff = y - round_y; + bool not_int = fabs(diff) >= 1e-9; + if (src1 == 0) { + if (src2 < 0) + return INFINITY; + return 0; + } + if (src1 < 0) { + if (not_int) + return nan; + } + if (fabs(src1) < 0.1) { + if (src2 > 38) + return 0; + if (src2 < -38) + return INFINITY; + } + double n, r, r2, p, q, xd0, xd1, x0, x1, tmpf2, tmpf0, tmpf1; + Perf_64suf z, tmp, tbl, tbl0; + Perf_32suf m; + unsigned long idx, xu; + + tmp.f = src1; + xu = tmp.u + ADD_BIT; + m.u = xu >> 32; + m.i = m.i >> 11; + xu = EXPMASK & xu; + tmp.u = tmp.u - xu; + idx = m.u; + m.u = m.u & 0xfffffe00; + idx = idx & 0x1fc; + m.i = m.i >> 9; + tbl.u = __L_tbl[idx]; + q = (double)m.i; + r = tbl.f * tmp.f - ONE; + tbl.u = __L_tbl[idx + 2]; + tmpf0 = q + tbl.f; + r2 = r * r; + tbl.u = __L_tbl[idx + 3]; + tmpf1 = r * c1l_c1h_8 + tbl.f; + q = r * Invln2_pow + tmpf0; + tmpf2 = tmpf0 - q; + xd1 = r * A0_ha + B0_ha; + xd0 = r * A3_ha + B3_ha; + x1 = r2 * r2; + x0 = r * A1_ha + B1_ha; + p = r * A2_ha + B2_ha; + r = r * Invln2_pow + tmpf2; + xd0 = r2 * x0 + xd0; + xd1 = r2 * xd1 + p; + xd1 = x1 * xd1 + xd0; + tmpf0 = r + tmpf1; + xd1 = xd1 * r2 + tmpf0; + + // exp + tmpf0 = q * src2; + x1 = xd1 * src2 + tmpf0; + q = src2 * q - tmpf0; + p = tmpf0 - x1; + z.f = EXP_SHIFTER + x1; + p = src2 * xd1 + p; + tmp.f = x1; + n = z.f - EXP_SHIFTER; + q = q + p; + idx = tmp.u >> 32; + idx = idx >> 16; + idx = idx & 0x7fff; + r = x1 - n + q; + r2 = r * r; + idx = z.u; + tmp.u = idx; + idx = idx & 0x7f; + idx = idx + idx; + z.f = r * C0_ha + D0_ha; + tmp.u = tmp.u & 0xffffff80; + q = r * C1_ha + D1_ha; + tbl0.u = __Exp_tbl[idx]; + tbl.u = __Exp_tbl[idx + 1]; + p = LN2_pow * r + tbl.f; + q = r2 * z.f + q; + q = q * r2 + p; + x0 = tmp.f; + tmp.u = tmp.u << 45; + tmp.u = tbl0.u | tmp.u; + x0 = tmp.f; + x0 = x0 * q + x0; + if (x0 > 3.4028235e38) + return INFINITY; + return (float)x0; +} diff --git a/third_party/wafer/crt/lib/Wafer/pow2.c b/third_party/wafer/crt/lib/Wafer/pow2.c new file mode 100755 index 00000000..34b28e22 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/pow2.c @@ -0,0 +1,33 @@ +//===------------------------ pow2.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Pow2 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Pow2(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmTranscendental *cmd = g_intrinsic()->transcendental_pointer; + TsmTranscendentalInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Pow2(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/print.c b/third_party/wafer/crt/lib/Wafer/print.c new file mode 100755 index 00000000..48eace81 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/print.c @@ -0,0 +1,27 @@ +// ===------------------------ print.c ------------------------------------===// + +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. + +// ===---------------------------------------------------------------------===// + +// Enable wafer kernel printf support + +#include "lib_log.h" +#include "wafer.h" +#include +#include +#include + +void __Print(const char *__restrict fmt, ...) { + va_list args; + va_start(args, fmt); + + // FIXME: va_list memory layout is specific to the platform. +#ifndef USE_SIM_MODE + _tsm_ep_log(__FILE__, __func__, __LINE__, KCORE_LOG_ERROR, fmt, args); +#else + vprintf(fmt, args); +#endif + va_end(args); +} diff --git a/third_party/wafer/crt/lib/Wafer/randgen.c b/third_party/wafer/crt/lib/Wafer/randgen.c new file mode 100755 index 00000000..9f515a26 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/randgen.c @@ -0,0 +1,39 @@ +//===------------------------ randgen.c -----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::RandGen see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __RandGen(uint64_t *src0, uint64_t *src1, uint64_t *dst0, uint64_t *dst1, + uint64_t *dst2, uint32_t src_elem_num, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmPeripheral *cmd = g_intrinsic()->peripheral_pointer; + TsmPeripheralInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + ; + + // The instruction receives SPM addresses, not seed words as addresses. + // src_elem_num is the SDK byte count (a multiple of 128). + cmd->RandGen(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst0, + (uint64_t)dst1, (uint64_t)dst2, src_elem_num, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/rdma.c b/third_party/wafer/crt/lib/Wafer/rdma.c new file mode 100755 index 00000000..5de4bb8b --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/rdma.c @@ -0,0 +1,140 @@ +//===------------------------ rdma.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Rdma, see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" +#include + +void __Rdma4d(void *restrict dest, const void *restrict src, + uint32_t elem_count, uint32_t stride0, uint32_t iteration0, + uint32_t stride1, uint32_t iteration1, uint32_t stride2, + uint32_t iteration2, uint32_t fmt) { + INTRNISIC_RUN_SWITCH; + TsmRdma *rdma = g_intrinsic()->rdma_pointer; + TsmRdmaInstr inst = {I_RDMA, + { + 0, + }, + { + 0, + }}; + + rdma->AddSrcDst(&inst, (uint64_t)src, (uint64_t)dest, (Data_Format)fmt); + rdma->ConfigStrideIteration(&inst, elem_count, stride0, iteration0, stride1, + iteration1, stride2, iteration2); + TsmExecute(&inst); + TsmWaitfinish(); +} + +void __Rdma1d(void *restrict dest, const void *restrict src, + uint32_t elem_count, uint32_t fmt) { + TsmRdma *rdma = g_intrinsic()->rdma_pointer; + TsmRdmaInstr inst = {I_RDMA, + { + 0, + }, + { + 0, + }}; + rdma->Rdma1d(&inst, (uint64_t)src, (uint64_t)dest, elem_count, + (Data_Format)fmt); + TsmExecute(&inst); + TsmWaitfinish(); +} + +// Rdma line by line. +void __RdmaVectorize(char *srcPtr, char *dstPtr, int *src_shape, + int *src_stride, int *dst_shape, int *dst_stride, int rank, + uint32_t elem_bytes, uint32_t fmt, int innermost_rank, + int inner_elem_count) { + INTRNISIC_RUN_SWITCH; + TsmRdma *rdma = g_intrinsic()->rdma_pointer; + TsmRdmaInstr inst = {I_RDMA, + { + 0, + }, + { + 0, + }}; + + int64_t readIndex = 0; + int64_t writeIndex = 0; + int64_t indices[rank], srcStrides[rank], dstStrides[rank]; + + // Initialize index and scale strides. + for (int rankp = 0; rankp < rank; ++rankp) { + indices[rankp] = 0; + srcStrides[rankp] = (int64_t)src_stride[rankp] * (int64_t)elem_bytes; + dstStrides[rankp] = (int64_t)dst_stride[rankp] * (int64_t)elem_bytes; + } + + for (;;) { + // Copy inner dim, line by line. + rdma->Rdma1d(&inst, (uint64_t)(srcPtr + readIndex), + (uint64_t)(dstPtr + writeIndex), inner_elem_count, + (Data_Format)fmt); + TsmExecute(&inst); + TsmWaitfinish(); + + // Advance index and read position. + // Start from the second-to-last dimension, copy one line at a time + for (int64_t axis = innermost_rank; axis >= 0; --axis) { + // Advance at current axis. + int64_t newIndex = ++indices[axis]; + readIndex += srcStrides[axis]; + writeIndex += dstStrides[axis]; + // If this is a valid index, we have our next index, so continue copying. + if (src_shape[axis] != newIndex) + break; + // We reached the end of this axis. If this is axis 0, we are done. + if (axis == 0) + return; + // Else, reset to 0 and undo the advancement of the linear index that + // this axis had. Then continue with the axis one outer. + indices[axis] = 0; + readIndex -= (int64_t)newIndex * srcStrides[axis]; + writeIndex -= (int64_t)newIndex * dstStrides[axis]; + } + } +} + +void __Rdma(uint64_t *src, uint64_t *dst, int *src_shape, int *src_stride, + int *dst_shape, int *dst_stride, int rank, uint32_t elem_bytes, + uint32_t fmt) { + INTRNISIC_RUN_SWITCH; + + // Dynamic shape, kernel implementation will cause shape equal to 0 + for (int i = 0; i < rank; i++) { + if (src_shape[i] == 0) { + return; + } + } + + // If inner dim stride is 1, use scalar rdma. + if (src_stride[rank - 1] != 1 || dst_stride[rank - 1] != 1) { + __RdmaVectorize((char *)src, (char *)dst, src_shape, src_stride, dst_shape, + dst_stride, rank, elem_bytes, Fmt_INT8, rank - 1, + elem_bytes); + return; + } + legalizeMemoryOpAttribute(src_shape, src_stride, dst_shape, dst_stride, rank, + &elem_bytes, &fmt); + + if (rank == 4 && no_reverse_memory_access(src_stride, rank) && + is_contiguous(dst_shape, dst_stride, rank)) { + __Rdma4d(dst, src, src_shape[3], src_stride[2], src_shape[2], src_stride[1], + src_shape[1], src_stride[0], src_shape[0], fmt); + return; + } + + __RdmaVectorize((char *)src, (char *)dst, src_shape, src_stride, dst_shape, + dst_stride, rank, elem_bytes, fmt, rank - 2, + src_shape[rank - 1]); +} diff --git a/third_party/wafer/crt/lib/Wafer/recip.c b/third_party/wafer/crt/lib/Wafer/recip.c new file mode 100755 index 00000000..31a78a9c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/recip.c @@ -0,0 +1,33 @@ +//===------------------------ recip.c--------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::recipVVOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __RecipVV(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->RecipVV(&inst, (uint64_t)src, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); +} diff --git a/third_party/wafer/crt/lib/Wafer/recv.c b/third_party/wafer/crt/lib/Wafer/recv.c new file mode 100755 index 00000000..0195f6b6 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/recv.c @@ -0,0 +1,26 @@ +//===------------------------ recv.c --------------------------------------===// +// +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Recv, see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +// #include "instr_adapter_plat.h" + +#include "direct_dte_and_fsm.h" +#include "wafer.h" +#include "tx81_spm.h" +#include +#include + +uint32_t __get_pid(uint32_t); + +// Blockingly receive data from a source tile into a destination buffer. +// Returns the destination buffer address. +void __Recv(int64_t chip_x, int64_t chip_y, int64_t die_id, int64_t tile_id, + void *dst, uint32_t elem_bytes, uint32_t data_size) { + // TODO + return; +} diff --git a/third_party/wafer/crt/lib/Wafer/reduce.c b/third_party/wafer/crt/lib/Wafer/reduce.c new file mode 100755 index 00000000..9cf2be3c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/reduce.c @@ -0,0 +1,112 @@ +//===---------------------- reduce.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TsmReduce, see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "op_reduce_mul_impl.h" +#include "wafer.h" +// The arguments list is aligned with TsmConv in WaferOps.td +void __ReduceSum(uint64_t *src, uint64_t *dst, uint32_t dim, uint16_t src_n, + uint16_t src_h, uint16_t src_w, uint16_t src_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create reduce command buffer. + TsmReduce *cmd = g_intrinsic()->reduce_pointer; + TsmReduceInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + // TODO + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + cmd->ReduceSum(&inst, (uint64_t)src, (uint64_t)dst, dim, shape1, + (Data_Format)fmt); + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} + +void __ReduceAvg(uint64_t *src, uint64_t *dst, uint32_t dim, uint16_t src_n, + uint16_t src_h, uint16_t src_w, uint16_t src_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create reduce command buffer. + TsmReduce *cmd = g_intrinsic()->reduce_pointer; + TsmReduceInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + // TODO + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + cmd->ReduceAvg(&inst, (uint64_t)src, (uint64_t)dst, dim, shape1, + (Data_Format)fmt); + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} + +void __ReduceMax(uint64_t *src, uint64_t *dst, uint32_t dim, uint16_t src_n, + uint16_t src_h, uint16_t src_w, uint16_t src_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create reduce command buffer. + TsmReduce *cmd = g_intrinsic()->reduce_pointer; + TsmReduceInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + // TODO + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + cmd->ReduceMax(&inst, (uint64_t)src, (uint64_t)dst, dim, shape1, + (Data_Format)fmt); + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} + +void __ReduceMin(uint64_t *src, uint64_t *dst, uint32_t dim, uint16_t src_n, + uint16_t src_h, uint16_t src_w, uint16_t src_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create reduce command buffer. + TsmReduce *cmd = g_intrinsic()->reduce_pointer; + TsmReduceInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + // TODO + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + cmd->ReduceMin(&inst, (uint64_t)src, (uint64_t)dst, dim, shape1, + (Data_Format)fmt); + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + + // Destroy the command buffer. +} +void __ReduceMul(uint64_t *src, uint64_t *dst, uint32_t dim, uint16_t src_n, + uint16_t src_h, uint16_t src_w, uint16_t src_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + + // TODO + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + op_reduce_mul_impl(src, dst, shape1, dim, fmt); +} diff --git a/third_party/wafer/crt/lib/Wafer/relation.c b/third_party/wafer/crt/lib/Wafer/relation.c new file mode 100755 index 00000000..252b6ce3 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/relation.c @@ -0,0 +1,562 @@ +//===------------------------ relation.c-----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::RelationOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __BoolEqualVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolEqualVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolUnEqualVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolUnEqualVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} + +void __BoolGreaterEqualVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolGreaterEqualVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + TsmWaitfinish(); + // Destroy the command buffer. +} + +void __BoolGreaterVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolGreaterVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolLessEqualVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolLessEqualVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolLessThenVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolLessThenVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __EqualVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->EqualVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __UnEqualVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->UnEqualVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __GreaterEqualVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->GreaterEqualVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __GreaterVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->GreaterVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __LessEqualVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->LessEqualVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __LessThenVV(uint64_t *src0, uint64_t *src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->LessThenVV(&inst, (uint64_t)src0, (uint64_t)src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolEqualVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolEqualVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolUnEqualVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolUnEqualVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolGreaterEqualVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolGreaterEqualVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, + elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolGreaterVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolGreaterVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolLessEqualVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolLessEqualVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __BoolLessThenVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->BoolLessThenVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __EqualVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->EqualVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __UnEqualVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->UnEqualVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __GreaterEqualVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->GreaterEqualVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __GreaterVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->GreaterVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __LessEqualVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->LessEqualVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} + +void __LessThenVS(uint64_t *src0, uint32_t src1, uint64_t *dst, + uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmRelation *cmd = g_intrinsic()->relation_pointer; + TsmRelationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->LessThenVS(&inst, (uint64_t)src0, src1, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/relu.c b/third_party/wafer/crt/lib/Wafer/relu.c new file mode 100755 index 00000000..e58e66bd --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/relu.c @@ -0,0 +1,33 @@ +//===------------------------ relu.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Relu see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Relu(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmActivation *cmd = g_intrinsic()->activation_pointer; + TsmActivationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Relu(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/rotate180.c b/third_party/wafer/crt/lib/Wafer/rotate180.c new file mode 100755 index 00000000..88fbec16 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/rotate180.c @@ -0,0 +1,38 @@ +//===------------------------ rotate180.c ---------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Rotate180 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Rotate180(uint64_t *src, uint16_t src_n, uint16_t src_h, uint16_t src_w, + uint16_t src_c, uint64_t *dst, uint16_t dst_n, uint16_t dst_h, + uint16_t dst_w, uint16_t dst_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + Data_Shape shape2 = {dst_n, dst_h, dst_w, dst_c}; + cmd->Rotate180(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/rotate270.c b/third_party/wafer/crt/lib/Wafer/rotate270.c new file mode 100755 index 00000000..54518307 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/rotate270.c @@ -0,0 +1,38 @@ +//===------------------------ rotate270.c ---------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Rotate270 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Rotate270(uint64_t *src, uint16_t src_n, uint16_t src_h, uint16_t src_w, + uint16_t src_c, uint64_t *dst, uint16_t dst_n, uint16_t dst_h, + uint16_t dst_w, uint16_t dst_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + Data_Shape shape2 = {dst_n, dst_h, dst_w, dst_c}; + cmd->Rotate270(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/rotate90.c b/third_party/wafer/crt/lib/Wafer/rotate90.c new file mode 100755 index 00000000..94499abf --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/rotate90.c @@ -0,0 +1,38 @@ +//===------------------------ rotate90.c ----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Rotate90 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Rotate90(uint64_t *src, uint16_t src_n, uint16_t src_h, uint16_t src_w, + uint16_t src_c, uint64_t *dst, uint16_t dst_n, uint16_t dst_h, + uint16_t dst_w, uint16_t dst_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + Data_Shape shape2 = {dst_n, dst_h, dst_w, dst_c}; + cmd->Rotate90(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/rsqrt.c b/third_party/wafer/crt/lib/Wafer/rsqrt.c new file mode 100755 index 00000000..622bca9b --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/rsqrt.c @@ -0,0 +1,35 @@ +//===------------------------ rsqrt.c -------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::RsqrtVVOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __RsqrtVV(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->RsqrtVV(&inst, (uint64_t)src, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/satrelu.c b/third_party/wafer/crt/lib/Wafer/satrelu.c new file mode 100755 index 00000000..a71339af --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/satrelu.c @@ -0,0 +1,35 @@ +//===------------------------ satrelu.c -----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Satrelu see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Satrelu(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmActivation *cmd = g_intrinsic()->activation_pointer; + TsmActivationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Satrelu(&inst, (uint64_t)src, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/send.c b/third_party/wafer/crt/lib/Wafer/send.c new file mode 100755 index 00000000..80a81ab6 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/send.c @@ -0,0 +1,199 @@ +//===------------------------ send.c --------------------------------------===// +// +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Send, see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +// #include "instr_adapter_plat.h" +#include "direct_dte_and_fsm.h" +#include "wafer.h" +#include "tx81_spm.h" +#include +#include +#include + +#define MAX_TILE_NUM 16 + +const uint32_t tile_physical_relation[MAX_TILE_NUM] = { + 0, 1, 2, 3, 7, 11, 15, 14, 13, 12, 8, 9, 10, 6, 5, 4}; + +// 获取物理连接最近的下一个tile id +int32_t getNextNearestTileId(uint32_t tileId) { + for (int i = 0; i < MAX_TILE_NUM; i++) { + if (tile_physical_relation[i] == tileId) { + if (i != (MAX_TILE_NUM - 1)) { + return tile_physical_relation[++i]; + } else { + return tile_physical_relation[0]; + } + } + } + + return -1; +} + +// 获取物理连接最近的上一个tile id +int32_t getPrevNearestTileId(uint32_t tileId) { + for (int i = 0; i < MAX_TILE_NUM; i++) { + if (tile_physical_relation[i] == tileId) { + if (i != 0) { + return tile_physical_relation[--i]; + } else { + return tile_physical_relation[MAX_TILE_NUM - 1]; + } + } + } + + return -1; +} + +void tile_sync_by_spm_single_direction(int32_t tile_this, int32_t tile_a, + int32_t tile_x, int32_t tile_y, + uint32_t this_sync_spm_index, + uint32_t other_sync_spm_index) { + // this_debug_spm_ptr 只用于板端记录调试信息 + volatile uint32_t *this_debug_spm_ptr = + (volatile uint32_t *)get_spm_memory_mapping(SINGLE_SPM_SYNC_DEBUG_ADDR); + this_debug_spm_ptr[0] = tile_this; + this_debug_spm_ptr[1] = tile_a; + this_debug_spm_ptr[2] = tile_x; + this_debug_spm_ptr[3] = tile_y; + this_debug_spm_ptr[4] = this_sync_spm_index; + this_debug_spm_ptr[5] = other_sync_spm_index; + + uint64_t tile_a_spm = + get_tile_spm_addr_base(tile_a, tile_x, tile_y) + SINGLE_SPM_SYNC_ADDR; + *(uint64_t *)(this_debug_spm_ptr + 6) = tile_a_spm; + volatile uint32_t *this_spm_ptr = + (volatile uint32_t *)get_spm_memory_mapping(SINGLE_SPM_SYNC_ADDR); + *(uint64_t *)(this_debug_spm_ptr + 8) = (uint64_t)this_spm_ptr; + volatile uint32_t *tile_a_spm_ptr = (volatile uint32_t *)(tile_a_spm); + + this_debug_spm_ptr[10] = 0; // 写对端开始 + tile_a_spm_ptr[other_sync_spm_index] = 1; // forward + + this_debug_spm_ptr[10] = 1; // 写对端结束 + + while (!this_spm_ptr[this_sync_spm_index]) { + } + + this_spm_ptr[this_sync_spm_index] = 0; // forward +} + +#define SCFG_TILE_ID_ADDR 0x6A0058 // KUIPER_ADDR_MAP_REG_BASE 0x6A0000 + +uint32_t __get_pid(uint32_t); + +int initTileId(uint32_t tileId, uint32_t rowLength) { + // init 1D tile-id + *(volatile uint32_t *)(get_spm_memory_mapping(TILE_ID_ADDR)) = tileId; + // init 2D logic id + *(volatile uint32_t *)(get_spm_memory_mapping(LOGIC_ID_ADDR)) = + *(volatile uint32_t *)(SCFG_TILE_ID_ADDR); + // init row-length + *(volatile uint32_t *)(get_spm_memory_mapping(ROW_LENGTH_ADDR)) = rowLength; + *(volatile uint32_t *)(get_spm_memory_mapping(INNER_CHIP_ERROR_CODE)) = 0; + + // __LOG__(KCORE_LOG_DEBUG, "logic_id:0x%x, tileId:%u, rowLength:%u\n", + // *(volatile uint32_t *)(get_spm_memory_mapping(LOGIC_ID_ADDR)), + // tileId, rowLength); + return 0; +} + +static void noc_memory_fence(void) { +#ifdef __riscv + __asm__ __volatile__("fence iorw, iorw" ::: "memory"); +#else + __sync_synchronize(); +#endif +} + +// SINGLE_SPM_SYNC reserves 0x320..0x370 for the ring protocol. Use its first +// two words for request/acknowledgement. A sender cannot publish its next +// request until its previous request has been consumed and acknowledged. +// This prevents a fast tile from overwriting an unconsumed boolean signal. +static void noc_ring_sync(uint32_t previous, uint32_t next) { + volatile uint32_t *local = + (volatile uint32_t *)get_spm_memory_mapping(SINGLE_SPM_SYNC_ADDR); + volatile uint32_t *previous_spm = (volatile uint32_t *)( + get_tile_spm_addr_base(previous, 4, 4) + SINGLE_SPM_SYNC_ADDR); + volatile uint32_t *next_spm = (volatile uint32_t *)( + get_tile_spm_addr_base(next, 4, 4) + SINGLE_SPM_SYNC_ADDR); + + noc_memory_fence(); + previous_spm[0] = 1; + while (!local[0]) { + } + noc_memory_fence(); + local[0] = 0; + noc_memory_fence(); + next_spm[1] = 1; + while (!local[1]) { + } + noc_memory_fence(); + local[1] = 0; + noc_memory_fence(); +} + +// Send to the next tile and wait for both transmission and ring reception. +void __Send(int64_t chipX, int64_t chipY, int64_t dieId, int64_t tileId, + void *restrict dst, void *restrict src, uint32_t elem_bytes, + uint64_t data_size) { + uint32_t coreIndex = __get_pid(0); // 全局tile id + initTileId(coreIndex, 4); + // __EP_LOG__(0, "+++++++++ Send444 dst:%lx, src: %lx, cur_tileId:%d, + // nextTileId: %d, data_size: %d\n", dst, src, coreIndex, tileId, data_size); + (void)chipX; + (void)chipY; + (void)dieId; + int64_t nextTileId = getNextNearestTileId(coreIndex); + int64_t preTileId = getPrevNearestTileId(coreIndex); + // tile_sync_by_spm_single_direction(coreIndex, preTileId, 4, 4, 0, 0); + const TsmOperatorPointer *intrinsic = g_intrinsic(); + + int fringFsmId = DIRECT_DTE_FSM_ID_0; + int remottFringFsmId = DIRECT_DTE_FSM_ID_0; + + void *fdteNode = direct_dte_attach(0); + void *fringFsmHd = direct_fsm_monitor_init(fringFsmId, 0, data_size, 1); + + TsmStream *stream = (TsmStream *)(intrinsic->stream_pointer); + uint64_t nextTileBaseAddr = get_tile_spm_addr_base(nextTileId, 4, 4); + + stream->wait_finish(); + + // __EP_LOG__(0, "fdte info base: %ld, dst: %ld, src: %ld, remote_fsm_id: %d, + // data_size: %d, dst_tileId: %d, this_tileId: %d\n", + // nextTileBaseAddr, nextTileBaseAddr + (uint64_t)dst, (uint64_t)src, + // remottFringFsmId, data_size, tileId, coreIndex); + DirectDTESendInfo fdteInfo = {.src_addr = (uint64_t)src, + .dst_addr = nextTileBaseAddr + (uint64_t)dst, + .length = data_size, + .remote_fsm_id = remottFringFsmId, + .mode = 0, // unicast + .dst_tile = nextTileId, + .tile_this = coreIndex, + .dte_node = fdteNode}; + set_direct_fsm_monitor_dst_addr(fringFsmId, nextTileBaseAddr + (uint64_t)dst); + noc_ring_sync(preTileId, nextTileId); + direct_dte_send_async(&fdteInfo); // 把当前数据异步发送给下一个tile + // __EP_LOG__(KCORE_LOG_DEBUG, "send data to next tile: %u, current + // tile:%u.\n", + // getNextNearestTileId(coreIndex), coreIndex); + + direct_fsm_monitor_receive(coreIndex, preTileId, + fringFsmHd); // 阻塞接收前一个tile发送的数据 + // __EP_LOG__(KCORE_LOG_DEBUG, + // "receive data from coreIndex: %u, current tile:%u.\n", + // getPrevNearestTileId(coreIndex), coreIndex); + direct_dte_wait_done(&fdteInfo); // 等待异步发送完成 + TsmWaitfinish(); + + noc_ring_sync(preTileId, nextTileId); + direct_dte_release(fdteNode); + direct_fsm_monitor_deinit(fringFsmHd); + // __EP_LOG__(0, "-------- Send\n") +} diff --git a/third_party/wafer/crt/lib/Wafer/sigmoid.c b/third_party/wafer/crt/lib/Wafer/sigmoid.c new file mode 100755 index 00000000..da1ae9a5 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/sigmoid.c @@ -0,0 +1,35 @@ +//===------------------------ sigmoid.c -----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Sigmoid see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Sigmoid(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmActivation *cmd = g_intrinsic()->activation_pointer; + TsmActivationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Sigmoid(&inst, (uint64_t)src, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/sin.c b/third_party/wafer/crt/lib/Wafer/sin.c new file mode 100755 index 00000000..350fefb9 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/sin.c @@ -0,0 +1,33 @@ +//===------------------------ Sin.c ---------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Sin see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Sin(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmTranscendental *cmd = g_intrinsic()->transcendental_pointer; + TsmTranscendentalInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Sin(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/softplus.c b/third_party/wafer/crt/lib/Wafer/softplus.c new file mode 100755 index 00000000..4d3b3a46 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/softplus.c @@ -0,0 +1,36 @@ +//===------------------------ softplus.cpp +//------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Softplus see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Softplus(uint64_t *src, uint64_t *dst, uint32_t elem_count, + uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmActivation *cmd = g_intrinsic()->activation_pointer; + TsmActivationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Softplus(&inst, (uint64_t)src, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/sqrt.c b/third_party/wafer/crt/lib/Wafer/sqrt.c new file mode 100755 index 00000000..ebd1bc9d --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/sqrt.c @@ -0,0 +1,34 @@ +//===------------------------ sqrt.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::SqrtVVOp see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __SqrtVV(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmArith *cmd = g_intrinsic()->arith_pointer; + TsmArithInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->SqrtVV(&inst, (uint64_t)src, (uint64_t)dst, elem_count, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/tanh.c b/third_party/wafer/crt/lib/Wafer/tanh.c new file mode 100755 index 00000000..2129157c --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/tanh.c @@ -0,0 +1,33 @@ +//===------------------------ tanh.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Tanh see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Tanh(uint64_t *src, uint64_t *dst, uint32_t elem_count, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmActivation *cmd = g_intrinsic()->activation_pointer; + TsmActivationInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->Tanh(&inst, (uint64_t)src, (uint64_t)dst, elem_count, (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/tensornorm.c b/third_party/wafer/crt/lib/Wafer/tensornorm.c new file mode 100755 index 00000000..7003b5a7 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/tensornorm.c @@ -0,0 +1,38 @@ +//===------------------------ tensornorm.c --------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TensorNorm see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __TensorNorm(uint64_t *src, uint16_t src_n, uint16_t src_h, uint16_t src_w, + uint16_t src_c, uint64_t *dst, uint16_t dst_n, uint16_t dst_h, + uint16_t dst_w, uint16_t dst_c, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_n, src_h, src_w, src_c}; + Data_Shape shape2 = {dst_n, dst_h, dst_w, dst_c}; + cmd->TensorNom(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/tf32_bf16.c b/third_party/wafer/crt/lib/Wafer/tf32_bf16.c new file mode 100755 index 00000000..d0964bb3 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/tf32_bf16.c @@ -0,0 +1,33 @@ +//===------------------------ tf32_bf16.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TF32_BF16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __TF32_BF16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->TF32_BF16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/tf32_fp16.c b/third_party/wafer/crt/lib/Wafer/tf32_fp16.c new file mode 100755 index 00000000..eff7856e --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/tf32_fp16.c @@ -0,0 +1,32 @@ +//===------------------------ tf32_fp16.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TF32_FP16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __TF32_FP16(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->TF32_FP16(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/tf32_fp32.c b/third_party/wafer/crt/lib/Wafer/tf32_fp32.c new file mode 100755 index 00000000..100cc4fe --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/tf32_fp32.c @@ -0,0 +1,32 @@ +//===------------------------ tf32_fp32.c ---------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TF32_FP32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __TF32_FP32(uint64_t *src, uint64_t *dst, uint32_t elem_count) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->TF32_FP32(&inst, (uint64_t)src, (uint64_t)dst, elem_count); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/tf32_int16.c b/third_party/wafer/crt/lib/Wafer/tf32_int16.c new file mode 100755 index 00000000..87557cfc --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/tf32_int16.c @@ -0,0 +1,33 @@ +//===------------------------ tf32_int16.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TF32_INT16 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __TF32_INT16(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->TF32_INT16(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/tf32_int32.c b/third_party/wafer/crt/lib/Wafer/tf32_int32.c new file mode 100755 index 00000000..a86bed1f --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/tf32_int32.c @@ -0,0 +1,33 @@ +//===------------------------ tf32_int32.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TF32_INT32 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __TF32_INT32(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->TF32_INT32(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/tf32_int8.c b/third_party/wafer/crt/lib/Wafer/tf32_int8.c new file mode 100755 index 00000000..cb5c7320 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/tf32_int8.c @@ -0,0 +1,33 @@ +//===------------------------ tf32_int8.c --------------------------------===// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::TF32_INT8 see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __TF32_INT8(uint64_t *src, uint64_t *dst, uint32_t elem_count, + RND_MODE round) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmConvert *cmd = g_intrinsic()->convert_pointer; + TsmConvertInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + cmd->TF32_INT8(&inst, (uint64_t)src, (uint64_t)dst, elem_count, round); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/transpose.c b/third_party/wafer/crt/lib/Wafer/transpose.c new file mode 100755 index 00000000..cbd3a734 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/transpose.c @@ -0,0 +1,37 @@ +//===------------------------ transpose.c ---------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Transpose see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" + +void __Transpose(uint64_t *src, uint64_t *dst, int32_t *src_shape, + int32_t *dst_shape, uint16_t fmt) { + INTRNISIC_RUN_SWITCH; + // Create command buffer. + TsmDataMove *cmd = g_intrinsic()->datamove_pointer; + TsmDataMoveInstr inst = {I_CGRA, + { + 0, + }, + { + 0, + }}; + + Data_Shape shape1 = {src_shape[0], src_shape[1], src_shape[2], src_shape[3]}; + Data_Shape shape2 = {dst_shape[0], dst_shape[1], dst_shape[2], dst_shape[3]}; + cmd->Transpose(&inst, (uint64_t)src, shape1, (uint64_t)dst, shape2, + (Data_Format)fmt); + + // Dispatch the command to accelerator + TsmExecute(&inst); + SYNCHRONOUS_INTRINSIC_SWITCH; + + // Destroy the command buffer. +} diff --git a/third_party/wafer/crt/lib/Wafer/wafer.c b/third_party/wafer/crt/lib/Wafer/wafer.c new file mode 100755 index 00000000..1536a764 --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/wafer.c @@ -0,0 +1,178 @@ +//===------------------------- wafer.c--------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" +#include + +#ifdef __cplusplus +extern "C" { +#endif + +bool is_contiguous(int *shape, int *strides, int rank) { + int expected_stride = 1; + for (int i = rank - 1; i >= 0; i--) { + if (shape[i] != 1 && strides[i] != expected_stride) { + return false; + } + expected_stride *= shape[i]; + } + return true; +} +uint64_t next_power_of_two_64(uint64_t x) { + if (x == 0) { + return 1; + } + x--; + x |= x >> 1; + x |= x >> 2; + x |= x >> 4; + x |= x >> 8; + x |= x >> 16; + x |= x >> 32; + return x + 1; +} + +uint32_t get_dtype_size_new(Data_Format fmt) { + switch (fmt) { + case Fmt_INT8: + return sizeof(int8_t); + case Fmt_INT16: + case Fmt_FP16: + case Fmt_BF16: + return sizeof(int16_t); + case Fmt_INT32: + case Fmt_FP32: + case Fmt_TF32: + return sizeof(int32_t); + case Fmt_INT64: + return sizeof(int64_t); + default: + assert(false && "Unsupported format\n"); + return 0; + } +} + +uint32_t get_cx_align_base_new(uint32_t c, Data_Format fmt) { + switch (fmt) { + case Fmt_INT8: + return c < 128 ? (c < 4 ? 4 : next_power_of_two_64(c)) : 128; + case Fmt_INT16: + case Fmt_FP16: + case Fmt_BF16: + case Fmt_INT32: + case Fmt_FP32: + case Fmt_TF32: + return c < 64 ? (c < 4 ? 4 : next_power_of_two_64(c)) : 64; + default: + assert(false && "Unsupported format\n"); + return 0; + } +} + +bool no_reverse_memory_access(int *stride, int rank) { + for (int i = 1; i < rank; i++) { + if (stride[i] < 0) { + return false; + } + } + return true; +} + +void wafer_memcpy(char *srcPtr, char *dstPtr, int *src_shape, int *src_stride, + int *dst_shape, int *dst_stride, int rank, + uint32_t elem_bytes) { + int64_t readIndex = 0; + int64_t writeIndex = 0; + int64_t indices[rank], srcStrides[rank], dstStrides[rank]; + + // Initialize index and scale strides. + for (int rankp = 0; rankp < rank; ++rankp) { + indices[rankp] = 0; + srcStrides[rankp] = (int64_t)src_stride[rankp] * (int64_t)elem_bytes; + dstStrides[rankp] = (int64_t)dst_stride[rankp] * (int64_t)elem_bytes; + } + + for (;;) { + // Copy over the element, byte by byte. + for (int i = 0; i < elem_bytes; i++) + dstPtr[writeIndex + i] = srcPtr[readIndex + i]; + + // Advance index and read position. + // Loop from innermost dimension + for (int64_t axis = rank - 1; axis >= 0; --axis) { + // Advance at current axis. + int64_t newIndex = ++indices[axis]; + readIndex += srcStrides[axis]; + writeIndex += dstStrides[axis]; + // If this is a valid index, we have our next index, so continue copying. + if (src_shape[axis] != newIndex) + break; + // We reached the end of this axis. If this is axis 0, we are done. + if (axis == 0) + return; + // Else, reset to 0 and undo the advancement of the linear index that + // this axis had. Then continue with the axis one outer. + indices[axis] = 0; + readIndex -= newIndex * srcStrides[axis]; + writeIndex -= newIndex * dstStrides[axis]; + } + } +} + +void legalizeMemoryOpAttribute(int *src_shape, int *src_stride, int *dst_shape, + int *dst_stride, int rank, uint32_t *elem_bytes, + uint32_t *fmt) { + switch (*fmt) { + case Fmt_INT8: { + break; + } + case Fmt_INT16: + case Fmt_FP16: + case Fmt_BF16: { + *fmt = Fmt_FP16; + break; + } + case Fmt_INT32: + case Fmt_FP32: + case Fmt_TF32: { + *fmt = Fmt_FP32; + break; + } + case Fmt_INT64: { + *fmt = Fmt_FP32; + src_shape[rank - 1] *= sizeof(int64_t) / sizeof(int32_t); + dst_shape[rank - 1] *= sizeof(int64_t) / sizeof(int32_t); + *elem_bytes = sizeof(int32_t); + // Last stride is always 1 + for (int i = 0; i < rank - 1; i++) { + src_stride[i] *= 2; + dst_stride[i] *= 2; + } + break; + } + default: { + // Other formats are not supported. + assert(false && "Unsupported format\n"); + break; + } + } +} + +// Used for kcore load/store data from/to spm +const int64_t spmMappingOffset = 0x30400000; + +int8_t *get_spm_memory_mapping_wrapper(uint64_t ptr) { +#ifdef USE_SIM_MODE + return get_spm_memory_mapping(ptr); +#else + return (int8_t *)(ptr + spmMappingOffset); +#endif +} + +#ifdef __cplusplus +} +#endif diff --git a/third_party/wafer/crt/lib/Wafer/wdma.c b/third_party/wafer/crt/lib/Wafer/wdma.c new file mode 100755 index 00000000..7fb6179f --- /dev/null +++ b/third_party/wafer/crt/lib/Wafer/wdma.c @@ -0,0 +1,139 @@ +//===------------------------ wdma.c --------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Runtime API of MLIR operation tx::Wdma, see WaferOps.td for detail. +// +//===----------------------------------------------------------------------===// + +#include "wafer.h" +#include + +void __Wdma4d(void *restrict dest, const void *restrict src, + uint32_t elem_count, uint32_t stride0, uint32_t iteration0, + uint32_t stride1, uint32_t iteration1, uint32_t stride2, + uint32_t iteration2, uint32_t fmt) { + INTRNISIC_RUN_SWITCH; + TsmWdma *wdma = g_intrinsic()->wdma_pointer; + TsmWdmaInstr inst = {I_WDMA, + { + 0, + }, + { + 0, + }}; + + wdma->AddSrcDst(&inst, (uint64_t)src, (uint64_t)dest, (Data_Format)fmt); + wdma->ConfigStrideIteration(&inst, elem_count, stride0, iteration0, stride1, + iteration1, stride2, iteration2); + TsmExecute(&inst); + TsmWaitfinish(); +} + +void __Wdma1d(void *restrict dest, const void *restrict src, + uint32_t elem_count, uint32_t fmt) { + TsmWdma *wdma = g_intrinsic()->wdma_pointer; + TsmWdmaInstr inst = {I_WDMA, + { + 0, + }, + { + 0, + }}; + wdma->Wdma1d(&inst, (uint64_t)src, (uint64_t)dest, elem_count, + (Data_Format)fmt); + TsmExecute(&inst); + TsmWaitfinish(); +} + +// Wdma line by line. +void __WdmaVectorize(char *srcPtr, char *dstPtr, int *src_shape, + int *src_stride, int *dst_shape, int *dst_stride, int rank, + uint32_t elem_bytes, uint32_t fmt, int innermost_rank, + int inner_elem_count) { + INTRNISIC_RUN_SWITCH; + TsmWdma *wdma = g_intrinsic()->wdma_pointer; + TsmWdmaInstr inst = {I_WDMA, + { + 0, + }, + { + 0, + }}; + + int64_t readIndex = 0; + int64_t writeIndex = 0; + int64_t indices[rank], srcStrides[rank], dstStrides[rank]; + + // Initialize index and scale strides. + for (int rankp = 0; rankp < rank; ++rankp) { + indices[rankp] = 0; + srcStrides[rankp] = (int64_t)src_stride[rankp] * (int64_t)elem_bytes; + dstStrides[rankp] = (int64_t)dst_stride[rankp] * (int64_t)elem_bytes; + } + + for (;;) { + // Copy inner dim, line by line. + wdma->Wdma1d(&inst, (uint64_t)(srcPtr + readIndex), + (uint64_t)(dstPtr + writeIndex), inner_elem_count, + (Data_Format)fmt); + TsmExecute(&inst); + TsmWaitfinish(); + + // Advance index and read position. + // Start from the second-to-last dimension, copy one line at a time + for (int64_t axis = innermost_rank; axis >= 0; --axis) { + // Advance at current axis. + int64_t newIndex = ++indices[axis]; + readIndex += srcStrides[axis]; + writeIndex += dstStrides[axis]; + // If this is a valid index, we have our next index, so continue copying. + if (src_shape[axis] != newIndex) + break; + // We reached the end of this axis. If this is axis 0, we are done. + if (axis == 0) + return; + // Else, reset to 0 and undo the advancement of the linear index that + // this axis had. Then continue with the axis one outer. + indices[axis] = 0; + readIndex -= (int64_t)newIndex * srcStrides[axis]; + writeIndex -= (int64_t)newIndex * dstStrides[axis]; + } + } +} + +void __Wdma(uint64_t *src, uint64_t *dst, int *src_shape, int *src_stride, + int *dst_shape, int *dst_stride, int rank, uint32_t elem_bytes, + uint32_t fmt) { + INTRNISIC_RUN_SWITCH; + // Dynamic shape, kernel implementation will cause shape equal to 0 + for (int i = 0; i < rank; i++) { + if (src_shape[i] == 0) { + return; + } + } + + // If inner dim stride is 1, use scalar wdma. + if (src_stride[rank - 1] != 1 || dst_stride[rank - 1] != 1) { + __WdmaVectorize((char *)src, (char *)dst, src_shape, src_stride, dst_shape, + dst_stride, rank, elem_bytes, Fmt_INT8, rank - 1, + elem_bytes); + return; + } + legalizeMemoryOpAttribute(src_shape, src_stride, dst_shape, dst_stride, rank, + &elem_bytes, &fmt); + + if (rank == 4 && no_reverse_memory_access(dst_stride, rank) && + is_contiguous(src_shape, src_stride, rank)) { + __Wdma4d(dst, src, dst_shape[3], dst_stride[2], dst_shape[2], dst_stride[1], + dst_shape[1], dst_stride[0], dst_shape[0], fmt); + return; + } + + __WdmaVectorize((char *)src, (char *)dst, src_shape, src_stride, dst_shape, + dst_stride, rank, elem_bytes, fmt, rank - 2, + src_shape[rank - 1]); +} diff --git a/third_party/wafer/examples/_wafer_reference.py b/third_party/wafer/examples/_wafer_reference.py new file mode 100644 index 00000000..1ae54ec9 --- /dev/null +++ b/third_party/wafer/examples/_wafer_reference.py @@ -0,0 +1,27 @@ +"""Independent CPU decoding for the scaled-dot numerical oracle.""" + +import torch + + +def upcast_mxfp_cpu(packed, scale, format_name, dtype): + """Decode E2M1/E4M3/E5M2 plus E8M0 scales, rounding like the original oracle. + + packed is contiguous and packed along its last dimension. For E2M1, the low + nibble precedes the high nibble. A scale byte of 0 maps to zero and 255 to NaN, + matching the bitcast/where logic in the original Triton reference kernel. + """ + if packed.device.type != "cpu" or scale.device.type != "cpu": + raise ValueError("The Wafer numerical oracle must run on CPU") + packed = packed.contiguous() + if format_name == "e2m1": + codes = torch.stack((packed & 15, packed >> 4), dim=-1).flatten(-2) + magnitudes = torch.tensor([0, 0.5, 1, 1.5, 2, 3, 4, 6], dtype=dtype) + values = magnitudes[(codes & 7).long()] + values = torch.where((codes & 8) != 0, -values, values) + else: + fp8 = {"e4m3": torch.float8_e4m3fn, "e5m2": torch.float8_e5m2}[format_name] + values = packed.view(fp8).to(dtype) + scales = (scale.to(torch.int32) << 23).view(torch.float32).to(dtype) + scales = scales.repeat_interleave(32, dim=-1) + result = values * scales + return torch.where(scale.repeat_interleave(32, dim=-1) == 255, float("nan"), result) diff --git a/third_party/wafer/examples/bare_matmul.py b/third_party/wafer/examples/bare_matmul.py new file mode 100755 index 00000000..24201510 --- /dev/null +++ b/third_party/wafer/examples/bare_matmul.py @@ -0,0 +1,52 @@ +# this is a benchmark which multiplies square matrices with maximum block size +# to check the performance of tl.dot operation + +import torch +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def bare_matmul(X, Y, Z, M, N, K, BLOCK_SIZE: tl.constexpr): + pid_x = tl.program_id(0) # block row id + pid_y = tl.program_id(1) # block column id + + offs_x = pid_x * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + offs_y = pid_y * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + + x = tl.load(X + offs_x[:, None] * K + offs_y[None, :]) + y = tl.load(Y + offs_x[:, None] * N + offs_y[None, :]) + + z = tl.dot(x, y) + + tl.store(Z + offs_x[:, None] * N + offs_y[None, :], z) + + +# @benchmark.measure() +def bench_matmul(N, provider): + device = 'cpu' + dtype = torch.float32 + a = torch.randint(0, 10, (N, N), dtype=torch.int32).to(dtype) + b = torch.randint(0, 10, (N, N), dtype=torch.int32).to(dtype) + # a = torch.randn((N, N), device=device, dtype=dtype) + # b = torch.randn((N, N), device=device, dtype=dtype) + c = torch.empty((N, N), device=device, dtype=dtype) + if provider == 'torch' or provider == 'test': + c_ref = torch.matmul(a, b) + # print("====cref:",c_ref) + if provider == 'triton' or provider == 'test': + bare_matmul[(1, )](a, b, c, N, N, N, N) + if provider == 'test': + torch.testing.assert_close(c, c_ref, atol=1e-2, rtol=0) + print("======test====") + + +if __name__ == "__main__": + for provider in ['test']: + bench_matmul(16, provider) + # for X in [2**i for i in range(7, 10, 1)]: + # for provider in ['test', 'torch', 'triton']: + # bench_matmul(X, provider) diff --git a/third_party/wafer/examples/bare_matmul_acc.py b/third_party/wafer/examples/bare_matmul_acc.py new file mode 100755 index 00000000..28a2c5f6 --- /dev/null +++ b/third_party/wafer/examples/bare_matmul_acc.py @@ -0,0 +1,46 @@ +# this is a benchmark which multiplies square matrices with maximum block size +# and additional accumulation to check the performance of tl.dot operation + +import torch +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def bare_matmul_acc(X, Y, Z, C, M, N, K, BLOCK_SIZE: tl.constexpr): + pid_x = tl.program_id(0) # block row id + pid_y = tl.program_id(1) # block column id + + offs_x = pid_x * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + offs_y = pid_y * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + + x = tl.load(X + offs_x[:, None] * K + offs_y[None, :]) + y = tl.load(Y + offs_x[:, None] * N + offs_y[None, :]) + c = tl.load(C + offs_x[:, None] * N + offs_y[None, :]) + + z = tl.dot(x, y, c) + + tl.store(Z + offs_x[:, None] * N + offs_y[None, :], z) + + +@benchmark.measure() +def bench_matmul(N, provider): + device = 'cpu' + dtype = torch.float32 + a = torch.randn((N, N), device=device, dtype=dtype) + b = torch.randn((N, N), device=device, dtype=dtype) + c = torch.randn((N, N), device=device, dtype=dtype) + z = torch.empty((N, N), device=device, dtype=dtype) + if provider == 'torch' or provider == 'test': + z_ref = torch.matmul(a, b) + c + if provider == 'triton' or provider == 'test': + bare_matmul_acc[(1, )](a, b, z, c, N, N, N, N) + if provider == 'test': + torch.testing.assert_close(z, z_ref, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for X in [2**i for i in range(7, 10, 1)]: + for provider in ['test', 'torch', 'triton']: + bench_matmul(X, provider) diff --git a/third_party/wafer/examples/bare_matmul_autotune.py b/third_party/wafer/examples/bare_matmul_autotune.py new file mode 100755 index 00000000..740f2629 --- /dev/null +++ b/third_party/wafer/examples/bare_matmul_autotune.py @@ -0,0 +1,69 @@ +# this is a benchmark which multiplies square matrices with maximum block size +# to check the performance of tl.dot operation +import os + +os.environ["TRITON_PRINT_AUTOTUNING"] = "1" + +import torch +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +def prune_invalid_configs(configs, named_args, **kwargs): + """过滤掉会导致 tensor 过大的配置""" + MAX_NUMEL = 1048576 + valid_configs = [] + for config in configs: + block_size = config.kwargs.get("BLOCK_SIZE", 64) + # 检查 2D tensor 大小 + if block_size * block_size <= MAX_NUMEL: + valid_configs.append(config) + return valid_configs if valid_configs else [configs[0]] + + +@triton.autotune( + configs=[ + triton.Config(kwargs={"BLOCK_SIZE": 4096}, num_stages=1, num_warps=32), + triton.Config(kwargs={"BLOCK_SIZE": 2048}, num_stages=1, num_warps=32), + triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_stages=1, num_warps=32), + triton.Config(kwargs={"BLOCK_SIZE": 512}, num_stages=1, num_warps=32), + triton.Config(kwargs={"BLOCK_SIZE": 256}, num_stages=1, num_warps=32), + triton.Config(kwargs={"BLOCK_SIZE": 128}, num_stages=1, num_warps=32), + triton.Config(kwargs={"BLOCK_SIZE": 64}, num_stages=1, num_warps=16), + triton.Config(kwargs={"BLOCK_SIZE": 64}, num_stages=1, num_warps=32), + ], key=["M", "N", "K"], prune_configs_by={"early_config_prune": prune_invalid_configs}, warmup=25, rep=6000) +@triton.jit +def bare_matmul(X, Y, Z, M, N, K, BLOCK_SIZE: tl.constexpr): + pid_x = tl.program_id(0) # block row id + pid_y = tl.program_id(1) # block column id + + offs_x = pid_x * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + offs_y = pid_y * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + + x = tl.load(X + offs_x[:, None] * K + offs_y[None, :]) + y = tl.load(Y + offs_x[:, None] * N + offs_y[None, :]) + + z = tl.dot(x, y) + + tl.store(Z + offs_x[:, None] * N + offs_y[None, :], z) + + +# @benchmark.measure() +def bench_matmul(M, N, K): + device = 'cpu' + dtype = torch.float32 + a = torch.randint(0, 10, (M, K), dtype=torch.int32).to(dtype) + b = torch.randint(0, 10, (K, N), dtype=torch.int32).to(dtype) + c = torch.empty((M, N), device=device, dtype=dtype) + + grid = lambda meta: (triton.cdiv(M, meta['BLOCK_SIZE']), triton.cdiv(N, meta['BLOCK_SIZE'])) + bare_matmul[grid](a, b, c, M, N, K) + + print(bare_matmul.best_config) + + +if __name__ == "__main__": + bench_matmul(16, 16, 16) diff --git a/third_party/wafer/examples/benchmark.py b/third_party/wafer/examples/benchmark.py new file mode 100755 index 00000000..4a2284dd --- /dev/null +++ b/third_party/wafer/examples/benchmark.py @@ -0,0 +1,65 @@ +import time +import numpy as np +from functools import wraps +import triton + +# Unfortunately, we can't use triton.testing.perf_report and triton.testing.do_bench for CPU backend because +# they are very specific to cuda + + +def measure(repeats=20, percentiles=(), timers={'Wall': time.perf_counter, 'CPU': time.process_time}): + """ + Decorator to benchmark a function. + + Parameters: + - repeats (int): The number of times the function should be executed for each set of parameters. + - percentiles (tuple): The percentiles to compute on the execution times (e.g., (50, 90, 99)). + - timers (dict): A dictionary where keys are timer names (e.g., 'Wall', 'CPU') and values are timer functions + that measure elapsed time. By default: + * 'Wall': Uses time.perf_counter for high-resolution wall-clock time. + * 'CPU': Uses time.process_time for CPU time spent by the process. + + Returns: + - A decorated function that prints: + * Average execution time. + * Standard deviation time. + * Minimum and maximum times. + * Computed percentiles for each timer. + """ + + def decorator(func): + + @wraps(func) + def wrapper(*args, **kwargs): + print(f"{func.__name__}{args} {kwargs}, {repeats} times, all results in seconds") + times = {} + for t, _ in timers.items(): + times[t] = [] + + for _ in range(repeats): + starts = {} + for t, f in timers.items(): + starts[t] = f() + + result = func(*args, **kwargs) + + for t, f in timers.items(): + times[t].append(f() - starts[t]) + + for t, _ in timers.items(): + average_time = np.mean(times[t]) + min_time = np.min(times[t]) + max_time = np.max(times[t]) + computed_percentiles = np.percentile(times[t], percentiles) + std_dev_time = np.std(times[t]) + + print(f"{t}: Avg={average_time:.6f}, min={min_time:.6f}, std={std_dev_time:.6f},", end=" ") + for p, value in zip(percentiles, computed_percentiles): + print(f"{p}pp={value:.6f},", end=" ") + print(f"max={max_time:.6f}") + + return result + + return wrapper + + return decorator diff --git a/third_party/wafer/examples/conftest.py b/third_party/wafer/examples/conftest.py new file mode 100755 index 00000000..54979dfe --- /dev/null +++ b/third_party/wafer/examples/conftest.py @@ -0,0 +1,12 @@ +"""Native Wafer examples retain their original deterministic input baseline.""" +import pytest + + +@pytest.fixture(autouse=True) +def require_device(wafer_device): + import numpy as np + import torch + # Preserve initialization formerly supplied by _wafer_harness. + np.random.seed(0) + torch.manual_seed(0) + return wafer_device diff --git a/third_party/wafer/examples/dump_vec_add_ir.sh b/third_party/wafer/examples/dump_vec_add_ir.sh new file mode 100755 index 00000000..ec50a8cf --- /dev/null +++ b/third_party/wafer/examples/dump_vec_add_ir.sh @@ -0,0 +1,47 @@ +#!/bin/bash + +set -euo pipefail + +SCRIPT_DIR=$(cd "$(dirname "$0")" && pwd) +DLC_ROOT=$(cd "$SCRIPT_DIR/../../.." && pwd) + +if [[ -z ${LLVM_BINARY_DIR:-} ]]; then + : "${LLVM_SYSPATH:?Set LLVM_SYSPATH or LLVM_BINARY_DIR before running this script}" + LLVM_BINARY_DIR="$LLVM_SYSPATH/bin" +fi +TRITON_DUMP_PATH=${TRITON_DUMP_PATH:-/tmp/tsm_dump} +DUMP_INDEX=${DUMP_INDEX:-1} +DUMP_DIR="$TRITON_DUMP_PATH/dump$DUMP_INDEX" + +rm -rf "$TRITON_DUMP_PATH" +mkdir -p "$TRITON_DUMP_PATH" + +export DICP_BACKEND=${DICP_BACKEND:-wafer} +export USE_SIM_MODE=${USE_SIM_MODE:-1} +export LLVM_BINARY_DIR +export TRITON_DUMP_PATH +export TRITON_ALWAYS_COMPILE=1 +export MLIR_ENABLE_DUMP=1 + +cd "$SCRIPT_DIR" + +python3 - <<'PY' +import torch +import test_vec_add as v + +x = torch.rand(1024, device="cpu") +y = torch.rand(1024, device="cpu") +v.add(x, y) +PY + +echo "" +echo "Dump directory: $DUMP_DIR" +echo "" +ls -lah "$DUMP_DIR" +echo "" +echo "Common files:" +for f in tt_0.mlir core_0.mlir wafer_0.mlir ll_0.mlir ll_0.ir kernel_0.ll kernel_0.o cmds.txt; do + if [ -e "$DUMP_DIR/$f" ]; then + echo " $DUMP_DIR/$f" + fi +done diff --git a/third_party/wafer/examples/embedding.py b/third_party/wafer/examples/embedding.py new file mode 100755 index 00000000..8f3c2cad --- /dev/null +++ b/third_party/wafer/examples/embedding.py @@ -0,0 +1,89 @@ +import torch +import math + +import triton +import triton.language as tl + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def embedding_kernel( + out_ptr, # pointer to the output + in_ptr, # pointer to the input + weight_ptr, # pointer to the weights + N: tl.constexpr, # number of columns in X + BLOCK_SIZE: tl.constexpr, +): + pid = tl.program_id(0) + out_ptr += pid * N + in_ptr += pid + + mask = tl.arange(0, BLOCK_SIZE) < N + cols = tl.arange(0, BLOCK_SIZE) + + row_idx = tl.load(in_ptr) + weight_ptr += row_idx * N + embedding_weight = tl.load(weight_ptr + cols, mask, other=0.0) + tl.store(out_ptr + cols, embedding_weight, mask) + + +class Embedding(torch.autograd.Function): + + @staticmethod + def forward(ctx, weight, indices, padding_idx=-1, scale_grad_by_freq=False, sparse=False): + + assert not sparse, "Currently do not support sparse format" + + M = math.prod(indices.shape) + N = weight.shape[-1] + + BLOCK_SIZE = triton.next_power_of_2(N) + indices = indices.contiguous() + weight = weight.contiguous() + output = torch.empty((*indices.shape, N), device=indices.device, dtype=weight.dtype) + + embedding_kernel[ + M, + ](output, indices, weight, N, BLOCK_SIZE) + + ctx.M = M + ctx.N = N + ctx.num_weights = weight.shape[0] + ctx.padding_idx = padding_idx + ctx.scale_grad_by_freq = scale_grad_by_freq + ctx.sparse = sparse + ctx.indices = indices + + return output + + +def embedding(weight, indices, padding_idx=-1, scale_grad_by_freq=False, sparse=False): + return Embedding.apply(weight, indices, padding_idx, scale_grad_by_freq, sparse) + + +def test(M=1151, N=8192, dtype=torch.float32): + torch.manual_seed(0) + + weight = torch.rand((M, N), dtype=dtype, device="cpu") + indices = torch.randint(0, M, [M], dtype=torch.int32, device="cpu") + + # pytorch + torch_embedding = torch.nn.Embedding(M, N, _weight=weight) + torch_output = torch_embedding(indices) + + weight = weight.to(DEVICE) + indices = indices.to(DEVICE) + triton_output = embedding(weight, indices) + + triton_output = triton_output.to("cpu") + + print("triton_output:\n", triton_output) + print("torch_output:\n", torch_output) + + # 验证结果一致性 + assert torch.allclose(torch_output, triton_output, atol=1e-6), "verification failure!\n" + print("verification success!\n") + + +test() diff --git a/third_party/wafer/examples/flagtree/test_tle_cumsum.py b/third_party/wafer/examples/flagtree/test_tle_cumsum.py new file mode 100644 index 00000000..ab040872 --- /dev/null +++ b/third_party/wafer/examples/flagtree/test_tle_cumsum.py @@ -0,0 +1,135 @@ +# Adapted from FlagTree 22f4ff0: target name and host reference device only. +# flagtree tle +"""Generic tle.cumsum (tle-lite) integration tests on tsingmicro (txda). + +tle.cumsum is a generic tle-lite primitive; on tsingmicro it lowers through +the shared tle dialect to the hardware scan. Current constraints: +float input only (integer scan unsupported by hardware), forward scan only +(reverse=True unsupported), rank-1 tensors only (shared op constraint). +""" + +import pytest +import torch + +import triton +import triton.experimental.tle.language as tle +import triton.language as tl + + +def _is_txda(): + try: + import torch_txda # noqa: F401 + except ImportError: + return False + target = triton.runtime.driver.active.get_current_target() + return getattr(target, "backend", None) == "wafer" + + +pytestmark = pytest.mark.skipif(not _is_txda(), reason="requires TsingMicro (txda) backend") + + +@triton.jit +def _cumsum_1d(x_ptr, exclusive_ptr, total_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n + x = tl.load(x_ptr + offs, mask=mask, other=0) + exclusive, total = tle.cumsum(x, axis=0) + tl.store(exclusive_ptr + offs, exclusive, mask=mask) + tl.store(total_ptr, total) + + +@triton.jit +def _cumsum_1d_reverse(x_ptr, exclusive_ptr, total_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + x = tl.load(x_ptr + offs) + exclusive, total = tle.cumsum(x, axis=0, reverse=True) + tl.store(exclusive_ptr + offs, exclusive, mask=offs < n) + tl.store(total_ptr, total) + + +@triton.jit +def _cumsum_2d(x_ptr, out_ptr, M: tl.constexpr, N: tl.constexpr): + m = tl.arange(0, M) + n = tl.arange(0, N) + x = tl.load(x_ptr + m[:, None] * N + n[None, :]) + exclusive, total = tle.cumsum(x, axis=1) + tl.store(out_ptr + m[:, None] * N + n[None, :], exclusive) + + +def _exclusive_expected(x, dtype): + cs = torch.cumsum(x, dim=0, dtype=dtype) + zero = torch.zeros(1, device="cpu", dtype=dtype) + return torch.cat([zero, cs[:-1]]) + + +def test_cumsum_1d_masked(): + torch.manual_seed(42) + n, block = 100, 128 + x = torch.randn(n, device="cpu", dtype=torch.float32) + + exclusive = torch.zeros(block, device="cpu", dtype=torch.float32) + total = torch.zeros(1, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + exclusive_txda = exclusive.to("txda") + total_txda = total.to("txda") + _cumsum_1d[(1, )](x_txda, exclusive_txda, total_txda, n, BLOCK=block) + with torch.no_grad(): + exclusive.copy_(exclusive_txda.cpu()) + total.copy_(total_txda.cpu()) + + expected = _exclusive_expected(x, torch.float32) + torch.testing.assert_close(exclusive[:n], expected) + torch.testing.assert_close(total[0], x.sum(dim=0, dtype=torch.float32)) + + +def test_cumsum_1d_full_block(): + torch.manual_seed(43) + n = 64 + x = torch.randn(n, device="cpu", dtype=torch.float32) + + exclusive = torch.zeros(n, device="cpu", dtype=torch.float32) + total = torch.zeros(1, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + exclusive_txda = exclusive.to("txda") + total_txda = total.to("txda") + _cumsum_1d[(1, )](x_txda, exclusive_txda, total_txda, n, BLOCK=n) + with torch.no_grad(): + exclusive.copy_(exclusive_txda.cpu()) + total.copy_(total_txda.cpu()) + + expected = _exclusive_expected(x, torch.float32) + torch.testing.assert_close(exclusive, expected, atol=2e-6, rtol=1e-5) + torch.testing.assert_close(total[0], x.sum(dim=0), atol=2e-6, rtol=1e-5) + + +def test_cumsum_2d_unsupported(): + """The shared tle.exclusive_cumsum op only accepts rank-1 tensors.""" + torch.manual_seed(46) + m, n = 8, 32 + x = torch.randn(m, n, device="cpu", dtype=torch.float32) + out = torch.zeros(m, n, device="cpu", dtype=torch.float32) + with pytest.raises(Exception): + x_txda = x.to("txda") + out_txda = out.to("txda") + _cumsum_2d[(1, )](x_txda, out_txda, M=m, N=n) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + +def test_cumsum_reverse_unsupported(): + torch.manual_seed(45) + x = torch.randn(64, device="cpu", dtype=torch.float32) + exclusive = torch.zeros(64, device="cpu", dtype=torch.float32) + total = torch.zeros(1, device="cpu", dtype=torch.float32) + with pytest.raises(Exception): + x_txda = x.to("txda") + exclusive_txda = exclusive.to("txda") + total_txda = total.to("txda") + _cumsum_1d_reverse[(1, )](x_txda, exclusive_txda, total_txda, 64, BLOCK=64) + with torch.no_grad(): + exclusive.copy_(exclusive_txda.cpu()) + total.copy_(total_txda.cpu()) + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) diff --git a/third_party/wafer/examples/flagtree/test_tle_dsa_arith.py b/third_party/wafer/examples/flagtree/test_tle_dsa_arith.py new file mode 100644 index 00000000..2d3c2fd8 --- /dev/null +++ b/third_party/wafer/examples/flagtree/test_tle_dsa_arith.py @@ -0,0 +1,102 @@ +# Adapted from FlagTree 22f4ff0: target name and host reference device only. +import pytest +import torch + +import triton +import triton.language as tl +import triton.experimental.tle.language as tle + + +def _is_txda(): + try: + import torch_txda # noqa: F401 + except ImportError: + return False + target = triton.runtime.driver.active.get_current_target() + return getattr(target, "backend", None) == "wafer" + + +pytestmark = pytest.mark.skipif(not _is_txda(), reason="TLE DSA tests require TsingMicro (txda) backend") + + +@triton.jit +def dsa_arith_kernel(x_ptr, y_ptr, out_ptr, M, N, P, Q, BM: tl.constexpr, BN: tl.constexpr, BP: tl.constexpr, + BQ: tl.constexpr, OP: tl.constexpr): + """Four-dimensional three-operand buffer arithmetic (tiled).""" + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + pid_p = tl.program_id(2) + offs_m = pid_m * BM + tl.arange(0, BM) + offs_n = pid_n * BN + tl.arange(0, BN) + offs_p = pid_p * BP + tl.arange(0, BP) + offs_q = tl.arange(0, BQ) + idx = (offs_m[:, None, None, None] * N * P * Q + offs_n[None, :, None, None] * P * Q + + offs_p[None, None, :, None] * Q + offs_q[None, None, None, :]) + mask = (offs_m[:, None, None, None] < M) & (offs_n[None, :, None, None] < N) & \ + (offs_p[None, None, :, None] < P) & (offs_q[None, None, None, :] < Q) + + lhs = tl.load(x_ptr + idx, mask=mask) + rhs = tl.load(y_ptr + idx, mask=mask) + + lhs_buf = tle.dsa.to_buffer(lhs, tle.dsa.spm) + rhs_buf = tle.dsa.to_buffer(rhs, tle.dsa.spm) + out_buf = tle.dsa.alloc((BM, BN, BP, BQ), tl.float32) + + if OP == "add": + tle.dsa.add(lhs_buf, rhs_buf, out_buf) + elif OP == "sub": + tle.dsa.sub(lhs_buf, rhs_buf, out_buf) + elif OP == "mul": + tle.dsa.mul(lhs_buf, rhs_buf, out_buf) + elif OP == "max": + tle.dsa.max(lhs_buf, rhs_buf, out_buf) + elif OP == "min": + tle.dsa.min(lhs_buf, rhs_buf, out_buf) + elif OP == "div": + tle.dsa.div(lhs_buf, rhs_buf, out_buf) + + result = tle.dsa.to_tensor(out_buf) + tl.store(out_ptr + idx, result, mask=mask) + + +class TestTLEDsaArith: + """Three-operand DSA buffer arithmetic (add/sub/mul/max/min/div).""" + + @pytest.mark.parametrize( + "op,ref", + [ + ("add", lambda a, b: a + b), + ("sub", lambda a, b: a - b), + ("mul", lambda a, b: a * b), + ("max", lambda a, b: torch.maximum(a, b)), + ("min", lambda a, b: torch.minimum(a, b)), + ("div", lambda a, b: a / b), + ], + ) + @pytest.mark.parametrize( + "shape,block", + [((17, 13, 9, 8), (16, 8, 8, 8)), # tails on m/n/p; q fully covered (3-axis grid) + ], + ) + def test_arith(self, op, ref, shape, block): + torch.manual_seed(42) + m, n, p, q = shape + bm, bn, bp, bq = block + a = torch.randn(*shape, device="cpu", dtype=torch.float32) + b = torch.randn(*shape, device="cpu", dtype=torch.float32) + out = torch.empty_like(a) + + grid = (triton.cdiv(m, bm), triton.cdiv(n, bn), triton.cdiv(p, bp)) + a_txda = a.to("txda") + b_txda = b.to("txda") + out_txda = out.to("txda") + dsa_arith_kernel[grid](a_txda, b_txda, out_txda, m, n, p, q, BM=bm, BN=bn, BP=bp, BQ=bq, num_ctas=1, OP=op) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + expected = ref(a, b) + torch.testing.assert_close(out, expected, atol=1e-4, rtol=1e-4) + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) diff --git a/third_party/wafer/examples/flagtree/test_tle_dsa_bridge.py b/third_party/wafer/examples/flagtree/test_tle_dsa_bridge.py new file mode 100644 index 00000000..c75f5e6f --- /dev/null +++ b/third_party/wafer/examples/flagtree/test_tle_dsa_bridge.py @@ -0,0 +1,84 @@ +# Adapted from FlagTree 22f4ff0: target name and host reference device only. +import pytest +import torch + +import triton +import triton.language as tl +import triton.experimental.tle.language as tle + + +def _is_txda(): + try: + import torch_txda # noqa: F401 + except ImportError: + return False + target = triton.runtime.driver.active.get_current_target() + return getattr(target, "backend", None) == "wafer" + + +pytestmark = pytest.mark.skipif(not _is_txda(), reason="TLE DSA tests require TsingMicro (txda) backend") + + +@triton.jit +def to_buffer_to_tensor_kernel(x_ptr, y_ptr, out_ptr, M, N, P, Q, BM: tl.constexpr, BN: tl.constexpr, BP: tl.constexpr, + BQ: tl.constexpr): + """4D round-trip: to_buffer -> to_tensor -> compute -> to_buffer -> store.""" + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + pid_p = tl.program_id(2) + offs_m = pid_m * BM + tl.arange(0, BM) + offs_n = pid_n * BN + tl.arange(0, BN) + offs_p = pid_p * BP + tl.arange(0, BP) + offs_q = tl.arange(0, BQ) + idx = (offs_m[:, None, None, None] * N * P * Q + offs_n[None, :, None, None] * P * Q + + offs_p[None, None, :, None] * Q + offs_q[None, None, None, :]) + mask = (offs_m[:, None, None, None] < M) & (offs_n[None, :, None, None] < N) & \ + (offs_p[None, None, :, None] < P) & (offs_q[None, None, None, :] < Q) + + x = tl.load(x_ptr + idx, mask=mask) + y = tl.load(y_ptr + idx, mask=mask) + + # tl.tensor -> fresh SPM buffer. + buf_x = tle.dsa.to_buffer(x, tle.dsa.spm) + buf_y = tle.dsa.to_buffer(y, tle.dsa.spm) + + # SPM buffer -> zero-copy tl.tensor view, then compute. + tx = tle.dsa.to_tensor(buf_x) + ty = tle.dsa.to_tensor(buf_y) + z = tx * ty + + # Result into a fresh SPM buffer, then read back out. + buf_z = tle.dsa.to_buffer(z, tle.dsa.spm) + zz = tle.dsa.to_tensor(buf_z) + tl.store(out_ptr + idx, zz, mask=mask) + + +class TestTLEDsaBridge: + """to_tensor / to_buffer bridge between tl.tensor and SPM buffers.""" + + @pytest.mark.parametrize( + "shape,block", + [((17, 13, 9, 8), (16, 8, 8, 8)), # tails on m/n/p; q fully covered (3-axis grid) + ], + ) + def test_bridge(self, shape, block): + torch.manual_seed(42) + m, n, p, q = shape + bm, bn, bp, bq = block + a = torch.randn(*shape, device="cpu", dtype=torch.float32) + b = torch.randn(*shape, device="cpu", dtype=torch.float32) + out = torch.empty_like(a) + + grid = (triton.cdiv(m, bm), triton.cdiv(n, bn), triton.cdiv(p, bp)) + a_txda = a.to("txda") + b_txda = b.to("txda") + out_txda = out.to("txda") + to_buffer_to_tensor_kernel[grid](a_txda, b_txda, out_txda, m, n, p, q, BM=bm, BN=bn, BP=bp, BQ=bq, num_ctas=1) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + torch.testing.assert_close(out, a * b, atol=1e-4, rtol=1e-4) + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) diff --git a/third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py b/third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py new file mode 100644 index 00000000..cf58b472 --- /dev/null +++ b/third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py @@ -0,0 +1,234 @@ +# Adapted from FlagTree 22f4ff0: target name and host reference device only. +# flagtree tle +""" +TLE End-to-End Integration Tests + +Tests complete workflow of TLE pipeline functionality in real GPU environment: +- Memory allocation (tle.gpu.alloc) +- Copying between GM and shared memory (tle.gpu.copy) +- Shared-memory pointer materialization (tle.gpu.local_ptr + tl.load/store) +- Pipeline iterator (tle.gpu.pipeline) +- Integration with Triton JIT +""" + +import pytest +import torch + +import triton +import triton.experimental.tle.language as tle +import triton.language as tl + + +def _is_txda(): + try: + import torch_txda # noqa: F401 + except ImportError: + return False + target = triton.runtime.driver.active.get_current_target() + return getattr(target, "backend", None) == "wafer" + + +pytestmark = pytest.mark.skipif(not _is_txda(), reason="TLE DSA tests require TsingMicro (txda) backend") + + +@triton.jit +def elementwise_add_kernel( + a_ptr, + b_ptr, + c_ptr, + xnumel, + ynumel, + xstride_a, + ystride_a, + xstride_b, + ystride_b, + xstride_c, + ystride_c, + XBLOCK: tl.constexpr, + YBLOCK: tl.constexpr, +): + """ + Element-wise addition kernel using TLE pipeline + + This kernel demonstrates the complete TLE workflow: + 1. Allocate shared memory buffers + 2. Use pipeline for for-loop iteration + 3. copy data to shared memory + 4. Load data from shared memory for computation + 5. Store results back to global memory + """ + pid = tl.program_id(0) + + # Calculate row offset for current program + xoffs = pid * XBLOCK + tl.arange(0, XBLOCK) + + # Calculate global memory pointers + a_ptrs = a_ptr + xstride_a * xoffs[:, None] + b_ptrs = b_ptr + xstride_b * xoffs[:, None] + c_ptrs = c_ptr + xstride_c * xoffs[:, None] + + # Allocate shared memory buffers + a_smem = tle.dsa.alloc([XBLOCK, YBLOCK], dtype=tl.float32) + b_smem = tle.dsa.alloc([XBLOCK, YBLOCK], dtype=tl.float32) + row_ids = tl.arange(0, XBLOCK)[:, None] + col_ids = tl.arange(0, YBLOCK)[None, :] + row_ids = tl.broadcast_to(row_ids, (XBLOCK, YBLOCK)) + col_ids = tl.broadcast_to(col_ids, (XBLOCK, YBLOCK)) + a_smem_ptrs = tle.dsa.local_ptr(a_smem, (row_ids, col_ids)) + b_smem_ptrs = tle.dsa.local_ptr(b_smem, (row_ids, col_ids)) + + # Use TLE pipeline for block-wise processing + # for yoff in range(0, ynumel, YBLOCK): + for yoff in tle.dsa.pipeline(0, ynumel, YBLOCK, num_stages=2): + # Calculate column offset for current block + yoffs = tl.arange(0, YBLOCK) + yoff + mask = (xoffs < xnumel)[:, None] & (yoffs < ynumel)[None, :] + + # copy data to shared memory + tle.dsa.copy(a_ptrs + ystride_a * yoffs[None, :], a_smem, [XBLOCK, YBLOCK]) + tle.dsa.copy(b_ptrs + ystride_b * yoffs[None, :], b_smem, [XBLOCK, YBLOCK]) + + # Load data from shared memory + aval = tl.load(a_smem_ptrs) + bval = tl.load(b_smem_ptrs) + + # Perform computation + c_val = aval + bval + + # Store results + tl.store(c_ptrs + ystride_c * yoffs[None, :], c_val, mask=mask) + + +def elementwise_add(A, B, C, XBLOCK=32, YBLOCK=64): + """ + Wrapper function to execute element-wise addition using TLE pipeline + + Args: + A: Input tensor A (CUDA tensor) + B: Input tensor B (CUDA tensor) + C: Output tensor C (CUDA tensor) + XBLOCK: Block size for X dimension + YBLOCK: Block size for Y dimension + """ + assert A.shape == B.shape == C.shape, "Input and output tensor shapes must match" + xnumel, ynumel = A.shape + grid = (triton.cdiv(xnumel, XBLOCK), ) + + A_txda = A.to("txda") + B_txda = B.to("txda") + C_txda = C.to("txda") + compiled = elementwise_add_kernel[grid](A_txda, B_txda, C_txda, xnumel, ynumel, *A_txda.stride(), *B_txda.stride(), *C_txda.stride(), XBLOCK, YBLOCK) + with torch.no_grad(): + C.copy_(C_txda.cpu()) + return compiled + + +class TestTLEPipelineEndToEnd: + """TLE Pipeline End-to-End Integration Tests""" + + def test_elementwise_add_basic(self): + """Test basic element-wise addition functionality""" + torch.manual_seed(42) # Ensure reproducibility + + xnumel, ynumel = 512, 512 + XBLOCK, YBLOCK = 64, 64 + + # Create test data + a = torch.randn(xnumel, ynumel, device="cpu", dtype=torch.float32) + b = torch.randn(xnumel, ynumel, device="cpu", dtype=torch.float32) + c = torch.empty_like(a, device="cpu", dtype=torch.float32) + + # Execute TLE pipeline computation + elementwise_add(a, b, c, XBLOCK, YBLOCK) + + # Verify results + expected = a + b + torch.testing.assert_close(c, expected, atol=1e-5, rtol=1e-5) + + def test_elementwise_add_different_sizes(self): + """Test different tensor sizes""" + torch.manual_seed(123) + + test_cases = [ + (256, 256, 32, 32), + (1024, 1024, 64, 64), + (2048, 512, 128, 32), + (512, 2048, 32, 128), + ] + + for xnumel, ynumel, XBLOCK, YBLOCK in test_cases: + # Create test data + a = torch.randn(xnumel, ynumel, device="cpu", dtype=torch.float32) + b = torch.randn(xnumel, ynumel, device="cpu", dtype=torch.float32) + c = torch.empty_like(a, device="cpu", dtype=torch.float32) + + # Execute computation + elementwise_add(a, b, c, XBLOCK, YBLOCK) + + # Verify results + expected = a + b + torch.testing.assert_close(c, expected, atol=1e-5, rtol=1e-5) + assert c.shape == expected.shape, f"Shape mismatch: {xnumel}x{ynumel}, block size: {XBLOCK}x{YBLOCK}" + + def test_elementwise_add_different_dtypes(self): + """Test different data types""" + torch.manual_seed(456) + + dtypes = [torch.float32] # Skip float16 due to type conversion issues in TLE + xnumel, ynumel = 512, 512 + XBLOCK, YBLOCK = 64, 64 + + for dtype in dtypes: + # Create test data + a = torch.randn(xnumel, ynumel, device="cpu", dtype=dtype) + b = torch.randn(xnumel, ynumel, device="cpu", dtype=dtype) + c = torch.empty_like(a, device="cpu", dtype=dtype) + + # Execute computation + elementwise_add(a, b, c, XBLOCK, YBLOCK) + + # Verify results + expected = a + b + torch.testing.assert_close(c, expected, atol=1e-3, rtol=1e-3) + assert c.shape == expected.shape, f"Shape mismatch for data type {dtype}" + + def test_elementwise_add_edge_cases(self): + """Test edge cases""" + torch.manual_seed(789) + + # Test minimum size + a = torch.randn(1, 1, device="cpu", dtype=torch.float32) + b = torch.randn(1, 1, device="cpu", dtype=torch.float32) + c = torch.empty_like(a, device="cpu", dtype=torch.float32) + + elementwise_add(a, b, c, 1, 1) + torch.testing.assert_close(c, a + b, atol=1e-5, rtol=1e-5) + + # Test non-square tensors + a = torch.randn(128, 1024, device="cpu", dtype=torch.float32) + b = torch.randn(128, 1024, device="cpu", dtype=torch.float32) + c = torch.empty_like(a, device="cpu", dtype=torch.float32) + + elementwise_add(a, b, c, 32, 128) + torch.testing.assert_close(c, a + b, atol=1e-5, rtol=1e-5) + + def test_tle_module_import(self): + """Test TLE module import (no GPU required)""" + # Verify all necessary functions and types can be imported + assert hasattr(tle, "dsa") + assert hasattr(tle.dsa, "alloc") + assert hasattr(tle.dsa, "copy") + assert hasattr(tle.dsa, "local_ptr") + assert hasattr(tle.dsa, "pipeline") + assert hasattr(tle.dsa, "scope") + assert hasattr(tle.dsa, "buffered_tensor") + + # Verify functions have docstrings + assert tle.dsa.alloc.__doc__ is not None + assert tle.dsa.copy.__doc__ is not None + assert tle.dsa.local_ptr.__doc__ is not None + assert tle.dsa.pipeline.__doc__ is not None + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) diff --git a/third_party/wafer/examples/flagtree/test_tle_dsa_rand.py b/third_party/wafer/examples/flagtree/test_tle_dsa_rand.py new file mode 100644 index 00000000..7e6272f9 --- /dev/null +++ b/third_party/wafer/examples/flagtree/test_tle_dsa_rand.py @@ -0,0 +1,155 @@ +# Adapted from FlagTree 22f4ff0: target/vendor name and host reference device only. +# flagtree tle +"""TsingMicro vendor DSA rand-family tests (tle.dsa.wafer). + +Covers ``randgen`` (raw xorshift128+ i64 stream), ``rand`` (Uniform(0,1)) +and ``randn`` (Normal(0,1) via Box-Muller) on the Wafer hardware TRNG. + +Constraints exercised by the API: + - ``randgen``: n_out multiple of 16, seeds are ``[16]`` i64 blocks + - ``rand`` / ``randn``: n_out multiple of 32 +""" + +import pytest +import torch + +import triton +import triton.language as tl +import triton.experimental.tle.language as tle + + +def _is_txda(): + try: + import torch_txda # noqa: F401 + except ImportError: + return False + target = triton.runtime.driver.active.get_current_target() + return getattr(target, "backend", None) == "wafer" + + +pytestmark = pytest.mark.skipif(not _is_txda(), reason="requires TsingMicro (txda) backend") + + +@triton.jit +def _randgen_kernel(seed_val, out_ptr, s0_ptr, s1_ptr, N: tl.constexpr): + s0 = tl.arange(0, 16).to(tl.int64) * 0x2545F4914F6CDD1D + seed_val + s1 = tl.arange(0, 16).to(tl.int64) * 0x1E3779B97F4A7C15 + seed_val + 1 + out, s0o, s1o = tle.dsa.wafer.randgen(s0, s1, N) + offs = tl.arange(0, N) + tl.store(out_ptr + offs, out) + tl.store(s0_ptr + tl.arange(0, 16), s0o) + tl.store(s1_ptr + tl.arange(0, 16), s1o) + + +@triton.jit +def _rand_kernel(seed_val, out_ptr, N: tl.constexpr): + s0 = tl.arange(0, 16).to(tl.int64) * 0x2545F4914F6CDD1D + seed_val + s1 = tl.arange(0, 16).to(tl.int64) * 0x1E3779B97F4A7C15 + seed_val + 1 + u, s0o, s1o = tle.dsa.wafer.rand(s0, s1, N) + tl.store(out_ptr + tl.arange(0, N), u) + + +@triton.jit +def _randn_kernel(seed_val, out_ptr, N: tl.constexpr): + s0 = tl.arange(0, 16).to(tl.int64) * 0x2545F4914F6CDD1D + seed_val + s1 = tl.arange(0, 16).to(tl.int64) * 0x1E3779B97F4A7C15 + seed_val + 1 + n, s0o, s1o = tle.dsa.wafer.randn(s0, s1, N) + tl.store(out_ptr + tl.arange(0, N), n) + + +def test_randgen_shape_and_determinism(): + n = 32 + out = torch.empty(n, device="cpu", dtype=torch.int64) + s0o = torch.zeros(16, device="cpu", dtype=torch.int64) + s1o = torch.zeros(16, device="cpu", dtype=torch.int64) + out_txda = out.to("txda") + s0o_txda = s0o.to("txda") + s1o_txda = s1o.to("txda") + _randgen_kernel[(1, )](42, out_txda, s0o_txda, s1o_txda, N=n) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + s0o.copy_(s0o_txda.cpu()) + s1o.copy_(s1o_txda.cpu()) + + # determinism: same seed -> same stream + out2 = torch.empty_like(out) + out2_txda = out2.to("txda") + s0o_txda = s0o.to("txda") + s1o_txda = s1o.to("txda") + _randgen_kernel[(1, )](42, out2_txda, s0o_txda, s1o_txda, N=n) + with torch.no_grad(): + out2.copy_(out2_txda.cpu()) + s0o.copy_(s0o_txda.cpu()) + s1o.copy_(s1o_txda.cpu()) + torch.testing.assert_close(out, out2) + + # stream advances: different seed -> different values + out3 = torch.empty_like(out) + out3_txda = out3.to("txda") + s0o_txda = s0o.to("txda") + s1o_txda = s1o.to("txda") + _randgen_kernel[(1, )](43, out3_txda, s0o_txda, s1o_txda, N=n) + with torch.no_grad(): + out3.copy_(out3_txda.cpu()) + s0o.copy_(s0o_txda.cpu()) + s1o.copy_(s1o_txda.cpu()) + assert not torch.equal(out, out3) + + +def test_randgen_invalid_n_out(): + out = torch.empty(8, device="cpu", dtype=torch.int64) + s0o = torch.zeros(16, device="cpu", dtype=torch.int64) + s1o = torch.zeros(16, device="cpu", dtype=torch.int64) + with pytest.raises(Exception): + out_txda = out.to("txda") + s0o_txda = s0o.to("txda") + s1o_txda = s1o.to("txda") + _randgen_kernel[(1, )](42, out_txda, s0o_txda, s1o_txda, N=8) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + s0o.copy_(s0o_txda.cpu()) + s1o.copy_(s1o_txda.cpu()) # not a multiple of 16 + + +def test_rand_uniform_stats(): + n = 16384 + out = torch.empty(n, device="cpu", dtype=torch.float32) + out_txda = out.to("txda") + _rand_kernel[(1, )](7, out_txda, N=n) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + assert out.min().item() >= 0.0 and out.max().item() < 1.0 + mean, std = out.mean().item(), out.std().item() + # Uniform(0,1): mean 0.5, std sqrt(1/12) ~= 0.2887; 16K samples of a + # single fixed-seed hardware stream, tolerance ~3 sigma. + assert abs(mean - 0.5) < 0.03, f"uniform mean off: {mean}" + assert abs(std - 0.2887) < 0.01, f"uniform std off: {std}" + + +def test_randn_normal_stats(): + n = 16384 + out = torch.empty(n, device="cpu", dtype=torch.float32) + out_txda = out.to("txda") + _randn_kernel[(1, )](11, out_txda, N=n) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + mean, std = out.mean().item(), out.std().item() + assert abs(mean) < 0.05, f"normal mean off: {mean}" + assert abs(std - 1.0) < 0.05, f"normal std off: {std}" + # heavier tails than uniform + assert out.abs().max().item() > 3.0 + + +def test_rand_invalid_n_out(): + out = torch.empty(16, device="cpu", dtype=torch.float32) + with pytest.raises(Exception): + out_txda = out.to("txda") + _rand_kernel[(1, )](7, out_txda, N=16) + with torch.no_grad(): + out.copy_(out_txda.cpu()) # not a multiple of 32 + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) diff --git a/third_party/wafer/examples/flagtree/test_tle_dsa_slice.py b/third_party/wafer/examples/flagtree/test_tle_dsa_slice.py new file mode 100644 index 00000000..17e19c8b --- /dev/null +++ b/third_party/wafer/examples/flagtree/test_tle_dsa_slice.py @@ -0,0 +1,338 @@ +# Adapted from FlagTree 22f4ff0: target name and host reference device only. +# flagtree tle +"""TLE DSA slice family tests on tsingmicro (txda). + +Covers `tle.dsa.extract_slice`/`insert_slice` (element-level strided slicing) +and the generic tle-lite `tle.extract_tile`/`tle.insert_tile` grid-coordinate +forms, which lower through the shared tle dialect and resolve to the same +dsa.extract_slice / dsa.insert_slice IR ops on this backend. +""" + +import pytest +import torch + +import triton +import triton.experimental.tle.language as tle +import triton.language as tl + + +def _is_txda(): + try: + import torch_txda # noqa: F401 + except ImportError: + return False + target = triton.runtime.driver.active.get_current_target() + return getattr(target, "backend", None) == "wafer" + + +pytestmark = pytest.mark.skipif(not _is_txda(), reason="TLE DSA tests require TsingMicro (txda) backend") + + +@triton.jit +def _idx(shape: tl.constexpr): + return (tl.arange(0, shape[0])[:, None, None, None] * shape[1] * shape[2] * shape[3] + + tl.arange(0, shape[1])[None, :, None, None] * shape[2] * shape[3] + + tl.arange(0, shape[2])[None, None, :, None] * shape[3] + tl.arange(0, shape[3])[None, None, None, :]) + + +TILE_SRC = (32, 32, 32, 16) # tile-test source shape +TILE = (16, 16, 16, 8) # tile shape -> grid [2, 2, 2, 2] +SLICE = (16, 16, 16, 16) # slice-test source shape + +# -------------------------------- tile kernels ------------------------------- + + +@triton.jit +def extract_scalar(x_ptr, out_ptr, shape: tl.constexpr, tile_shape: tl.constexpr, LIN: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + tile = tle.extract_tile(x, index=LIN, tile_shape=tile_shape) + tl.store(out_ptr + _idx(tile_shape), tile) + + +@triton.jit +def extract_dyn_multi(x_ptr, out_ptr, shape: tl.constexpr, tile_shape: tl.constexpr, I0: tl.constexpr, I1: tl.constexpr, + I2: tl.constexpr, I3: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + r = tl.full((), I0, tl.int32) + c = tl.full((), I1, tl.int32) + tile = tle.extract_tile(x, index=[r, c, I2, I3], tile_shape=tile_shape) + tl.store(out_ptr + _idx(tile_shape), tile) + + +@triton.jit +def insert_multi(x_ptr, out_ptr, shape: tl.constexpr, tile_shape: tl.constexpr, SI0: tl.constexpr, SI1: tl.constexpr, + SI2: tl.constexpr, SI3: tl.constexpr, DI0: tl.constexpr, DI1: tl.constexpr, DI2: tl.constexpr, + DI3: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + tile = tle.extract_tile(x, index=[SI0, SI1, SI2, SI3], tile_shape=tile_shape) + tile = tile + 1.0 + y = tle.insert_tile(x, tile, index=[DI0, DI1, DI2, DI3]) + tl.store(out_ptr + _idx(shape), y) + + +@triton.jit +def insert_oop(x_ptr, out_ptr, shape: tl.constexpr, tile_shape: tl.constexpr, DI0: tl.constexpr, DI1: tl.constexpr, + DI2: tl.constexpr, DI3: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + tile = tle.extract_tile(x, index=[0, 0, 0, 0], tile_shape=tile_shape) + t2 = tile + 1.0 + s = tl.sum(tile) # forces out-of-place mk.addvs: original tile still needed + y = tle.insert_tile(x, t2, index=[DI0, DI1, DI2, DI3]) + y = y + s + tl.store(out_ptr + _idx(shape), y) + + +@triton.jit +def insert_dyn_scalar(x_ptr, idx_src_ptr, idx_dst_ptr, out_ptr, shape: tl.constexpr, tile_shape: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + idx_src = tl.load(idx_src_ptr) + idx_dst = tl.load(idx_dst_ptr) + tile = tle.extract_tile(x, index=idx_src, tile_shape=tile_shape) + tile = tile + 1.0 + y = tle.insert_tile(x, tile, index=idx_dst) + tl.store(out_ptr + _idx(shape), y) + + +# -------------------------------- slice kernels ------------------------------ + + +@triton.jit +def extract_static(x_ptr, out_ptr, shape: tl.constexpr, offsets: tl.constexpr, sizes: tl.constexpr, + strides: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + sub = tle.dsa.extract_slice(x, offsets=offsets, sizes=sizes, strides=strides) + tl.store(out_ptr + _idx(sizes), sub) + + +@triton.jit +def extract_dyn(x_ptr, o0_ptr, o1_ptr, out_ptr, shape: tl.constexpr, sizes: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + o0 = tl.load(o0_ptr) + o1 = tl.load(o1_ptr) + sub = tle.dsa.extract_slice(x, offsets=(o0, o1, 0, 0), sizes=sizes, strides=(1, 1, 1, 1)) + tl.store(out_ptr + _idx(sizes), sub) + + +@triton.jit +def extract_mixed(x_ptr, o0_ptr, out_ptr, shape: tl.constexpr, sizes: tl.constexpr, O1: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + o0 = tl.load(o0_ptr) + sub = tle.dsa.extract_slice(x, offsets=(o0, O1, 0, 0), sizes=sizes, strides=(1, 1, 1, 1)) + tl.store(out_ptr + _idx(sizes), sub) + + +@triton.jit +def insert_default(x_ptr, out_ptr, shape: tl.constexpr, offsets: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + tile = tle.dsa.extract_slice(x, offsets=(0, 0, 0, 0), sizes=(4, 4, 4, 4), strides=(1, 1, 1, 1)) + tile = tile + 1.0 + y = tle.dsa.insert_slice(x, tile, offsets=offsets) + tl.store(out_ptr + _idx(shape), y) + + +@triton.jit +def insert_strided(x_ptr, out_ptr, shape: tl.constexpr, offsets: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + tile = tle.dsa.extract_slice(x, offsets=(0, 0, 0, 0), sizes=(4, 4, 4, 4), strides=(1, 1, 1, 1)) + tile = tile + 1.0 + y = tle.dsa.insert_slice(x, tile, offsets=offsets, sizes=(4, 4, 4, 4), strides=(2, 2, 2, 2)) + tl.store(out_ptr + _idx(shape), y) + + +@triton.jit +def member_roundtrip(x_ptr, out_ptr, shape: tl.constexpr, offsets: tl.constexpr): + x = tl.load(x_ptr + _idx(shape)) + sub = x.extract_slice(offsets=(0, 0, 0, 0), sizes=(8, 8, 8, 8), strides=(1, 1, 1, 1)) + sub = sub + 1.0 + y = x.insert_slice(sub, offsets=offsets) + tl.store(out_ptr + _idx(shape), y) + + +class TestTile: + + def test_extract_scalar(self): + torch.manual_seed(42) + x = torch.randn(*TILE_SRC, device="cpu", dtype=torch.float32) + + # LIN=15 == linear id of [1,1,1,1] + out = torch.zeros(*TILE, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + extract_scalar[(1, )](x_txda, out_txda, TILE_SRC, TILE, 15) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + torch.testing.assert_close(out, x[16:32, 16:32, 16:32, 8:16]) + + def test_extract_dyn_multi(self): + torch.manual_seed(42) + x = torch.randn(*TILE_SRC, device="cpu", dtype=torch.float32) + + # mixed: dims 0,1 dynamic, dims 2,3 static + out = torch.zeros(*TILE, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + extract_dyn_multi[(1, )](x_txda, out_txda, TILE_SRC, TILE, 1, 1, 0, 0) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + torch.testing.assert_close(out, x[16:32, 16:32, 0:16, 0:8]) + + def test_insert_multi(self): + torch.manual_seed(44) + shape, tile = (16, 16, 16, 16), (8, 8, 8, 8) + x = torch.randn(*shape, device="cpu", dtype=torch.float32) + + # identity: extract [0,0,0,0], +1.0, insert back at [0,0,0,0] + out = torch.zeros(*shape, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + insert_multi[(1, )](x_txda, out_txda, shape, tile, 0, 0, 0, 0, 0, 0, 0, 0) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + expected = x.clone() + expected[0:8, 0:8, 0:8, 0:8] += 1.0 + torch.testing.assert_close(out, expected) + + # relocate: extract [1,1,1,1], +1.0, insert at [0,0,0,0] + out = torch.zeros(*shape, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + insert_multi[(1, )](x_txda, out_txda, shape, tile, 1, 1, 1, 1, 0, 0, 0, 0) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + expected = x.clone() + expected[0:8, 0:8, 0:8, 0:8] = x[8:16, 8:16, 8:16, 8:16] + 1.0 + torch.testing.assert_close(out, expected) + + def test_insert_oop(self): + torch.manual_seed(47) + shape, tile = (16, 16, 16, 16), (8, 8, 8, 8) + x2 = torch.randn(*shape, device="cpu", dtype=torch.float32) + out = torch.zeros(*shape, device="cpu", dtype=torch.float32) + x2_txda = x2.to("txda") + out_txda = out.to("txda") + insert_oop[(1, )](x2_txda, out_txda, shape, tile, 1, 1, 1, 1) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + src = x2[0:8, 0:8, 0:8, 0:8] + expected = x2.clone() + expected[8:16, 8:16, 8:16, 8:16] = src + 1.0 + expected = expected + src.sum() + torch.testing.assert_close(out, expected) + + def test_insert_dyn_scalar(self): + torch.manual_seed(50) + shape, tile = (16, 16, 16, 16), (8, 8, 8, 8) + + # Runtime indices (5=[0,1,0,1], 10=[1,0,1,0]) exercise the div/mod path. + x = torch.randn(*shape, device="cpu", dtype=torch.float32) + idx_src = torch.tensor(5, device="cpu", dtype=torch.int32) + idx_dst = torch.tensor(10, device="cpu", dtype=torch.int32) + out = torch.zeros(*shape, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + idx_src_txda = idx_src.to("txda") + idx_dst_txda = idx_dst.to("txda") + out_txda = out.to("txda") + insert_dyn_scalar[(1, )](x_txda, idx_src_txda, idx_dst_txda, out_txda, shape, tile) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + expected = x.clone() + expected[8:16, 0:8, 8:16, 0:8] = x[0:8, 8:16, 0:8, 8:16] + 1.0 + torch.testing.assert_close(out, expected) + + +class TestSlice: + + def test_extract_static(self): + torch.manual_seed(42) + x = torch.randn(*SLICE, device="cpu", dtype=torch.float32) + + # stride 1 + out = torch.zeros(8, 8, 8, 8, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + extract_static[(1, )](x_txda, out_txda, SLICE, (4, 4, 4, 4), (8, 8, 8, 8), (1, 1, 1, 1)) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + torch.testing.assert_close(out, x[4:12, 4:12, 4:12, 4:12]) + + # stride 2 + out = torch.zeros(8, 8, 8, 8, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + extract_static[(1, )](x_txda, out_txda, SLICE, (0, 0, 0, 0), (8, 8, 8, 8), (2, 2, 2, 2)) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + torch.testing.assert_close(out, x[0:16:2, 0:16:2, 0:16:2, 0:16:2]) + + def test_extract_dyn(self): + torch.manual_seed(42) + x = torch.randn(*SLICE, device="cpu", dtype=torch.float32) + o0 = torch.tensor(4, device="cpu", dtype=torch.int32) + o1 = torch.tensor(4, device="cpu", dtype=torch.int32) + + # all-dynamic offsets on dims 0,1 + out = torch.zeros(8, 8, 8, 8, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + o0_txda = o0.to("txda") + o1_txda = o1.to("txda") + out_txda = out.to("txda") + extract_dyn[(1, )](x_txda, o0_txda, o1_txda, out_txda, SLICE, (8, 8, 8, 8)) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + torch.testing.assert_close(out, x[4:12, 4:12, 0:8, 0:8]) + + # mixed: dynamic dim0, static dim1 (=8) + out = torch.zeros(8, 8, 8, 8, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + o0_txda = o0.to("txda") + out_txda = out.to("txda") + extract_mixed[(1, )](x_txda, o0_txda, out_txda, SLICE, (8, 8, 8, 8), 8) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + torch.testing.assert_close(out, x[4:12, 8:16, 0:8, 0:8]) + + def test_insert_default(self): + torch.manual_seed(44) + x = torch.randn(*SLICE, device="cpu", dtype=torch.float32) + + out = torch.zeros(*SLICE, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + insert_default[(1, )](x_txda, out_txda, SLICE, (8, 8, 8, 8)) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + expected = x.clone() + expected[8:12, 8:12, 8:12, 8:12] = x[0:4, 0:4, 0:4, 0:4] + 1.0 + torch.testing.assert_close(out, expected) + + def test_insert_strided(self): + torch.manual_seed(44) + x = torch.randn(*SLICE, device="cpu", dtype=torch.float32) + + out = torch.zeros(*SLICE, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + insert_strided[(1, )](x_txda, out_txda, SLICE, (4, 4, 4, 4)) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + expected = x.clone() + expected[4:12:2, 4:12:2, 4:12:2, 4:12:2] = x[0:4, 0:4, 0:4, 0:4] + 1.0 + torch.testing.assert_close(out, expected) + + def test_member(self): + torch.manual_seed(44) + x = torch.randn(*SLICE, device="cpu", dtype=torch.float32) + + out = torch.zeros(*SLICE, device="cpu", dtype=torch.float32) + x_txda = x.to("txda") + out_txda = out.to("txda") + member_roundtrip[(1, )](x_txda, out_txda, SLICE, (8, 8, 8, 8)) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + expected = x.clone() + expected[8:16, 8:16, 8:16, 8:16] = x[0:8, 0:8, 0:8, 0:8] + 1.0 + torch.testing.assert_close(out, expected) + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) diff --git a/third_party/wafer/examples/mult_ir.py b/third_party/wafer/examples/mult_ir.py new file mode 100755 index 00000000..28565486 --- /dev/null +++ b/third_party/wafer/examples/mult_ir.py @@ -0,0 +1,194 @@ +import torch + +import triton +import triton.language as tl +# import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes: +# - A list of `triton.Config` objects that define different configurations of +# meta-parameters (e.g., `BLOCK_SIZE_M`) and compilation options (e.g., `num_warps`) to try +# - An auto-tuning *key* whose change in values will trigger evaluation of all the +# provided configs +@triton.jit +def matmul_kernel( + # Pointers to matrices + a_ptr, b_ptr, c_ptr, + # Matrix dimensions + M, N, K, + # The stride variables represent how much to increase the ptr by when moving by 1 + # element in a particular dimension. E.g. `stride_am` is how much to increase `a_ptr` + # by to get the element one row down (A has M rows). + stride_am, stride_ak, # + stride_bk, stride_bn, # + stride_cm, stride_cn, + # Meta-parameters + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, # + GROUP_SIZE_M: tl.constexpr, BLOCK_SIZE_N1: tl.constexpr, BLOCK_SIZE_N2: tl.constexpr, # + BLOCK_SIZE_K1: tl.constexpr, BLOCK_SIZE_K2: tl.constexpr, ACTIVATION: tl.constexpr # +): + """Kernel for computing the matmul C = A x B. + A has shape (M, K), B has shape (K, N) and C has shape (M, N) + """ + # ----------------------------------------------------------- + # Map program ids `pid` to the block of C it should compute. + # This is done in a grouped ordering to promote L2 data reuse. + # See above `L2 Cache Optimizations` section for details. + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + # ---------------------------------------------------------- + # Create pointers for the first blocks of A and B. + # We will advance this pointer as we move in the K direction + # and accumulate + # `a_ptrs` is a block of [BLOCK_SIZE_M, BLOCK_SIZE_K] pointers + # `b_ptrs` is a block of [BLOCK_SIZE_K, BLOCK_SIZE_N] pointers + # See above `Pointer Arithmetics` section for details + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + mask_m = offs_m < M + mask_n = offs_n < N + offs_nn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N1)[:, None] * BLOCK_SIZE_N2 + tl.arange( + 0, BLOCK_SIZE_N2)[None, :] + offs_kk = tl.arange(0, BLOCK_SIZE_K1)[:, None] * BLOCK_SIZE_K2 + tl.arange(0, BLOCK_SIZE_K2)[None, :] + mask_nn = offs_nn < N + a_ptrs = a_ptr + (offs_m[None, :, None] * stride_am + offs_kk[:, None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[None, :, None] * stride_bk + offs_nn[:, None, :] * stride_bn) + + # ----------------------------------------------------------- + # Iterate to compute a block of the C matrix. + # We accumulate into a `[BLOCK_SIZE_M, BLOCK_SIZE_N]` block + # of fp32 values for higher accuracy. + # `accumulator` will be converted back to fp16 after the loop. + # accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float16) + for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + # Load the next block of A and B, generate a mask by checking the K dimension. + # If it is out of bounds, set it to 0. + mask_k = offs_k < K - k * BLOCK_SIZE_K + mask_kk = offs_kk < K - k * BLOCK_SIZE_K + # a = tl.load(a_ptrs, mask=(mask_m[None, :, None] & mask_kk[:, None, :]), other=0.0) + # b = tl.load(b_ptrs, mask=(mask_k[None, :, None] & mask_nn[:, None, :]), other=0.0) + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + if BLOCK_SIZE_K1 != 1: + a = tl.trans(a, (1, 0, 2)) + a = tl.reshape(a, (BLOCK_SIZE_M, BLOCK_SIZE_K)) + if BLOCK_SIZE_N1 != 1: + b = tl.trans(b, (1, 0, 2)) + b = tl.reshape(b, (BLOCK_SIZE_K, BLOCK_SIZE_N)) + # We accumulate along the K dimension. + acc += tl.dot(a, b, out_dtype=tl.float16) + # Advance the ptrs to the next K block. + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + + if BLOCK_SIZE_N1 == 1: + acc = tl.reshape(acc, (BLOCK_SIZE_N1, BLOCK_SIZE_M, BLOCK_SIZE_N2)) + else: + acc = tl.reshape(acc, (BLOCK_SIZE_M, BLOCK_SIZE_N1, BLOCK_SIZE_N2)) + acc = tl.trans(acc, (1, 0, 2)) + + if ACTIVATION == "leaky_relu": + acc = leaky_relu(acc) + c = acc + + # You can fuse arbitrary activation functions here + # while the accumulator is still in FP32! + # if ACTIVATION == "leaky_relu": + # accumulator = leaky_relu(accumulator) + # c = accumulator.to(tl.float32) + + # ----------------------------------------------------------- + # Write back the block of the output matrix C with masks. + c_ptrs = c_ptr + stride_cm * offs_m[None, :, None] + stride_cn * offs_nn[:, None, :] + tl.store(c_ptrs, c) + # tl.store(c_ptrs, c, mask=(mask_m[None, :, None] & mask_nn[:, None, :])) + + +# We can fuse `leaky_relu` by providing it as an `ACTIVATION` meta-parameter in `_matmul`. +@triton.jit +def leaky_relu(x): + x = x + 1 + return tl.where(x >= 0, x, 0.01 * x) + + +def matmul(a, b, activation=""): + BLOCK_M = 1024 + BLOCK_N = 1024 + BLOCK_K = 128 + + # Check constraints. + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + assert a.is_contiguous(), "Matrix A must be contiguous" + assert b.is_contiguous(), "Matrix B must be contiguous" + M, K = a.shape + K, N = b.shape + + assert M % BLOCK_M == 0 and N % BLOCK_N == 0 and K % BLOCK_K == 0 + + ALIGN = 64 + BLOCK_K1 = max(1, BLOCK_K // ALIGN) + BLOCK_K2 = min(ALIGN, BLOCK_K) + BLOCK_N1 = max(1, BLOCK_N // ALIGN) + BLOCK_N2 = min(ALIGN, BLOCK_N) + # Allocates output. + c = torch.empty((M, N), device=a.device, dtype=a.dtype) + # 1D launch kernel where each block gets its own program. + grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), ) + matmul_kernel[grid]( + a, + b, + c, # + M, + N, + K, # + a.stride(0), + a.stride(1), # + b.stride(0), + b.stride(1), # + c.stride(0), + c.stride(1), # + ACTIVATION=activation, # + BLOCK_SIZE_M=BLOCK_M, + BLOCK_SIZE_N=BLOCK_N, + BLOCK_SIZE_K=BLOCK_K, + GROUP_SIZE_M=8, + BLOCK_SIZE_N1=BLOCK_N1, + BLOCK_SIZE_N2=BLOCK_N2, + BLOCK_SIZE_K1=BLOCK_K1, + BLOCK_SIZE_K2=BLOCK_K2, + ) + return c + + +# @benchmark.measure(repeats=20) +def bench_matmul(a, b): + x = a.to(DEVICE) + y = b.to(DEVICE) + z = matmul(x, y) + z = z.to("cpu") + return z + + +if __name__ == "__main__": + M = 1024 + K = 1024 + N = 1024 + a = torch.randn((M, K), device='cpu', dtype=torch.float16) + b = torch.randn((K, N), device='cpu', dtype=torch.float16) + out = bench_matmul(a, b) + print(out) + ref = torch.matmul(a, b) + print(ref) + print(ref - out) diff --git a/third_party/wafer/examples/profile_matmul.py b/third_party/wafer/examples/profile_matmul.py new file mode 100755 index 00000000..99568b70 --- /dev/null +++ b/third_party/wafer/examples/profile_matmul.py @@ -0,0 +1,193 @@ +import torch + +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes: +# - A list of `triton.Config` objects that define different configurations of +# meta-parameters (e.g., `BLOCK_SIZE_M`) and compilation options (e.g., `num_warps`) to try +# - An auto-tuning *key* whose change in values will trigger evaluation of all the +# provided configs +# @triton.autotune( +# configs=[ +# triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=3, +# num_warps=8), +# triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=5, +# num_warps=2), +# triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=5, +# num_warps=2), +# ], +# key=['M', 'N', 'K'], +# ) +@triton.jit +def matmul_kernel( + # Pointers to matrices + a_ptr, b_ptr, c_ptr, + # Matrix dimensions + M, N, K, + # The stride variables represent how much to increase the ptr by when moving by 1 + # element in a particular dimension. E.g. `stride_am` is how much to increase `a_ptr` + # by to get the element one row down (A has M rows). + stride_am, stride_ak, # + stride_bk, stride_bn, # + stride_cm, stride_cn, + # Meta-parameters + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, # + GROUP_SIZE_M: tl.constexpr, # + ACTIVATION: tl.constexpr # +): + """Kernel for computing the matmul C = A x B. + A has shape (M, K), B has shape (K, N) and C has shape (M, N) + """ + # ----------------------------------------------------------- + # Map program ids `pid` to the block of C it should compute. + # This is done in a grouped ordering to promote L2 data reuse. + # See above `L2 Cache Optimizations` section for details. + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + # ---------------------------------------------------------- + # Create pointers for the first blocks of A and B. + # We will advance this pointer as we move in the K direction + # and accumulate + # `a_ptrs` is a block of [BLOCK_SIZE_M, BLOCK_SIZE_K] pointers + # `b_ptrs` is a block of [BLOCK_SIZE_K, BLOCK_SIZE_N] pointers + # See above `Pointer Arithmetics` section for details + offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + + # ----------------------------------------------------------- + # Iterate to compute a block of the C matrix. + # We accumulate into a `[BLOCK_SIZE_M, BLOCK_SIZE_N]` block + # of fp32 values for higher accuracy. + # `accumulator` will be converted back to fp16 after the loop. + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + # Load the next block of A and B, generate a mask by checking the K dimension. + # If it is out of bounds, set it to 0. + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) + # We accumulate along the K dimension. + accumulator += tl.dot(a, b) + # Advance the ptrs to the next K block. + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + # You can fuse arbitrary activation functions here + # while the accumulator is still in FP32! + if ACTIVATION == "leaky_relu": + accumulator = leaky_relu(accumulator) + c = accumulator.to(tl.float32) + + # ----------------------------------------------------------- + # Write back the block of the output matrix C with masks. + offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, c, mask=c_mask) + + +# We can fuse `leaky_relu` by providing it as an `ACTIVATION` meta-parameter in `_matmul`. +@triton.jit +def leaky_relu(x): + x = x + 1 + return tl.where(x >= 0, x, 0.01 * x) + + +def matmul(a, b, activation=""): + # Check constraints. + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + assert a.is_contiguous(), "Matrix A must be contiguous" + assert b.is_contiguous(), "Matrix B must be contiguous" + M, K = a.shape + K, N = b.shape + # Allocates output. + c = torch.empty((M, N), device=a.device, dtype=a.dtype) + a = a.to(DEVICE) + b = b.to(DEVICE) + c = c.to(DEVICE) + # 1D launch kernel where each block gets its own program. + grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), ) + matmul_kernel[grid]( + a, b, c, # + M, N, K, # + a.stride(0), a.stride(1), # + b.stride(0), b.stride(1), # + c.stride(0), c.stride(1), # + ACTIVATION=activation, # + BLOCK_SIZE_M=512, BLOCK_SIZE_N=128, BLOCK_SIZE_K=128, GROUP_SIZE_M=8) + c = c.to('cpu') + return c + + +def test_matmul(device): + torch.manual_seed(0) + rows1 = 179 + cols1 = 167 + rows2 = 167 + cols2 = 321 + a = torch.randn((rows1, cols1), device=device, dtype=torch.float32) + b = torch.randn((rows2, cols2), device=device, dtype=torch.float32) + # a = torch.full((rows1, cols1), 1, device='cpu', dtype=torch.float32) + # b = torch.full((rows2, cols2), 1, device='cpu', dtype=torch.float32) + triton_output = matmul(a, b) + triton_output = triton_output.to("cpu") + + torch_output = torch.matmul(a, b) + torch.testing.assert_close(triton_output, torch_output, atol=1e-2, rtol=0) + + +@benchmark.measure() +def bench_matmul(M, N, K, provider): + a = torch.randn((M, K), device='cpu', dtype=torch.float16) + b = torch.randn((K, N), device='cpu', dtype=torch.float16) + + torch_output = torch.matmul(a.to(torch.float32), b.to(torch.float32)) + triton_output = matmul(a, b) + # max_deff = torch.max(torch.abs(torch_output - triton_output)) + + abs_diff = torch.abs(torch_output - triton_output) + max_diff = torch.max(abs_diff) + max_diff_index = torch.argmax(abs_diff) + + print(max_diff) + print(max_diff_index) + print(torch_output.reshape(-1)[max_diff_index]) + print(triton_output.reshape(-1)[max_diff_index]) + print(f"The maximum difference between torch and triton is " + f"{max_diff}") + if max_diff > 0.2: + print(torch_output) + print(triton_output) + assert max_diff < 0.2 + + +if __name__ == "__main__": + bench_matmul(4096, 4096, 4096, 'triton') + # for X in [128 * i for i in range(2, 7)]: + # for provider in ['torch', 'triton']: + # bench_matmul(X, X, X, provider) diff --git a/third_party/wafer/examples/quant_gptq.py b/third_party/wafer/examples/quant_gptq.py new file mode 100755 index 00000000..13b622c4 --- /dev/null +++ b/third_party/wafer/examples/quant_gptq.py @@ -0,0 +1,280 @@ +import sys +import itertools +import logging + +import numpy as np +import torch +from torch import nn + +import triton +import triton.language as tl + +logger = logging.getLogger(__name__) +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +def make_dequant_configs(block_sizes, num_warps): + configs = [] + for bs, ws in itertools.product(block_sizes, num_warps): + configs.append(triton.Config({"x_block": bs}, num_warps=ws)) + return configs + + +DEFAULT_DEQUANT_CONFIGS = make_dequant_configs([128], [4, 8]) + + +@triton.autotune(DEFAULT_DEQUANT_CONFIGS, key=["numels"]) +@triton.jit +def dequant_kernel_248( + g_idx_ptr, + scales_ptr, + qweight_ptr, + # qzeros_ptr, + out_ptr, + numels, + # maxq: tl.constexpr, + bits: tl.constexpr, + outfeatures: tl.constexpr, + num_groups: tl.constexpr, + out_type: tl.constexpr, + x_block: tl.constexpr, +): + # Block indexing + xoffset = tl.program_id(0) * x_block + x_index = xoffset + tl.arange(0, x_block) + xmask = x_index < numels + + row_idx = x_index // outfeatures + col_idx = x_index % outfeatures + + elements_per_feature: tl.constexpr = 32 // bits + g_idx = tl.load(g_idx_ptr + (row_idx), None, eviction_policy="evict_last") + qweights = tl.load( + qweight_ptr + (col_idx // elements_per_feature + (outfeatures // elements_per_feature * row_idx)), + None, + ) + + wf_weights = (col_idx % elements_per_feature) * bits + + tmp1 = g_idx + num_groups + tmp2 = g_idx < 0 + tl.device_assert(g_idx >= 0, "index out of bounds: 0 <= tmp0 < 0") + groups = tl.where(tmp2, tmp1, g_idx) # tmp3 are g_idx + + scales = tl.load(scales_ptr + (groups * outfeatures + col_idx), None).to(out_type) + # Unpack weights + weights = qweights >> wf_weights # bit shift qweight + + weights = weights & 0xFF + weights = weights.to(tl.int8).to(out_type) + weights = scales * weights + + tl.store(out_ptr + (x_index), weights, mask=xmask) + + +def dequant248(qweight, scales, qzeros, g_idx, bits, maxq=None, dtype=torch.float16): + """ + Launcher for triton dequant kernel. Only valid for bits = 2, 4, 8 + # compress_ratio = 32 / quant_bit + # float_weight [IC, OC] + # qweight [IC, OC / compress_ratio] + # scales [IC / group_size, OC] + # qzeros [IC / group_size, OC / compress_ratio] + # g_idx [IC] + """ + _ = qzeros + num_groups = scales.shape[0] + outfeatures = scales.shape[1] + infeatures = g_idx.shape[0] + out = torch.empty((infeatures, outfeatures), device="cpu", dtype=dtype) + numels = out.numel() + maxq = 2**bits - 1 if maxq is None else maxq + if dtype == torch.float16: + out_type = tl.float16 + elif dtype == torch.bfloat16: + out_type = tl.bfloat16 + else: + raise NotImplementedError(f"dtype: {dtype}") + + # grid = lambda meta: (triton.cdiv(numels, meta["x_block"]),) # noqa: E731 + def grid(meta): + return (triton.cdiv(numels, meta["x_block"]), ) + + g_idx = g_idx.to(DEVICE) + scales = scales.to(DEVICE) + qweight = qweight.to(DEVICE) + out = out.to(DEVICE) + dequant_kernel_248[grid]( + g_idx, + scales, + qweight, + # qzeros, + out, + numels, + # maxq=maxq, + bits=bits, + outfeatures=outfeatures, + num_groups=num_groups, + out_type=out_type, + ) + out = out.to('cpu') + return out + + +# Copied from https://github.com/IST-DASLab/marlin/pull/1 +def _unpack_weight(qweight, qzeros, weight_width=4, row=False): + # Unpack 4-bit values and interpret them as signed integers + assert weight_width in [4, 8], "weight only support 4 or 8." + compress_ratio = int(32 / weight_width) + if row: + unpacked_shape = (qweight.shape[0] * compress_ratio, qweight.shape[1]) + idx_list = list(range(qweight.shape[0])) + else: + unpacked_shape = (qweight.shape[0], qweight.shape[1] * compress_ratio) + idx_list = list(range(qweight.shape[1])) + + unpacked_weights = torch.zeros( + unpacked_shape, + dtype=torch.int8, + device=qweight.device, + requires_grad=False, + ) + + unpacked_zeros = torch.zeros( + (qzeros.shape[0], qzeros.shape[1] * compress_ratio), + dtype=torch.int8, + device=qzeros.device, + requires_grad=False, + ) + + if weight_width == 4: + mask = 0xF + elif weight_width == 8: + mask = 0xFF + else: + raise ValueError("weight_width must be 4 or 8") + + for i in range(compress_ratio): + idx = (np.array(idx_list) * compress_ratio + i).tolist() + if row: + unpacked_weights[idx, :] = ((qweight >> (weight_width * i)) & mask).to(torch.int8) + else: + unpacked_weights[:, idx] = ((qweight >> (weight_width * i)) & mask).to(torch.int8) + + idx_list = list(range(qzeros.shape[1])) + for i in range(compress_ratio): + idx = (np.array(idx_list) * compress_ratio + i).tolist() + unpacked_zeros[:, idx] = ((qzeros >> (weight_width * i)) & mask).to(torch.int8) + + return unpacked_weights, unpacked_zeros, compress_ratio + + +# Copied from https://github.com/IST-DASLab/marlin/pull/1 +def dequantize_weight(qweight, qzeros, scales, weight_width=4, dtype=torch.float16, row=False): + is_dtensor = False + device_mesh = None + placements = None + if hasattr(qweight, "to_local"): + is_dtensor = True + device_mesh = qweight.device_mesh + placements = qweight.placements + qweight = qweight.to_local() + qweight = qweight.view(torch.int32) + + if hasattr(qzeros, "to_local"): + qzeros = qzeros.to_local() + qzeros = qzeros.view(torch.int32) + if hasattr(scales, "to_local"): + scales = scales.to_local() + + unpacked_qweight, unpacked_qzeros, compress_ratio = _unpack_weight(qweight, qzeros, weight_width, row) + group_size = unpacked_qweight.shape[0] // scales.shape[0] + scales = scales.repeat_interleave(group_size, dim=0) + unpacked_qzeros = unpacked_qzeros.repeat_interleave(group_size, dim=0) + unpacked_dqweight = (unpacked_qweight.to(dtype) - unpacked_qzeros.to(dtype)) * scales + if is_dtensor: + from torch.distributed.tensor import DTensor + + unpacked_dqweight = DTensor.from_local(unpacked_dqweight, device_mesh=device_mesh, placements=placements) + + return unpacked_dqweight, group_size, compress_ratio + + +def dequantize_weight_triton(qweight, qzeros, scales, weight_width=4, dtype=torch.float16): + if len(qweight.shape) == 4: + group_size = qweight.shape[2] // scales.shape[2] + else: + group_size = qweight.shape[0] // scales.shape[0] + compress_ratio = int(32 / weight_width) + if len(qweight.shape) == 4: + infeatures = qweight.shape[2] + else: + infeatures = qweight.shape[0] + g_idx = torch.tensor([i // group_size for i in range(infeatures)], dtype=torch.int32, device=qweight.device) + unpacked_dqweight_triton = dequant248(qweight, scales, qzeros, g_idx, weight_width, dtype=dtype) + return unpacked_dqweight_triton, group_size, compress_ratio + + +class GPTQ(nn.Module): + + def __init__( + self, + mod, + node_name, + wq_params, + prefix, + # weight_dtype=torch.float16, + ): + super().__init__() + self.bits = 4 + self.group_size = 128 + self.outfeatures, self.infeatures = mod.weight.shape + self.weight = mod.weight + self.bias = mod.bias + self.maxq = 2**self.bits - 1 + + # if node_name[-7:] == '_MatMul': + # prefix = node_name[:-7] + # else: + # raise ValueError(f'can not identify prefix for {node_name}') + + prefix = prefix + node_name + # self.g_idx = wq_params[prefix+'.g_idx'] + if prefix + ".qweight" in wq_params and prefix + ".qzeros" in wq_params and prefix + ".scales" in wq_params: + self.qweight = wq_params[prefix + ".qweight"] + self.qzeros = wq_params[prefix + ".qzeros"] + self.scales = wq_params[prefix + ".scales"] + elif prefix.startswith("model_") and prefix[len("model_"):] + ".qweight" in wq_params: + prefix = prefix[len("model_"):] + self.qweight = wq_params[prefix + ".qweight"] + self.qzeros = wq_params[prefix + ".qzeros"] + self.scales = wq_params[prefix + ".scales"] + else: + logger.error(f"GPTQ: {prefix} not found, please check or set this node in noquant_layers!") + sys.exit(-1) + + def forward(self, x: torch.Tensor): + raise NotImplementedError() + # out_shape = x.shape[:-1] + (self.outfeatures,) + # x = x.reshape(-1, x.shape[-1]) + + # out = torch.matmul(x, self.weight) + # out = out.to(x_dtype) + # out = out + self.bias if self.bias is not None else out + # return out + + +__all__ = ["GPTQ"] + +if __name__ == "__main__": + batch_size, num_channels, height, width = 1, 16, 64, 64 + qweight = torch.randint(0, 255, (batch_size, num_channels, height, width), dtype=torch.uint8) + qzeros = torch.randint(0, 255, (batch_size, num_channels, height // 4, width // 4), dtype=torch.uint8) # 假设零值压缩 + scales = torch.randn(batch_size, num_channels, height // 4, width // 4, dtype=torch.float16) + + unpacked_weight, group_size, compress_ratio = dequantize_weight_triton(qweight, qzeros, scales, weight_width=4, + dtype=torch.float16) + + print(f"反量化权重形状: {unpacked_weight.shape}") + print(f"分组大小: {group_size}") + print(f"压缩比: {compress_ratio}") diff --git a/third_party/wafer/examples/quant_kernel.py b/third_party/wafer/examples/quant_kernel.py new file mode 100755 index 00000000..bfa7f071 --- /dev/null +++ b/third_party/wafer/examples/quant_kernel.py @@ -0,0 +1,623 @@ +import torch +import triton +import triton.language as tl + +# pylint: disable=invalid-name +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.autotune( + configs=[ + # A100优化配置 + triton.Config({"BLOCK_SIZE_M": 256, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 8}, num_stages=4, + num_warps=4), # 更适合A100的SM架构 + triton.Config({"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 128, "GROUP_SIZE_M": 8}, num_stages=5, + num_warps=8), # 增大K维度分块 + ], + key=["M", "N", "K"], +) +@triton.jit +def matmul_kernel( + # Pointers to matrices + a_ptr, b_ptr, c_ptr, + # Matrix dimensions + M, N, K, + # The stride variables represent how much to increase the ptr by when moving by 1 + # element in a particular dimension. E.g. `stride_am` is how much to increase `a_ptr` + # by to get the element one row down (A has M rows). + stride_am, stride_ak, # + stride_bk, stride_bn, # + stride_cm, stride_cn, + # Meta-parameters + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, # + GROUP_SIZE_M: tl.constexpr, # +): + """Kernel for computing the matmul C = A x B. + A has shape (M, K), B has shape (K, N) and C has shape (M, N) + """ + # ----------------------------------------------------------- + # Map program ids `pid` to the block of C it should compute. + # This is done in a grouped ordering to promote L2 data reuse. + # See above `L2 Cache Optimizations` section for details. + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + # ---------------------------------------------------------- + # Create pointers for the first blocks of A and B. + # We will advance this pointer as we move in the K direction + # and accumulate + # `a_ptrs` is a block of [BLOCK_SIZE_M, BLOCK_SIZE_K] pointers + # `b_ptrs` is a block of [BLOCK_SIZE_K, BLOCK_SIZE_N] pointers + # See above `Pointer Arithmetics` section for details + offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + + # ----------------------------------------------------------- + # Iterate to compute a block of the C matrix. + # We accumulate into a `[BLOCK_SIZE_M, BLOCK_SIZE_N]` block + # of fp32 values for higher accuracy. + # `accumulator` will be converted back to fp16 after the loop. + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.int32) + for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + # Load the next block of A and B, generate a mask by checking the K dimension. + # If it is out of bounds, set it to 0. + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) + # We accumulate along the K dimension. + accumulator += tl.dot(a, b, out_dtype=tl.int32) + # Advance the ptrs to the next K block. + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + + c = accumulator.to(tl.int32) + + # ----------------------------------------------------------- + # Write back the block of the output matrix C with masks. + offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, c, mask=c_mask) + + +# %% +# We can now create a convenience wrapper function that only takes two input tensors, +# and (1) checks any shape constraint; (2) allocates the output; (3) launches the above kernel. + + +def int8_matmul(a, b): + assert a.shape[-1] == b.shape[0], f"Incompatible dimensions: A {a.shape} vs B {b.shape}" + assert a.is_contiguous() and b.is_contiguous() + + # Save original batch dimensions and flatten + original_shape = a.shape[:-1] + a_flat = a.view(-1, a.shape[-1]) # (B*M, K) + + # Get dimensions for kernel + M_flat, K = a_flat.shape + N = b.shape[1] + + # Allocate output tensor + c = torch.empty((M_flat, N), device=a.device, dtype=torch.int32) + + # Configure kernel grid + def grid(META): + return (triton.cdiv(M_flat, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]), ) + + # Launch kernel with flattened dimensions + matmul_kernel[grid]( + a_flat, + b, + c, + M_flat, + N, + K, # Updated dimensions + a_flat.stride(0), + a_flat.stride(1), + b.stride(0), + b.stride(1), + c.stride(0), + c.stride(1), + ) + + # Unflatten output to (B, M, N) + return c.view(*original_shape, N) + + +def init_to_zero(name): + return lambda nargs: nargs[name].zero_() + + # """ + # A copy of triton.autotune that calls our subclass above + # """ + + # def decorator(fn): + # def wrapper(kernel): + # return Autotuner( + # kernel, fn.arg_names, configs, key, reset_to_zero, prune_configs_by + # ) + + # fn.kernel_decorators.append(wrapper) + # return fn + + # return decorator + + +# reference https://github.com/pytorch/torchdynamo/pull/971 +def conv_heuristics(pre_hook=None): + configs = [ + triton.Config( + {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 32}, + num_stages=2, + num_warps=8, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 256, "BLOCK_N": 64, "BLOCK_K": 32}, + num_stages=2, + num_warps=8, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 256, "BLOCK_N": 32, "BLOCK_K": 32}, + num_stages=4, + num_warps=4, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 256, "BLOCK_N": 32, "BLOCK_K": 64}, + num_stages=4, + num_warps=4, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 256, "BLOCK_N": 16, "BLOCK_K": 32}, + num_stages=4, + num_warps=2, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 32}, + num_stages=4, + num_warps=8, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 32}, + num_stages=4, + num_warps=4, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 64, "BLOCK_N": 64, "BLOCK_K": 32}, + num_stages=4, + num_warps=4, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 128, "BLOCK_N": 16, "BLOCK_K": 32}, + num_stages=4, + num_warps=4, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 128}, + num_stages=3, + num_warps=8, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 256, "BLOCK_N": 64, "BLOCK_K": 128}, + num_stages=3, + num_warps=8, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 256, "BLOCK_N": 32, "BLOCK_K": 128}, + num_stages=4, + num_warps=4, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 128}, + num_stages=4, + num_warps=4, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 128}, + num_stages=4, + num_warps=4, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 128, "BLOCK_N": 32, "BLOCK_K": 64}, + num_stages=4, + num_warps=2, + pre_hook=pre_hook, + ), + triton.Config( + {"BLOCK_M": 64, "BLOCK_N": 64, "BLOCK_K": 64}, + num_stages=4, + num_warps=2, + pre_hook=pre_hook, + ), + # triton.Config( + # {"BLOCK_M": 128, "BLOCK_N": 16, "BLOCK_K": 64}, num_stages=4, num_warps=2, + # ), + ] + key = [ + "BATCH", + "IN_C", + "IN_H", + "IN_W", + "KERNEL_N", + "KERNEL_H", + "KERNEL_W", + "OUT_H", + "OUT_W", + # parameters of conv + "stride_h", + "stride_w", + "padding_h", + "padding_w", + "dilation_h", + "dilation_w", + "output_padding_h", + "output_padding_w", + "groups", + ] + prune_configs_by = { + "top_k": 10, + } + return triton.autotune(configs, key, prune_configs_by=prune_configs_by) + + +@conv_heuristics(pre_hook=init_to_zero("y")) +@triton.jit +def _kernel( + x, + w, + bias, # pylint: disable=unused-argument + y, + # stride of tensor + stride_xn, + stride_xc, + stride_xh, + stride_xw, + stride_wn, + stride_wc, + stride_wh, + stride_ww, + stride_yn, + stride_yc, + stride_yh, + stride_yw, + stride_biasn, # pylint: disable=unused-argument + # Tensor dimensions + BATCH, + IN_C, + IN_H, + IN_W, + KERNEL_N, + KERNEL_H, # pylint: disable=unused-argument + KERNEL_W, + OUT_H, + OUT_W, + # parameters of conv + stride_h, + stride_w, + padding_h, + padding_w, + dilation_h, + dilation_w, + output_padding_h, + output_padding_w, + groups, # pylint: disable=unused-argument + # Metaparameters + ACC_TYPE: tl.constexpr, + # blocks in different dimension + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + # reduction tiling parameter for matmul + BLOCK_K: tl.constexpr, + SPLIT_K: tl.constexpr, +): + """ + each program instance computes a [BLOCK_BATCH, BLOCK_N, BLOCK_H, BLOCK_W] block of y + """ + # ----------------------------------------------------------- + # Map program ids `pid` to the block of y it should compute. + pid_nhw = tl.program_id(0) + pid_k = tl.program_id(1) + pid_window = tl.program_id(2) + + # offset for output y + off_y_k = pid_k * BLOCK_N + tl.arange(0, BLOCK_N) + off_y_nhw = pid_nhw * BLOCK_M + tl.arange(0, BLOCK_M) + off_y_n = off_y_nhw // (OUT_H * OUT_W) + off_y_hw = off_y_nhw % (OUT_H * OUT_W) + off_y_h = off_y_hw // OUT_W + off_y_w = off_y_hw % OUT_W + + # offset for the initial ptr for x + off_x_n = off_y_n + off_x_h = off_y_h * stride_h - padding_h + off_x_w = off_y_w * stride_w - padding_w + off_x_nhw = off_x_n * stride_xn + off_x_h * stride_xh + off_x_w * stride_xw + off_x_inc = tl.arange(0, BLOCK_K) + + # load inc ptr of x, upade x_ptrs + delta_xh = pid_window // KERNEL_W + delta_xw = pid_window % KERNEL_W + delta_xc = off_x_inc + # c, h, w: IN_C, KERNEL_H, KERNEL_W + off_x_crs_unpacked = delta_xh * dilation_h * stride_xh + delta_xw * dilation_w * stride_xw + delta_xc * stride_xc + x_ptrs = x + off_x_nhw[:, None] + off_x_crs_unpacked[None, :] + + mask_x = ((off_x_n < BATCH)[:, None] + & (off_x_inc < IN_C)[None, :] + & (off_x_h + (delta_xh * dilation_h) >= 0)[:, None] + & (off_x_h + (delta_xh * dilation_h) < IN_H)[:, None] + & (off_x_w + (delta_xw * dilation_w) >= 0)[:, None] + & (off_x_w + (delta_xw * dilation_w) < IN_W)[:, None]) + + # offset for the inital ptr for w + off_w_crs = delta_xh * stride_wh + delta_xw * stride_ww + delta_xc * stride_wc + off_w_k = off_y_k + w_ptrs = w + off_w_crs[:, None] + off_w_k[None, :] * stride_wn + # tell triton not to vectorize, otherwise misaligned address error + # w_ptrs = tl.multiple_of(w_ptrs, [1, 1]) + mask_w = (off_x_inc < IN_C)[:, None] & (off_w_k < KERNEL_N)[None, :] + + # ------ load x ------ + matrix_x = tl.load(x_ptrs, mask=mask_x) # BLOCK_M * crs_mul_of_KERNEL + # ------ load w ------ + matrix_w = tl.load(w_ptrs, mask=mask_w) # crs_mul_of_KERNEL * BLOCK_N + + # ----------------------------------------------------------- + # allocate accumulator + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_TYPE) + # acc += tl.dot(matrix_x, matrix_w) + for inc in range(0, IN_C, BLOCK_K): + + # ------ matrix multiplication ------ + acc += tl.dot(matrix_x, matrix_w, out_dtype=ACC_TYPE) + # ------ update ptrs ------ + off_x_inc = inc + BLOCK_K + tl.arange(0, BLOCK_K) + x_ptrs += BLOCK_K * stride_xc + w_ptrs += BLOCK_K * stride_wc + + mask_x = ((off_x_n < BATCH)[:, None] + & (off_x_inc < IN_C)[None, :] + & (off_x_h[:, None] + (delta_xh * dilation_h) >= 0) + & (off_x_h[:, None] + (delta_xh * dilation_h) < IN_H) + & (off_x_w[:, None] + (delta_xw * dilation_w) >= 0) + & (off_x_w[:, None] + (delta_xw * dilation_w) < IN_W)) + mask_w = (off_x_inc < IN_C)[:, None] & (off_w_k < KERNEL_N)[None, :] + # ------ prefetch ------ + # ------ load x ------ + matrix_x = tl.load(x_ptrs, mask=mask_x) + # ------ load w ------ + matrix_w = tl.load(w_ptrs, mask=mask_w) + + # add bias if is not None + # if bias is not None and pid_window == 0: + # off_bias_k = pid_k * BLOCK_N + tl.arange(0, BLOCK_N) + # bias_ptrs = bias + off_bias_k * stride_biasn + # mask_bias = off_bias_k < KERNEL_N + # _bias = tl.load(bias_ptrs, mask=mask_bias) + # acc += _bias[None, :] + + acc = acc.to(ACC_TYPE) + + # rematerialize -- this saves some registers + # offset for output y + off_y_k = pid_k * BLOCK_N + tl.arange(0, BLOCK_N) + off_y_nhw = pid_nhw * BLOCK_M + tl.arange(0, BLOCK_M) + off_y_n = off_y_nhw // (OUT_H * OUT_W) + off_y_hw = off_y_nhw % (OUT_H * OUT_W) + # consider output padding + off_y_h = off_y_hw // OUT_W + output_padding_h + off_y_w = off_y_hw % OUT_W + output_padding_w + + # y ptrs in the block of [BLOCK_M, BLOCK_N] + y_ptrs = (y + off_y_n[:, None] * stride_yn + off_y_h[:, None] * stride_yh + off_y_w[:, None] * stride_yw + + off_y_k[None, :] * stride_yc) + + # out-of-bounds check + mask_y = ((off_y_n < BATCH)[:, None] + & (off_y_h < OUT_H + output_padding_h)[:, None] + & (off_y_w < OUT_W + output_padding_w)[:, None] + & (off_y_k < KERNEL_N)[None, :]) + + if SPLIT_K == 1: + tl.store(y_ptrs, acc, mask=mask_y) + # TODO: 暂时不支持atomic操作,等支持后放开 + # else: + # tl.atomic_add(y_ptrs, acc, mask=mask_y) + + +class _ConvSplit: + kernel = _kernel + + @staticmethod + def _call( + x, + w, + bias, + stride, + padding, + dilation, + transposed, # pylint: disable=unused-argument + output_padding, + groups, + ): + # Q: should we check x, w, bias dtypes? + device = x.device + # input shapes + shape_x = x.shape + shape_w = w.shape + shape_bias = bias.shape if bias is not None else None + + # indicies for the layeout + xn, xc, xh, xw = 0, 1, 2, 3 + yn, yc, yh, yw = 0, 1, 2, 3 + wn, wc, wh, ww = 0, 1, 2, 3 + + # out_channel, in_channel, kernel_height, kernel_width + kernel_size = [shape_w[wh], shape_w[ww]] + input_size = [shape_x[xh], shape_x[xw]] + assert not shape_bias or shape_bias[0] == shape_w[wn], f"bias shape did not match{shape_bias} != {shape_w[wn]}" + in_channel = shape_w[wc] * groups + + assert shape_x[xc] % groups == 0, "in_channels must be divisible by groups" + assert shape_w[wn] % groups == 0, "out_channels must be divisible by groups" + assert shape_x[xc] == in_channel, f"in_channel did not match {shape_x[xc]} != {in_channel}" + # assert kernel_size == [3, 3], "should be _split kernel" + + assert (len(stride) == len(padding) == len(dilation) == len(output_padding) == len(kernel_size) == + len(input_size)) + + # output shape + shape_y = [0] * 4 + shape_y[yn] = shape_x[xn] + shape_y[yc] = shape_w[wn] + shape_y[yh] = (input_size[0] + 2 * padding[0] - dilation[0] * + (kernel_size[0] - 1) - 1 + stride[0]) // stride[0] + 2 * output_padding[0] + shape_y[yw] = (input_size[1] + 2 * padding[1] - dilation[1] * + (kernel_size[1] - 1) - 1 + stride[1]) // stride[1] + 2 * output_padding[1] + + BATCH = shape_x[xn] + IN_C = shape_x[xc] + IN_H = shape_x[xh] + IN_W = shape_x[xw] + KERNEL_N = shape_w[wn] + KERNEL_H = shape_w[wh] + KERNEL_W = shape_w[ww] + OUT_H = shape_y[yh] + OUT_W = shape_y[yw] + + # allocate output + y = torch.empty(shape_y, device=device, dtype=torch.float32) + + # get strides for tensors + stride_x = x.stride() + stride_w = w.stride() + stride_bias = bias.stride() if shape_bias else None + stride_biasn = stride_bias[0] if stride_bias else None + + # output layout should be the same as x + if stride_x[xc] < stride_x[xh] and stride_x[xc] < stride_x[xw]: + y = y.to(memory_format=torch.channels_last) + stride_y = y.stride() + + # accumulator types + ACC_TYPE = torch.float32 + + # launch kernel, 2-dim, batch*h*w, kernel + def grid(META): + return ( + triton.cdiv(BATCH * OUT_H * OUT_W, META["BLOCK_M"]), + triton.cdiv(KERNEL_N, META["BLOCK_N"]), + # split over sliding window KERNEL_H * KERNEL_W; atomic_add + KERNEL_H * KERNEL_W, + ) + + x = x.to(DEVICE) + w = w.to(DEVICE) + # bias = bias.to(DEVICE) + y = y.to(DEVICE) + _kernel[grid]( + x, + w, + bias, + y, + # stride nchw for x,w,y tensor + stride_x[xn], + stride_x[xc], + stride_x[xh], + stride_x[xw], + stride_w[wn], + stride_w[wc], + stride_w[wh], + stride_w[ww], + stride_y[yn], + stride_y[yc], + stride_y[yh], + stride_y[yw], + stride_biasn, + # Tensor dimensions + BATCH, + IN_C, + IN_H, + IN_W, + KERNEL_N, + KERNEL_H, + KERNEL_W, + OUT_H, + OUT_W, + # conv parameters + stride[0], + stride[1], + padding[0], + padding[1], + dilation[0], + dilation[1], + output_padding[0], + output_padding[1], + groups, + # Metaparameters + ACC_TYPE=ACC_TYPE, + # BLOCK_M=128, + # BLOCK_N=32, + # BLOCK_K=BLOCK_K, + SPLIT_K=KERNEL_H * KERNEL_W, + ) + y = y.to('cpu') + return y + + @staticmethod + def forward( + x, + w, + bias=None, + stride=(1, 1), + padding=(0, 0), + dilation=(1, 1), + transposed=False, + output_padding=(0, 0), + groups=1, + ): + if groups != 1: + print(f"Do not support groups = {groups}") + return + if transposed: + print("Do not support transposed") + return _ConvSplit._call( + x, + w, + bias, + stride, + padding, + dilation, + transposed, + output_padding, + groups, + ) + + +if __name__ == "__main__": + x = torch.randn((4, 3, 1024, 1024), device='cpu', dtype=torch.float32) + w = torch.randn((64, 3, 3, 3), device='cpu', dtype=torch.float32) + int8_conv = _ConvSplit.forward(x, w) diff --git a/third_party/wafer/examples/single_conv2d.py b/third_party/wafer/examples/single_conv2d.py new file mode 100755 index 00000000..8c532fd8 --- /dev/null +++ b/third_party/wafer/examples/single_conv2d.py @@ -0,0 +1,261 @@ +import torch +import triton +import triton.language as tl +import pytest + + +# Triton Conv2D Kernel实现 +@triton.autotune( + configs=[ + triton.Config({'BLOCK_NI_HO_WO': 128, 'BLOCK_CI': 32, 'BLOCK_CO': 32}, num_warps=4), + triton.Config({'BLOCK_NI_HO_WO': 256, 'BLOCK_CI': 64, 'BLOCK_CO': 32}, num_warps=8), + ], + key=[ + 'in_n', 'weight_c', 'input_h', 'input_w', 'out_c', 'out_h', 'out_w', 'weight_h', 'weight_w', 'stride', + 'padding', 'groups' + ], +) +@triton.jit +def conv2d_kernel( + input_ptr, + weight_ptr, + output_ptr, + bias_ptr, + in_n, + input_h, + input_w, + out_c, + out_h, + out_w, + input_n_stride, + input_c_stride, + input_h_stride, + input_w_stride, + weight_n_stride, + weight_c_stride, + weight_h_stride, + weight_w_stride, + output_n_stride, + output_c_stride, + output_h_stride, + output_w_stride, + weight_c: tl.constexpr, + weight_h: tl.constexpr, + weight_w: tl.constexpr, + stride: tl.constexpr, + padding: tl.constexpr, + dilation: tl.constexpr, + groups: tl.constexpr, + BLOCK_NI_HO_WO: tl.constexpr, + BLOCK_CI: tl.constexpr, + BLOCK_CO: tl.constexpr, +): + pid_ni_ho_wo = tl.program_id(0) + pid_co = tl.program_id(1) + pid_group = tl.program_id(2) + + # 计算位置索引 + ni_ho_wo_offset = pid_ni_ho_wo * BLOCK_NI_HO_WO + tl.arange(0, BLOCK_NI_HO_WO) + ni_ho_offset = ni_ho_wo_offset // out_w + in_n_idx = ni_ho_offset // out_h + out_h_idx = ni_ho_offset % out_h + out_w_idx = ni_ho_wo_offset % out_w + + # 计算指针偏移 + out_per_group_c = out_c // groups + output_c_offset = pid_co * BLOCK_CO + tl.arange(0, BLOCK_CO) + + input_ptr += (input_n_stride * in_n_idx + input_c_stride * pid_group * weight_c)[:, None] + weight_ptr += (weight_n_stride * output_c_offset + weight_n_stride * pid_group * out_per_group_c)[None, :] + + # 累加器初始化 + accum = tl.zeros((BLOCK_NI_HO_WO, BLOCK_CO), dtype=tl.float32) + + # 计算循环次数 + BLOCK_CI_COUNT = (weight_c + BLOCK_CI - 1) // BLOCK_CI + + for hwc in range(weight_h * weight_w * BLOCK_CI_COUNT): + c = (hwc % BLOCK_CI_COUNT) * BLOCK_CI + hw = hwc // BLOCK_CI_COUNT + h = hw // weight_w + w = hw % weight_w + + input_c_offset = c + tl.arange(0, BLOCK_CI) + input_h_offset = h * dilation - padding + stride * out_h_idx + input_w_offset = w * dilation - padding + stride * out_w_idx + + # 计算输入指针 + curr_input_ptr = (input_ptr + (input_c_stride * input_c_offset)[None, :] + + (input_h_stride * input_h_offset)[:, None] + (input_w_stride * input_w_offset)[:, None]) + + # 计算权重指针 + curr_weight_ptr = (weight_ptr + (weight_c_stride * input_c_offset)[:, None] + (weight_h_stride * h) + + (weight_w_stride * w)) + + # 掩码计算 + input_mask = ((in_n_idx < in_n)[:, None] & (input_c_offset < weight_c)[None, :] & (0 <= input_h_offset)[:, None] + & (input_h_offset < input_h)[:, None] & (0 <= input_w_offset)[:, None] & + (input_w_offset < input_w)[:, None]) + + weight_mask = (input_c_offset < weight_c)[:, None] & (output_c_offset < out_per_group_c)[None, :] + + # 加载数据 + input_block = tl.load(curr_input_ptr, mask=input_mask) + weight_block = tl.load(curr_weight_ptr, mask=weight_mask) + + # 矩阵乘法累加 + accum += tl.dot(input_block, weight_block, allow_tf32=False) + + # 处理偏置 + bias_ptr += (pid_group * out_per_group_c)[None, :] + output_c_offset[None, :] + mask_bias = (output_c_offset < out_per_group_c)[None, :] + bias = tl.load(bias_ptr, mask_bias).to(tl.float32) + accum += bias + + # 计算结果存储位置 + output_ptr += ((output_n_stride * in_n_idx)[:, None] + (output_c_stride * + (pid_group * out_per_group_c + output_c_offset))[None, :] + + (output_h_stride * out_h_idx)[:, None] + (output_w_stride * out_w_idx)[:, None]) + + # 输出掩码 + output_mask = ((in_n_idx < in_n)[:, None] & (output_c_offset < out_per_group_c)[None, :] & + (out_h_idx < out_h)[:, None] & (out_w_idx < out_w)[:, None]) + + # 存储结果 + tl.store(output_ptr, accum, mask=output_mask) + + +# 计算输出尺寸的辅助函数 +def conv2d_output_size(in_size, kernel_size, stride, padding, dilation): + return (in_size + 2 * padding - dilation * (kernel_size - 1) - 1) // stride + 1 + + +# Triton Conv2D函数封装 +def triton_conv2d(input, weight, bias=None, stride=1, padding=0, dilation=1, groups=1): + assert weight.ndim == 4, "Weights must be 4D" + assert bias is None or bias.ndim == 1, "Bias must be 1D" + assert input.shape[1] == groups * weight.shape[1], "Incompatible input and weights shape" + assert bias is None or weight.shape[0] == bias.shape[0], "Incompatible weights and bias shape" + + # 统一参数格式 + stride = (stride, stride) if isinstance(stride, int) else stride + padding = (padding, padding) if isinstance(padding, int) else padding + dilation = (dilation, dilation) if isinstance(dilation, int) else dilation + + in_n, _, input_h, input_w = input.shape + out_c, weight_c, weight_h, weight_w = weight.shape + + # 计算输出尺寸 + out_h = conv2d_output_size(input_h, weight_h, stride[0], padding[0], dilation[0]) + out_w = conv2d_output_size(input_w, weight_w, stride[1], padding[1], dilation[1]) + + # 准备输出张量 + output = torch.empty((in_n, out_c, out_h, out_w), device=input.device, dtype=input.dtype) + + # 如果没有偏置,创建零偏置 + if bias is None: + bias = torch.zeros(out_c, device=input.device, dtype=input.dtype) + + # 计算网格大小 + grid = lambda META: ( + triton.cdiv(in_n * out_h * out_w, META['BLOCK_NI_HO_WO']), + triton.cdiv(int(out_c // groups), META['BLOCK_CO']), + groups, + ) + + # 调用kernel + conv2d_kernel[grid]( + input, + weight, + output, + bias, + in_n, + input_h, + input_w, + out_c, + out_h, + out_w, + *input.stride(), + *weight.stride(), + *output.stride(), + weight_c, + weight_h, + weight_w, + stride[0], + padding[0], + dilation[0], + groups, + ) + + return output + + +# 测试用例 +@pytest.mark.parametrize("shape,kernel,groups", [ + ((1, 2, 5, 5), (1, 2, 3, 3), 1), + ((2, 3, 9, 9), (1, 3, 3, 3), 1), + ((32, 8, 8, 8), (32, 8, 2, 2), 1), + ((18, 16, 4, 4), (16, 16, 2, 2), 1), + ((9, 16, 4, 4), (128, 4, 2, 2), 4), +]) +@pytest.mark.parametrize("stride", [1, 2]) +@pytest.mark.parametrize("padding", [0, 1]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16]) +@pytest.mark.parametrize("bias", [True, False]) +def test_triton_conv2d(shape, kernel, groups, stride, padding, dtype, bias): + # 跳过float16测试如果设备不支持 + if dtype == torch.float16 and not torch.cuda.is_available(): + pytest.skip("CUDA not available for float16 test") + + # 创建输入数据 + torch.manual_seed(42) + input = torch.randn(shape, dtype=dtype, device='cuda' if torch.cuda.is_available() else 'cpu') + weight = torch.randn(kernel, dtype=dtype, device=input.device) + + # 创建偏置 + if bias: + bias_tensor = torch.randn(kernel[0], dtype=dtype, device=input.device) + else: + bias_tensor = None + + # 计算PyTorch参考输出 + torch_out = torch.nn.functional.conv2d(input, weight, bias=bias_tensor, stride=stride, padding=padding, dilation=1, + groups=groups) + + # 计算Triton输出 + triton_out = triton_conv2d(input, weight, bias=bias_tensor, stride=stride, padding=padding, dilation=1, + groups=groups) + + # 验证结果 + if dtype == torch.float32: + rtol, atol = 1e-4, 1e-5 + else: + rtol, atol = 1e-2, 1e-3 + + torch.testing.assert_close(triton_out, torch_out, rtol=rtol, atol=atol) + + +if __name__ == "__main__": + # 简单演示 + if torch.cuda.is_available(): + device = 'cuda' + print("Running demo on CUDA") + else: + device = 'cpu' + print("Running demo on CPU") + + # 创建测试数据 + input = torch.randn(1, 3, 16, 16, device=device) + weight = torch.randn(6, 3, 3, 3, device=device) + bias = torch.randn(6, device=device) + + # 运行Triton实现 + output = triton_conv2d(input, weight, bias=bias, stride=1, padding=1) + print("Triton Conv2D output shape:", output.shape) + + # 运行PyTorch实现 + torch_output = torch.nn.functional.conv2d(input, weight, bias=bias, stride=1, padding=1) + print("PyTorch Conv2D output shape:", torch_output.shape) + + # 比较结果 + print("Max difference:", (output - torch_output).abs().max().item()) diff --git a/third_party/wafer/examples/test_abs.py b/third_party/wafer/examples/test_abs.py new file mode 100755 index 00000000..f41d4063 --- /dev/null +++ b/third_party/wafer/examples/test_abs.py @@ -0,0 +1,103 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def abs_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.abs(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def abs_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + x = x.to(DEVICE) + output = output.to(DEVICE) + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + abs_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output = output.to('cpu') + return output + + +@benchmark.measure() +def benchmark_abs_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input tensor + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = abs_triton(x) + + # Verify the result + expected = torch.abs(x) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_abs(size, dtype, device="cpu"): + # Generate random input tensor + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = abs_triton(x) + + # Verify the result + expected = torch.abs(x) + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(expected - output))}") + assert torch.allclose(output, expected, atol=1e-5, rtol=0) + + +if __name__ == "__main__": + for X in [2**i for i in range(22, 25, 1)]: + benchmark_abs_triton(X, torch.float32, provider="triton") diff --git a/third_party/wafer/examples/test_addptr.py b/third_party/wafer/examples/test_addptr.py new file mode 100755 index 00000000..c368edcd --- /dev/null +++ b/third_party/wafer/examples/test_addptr.py @@ -0,0 +1,45 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + +@triton.jit +def addptr(in0, out0): + for i in range(0, 10, 2): + in1 = in0 + 1 + i + in2 = in1 + 1 + + out1 = out0 + 1 + i + out2 = out1 + 1 + + a1 = tl.load(in1) + a2 = tl.load(in2) + + tl.store(out1, a1) + tl.store(out2, a2) + + +def test(device): + input = torch.arange(0, 11, device="cpu", dtype=torch.float32) + output = torch.full((11, ), 0, device="cpu", dtype=torch.float32) + grid = lambda meta: (1, ) + + print(output) + input_txda = input.to("txda") + output_txda = output.to("txda") + addptr[grid](input_txda, output_txda) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(input) + print(output) + assert torch.equal(input, output) + + # TODO: need to check some conditions otherwise the code below does not make any difference for the test + src = triton.compiler.ASTSource( + fn=addptr, + signature={'in0': '*fp32', 'out0': '*fp32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) diff --git a/third_party/wafer/examples/test_argmax2d.py b/third_party/wafer/examples/test_argmax2d.py new file mode 100755 index 00000000..e93d10a5 --- /dev/null +++ b/third_party/wafer/examples/test_argmax2d.py @@ -0,0 +1,53 @@ +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import pytest +import benchmark + + +@triton.jit +def argmax_kernel_2d( + x_ptr, + output_ptr, + N: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + # Get current block's row index + pid_x = tl.program_id(0) + + # Calculate data offsets for current block + offs_x = pid_x * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + offs_y = tl.arange(0, N) # Iterate through all columns + + # Load input data + x = tl.load(x_ptr + offs_x[:, None] * N + offs_y[None, :]) + + # Calculate argmax for each row + max_idx = tl.argmax(x, axis=1) + + # Convert result to int32 and store + result = max_idx.to(tl.int32) + tl.store(output_ptr + offs_x, result) + + +@pytest.mark.parametrize("N", [16, 32, 64]) +def test_argmax(N, device): + # Set input size + x = torch.rand([N, N], device="cpu", dtype=torch.float32) + output = torch.empty([N], device="cpu", dtype=torch.int32) + + # Run kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + argmax_kernel_2d[(1, )](x_txda, output_txda, N, BLOCK_SIZE=N) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + # Calculate reference result and verify + ans = torch.argmax(x, dim=1).to(torch.int32) + torch.testing.assert_close(output, ans, rtol=0, atol=0) + + +if __name__ == "__main__": + test_argmax(16, "cpu") diff --git a/third_party/wafer/examples/test_assert.py b/third_party/wafer/examples/test_assert.py new file mode 100755 index 00000000..46a4ba92 --- /dev/null +++ b/third_party/wafer/examples/test_assert.py @@ -0,0 +1,56 @@ +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import pytest + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def kernel_device_assert_scalar(COND, BLOCK: tl.constexpr): + tl.device_assert(COND, "test scalar") + + +@triton.jit +def kernel_device_assert_tensor(COND, n_elements, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offsets = pid * BLOCK + tl.arange(0, BLOCK) + mask = offsets < n_elements + cond_value = tl.load(COND + offsets, mask=mask, other=1) + tl.device_assert(cond_value, "test tensor") + + +@pytest.mark.parametrize('cond', [True, False]) +def test_assert_scalar(cond, request): + if not cond and request.config.getoption("--wafer-hardware"): + pytest.skip("False device assertions terminate through firmware __assert_func; compile coverage only on this shared card") + kernel_device_assert_scalar[(1, )](cond, BLOCK=16, debug=True) + + +@pytest.mark.parametrize('cond_list', [ + [True, True, True], + [True, False, True], + [False, False, False], + [True], + [False], +]) +def test_assert_tensor(cond_list, request): + if not all(cond_list) and request.config.getoption("--wafer-hardware"): + pytest.skip("False device assertions terminate through firmware __assert_func; compile coverage only on this shared card") + cond_tensor = torch.tensor(cond_list, dtype=torch.bool, device="cpu") + n_elements = cond_tensor.numel() + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK']), ) + cond_tensor_txda = cond_tensor.to("txda") + kernel_device_assert_tensor[grid](cond_tensor_txda, n_elements, BLOCK=16, debug=True) + + +def run_all_tests(): + test_assert_scalar() + test_assert_tensor() + print("Manually check that the printout is correct!") + + +if __name__ == "__main__": + # Run the test with pytest + run_all_tests() diff --git a/third_party/wafer/examples/test_autotune.py b/third_party/wafer/examples/test_autotune.py new file mode 100644 index 00000000..4a9b4ac3 --- /dev/null +++ b/third_party/wafer/examples/test_autotune.py @@ -0,0 +1,39 @@ +"""Benchmark a simple kernel so timing is tested independently of GEMM lowering.""" +import math + +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +from triton.testing import do_bench + + +def test_autotune_vector(device): + measured = [] + + def benchmark(fn, quantiles): + times = do_bench(fn, warmup=1, rep=3, quantiles=quantiles) + assert all(math.isfinite(t) and t > 0 for t in times) + measured.append(times) + return times + + @triton.autotune(configs=[triton.Config({'BLOCK': 64}), triton.Config({'BLOCK': 128})], + key=['N'], do_bench=benchmark) + @triton.jit + def add_kernel(X, Y, Out, N: tl.constexpr, BLOCK: tl.constexpr): + offset = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + value = tl.load(X + offset, offset < N, other=0) + tl.load(Y + offset, offset < N, other=0) + tl.store(Out + offset, value, offset < N) + + x = torch.arange(257, dtype=torch.float32, device="cpu") + y = torch.full_like(x, 1.25) + out = torch.empty_like(x) + x_txda = x.to("txda") + y_txda = y.to("txda") + out_txda = out.to("txda") + add_kernel[lambda meta: (triton.cdiv(x_txda.numel(), meta['BLOCK']),)](x_txda, y_txda, out_txda, N=x_txda.numel()) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + assert len(measured) == 2 + assert add_kernel.best_config.kwargs['BLOCK'] in (64, 128) + torch.testing.assert_close(out, x + y) diff --git a/third_party/wafer/examples/test_bare_matmul.py b/third_party/wafer/examples/test_bare_matmul.py new file mode 100755 index 00000000..493976d5 --- /dev/null +++ b/third_party/wafer/examples/test_bare_matmul.py @@ -0,0 +1,75 @@ +# this is a benchmark which multiplies square matrices with maximum block size +# to check the performance of tl.dot operation + +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def bare_matmul(X, Y, Z, M, N, K, BLOCK_SIZE: tl.constexpr): + pid_x = tl.program_id(0) # block row id + pid_y = tl.program_id(1) # block column id + + offs_x = pid_x * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + offs_y = pid_y * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + + x = tl.load(X + offs_x[:, None] * K + offs_y[None, :]) + y = tl.load(Y + offs_x[:, None] * N + offs_y[None, :]) + + z = tl.dot(x, y) + + tl.store(Z + offs_x[:, None] * N + offs_y[None, :], z) + + +@benchmark.measure() +def bench_matmul(N, provider): + device = 'cpu' + dtype = torch.float32 + a = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + b = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + c = torch.empty((N, N), device="cpu", dtype=dtype) + if provider == 'torch' or provider == 'test': + c_ref = torch.matmul(a, b) + if provider == 'triton' or provider == 'test': + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + bare_matmul[(1, )](a_txda, b_txda, c_txda, N, N, N, N) + with torch.no_grad(): + c.copy_(c_txda.cpu()) + if provider == 'test': + torch.testing.assert_close(c, c_ref, atol=1e-2, rtol=0) + + +@pytest.mark.parametrize("N, dtype", [ # + (N, dtype) for N in [64, 128, 256] for dtype in [torch.float32] +]) +def test_bare_matmul(N, dtype, device='cpu'): + a = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + b = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + c = torch.empty((N, N), device="cpu", dtype=dtype) + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + bare_matmul[(1, )](a_txda, b_txda, c_txda, N, N, N, N) + with torch.no_grad(): + c.copy_(c_txda.cpu()) + c_ref = torch.matmul(a, b) + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(c_ref - c))}") + torch.testing.assert_close(c, c_ref, atol=1e-5, rtol=0) + + +if __name__ == "__main__": + test_bare_matmul(128, torch.float32) + # for X in [2**i for i in range(7, 10, 1)]: + # for provider in ['test', 'torch', 'triton']: + # bench_matmul(X, provider) diff --git a/third_party/wafer/examples/test_bare_matmul_acc.py b/third_party/wafer/examples/test_bare_matmul_acc.py new file mode 100755 index 00000000..ae4bb6e1 --- /dev/null +++ b/third_party/wafer/examples/test_bare_matmul_acc.py @@ -0,0 +1,77 @@ +# this is a benchmark which multiplies square matrices with maximum block size +# and additional accumulation to check the performance of tl.dot operation + +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def bare_matmul_acc(X, Y, Z, C, M, N, K, BLOCK_SIZE: tl.constexpr): + pid_x = tl.program_id(0) # block row id + pid_y = tl.program_id(1) # block column id + + offs_x = pid_x * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + offs_y = pid_y * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + + x = tl.load(X + offs_x[:, None] * K + offs_y[None, :]) + y = tl.load(Y + offs_x[:, None] * N + offs_y[None, :]) + c = tl.load(C + offs_x[:, None] * N + offs_y[None, :]) + + z = tl.dot(x, y, c) + + tl.store(Z + offs_x[:, None] * N + offs_y[None, :], z) + + +@benchmark.measure() +def bench_matmul(N, provider): + device = 'cpu' + dtype = torch.float32 + a = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + b = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + c = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + z = torch.empty((N, N), device="cpu", dtype=dtype) + if provider == 'torch' or provider == 'test': + z_ref = torch.matmul(a, b) + c + if provider == 'triton' or provider == 'test': + a_txda = a.to("txda") + b_txda = b.to("txda") + z_txda = z.to("txda") + c_txda = c.to("txda") + bare_matmul_acc[(1, )](a_txda, b_txda, z_txda, c_txda, N, N, N, N) + with torch.no_grad(): + z.copy_(z_txda.cpu()) + if provider == 'test': + torch.testing.assert_close(z, z_ref, atol=1e-2, rtol=0) + + +@pytest.mark.parametrize("N, dtype", [ # + (N, dtype) for N in [64, 128, 256] for dtype in [torch.float32] +]) +def test_bare_matmul_acc(N, dtype, device='cpu'): + a = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + b = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + c = torch.randint(0, 100, (N, N), device="cpu", dtype=dtype) + z = torch.empty((N, N), device="cpu", dtype=dtype) + a_txda = a.to("txda") + b_txda = b.to("txda") + z_txda = z.to("txda") + c_txda = c.to("txda") + bare_matmul_acc[(1, )](a_txda, b_txda, z_txda, c_txda, N, N, N, N) + with torch.no_grad(): + z.copy_(z_txda.cpu()) + z_ref = torch.matmul(a, b) + c + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(z_ref - z))}") + torch.testing.assert_close(z, z_ref, atol=1e-5, rtol=0) + + +if __name__ == "__main__": + for X in [2**i for i in range(7, 10, 1)]: + for provider in ['test', 'torch', 'triton']: + bench_matmul(X, provider) diff --git a/third_party/wafer/examples/test_blockptr_complex_offset.py b/third_party/wafer/examples/test_blockptr_complex_offset.py new file mode 100755 index 00000000..b61f13f1 --- /dev/null +++ b/third_party/wafer/examples/test_blockptr_complex_offset.py @@ -0,0 +1,41 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + +@triton.jit +def block_copy_kernel(a_ptr, b_ptr): + a_block_ptr = tl.make_block_ptr( + base=a_ptr + 8, + shape=(2, 2), + strides=(2, 1), + offsets=(0, 0), + block_shape=(2, 2), + order=(1, 0), + ) + b_block_ptr = tl.make_block_ptr( + base=b_ptr, + shape=(2, 2), + strides=(2, 1), + offsets=(0, 0), + block_shape=(2, 2), + order=(1, 0), + ) + a = tl.load(a_block_ptr, boundary_check=(0, )) + tl.store(b_block_ptr, a, boundary_check=(0, )) + + +def test(device): + input = torch.arange(0, 16, device="cpu", dtype=torch.float32) + output = torch.full((4, ), -1, device="cpu", dtype=torch.float32) + expected = torch.arange(8, 12, device="cpu") + grid = lambda meta: (1, ) + + input_txda = input.to("txda") + output_txda = output.to("txda") + block_copy_kernel[grid](input_txda, output_txda) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + assert torch.equal(expected, output) diff --git a/third_party/wafer/examples/test_cdiv.py b/third_party/wafer/examples/test_cdiv.py new file mode 100755 index 00000000..b1faea52 --- /dev/null +++ b/third_party/wafer/examples/test_cdiv.py @@ -0,0 +1,111 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def cdiv_kernel( + x_ptr, + y_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.cdiv(x, y) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def cdiv_triton(x, y): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + x = x.to(DEVICE) + y = y.to(DEVICE) + output = output.to(DEVICE) + # Launch the kernel + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + cdiv_kernel[grid]( + x_txda, + y_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output = output.to('cpu') + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.int32] +]) +def test_cdiv(size, dtype, device="cpu"): + # Generate random input tensor + x = torch.randint(1, 100, (size, ), device="cpu", dtype=dtype) + y = torch.randint(1, 100, (size, ), device="cpu", dtype=dtype) + + # Call the Triton kernel + output = cdiv_triton(x, y) + + # Verify the result + expected = (x + y - 1) // y + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(expected - output))}") + # Verify the result + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_cdiv_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randint(1, 100, (size, ), device="cpu", dtype=dtype) + y = torch.randint(1, 100, (size, ), device="cpu", dtype=dtype) + + # Run the Triton kernel + output = cdiv_triton(x, y) + + # Verify the result + torch.testing.assert_close(output, (x + y - 1) // y, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for X in [2**i for i in range(22, 25, 1)]: + benchmark_cdiv_triton(X, torch.int32, "triton") diff --git a/third_party/wafer/examples/test_ceil.py b/third_party/wafer/examples/test_ceil.py new file mode 100755 index 00000000..9bac5e43 --- /dev/null +++ b/third_party/wafer/examples/test_ceil.py @@ -0,0 +1,100 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def ceil_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.ceil(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def ceil_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + ceil_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_ceil(size, dtype, device="cpu"): + # Generate random input tensor + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = ceil_triton(x) + + # Verify the result + expected = torch.ceil(x) + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(expected - output))}") + assert torch.allclose(output, expected, atol=1e-5, rtol=0) + + +@benchmark.measure() +def benchmark_ceil_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input tensor + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = ceil_triton(x) + + # Verify the output + expected = torch.ceil(x) + assert torch.allclose(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [2**i for i in range(22, 25, 1)]: + benchmark_ceil_triton(size, torch.float32, provider="triton") diff --git a/third_party/wafer/examples/test_clamp.py b/third_party/wafer/examples/test_clamp.py new file mode 100755 index 00000000..967965d6 --- /dev/null +++ b/third_party/wafer/examples/test_clamp.py @@ -0,0 +1,106 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark as benchmark + + +@triton.jit +def clamp_kernel( + x_ptr, + output_ptr, + min_val, + max_val, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.clamp(x, min_val, max_val) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def clamp_triton(x, min_val, max_val): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + clamp_kernel[grid]( + x_txda, + output_txda, + min_val, + max_val, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_clamp(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + min_val = -1.0 + max_val = 1.0 + + # Call the Triton kernel + output = clamp_triton(x, min_val, max_val) + + # Verify the output + expected = torch.clamp(x, min_val, max_val) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_clamp_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + min_val = -1.0 + max_val = 1.0 + + # Call the Triton kernel + output = clamp_triton(x, min_val, max_val) + + # Verify the output + expected = torch.clamp(x, min_val, max_val) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + return output + + +if __name__ == "__main__": + for size in [2**i for i in range(22, 25, 1)]: + benchmark_clamp_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_cos.py b/third_party/wafer/examples/test_cos.py new file mode 100755 index 00000000..5155a120 --- /dev/null +++ b/third_party/wafer/examples/test_cos.py @@ -0,0 +1,103 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def cos_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.cos(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def cos_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + x = x.to(DEVICE) + output = output.to(DEVICE) + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + cos_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output = output.to('cpu') + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_cos(size, dtype, device="cpu"): + # Create a random tensor + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = cos_triton(x) + + # Verify the output + expected = torch.cos(x) + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(expected - output))}") + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_cos_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + # Create a random tensor + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = cos_triton(x) + + # Verify the output + expected = torch.cos(x) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [2**i for i in range(22, 25, 1)]: + benchmark_cos_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_debug_barrier.py b/third_party/wafer/examples/test_debug_barrier.py new file mode 100755 index 00000000..1dd2c182 --- /dev/null +++ b/third_party/wafer/examples/test_debug_barrier.py @@ -0,0 +1,30 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + +@triton.jit +def barrier(data): + ptrs = data + tl.arange(0, 128) + + tl.debug_barrier() + tl.store(ptrs, tl.load(ptrs) + 1.0) + tl.debug_barrier() + + +def test_barrier(): + data = torch.zeros(128, dtype=torch.float32, device="cpu") + + # Launch the kernel + grid = (1, ) + data_txda = data.to("txda") + barrier[grid](data_txda) + with torch.no_grad(): + data.copy_(data_txda.cpu()) + torch.testing.assert_close(data, torch.ones_like(data)) + + +if __name__ == "__main__": + test_barrier() diff --git a/third_party/wafer/examples/test_div_rn.py b/third_party/wafer/examples/test_div_rn.py new file mode 100755 index 00000000..bae7ce50 --- /dev/null +++ b/third_party/wafer/examples/test_div_rn.py @@ -0,0 +1,107 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark +from util import gems_assert_cosine_similarity + + +@triton.jit +def div_rn_kernel( + x_ptr, + y_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.div_rn(x, y) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def div_rn_triton(x, y): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + div_rn_kernel[grid]( + x_txda, + y_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_div_rn(size, dtype, device="cpu"): + # Generate random input tensors + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = div_rn_triton(x, y) + + # Verify the output + # TODO:rounding_mode need double check + expected = torch.div(x, y, rounding_mode=None) + # torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + # FIXME: div int has low precision, we use cosine similarity to verify + gems_assert_cosine_similarity(output, expected, dtype=dtype) + + +@benchmark.measure() +def benchmark_div_rn_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input tensors + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = div_rn_triton(x, y) + + # Verify the output + # TODO:rounding_mode need double check + expected = torch.div(x, y, rounding_mode=None) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for i in [2**i for i in range(22, 25, 1)]: + benchmark_div_rn_triton(i, torch.float32, provider="triton") diff --git a/third_party/wafer/examples/test_dot_scaled.py b/third_party/wafer/examples/test_dot_scaled.py new file mode 100755 index 00000000..85dd4943 --- /dev/null +++ b/third_party/wafer/examples/test_dot_scaled.py @@ -0,0 +1,393 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import itertools +import benchmark +from _wafer_reference import upcast_mxfp_cpu + +from triton._internal_testing import ( + integral_dtypes, + int_dtypes, + str_to_triton_dtype, + uint_dtypes, + float_dtypes, + float_dtypes_with_bfloat16, + dtypes, + dtypes_with_bfloat16, + is_cuda, + is_interpreter, + is_hopper, + is_hip, + is_hip_cdna, + is_hip_cdna2, + is_hip_cdna3, + is_hip_cdna4, + is_xpu, + get_arch, + torch_float8_dtypes, + torch_dtypes, + numpy_random, + to_triton, + torch_dtype_name, + to_numpy, +) +from triton.runtime.errors import InterpreterError + +mma_nonk_sizes = [] + +GPU_DIALECT = "ttg" +if is_interpreter(): + THREADS_PER_WARP = 1 +elif is_hip(): + THREADS_PER_WARP = triton.runtime.driver.active.get_current_target().warp_size + # for CDNA multiple variants of mma instructions are supported: + # mfma 16x16/mfma 32x32 + # 0 is a special value for automatic heuristic + if is_hip_cdna(): + mma_nonk_sizes = [0, 16, 32] +else: + THREADS_PER_WARP = 32 + +RESOLUTION = { + torch.bool: + 0, + torch.int16: + 0, + torch.int32: + 0, + torch.int64: + 0, + # torch.float16: 1e-3, # FIXME: test_scaled_dot[32-32-64-True-True-False-e2m1-fp16-4-16-1]... (fp16 cases) Failed + torch.float16: + 1e-2, + torch.float32: + 1.3e-6, + # torch.bfloat16: 0.016, #FIXME: test_scaled_dot[128-128-64-True-True-False-e2m1-e5m2-4-16-1] Failed + torch.bfloat16: + 0.018, + torch.float64: + 1e-7, + torch.complex32: + 1e-3, + torch.complex64: + 1.3e-6, +} + + +def flaggems_assert_close(res, ref, dtype, equal_nan=False, reduce_dim=1): + assert res.dtype == dtype + ref = ref.to(dtype) + atol = 1e-4 * reduce_dim + rtol = RESOLUTION[dtype] + torch.testing.assert_close(res, ref, atol=atol, rtol=rtol, equal_nan=equal_nan) + + +@pytest.mark.parametrize("M, N, K, col_a, col_b, rhs_scale, mxfp_type, normal_type, num_warps, mma, kpack", + [(M, N, K, col_a, col_b, rhs_scale, mxfp_type, normal_type, 4, mma, kpack) + for M, N, K in itertools.product([32, 64, 128], [32, 64, 128], [64, 128]) + for col_a, col_b in itertools.product([True, False], repeat=2) + for rhs_scale in [False, True] + for mxfp_type in ["e2m1", "e4m3", "e5m2"] + for normal_type in ["e4m3", "e5m2", "bf16", "fp16"] + for mma in (mma_nonk_sizes if is_hip() else [16]) + for kpack in ([1, 2] if is_hip() else [1])]) +def test_scaled_dot(M, N, K, col_a, col_b, rhs_scale, mxfp_type, normal_type, num_warps, mma, kpack, device): + if is_cuda(): + cc = torch.cuda.get_device_capability() + if cc < (8, 9): + pytest.skip("float8e4nv not supported on CUDA < 8.9") + if is_hip(): + if not is_hip_cdna(): + pytest.skip("scaled_dot only implemented for HIP CDNA") + if "e4m3" in (mxfp_type, normal_type): + if not (is_hip_cdna3() or is_hip_cdna4()): + pytest.skip(f"scaled_dot({mxfp_type}, {normal_type}) only implemented for MI300 and MI350") + if mma == 16 and K == 64: + pytest.skip(f"K == {K} too small for mfma {mma} in scaled_dot") + + @triton.jit + def dot_scale_kernel(a_base, stride_a0, stride_a1, a_scale, b_base, stride_b0, stride_b1, b_scale, out, + BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, type_a: tl.constexpr, + type_b: tl.constexpr): + DIV_FACTOR_A: tl.constexpr = 2 if type_a == "e2m1" else 1 + DIV_FACTOR_B: tl.constexpr = 2 if type_b == "e2m1" else 1 + PACKED_BLOCK_K_A: tl.constexpr = BLOCK_K // DIV_FACTOR_A + PACKED_BLOCK_K_B: tl.constexpr = BLOCK_K // DIV_FACTOR_B + a_ptr = a_base + tl.arange(0, BLOCK_M)[:, None] * stride_a0 + tl.arange(0, + PACKED_BLOCK_K_A)[None, :] * stride_a1 + b_ptr = b_base + tl.arange(0, PACKED_BLOCK_K_B)[:, None] * stride_b0 + tl.arange(0, + BLOCK_N)[None, :] * stride_b1 + + a = tl.load(a_ptr) + b = tl.load(b_ptr) + SCALE_BLOCK_K: tl.constexpr = BLOCK_K // 32 + if a_scale is not None: + scale_a_ptr = a_scale + tl.arange(0, BLOCK_M)[:, None] * SCALE_BLOCK_K + tl.arange(0, + SCALE_BLOCK_K)[None, :] + a_scale = tl.load(scale_a_ptr) + if b_scale is not None: + scale_b_ptr = b_scale + tl.arange(0, BLOCK_N)[:, None] * SCALE_BLOCK_K + tl.arange(0, + SCALE_BLOCK_K)[None, :] + b_scale = tl.load(scale_b_ptr) + c = tl.dot_scaled(a, a_scale, type_a, b, b_scale, type_b) + out_ptr = out + tl.arange(0, BLOCK_M)[:, None] * BLOCK_N + tl.arange(0, BLOCK_N)[None, :] + tl.store(out_ptr, c.to(tl.bfloat16)) + + @triton.jit + def mxfp_upcast_kernel( + x_ptr, + scale_ptr, + mxfp_ptr, + N, + e_bits: tl.constexpr, + m_bits: tl.constexpr, + to_type: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + ): + # x.shape == (N, 32) for fp8 or (N, 16) for fp4 + # scale.shape == (N,) + # out.shape == (N, 32) + is_fp8: tl.constexpr = e_bits + m_bits == 7 + # fp8: BLOCK_SIZE -> BLOCK_SIZE // 32, 32 + # fp4: BLOCK_SIZE // 2 -> BLOCK_SIZE // 32 , 16 + PARALLEL_DIM: tl.constexpr = BLOCK_SIZE // 32 + LAST_DIM: tl.constexpr = 32 if is_fp8 else 16 + LOAD_SIZE: tl.constexpr = LAST_DIM * PARALLEL_DIM + + offsets = (tl.program_id(0) * LOAD_SIZE + tl.arange(0, PARALLEL_DIM)[:, None] * LAST_DIM + + tl.arange(0, LAST_DIM)[None, :]) + x = tl.load(x_ptr + offsets, mask=offsets < N * LAST_DIM) + + offsets = tl.program_id(0) * PARALLEL_DIM + tl.arange(0, PARALLEL_DIM)[:, None] + scale = tl.load(scale_ptr + offsets, mask=offsets < N) + tl.static_assert(scale.dtype == tl.uint8) + tl.static_assert(x.dtype == tl.uint8) + + if to_type == tl.bfloat16: + upcasted_scale = (scale.to(tl.uint16) << 7).to(tl.bfloat16, bitcast=True) + else: + tl.static_assert(to_type == tl.float16) + scale_fp32 = (scale.to(tl.uint32) << 23).to(tl.float32, bitcast=True) + upcasted_scale = scale_fp32.to(tl.float16) + + to_e_bits: tl.constexpr = 8 if to_type == tl.bfloat16 else 5 + to_m_bits: tl.constexpr = 7 if to_type == tl.bfloat16 else 10 + if is_fp8: + if e_bits == 5 and m_bits == 2: + x_f8 = x.to(tl.float8e5, bitcast=True) + upcasted_x = x_f8.to(to_type) + # Preserve infs and nans. FIXME Fp8E5M2_to_Bf16 doesn't preserve them! + non_finite_mask: tl.constexpr = ((1 << e_bits) - 1) << m_bits + non_finite_mask_16bit: tl.constexpr = ((1 << to_e_bits) - 1) << to_m_bits + upcasted_x = tl.where( + x & non_finite_mask == non_finite_mask, + (upcasted_x.to(tl.uint16, bitcast=True) | non_finite_mask_16bit).to(to_type, bitcast=True), + upcasted_x, + ) + else: + tl.static_assert(e_bits == 4 and m_bits == 3) + x_f8 = x.to(tl.float8e4nv, bitcast=True) + upcasted_x = x_f8.to(to_type) + else: + to_bias: tl.constexpr = 127 if to_type == tl.bfloat16 else 15 + to_point5: tl.constexpr = 16128 if to_type == tl.bfloat16 else 0x3800 + # e2m1 + em0 = x & 0x7 + em1 = x & 0x70 + x0 = (em0.to(tl.uint16) << (to_m_bits - 1)) | ((x & 0x8).to(tl.uint16) << 12) + x1 = (em1.to(tl.uint16) << (to_m_bits - 1 - 4)) | ((x & 0x80).to(tl.uint16) << 8) + # Three cases: + # 1) x is normal and non-zero: Correct bias + x0 = tl.where((em0 & 0x6) != 0, x0 + ((to_bias - 1) << to_m_bits), x0) + x1 = tl.where((em1 & 0x60) != 0, x1 + ((to_bias - 1) << to_m_bits), x1) + # 2) x is subnormal (x == 0bs001 where s is the sign): Map to +-0.5 in bf16 + x0 = tl.where(em0 == 0x1, to_point5 | (x0 & 0x8000), x0) + x1 = tl.where(em1 == 0x10, to_point5 | (x1 & 0x8000), x1) + # 3) x is zero, do nothing + upcasted_x = tl.interleave(x0, x1).to(to_type, bitcast=True) + # Multiplication preserves infs and NaNs in upcasted_x + mxfp = upcasted_x * upcasted_scale + # If scale is NaN, we encode it as an inf, so we need to correct for that + mxfp = tl.where(scale == 0xFF, float("nan"), mxfp) + + offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + tl.store(mxfp_ptr + offsets, tl.ravel(mxfp), mask=offsets < N * 32) + + def dot_scale_ref(x, scale_x, y, scale_y, type_x, type_y): + + def upcast(v, scale, type, comp_dtype, transposed): + if scale is None: + type = { + "e4m3": torch.float8_e4m3fn, + "e5m2": torch.float8_e5m2, + "bf16": torch.bfloat16, + "fp16": torch.float16, + }[type] + return v.view(type).to(comp_dtype) + e_bits, m_bits = {"e2m1": (2, 1), "e4m3": (4, 3), "e5m2": (5, 2)}[type] + # Packing is always on the K dimension so we transpose before upcasting then transpose back. + if transposed: + v = v.mT.contiguous() + v = v.contiguous() + if v.device.type == "cpu": + v_upcast = upcast_mxfp_cpu(v, scale, type, comp_dtype) + assert v_upcast.isfinite().all() + return v_upcast.mT if transposed else v_upcast + v_upcast = v.new_empty(scale.shape[:-1] + (32 * scale.shape[-1], ), dtype=comp_dtype) + N = v_upcast.numel() + BLOCK_SIZE = 512 + grid = ((N + BLOCK_SIZE - 1) // BLOCK_SIZE, ) + comp_dtype = tl.float16 if comp_dtype == torch.float16 else tl.bfloat16 + v_txda = v.to("txda") + scale_txda = scale.to("txda") + v_upcast_txda = v_upcast.to("txda") + mxfp_upcast_kernel[grid](v_txda, scale_txda, v_upcast_txda, scale_txda.numel(), e_bits, m_bits, comp_dtype, BLOCK_SIZE, + num_warps=num_warps) + with torch.no_grad(): + v_upcast.copy_(v_upcast_txda.cpu()) + + assert v_upcast.isfinite().all() + if transposed: + v_upcast = v_upcast.mT + return v_upcast + + # Upcast to fp16 if one of the input is fp16 + comp_dtype = torch.float16 if "fp16" in (type_x, type_y) else torch.bfloat16 + + x_upcast = upcast(x, scale_x, type_x, comp_dtype, False) + y_upcast = upcast(y, scale_y, type_y, comp_dtype, True) + + class AccumulateInFp32: + + def __enter__(self): + self.prev_value = torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction + torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = False + + def __exit__(self, exc_type, exc_val, exc_tb): + torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = self.prev_value + + with AccumulateInFp32(): + return torch.matmul(x_upcast, y_upcast) + + comp_dtype = torch.float16 if normal_type == "fp16" else torch.bfloat16 + # The max exponent we use to initialize data in the x/y and associated scale tensor to avoid + # overflow when scaling. + comp_dtype_max_exp = 6 if normal_type == "fp16" else 15 + + torch.manual_seed(0) + + def make_arg(shape, ty, col_major=False): + if col_major: + shape = shape[:-2] + (shape[-1], shape[-2]) + if ty == "bf16" or ty == "fp16": + ret = torch.randn(shape, dtype=comp_dtype, device="cpu") + # Clamp to avoid relative error issues + ret.clamp_(-2**comp_dtype_max_exp, 2**comp_dtype_max_exp - 1) + else: + if is_hip_cdna4(): + # On other chips, the A/B operands are upcasted to fp16/bf16 + # before matmul, which has larger range to avoid overflow. + # On MI350, we use the V_MFMA_*_F8F6F4 instructions to + # directly calculate matmul on F8F6F4 data. So we need + # to narrow down the range of input to avoid overflow. + ret = torch.randint(20, 40, shape, dtype=torch.uint8, device="cpu") + else: + ret = torch.randint(256, shape, dtype=torch.uint8, device="cpu") + if col_major: + ret = ret.mT + return ret + + type_a = normal_type if rhs_scale else mxfp_type + type_b = mxfp_type if rhs_scale else normal_type + + DIV_FACTOR_A = 2 if type_a == "e2m1" else 1 + DIV_FACTOR_B = 2 if type_b == "e2m1" else 1 + x = make_arg((M, K // DIV_FACTOR_A), type_a, col_major=col_a) + y = make_arg((K // DIV_FACTOR_B, N), type_b, col_major=col_b) + + min_scale, max_scale = (0, 142) if comp_dtype == torch.bfloat16 else (124, 131) + scale_x = torch.randint(min_scale, max_scale + 1, (M, K // 32), dtype=torch.uint8, device="cpu") + scale_y = torch.randint(min_scale, max_scale + 1, (N, K // 32), dtype=torch.uint8, device="cpu") + if rhs_scale: + scale_x = None + else: + scale_y = None + + def make_finite(x, dtype): + # e5m2 has too many non-finite values when sampled uniformly (1 / 32) and + # Fp8E5M2_to_Bf16 doesn't preserve NaNs (fixme) + if dtype not in ("e5m2", "e4m3"): + return x + if dtype == "e5m2" and comp_dtype == torch.float16: + x = x & 0xB + mask = 0x7C if dtype == "e5m2" else 0x7F + finite = torch.arange(x.numel(), device="cpu", dtype=torch.uint8).reshape_as(x) % mask + x_finite = torch.where(x & mask == mask, finite | (0x80 & x), x) + x.copy_(x_finite) + return x + + x = make_finite(x, type_a) + y = make_finite(y, type_b) + + kernel_kwargs = {"num_warps": num_warps} + if is_hip(): + kernel_kwargs["kpack"] = kpack + kernel_kwargs["matrix_instr_nonkdim"] = mma + z = x.new_empty((M, N), dtype=comp_dtype) + x_txda = x.to("txda") + y_txda = y.to("txda") + scale_x_txda = None if scale_x is None else scale_x.to("txda") + scale_y_txda = None if scale_y is None else scale_y.to("txda") + z_txda = z.to("txda") + pgm = dot_scale_kernel[(1, )](x_txda, *x_txda.stride(), scale_x_txda, y_txda, *y_txda.stride(), scale_y_txda, z_txda, M, N, K, type_a, type_b, + **kernel_kwargs) + with torch.no_grad(): + z.copy_(z_txda.cpu()) + z_ref = dot_scale_ref(x, scale_x, y, scale_y, type_a, type_b) + # Bigger tolerance for AMD MI200 devices. + # MI200 devices use reduced precision fp16 and bf16 and flush input and output denormal values + # to zero. Detailed info is at: + # https://pytorch.org/docs/stable/notes/numerical_accuracy.html#reduced-precision-fp16-and-bf16-gemms-and-convolutions-on-amd-instinct-mi200-devices + + # FIXME: If enable orgin assertion, these cases failed due to matmul precision: + # test_scaled_dot[32-64-128-True-True-False-e2m1-e5m2-4-16-1] + # test_scaled_dot[64-64-64-False-False-False-e4m3-e4m3-4-16-1] + # test_scaled_dot[64-64-128-True-True-False-e4m3-fp16-4-16-1] + # test_scaled_dot[128-32-128-True-False-False-e4m3-fp16-4-16-1] + # test_scaled_dot[128-64-128-False-False-True-e4m3-fp16-4-16-1] + # test_scaled_dot[128-128-64-True-True-False-e2m1-e5m2-4-16-1] + + # atol = 2e-4 if is_hip_cdna2() else 1e-5 + # rtol = 2e-2 if is_hip_cdna2() else 1e-2 + # torch.testing.assert_close(z, z_ref, atol=atol, rtol=rtol) + + flaggems_assert_close(z, z_ref, dtype=comp_dtype, reduce_dim=K) + + # make sure ld/st are vectorized + if is_cuda(): + ptx = pgm.asm['ptx'] + if (max(M, N) * K) // (num_warps * 32) >= 4: + assert 'ld.global.v4' in ptx + if M * N // (num_warps * 32) >= 4: + assert 'st.global.v4' in ptx + assert (re.search(r'(mma|wgmma.mma_async).sync.aligned.m\d+n\d+k16(?:.row.col)?.f32.(f|bf)16.(f|bf)16', ptx) + or "tcgen05.mma.cta_group::1.kind::f16" in ptx) + + +if __name__ == "__main__": + M = 32 + N = 64 + K = 128 + col_a = True + col_b = False + rhs_scale = False + mxfp_type = "e5m2" + normal_type = "e4m3" + num_warps = 4 + mma = 16 + kpack = 1 + device = "cpu" + + test_scaled_dot(M, N, K, col_a, col_b, rhs_scale, mxfp_type, normal_type, num_warps, mma, kpack, device) diff --git a/third_party/wafer/examples/test_early_return.py b/third_party/wafer/examples/test_early_return.py new file mode 100755 index 00000000..857ca642 --- /dev/null +++ b/third_party/wafer/examples/test_early_return.py @@ -0,0 +1,64 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + + +@triton.jit +def early_return(in0, out0): + pid = tl.program_id(0) + id = tl.load(in0 + pid) + if id == -1: + return + offs = 1 + tl.arange(0, 4) + out_offs = tl.arange(0, 4) + tl.store(out0 + out_offs, offs) + + +def compile(device): + src = triton.compiler.ASTSource( + fn=early_return, + signature={'in0': '*fp32', 'out0': '*fp32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) + + +def test_return_case(device): + if device == 'cpu': + pass # Wafer driver is selected by conftest.py. + + SIZE = 8 + input = torch.full((SIZE, ), -1, device="cpu", dtype=torch.int32) + output = torch.full((SIZE, ), -1, device="cpu", dtype=torch.int32) + grid = lambda meta: (1, ) + print(output) + input_txda = input.to("txda") + output_txda = output.to("txda") + early_return[grid](input_txda, output_txda) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(input) + print(output) + torch.testing.assert_close(torch.tensor([-1, -1, -1, -1, -1, -1, -1, -1], dtype=torch.int32), output) + + +def test_normal_case(device): + if device == 'cpu': + pass # Wafer driver is selected by conftest.py. + + SIZE = 8 + input = torch.arange(0, SIZE, device="cpu", dtype=torch.int32) + output = torch.full((SIZE, ), -1, device="cpu", dtype=torch.int32) + grid = lambda meta: (1, ) + print(output) + input_txda = input.to("txda") + output_txda = output.to("txda") + early_return[grid](input_txda, output_txda) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(input) + print(output) + torch.testing.assert_close(torch.tensor([1, 2, 3, 4, -1, -1, -1, -1], dtype=torch.int32), output) diff --git a/third_party/wafer/examples/test_embedding.py b/third_party/wafer/examples/test_embedding.py new file mode 100755 index 00000000..8fe93ee4 --- /dev/null +++ b/third_party/wafer/examples/test_embedding.py @@ -0,0 +1,94 @@ +import torch +import torch_txda # noqa: F401 +import math + +import triton +import triton.language as tl + +import pytest +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def embedding_kernel( + out_ptr, # pointer to the output + in_ptr, # pointer to the input + weight_ptr, # pointer to the weights + N: tl.constexpr, # number of columns in X + BLOCK_SIZE: tl.constexpr, +): + pid = tl.program_id(0) + out_ptr += pid * N + in_ptr += pid + + mask = tl.arange(0, BLOCK_SIZE) < N + cols = tl.arange(0, BLOCK_SIZE) + + row_idx = tl.load(in_ptr) + weight_ptr += row_idx * N + embedding_weight = tl.load(weight_ptr + cols, mask, other=0.0) + tl.store(out_ptr + cols, embedding_weight, mask) + + +class Embedding(torch.autograd.Function): + + @staticmethod + def forward(ctx, weight, indices, padding_idx=-1, scale_grad_by_freq=False, sparse=False): + + assert not sparse, "Currently do not support sparse format" + + M = math.prod(indices.shape) + N = weight.shape[-1] + + BLOCK_SIZE = triton.next_power_of_2(N) + indices = indices.contiguous() + weight = weight.contiguous() + output = torch.empty((*indices.shape, N), device=indices.device, dtype=weight.dtype) + + output = output.to(DEVICE) + indices = indices.to(DEVICE) + weight = weight.to(DEVICE) + embedding_kernel[ + M, + ](output, indices, weight, N, BLOCK_SIZE) + output = output.to("cpu") + ctx.M = M + ctx.N = N + ctx.num_weights = weight.shape[0] + ctx.padding_idx = padding_idx + ctx.scale_grad_by_freq = scale_grad_by_freq + ctx.sparse = sparse + ctx.indices = indices + + return output + + +def embedding(weight, indices, padding_idx=-1, scale_grad_by_freq=False, sparse=False): + return Embedding.apply(weight, indices, padding_idx, scale_grad_by_freq, sparse) + + +@pytest.mark.parametrize("M, N, dtype", [ # + (M, N, dtype) for M in [1152] for N in [2048] for dtype in [torch.float32] +]) +def test_embedding(M, N, dtype, device='cpu'): + torch.manual_seed(0) + + weight = torch.rand((M, N), dtype=dtype, device=device) + indices = torch.randint(0, M, [M], dtype=torch.int32, device=device) + + triton_output = embedding(weight, indices) + + # pytorch + torch_embedding = torch.nn.Embedding(M, N, _weight=weight) + torch_output = torch_embedding(indices) + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(torch_output - triton_output))}") + assert torch.allclose(triton_output, torch_output, atol=1e-5, rtol=0) + + +if __name__ == "__main__": + test_embedding(1151, 8192, torch.float32) diff --git a/third_party/wafer/examples/test_exp.py b/third_party/wafer/examples/test_exp.py new file mode 100755 index 00000000..875a23b3 --- /dev/null +++ b/third_party/wafer/examples/test_exp.py @@ -0,0 +1,96 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def exp_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.exp(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def exp_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + exp_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_exp(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton Kernel + output = exp_triton(x) + + # Verify the output + expected = torch.exp(x) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_exp_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton Kernel + output = exp_triton(x) + + # Verify the output + expected = torch.exp(x) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_exp_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_exp2.py b/third_party/wafer/examples/test_exp2.py new file mode 100755 index 00000000..43a72d76 --- /dev/null +++ b/third_party/wafer/examples/test_exp2.py @@ -0,0 +1,98 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def exp2_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.exp2(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def exp2_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + exp2_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_exp2(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = exp2_triton(x) + + # Verify the output + expected = torch.exp2(x) + + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_exp2_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = exp2_triton(x) + + # Verify the output + expected = torch.exp2(x) + + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_exp2_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_fdiv.py b/third_party/wafer/examples/test_fdiv.py new file mode 100755 index 00000000..8ee07965 --- /dev/null +++ b/third_party/wafer/examples/test_fdiv.py @@ -0,0 +1,104 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def fdiv_kernel( + x_ptr, + y_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.fdiv(x, y) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def fdiv_triton(x, y): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + fdiv_kernel[grid]( + x_txda, + y_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_fdiv(size, dtype, device="cpu"): + # Generate random input tensors + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = fdiv_triton(x, y) + + # Verify the output + expected = torch.div(x, y) + + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_fdiv_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input tensors + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = fdiv_triton(x, y) + + # Verify the output + expected = torch.div(x, y) + + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_fdiv_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_flip.py b/third_party/wafer/examples/test_flip.py new file mode 100755 index 00000000..a0bd2329 --- /dev/null +++ b/third_party/wafer/examples/test_flip.py @@ -0,0 +1,51 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import benchmark + +from triton._internal_testing import numpy_random + + +# @pytest.mark.interpreter +# @pytest.mark.parametrize("M, N", [[1, 512], [8, 64], [256, 16], [512, 8]]) +@pytest.mark.parametrize("M, N", [[8, 64]]) +# @pytest.mark.parametrize("dtype_str", ['int32', 'float16', 'float32', 'bfloat16']) +@pytest.mark.parametrize("dtype_str", ['float32']) +def test_flip(M, N, dtype_str, device): + + @triton.jit + def flip_kernel(X, Z, N: tl.constexpr, M: tl.constexpr): + offx = tl.arange(0, M) + offy = tl.arange(0, N) * M + off2d = offx[None, :] + offy[:, None] + x = tl.load(X + off2d) + x = tl.flip(x, dim=1) + tl.store(Z + off2d, x) + + x = numpy_random((N, M), dtype_str=dtype_str) + x = torch.from_numpy(x) + y = torch.flip(x, (1, )) + z = torch.empty_like(x, device="cpu") + x_txda = x.to("txda") + z_txda = z.to("txda") + flip_kernel[(1, )](x_txda, z_txda, N, M, num_warps=8) + with torch.no_grad(): + z.copy_(z_txda.cpu()) + assert (y == z).all(), (y, z) + + +if __name__ == "__main__": + # Test with different sizes and data types + test_sizes = [(8, 64)] + dtypes = ['float32'] + + for M, N in test_sizes: + for dtype in dtypes: + print(f"\nTesting flip with M={M}, N={N}, dtype={dtype}") + try: + test_flip(M, N, dtype, "cpu") + print("Test passed!") + except Exception as e: + print(f"Test failed: {str(e)}") diff --git a/third_party/wafer/examples/test_floor.py b/third_party/wafer/examples/test_floor.py new file mode 100755 index 00000000..7ac83f27 --- /dev/null +++ b/third_party/wafer/examples/test_floor.py @@ -0,0 +1,100 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def floor_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.floor(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def floor_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + floor_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_floor(size, dtype, device="cpu"): + # Generate random input tensor + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = floor_triton(x) + + # Verify the result + expected = torch.floor(x) + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(expected - output))}") + assert torch.allclose(output, expected, atol=1e-5, rtol=0) + + +@benchmark.measure() +def benchmark_floor_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = floor_triton(x) + + # Verify the output + expected = torch.floor(x) + assert torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_floor_triton(size, torch.float32, provider="triton") diff --git a/third_party/wafer/examples/test_fma.py b/third_party/wafer/examples/test_fma.py new file mode 100755 index 00000000..eea8a917 --- /dev/null +++ b/third_party/wafer/examples/test_fma.py @@ -0,0 +1,112 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def fma_kernel( + x_ptr, + y_ptr, + z_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + z = tl.load(z_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.fma(x, y, z) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def fma_triton(x, y, z): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + y_txda = y.to("txda") + z_txda = z.to("txda") + output_txda = output.to("txda") + fma_kernel[grid]( + x_txda, + y_txda, + z_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [1024] for dtype in [torch.float32] +]) +def test_fma(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + z = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton implementation + output = fma_triton(x, y, z) + + # Verify the output + expected = torch.mul(x, y) + z + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(expected - output))}") + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_fma_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + z = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton implementation + output = fma_triton(x, y, z) + + # Verify the output + expected = torch.mul(x, y) + z + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_fma_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_fp8_conversion.py b/third_party/wafer/examples/test_fp8_conversion.py new file mode 100644 index 00000000..f3c846a9 --- /dev/null +++ b/third_party/wafer/examples/test_fp8_conversion.py @@ -0,0 +1,29 @@ +"""Exercise the hardware E5M2 -> FP16 CRT path for all raw encodings.""" +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def convert_e5m2_kernel(src, dst): + indices = tl.arange(0, 256) + tl.store(dst + indices, tl.load(src + indices).to(tl.float16)) + + +def test_e5m2_to_fp16_all_encodings(device): + raw = torch.arange(256, dtype=torch.uint8, device="cpu") + out = torch.empty(256, dtype=torch.float16, device="cpu") + raw_txda = raw.to("txda") + out_txda = out.to("txda") + convert_e5m2_kernel[(1,)](triton.reinterpret(raw_txda, tl.float8e5), out_txda) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + expected = raw.cpu().view(torch.float8_e5m2).to(torch.float16) + torch.testing.assert_close(out.cpu(), expected, rtol=0, atol=0, equal_nan=True) + # Torch may quiet NaNs. Test the CRT's payload/sign preservation separately. + bits = out.cpu().view(torch.uint16).to(torch.int32) + codes = raw.cpu().to(torch.int32) + assert torch.equal(bits >> 15, codes >> 7) + nan = ((codes & 0x7c) == 0x7c) & ((codes & 3) != 0) + assert torch.equal(bits[nan] & 1023, (codes[nan] & 3) * 256) diff --git a/third_party/wafer/examples/test_gather.py b/third_party/wafer/examples/test_gather.py new file mode 100755 index 00000000..5f06aa7e --- /dev/null +++ b/third_party/wafer/examples/test_gather.py @@ -0,0 +1,53 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def gather_test_kernel(src_ptr, idx_ptr, out_ptr, axis: tl.constexpr, src_dim0: tl.constexpr, src_dim1: tl.constexpr, + src_stride0: tl.constexpr, src_stride1: tl.constexpr, idx_dim0: tl.constexpr, + idx_dim1: tl.constexpr, idx_stride0: tl.constexpr, idx_stride1: tl.constexpr, + out_dim0: tl.constexpr, out_dim1: tl.constexpr, out_stride0: tl.constexpr, + out_stride1: tl.constexpr): + src_offs = (tl.arange(0, src_dim0)[:, None] * src_stride0 + tl.arange(0, src_dim1)[None, :] * src_stride1) + src = tl.load(src_ptr + src_offs) + + idx_offs = (tl.arange(0, idx_dim0)[:, None] * idx_stride0 + tl.arange(0, idx_dim1)[None, :] * idx_stride1) + idx = tl.load(idx_ptr + idx_offs) + + out = tl.gather(src, idx, axis) + + out_offs = (tl.arange(0, out_dim0)[:, None] * out_stride0 + tl.arange(0, out_dim1)[None, :] * out_stride1) + tl.store(out_ptr + out_offs, out) + + +@pytest.mark.interpreter +@pytest.mark.parametrize("src_shape, indices_shape, axis", [ + ([4, 4], [8, 4], 0), + # ([128, 64], [256, 64], 0), + # ([128, 64], [128, 128], 1), +]) +def test_gather(src_shape, indices_shape, axis, device): + + def triton_gather(src: torch.Tensor, axis: int, indices: torch.Tensor): + output = torch.empty(indices.shape, dtype=src.dtype, device="cpu") + + src_txda = src.to("txda") + indices_txda = indices.to("txda") + output_txda = output.to("txda") + gather_test_kernel[(1, )](src_txda, indices_txda, output_txda, axis, src_txda.shape[0], + src_txda.shape[1], src_txda.stride(0), src_txda.stride(1), indices_txda.shape[0], indices_txda.shape[1], + indices_txda.stride(0), indices_txda.stride(1), output_txda.shape[0], output_txda.shape[1], + output_txda.stride(0), output_txda.stride(1)) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + src = torch.randn(src_shape, device="cpu") + indices = torch.randint(0, src.shape[axis], indices_shape, device="cpu") + ref = torch.gather(src, axis, indices) + result = triton_gather(src, axis, indices) + torch.testing.assert_close(result, ref, rtol=0, atol=0) diff --git a/third_party/wafer/examples/test_histogram.py b/third_party/wafer/examples/test_histogram.py new file mode 100755 index 00000000..4a375040 --- /dev/null +++ b/third_party/wafer/examples/test_histogram.py @@ -0,0 +1,47 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +# @pytest.mark.interpreter +# @pytest.mark.parametrize("M, N", [[2048, 2], [1024, 8], [1024, 128], [256, 512], [32, 512], [8, 512], [8, 2]]) +@pytest.mark.parametrize("M, N", [[1024, 8]]) +def test_histogram(M, N, device): + + @triton.jit + def histogram_kernel(x_ptr, z_ptr, M: tl.constexpr, N: tl.constexpr): + offset1 = tl.arange(0, M) + offset2 = tl.arange(0, N) + x = tl.load(x_ptr + offset1) + z = tl.histogram(x, N) + bias = tl.full([M, N], 1, dtype=tl.int32) + # check that histogram produces object compatible with broadcasting + biased = z + bias + tl.store(z_ptr + offset2, z) + + torch.manual_seed(17) + x = torch.randint(0, N, (M, ), device="cpu", dtype=torch.int32) + z = torch.empty(N, dtype=torch.int32, device="cpu") + # torch.histc does not work when the input type is not float and the device is CPU + # https://github.com/pytorch/pytorch/issues/74236 + # This is a workload by converting the input to float + z_torch = torch.histc(x.float(), bins=N, min=0, max=N - 1) + x_txda = x.to("txda") + z_txda = z.to("txda") + histogram_kernel[(1, )](x_txda, z_txda, M=M, N=N) + with torch.no_grad(): + z.copy_(z_txda.cpu()) + assert (z_torch == z).all() + + +if __name__ == "__main__": + test_histogram(1024, 8, 'cpu') + # test_histogram_2d(32, 16, 8, 'cpu') + # test_histogram(2048, 2, 'cpu') + # test_histogram(1024, 128, 'cpu') + # test_histogram(256, 512, 'cpu') + # test_histogram(32, 512, 'cpu') + # test_histogram(8, 512, 'cpu') + # test_histogram(8, 2, 'cpu') diff --git a/third_party/wafer/examples/test_layernorm.py b/third_party/wafer/examples/test_layernorm.py new file mode 100755 index 00000000..a1e7323d --- /dev/null +++ b/third_party/wafer/examples/test_layernorm.py @@ -0,0 +1,183 @@ +# This is the Layer Norm forward pass from the Triton tutorial found here: +# https://github.com/triton-lang/triton/blob/main/python/tutorials/05-layer-norm.py + +# %% +# Motivations +# ----------- +# +# The *LayerNorm* operator was first introduced in [BA2016]_ as a way to improve the performance +# of sequential models (e.g., Transformers) or neural networks with small batch size. +# It takes a vector :math:`x` as input and produces a vector :math:`y` of the same shape as output. +# The normalization is performed by subtracting the mean and dividing by the standard deviation of :math:`x`. +# After the normalization, a learnable linear transformation with weights :math:`w` and biases :math:`b` is applied. +# The forward pass can be expressed as follows: +# +# .. math:: +# y = \frac{ x - \text{E}[x] }{ \sqrt{\text{Var}(x) + \epsilon} } * w + b +# +# where :math:`\epsilon` is a small constant added to the denominator for numerical stability. +# Let’s first take a look at the forward pass implementation. + +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import pytest +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def _layer_norm_fwd_fused( + X, # pointer to the input + Y, # pointer to the output + W, # pointer to the weights + B, # pointer to the biases + Mean, # pointer to the mean + Rstd, # pointer to the 1/std + stride, # how much to increase the pointer when moving by 1 row + N, # number of columns in X + eps, # epsilon to avoid division by zero + BLOCK_SIZE: tl.constexpr, +): + # Map the program id to the row of X and Y it should compute. + row = tl.program_id(0) + Y += row * stride + X += row * stride + # Compute mean + mean = 0 + _mean = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + a = tl.load(X + cols, mask=cols < N, other=0.).to(tl.float32) + _mean += a + mean = tl.sum(_mean, axis=0) / N + # Compute variance + _var = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + x = tl.load(X + cols, mask=cols < N, other=0.).to(tl.float32) + x = tl.where(cols < N, x - mean, 0.) + _var += x * x + var = tl.sum(_var, axis=0) / N + rstd = 1 / tl.sqrt(var + eps) + # Write mean / rstd + tl.store(Mean + row, mean) + tl.store(Rstd + row, rstd) + # Normalize and apply linear transformation + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + mask = cols < N + w = tl.load(W + cols, mask=mask) + b = tl.load(B + cols, mask=mask) + x = tl.load(X + cols, mask=mask, other=0.).to(tl.float32) + x_hat = (x - mean) * rstd + y = x_hat * w + b + # Write output + tl.store(Y + cols, y, mask=mask) + + +class LayerNorm(torch.autograd.Function): + + @staticmethod + def forward(ctx, x, normalized_shape, weight, bias, eps, device): + # allocate output + y = torch.empty_like(x) + # reshape input data into 2D tensor + x_arg = x.reshape(-1, x.shape[-1]) + M, N = x_arg.shape + mean = torch.empty((M, ), dtype=torch.float32, device=device) + rstd = torch.empty((M, ), dtype=torch.float32, device=device) + # Less than 64KB per feature: enqueue fused kernel + MAX_FUSED_SIZE = 65536 // x.element_size() + BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(N)) + if N > BLOCK_SIZE: + raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.") + # heuristics for number of warps + num_warps = min(max(BLOCK_SIZE // 256, 1), 8) + + x_arg_dev = x_arg.to(DEVICE) + y_dev = y.to(DEVICE) + weight_dev = weight.to(DEVICE) + bias_dev = bias.to(DEVICE) + mean_dev = mean.to(DEVICE) + rstd_dev = rstd.to(DEVICE) + # enqueue kernel + # _layer_norm_fwd_fused[(M, )]( # + # x_arg, y, weight, bias, mean, rstd, # + # x_arg.stride(0), N, eps, # + # BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps, num_ctas=1) + + _layer_norm_fwd_fused[(M, )]( # + x_arg_dev, y_dev, weight_dev, bias_dev, mean_dev, rstd_dev, # + x_arg.stride(0), N, eps, # + BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps, num_ctas=1) + x = x_arg_dev.to("cpu") + y = y_dev.to("cpu") + mean = mean_dev.cpu() + rstd = rstd_dev.cpu() + + ctx.save_for_backward(x, weight, bias, mean, rstd) + ctx.BLOCK_SIZE = BLOCK_SIZE + ctx.num_warps = num_warps + ctx.eps = eps + return y + + +@pytest.mark.parametrize("M, N, dtype, eps", [ # + (M, N, dtype, eps) for M in [1151] for N in [8192] for dtype in [torch.float16] for eps in [1e-5] +]) +def test_layer_norm(M, N, dtype, eps, device): + layer_norm = LayerNorm.apply + # create data + x_shape = (M, N) + w_shape = (x_shape[-1], ) + weight = torch.rand(w_shape, dtype=dtype, device="cpu", requires_grad=False) + bias = torch.rand(w_shape, dtype=dtype, device="cpu", requires_grad=False) + x = -2.3 + 0.5 * torch.randn(x_shape, dtype=dtype, device="cpu") + dy = .1 * torch.randn_like(x) + x.requires_grad_(False) + + # forward pass + y_tri = layer_norm(x, w_shape, weight, bias, eps, device) + # TODO We can't compare against Torch layer_norm since it doesn't support float16 on CPU + #y_ref = torch.nn.functional.layer_norm(x, w_shape, weight, bias, eps).to(dtype) + + print(y_tri) + #print(y_ref) + + # compare + #assert torch.allclose(y_tri, y_ref, atol=1e-2, rtol=0) + + +@benchmark.measure() +def bench_layernorm(size, provider): + layer_norm = LayerNorm.apply + device = 'cpu' + eps = 1e-5 + # dtype = torch.float16 + dtype = torch.float32 + x_shape = (size, size) + w_shape = (x_shape[-1], ) + weight = torch.rand(w_shape, dtype=dtype, device="cpu", requires_grad=False) + bias = torch.rand(w_shape, dtype=dtype, device="cpu", requires_grad=False) + x = -2.3 + 0.5 * torch.randn(x_shape, dtype=dtype, device="cpu") + dy = .1 * torch.randn_like(x) + x.requires_grad_(False) + # forward pass + y_tri = layer_norm(x, w_shape, weight, bias, eps, device) + y_ref = torch.nn.functional.layer_norm(x, w_shape, weight, bias, eps).to(dtype) + + print(y_tri) + print(y_ref) + + # compare + assert torch.allclose(y_tri, y_ref, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for X in [2**i for i in range(10, 13, 1)]: + for provider in ['triton']: + bench_layernorm(X, provider) diff --git a/third_party/wafer/examples/test_libdevice.py b/third_party/wafer/examples/test_libdevice.py new file mode 100755 index 00000000..9272cb02 --- /dev/null +++ b/third_party/wafer/examples/test_libdevice.py @@ -0,0 +1,207 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + +from triton.language.extra import libdevice + + +@pytest.mark.parametrize("dtype_str", ["float32"]) +@pytest.mark.parametrize("size", [128, 4]) +@pytest.mark.parametrize( + "libdevice_fn, torch_special_fn", + [("tanh", "tanh"), ("pow", "pow"), ("fmod", "fmod"), ("isnan", "isnan"), + ("isinf", "isinf"), ("finitef", "isfinite"), ("ceil", "ceil"), ("floor", "floor"), ("rint", "round"), + ("trunc", "trunc"), # Add the new function test + ], +) +def test_special(dtype_str, size, libdevice_fn, torch_special_fn, device): + SIZE = size + dtype = getattr(torch, dtype_str) + + if torch_special_fn in ["pow", "fmod"]: + x = torch.randn((SIZE, ), dtype=dtype, device="cpu") + y = torch.randn((SIZE, ), dtype=dtype, device="cpu") + y_exp = torch.empty((SIZE, ), dtype=dtype, device="cpu") + if torch_special_fn == "pow": + y_ref = torch.pow(x, y) + elif torch_special_fn == "fmod": + y_ref = torch.fmod(x, y) + elif torch_special_fn in ["isnan", "isinf", "isfinite"]: # Add isfinite to this condition + x = torch.randn((SIZE, ), dtype=dtype, device="cpu") + # Set some element as nan&-nan + x[SIZE // 4] = float('inf') + x[SIZE // 2] = float('-inf') + x[3 * SIZE // 4] = float('nan') + + # Use bool as return value + y_exp = torch.empty((SIZE, ), dtype=torch.bool, device="cpu") + if torch_special_fn == "isnan": + y_ref = torch.isnan(x) + elif torch_special_fn == "isinf": + y_ref = torch.isinf(x) + else: # isfinite + y_ref = torch.isfinite(x) + elif torch_special_fn in ["ceil", "floor", "trunc", "round"]: + # For ceil, floor, and trunc, we can use the same input + # as they are unary operations. + x = torch.randn((SIZE, ), dtype=dtype, device="cpu") + y_exp = torch.empty((SIZE, ), dtype=dtype, device="cpu") + if torch_special_fn == "ceil": + y_ref = torch.ceil(x) + elif torch_special_fn == "floor": + y_ref = torch.floor(x) + elif torch_special_fn == "trunc": + y_ref = torch.trunc(x) + elif torch_special_fn == "round": + y_ref = torch.round(x) + else: + x = torch.randn((SIZE, ), dtype=dtype, device="cpu") + y_exp = torch.empty((SIZE, ), dtype=dtype, device="cpu") + if torch_special_fn == "tanh": + y_ref = torch.tanh(x) + else: + y_ref = getattr(torch.special, torch_special_fn)(x) + + @triton.jit + def kernel_pow(x_ptr, y_ptr, out_ptr, SIZE: tl.constexpr): + off = tl.arange(0, SIZE) + x = tl.load(x_ptr + off) + y = tl.load(y_ptr + off) + res = libdevice.pow(x, y) + tl.store(out_ptr + off, res) + + @triton.jit + def kernel_fmod(x_ptr, y_ptr, out_ptr, SIZE: tl.constexpr): + off = tl.arange(0, SIZE) + x = tl.load(x_ptr + off) + y = tl.load(y_ptr + off) + res = libdevice.fmod(x, y) + tl.store(out_ptr + off, res) + + @triton.jit + def kernel_rint(in_p, out_p, SIZE: tl.constexpr): + off = tl.arange(0, SIZE) + x = tl.load(in_p + off) + # Get rounded result + res = libdevice.rint(x) + tl.store(out_p + off, res) + + @triton.jit + def kernel_unary(in_p, out_p, fn: tl.constexpr, SIZE: tl.constexpr): + off = tl.arange(0, SIZE) + x = tl.load(in_p + off) + res = getattr(libdevice, fn)(x) + tl.store(out_p + off, res) + + @triton.jit + def kernel_isnan(in_p, out_p, SIZE: tl.constexpr): + off = tl.arange(0, SIZE) + x = tl.load(in_p + off) + # Get bool result + res = libdevice.isnan(x) + tl.store(out_p + off, res) + + @triton.jit + def kernel_isinf(in_p, out_p, SIZE: tl.constexpr): + off = tl.arange(0, SIZE) + x = tl.load(in_p + off) + # Get bool result + res = libdevice.isinf(x) + tl.store(out_p + off, res) + + @triton.jit + def kernel_finitef(in_p, out_p, SIZE: tl.constexpr): + off = tl.arange(0, SIZE) + x = tl.load(in_p + off) + # Get bool result + res = libdevice.finitef(x) + tl.store(out_p + off, res) + + if torch_special_fn == "pow": + x_txda = x.to("txda") + y_txda = y.to("txda") + y_exp_txda = y_exp.to("txda") + kernel_pow[(1, )](x_txda, y_txda, y_exp_txda, SIZE=SIZE, num_warps=4, num_ctas=1) + with torch.no_grad(): + y_exp.copy_(y_exp_txda.cpu()) + elif torch_special_fn == "round": + x_txda = x.to("txda") + y_exp_txda = y_exp.to("txda") + kernel_rint[(1, )](x_txda, y_exp_txda, SIZE=SIZE, num_warps=4, num_ctas=1) + with torch.no_grad(): + y_exp.copy_(y_exp_txda.cpu()) + elif torch_special_fn == "fmod": + x_txda = x.to("txda") + y_txda = y.to("txda") + y_exp_txda = y_exp.to("txda") + kernel_fmod[(1, )](x_txda, y_txda, y_exp_txda, SIZE=SIZE, num_warps=4, num_ctas=1) + with torch.no_grad(): + y_exp.copy_(y_exp_txda.cpu()) + elif torch_special_fn == "isnan": + x_txda = x.to("txda") + y_exp_txda = y_exp.to("txda") + kernel_isnan[(1, )](x_txda, y_exp_txda, SIZE=SIZE, num_warps=4, num_ctas=1) + with torch.no_grad(): + y_exp.copy_(y_exp_txda.cpu()) + elif torch_special_fn == "isinf": + x_txda = x.to("txda") + y_exp_txda = y_exp.to("txda") + kernel_isinf[(1, )](x_txda, y_exp_txda, SIZE=SIZE, num_warps=4, num_ctas=1) + with torch.no_grad(): + y_exp.copy_(y_exp_txda.cpu()) + elif torch_special_fn == "isfinite": + x_txda = x.to("txda") + y_exp_txda = y_exp.to("txda") + kernel_finitef[(1, )](x_txda, y_exp_txda, SIZE=SIZE, num_warps=4, num_ctas=1) + with torch.no_grad(): + y_exp.copy_(y_exp_txda.cpu()) + else: + x_txda = x.to("txda") + y_exp_txda = y_exp.to("txda") + kernel_unary[(1, )](x_txda, y_exp_txda, fn=libdevice_fn, SIZE=SIZE, num_warps=4, num_ctas=1) + with torch.no_grad(): + y_exp.copy_(y_exp_txda.cpu()) + + torch.testing.assert_close(y_ref, y_exp, equal_nan=True) + + +def test_libdevice_rename(device): + + @triton.jit + def triton_copy(in_ptr, out_ptr, BLOCK_SIZE: tl.constexpr): + offsets = tl.arange(0, BLOCK_SIZE) + data = tl.load(in_ptr + offsets) + tl.store(out_ptr + offsets, data) + + BLOCK_SIZE = 256 + inp = torch.randn(BLOCK_SIZE, device="cpu") + out = torch.empty_like(inp) + + inp_txda = inp.to("txda") + out_txda = out.to("txda") + triton_copy[(1, )](inp_txda, out_txda, BLOCK_SIZE) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + torch.testing.assert_close(out, inp) + + +def test_libdevice_erf(device): + """Exercise the extern wrapper, independently of tl.erf's frontend entry.""" + @triton.jit + def erf_kernel(inp, out, BLOCK: tl.constexpr): + indices = tl.arange(0, BLOCK) + values = tl.load(inp + indices) + tl.store(out + indices, libdevice.erf(values)) + + values = torch.tensor([-3.0, -1.5, -0.5, -0.0, 0.0, 0.5, 1.5, 3.0], + dtype=torch.float32, device="cpu") + output = torch.empty_like(values) + values_txda = values.to("txda") + output_txda = output.to("txda") + erf_kernel[(1,)](values_txda, output_txda, BLOCK=8) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch.erf(values), atol=1e-3, rtol=1e-3) diff --git a/third_party/wafer/examples/test_load6d.py b/third_party/wafer/examples/test_load6d.py new file mode 100755 index 00000000..baa07b35 --- /dev/null +++ b/third_party/wafer/examples/test_load6d.py @@ -0,0 +1,76 @@ +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import pytest +import benchmark + + +@triton.jit +def six_dim_load(T_ptr, output_ptr, B, C, D1, D2, D3, D4, stride_b, stride_c, stride_d1, stride_d2, stride_d3, + stride_d4, N: tl.constexpr): + + pid = tl.program_id(0) + b = pid // (C * D1 * D2 * D3 * D4) + remaining = pid % (C * D1 * D2 * D3 * D4) + + c = remaining // (D1 * D2 * D3 * D4) + remaining %= (D1 * D2 * D3 * D4) + + d1 = remaining // (D2 * D3 * D4) + remaining %= (D2 * D3 * D4) + + d2 = remaining // (D3 * D4) + remaining %= (D3 * D4) + + d3 = remaining // D4 + d4 = remaining % D4 + + off_b = b * N + tl.arange(0, N)[:, None, None, None, None, None] + off_c = c * N + tl.arange(0, N)[None, :, None, None, None, None] + off_d1 = d1 * N + tl.arange(0, N)[None, None, :, None, None, None] + off_d2 = d2 * N + tl.arange(0, N)[None, None, None, :, None, None] + off_d3 = d3 * N + tl.arange(0, N)[None, None, None, None, :, None] + off_d4 = d4 * N + tl.arange(0, N)[None, None, None, None, None, :] + + mask_b = off_b < B # Shape: [B,1,1,1,1,1] + mask_c = off_c < C # Shape: [1,C,1,1,1,1] + mask_d1 = off_d1 < D1 # Shape: [1,1,D1,1,1,1] + mask_d2 = off_d2 < D2 # Shape: [1,1,1,D2,1,1] + mask_d3 = off_d3 < D3 # Shape: [1,1,1,1,D3,1] + mask_d4 = off_d4 < D4 # Shape: [1,1,1,1,1,D4] + + final_mask = (mask_b & mask_c & mask_d1 & mask_d2 & mask_d3 & mask_d4) + + global_idx = (off_b * stride_b + off_c * stride_c + off_d1 * stride_d1 + off_d2 * stride_d2 + off_d3 * stride_d3 + + off_d4 * stride_d4) + + data = tl.load(T_ptr + global_idx, mask=final_mask) + tl.store(output_ptr + global_idx, data, mask=final_mask) + + +@pytest.mark.parametrize("N", [2, 4]) +def test_triton_six_dim_load(N: int): + shape = (N, N, N, N, N, N) + input = torch.arange(0, N**6, device="cpu", dtype=torch.float32).reshape(shape).contiguous() + + B, C, D1, D2, D3, D4 = input.shape + output = torch.empty_like(input) + + strides = input.stride() + stride_b, stride_c, stride_d1, stride_d2, stride_d3, stride_d4 = strides[0], strides[1], strides[2], strides[ + 3], strides[4], strides[5] + + input_txda = input.to("txda") + output_txda = output.to("txda") + six_dim_load[(1, )](input_txda, output_txda, B, C, D1, D2, D3, D4, stride_b, stride_c, stride_d1, stride_d2, stride_d3, + stride_d4, N=N) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + torch.testing.assert_close(input, output, rtol=0, atol=0) + + +if __name__ == "__main__": + N = 2 + test_triton_six_dim_load(N) diff --git a/third_party/wafer/examples/test_load_2d_tensor_block.py b/third_party/wafer/examples/test_load_2d_tensor_block.py new file mode 100755 index 00000000..f0dceb66 --- /dev/null +++ b/third_party/wafer/examples/test_load_2d_tensor_block.py @@ -0,0 +1,79 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +""" + +|-----|-----|-----|-----| +| | | | | +|-----|-----|-----|-----| +| | | | | +|-----|-----|-----|-----| + +Each instance loads BLOCK_SIZE_ROW * BLOCK_SIZE_COL +""" + + +@triton.jit +def kernel( + x_ptr, + y_ptr, + n_rows, + n_cols, + stride_0, + stride_1, + BLOCK_SIZE_ROW: tl.constexpr, + BLOCK_SIZE_COL: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + pid1 = tl.program_id(axis=1) + + input_ptr = tl.make_block_ptr( + base=x_ptr, + shape=[n_rows, n_cols], + strides=[stride_0, stride_1], + offsets=[pid0 * BLOCK_SIZE_ROW, pid1 * BLOCK_SIZE_COL], + block_shape=[BLOCK_SIZE_ROW, BLOCK_SIZE_COL], + order=[1, 0], + ) + x = tl.load(input_ptr) + x = (2 * x) + 1 + output_ptr = tl.make_block_ptr( + base=y_ptr, + shape=[n_rows, n_cols], + strides=[stride_0, stride_1], + offsets=[pid0 * BLOCK_SIZE_ROW, pid1 * BLOCK_SIZE_COL], + block_shape=[BLOCK_SIZE_ROW, BLOCK_SIZE_COL], + order=[1, 0], + ) + tl.store(output_ptr, x) + + +def test(device): + n_rows = 512 + n_cols = 256 + x = torch.arange(0, n_rows * n_cols, 1, device="cpu", dtype=torch.float32).reshape([n_rows, n_cols]) + output = torch.full([n_rows, n_cols], -1, device="cpu", dtype=x.dtype) + BLOCK_SIZE_ROW = 4 + BLOCK_SIZE_COL = 2 + + grid = lambda meta: (n_rows // BLOCK_SIZE_ROW, n_cols // BLOCK_SIZE_COL) + + x_txda = x.to("txda") + output_txda = output.to("txda") + kernel[grid]( + x_txda, + output_txda, + n_rows, + n_cols, + x_txda.stride(0), + x_txda.stride(1), + BLOCK_SIZE_ROW=BLOCK_SIZE_ROW, + BLOCK_SIZE_COL=BLOCK_SIZE_COL, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + expected = (2 * x) + 1 + + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/examples/test_load_2d_tensor_col.py b/third_party/wafer/examples/test_load_2d_tensor_col.py new file mode 100755 index 00000000..15c7570f --- /dev/null +++ b/third_party/wafer/examples/test_load_2d_tensor_col.py @@ -0,0 +1,71 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +""" + +|-----|-----|-----|-----| +| | | | | +|-----|-----|-----|-----| +| | | | | +|-----|-----|-----|-----| + +Each instance loads the entire column +""" + + +@triton.jit +def kernel( + x_ptr, + y_ptr, + n_rows, + n_cols, + BLOCK_SIZE_ROW: tl.constexpr, + BLOCK_SIZE_COL: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + input_ptr = tl.make_block_ptr( + base=x_ptr, + shape=[n_rows, n_cols], + strides=[BLOCK_SIZE_COL, 1], + offsets=[0, pid0], + block_shape=[BLOCK_SIZE_ROW, 1], + order=[1, 0], + ) + x = tl.load(input_ptr) + output_ptr = tl.make_block_ptr( + base=y_ptr, + shape=[n_rows, n_cols], + strides=[BLOCK_SIZE_COL, 1], + offsets=[0, pid0], + block_shape=[BLOCK_SIZE_ROW, 1], + order=[1, 0], + ) + tl.store(output_ptr, x) + + +def test(device): + n_rows = 4 + n_cols = 2 + x = torch.arange(0, n_rows * n_cols, 1, device="cpu", dtype=torch.float32).reshape([n_rows, n_cols]) + output = torch.full([n_rows, n_cols], -1, device="cpu", dtype=x.dtype) + BLOCK_SIZE_ROW = n_rows + BLOCK_SIZE_COL = n_cols + + grid = lambda meta: (n_cols, ) + + x_txda = x.to("txda") + output_txda = output.to("txda") + kernel[grid]( + x_txda, + output_txda, + n_rows, + n_cols, + BLOCK_SIZE_ROW=BLOCK_SIZE_ROW, + BLOCK_SIZE_COL=BLOCK_SIZE_COL, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + torch.testing.assert_close(output, x, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/examples/test_load_store_mod.py b/third_party/wafer/examples/test_load_store_mod.py new file mode 100755 index 00000000..a29540f0 --- /dev/null +++ b/third_party/wafer/examples/test_load_store_mod.py @@ -0,0 +1,76 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + +@triton.jit +def stacked_load_2d_kernel(x_ptr, y_ptr, M, C, BLOCK_SIZE: tl.constexpr): + pid = tl.program_id(0) + m_offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)[:, None] + m_mask = m_offsets < M + c_offsets = m_offsets % C + w = tl.load(x_ptr + c_offsets, mask=m_mask) + tl.store(y_ptr + m_offsets, w, mask=m_mask) + + +@triton.jit +def sidebyside_load_2d_kernel(x_ptr, y_ptr, M, C, BLOCK_SIZE: tl.constexpr): + pid = tl.program_id(0) + m_offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + m_offsets = m_offsets[None, :] + m_mask = m_offsets < M + c_offsets = m_offsets % C + w = tl.load(x_ptr + c_offsets, mask=m_mask) + tl.store(y_ptr + m_offsets, w, mask=m_mask) + + +@triton.jit +def sidebyside_load_1d_kernel(x_ptr, y_ptr, M, C, BLOCK_SIZE: tl.constexpr): + pid = tl.program_id(0) + m_offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + m_mask = m_offsets < M + c_offsets = m_offsets % C + w = tl.load(x_ptr + c_offsets, mask=m_mask) + tl.store(y_ptr + m_offsets, w, mask=m_mask) + + +configs = [ + 8, # mask < N + 16, # N < mask = 2N < colsize + 32, # N < mask = 3N < colsize + 64, # N < mask = 4N = colsize + 128 # N < colsize < mask +] + + +@pytest.mark.parametrize("BLOCK_SIZE", configs) +def test(device, BLOCK_SIZE): + C = 16 + B = 4 + M = B * C + weight = torch.randn(size=(C, ), dtype=torch.float32, device="cpu", requires_grad=True) + output = torch.full([B, C], -1, device="cpu", dtype=torch.float32) + + indices = torch.arange(M, device="cpu") % C + torch_output = weight[indices].view(B, C) + + grid = lambda meta: (triton.cdiv(M, BLOCK_SIZE), ) + weight_txda = weight.to("txda") + output_txda = output.to("txda") + stacked_load_2d_kernel[grid](weight_txda, output_txda, M, C, BLOCK_SIZE) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch_output, rtol=0.001, atol=1e-5) + + sidebyside_load_2d_kernel[grid](weight_txda, output_txda, M, C, BLOCK_SIZE) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch_output, rtol=0.001, atol=1e-5) + + sidebyside_load_1d_kernel[grid](weight_txda, output_txda, M, C, BLOCK_SIZE) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, torch_output, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/examples/test_log.py b/third_party/wafer/examples/test_log.py new file mode 100755 index 00000000..444e9929 --- /dev/null +++ b/third_party/wafer/examples/test_log.py @@ -0,0 +1,96 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def log_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.log(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def log_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + log_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_log(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = log_triton(x) + + # Verify the output + expected = torch.log(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_log_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = log_triton(x) + + # Verify the output + expected = torch.log(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_log_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_log2.py b/third_party/wafer/examples/test_log2.py new file mode 100755 index 00000000..327e6db8 --- /dev/null +++ b/third_party/wafer/examples/test_log2.py @@ -0,0 +1,96 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def log2_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.log2(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def log2_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + log2_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_log2(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = log2_triton(x) + + # Verify the output + expected = torch.log2(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_log2_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = log2_triton(x) + + # Verify the output + expected = torch.log2(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_log2_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_mask.py b/third_party/wafer/examples/test_mask.py new file mode 100755 index 00000000..a7d48729 --- /dev/null +++ b/third_party/wafer/examples/test_mask.py @@ -0,0 +1,42 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + + +def test_mask(device): + + @triton.jit + def test(in0, out0): + offs = 100 + tl.arange(0, 4) + out_offs = tl.arange(0, 4) + a = tl.load(in0 + offs, mask=offs < 4, other=-1) + tl.store(out0 + out_offs, a) + + SIZE = 8 + input = torch.arange(0, SIZE, device="cpu", dtype=torch.int32) + output = torch.full((SIZE, ), -2, device="cpu", dtype=torch.int32) + + if device == 'cpu': + pass # Wafer driver is selected by conftest.py. + + grid = lambda meta: (1, ) + + src = triton.compiler.ASTSource( + fn=test, + signature={'in0': '*fp32', 'out0': '*fp32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) + + print(output) + input_txda = input.to("txda") + output_txda = output.to("txda") + test[grid](input_txda, output_txda) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(input) + print(output) + torch.testing.assert_close(output, torch.tensor([-1, -1, -1, -1, -2, -2, -2, -2], device="cpu", dtype=torch.int32)) diff --git a/third_party/wafer/examples/test_math_erf_op.py b/third_party/wafer/examples/test_math_erf_op.py new file mode 100755 index 00000000..fd979405 --- /dev/null +++ b/third_party/wafer/examples/test_math_erf_op.py @@ -0,0 +1,85 @@ +import textwrap + +import numpy as np +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + +import inspect +import benchmark + +from numpy.random import RandomState + +from triton._internal_testing import ( + integral_dtypes, + int_dtypes, + str_to_triton_dtype, + uint_dtypes, + float_dtypes, + float_dtypes_with_bfloat16, + dtypes, + dtypes_with_bfloat16, + is_cuda, + is_interpreter, + is_hopper, + is_hip, + is_hip_cdna, + is_hip_cdna2, + is_hip_cdna3, + is_hip_cdna4, + is_xpu, + get_arch, + torch_float8_dtypes, + torch_dtypes, + numpy_random, + to_triton, + torch_dtype_name, + to_numpy, +) + + +def check_type_supported(dtype, device): + ''' + skip test if dtype is not supported on the current device + ''' + if device in ['cuda']: + cc = torch.cuda.get_device_capability() + if cc[0] < 8 and (dtype is tl.bfloat16 or dtype == "bfloat16" or dtype is torch.bfloat16): + pytest.skip("bfloat16 is only supported on NVGPU with cc >= 80") + if cc[0] < 9 and dtype in {tl.float8e4nv, "float8e4nv", "float8_e4m3fn"}: + pytest.skip("float8e4nv is only supported on NVGPU with cc >= 90") + if is_interpreter(): + if dtype in [tl.bfloat16, "bfloat16", torch.bfloat16]: + pytest.skip("bfloat16 is not supported in the interpreter") + + +@pytest.mark.interpreter +@pytest.mark.parametrize("dtype", [dtype for dtype in ["float32"]]) +def test_math_erf_op(dtype, device): + check_type_supported(dtype, device) + SIZE = 128 + + @triton.jit + def kernel(Z, X, SIZE: tl.constexpr): + off = tl.arange(0, SIZE) + x = tl.load(X + off) + z = tl.math.erf(x) + tl.store(Z + off, z) + + torch_dtype = torch.float32 if dtype == "float32" else torch.float64 + x = torch.randn(SIZE, dtype=torch_dtype, device="cpu") + z_ref = torch.erf(x) + z_tri = torch.zeros_like(x) + z_tri_txda = z_tri.to("txda") + x_txda = x.to("txda") + kernel[(1, )](z_tri_txda, x_txda, SIZE=SIZE, num_warps=4) + with torch.no_grad(): + z_tri.copy_(z_tri_txda.cpu()) + torch.testing.assert_close(z_tri, z_ref) + + +if __name__ == "__main__": + test_math_erf_op("float32", "cpu") diff --git a/third_party/wafer/examples/test_matmul.py b/third_party/wafer/examples/test_matmul.py new file mode 100755 index 00000000..96fc9ae3 --- /dev/null +++ b/third_party/wafer/examples/test_matmul.py @@ -0,0 +1,187 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark +import pytest + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes: +# - A list of `triton.Config` objects that define different configurations of +# meta-parameters (e.g., `BLOCK_SIZE_M`) and compilation options (e.g., `num_warps`) to try +# - An auto-tuning *key* whose change in values will trigger evaluation of all the +# provided configs +@triton.autotune( + configs=[ + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=3, + num_warps=8), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, + num_warps=4), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, + num_warps=4), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, + num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, + num_warps=4), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, + num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=5, + num_warps=2), + triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=5, + num_warps=2), + ], + key=['M', 'N', 'K'], +) +@triton.jit +def matmul_kernel( + # Pointers to matrices + a_ptr, b_ptr, c_ptr, + # Matrix dimensions + M, N, K, + # The stride variables represent how much to increase the ptr by when moving by 1 + # element in a particular dimension. E.g. `stride_am` is how much to increase `a_ptr` + # by to get the element one row down (A has M rows). + stride_am, stride_ak, # + stride_bk, stride_bn, # + stride_cm, stride_cn, + # Meta-parameters + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, # + GROUP_SIZE_M: tl.constexpr, # + ACTIVATION: tl.constexpr # +): + """Kernel for computing the matmul C = A x B. + A has shape (M, K), B has shape (K, N) and C has shape (M, N) + """ + # ----------------------------------------------------------- + # Map program ids `pid` to the block of C it should compute. + # This is done in a grouped ordering to promote L2 data reuse. + # See above `L2 Cache Optimizations` section for details. + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + # ---------------------------------------------------------- + # Create pointers for the first blocks of A and B. + # We will advance this pointer as we move in the K direction + # and accumulate + # `a_ptrs` is a block of [BLOCK_SIZE_M, BLOCK_SIZE_K] pointers + # `b_ptrs` is a block of [BLOCK_SIZE_K, BLOCK_SIZE_N] pointers + # See above `Pointer Arithmetics` section for details + offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + + # ----------------------------------------------------------- + # Iterate to compute a block of the C matrix. + # We accumulate into a `[BLOCK_SIZE_M, BLOCK_SIZE_N]` block + # of fp32 values for higher accuracy. + # `accumulator` will be converted back to fp16 after the loop. + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + # Load the next block of A and B, generate a mask by checking the K dimension. + # If it is out of bounds, set it to 0. + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) + # We accumulate along the K dimension. + accumulator += tl.dot(a, b) + # Advance the ptrs to the next K block. + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + # You can fuse arbitrary activation functions here + # while the accumulator is still in FP32! + if ACTIVATION == "leaky_relu": + accumulator = leaky_relu(accumulator) + c = accumulator.to(tl.float32) + + # ----------------------------------------------------------- + # Write back the block of the output matrix C with masks. + offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, c, mask=c_mask) + + +# We can fuse `leaky_relu` by providing it as an `ACTIVATION` meta-parameter in `_matmul`. +@triton.jit +def leaky_relu(x): + x = x + 1 + return tl.where(x >= 0, x, 0.01 * x) + + +def matmul(a, b, activation=""): + # Check constraints. + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + assert a.is_contiguous(), "Matrix A must be contiguous" + assert b.is_contiguous(), "Matrix B must be contiguous" + M, K = a.shape + K, N = b.shape + # Allocates output. + c = torch.empty((M, N), device=a.device, dtype=a.dtype) + # 1D launch kernel where each block gets its own program. + grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), ) + matmul_kernel[grid]( + a, b, c, # + M, N, K, # + a.stride(0), a.stride(1), # + b.stride(0), b.stride(1), # + c.stride(0), c.stride(1), # + ACTIVATION=activation, # + ) + return c + + +@pytest.mark.parametrize("M, K, N, dtype", [ # + (M, K, N, dtype) + for M in [48, 64, 128] + for K in [128, 156, 512] + for N in [48, 64, 128] + for dtype in [torch.float32, torch.float16, torch.bfloat16] +]) +def test_matmul(M, K, N, dtype, device='cpu'): + # Generate random input tensors + torch.manual_seed(0) + rows1 = 179 + cols1 = 167 + rows2 = 167 + cols2 = 321 + # a = torch.randn((rows1, cols1), device=device, dtype=torch.float32) + # b = torch.randn((rows2, cols2), device=device, dtype=torch.float32) + a_ = torch.full((rows1, cols1), 1, device='cpu', dtype=torch.float32) + b_ = torch.full((rows2, cols2), 1, device='cpu', dtype=torch.float32) + a = a_.to(DEVICE) + b = b_.to(DEVICE) + triton_output = matmul(a, b) + triton_output = triton_output.to("cpu") + + torch_output = torch.matmul(a_, b_) + torch.testing.assert_close(triton_output, torch_output, atol=1e-2, rtol=0) + + +@benchmark.measure() +def bench_matmul(M, N, K, provider): + a = torch.randn((M, K), device='cpu', dtype=torch.float32) + b = torch.randn((K, N), device='cpu', dtype=torch.float32) + if provider == 'torch': + torch.matmul(a, b) + if provider == 'triton': + matmul(a.to("txda"), b.to("txda")) + + +if __name__ == "__main__": + # test_matmul(179,167,321,torch.float32) + # test_matmul("txda") + for X in [128 * i for i in range(2, 7)]: + for provider in ['torch', 'triton']: + bench_matmul(X, X, X, provider) diff --git a/third_party/wafer/examples/test_maximum.py b/third_party/wafer/examples/test_maximum.py new file mode 100755 index 00000000..bb4b1ca4 --- /dev/null +++ b/third_party/wafer/examples/test_maximum.py @@ -0,0 +1,102 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def maximum_kernel( + x_ptr, + y_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.maximum(x, y) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def maximum_triton(x, y): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + maximum_kernel[grid]( + x_txda, + y_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_maximum(size, dtype, device="cpu"): + # Generate random input tensors + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = maximum_triton(x, y) + + # Verify the output + expected = torch.maximum(x, y) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_maximum_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input tensors + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = maximum_triton(x, y) + + # Verify the output + expected = torch.maximum(x, y) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_maximum_triton(size, torch.float32, provider="triton") diff --git a/third_party/wafer/examples/test_minimum.py b/third_party/wafer/examples/test_minimum.py new file mode 100755 index 00000000..0fe34279 --- /dev/null +++ b/third_party/wafer/examples/test_minimum.py @@ -0,0 +1,102 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def minimum_kernel( + x_ptr, + y_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.minimum(x, y) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def minimum_triton(x, y): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + minimum_kernel[grid]( + x_txda, + y_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_minimum(size, dtype, device="cpu"): + # Generate random input tensors + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = minimum_triton(x, y) + + # Verify the output + expected = torch.minimum(x, y) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_minimum_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input tensors + x = torch.randn(size, device="cpu", dtype=dtype) + y = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = minimum_triton(x, y) + + # Verify the output + expected = torch.minimum(x, y) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_minimum_triton(size, torch.float32, provider="triton") diff --git a/third_party/wafer/examples/test_modulo.py b/third_party/wafer/examples/test_modulo.py new file mode 100755 index 00000000..1243513d --- /dev/null +++ b/third_party/wafer/examples/test_modulo.py @@ -0,0 +1,329 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + +def test_wrap_stacked(device): + + @triton.jit + def wrap_stacked(a_ptr, c_ptr, M, N, stride_am, stride_an, stride_cm, stride_cn, BLOCK_SIZE_K: tl.constexpr): + offs_am = (2 + tl.arange(0, 4)) % M + offs_an = tl.arange(0, 4) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_an[None, :] * stride_an) + + offs_cm = tl.arange(0, 4) + offs_cn = tl.arange(0, 4) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + + for k in range(0, 2): + a = tl.load(a_ptrs) + tl.store(c_ptrs, a) + a_ptrs += BLOCK_SIZE_K * stride_an + c_ptrs += BLOCK_SIZE_K * stride_an + + M = 4 + N = 8 + A = torch.arange(0, M * N, device="cpu", dtype=torch.float32).reshape((M, N)) + out = torch.full((M, N), 88888, device="cpu", dtype=torch.float32) + grid = lambda meta: (1, ) + + A_txda = A.to("txda") + out_txda = out.to("txda") + wrap_stacked[grid](A_txda, out_txda, M, N, A_txda.stride(0), A_txda.stride(1), out_txda.stride(0), out_txda.stride(1), BLOCK_SIZE_K=4) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor([[16, 17, 18, 19, 20, 21, 22, 23], [24, 25, 26, 27, 28, 29, 30, 31], + [0, 1, 2, 3, 4, 5, 6, 7], [8, 9, 10, 11, 12, 13, 14, 15]], device="cpu") + + assert torch.equal(expected_out.int(), out.int()) + + +def test_1d(device): + + @triton.jit + def mod_1d(a_ptr, c_ptr, M, N, stride_am, stride_an, stride_cm, stride_cn, BLOCK_SIZE_K: tl.constexpr): + row = 7 + offs_an = (6 + tl.arange(0, 4)) % N + a_ptrs = a_ptr + (row * stride_am) + offs_an[None, :] * stride_an + + offs_cn = tl.arange(0, 4) + c_ptrs = c_ptr + stride_cn * offs_cn[None, :] + + a = tl.load(a_ptrs) + tl.store(c_ptrs, a) + + M = 8 + N = 8 + A = torch.arange(0, M * N, device="cpu", dtype=torch.float32).reshape((M, N)) + out = torch.full((M, N), 88888, device="cpu", dtype=torch.float32) + grid = lambda meta: (1, ) + + A_txda = A.to("txda") + out_txda = out.to("txda") + mod_1d[grid](A_txda, out_txda, M, N, A_txda.stride(0), A_txda.stride(1), out_txda.stride(0), out_txda.stride(1), BLOCK_SIZE_K=4) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[62, 63, 56, 57, 88888, 88888, 88888, 88888], [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888]], device="cpu") + + assert torch.equal(expected_out.int(), out.int()) + + +def test_2d(device): + + @triton.jit + def mod_2d(a_ptr, c_ptr, M, N, stride_am, stride_an, stride_cm, stride_cn, BLOCK_SIZE_K: tl.constexpr): + offs_am = 2 + tl.arange(0, 4) + offs_an = (6 + tl.arange(0, 4)) % N + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_an[None, :] * stride_an) + + offs_cm = tl.arange(0, 4) + offs_cn = tl.arange(0, 4) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + + a = tl.load(a_ptrs) + tl.store(c_ptrs, a) + + M = 8 + N = 8 + A = torch.arange(0, M * N, device="cpu", dtype=torch.float32).reshape((M, N)) + out = torch.full((M, N), 88888, device="cpu", dtype=torch.float32) + grid = lambda meta: (1, ) + + A_txda = A.to("txda") + out_txda = out.to("txda") + mod_2d[grid](A_txda, out_txda, M, N, A_txda.stride(0), A_txda.stride(1), out_txda.stride(0), out_txda.stride(1), BLOCK_SIZE_K=4) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[22, 23, 16, 17, 88888, 88888, 88888, 88888], [30, 31, 24, 25, 88888, 88888, 88888, 88888], + [38, 39, 32, 33, 88888, 88888, 88888, 88888], [46, 47, 40, 41, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888]], device="cpu") + + assert torch.equal(expected_out.int(), out.int()) + + +def test_side_by_side_masked_loop(device): + + @triton.jit + def wrap_side_by_side_masked_loop(a_ptr, c_ptr, M, N, stride_am, stride_an, stride_cm, stride_cn, + BLOCK_SIZE_K: tl.constexpr): + offs_am = 2 + tl.arange(0, BLOCK_SIZE_K) + offs_an = (6 + tl.arange(0, BLOCK_SIZE_K)) % N + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_an[None, :] * stride_an) + + offs_k = tl.arange(0, BLOCK_SIZE_K) + + offs_cm = tl.arange(0, BLOCK_SIZE_K) + offs_cn = tl.arange(0, BLOCK_SIZE_K) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + + for k in range(0, 2): + a = tl.load(a_ptrs, mask=offs_k[:, None] < 2, other=-99) + tl.store(c_ptrs, a) + a_ptrs += BLOCK_SIZE_K * stride_am + c_ptrs += BLOCK_SIZE_K * stride_an + + M = 12 + N = 8 + A = torch.arange(0, M * N, device="cpu", dtype=torch.float32).reshape((M, N)) + out = torch.full((M, N), 88888, device="cpu", dtype=torch.float32) + print(out) + grid = lambda meta: (1, ) + + A_txda = A.to("txda") + out_txda = out.to("txda") + wrap_side_by_side_masked_loop[grid](A_txda, out_txda, M, N, A_txda.stride(0), A_txda.stride(1), out_txda.stride(0), out_txda.stride(1), + BLOCK_SIZE_K=4) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[22, 23, 16, 17, 54, 55, 48, 49], [30, 31, 24, 25, 62, 63, 56, 57], [-99, -99, -99, -99, -99, -99, -99, -99], + [-99, -99, -99, -99, -99, -99, -99, -99], [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888]], dtype=torch.int32) + + assert torch.equal(expected_out.int(), out.int()) + + +def test_stacked_masked_loop(device): + + @triton.jit + def wrap_stacked_masked_loop(a_ptr, c_ptr, M, N, stride_am, stride_an, stride_cm, stride_cn, + BLOCK_SIZE_K: tl.constexpr): + offs_am = (2 + tl.arange(0, BLOCK_SIZE_K)) % M + offs_an = 3 + tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_an[None, :] * stride_an) + + offs_cm = tl.arange(0, BLOCK_SIZE_K) + offs_cn = tl.arange(0, BLOCK_SIZE_K) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + + offs_k = tl.arange(0, BLOCK_SIZE_K) + + for k in range(0, 2): + a = tl.load(a_ptrs, mask=offs_k[None, :] < 3, other=-99) + tl.store(c_ptrs, a) + a_ptrs += BLOCK_SIZE_K * stride_an + c_ptrs += BLOCK_SIZE_K * stride_an + + M = 4 + N = 12 + BLOCK_SIZE_M = 4 + BLOCK_SIZE_N = 4 + A = torch.arange(0, M * N, device="cpu", dtype=torch.float32).reshape((M, N)) + out = torch.full((BLOCK_SIZE_M, N), 88888, device="cpu", dtype=torch.float32) + print(out) + grid = lambda meta: (1, ) + + A_txda = A.to("txda") + out_txda = out.to("txda") + wrap_stacked_masked_loop[grid](A_txda, out_txda, M, N, A_txda.stride(0), A_txda.stride(1), out_txda.stride(0), out_txda.stride(1), BLOCK_SIZE_K=4) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor([ + [ + 27.0, + 28.0, + 29.0, + -99.0, + 31.0, + 32.0, + 33.0, + -99.0, + 88888, + 88888, + 88888, + 88888, + ], + [ + 39.0, + 40.0, + 41.0, + -99.0, + 43.0, + 44.0, + 45.0, + -99.0, + 88888, + 88888, + 88888, + 88888, + ], + [ + 3.0, + 4.0, + 5.0, + -99.0, + 7.0, + 8.0, + 9.0, + -99.0, + 88888, + 88888, + 88888, + 88888, + ], + [ + 15.0, + 16.0, + 17.0, + -99.0, + 19.0, + 20.0, + 21.0, + -99.0, + 88888, + 88888, + 88888, + 88888, + ], + ], ) + + assert torch.equal(expected_out.int(), out.int()) + + +def test_torch_inductor_pattern(): + + @triton.jit + def triton_(in_ptr2, out_ptr2, rnumel, XBLOCK: tl.constexpr, RBLOCK: tl.constexpr): + xnumel = 128 + rnumel = 32 + xoffset = tl.program_id(0) * XBLOCK + xindex = xoffset + tl.arange(0, XBLOCK)[:, None] + rbase = tl.arange(0, RBLOCK)[None, :] + x0 = xindex % 7 + x0 = xindex + roffset = 0 + rindex = roffset + rbase + rmask = rindex < rnumel + r2 = rindex + tmp3 = tl.load(in_ptr2 + (r2 + (xnumel * x0)), rmask, other=77) + tl.store(out_ptr2 + (XBLOCK * tl.arange(0, RBLOCK)[None, :] + tl.arange(0, XBLOCK)[:, None]), tmp3) + + device = "cpu" + xnumel = 128 + rnumel = 32 + + XBLOCK = 4 + RBLOCK = 64 + A = torch.arange(0, xnumel * rnumel, device="cpu", dtype=torch.int32).reshape((xnumel, rnumel)) + out = torch.full((XBLOCK, RBLOCK), 88888, device="cpu", dtype=torch.int32) + grid = lambda meta: (1, ) + + A_txda = A.to("txda") + out_txda = out.to("txda") + triton_[grid](A_txda, out_txda, rnumel, XBLOCK=XBLOCK, RBLOCK=RBLOCK) + with torch.no_grad(): + out.copy_(out_txda.cpu()) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[ + 0, 128, 256, 384, 1, 129, 257, 385, 2, 130, 258, 386, 3, 131, 259, 387, 4, 132, 260, 388, 5, 133, 261, 389, + 6, 134, 262, 390, 7, 135, 263, 391, 8, 136, 264, 392, 9, 137, 265, 393, 10, 138, 266, 394, 11, 139, 267, + 395, 12, 140, 268, 396, 13, 141, 269, 397, 14, 142, 270, 398, 15, 143, 271, 399 + ], + [ + 16, 144, 272, 400, 17, 145, 273, 401, 18, 146, 274, 402, 19, 147, 275, 403, 20, 148, 276, 404, 21, 149, + 277, 405, 22, 150, 278, 406, 23, 151, 279, 407, 24, 152, 280, 408, 25, 153, 281, 409, 26, 154, 282, 410, + 27, 155, 283, 411, 28, 156, 284, 412, 29, 157, 285, 413, 30, 158, 286, 414, 31, 159, 287, 415 + ], + [ + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77 + ], + [ + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77 + ]], device="cpu", dtype=torch.int32) + + assert torch.equal(expected_out.int(), out.int()) diff --git a/third_party/wafer/examples/test_nested_loops.py b/third_party/wafer/examples/test_nested_loops.py new file mode 100755 index 00000000..8ba82a1b --- /dev/null +++ b/third_party/wafer/examples/test_nested_loops.py @@ -0,0 +1,399 @@ +import torch +import torch_txda # noqa: F401 + +import triton +from triton.backends.compiler import GPUTarget +import triton.language as tl + + +# Not used for testing but serves as a template to generate the lit test at +# test/Conversion/TritonToStructured/ridiculously_nested_loops.mlir +@triton.jit +def nested_who_knows_how_many_levels(in_ptr, out_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = in_ptr + (offs_am[:, None] * stride_m + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = out_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[None, :] + + for i1 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j1 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k1 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + for i2 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j2 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k2 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + for i3 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j3 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k3 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + for i4 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j4 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k4 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + for i5 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j5 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k5 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + a_ptrs += 2 * stride_n + + for i6 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j6 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k6 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + a_ptrs += 2 * stride_n + + a_ptrs += 2 * stride_n + + +@triton.jit +def nested_use_same_level_loop_results(in_ptr, out_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = in_ptr + (offs_am[:, None] * stride_m + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = out_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[None, :] + + for i1 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j1 in range(0, 2): + a_ptrs += 2 * stride_n + + for i6 in range(0, 2): + a1 = tl.load(a_ptrs) + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + a_ptrs += 2 * stride_n + + a_ptrs += 2 * stride_n + + +@triton.jit +def nested2_complex_body(a_ptr, c_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = a_ptr + (offs_am[:, None] * stride_m + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = c_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[None, :] + + for i in range(0, 2): + a_ptrs_copy = a_ptrs + c_ptrs_copy = c_ptrs + + a_ptrs += 1 + c_ptrs += 1 + + for j in range(0, 2): + a2 = tl.load(a_ptrs) + tl.store(c_ptrs, a2) + a_ptrs += 3 + c_ptrs += 3 + + a_ptrs = a_ptrs_copy + 2 * stride_m + 1 + c_ptrs = c_ptrs_copy + 2 * stride_m + 1 + + +@triton.jit +def nested2_use_loop_results(in_ptr, out_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = in_ptr + (offs_am[:, None] * stride_m + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = out_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[None, :] + + for i in range(0, 2): + a2 = tl.load(a_ptrs) + tl.store(c_ptrs, a2) + + a_ptrs += 4 * stride_n + c_ptrs += 4 * stride_n + + for j in range(0, 2): + a2 = tl.load(a_ptrs) + tl.store(c_ptrs, a2) + a_ptrs += 4 * stride_n + c_ptrs += 4 * stride_n + + +@triton.jit +def nested3(in_ptr, out_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = in_ptr + (offs_am[:, None] * stride_m + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = out_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[None, :] + + for i in range(0, 2): + a1 = tl.load(a_ptrs) + + for j in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + a_ptrs += 2 * stride_n + + +def test_nested3(): + n_rows = 4 + n_cols = 48 + expected = torch.tensor( + [[ + 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 6, 7, 0, 1, 8, 9, 10, 11, 0, 1, 8, 9, 12, 13, 14, 15, 16, 17, 18, 19, 14, 15, + 16, 17, 20, 21, 14, 15, 22, 23, 24, 25, 14, 15, 22, 23, 26, 27 + ], + [ + 48, 49, 50, 51, 52, 53, 48, 49, 50, 51, 54, 55, 48, 49, 56, 57, 58, 59, 48, 49, 56, 57, 60, 61, 62, 63, 64, + 65, 66, 67, 62, 63, 64, 65, 68, 69, 62, 63, 70, 71, 72, 73, 62, 63, 70, 71, 74, 75 + ], + [ + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 + ], + [ + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 + ]], dtype=torch.int32, device="cpu") + pass # Wafer driver is selected by conftest.py. + x = torch.arange(0, n_rows * n_cols, device="cpu", dtype=torch.int32).reshape([n_rows, n_cols]) + output = torch.zeros([n_rows, n_cols], device="cpu", dtype=x.dtype) + grid = lambda meta: (n_cols // 4, ) + + print('before:') + print(x) + print(output) + + x_txda = x.to("txda") + output_txda = output.to("txda") + nested3[grid](x_txda, output_txda, x_txda.stride(0), x_txda.stride(1)) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) + print("Pass!") + + src = triton.compiler.ASTSource( + fn=nested3, + signature={'in_ptr': '*fp32', 'out_ptr': '*fp32', 'stride_m': 'i32', 'stride_n': 'i32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) + print('Pass') + + +def test_nested2_use_loop_results(): + n_rows = 4 + n_cols = 32 + expected = torch.tensor( + [[0, 1, 0, 0, 4, 5, 0, 0, 8, 9, 0, 0, 12, 13, 0, 0, 16, 17, 0, 0, 20, 21, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [32, 33, 0, 0, 36, 37, 0, 0, 40, 41, 0, 0, 44, 45, 0, 0, 48, 49, 0, 0, 52, 53, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]], + device="cpu", dtype=torch.int32) + # x = torch.arange(0, n_rows * n_cols, device="cuda", dtype=torch.int32).reshape([n_rows, n_cols]) + pass # Wafer driver is selected by conftest.py. + x = torch.arange(0, n_rows * n_cols, device="cpu", dtype=torch.int32).reshape([n_rows, n_cols]) + output = torch.zeros([n_rows, n_cols], device="cpu", dtype=x.dtype) + grid = lambda meta: (n_cols // 4, ) + + print('before:') + print(x) + print(output) + + x_txda = x.to("txda") + output_txda = output.to("txda") + nested2_use_loop_results[grid](x_txda, output_txda, x_txda.stride(0), x_txda.stride(1)) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) + print("Pass!") + + src = triton.compiler.ASTSource( + fn=nested2_use_loop_results, + signature={'in_ptr': '*fp32', 'out_ptr': '*fp32', 'stride_m': 'i32', 'stride_n': 'i32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) + print('Pass') + + +def test_nested2_complex_body(): + n_rows = 4 + n_cols = 8 + grid = lambda meta: (n_cols // 4, ) + expected = torch.tensor([[0, 1, 2, 0, 4, 5, 0, 0], [0, 9, 10, 0, 12, 13, 0, 0], [0, 0, 18, 19, 0, 21, 22, 0], + [0, 0, 26, 27, 0, 29, 30, 0]], device="cpu", dtype=torch.int32) + + x = torch.arange(0, n_rows * n_cols, device="cpu", dtype=torch.int32).reshape([n_rows, n_cols]) + pass # Wafer driver is selected by conftest.py. + output = torch.zeros([n_rows, n_cols], device="cpu", dtype=x.dtype) + + print('before:') + print(x) + print(output) + + x_txda = x.to("txda") + output_txda = output.to("txda") + nested2_complex_body[grid](x_txda, output_txda, x_txda.stride(0), x_txda.stride(1)) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) + print("Pass!") + + src = triton.compiler.ASTSource( + fn=nested2_complex_body, + signature={'a_ptr': '*fp32', 'c_ptr': '*fp32', 'stride_m': 'i32', 'stride_n': 'i32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) + print('Pass') + + +def test_nested2_use_same_level_loop_result(): + n_rows = 4 + n_cols = 32 + grid = lambda meta: (n_cols // 4, ) + expected = torch.tensor([[ + 4, 5, 0, 0, 6, 7, 8, 9, 0, 0, 10, 11, 18, 19, 0, 0, 20, 21, 22, 23, 0, 0, 24, 25, 0, 0, 0, 0, 0, 0, 0, 0 + ], [36, 37, 0, 0, 38, 39, 40, 41, 0, 0, 42, 43, 50, 51, 0, 0, 52, 53, 54, 55, 0, 0, 56, 57, 0, 0, 0, 0, 0, 0, 0, 0 + ], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 + ], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]], + device="cpu", dtype=torch.int32) + + x = torch.arange(0, n_rows * n_cols, device="cpu", dtype=torch.int32).reshape([n_rows, n_cols]) + pass # Wafer driver is selected by conftest.py. + output = torch.zeros([n_rows, n_cols], device="cpu", dtype=x.dtype) + + print('before:') + print(x) + print(output) + + x_txda = x.to("txda") + output_txda = output.to("txda") + nested_use_same_level_loop_results[grid](x_txda, output_txda, x_txda.stride(0), x_txda.stride(1)) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(output) + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) + print("Pass!") + + src = triton.compiler.ASTSource( + fn=nested_use_same_level_loop_results, + signature={'in_ptr': '*fp32', 'out_ptr': '*fp32', 'stride_m': 'i32', 'stride_n': 'i32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) + print('Pass') diff --git a/third_party/wafer/examples/test_pipeline.py b/third_party/wafer/examples/test_pipeline.py new file mode 100644 index 00000000..3ba3c52a --- /dev/null +++ b/third_party/wafer/examples/test_pipeline.py @@ -0,0 +1,36 @@ +"""Exercise software pipelining across full tiles, tails and short loops.""" +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def pipelined_gemm(A, B, C, K: tl.constexpr, BK: tl.constexpr): + m = tl.arange(0, 32) + n = tl.arange(0, 32) + k = tl.arange(0, BK) + acc = tl.full((32, 32), 0, tl.float32) + for block in tl.range(0, tl.cdiv(K, BK), num_stages=2): + offsets = block * BK + k + a = tl.load(A + m[:, None] * K + offsets[None, :], offsets[None, :] < K, other=0) + b = tl.load(B + offsets[:, None] * 32 + n[None, :], offsets[:, None] < K, other=0) + acc += tl.dot(a, b) + tl.store(C + m[:, None] * 32 + n[None, :], acc) + + +@pytest.mark.parametrize("k", [16, 32, 33, 64, 96, 128]) +def test_pipeline_gemm(device, k): + a = torch.randn((32, k), dtype=torch.float16, device="cpu") + b = torch.randn((k, 32), dtype=torch.float16, device="cpu") + expected = a.float() @ b.float() + for enabled in (False, True): + output = torch.empty((32, 32), dtype=torch.float32, device="cpu") + a_txda = a.to("txda") + b_txda = b.to("txda") + output_txda = output.to("txda") + pipelined_gemm[(1,)](a_txda, b_txda, output_txda, k, 32, num_stages=2, enable_pipeline=enabled) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + torch.testing.assert_close(output, expected, rtol=2e-3, atol=2e-3) diff --git a/third_party/wafer/examples/test_precision_modes.py b/third_party/wafer/examples/test_precision_modes.py new file mode 100644 index 00000000..50f49d40 --- /dev/null +++ b/third_party/wafer/examples/test_precision_modes.py @@ -0,0 +1,38 @@ +"""Integer precision modes and reciprocal correction on the real device.""" +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def integer_math(X, Y, A, Q, R, BLOCK: tl.constexpr): + i = tl.arange(0, BLOCK) + x = tl.load(X + i) + y = tl.load(Y + i) + tl.store(A + i, x + y) + tl.store(Q + i, x // y) + tl.store(R + i, x % y) + + +@pytest.mark.parametrize("mode,dtype", [(0, torch.int32), (1, torch.int64), (2, torch.int32)]) +def test_integer_modes(device, mode, dtype): + # Mode 0 checks exactly representable inputs, including reciprocal rounding + # at equal operands. Modes 1/2 additionally exercise integers beyond 2**24. + values = [7, -7, 14, -14, 15, -15, 0, 1] + if mode: + values[4:6] = [2**24 + 7, -(2**24 + 7)] + x = torch.tensor(values, dtype=dtype, device="cpu") + y = torch.tensor([7, 7, -7, -7, 7, 7, 7, 7], dtype=dtype, device="cpu") + outputs = [torch.empty_like(x) for _ in range(3)] + x_txda = x.to("txda") + y_txda = y.to("txda") + outputs_txda = [value.to("txda") for value in outputs] + integer_math[(1,)](x_txda, y_txda, *outputs_txda, 8, precision_mode=mode) + with torch.no_grad(): + for host, native in zip(outputs, outputs_txda): + host.copy_(native.cpu()) + references = (x + y, torch.div(x, y, rounding_mode="trunc"), x - torch.div(x, y, rounding_mode="trunc") * y) + for actual, expected in zip(outputs, references): + torch.testing.assert_close(actual, expected, rtol=0, atol=0) diff --git a/third_party/wafer/examples/test_print.py b/third_party/wafer/examples/test_print.py new file mode 100755 index 00000000..e10e5a2f --- /dev/null +++ b/third_party/wafer/examples/test_print.py @@ -0,0 +1,30 @@ +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +@triton.jit +def kernel_device_print(X, Y, BLOCK: tl.constexpr): + x = tl.load(X + tl.arange(0, BLOCK)) + y = tl.load(Y + tl.arange(0, BLOCK)) + tl.device_print("x: ", x, hex=True) + tl.device_print("constant : ", tl.constexpr(42)) + tl.store(Y + tl.arange(0, BLOCK), x) + + +def test_print(): + x = torch.arange(16, dtype=torch.int32) + x.reshape(4, 4) + y = torch.zeros_like(x) + x_txda = x.to("txda") + y_txda = y.to("txda") + kernel_device_print[(1, )](x_txda, y_txda, BLOCK=16) + with torch.no_grad(): + y.copy_(y_txda.cpu()) + torch.testing.assert_close(y, x) + + +if __name__ == "__main__": + # Run the test with pytest + test_print() diff --git a/third_party/wafer/examples/test_reduce.py b/third_party/wafer/examples/test_reduce.py new file mode 100755 index 00000000..37439720 --- /dev/null +++ b/third_party/wafer/examples/test_reduce.py @@ -0,0 +1,55 @@ +import torch +import torch_txda # noqa: F401 + +import triton +from triton.backends.compiler import GPUTarget +import triton.language as tl + + +@triton.jit +def reduce_kernel_2d( + x_ptr, + output_ptr, + stride, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + x = tl.load( + tl.make_block_ptr( + base=x_ptr, + shape=[n_elements * tl.num_programs(0)], + strides=[1], + offsets=[stride * pid0], + block_shape=[BLOCK_SIZE], + order=[0], + ), + boundary_check=[0], + ) + output = triton.language.sum(x, axis=0).to(dtype=x.dtype) + tl.store(output_ptr + pid0, output) + + +def test(device): + n_rows = 16 + n_cols = 32 + x = torch.rand([n_cols, n_rows], device="cpu", dtype=torch.float32) + output = torch.empty([n_cols], device="cpu", dtype=x.dtype) + BLOCK_SIZE = n_rows + grid = lambda meta: (n_cols, ) + + x_txda = x.to("txda") + output_txda = output.to("txda") + reduce_kernel_2d[grid](x_txda, output_txda, x_txda.stride(0), n_rows, BLOCK_SIZE=BLOCK_SIZE) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + ans = torch.sum(x, dim=1) + torch.testing.assert_close(output, ans, rtol=0.001, atol=1e-5) + + # TODO: need to check some conditions otherwise the code below does not make any difference for the test + src = triton.compiler.ASTSource(fn=reduce_kernel_2d, signature={'x_ptr': '*fp32', 'output_ptr': '*fp32', 'stride': 'i32', 'n_elements': 'i32', 'BLOCK_SIZE': 'constexpr'}, constexprs={"BLOCK_SIZE": 32}) + ret = triton.compile(src, target=GPUTarget("wafer", "wafer", 32)) + print(ret.asm["ttir"]) + print(ret.asm["coreir"]) + print(ret.asm["llir"]) + print(ret.asm["wafer_ir"]) diff --git a/third_party/wafer/examples/test_reduce1d.py b/third_party/wafer/examples/test_reduce1d.py new file mode 100755 index 00000000..8f66ca37 --- /dev/null +++ b/third_party/wafer/examples/test_reduce1d.py @@ -0,0 +1,45 @@ +import torch +import torch_txda # noqa: F401 + +import triton +from triton.backends.compiler import GPUTarget +import triton.language as tl +import benchmark + + +@triton.jit +def reduce_kernel_1d( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.arange(0, BLOCK_SIZE) + mask = offsets < BLOCK_SIZE + x = tl.load(x_ptr + offsets, mask=mask, other=0.0) + acc = tl.sum(x, axis=0) + tl.store(output_ptr, acc) + + +def test_1d_reduce_sum(device): + BLOCK_SIZE = 32768 + x = torch.ones([BLOCK_SIZE], device="cpu", dtype=torch.float32) + output = torch.empty([1], device="cpu", dtype=x.dtype) + grid = lambda meta: (1, ) + + x_txda = x.to("txda") + output_txda = output.to("txda") + reduce_kernel_1d[grid](x_txda, output_txda, BLOCK_SIZE, BLOCK_SIZE=BLOCK_SIZE) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + # CPU reference + ref = x.sum().unsqueeze(0) + + print(f"The maximum difference between ref and triton is " + f"{torch.max(torch.abs(ref - output))}") + torch.testing.assert_close(output, ref, rtol=0.001, atol=1e-5) + + +if __name__ == "__main__": + device = "cpu" + test_1d_reduce_sum(device) diff --git a/third_party/wafer/examples/test_rsqrt.py b/third_party/wafer/examples/test_rsqrt.py new file mode 100755 index 00000000..ef9b3b38 --- /dev/null +++ b/third_party/wafer/examples/test_rsqrt.py @@ -0,0 +1,96 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def rsqrt_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.rsqrt(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def rsqrt_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + rsqrt_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_rsqrt(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = rsqrt_triton(x) + + # Verify the output + expected = torch.rsqrt(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_rsqrt_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = rsqrt_triton(x) + + # Verify the output + expected = torch.rsqrt(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_rsqrt_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_scalar_store.py b/third_party/wafer/examples/test_scalar_store.py new file mode 100755 index 00000000..778d63dd --- /dev/null +++ b/third_party/wafer/examples/test_scalar_store.py @@ -0,0 +1,32 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + +@triton.jit +def reduce_kernel_2d( + output_ptr, + BLOCK_SIZE: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + base_ptr = output_ptr + pid0 + for i in range(0, BLOCK_SIZE): + output = i + tl.store(base_ptr, output) + base_ptr += 1 + + +def test(device): + BLOCK_SIZE = 8 + x = torch.full([BLOCK_SIZE], -1, device="cpu", dtype=torch.float32) + output = torch.full((BLOCK_SIZE, ), -99, device="cpu", dtype=x.dtype) + grid = lambda meta: (1, ) + + output_txda = output.to("txda") + reduce_kernel_2d[grid](output_txda, BLOCK_SIZE=BLOCK_SIZE) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + ans = torch.arange(BLOCK_SIZE, device="cpu", dtype=torch.float32) + torch.testing.assert_close(output, ans, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/examples/test_scan.py b/third_party/wafer/examples/test_scan.py new file mode 100755 index 00000000..39116113 --- /dev/null +++ b/third_party/wafer/examples/test_scan.py @@ -0,0 +1,72 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import benchmark + + +# @pytest.mark.parametrize("M, N", [(1, 64), (2, 32), (4, 16), (8, 8), (16, 4), (32, 2), (64, 1)]) +@pytest.mark.parametrize("M, N", [(32, 2)]) +@pytest.mark.parametrize("reversed", [True, False]) +def test_scan_1d(M, N, reversed, device): + + @triton.jit + def scan_kernel(out_ptr, in_ptr, n_elements, M: tl.constexpr, N: tl.constexpr): + offsets = tl.arange(0, M) + mask = offsets < n_elements + input = tl.load(in_ptr + offsets, mask=mask, other=0.0) + output = tl.cumsum(input).reshape([1, M]).broadcast_to([N, M]) + # tl.store(out_ptr + tl.arange(0, M * N), output.reshape([M * N])) + offs_cm = tl.arange(0, N) + offs_cn = tl.arange(0, M) + c_ptrs = out_ptr + M * offs_cm[:, None] + offs_cn[None, :] + c_mask = (offs_cn[None, :] < n_elements) + tl.store(c_ptrs, output, mask=c_mask) + + @triton.jit + def scan_kernel_reverse(out_ptr, in_ptr, n_elements, M: tl.constexpr, N: tl.constexpr): + offsets = tl.arange(0, M) + mask = offsets < n_elements + input = tl.load(in_ptr + offsets, mask=mask, other=0.0) + output = tl.cumsum(input, reverse=True).reshape([1, M]).broadcast_to([N, M]) + # tl.store(out_ptr + tl.arange(0, M * N), output.reshape([M * N])) + offs_cm = tl.arange(0, N) + offs_cn = tl.arange(0, M) + c_ptrs = out_ptr + M * offs_cm[:, None] + offs_cn[None, :] + c_mask = (offs_cn[None, :] < n_elements) + tl.store(c_ptrs, output, mask=c_mask) + + # x = torch.randint(-100, 100, (M, ), dtype=torch.int32, device=device) + # output = torch.empty(M * N, dtype=torch.int32, device=device) + + # x = torch.rand((M, ), dtype=torch.float32, device=device + n_elements = 32 + + x = torch.arange(0, n_elements, dtype=torch.float32, device="cpu") + output = torch.empty(M * N, dtype=torch.float32, device="cpu") + + if reversed: + scan_kernel = scan_kernel_reverse + ref_x = torch.flip(x, dims=[0]) + else: + scan_kernel = scan_kernel + ref_x = x + output_txda = output.to("txda") + x_txda = x.to("txda") + scan_kernel[(1, )](output_txda, x_txda, n_elements, M, N) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + ref = torch.cumsum(ref_x, dim=0).reshape([1, M]).broadcast_to([N, M]).reshape([M * N]) + if reversed: + ref = torch.flip(ref, dims=[0]) + + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(ref - output))}") + torch.testing.assert_close(ref.to(torch.float32), output, atol=0, rtol=0) + + +if __name__ == "__main__": + test_scan_1d(32, 2, False, "cpu") + test_scan_1d(32, 2, True, "cpu") diff --git a/third_party/wafer/examples/test_scan2d.py b/third_party/wafer/examples/test_scan2d.py new file mode 100755 index 00000000..bda2fca4 --- /dev/null +++ b/third_party/wafer/examples/test_scan2d.py @@ -0,0 +1,222 @@ +import textwrap + +import numpy as np +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + +import inspect +from numpy.random import RandomState +import benchmark + +from triton._internal_testing import ( + int_dtypes, + is_interpreter, + numpy_random, + to_triton, + to_numpy, +) + + +def patch_kernel(template, to_replace): + if is_interpreter(): + local_namespace = {} + src = textwrap.dedent(inspect.getsource(template.fn)) + for k, v in to_replace.items(): + src = src.replace(k, v) + exec(src, globals(), local_namespace) + return local_namespace[template.fn.__name__] + else: + kernel = triton.JITFunction(template.fn) + for key, value in to_replace.items(): + kernel._unsafe_update_src(kernel.src.replace(key, value)) + return kernel + + +def check_type_supported(dtype, device): + ''' + skip test if dtype is not supported on the current device + ''' + if device in ['cuda']: + cc = torch.cuda.get_device_capability() + if cc[0] < 8 and (dtype is tl.bfloat16 or dtype == "bfloat16" or dtype is torch.bfloat16): + pytest.skip("bfloat16 is only supported on NVGPU with cc >= 80") + if cc[0] < 9 and dtype in {tl.float8e4nv, "float8e4nv", "float8_e4m3fn"}: + pytest.skip("float8e4nv is only supported on NVGPU with cc >= 90") + if is_interpreter(): + if dtype in [tl.bfloat16, "bfloat16", torch.bfloat16]: + pytest.skip("bfloat16 is not supported in the interpreter") + + +# scan2d_shapes = [(8, 32), (16, 32), (32, 16), (2, 1024), (1024, 2), (32, 32), (1, 1024)] +scan2d_shapes = [(8, 32)] + +scan_configs = [(op, type, shape, axis, reverse) + # for type in ['int32', 'float32', 'bfloat16'] + for type in ['float32'] + for axis in [1, 0] + for reverse in [True, False] + for shape in scan2d_shapes + # for op in ['cumsum', 'cumprod', 'get_first_element', 'linear_recurrence', 'cummax', 'roll']] + for op in ['cumsum']] +# negative_config = [('cumsum', 'float32', (32, 32), -1, False)] +negative_config = [] + + +@pytest.mark.interpreter +@pytest.mark.parametrize("op, dtype_str, shape, axis, reverse", scan_configs + negative_config) +def test_scan2d(op, dtype_str, shape, axis, reverse, device): + check_type_supported(dtype_str, device) + if dtype_str == 'bfloat16': + if op == 'cummax': + pytest.skip("bfloat16 compare not supported before sm90") + if op == 'linear_recurrence': + pytest.skip("Skipping linear_recurrence scan on bfloat16 due to accuracy issues") + numpy_dtype_str = 'float32' if dtype_str == 'bfloat16' else dtype_str + + # triton kernel + @triton.jit + def kernel(X, Y, Z, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, AXIS: tl.constexpr): + range_m = tl.arange(0, BLOCK_M) + range_n = tl.arange(0, BLOCK_N) + x = tl.load(X + range_m[:, None] * BLOCK_N + range_n[None, :]) + y = tl.load(Y + range_m[:, None] * BLOCK_N + range_n[None, :]) + GENERATE_TEST_HERE + tl.store(Z + range_m[:, None] * BLOCK_N + range_n[None, :], z) + + if op == 'cumsum' or op == 'cumprod': + kernel = patch_kernel(kernel, {'GENERATE_TEST_HERE': f'z = tl.{op}(x, axis={axis}, reverse={reverse})'}) + elif op == 'get_first_element': + kernel = patch_kernel( + kernel, + {'GENERATE_TEST_HERE': f'z = tl.associative_scan(x, axis={axis}, combine_fn={op}, reverse={reverse})'}) + elif op == 'cummax': + rg = "range_m[:, None]" if axis == 0 else "range_n[None, :]" + rg = f"tl.broadcast_to({rg}.to(tl.int64), [BLOCK_M, BLOCK_N])" + kernel = patch_kernel(kernel, { + 'GENERATE_TEST_HERE': + f'_, z = tl.associative_scan((x, {rg}), axis={axis}, combine_fn={op}, reverse={reverse})' + }) + elif op == 'roll': + assert op == 'roll' + kernel = patch_kernel( + kernel, { + 'GENERATE_TEST_HERE': + f'_, z, _ = tl.associative_scan((1 + 0* x, 0 * x, x), axis={axis}, combine_fn={op}, reverse={reverse})' + }) + else: + assert op == 'linear_recurrence' + kernel = patch_kernel(kernel, { + 'GENERATE_TEST_HERE': + f'_, z = tl.associative_scan((x, y), axis={axis}, combine_fn={op}, reverse={reverse})' + }) + # input + rs = RandomState(17) + if op == 'linear_recurrence' and dtype_str in int_dtypes: + # If the numbers are too large the op will overflow + # We sample numbers in -1, 0, 1 + x = rs.randint(-1, 2, shape, dtype=dtype_str) + y = rs.randint(-1, 2, shape, dtype=dtype_str) + else: + # x = numpy_random(shape, dtype_str=dtype_str, rs=rs) + x = np.arange(0, np.prod(shape), dtype=np.float32).reshape(shape) + print(x) + # y is just used in linear_recurrence + y = numpy_random(shape, dtype_str=dtype_str, rs=rs) + x_in = x + if reverse: + x_in = np.flip(x, axis) + z = np.empty_like(x) + x_tri = to_triton(x, device=device, dst_type=dtype_str) + y_tri = to_triton(y, device=device, dst_type=dtype_str) + if op == 'cumsum' or op == 'cumprod': + numpy_op = {'cumsum': np.cumsum, 'cumprod': np.cumprod}[op] + z_ref = numpy_op(x_in, axis=axis).astype(getattr(np, numpy_dtype_str)) + if reverse: + z_ref = np.flip(z_ref, axis) + + elif op == 'cummax': + # NumPy does not have cummax + z = z.astype(np.int64) + z_ref = torch.cummax(torch.from_numpy(x_in.copy()), axis=axis).indices.numpy() + if reverse: + z_ref = x_in.shape[axis] - np.flip(z_ref, axis) - 1 + elif op == 'roll': + ROLL = 1 + z_ref = np.roll(x_in.copy(), ROLL, axis=axis) + if axis == 0: + z_ref[:ROLL] = 0 + else: + z_ref[:, :ROLL] = 0 + + if reverse: + z_ref = np.flip(z_ref, axis) + elif op == 'linear_recurrence': + # Simplify to the axis=1 case + x_ref = x.T if axis == 0 else x + y_ref = y.T if axis == 0 else y + if reverse: + x_ref = np.flip(x_ref, 1) + y_ref = np.flip(y_ref, 1) + + result = [] + for x_refi, y_refi in zip(x_ref, y_ref): + li = [] + acc = 0 + for xi, yi in zip(x_refi, y_refi): + acc = xi * acc + yi + li.append(acc) + result.append(li) + z_ref = np.array(result) + if reverse: + z_ref = np.flip(z_ref, 1) + + if axis == 0: + z_ref = z_ref.T + else: + assert op == 'get_first_element' + z_ref = x + if axis == 0: + if reverse: + z_ref[:-1] = x[-1] + else: + z_ref[1:] = x[0] + else: + if reverse: + z_ref[:, :-1] = x[:, -1:] + else: + z_ref[:, 1:] = x[:, 0:1] + + # triton result + # we don't cast the `fp32 = bf16 op bf16` result to bfloat16 to alleviate accuracy issues + z_tri = to_triton(z, device=device) + x_tri_txda = x_tri.to("txda") + y_tri_txda = y_tri.to("txda") + z_tri_txda = z_tri.to("txda") + kernel[(1, )](x_tri_txda, y_tri_txda, z_tri_txda, BLOCK_M=shape[0], BLOCK_N=shape[1], AXIS=axis) + with torch.no_grad(): + z_tri.copy_(z_tri_txda.cpu()) + + z_tri = to_numpy(z_tri) + # compare + if dtype_str not in int_dtypes: + if op == 'cumprod': + np.testing.assert_allclose(z_ref, z_tri, rtol=0.01, atol=1e-3) + else: + np.testing.assert_allclose(z_ref, z_tri, rtol=0.01) + else: + np.testing.assert_equal(z_ref, z_tri) + + +if __name__ == "__main__": + test_scan2d('cumsum', 'float32', (8, 32), 1, False, 'cpu') + test_scan2d('cumsum', 'float32', (8, 32), 0, True, 'cpu') + test_scan2d('cumsum', 'float32', (8, 32), 0, False, 'cpu') + # test_scan2d('cumprod', 'float32', (8, 32), 1, False, 'cpu') + # test_scan2d('get_first_element', 'float32', (8, 32), 1, False, 'cpu') + # test_scan2d('linear_recurrence', 'float32', (8, 32), 1, False, 'cpu') + # test_scan2d('cummax', 'float32', (8, 32), 1, False, 'cpu') + # test_scan2d('roll', 'float32', (8, 32), 1, False, 'cpu') diff --git a/third_party/wafer/examples/test_scan3d.py b/third_party/wafer/examples/test_scan3d.py new file mode 100755 index 00000000..cc9c428a --- /dev/null +++ b/third_party/wafer/examples/test_scan3d.py @@ -0,0 +1,168 @@ +import textwrap + +import numpy as np +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + +import inspect +from numpy.random import RandomState +import benchmark + +from triton._internal_testing import ( + int_dtypes, + is_interpreter, + numpy_random, + to_triton, + to_numpy, +) + + +def patch_kernel(template, to_replace): + if is_interpreter(): + local_namespace = {} + src = textwrap.dedent(inspect.getsource(template.fn)) + for k, v in to_replace.items(): + src = src.replace(k, v) + exec(src, globals(), local_namespace) + return local_namespace[template.fn.__name__] + else: + kernel = triton.JITFunction(template.fn) + for key, value in to_replace.items(): + kernel._unsafe_update_src(kernel.src.replace(key, value)) + return kernel + + +def check_type_supported(dtype, device): + ''' + skip test if dtype is not supported on the current device + ''' + if device in ['cuda']: + cc = torch.cuda.get_device_capability() + if cc[0] < 8 and (dtype is tl.bfloat16 or dtype == "bfloat16" or dtype is torch.bfloat16): + pytest.skip("bfloat16 is only supported on NVGPU with cc >= 80") + if cc[0] < 9 and dtype in {tl.float8e4nv, "float8e4nv", "float8_e4m3fn"}: + pytest.skip("float8e4nv is only supported on NVGPU with cc >= 90") + if is_interpreter(): + if dtype in [tl.bfloat16, "bfloat16", torch.bfloat16]: + pytest.skip("bfloat16 is not supported in the interpreter") + + +# scan2d_shapes = [(8, 32), (16, 32), (32, 16), (2, 1024), (1024, 2), (32, 32), (1, 1024)] +scan3d_shapes = [((2, 2, 32))] + +scan_configs = [(op, type, shape, axis, reverse) + # for type in ['int32', 'float32', 'bfloat16'] + for type in ['float32'] + # for axis in [1, 0] + for axis in [0, 1, 2] + for reverse in [True, False] + # for reverse in [True] + for shape in scan3d_shapes + # for op in ['cumsum', 'cumprod', 'cummax', 'linear_recurrence'] + for op in ['cumsum']] +# negative_config = [('cumsum', 'float32', (32, 32), -1, False)] +negative_config = [] + + +@pytest.mark.interpreter +@pytest.mark.parametrize("op, dtype_str, shape, axis, reverse", scan_configs + negative_config) +def test_scan3d(op, dtype_str, shape, axis, reverse, device): + check_type_supported(dtype_str, device) + + numpy_dtype_str = 'float32' if dtype_str == 'bfloat16' else dtype_str + + # triton kernel + @triton.jit + def kernel(X, Y, Z, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, AXIS: tl.constexpr): + range_m = tl.arange(0, BLOCK_M) + range_n = tl.arange(0, BLOCK_N) + range_k = tl.arange(0, BLOCK_K) + x = tl.load(X + range_m[:, None, None] * BLOCK_N * BLOCK_K + range_n[None, :, None] * BLOCK_K + + range_k[None, None, :]) + y = tl.load(Y + range_m[:, None, None] * BLOCK_N * BLOCK_K + range_n[None, :, None] * BLOCK_K + + range_k[None, None, :]) + GENERATE_TEST_HERE + tl.store( + Z + range_m[:, None, None] * BLOCK_N * BLOCK_K + range_n[None, :, None] * BLOCK_K + range_k[None, None, :], + z) + + if op == 'cumsum' or op == 'cumprod': + kernel = patch_kernel(kernel, {'GENERATE_TEST_HERE': f'z = tl.{op}(x, axis={axis}, reverse={reverse})'}) + elif op == 'cummax': + rg = "range_m[:, None]" if axis == 0 else "range_n[None, :]" + rg = f"tl.broadcast_to({rg}.to(tl.int64), [BLOCK_M, BLOCK_N])" + kernel = patch_kernel(kernel, { + 'GENERATE_TEST_HERE': + f'_, z = tl.associative_scan((x, {rg}), axis={axis}, combine_fn={op}, reverse={reverse})' + }) + else: + assert 1 == 0 + + # input + rs = RandomState(17) + + # x = numpy_random(shape, dtype_str=dtype_str, rs=rs) + x = np.arange(0, np.prod(shape), dtype=np.float32).reshape(shape) + print(x) + # y is just used in linear_recurrence + y = numpy_random(shape, dtype_str=dtype_str, rs=rs) + + x_in = x + if reverse: + x_in = np.flip(x, axis) + z = np.empty_like(x) + x_tri = to_triton(x, device=device, dst_type=dtype_str) + y_tri = to_triton(y, device=device, dst_type=dtype_str) + if op == 'cumsum' or op == 'cumprod': + numpy_op = {'cumsum': np.cumsum, 'cumprod': np.cumprod}[op] + z_ref = numpy_op(x_in, axis=axis).astype(getattr(np, numpy_dtype_str)) + if reverse: + z_ref = np.flip(z_ref, axis) + + elif op == 'cummax': + # NumPy does not have cummax + z = z.astype(np.int64) + z_ref = torch.cummax(torch.from_numpy(x_in.copy()), axis=axis).indices.numpy() + if reverse: + z_ref = x_in.shape[axis] - np.flip(z_ref, axis) - 1 + + else: + assert 1 == 0 + + # triton result + # we don't cast the `fp32 = bf16 op bf16` result to bfloat16 to alleviate accuracy issues + z_tri = to_triton(z, device=device) + x_tri_txda = x_tri.to("txda") + y_tri_txda = y_tri.to("txda") + z_tri_txda = z_tri.to("txda") + kernel[(1, )](x_tri_txda, y_tri_txda, z_tri_txda, BLOCK_M=shape[0], BLOCK_N=shape[1], BLOCK_K=shape[2], AXIS=axis) + with torch.no_grad(): + z_tri.copy_(z_tri_txda.cpu()) + + z_tri = to_numpy(z_tri) + + print(f"The maximum difference between torch and triton is " + f"{np.max(np.abs(z_ref - z_tri))}") + # compare + if dtype_str not in int_dtypes: + if op == 'cumprod': + np.testing.assert_allclose(z_ref, z_tri, rtol=0.01, atol=1e-3) + else: + np.testing.assert_allclose(z_ref, z_tri, rtol=0.01) + else: + np.testing.assert_equal(z_ref, z_tri) + + +if __name__ == "__main__": + test_scan3d('cumsum', 'float32', (2, 2, 32), 2, False, 'cpu') + test_scan3d('cumsum', 'float32', (2, 2, 32), 0, False, 'cpu') + test_scan3d('cumsum', 'float32', (2, 2, 32), 1, False, 'cpu') + # test_scan2d('cumprod', 'float32', (8, 32), 1, False, 'cpu') + # test_scan2d('get_first_element', 'float32', (8, 32), 1, False, 'cpu') + # test_scan2d('linear_recurrence', 'float32', (8, 32), 1, False, 'cpu') + # test_scan2d('cummax', 'float32', (8, 32), 1, False, 'cpu') + # test_scan2d('roll', 'float32', (8, 32), 1, False, 'cpu') diff --git a/third_party/wafer/examples/test_sigmoid.py b/third_party/wafer/examples/test_sigmoid.py new file mode 100755 index 00000000..b7840d47 --- /dev/null +++ b/third_party/wafer/examples/test_sigmoid.py @@ -0,0 +1,96 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def sigmoid_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.sigmoid(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def sigmoid_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + sigmoid_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_sigmoid(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = sigmoid_triton(x) + + # Verify the output + expected = torch.sigmoid(x) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_sigmoid_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = sigmoid_triton(x) + + # Verify the output + expected = torch.sigmoid(x) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_sigmoid_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_sign_extend.py b/third_party/wafer/examples/test_sign_extend.py new file mode 100755 index 00000000..b4dd3596 --- /dev/null +++ b/third_party/wafer/examples/test_sign_extend.py @@ -0,0 +1,45 @@ +import torch +import torch_txda # noqa: F401 + +import triton + +import triton.language as tl + + + +@triton.jit +def sign_extend(off, in0, out0, in0_size): + offset = tl.load(off).to(tl.int64) + offsets = offset + tl.arange(0, 4) + a = tl.load(in0 + offsets, mask=offsets < in0_size, other=11) + tl.store(out0 + tl.arange(0, 4), a) + + +def compile(): + src = triton.compiler.ASTSource( + fn=sign_extend, + signature={'off': '*i32', 'in0': '*fp32', 'out0': '*fp32', 'in0_size': 'i32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) + + +def test_sign_extend(device): + if device == 'cpu': + pass # Wafer driver is selected by conftest.py. + + SIZE = 4 + offsets = torch.full((1, ), 1, device="cpu", dtype=torch.int32) + input = torch.arange(0, SIZE, device="cpu", dtype=torch.int32) + output = torch.full((SIZE, ), -1, device="cpu", dtype=torch.int32) + grid = lambda meta: (1, ) + print(output) + offsets_txda = offsets.to("txda") + input_txda = input.to("txda") + output_txda = output.to("txda") + sign_extend[grid](offsets_txda, input_txda, output_txda, SIZE) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(input) + print(output) + torch.testing.assert_close(torch.tensor([1, 2, 3, 11], device="cpu", dtype=torch.int32), output) diff --git a/third_party/wafer/examples/test_sin.py b/third_party/wafer/examples/test_sin.py new file mode 100755 index 00000000..fae1202d --- /dev/null +++ b/third_party/wafer/examples/test_sin.py @@ -0,0 +1,104 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def sin_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.sin(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def sin_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + x = x.to(DEVICE) + output = output.to(DEVICE) + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + sin_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + output = output.to('cpu') + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_sin(size, dtype, device="cpu"): + # Create a random tensor + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = sin_triton(x) + + # Verify the output + expected = torch.sin(x) + + # compare + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(expected - output))}") + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_sin_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = sin_triton(x) + + # Verify the output + expected = torch.sin(x) + torch.testing.assert_close(output, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_sin_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_softmax.py b/third_party/wafer/examples/test_softmax.py new file mode 100755 index 00000000..95a4a35f --- /dev/null +++ b/third_party/wafer/examples/test_softmax.py @@ -0,0 +1,93 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def softmax_kernel(output_ptr, input_ptr, input_row_stride, output_row_stride, n_cols, BLOCK_SIZE: tl.constexpr): + # The rows of the softmax are independent, so we parallelize across those + row_idx = tl.program_id(0) + # The stride represents how much we need to increase the pointer to advance 1 row + row_start_ptr = input_ptr + row_idx * input_row_stride + # The block size is the next power of two greater than n_cols, so we can fit each + # row in a single block + col_offsets = tl.arange(0, BLOCK_SIZE) + input_ptrs = row_start_ptr + col_offsets + # Load the row into SRAM, using a mask since BLOCK_SIZE may be > than n_cols + row = tl.load(input_ptrs, mask=col_offsets < n_cols, other=-float('inf')) + # Subtract maximum for numerical stability + row_minus_max = row - tl.max(row, axis=0) + # Note that exponentiation in Triton is fast but approximate (i.e., think __expf in CUDA) + numerator = tl.exp(row_minus_max) + denominator = tl.sum(numerator, axis=0) + softmax_output = numerator / denominator + # Write back output to DRAM + output_row_start_ptr = output_ptr + row_idx * output_row_stride + output_ptrs = output_row_start_ptr + col_offsets + tl.store(output_ptrs, softmax_output, mask=col_offsets < n_cols) + + +def softmax(x): + n_rows, n_cols = x.shape + # The block size is the smallest power of two greater than the number of columns in `x` + BLOCK_SIZE = triton.next_power_of_2(n_cols) + # Another trick we can use is to ask the compiler to use more threads per row by + # increasing the number of warps (`num_warps`) over which each row is distributed. + # You will see in the next tutorial how to auto-tune this value in a more natural + # way so you don't have to come up with manual heuristics yourself. + num_warps = 4 + if BLOCK_SIZE >= 2048: + num_warps = 8 + if BLOCK_SIZE >= 4096: + num_warps = 16 + # Allocate output + y = torch.empty_like(x) + # Enqueue kernel. The 1D launch grid is simple: we have one kernel instance per row o + # f the input matrix + + x = x.to(DEVICE) + y = y.to(DEVICE) + y_txda = y.to("txda") + x_txda = x.to("txda") + softmax_kernel[(n_rows, )]( + y_txda, + x_txda, + x_txda.stride(0), + y_txda.stride(0), + n_cols, + num_warps=num_warps, + BLOCK_SIZE=BLOCK_SIZE, + ) + with torch.no_grad(): + y.copy_(y_txda.cpu()) + y = y.to('cpu') + return y + + +def test_softmax(device): + torch.manual_seed(0) + x = torch.randn(1823, 781, device="cpu") + y_triton = softmax(x) + y_torch = torch.softmax(x, axis=1) + assert torch.allclose(y_triton, y_torch), (y_triton, y_torch) + + +@benchmark.measure() +def bench_softmax(size, provider): + torch.manual_seed(0) + x = torch.randn(size, size, device="cpu") + if provider == 'torch': + torch.softmax(x, axis=1) + if provider == 'triton': + softmax(x) + + +if __name__ == "__main__": + for X in [2**i for i in range(10, 14, 1)]: + for provider in ['torch', 'triton']: + bench_softmax(X, provider) diff --git a/third_party/wafer/examples/test_sort.py b/third_party/wafer/examples/test_sort.py new file mode 100755 index 00000000..d13b7b0b --- /dev/null +++ b/third_party/wafer/examples/test_sort.py @@ -0,0 +1,36 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +from triton._internal_testing import numpy_random + + +# @pytest.mark.interpreter +# @pytest.mark.parametrize("M, N", [[1, 512], [8, 64], [256, 16], [512, 8]]) +@pytest.mark.parametrize("M, N", [[8, 64]]) +# @pytest.mark.parametrize("descending", [False, True]) +@pytest.mark.parametrize("descending", [False]) +# @pytest.mark.parametrize("dtype_str", ['int32', 'float16', 'float32', 'bfloat16']) +@pytest.mark.parametrize("dtype_str", ['float32']) +def test_sort(M, N, descending, dtype_str, device): + + @triton.jit + def sort_kernel(X, Z, N: tl.constexpr, M: tl.constexpr, descending: tl.constexpr): + offx = tl.arange(0, M) + offy = tl.arange(0, N) * M + off2d = offx[None, :] + offy[:, None] + x = tl.load(X + off2d) + x = tl.sort(x, descending=descending) + tl.store(Z + off2d, x) + + x = numpy_random((N, M), dtype_str=dtype_str) + x = torch.from_numpy(x) + y = torch.sort(x, descending=descending)[0] + z = torch.empty_like(x) + x_txda = x.to("txda") + z_txda = z.to("txda") + sort_kernel[(1, )](x_txda, z_txda, N, M, descending, num_warps=8) + with torch.no_grad(): + z.copy_(z_txda.cpu()) + assert (y == z).all(), (y, z) diff --git a/third_party/wafer/examples/test_splat.py b/third_party/wafer/examples/test_splat.py new file mode 100755 index 00000000..0c3ac496 --- /dev/null +++ b/third_party/wafer/examples/test_splat.py @@ -0,0 +1,45 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + +@triton.jit +def splat( + f32_val, + f32_out, + stride_row, + stride_col, + BLOCK_SIZE_ROW: tl.constexpr, + BLOCK_SIZE_COL: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + x = tl.full((2, BLOCK_SIZE_COL), f32_val, dtype=tl.float32) + offs_row = 2 * pid0 + tl.arange(0, 2) + offs_col = tl.arange(0, BLOCK_SIZE_COL) + a_ptrs = f32_out + (offs_row[:, None] * stride_row + offs_col[None, :] * stride_col) + tl.store(a_ptrs, x) + + +def test(device): + n_rows = 256 + n_cols = 512 + fill_value = 123.456 + expected_result = torch.full((n_rows, n_cols), fill_value, dtype=torch.float32) + output = torch.empty([n_rows, n_cols], device="cpu", dtype=expected_result.dtype) + grid = lambda meta: (n_rows // 2, ) + + output_txda = output.to("txda") + splat[grid]( + fill_value, + output_txda, + output_txda.stride(0), + output_txda.stride(1), + BLOCK_SIZE_ROW=n_rows, + BLOCK_SIZE_COL=n_cols, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + torch.testing.assert_close(output, expected_result, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/examples/test_sqrt.py b/third_party/wafer/examples/test_sqrt.py new file mode 100755 index 00000000..0989846c --- /dev/null +++ b/third_party/wafer/examples/test_sqrt.py @@ -0,0 +1,98 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def sqrt_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.sqrt(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def sqrt_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + sqrt_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_sqrt(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = sqrt_triton(x) + + # Verify the output + expected = torch.sqrt(x) + print(output) + print(expected) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_sqrt_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = sqrt_triton(x) + + # Verify the output + expected = torch.sqrt(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_sqrt_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_sqrt_rn.py b/third_party/wafer/examples/test_sqrt_rn.py new file mode 100755 index 00000000..5cc455af --- /dev/null +++ b/third_party/wafer/examples/test_sqrt_rn.py @@ -0,0 +1,97 @@ +import pytest +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def sqrt_rn_kernel( + x_ptr, + output_ptr, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + pid = tl.program_id(0) + + # Calculate the start and offsets + start = pid * BLOCK_SIZE + offsets = start + tl.arange(0, BLOCK_SIZE) + + # Create a mask to avoid out-of-bounds access + mask = offsets < n_elements + + # Load the input data + x = tl.load(x_ptr + offsets, mask=mask) + + # Compute the absolute value + out = tl.sqrt_rn(x) + + # Store the result + tl.store(output_ptr + offsets, out, mask=mask) + + +def sqrt_rn_triton(x): + # Get the number of elements + n_elements = x.numel() + + # Allocate output tensor + output = torch.empty_like(x) + + # Define block size + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']), ) + print("grid value is ", grid) + + # Launch the kernel + x_txda = x.to("txda") + output_txda = output.to("txda") + sqrt_rn_kernel[grid]( + x_txda, + output_txda, + n_elements, + BLOCK_SIZE=1024, + ) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + + return output + + +@pytest.mark.parametrize("size, dtype", [ # + (size, dtype) for size in [98432] for dtype in [torch.float32] +]) +def test_sqrt_rn(size, dtype, device="cpu"): + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = sqrt_rn_triton(x) + + # Verify the output + expected = torch.sqrt(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +@benchmark.measure() +def benchmark_sqrt_rn_triton(size, dtype, provider): + if provider != "triton": + raise ValueError("This benchmark is only for the Triton provider.") + + # Generate random input data + x = torch.randn(size, device="cpu", dtype=dtype) + + # Call the Triton kernel + output = sqrt_rn_triton(x) + + # Verify the output + # TODO:rounding_mode need double check + expected = torch.sqrt(x) + torch.testing.assert_close(output, expected, equal_nan=True, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + for size in [i**2 for i in range(22, 25, 1)]: + benchmark_sqrt_rn_triton(size, torch.float32, "triton") diff --git a/third_party/wafer/examples/test_swap.py b/third_party/wafer/examples/test_swap.py new file mode 100755 index 00000000..5c690030 --- /dev/null +++ b/third_party/wafer/examples/test_swap.py @@ -0,0 +1,50 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + +# The purpose of this kernel and test is to catch incorrectly optimized kernels +# where copy elimination happens erroneously in the absence of explicit memory allocation. +# Such optimization bugs can result in incorrect behavior when swapping two arrays, +# particularly when both arrays unintentionally end up with the same data due to +# missing intermediate storage or mismanaged memory access. + + +@triton.jit +def swap_kernel(x_ptr, # *Pointer* to first inout vector. + y_ptr, # *Pointer* to second inout vector. + BLOCK_SIZE: tl.constexpr, # Number of elements each program should process. + # NOTE: `constexpr` so it can be used as a shape value. + ): + pid = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0. + block_start = pid * BLOCK_SIZE + offsets = block_start + tl.arange(0, BLOCK_SIZE) + x = tl.load(x_ptr + offsets) + y = tl.load(y_ptr + offsets) + tl.store(x_ptr + offsets, y) + tl.store(y_ptr + offsets, x) + + +def swap(x: torch.Tensor, y: torch.Tensor): + n_elements = x.numel() + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]), ) + x_txda = x.to("txda") + y_txda = y.to("txda") + swap_kernel[grid](x_txda, y_txda, BLOCK_SIZE=1024) + with torch.no_grad(): + x.copy_(x_txda.cpu()) + y.copy_(y_txda.cpu()) + + +def test(device): + torch.manual_seed(0) + size = 10240 + x = torch.rand(size, device="cpu") + y = torch.rand(size, device="cpu") + assert not torch.equal(x, y) + x_ = x.clone() + y_ = y.clone() + swap(x, y) + assert torch.equal(x, y_) + assert torch.equal(y, x_) diff --git a/third_party/wafer/examples/test_swizzle2d.py b/third_party/wafer/examples/test_swizzle2d.py new file mode 100755 index 00000000..6e6137a7 --- /dev/null +++ b/third_party/wafer/examples/test_swizzle2d.py @@ -0,0 +1,26 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl + + +# @pytest.mark.interpreter +@pytest.mark.parametrize("size_i, size_j, size_g", [[5, 7, 3]]) +def test_swizzle2d(size_i, size_j, size_g, device): + + @triton.jit + def swizzle2d_kernel(output, size_i, size_j, size_g): + for i in tl.range(0, size_i, 1): + for j in tl.range(0, size_j, 1): + new_i, new_j = tl.swizzle2d(i, j, size_i, size_j, size_g) + tl.store(output + new_i * size_j + new_j, i * size_j + j) + + output = torch.zeros(size_i, size_j) + output_txda = output.to("txda") + swizzle2d_kernel[(1, )](output_txda, size_i, size_j, size_g) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + expected_order = torch.tensor([[0, 3, 6, 9, 12, 15, 18], [1, 4, 7, 10, 13, 16, 19], [2, 5, 8, 11, 14, 17, 20], + [21, 23, 25, 27, 29, 31, 33], [22, 24, 26, 28, 30, 32, 34]]) + assert (output == expected_order).all(), (output, expected_order) diff --git a/third_party/wafer/examples/test_tensor_index_iterargs.py b/third_party/wafer/examples/test_tensor_index_iterargs.py new file mode 100755 index 00000000..f9a3afeb --- /dev/null +++ b/third_party/wafer/examples/test_tensor_index_iterargs.py @@ -0,0 +1,126 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl + + + +def test_tensor_indices_nested_with_mask(device): + + @triton.jit + def addptr_with_masks(in0, out0, mask_bound): + offs = tl.arange(0, 4) + out_offs = tl.arange(0, 4) + # We're loading 16 elements here, the bound is set to 14 so that + # the mask only applies to the last iteration's load + # TODO: The current mask implementation in triton-shared does not seem + # to work when the mask applies to the entire tensor load, perhaps + # the lowerings for subviews with 0-dimensions do not work? + for i in range(0, 4): + mask = offs < mask_bound + a = tl.load(in0 + offs, mask=mask, other=-11) + tl.store(out0 + out_offs, a) + offs += 4 + out_offs += 4 + + SIZE = 17 + input = torch.arange(0, SIZE, device="cpu", dtype=torch.int32) + output = torch.full((SIZE, ), -1, device="cpu", dtype=torch.int32) + + if device == 'cpu': + pass # Wafer driver is selected by conftest.py. + + grid = lambda meta: (1, ) + + print(output) + input_txda = input.to("txda") + output_txda = output.to("txda") + addptr_with_masks[grid](input_txda, output_txda, 14) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + expected_output = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, -11, -11, -1], dtype=torch.int32, + device="cpu") + torch.testing.assert_close(output, expected_output) + print(input) + print(output) + + +def test_tensor_indices_nested(device): + + @triton.jit + def tensor_indices_nested(in0, out0): + offs = tl.arange(0, 4) + out_offs = tl.arange(0, 4) + for i in range(0, 2): + offs += i * 2 + a = tl.load(in0 + offs) + tl.store(out0 + out_offs, a) + offs += 4 + out_offs += 4 + for j in range(0, 3): + offs += j * 3 + a = tl.load(in0 + offs) + tl.store(out0 + out_offs, a) + offs += 4 + out_offs += 4 + + SIZE = 64 + input = torch.arange(0, SIZE, device="cpu", dtype=torch.int32) + output = torch.full((SIZE, ), -1, device="cpu", dtype=torch.int32) + + if device == 'cpu': + pass # Wafer driver is selected by conftest.py. + + grid = lambda meta: (1, ) + + print(output) + input_txda = input.to("txda") + output_txda = output.to("txda") + tensor_indices_nested[grid](input_txda, output_txda) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + expected_output = torch.tensor([ + 0, 1, 2, 3, 4, 5, 6, 7, 11, 12, 13, 14, 21, 22, 23, 24, 27, 28, 29, 30, 31, 32, 33, 34, 38, 39, 40, 41, 48, 49, + 50, 51, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1 + ], device="cpu", dtype=torch.int32) + torch.testing.assert_close(output, expected_output) + print(input) + print(output) + + +def test_integer_tensor(device): + + @triton.jit + def test_1(out0): + offs = tl.arange(0, 4) + out_offs = tl.arange(0, 4) + for i in range(0, 2): + tl.store(out0 + out_offs, offs) + out_offs += 4 + offs += 4 + + SIZE = 8 + input = torch.arange(0, SIZE, device="cpu", dtype=torch.int32) + output = torch.full((SIZE, ), -1, device="cpu", dtype=torch.int32) + + if device == 'cpu': + pass # Wafer driver is selected by conftest.py. + + grid = lambda meta: (1, ) + + print(output) + output_txda = output.to("txda") + test_1[grid](output_txda) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + print(input) + print(output) + torch.testing.assert_close(input, output) + src = triton.compiler.ASTSource( + fn=test_1, + signature={'out0': '*fp32'}, + ) + ret = triton.compile(src, ) + print(ret.asm["ttir"]) diff --git a/third_party/wafer/examples/test_trans2d.py b/third_party/wafer/examples/test_trans2d.py new file mode 100755 index 00000000..bcb88c8a --- /dev/null +++ b/third_party/wafer/examples/test_trans2d.py @@ -0,0 +1,38 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import benchmark +import itertools +import math + + +@pytest.mark.parametrize("dtype_str", ["int32", "int8"]) +@pytest.mark.parametrize("shape", [(2, 4), (16, 16)]) +@pytest.mark.parametrize("perm", list(itertools.permutations([0, 1]))) +def test_trans_2d(dtype_str, shape, perm, device): + + @triton.jit + def kernel(In, Out, in_shape1: tl.constexpr, in_shape2: tl.constexpr, ou_shape1: tl.constexpr, + ou_shape2: tl.constexpr, trans1: tl.constexpr, trans2: tl.constexpr): + in_offs = tl.arange(0, in_shape1)[:, None] * in_shape2 + tl.arange(0, in_shape2)[None, :] + ou_offs = tl.arange(0, ou_shape1)[:, None] * ou_shape2 + tl.arange(0, ou_shape2)[None, :] + tl.store(Out + ou_offs, tl.permute(tl.load(In + in_offs), (trans1, trans2))) + + input = torch.arange(math.prod(shape), dtype=getattr(torch, dtype_str), device="cpu").reshape(shape) + expected = torch.permute(input, perm) + # Don't do zeros_like -- that copies the layout, which we don't want. + actual = torch.zeros(expected.shape, dtype=getattr(torch, dtype_str), device="cpu") + + input_txda = input.to("txda") + actual_txda = actual.to("txda") + kernel[(1, )](input_txda, actual_txda, *shape, *[shape[i] for i in perm], *perm) + with torch.no_grad(): + actual.copy_(actual_txda.cpu()) + + torch.testing.assert_close(actual, expected, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + test_trans_2d('float32', (32, 16), (1, 0), 'cpu') diff --git a/third_party/wafer/examples/test_umulhi.py b/third_party/wafer/examples/test_umulhi.py new file mode 100755 index 00000000..744a6613 --- /dev/null +++ b/third_party/wafer/examples/test_umulhi.py @@ -0,0 +1,82 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import numpy as np +from numpy.random import RandomState + +from triton._internal_testing import ( + integral_dtypes, + int_dtypes, + str_to_triton_dtype, + uint_dtypes, + float_dtypes, + float_dtypes_with_bfloat16, + dtypes, + dtypes_with_bfloat16, + is_cuda, + is_interpreter, + is_hopper, + is_hip, + is_hip_cdna, + is_hip_cdna2, + is_hip_cdna3, + is_hip_cdna4, + is_xpu, + get_arch, + torch_float8_dtypes, + torch_dtypes, + numpy_random, + to_triton, + torch_dtype_name, + to_numpy, +) + + +@pytest.mark.interpreter +@pytest.mark.parametrize("dtype_str", ['int32']) +def test_umulhi(dtype_str, device): + + @triton.jit + def kernel(X, Y, Z, N: tl.constexpr): + offs = tl.arange(0, N) + x = tl.load(X + offs) + y = tl.load(Y + offs) + z = tl.umulhi(x, y) + tl.store(Z + tl.arange(0, N), z) + + def umulhi32(a, b): + # Convert to 64-bit unsigned integers to prevent overflow + a_64 = a.astype(np.int64) + b_64 = b.astype(np.int64) + + # Perform the multiplication in 64-bit + product_64 = a_64 * b_64 + + # Shift right by 32 bits to get the high part of the product + result_high_32 = product_64 >> 32 + return result_high_32 + + rs = RandomState(17) + N = 128 + x = numpy_random((N, ), dtype_str=dtype_str, rs=rs, low=0) + x_tri = to_triton(x, device=device) + y = numpy_random((N, ), dtype_str=dtype_str, rs=rs, low=0) + y_tri = to_triton(y, device=device) + z_tri = torch.zeros_like(x_tri) + x_tri_txda = x_tri.to("txda") + y_tri_txda = y_tri.to("txda") + z_tri_txda = z_tri.to("txda") + kernel[(1, )](x_tri_txda, y_tri_txda, z_tri_txda, N=N) + with torch.no_grad(): + z_tri.copy_(z_tri_txda.cpu()) + + z_ref = umulhi32(x, y) + np.testing.assert_equal(z_ref, to_numpy(z_tri)) + + +if __name__ == "__main__": + # Run the test + # pytest.main([__file__, "-v", "-s", "--tb=short"]) + test_umulhi('int32', 'cpu') diff --git a/third_party/wafer/examples/test_vec_add.py b/third_party/wafer/examples/test_vec_add.py new file mode 100755 index 00000000..100ec736 --- /dev/null +++ b/third_party/wafer/examples/test_vec_add.py @@ -0,0 +1,98 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +@triton.jit +def add_kernel(x_ptr, # *Pointer* to first input vector. + y_ptr, # *Pointer* to second input vector. + output_ptr, # *Pointer* to output vector. + n_elements, # Size of the vector. + BLOCK_SIZE: tl.constexpr, # Number of elements each program should process. + # NOTE: `constexpr` so it can be used as a shape value. + ): + # There are multiple 'programs' processing different data. We identify which program + # we are here: + pid = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0. + # This program will process inputs that are offset from the initial data. + # For instance, if you had a vector of length 256 and block_size of 64, the programs + # would each access the elements [0:64, 64:128, 128:192, 192:256]. + # Note that offsets is a list of pointers: + block_start = pid * BLOCK_SIZE + offsets = block_start + tl.arange(0, BLOCK_SIZE) + # Create a mask to guard memory operations against out-of-bounds accesses. + mask = offsets < n_elements + # Load x and y from DRAM, masking out any extra elements in case the input is not a + # multiple of the block size. + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + output = x + y + # Write x + y back to DRAM. + tl.store(output_ptr + offsets, output, mask=mask) + + +def add(x: torch.Tensor, y: torch.Tensor): + output_torch = x.cpu() + y.cpu() + x = x.to(DEVICE) + y = y.to(DEVICE) + # We need to preallocate the output. + output = torch.empty_like(x) + # assert x.is_cuda and y.is_cuda and output.is_cuda + n_elements = output.numel() + # The SPMD launch grid denotes the number of kernel instances that run in parallel. + # It is analogous to CUDA launch grids. It can be either Tuple[int], or Callable(metaparameters) -> Tuple[int]. + # In this case, we use a 1D grid where the size is the number of blocks: + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]), ) + # NOTE: + # - Each torch.tensor object is implicitly converted into a pointer to its first element. + # - `triton.jit`'ed functions can be indexed with a launch grid to obtain a callable GPU kernel. + # - Don't forget to pass meta-parameters as keywords arguments. + add_kernel[grid](x, y, output, n_elements, BLOCK_SIZE=1024) + # The production Wafer launcher synchronizes; read back for the CPU oracle. + output = output.to("cpu") + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(output_torch - output))}") + return output + + +def test(device): + torch.manual_seed(0) + size = 1024 + x = torch.rand(size, device="cpu") + y = torch.rand(size, device="cpu") + print("x: ", x) + print("y: ", y) + output_torch = x + y + x = x.to(device) + y = y.to(device) + output_triton = add(x, y) + print("output_triton device: ", output_triton.device) + # TODO: need to check some conditions otherwise the code below does not make any difference for the test + output_triton = output_triton.to("cpu") + torch.testing.assert_close(output_triton, output_torch, atol=1e-5, rtol=0) + print("expected", output_torch) + print("actual", output_triton) + print(f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(output_torch - output_triton))}") + + +@benchmark.measure() +def bench_vecadd(size, provider): + a = torch.rand(size, device="cpu", dtype=torch.float32) + b = torch.rand(size, device="cpu", dtype=torch.float32) + if provider == 'torch': + a + b + if provider == 'triton': + add(a, b) + + +if __name__ == "__main__": + # test(DEVICE) + for X in [2**i for i in range(8, 25, 1)]: + for provider in ['torch', 'triton']: + bench_vecadd(X, provider) diff --git a/third_party/wafer/examples/test_where.py b/third_party/wafer/examples/test_where.py new file mode 100755 index 00000000..763b338c --- /dev/null +++ b/third_party/wafer/examples/test_where.py @@ -0,0 +1,69 @@ +import torch +import torch_txda # noqa: F401 + +import triton +import triton.language as tl +import benchmark +import numpy as np + + +@triton.jit +def where_kernel(a_ptr, b_ptr, output_ptr, n_elements, + BLOCK_SIZE: tl.constexpr, # Number of elements each program should process. + # NOTE: `constexpr` so it can be used as a shape value. + ): + offsets = tl.program_id(axis=0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_elements + + # Load x and y from DRAM, masking out any extra elements in case the input is not a + # multiple of the block size. + a = tl.load(a_ptr + offsets, mask=mask) + b = tl.load(b_ptr + offsets, mask=mask) + + decide = a > (b + 0) + output = tl.where(decide, a, b) + tl.store(output_ptr + offsets, output, mask=mask) + + +def where(x: torch.Tensor, y: torch.Tensor): + # We need to preallocate the output. + output = torch.empty_like(x) + # assert x.is_cuda and y.is_cuda and output.is_cuda + n_elements = output.numel() + # The SPMD launch grid denotes the number of kernel instances that run in parallel. + # It is analogous to CUDA launch grids. It can be either Tuple[int], or Callable(metaparameters) -> Tuple[int]. + # In this case, we use a 1D grid where the size is the number of blocks: + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]), ) + # NOTE: + # - Each torch.tensor object is implicitly converted into a pointer to its first element. + # - `triton.jit`'ed functions can be indexed with a launch grid to obtain a callable GPU kernel. + # - Don't forget to pass meta-parameters as keywords arguments. + x_txda = x.to("txda") + y_txda = y.to("txda") + output_txda = output.to("txda") + where_kernel[grid](x_txda, y_txda, output_txda, n_elements, BLOCK_SIZE=1024) + with torch.no_grad(): + output.copy_(output_txda.cpu()) + # We return a handle to z but, since `torch.cuda.synchronize()` hasn't been called, the kernel is still + # running asynchronously at this point. + return output + + +def test_where_1(device): + torch.manual_seed(0) + size = 98432 + x = torch.rand(size, device="cpu", dtype=torch.float32) + y = 1 - x + # y = torch.rand(size, device=device, dtype=torch.float32) + cond = x > y + output_torch = torch.where(cond, x, y) + + output_triton = where(x, y) + + max = torch.maximum(x, y) + + torch.testing.assert_close(output_torch, output_triton, atol=1e-5, rtol=0) + + torch.testing.assert_close(output_torch, max, atol=1e-5, rtol=0) + + # torch.testing.assert_close(output_triton, x, atol=1e-5, rtol=0) diff --git a/third_party/wafer/examples/time1.py b/third_party/wafer/examples/time1.py new file mode 100755 index 00000000..7a46b02d --- /dev/null +++ b/third_party/wafer/examples/time1.py @@ -0,0 +1,210 @@ +import torch + +import triton +import triton.language as tl +# import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes: +# - A list of `triton.Config` objects that define different configurations of +# meta-parameters (e.g., `BLOCK_SIZE_M`) and compilation options (e.g., `num_warps`) to try +# - An auto-tuning *key* whose change in values will trigger evaluation of all the +# provided configs +@triton.jit +def matmul_kernel( + # Pointers to matrices + a_ptr, b_ptr, c_ptr, + # Matrix dimensions + M, N, K, + # The stride variables represent how much to increase the ptr by when moving by 1 + # element in a particular dimension. E.g. `stride_am` is how much to increase `a_ptr` + # by to get the element one row down (A has M rows). + stride_am, stride_ak, # + stride_bk, stride_bn, # + stride_cm, stride_cn, + # Meta-parameters + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, # + GROUP_SIZE_M: tl.constexpr, BLOCK_SIZE_N1: tl.constexpr, BLOCK_SIZE_N2: tl.constexpr, # + BLOCK_SIZE_K1: tl.constexpr, BLOCK_SIZE_K2: tl.constexpr, ACTIVATION: tl.constexpr # +): + """Kernel for computing the matmul C = A x B. + A has shape (M, K), B has shape (K, N) and C has shape (M, N) + """ + # ----------------------------------------------------------- + # Map program ids `pid` to the block of C it should compute. + # This is done in a grouped ordering to promote L2 data reuse. + # See above `L2 Cache Optimizations` section for details. + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + # ---------------------------------------------------------- + # Create pointers for the first blocks of A and B. + # We will advance this pointer as we move in the K direction + # and accumulate + # `a_ptrs` is a block of [BLOCK_SIZE_M, BLOCK_SIZE_K] pointers + # `b_ptrs` is a block of [BLOCK_SIZE_K, BLOCK_SIZE_N] pointers + # See above `Pointer Arithmetics` section for details + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + mask_m = offs_m < M + mask_n = offs_n < N + offs_nn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N1)[:, None] * BLOCK_SIZE_N2 + tl.arange( + 0, BLOCK_SIZE_N2)[None, :] + offs_kk = tl.arange(0, BLOCK_SIZE_K1)[:, None] * BLOCK_SIZE_K2 + tl.arange(0, BLOCK_SIZE_K2)[None, :] + mask_nn = offs_nn < N + a_ptrs = a_ptr + (offs_m[None, :, None] * stride_am + offs_kk[:, None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[None, :, None] * stride_bk + offs_nn[:, None, :] * stride_bn) + + # ----------------------------------------------------------- + # Iterate to compute a block of the C matrix. + # We accumulate into a `[BLOCK_SIZE_M, BLOCK_SIZE_N]` block + # of fp32 values for higher accuracy. + # `accumulator` will be converted back to fp16 after the loop. + # accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float16) + for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + # Load the next block of A and B, generate a mask by checking the K dimension. + # If it is out of bounds, set it to 0. + mask_k = offs_k < K - k * BLOCK_SIZE_K + mask_kk = offs_kk < K - k * BLOCK_SIZE_K + # a = tl.load(a_ptrs, mask=(mask_m[None, :, None] & mask_kk[:, None, :]), other=0.0) + # b = tl.load(b_ptrs, mask=(mask_k[None, :, None] & mask_nn[:, None, :]), other=0.0) + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + if BLOCK_SIZE_K1 != 1: + a = tl.trans(a, (1, 0, 2)) + a = tl.reshape(a, (BLOCK_SIZE_M, BLOCK_SIZE_K)) + if BLOCK_SIZE_N1 != 1: + b = tl.trans(b, (1, 0, 2)) + b = tl.reshape(b, (BLOCK_SIZE_K, BLOCK_SIZE_N)) + # We accumulate along the K dimension. + acc += tl.dot(a, b, out_dtype=tl.float16) + # Advance the ptrs to the next K block. + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + + if BLOCK_SIZE_N1 == 1: + acc = tl.reshape(acc, (BLOCK_SIZE_N1, BLOCK_SIZE_M, BLOCK_SIZE_N2)) + else: + acc = tl.reshape(acc, (BLOCK_SIZE_M, BLOCK_SIZE_N1, BLOCK_SIZE_N2)) + acc = tl.trans(acc, (1, 0, 2)) + + if ACTIVATION == "leaky_relu": + acc = leaky_relu(acc) + c = acc + + # You can fuse arbitrary activation functions here + # while the accumulator is still in FP32! + # if ACTIVATION == "leaky_relu": + # accumulator = leaky_relu(accumulator) + # c = accumulator.to(tl.float32) + + # ----------------------------------------------------------- + # Write back the block of the output matrix C with masks. + c_ptrs = c_ptr + stride_cm * offs_m[None, :, None] + stride_cn * offs_nn[:, None, :] + tl.store(c_ptrs, c) + # tl.store(c_ptrs, c, mask=(mask_m[None, :, None] & mask_nn[:, None, :])) + + +# We can fuse `leaky_relu` by providing it as an `ACTIVATION` meta-parameter in `_matmul`. +@triton.jit +def leaky_relu(x): + x = x + 1 + return tl.where(x >= 0, x, 0.01 * x) + + +def matmul(a, b, activation=""): + BLOCK_M = 1024 + BLOCK_N = 1024 + BLOCK_K = 128 + + # Check constraints. + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + assert a.is_contiguous(), "Matrix A must be contiguous" + assert b.is_contiguous(), "Matrix B must be contiguous" + M, K = a.shape + K, N = b.shape + + assert M % BLOCK_M == 0 and N % BLOCK_N == 0 and K % BLOCK_K == 0 + + ALIGN = 64 + BLOCK_K1 = max(1, BLOCK_K // ALIGN) + BLOCK_K2 = min(ALIGN, BLOCK_K) + BLOCK_N1 = max(1, BLOCK_N // ALIGN) + BLOCK_N2 = min(ALIGN, BLOCK_N) + # Allocates output. + c = torch.empty((M, N), device=a.device, dtype=a.dtype) + # 1D launch kernel where each block gets its own program. + grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), ) + matmul_kernel[grid]( + a, + b, + c, # + M, + N, + K, # + a.stride(0), + a.stride(1), # + b.stride(0), + b.stride(1), # + c.stride(0), + c.stride(1), # + ACTIVATION=activation, # + BLOCK_SIZE_M=BLOCK_M, + BLOCK_SIZE_N=BLOCK_N, + BLOCK_SIZE_K=BLOCK_K, + GROUP_SIZE_M=8, + BLOCK_SIZE_N1=BLOCK_N1, + BLOCK_SIZE_N2=BLOCK_N2, + BLOCK_SIZE_K1=BLOCK_K1, + BLOCK_SIZE_K2=BLOCK_K2, + ) + return c + + +# @benchmark.measure(repeats=20) +def bench_matmul(a, b): + x = a.to(DEVICE) + y = b.to(DEVICE) + z = matmul(x, y) + z = z.to("cpu") + return z + + +if __name__ == "__main__": + M = 4096 + K = 4096 + N = 4096 + a = torch.randn((M, K), device='cpu', dtype=torch.float16) + b = torch.randn((K, N), device='cpu', dtype=torch.float16) + + # torch_output = torch.matmul(a, b) + + triton_output = bench_matmul(a, b) + + # print(torch_output) + print(triton_output) + + # abs_diff = torch.abs(torch_output - triton_output) + # print(abs_diff) + + # max_diff = torch.max(abs_diff) + # max_diff_index = torch.argmax(abs_diff) + # print(max_diff) + # print(max_diff_index) + # print(torch_output.reshape(-1)[max_diff_index]) + # print(triton_output.reshape(-1)[max_diff_index]) + # print( + # f"The maximum difference between torch and triton is " + # f"{max_diff}" + # ) diff --git a/third_party/wafer/examples/time_zs_opt2.py b/third_party/wafer/examples/time_zs_opt2.py new file mode 100755 index 00000000..cc72eee4 --- /dev/null +++ b/third_party/wafer/examples/time_zs_opt2.py @@ -0,0 +1,146 @@ +import torch + +import triton +import triton.language as tl +# import benchmark + +DEVICE = triton.runtime.driver.active.get_active_torch_device() + + +# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes: +# - A list of `triton.Config` objects that define different configurations of +# meta-parameters (e.g., `BLOCK_SIZE_M`) and compilation options (e.g., `num_warps`) to try +# - An auto-tuning *key* whose change in values will trigger evaluation of all the +# provided configs +@triton.jit +def matmul_kernel( + # Pointers to matrices + a_ptr, b_ptr, c_ptr, + # Matrix dimensions + M, N, K, + # The stride variables represent how much to increase the ptr by when moving by 1 + # element in a particular dimension. E.g. `stride_am` is how much to increase `a_ptr` + # by to get the element one row down (A has M rows). + stride_am, stride_ak, # + stride_bk, stride_bn, # + stride_cm, stride_cn, + # Meta-parameters + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, # + GROUP_SIZE_M: tl.constexpr, # + ACTIVATION: tl.constexpr # +): + """Kernel for computing the matmul C = A x B. + A has shape (M, K), B has shape (K, N) and C has shape (M, N) + """ + # ----------------------------------------------------------- + # Map program ids `pid` to the block of C it should compute. + # This is done in a grouped ordering to promote L2 data reuse. + # See above `L2 Cache Optimizations` section for details. + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + # ---------------------------------------------------------- + # Create pointers for the first blocks of A and B. + # We will advance this pointer as we move in the K direction + # and accumulate + # `a_ptrs` is a block of [BLOCK_SIZE_M, BLOCK_SIZE_K] pointers + # `b_ptrs` is a block of [BLOCK_SIZE_K, BLOCK_SIZE_N] pointers + # See above `Pointer Arithmetics` section for details + offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + + # ----------------------------------------------------------- + # Iterate to compute a block of the C matrix. + # We accumulate into a `[BLOCK_SIZE_M, BLOCK_SIZE_N]` block + # of fp32 values for higher accuracy. + # `accumulator` will be converted back to fp16 after the loop. + acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=c_ptr.dtype.element_ty) + for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + # Load the next block of A and B, generate a mask by checking the K dimension. + # If it is out of bounds, set it to 0. + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + # We accumulate along the K dimension. + acc += tl.dot(a, b, out_dtype=c_ptr.dtype.element_ty) + # Advance the ptrs to the next K block. + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + # You can fuse arbitrary activation functions here + # while the accumulator is still in FP32! + if ACTIVATION == "leaky_relu": + acc = leaky_relu(accumulator) + c = acc + + # ----------------------------------------------------------- + # Write back the block of the output matrix C with masks. + offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + tl.store(c_ptrs, c) + + +# We can fuse `leaky_relu` by providing it as an `ACTIVATION` meta-parameter in `_matmul`. +@triton.jit +def leaky_relu(x): + x = x + 1 + return tl.where(x >= 0, x, 0.01 * x) + + +def matmul(a, b, activation=""): + # Check constraints. + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + assert a.is_contiguous(), "Matrix A must be contiguous" + assert b.is_contiguous(), "Matrix B must be contiguous" + M, K = a.shape + K, N = b.shape + + BLOCK_M = 1024 + BLOCK_N = 1024 + BLOCK_K = 128 + assert M % BLOCK_M == 0 and N % BLOCK_N == 0 and K % BLOCK_K == 0 + + # Allocates output. + c = torch.empty((M, N), device=a.device, dtype=a.dtype) + # 1D launch kernel where each block gets its own program. + grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), ) + matmul_kernel[grid]( + a, b, c, # + M, N, K, # + a.stride(0), a.stride(1), # + b.stride(0), b.stride(1), # + c.stride(0), c.stride(1), # + ACTIVATION=activation, # + BLOCK_SIZE_M=BLOCK_M, BLOCK_SIZE_N=BLOCK_N, BLOCK_SIZE_K=BLOCK_K, GROUP_SIZE_M=8) + return c + + +# @benchmark.measure(repeats=20) +def bench_matmul(a, b): + x = a.to(DEVICE) + y = b.to(DEVICE) + z = matmul(x, y) + z = z.to("cpu") + return z + + +if __name__ == "__main__": + M = 4096 + N = 4096 + K = 4096 + a = torch.randn((M, K), device='cpu', dtype=torch.float16) + b = torch.randn((K, N), device='cpu', dtype=torch.float16) + out = bench_matmul(a, b) + print(out) + ref = torch.matmul(a.to(torch.float32), b.to(torch.float32)) + print(ref) + print(ref - out) diff --git a/third_party/wafer/examples/tle/test_tle_dsa_noc_gemm_4096.py b/third_party/wafer/examples/tle/test_tle_dsa_noc_gemm_4096.py new file mode 100755 index 00000000..7a3ac011 --- /dev/null +++ b/third_party/wafer/examples/tle/test_tle_dsa_noc_gemm_4096.py @@ -0,0 +1,153 @@ +import pytest +import torch +import torch_txda # noqa: F401 +import triton +import triton.language as tl +import triton.experimental.tle.language as tle + +TILE_NUM = 16 +M = 4096 +K = 1024 +N = 4096 +BLOCK_M = M // TILE_NUM +BLOCK_K = K +SUB_N = N // TILE_NUM + +TILE_PHYSICAL_RELATION = [0, 1, 2, 3, 7, 11, 15, 14, 13, 12, 8, 9, 10, 6, 5, 4] + +MESH = tle.device_mesh( + None, + _shape=(TILE_NUM, ), + _dim_names=("tile", ), + _physical_ids=tuple(TILE_PHYSICAL_RELATION), +) + + +@triton.jit +def dsa_shift_n_gemm_kernel( + A_ptr, + B_ptr, + C_ptr, + send_next_tile_lut_ptr, + ring_index_lut_ptr, + M: tl.constexpr, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_K: tl.constexpr, + SUB_N: tl.constexpr, + TILE_NUM: tl.constexpr, +): + pid = tl.program_id(0) + send_next_tile = tl.load(send_next_tile_lut_ptr + pid) + ring_index = tl.load(ring_index_lut_ptr + pid) + + offs_m = pid * BLOCK_M + tl.arange(0, BLOCK_M) + offs_k = tl.arange(0, BLOCK_K) + + a_ptrs = A_ptr + offs_m[:, None] * K + offs_k[None, :] + a = tl.load(a_ptrs) + + shard_idx = ring_index + offs_sub_n = shard_idx * SUB_N + tl.arange(0, SUB_N) + b_ptrs = B_ptr + offs_k[:, None] * N + offs_sub_n[None, :] + b_init = tl.load(b_ptrs) + + send_buf = tle.dsa.alloc((BLOCK_K, SUB_N), tl.float16) + recv_buf = tle.dsa.alloc((BLOCK_K, SUB_N), tl.float16) + + offs_buf_k = tl.arange(0, BLOCK_K)[:, None] + tl.zeros((1, SUB_N), dtype=tl.int32) + offs_buf_n = tl.arange(0, SUB_N)[None, :] + tl.zeros((BLOCK_K, 1), dtype=tl.int32) + + send_ptr = tle.dsa.local_ptr(send_buf, [offs_buf_k, offs_buf_n]) + recv_ptr = tle.dsa.local_ptr(recv_buf, [offs_buf_k, offs_buf_n]) + + remote_recv_buf = tle.remote(recv_buf, send_next_tile) + remote_recv_ptr = tle.dsa.local_ptr(remote_recv_buf, [offs_buf_k, offs_buf_n]) + + tl.store(send_ptr, b_init) + + for step in range(TILE_NUM): + b_cur = tl.load(send_ptr) + c_part = tl.dot(a, b_cur, out_dtype=tl.float32) + + offs_n = shard_idx * SUB_N + tl.arange(0, SUB_N) + c_ptrs = C_ptr + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ptrs, c_part.to(tl.float16)) + + if step < TILE_NUM - 1: + tl.store(remote_recv_ptr, tl.load(send_ptr)) + # tle.distributed_barrier(MESH) + tl.store(send_ptr, tl.load(recv_ptr)) + # tle.distributed_barrier(MESH) + + shard_idx = tl.where(shard_idx == 0, TILE_NUM - 1, shard_idx - 1) + + +def build_ring_luts(mesh, device): + phys = mesh.physical_ids + n = mesh.size + send_next = torch.empty(n, dtype=torch.int32) + ring_index = torch.empty(n, dtype=torch.int32) + for i in range(n): + cur = phys[i] + nxt = phys[(i + 1) % n] + send_next[cur] = nxt + ring_index[cur] = i + return send_next.to(device), ring_index.to(device) + + +def run(m=M, n=N, k=K, device="cpu", pattern="random", seed=0): + """Execute with explicit torch_txda tensors; the CRT ring has 16 members.""" + if m <= 0 or n <= 0 or m % TILE_NUM or n % TILE_NUM or k <= 0: + raise ValueError("M and N must be divisible by the 16-tile ring size") + torch.manual_seed(seed) + if pattern == "structured": + # Exact FP16 values exercise tile/shard routing without reduction noise. + a = torch.zeros((m, k), device="cpu", dtype=torch.float16) + a[torch.arange(m, device="cpu"), torch.arange(m, device="cpu") % k] = 1 + rows = torch.arange(k, device="cpu")[:, None] + cols = torch.arange(n, device="cpu")[None, :] + b = (((rows * 7 + cols * 11 + seed * 13) % 1024).float() / 1024).half() + elif pattern == "random": + a = torch.randn((m, k), device="cpu", dtype=torch.float16) + b = torch.randn((k, n), device="cpu", dtype=torch.float16) + else: + raise ValueError(f"Unknown input pattern: {pattern}") + c = torch.full((m, n), float("nan"), device="cpu", dtype=torch.float16) + send_next_lut, ring_index_lut = build_ring_luts(MESH, device) + from triton.backends.dicp_triton.wafer_runtime import initialize_noc + initialize_noc() + a_txda = a.to("txda") + b_txda = b.to("txda") + c_txda = c.to("txda") + send_next_lut_txda = send_next_lut.to("txda") + ring_index_lut_txda = ring_index_lut.to("txda") + dsa_shift_n_gemm_kernel[(TILE_NUM,)]( + a_txda, b_txda, c_txda, send_next_lut_txda, ring_index_lut_txda, + M=m, N=n, K=k, BLOCK_M=m // TILE_NUM, BLOCK_K=k, + SUB_N=n // TILE_NUM, TILE_NUM=TILE_NUM, + launch_mode="cluster", + ) + with torch.no_grad(): + c.copy_(c_txda.cpu()) + ref = a.cpu().float() @ b.cpu().float() + result = c.cpu().float() + tolerance = 0.0 if pattern == "structured" else 1e-1 + torch.testing.assert_close(result, ref, atol=tolerance, rtol=tolerance) + max_diff = (result - ref).abs().max().item() + print(f"PASS NoC ring GEMM: M={m}, N={n}, K={k}, tiles={TILE_NUM}, " + f"pattern={pattern}, seed={seed}, max_abs_diff={max_diff:.8g}", flush=True) + + +@pytest.mark.parametrize("m,n,k", [(256, 256, 64), (4096, 4096, 1024)]) +@pytest.mark.parametrize("pattern", ["structured", "random"]) +def test_noc_gemm(m, n, k, pattern, device): + for seed in (0, 1): + run(m, n, k, device=device, pattern=pattern, seed=seed) + + +if __name__ == "__main__": + raise SystemExit("Run this example with scripts/wafer/run_wafer_example_suite.py " + "--suite examples --select tle/test_tle_dsa_noc_gemm_4096.py " + "--output-dir /tmp/wafer-noc-results") diff --git a/third_party/wafer/examples/util.py b/third_party/wafer/examples/util.py new file mode 100755 index 00000000..2ee319a8 --- /dev/null +++ b/third_party/wafer/examples/util.py @@ -0,0 +1,52 @@ +import torch + + +def gems_assert_cosine_similarity(a, b, dtype, eps=1e-8): + a_cpu = a.to("cpu") + b_cpu = b.to(dtype) + a_cpu = a_cpu.to(dtype=torch.float64) + b_cpu = b_cpu.to(dtype=torch.float64) + + a_cpu = a_cpu.flatten() + b_cpu = b_cpu.flatten() + dim = 0 + + is_nan_a = torch.isnan(a_cpu) + is_nan_b = torch.isnan(b_cpu) + both_nan = is_nan_a & is_nan_b + + is_inf_a = torch.isinf(a_cpu) + is_inf_b = torch.isinf(b_cpu) + both_inf = is_inf_a & is_inf_b & (torch.sign(a_cpu) == torch.sign(b_cpu)) + + invalid_mask = both_nan | both_inf + valid_mask = ~invalid_mask + a_filtered = a_cpu[valid_mask] + b_filtered = b_cpu[valid_mask] + + if len(a_filtered) == 0 and len(b_filtered) == 0: + assert True + elif len(a_filtered) == 0 and len(b_filtered) != 0: + assert False, "The output of inf and nan results is misaligned" + elif len(a_filtered) != 0 and len(b_filtered) == 0: + assert False, "The output of inf and nan results is misaligned" + else: + dot_product = (a_filtered * b_filtered).sum(dim=dim) + a_norm = a_filtered.norm(p=2, dim=dim) + b_norm = b_filtered.norm(p=2, dim=dim) + + cosine_sim = dot_product / (a_norm * b_norm + eps) + + print(f"cosine_sim is: {cosine_sim*100:.4f}%") + print(f"dot_product is: {dot_product}") + print(f"X_norm*Y_norm is: {a_norm * b_norm + eps}") + + if torch.isnan(cosine_sim): + print(f"cosine_sim is: {cosine_sim}") + assert torch.isnan(cosine_sim) + elif torch.isinf(cosine_sim): + print(f"cosine_sim is: {cosine_sim}") + assert torch.isinf(cosine_sim) + elif cosine_sim < 0.9: + print(f"cosine_sim is: {cosine_sim*100:.4f}%") + assert False, f"cosine_sim < 90% ({cosine_sim*100:.4f}%)" diff --git a/third_party/wafer/examples/view_vec_add_ir.sh b/third_party/wafer/examples/view_vec_add_ir.sh new file mode 100755 index 00000000..ed72855e --- /dev/null +++ b/third_party/wafer/examples/view_vec_add_ir.sh @@ -0,0 +1,41 @@ +#!/bin/bash + +set -euo pipefail + +DUMP_ROOT=${TRITON_DUMP_PATH:-/tmp/tsm_dump} +DUMP_INDEX=${DUMP_INDEX:-1} +DUMP_DIR="$DUMP_ROOT/dump$DUMP_INDEX" +LINES=${LINES:-120} + +if [ ! -d "$DUMP_DIR" ]; then + echo "ERROR: dump directory not found: $DUMP_DIR" >&2 + echo "Run ./dump_vec_add_ir.sh first, or set TRITON_DUMP_PATH/DUMP_INDEX." >&2 + exit 1 +fi + +show_file() { + local title="$1" + local path="$2" + if [ -f "$path" ]; then + echo "========================================" + echo " $title" + echo "========================================" + echo "FILE: $path" + echo "" + sed -n "1,${LINES}p" "$path" + echo "" + fi +} + +echo "Dump directory: $DUMP_DIR" +echo "" +ls -lah "$DUMP_DIR" +echo "" + +show_file "Commands" "$DUMP_DIR/cmds.txt" +show_file "TTIR" "$DUMP_DIR/tt_0.mlir" +show_file "CoreIR" "$DUMP_DIR/core_0.mlir" +show_file "Wafer IR" "$DUMP_DIR/wafer_0.mlir" +show_file "LLVM Dialect MLIR" "$DUMP_DIR/ll_0.mlir" +show_file "LLVM IR" "$DUMP_DIR/ll_0.ir" +show_file "Kernel LLVM IR" "$DUMP_DIR/kernel_0.ll" diff --git a/third_party/wafer/experimental/README.md b/third_party/wafer/experimental/README.md new file mode 100644 index 00000000..48ba10d6 --- /dev/null +++ b/third_party/wafer/experimental/README.md @@ -0,0 +1,13 @@ +# Wafer experimental Python extensions + +`tle/` is installed as `triton.experimental.tle` by `setup_on_wafer.py`. +It is owned by this repository; the pinned Triton submodule is not modified. + +Imported from the user-provided `tle-migration-kit-20260910.tar.gz`: + +- archive SHA256: `35dc5742b589b9eb84f480499c03b6f4c5f79dc13864846f5520996060a32a85` +- FlagTree source HEAD: `bf9c62c3865e9ef9e467bdc1e77be53f0ebadf47` +- the archive records working-tree provenance; it is not asserted to be a clean checkout +- all eight Python files were imported; subsequent adaptations are recorded in Git + +The optional `raw` extension was not supplied. diff --git a/third_party/wafer/experimental/tle/__init__.py b/third_party/wafer/experimental/tle/__init__.py new file mode 100644 index 00000000..8ab812ee --- /dev/null +++ b/third_party/wafer/experimental/tle/__init__.py @@ -0,0 +1,14 @@ +# flagtree tle +from . import language + +try: + from . import raw +except (ModuleNotFoundError, ImportError): + raw = None + +__all__ = [ + "language", +] + +if raw is not None: + __all__.append("raw") diff --git a/third_party/wafer/experimental/tle/language/__init__.py b/third_party/wafer/experimental/tle/language/__init__.py new file mode 100644 index 00000000..6e11fe24 --- /dev/null +++ b/third_party/wafer/experimental/tle/language/__init__.py @@ -0,0 +1,40 @@ +# flagtree tle +from .core import ( + load, extract_tile, insert_tile, cumsum) +from .distributed import ( + B, + P, + S, + ShardedTensor, + ShardingSpec, + device_mesh, + distributed_barrier, + distributed_dot, + make_sharded_tensor, + remote, + reshard, + shard_id, + sharding, +) + +__all__ = [ + "load", + "extract_tile", "insert_tile", "cumsum", + "device_mesh", + "S", + "P", + "B", + "sharding", + "ShardingSpec", + "ShardedTensor", + "make_sharded_tensor", + "reshard", + "remote", + "shard_id", + "distributed_barrier", + "distributed_dot", + "distributed", + "dsa", +] + +from . import distributed, dsa diff --git a/third_party/wafer/experimental/tle/language/core.py b/third_party/wafer/experimental/tle/language/core.py new file mode 100644 index 00000000..a20c358a --- /dev/null +++ b/third_party/wafer/experimental/tle/language/core.py @@ -0,0 +1,140 @@ +# flagtree tle +import triton.language.core as tl +import builtins +import math +import triton +import triton.language as language + + +def _tile_offsets(x, index, tile_shape, semantic): + """Normalize FlagTree's row-major tile indices to 3.5 DSA slice offsets.""" + shape = tuple(tl._unwrap_if_constexpr(dim) for dim in x.shape) + tile_shape = tuple(tl._unwrap_if_constexpr(dim) for dim in tl._unwrap_if_constexpr(tile_shape)) + if len(shape) != len(tile_shape) or any(type(t) is not int or t <= 0 for t in tile_shape): + raise ValueError("tile_shape must contain positive integers and match source rank") + if any(s % t for s, t in zip(shape, tile_shape)): + raise ValueError("source dimensions must be divisible by tile dimensions") + grid = tuple(s // t for s, t in zip(shape, tile_shape)) + index = tl._unwrap_if_constexpr(index) + if isinstance(index, (tuple, list, tl.tuple)): + coords = [tl._unwrap_if_constexpr(i) for i in index] + if len(coords) != len(grid): + raise ValueError("tile index rank must match source rank") + else: + if isinstance(index, tl.tensor): + if len(index.shape) or not index.dtype.is_int(): + raise ValueError("dynamic tile index must be a scalar integer") + elif type(index) is not int or not 0 <= index < math.prod(grid): + raise ValueError("linear tile index out of range") + coords = [None] * len(grid) + for axis in builtins.range(len(grid) - 1, -1, -1): + if isinstance(index, tl.tensor): + coords[axis] = index.__mod__(grid[axis], _semantic=semantic) + index = index.__floordiv__(grid[axis], _semantic=semantic) + else: + coords[axis], index = index % grid[axis], index // grid[axis] + offsets = [] + for coord, extent, tile in zip(coords, grid, tile_shape): + if isinstance(coord, tl.tensor): + if len(coord.shape) or not coord.dtype.is_int(): + raise ValueError("dynamic tile indices must be scalar integers") + offsets.append(coord.__mul__(tile, _semantic=semantic)) + else: + if type(coord) is not int or not 0 <= coord < extent: + raise ValueError("tile index out of range") + offsets.append(coord * tile) + return offsets, tile_shape + + +@tl.builtin +def extract_tile(x, index, tile_shape, _semantic=None): + """Extract a tile using a scalar or per-axis row-major tile index. + + Dynamic indices must be in bounds. The official 3.5 integration lowers + directly through DSA slices rather than requiring FlagTree's shared IR. + """ + from .dsa import extract_slice + offsets, shape = _tile_offsets(x, index, tile_shape, _semantic) + return extract_slice(x, offsets, shape, (1,) * len(shape), _semantic=_semantic) + + +@tl.builtin +def insert_tile(x, tile, index, _semantic=None): + """Insert a tile using a scalar or per-axis row-major tile index.""" + from .dsa import insert_slice + offsets, shape = _tile_offsets(x, index, tile.shape, _semantic) + return insert_slice(x, tile, offsets, shape, (1,) * len(shape), _semantic=_semantic) + + +@triton.jit +def cumsum(x, axis: tl.constexpr = 0, reverse: tl.constexpr = False): + """Exclusive rank-one float scan and total, using the existing Wafer lowering. + + Shift the input before scanning to avoid cancellation in inclusive_sum - x. + The JIT wrapper lets Triton 3.5 inline its standard scan/reduce functions. + """ + tl.static_assert(len(x.shape) == 1 and (axis == 0 or axis == -1) and not reverse, + "Wafer TLE cumsum supports rank-one forward scans") + tl.static_assert(x.dtype.is_floating(), "Wafer TLE cumsum requires floating point") + indices = tl.arange(0, x.shape[0]) + previous = tl.maximum(indices - 1, 0) + shifted = tl.gather(x, previous, axis=0) + shifted = tl.where(indices > 0, shifted, 0) + return language.cumsum(shifted, 0), language.sum(x, 0) + + +# ----------------------- +# Non-Atomic Memory Operations +# ----------------------- + + +@tl.builtin +def load(pointer, mask=None, other=None, boundary_check=(), padding_option="", cache_modifier="", eviction_policy="", + volatile=False, is_async=False, _semantic=None): + """ + Return a tensor of data whose values are loaded from memory at location defined by `pointer`: + + (1) If `pointer` is a single element pointer, a scalar is be loaded. In + this case: + + - `mask` and `other` must also be scalars, + - `other` is implicitly typecast to `pointer.dtype.element_ty`, and + - `boundary_check` and `padding_option` must be empty. + + (2) If `pointer` is an N-dimensional tensor of pointers, an + N-dimensional tensor is loaded. In this case: + + - `mask` and `other` are implicitly broadcast to `pointer.shape`, + - `other` is implicitly typecast to `pointer.dtype.element_ty`, and + - `boundary_check` and `padding_option` must be empty. + + (3) If `pointer` is a block pointer defined by `make_block_ptr`, a + tensor is loaded. In this case: + + - `mask` and `other` must be `None`, and + - `boundary_check` and `padding_option` can be specified to control the behavior of out-of-bound access. + + :param pointer: Pointer to the data to be loaded + :type pointer: `triton.PointerType`, or block of `dtype=triton.PointerType` + :param mask: if `mask[idx]` is false, do not load the data at address `pointer[idx]` + (must be `None` with block pointers) + :type mask: Block of `triton.int1`, optional + :param other: if `mask[idx]` is false, return `other[idx]` + :type other: Block, optional + :param boundary_check: tuple of integers, indicating the dimensions which should do the boundary check + :type boundary_check: tuple of ints, optional + :param padding_option: should be one of {"", "zero", "nan"}, the padding value to use while out of bounds. "" means an undefined value. + :param cache_modifier: changes cache option in NVIDIA PTX + :type cache_modifier: str, optional, should be one of {"", ".ca", ".cg", ".cv"}, where ".ca" stands for + cache at all levels, ".cg" stands for cache at global level (cache in L2 and below, not L1), + and ".cv" means don’t cache and fetch again. see + `cache operator `_ for more details. + :param eviction_policy: changes eviction policy in NVIDIA PTX + :type eviction_policy: str, optional + :param volatile: changes volatile option in NVIDIA PTX + :type volatile: bool, optional + """ + x = tl.load(pointer, mask=mask, other=other, boundary_check=boundary_check, padding_option=padding_option, + cache_modifier=cache_modifier, eviction_policy=eviction_policy, volatile=volatile, _semantic=_semantic) + x.handle.set_attr("tt.load.async", _semantic.builder.get_bool_attr(is_async)) + return x diff --git a/third_party/wafer/experimental/tle/language/distributed.py b/third_party/wafer/experimental/tle/language/distributed.py new file mode 100644 index 00000000..3ea0bfc0 --- /dev/null +++ b/third_party/wafer/experimental/tle/language/distributed.py @@ -0,0 +1,669 @@ +# flagtree tle +from __future__ import annotations + +import copy +from dataclasses import dataclass +from itertools import product +from typing import Any, Iterable, Mapping, Sequence + +import triton.language.core as tl + + +def _prod(values: Iterable[int]) -> int: + result = 1 + for value in values: + result *= value + return result + + +def _as_positive_int(value: Any, label: str) -> int: + if not isinstance(value, int): + raise TypeError(f"{label} must be int, got {type(value).__name__}") + if value <= 0: + raise ValueError(f"{label} must be > 0, got {value}") + return value + + +class device_mesh: + """ + Logical view of a physical device topology. + """ + + def __init__( + self, + topology: Mapping[str, Any] | None = None, + *, + _shape: Sequence[int] | None = None, + _dim_names: Sequence[str] | None = None, + _physical_ids: Sequence[int] | None = None, + _launch_shape: Sequence[int] | None = None, + _launch_dim_names: Sequence[str] | None = None, + ): + if topology is None: + if _shape is None or _dim_names is None or _physical_ids is None: + raise ValueError("internal mesh constructor requires shape/names/physical ids") + self._shape = tuple(_shape) + self._dim_names = tuple(_dim_names) + self._physical_ids = tuple(_physical_ids) + self._launch_shape = tuple(_launch_shape if _launch_shape is not None else _shape) + self._launch_dim_names = tuple(_launch_dim_names if _launch_dim_names is not None else _dim_names) + return + + if not isinstance(topology, Mapping): + raise TypeError(f"topology must be a mapping, got {type(topology).__name__}") + if not topology: + raise ValueError("topology cannot be empty") + + shape = [] + dim_names = [] + for level_name, level_desc in topology.items(): + if not isinstance(level_name, str) or not level_name: + raise ValueError(f"invalid topology level name: {level_name!r}") + level_shape, level_names = self._parse_level(level_name, level_desc) + shape.extend(level_shape) + dim_names.extend(level_names) + + if len(set(dim_names)) != len(dim_names): + raise ValueError(f"dimension names must be unique, got {dim_names}") + + self._shape = tuple(shape) + self._dim_names = tuple(dim_names) + self._physical_ids = tuple(range(_prod(shape))) + self._launch_shape = self._shape + self._launch_dim_names = self._dim_names + + @staticmethod + def _parse_level(level_name: str, level_desc: Any) -> tuple[list[int], list[str]]: + if isinstance(level_desc, int): + return [_as_positive_int(level_desc, level_name)], [level_name] + if not isinstance(level_desc, (tuple, list)): + raise TypeError(f"topology[{level_name!r}] must be int or list/tuple of (name, size), " + f"got {type(level_desc).__name__}") + if not level_desc: + raise ValueError(f"topology[{level_name!r}] cannot be empty") + + shape = [] + names = [] + for item in level_desc: + if not isinstance(item, (tuple, list)) or len(item) != 2: + raise ValueError(f"topology[{level_name!r}] entries must be (name, size), got {item!r}") + dim_name, dim_size = item + if not isinstance(dim_name, str) or not dim_name: + raise ValueError(f"invalid dimension name in {level_name!r}: {dim_name!r}") + shape.append(_as_positive_int(dim_size, f"{level_name}.{dim_name}")) + names.append(dim_name) + return shape, names + + @property + def shape(self) -> tuple[int, ...]: + return self._shape + + @property + def ndim(self) -> int: + return len(self._shape) + + @property + def dim_names(self) -> tuple[str, ...]: + return self._dim_names + + @property + def physical_ids(self) -> tuple[int, ...]: + return self._physical_ids + + @property + def launch_shape(self) -> tuple[int, ...]: + return self._launch_shape + + @property + def launch_dim_names(self) -> tuple[str, ...]: + return self._launch_dim_names + + @property + def size(self) -> int: + return len(self._physical_ids) + + def flatten(self) -> "device_mesh": + return self.reshape(self.size) + + def reshape(self, *shape: int | Sequence[int]) -> "device_mesh": + if len(shape) == 1 and isinstance(shape[0], (tuple, list)): + new_shape = tuple(shape[0]) + else: + new_shape = tuple(shape) + if not new_shape: + raise ValueError("new shape cannot be empty") + new_shape = tuple(_as_positive_int(v, "shape dimension") for v in new_shape) + if _prod(new_shape) != self.size: + raise ValueError(f"cannot reshape mesh of size {self.size} into shape {new_shape}") + if len(new_shape) == self.ndim: + new_dim_names = self._dim_names + elif len(new_shape) == 1: + new_dim_names = ("flat", ) + else: + new_dim_names = tuple(f"dim{i}" for i in range(len(new_shape))) + return device_mesh( + None, + _shape=new_shape, + _dim_names=new_dim_names, + _physical_ids=self._physical_ids, + _launch_shape=self._launch_shape, + _launch_dim_names=self._launch_dim_names, + ) + + def _normalize_key(self, key: Any) -> tuple[Any, ...]: + if not isinstance(key, tuple): + key = (key, ) + + if any(item is Ellipsis for item in key): + if key.count(Ellipsis) > 1: + raise IndexError("an index can only have a single ellipsis") + ellipsis_pos = key.index(Ellipsis) + missing = self.ndim - (len(key) - 1) + if missing < 0: + raise IndexError("too many indices for device_mesh") + key = key[:ellipsis_pos] + (slice(None), ) * missing + key[ellipsis_pos + 1:] + + if len(key) > self.ndim: + raise IndexError("too many indices for device_mesh") + + return key + (slice(None), ) * (self.ndim - len(key)) + + def _linear_index(self, coords: Sequence[int]) -> int: + index = 0 + for coord, dim_size in zip(coords, self._shape): + index = index * dim_size + coord + return index + + def __getitem__(self, key: Any) -> "device_mesh": + key = self._normalize_key(key) + selected_per_dim: list[list[int]] = [] + keep_dim: list[bool] = [] + + for dim_size, dim_key in zip(self._shape, key): + if isinstance(dim_key, int): + idx = dim_key + dim_size if dim_key < 0 else dim_key + if idx < 0 or idx >= dim_size: + raise IndexError(f"index {dim_key} out of range for dim size {dim_size}") + selected_per_dim.append([idx]) + keep_dim.append(False) + elif isinstance(dim_key, slice): + indices = list(range(*dim_key.indices(dim_size))) + if not indices: + raise ValueError("empty sub-mesh is not supported") + selected_per_dim.append(indices) + keep_dim.append(True) + else: + raise TypeError(f"device_mesh indices must be int/slice/ellipsis, got {type(dim_key).__name__}") + + new_shape = tuple(len(indices) for indices, keep in zip(selected_per_dim, keep_dim) if keep) + new_dim_names = tuple(dim_name for dim_name, keep in zip(self._dim_names, keep_dim) if keep) + + new_physical_ids = [] + for coords in product(*selected_per_dim): + new_physical_ids.append(self._physical_ids[self._linear_index(coords)]) + + return device_mesh( + None, + _shape=new_shape, + _dim_names=new_dim_names, + _physical_ids=tuple(new_physical_ids), + _launch_shape=self._launch_shape, + _launch_dim_names=self._launch_dim_names, + ) + + def __repr__(self): + return f"DeviceMesh(shape={self._shape}, names={self._dim_names})" + + +class _BroadcastSpec: + + def __repr__(self) -> str: + return "B" + + +B = _BroadcastSpec() + + +@dataclass(frozen=True) +class S: + axis: str | Sequence[str] + + +@dataclass(frozen=True) +class P: + axis: str | Sequence[str] + + +def _normalize_axis_group(spec: Any, label: str) -> tuple[str, ...]: + if spec is None or spec is B: + return tuple() + + if isinstance(spec, S): + spec = spec.axis + if isinstance(spec, P): + spec = spec.axis + + if isinstance(spec, str): + if not spec: + raise ValueError(f"{label} axis name cannot be empty") + return (spec, ) + + if isinstance(spec, (tuple, list)): + if not spec: + return tuple() + axes = [] + for axis in spec: + if not isinstance(axis, str) or not axis: + raise ValueError(f"{label} axis name must be non-empty str, got {axis!r}") + axes.append(axis) + if len(set(axes)) != len(axes): + raise ValueError(f"{label} axis names must be unique, got {axes}") + return tuple(axes) + + raise TypeError(f"{label} axis spec must be str/list/tuple/S/P/B, got {type(spec).__name__}") + + +def _normalize_partial_specs(partial: Any) -> tuple[str, ...]: + if partial is None: + return tuple() + if isinstance(partial, (str, S, P)): + partial = [partial] + if not isinstance(partial, (tuple, list)): + raise TypeError(f"partial must be a list/tuple, got {type(partial).__name__}") + + axes = [] + for item in partial: + axes.extend(_normalize_axis_group(item, "partial")) + if len(set(axes)) != len(axes): + raise ValueError(f"partial axes must be unique, got {axes}") + return tuple(axes) + + +@dataclass(frozen=True) +class ShardingSpec: + mesh: device_mesh + split: tuple[tuple[str, ...], ...] + partial: tuple[str, ...] + broadcast: tuple[str, ...] + + def axis_state(self, axis: str) -> str: + if axis in self.partial: + return "P" + for split_axes in self.split: + if axis in split_axes: + return "S" + return "B" + + +@dataclass(frozen=True) +class ShardedTensor: + handle: Any + sharding: ShardingSpec + shape: tuple[int, ...] | None = None + + +def sharding( + mesh: device_mesh, + split: Sequence[Any] | None = None, + partial: Sequence[Any] | None = None, +) -> ShardingSpec: + """ + Construct a sharding spec bound to a device mesh. + + This is annotation metadata today. Communication lowering is added in later + phases. + """ + if not isinstance(mesh, device_mesh): + raise TypeError(f"mesh must be device_mesh, got {type(mesh).__name__}") + + split_specs: list[tuple[str, ...]] = [] + if split is None: + split = tuple() + if not isinstance(split, (tuple, list)): + raise TypeError(f"split must be a list/tuple, got {type(split).__name__}") + for split_item in split: + split_specs.append(_normalize_axis_group(split_item, "split")) + + partial_axes = _normalize_partial_specs(partial) + + split_axes = [axis for split_item in split_specs for axis in split_item] + if len(set(split_axes)) != len(split_axes): + raise ValueError(f"split axes must be unique across tensor dims, got {split_axes}") + + split_set = set(split_axes) + partial_set = set(partial_axes) + + unknown = [axis for axis in split_axes + list(partial_axes) if axis not in mesh.dim_names] + if unknown: + raise ValueError(f"unknown mesh axis names: {unknown}; mesh axes are {mesh.dim_names}") + + overlap = split_set.intersection(partial_set) + if overlap: + raise ValueError(f"mesh axis cannot be both split and partial: {sorted(overlap)}") + + broadcast = tuple(axis for axis in mesh.dim_names if axis not in split_set and axis not in partial_set) + return ShardingSpec( + mesh=mesh, + split=tuple(split_specs), + partial=tuple(axis for axis in mesh.dim_names if axis in partial_set), + broadcast=broadcast, + ) + + +def make_sharded_tensor( + handle: Any, + sharding: ShardingSpec, + shape: Sequence[int] | None = None, +) -> ShardedTensor: + if not isinstance(sharding, ShardingSpec): + raise TypeError(f"sharding must be ShardingSpec, got {type(sharding).__name__}") + normalized_shape = None + if shape is not None: + if not isinstance(shape, (tuple, list)): + raise TypeError(f"shape must be list/tuple, got {type(shape).__name__}") + normalized_shape = tuple(_as_positive_int(v, "tensor shape") for v in shape) + if sharding.split and len(sharding.split) != len(normalized_shape): + raise ValueError(f"split rank ({len(sharding.split)}) must match tensor rank ({len(normalized_shape)})") + return ShardedTensor(handle=handle, sharding=sharding, shape=normalized_shape) + + +def reshard(tensor: ShardedTensor, spec: ShardingSpec) -> ShardedTensor: + """ + M4 entrypoint. Deferred by roadmap priority. + """ + raise NotImplementedError("reshard is deferred to M4") + + +def _shape_to_cluster_dims(shape: Sequence[int]) -> tuple[int, int, int]: + if not shape: + return (1, 1, 1) + dims = tuple(int(v) for v in shape) + if len(dims) == 1: + return (dims[0], 1, 1) + if len(dims) == 2: + return (dims[0], dims[1], 1) + if len(dims) == 3: + return dims + return (_prod(dims), 1, 1) + + +def _mesh_to_cluster_dims(mesh: device_mesh) -> tuple[int, int, int]: + # Prefer explicit cluster axes, then block axes, then fallback to full mesh. + cluster_axes = [size for name, size in zip(mesh.launch_dim_names, mesh.launch_shape) if "cluster" in name] + if not cluster_axes: + cluster_axes = [size for name, size in zip(mesh.launch_dim_names, mesh.launch_shape) if "block" in name] + if not cluster_axes: + cluster_axes = list(mesh.launch_shape) + return _shape_to_cluster_dims(cluster_axes) + + +@dataclass(frozen=True) +class _BarrierGroupDescriptor: + kind: str + rank: int + shape: tuple[int, ...] + axes: tuple[int, ...] + mask: tuple[int, ...] + + +def _infer_submesh_barrier_group( + mesh: device_mesh, + cluster_dims: Sequence[int], +) -> _BarrierGroupDescriptor | None: + cluster_size = _prod(cluster_dims) + if mesh.size == cluster_size: + return None + if mesh.size > cluster_size: + raise ValueError(f"mesh size ({mesh.size}) exceeds inferred cluster size ({cluster_size})") + + launch_size = _prod(mesh.launch_shape) + if launch_size != cluster_size: + raise NotImplementedError( + "sub-mesh distributed_barrier currently requires launch mesh domain " + f"to match inferred cluster size; launch_size={launch_size}, cluster_size={cluster_size}") + + if not mesh.dim_names: + raise NotImplementedError("scalar sub-mesh barrier is not implemented yet; provide at least one sliced axis") + + launch_name_to_axis = {name: i for i, name in enumerate(mesh.launch_dim_names)} + if any(name not in launch_name_to_axis for name in mesh.dim_names): + raise NotImplementedError("sub-mesh barrier currently supports slicing-derived meshes with " + "axis names inherited from launch mesh") + + axes = tuple(int(launch_name_to_axis[name]) for name in mesh.dim_names) + if len(set(axes)) != len(axes): + raise ValueError(f"invalid subgroup axes (duplicate launch axes): {axes}") + + shape = tuple(int(v) for v in mesh.shape) + if not shape or any(v <= 0 for v in shape): + raise ValueError(f"invalid subgroup shape inferred from mesh: {shape}") + + mask = tuple(int(v) for v in mesh.physical_ids) + if not mask: + raise ValueError("sub-mesh barrier group mask cannot be empty") + if any(v < 0 or v >= cluster_size for v in mask): + raise ValueError("sub-mesh barrier group mask contains out-of-range cluster member ids: " + f"mask={mask}, cluster_size={cluster_size}") + + return _BarrierGroupDescriptor( + kind="submesh", + rank=len(shape), + shape=shape, + axes=axes, + mask=mask, + ) + + +def _apply_mesh_cluster_launch(mesh: device_mesh, _semantic) -> tuple[int, int, int]: + cluster_dims = _mesh_to_cluster_dims(mesh) + options = getattr(_semantic.builder, "options", None) + if options is None: + return cluster_dims + + # The num_ctas == 1 constraint is NVIDIA CTA-cluster specific. + # On backends like TsingMicro the mesh describes tile communication + # topology, not CUDA CTA clusters, so only enforce when num_ctas > 1 + # (i.e. the backend actively uses multi-CTA grouping). + num_ctas = int(getattr(options, "num_ctas", 1)) + if num_ctas > 1: + raise ValueError("mesh-driven cluster launch requires num_ctas=1; cluster size is inferred from mesh") + + if hasattr(options, "cluster_dims"): + existing = tuple(getattr(options, "cluster_dims", (1, 1, 1))) + if existing != (1, 1, 1) and existing != cluster_dims: + raise ValueError(f"conflicting cluster_dims: existing={existing}, inferred_from_mesh={cluster_dims}") + object.__setattr__(options, "cluster_dims", cluster_dims) + return cluster_dims + + +def _resolve_launch_axis(mesh: device_mesh, axis: str | int) -> int: + if isinstance(axis, int): + ndim = len(mesh.launch_shape) + axis_idx = axis + ndim if axis < 0 else axis + if axis_idx < 0 or axis_idx >= ndim: + raise IndexError(f"axis index {axis} out of range for launch ndim {ndim}") + return axis_idx + + if isinstance(axis, str): + if axis not in mesh.launch_dim_names: + raise ValueError(f"unknown mesh axis {axis!r}; available launch axes: {mesh.launch_dim_names}") + return mesh.launch_dim_names.index(axis) + + raise TypeError(f"axis must be int or str, got {type(axis).__name__}") + + +@tl.builtin +def shard_id( + mesh: device_mesh, + axis: str | int, + _semantic=None, +): + """ + Return current shard coordinate on the given launch mesh axis. + + `axis` can be axis name (`str`) or axis index (`int`, supports negative). + The returned value is a scalar int32 tensor. + """ + mesh = tl._unwrap_if_constexpr(mesh) + axis = tl._unwrap_if_constexpr(axis) + + if not isinstance(mesh, device_mesh): + raise TypeError(f"mesh must be device_mesh, got {type(mesh).__name__}") + axis_idx = _resolve_launch_axis(mesh, axis) + launch_shape = tuple(int(v) for v in mesh.launch_shape) + launch_size = _prod(launch_shape) + if launch_size <= 0: + raise ValueError(f"invalid launch mesh shape: {launch_shape}") + + _apply_mesh_cluster_launch(mesh, _semantic) + linear = tl.program_id(0, _semantic=_semantic) + if launch_size > 1: + launch_size_t = _semantic.to_tensor(launch_size) + linear = _semantic.mod(linear, launch_size_t) + + stride = _prod(launch_shape[axis_idx + 1:]) if axis_idx + 1 < len(launch_shape) else 1 + coord = linear + if stride > 1: + stride_t = _semantic.to_tensor(stride) + coord = _semantic.floordiv(coord, stride_t) + dim = launch_shape[axis_idx] + if dim > 1: + dim_t = _semantic.to_tensor(dim) + coord = _semantic.mod(coord, dim_t) + return coord + + +@tl.builtin +def distributed_barrier(mesh: device_mesh | None = None, _semantic=None): + """A cross-tile barrier is not implemented by the current Wafer CRT. + + TsmWaitfinish only drains the local stream. The supported 16-tile ring + exchange synchronizes inside __Send; it does not rely on this operation. + """ + raise NotImplementedError("Wafer distributed_barrier needs a cross-tile CRT implementation") + + +def _normalize_remote_shard_id( + shard_id: Any, + scope: device_mesh | None, +) -> int: + shard_id = tl._unwrap_if_constexpr(shard_id) + scope = tl._unwrap_if_constexpr(scope) + + if isinstance(shard_id, int): + if shard_id < 0: + raise ValueError(f"shard_id must be >= 0, got {shard_id}") + return shard_id + + if not isinstance(shard_id, (tuple, list)): + raise TypeError(f"shard_id must be int or tuple/list of ints, got {type(shard_id).__name__}") + if not shard_id: + raise ValueError("shard_id tuple cannot be empty") + if not all(isinstance(v, int) for v in shard_id): + raise TypeError(f"shard_id tuple must contain ints, got {shard_id!r}") + + if scope is None: + raise ValueError("tuple shard_id requires scope=device_mesh to linearize coordinates") + if not isinstance(scope, device_mesh): + raise TypeError(f"scope must be device_mesh when shard_id is tuple, got {type(scope).__name__}") + if len(shard_id) != scope.ndim: + raise ValueError(f"tuple shard_id rank mismatch: got {len(shard_id)}, expected {scope.ndim}") + + linear = 0 + for idx, dim in zip(shard_id, scope.shape): + if idx < 0 or idx >= dim: + raise ValueError(f"shard_id coordinate {idx} out of range for dim size {dim}") + linear = linear * dim + idx + return linear + + +def _is_buffered_tensor_like(value: Any) -> bool: + return (not isinstance(value, tl.tensor) and value.__class__.__name__ == "buffered_tensor" + and hasattr(value, "handle") and hasattr(value, "type")) + + +def _normalize_compile_time_remote_shard_id( + shard_id: int | tuple[int, ...] | list[int], + scope: device_mesh | None, +) -> int: + linear_shard_id = _normalize_remote_shard_id(shard_id, scope) + if linear_shard_id > 0x7FFFFFFF: + raise ValueError(f"linearized shard_id {linear_shard_id} exceeds int32 range") + return linear_shard_id + + +def _normalize_runtime_remote_shard_id_tensor(shard_id_tensor: tl.tensor) -> tl.tensor: + if not shard_id_tensor.dtype.is_int() or shard_id_tensor.dtype.primitive_bitwidth != 32: + raise TypeError("runtime shard_id must be a scalar int32 tensor/value") + if shard_id_tensor.shape: + raise ValueError("runtime shard_id must be scalar (shape=())") + return shard_id_tensor + + +@tl.builtin +def remote( + tensor, + shard_id, + scope: device_mesh | None = None, + _semantic=None, +): + """ + M3 entrypoint: mark distributed access target. + + Supported input: + - tle buffered_tensor: returns a remote-marked buffered tensor; caller + should then use `tle.dsa.local_ptr(...)` to materialize remote pointers. + + `shard_id` is the target block id inside the current thread block cluster. + When `scope` is provided, launch cluster dimensions are inferred from that + mesh and this mode requires `num_ctas=1` (one program maps to one block). + """ + shard_id = tl._unwrap_if_constexpr(shard_id) + scope = tl._unwrap_if_constexpr(scope) + if scope is not None and not isinstance(scope, device_mesh): + raise TypeError(f"scope must be device_mesh or None, got {type(scope).__name__}") + if scope is not None: + _apply_mesh_cluster_launch(scope, _semantic) + + # Buffered tensor path: carry remote metadata and let `local_ptr` materialize + # remote pointers later. + if _is_buffered_tensor_like(tensor): + if (hasattr(tensor, "_tle_remote_shard_id") or hasattr(tensor, "_tle_remote_scope") + or hasattr(tensor.type, "_tle_remote_shard_id") or hasattr(tensor.type, "_tle_remote_scope")): + raise ValueError("remote(buffered_tensor, ...) cannot be applied twice; " + "materialize pointer views with tle.dsa.local_ptr(remote_buffer, indices)") + if isinstance(shard_id, (int, tuple, list)): + shard_id = _normalize_compile_time_remote_shard_id(shard_id, scope) + else: + shard_id_tensor = shard_id if isinstance(shard_id, tl.tensor) else None + if shard_id_tensor is None: + if _semantic is None: + raise TypeError("runtime shard_id for remote(buffered_tensor, ...) must be scalar int32 " + "and requires JIT context for materialization") + shard_id_tensor = _semantic.to_tensor(shard_id) + shard_id = _normalize_runtime_remote_shard_id_tensor(shard_id_tensor) + # Keep remote metadata on buffered_tensor.type so it survives value + # reconstruction in JIT interpreter paths (value-level attrs can drop). + remote_buffer = copy.copy(tensor) + remote_type = copy.copy(tensor.type) + try: + setattr(remote_type, "_tle_remote_shard_id", shard_id) + setattr(remote_type, "_tle_remote_scope", scope) + remote_buffer.type = remote_type + except AttributeError: + # Type object may be immutable for unit-test stubs. + pass + # Keep value-level metadata as a secondary carrier to maximize + # compatibility with existing JIT object reconstruction paths. + setattr(remote_buffer, "_tle_remote_shard_id", shard_id) + setattr(remote_buffer, "_tle_remote_scope", scope) + return remote_buffer + + if isinstance(tensor, tl.tensor): + raise TypeError("remote(...) only accepts tle.buffered_tensor; " + "use remote(buffered_tensor, shard_id, scope) + local_ptr(...)") + raise TypeError(f"tensor must be tle.buffered_tensor, got {type(tensor).__name__}") + + +def distributed_dot(a: ShardedTensor, b: ShardedTensor, c: ShardedTensor | None = None): + raise NotImplementedError("distributed_dot is deferred to M5") diff --git a/third_party/wafer/experimental/tle/language/dsa/__init__.py b/third_party/wafer/experimental/tle/language/dsa/__init__.py new file mode 100644 index 00000000..7a8e8681 --- /dev/null +++ b/third_party/wafer/experimental/tle/language/dsa/__init__.py @@ -0,0 +1,36 @@ +# flagtree tle +from .core import ( + pipeline, + alloc, + copy, + memory_space, + local_ptr, + to_tensor, to_buffer, add, sub, mul, max, min, div, extract_slice, insert_slice, +) +from .types import ( + scope, + local, + spm, + buffered_tensor, + buffered_tensor_type, +) +from .semantic import DSASemantic, DSASemanticError + +__all__ = [ + "pipeline", + "alloc", + "copy", + "memory_space", + "local_ptr", + "to_tensor", "to_buffer", "add", "sub", "mul", "max", "min", "div", "extract_slice", "insert_slice", + "scope", + "local", + "spm", + "buffered_tensor", + "buffered_tensor_type", + "DSASemantic", + "DSASemanticError", +] + +from . import wafer +__all__.append("wafer") diff --git a/third_party/wafer/experimental/tle/language/dsa/core.py b/third_party/wafer/experimental/tle/language/dsa/core.py new file mode 100644 index 00000000..be7ea66d --- /dev/null +++ b/third_party/wafer/experimental/tle/language/dsa/core.py @@ -0,0 +1,586 @@ +# flagtree tle +import builtins +import triton.language.core as tl +from triton._C.libtriton.wafer import tle as _dsa_ir +from typing import Optional, Sequence +from enum import Enum +from . import types as tle + +from triton.language.core import ( + constexpr, + tensor, + range, +) + +# Address space 3 matches the shared-memory space used in TritonGPU lowering. +SHARED_MEMORY_ADDRESS_SPACE = 3 + + +# Triton 3.5 recognizes loop iterators by class identity, so the otherwise empty +# FlagTree range subclass cannot compile here. The alias preserves num_stages +# and loop_unroll_factor and emits the same tt.num_stages loop annotation. +# Actual software pipelining is selected by the Wafer enable_pipeline option. +pipeline = range + + +@tl.builtin +def memory_space(input, space, _semantic=None): + """ + Annotate a tensor with a target memory-space tag. + + The attribute ``tt.memory_space`` is propagated through the IR and can be + consumed by downstream DSA passes (e.g. ``--dsa-memory-to-core``) to make + allocation / placement decisions. + + Args: + input: Tensor to annotate. + space: Memory-space name string, e.g. ``"spm"`` or ``"shared_memory"``. + """ + space = tl._unwrap_if_constexpr(space) + if _semantic is not None and hasattr(input, 'handle') and hasattr(input.handle, 'set_attr'): + input.handle.set_attr("tt.memory_space", _semantic.builder.get_string_attr(str(space))) + return input + + +@tl.builtin +def alloc( + shape: tuple, + dtype: tl.dtype, + layout: Optional[object] = None, + scope: tle.scope = None, + _semantic=None, +) -> tle.buffered_tensor: + """ + Allocate local memory buffer + + Args: + shape: Buffer shape + dtype: Data type + layout: Memory layout encoding (optional) + scope: Storage type (default to shared memory) + _semantic: Semantic analyzer (internal use) + + Returns: + Allocated buffer tensor + + Raises: + ValueError: When parameters are invalid + RuntimeError: When allocation fails + """ + from .semantic import DSASemantic + + if _semantic is None: + raise ValueError("alloc must be used inside @triton.jit") + if layout is not None: + raise ValueError("alloc(): layout parameter is not yet support for DSA backend") + + # --- Validate inputs via semantic layer --- + unwrapped_shape = DSASemantic.validate_alloc_shape(shape) + elem_dtype = DSASemantic.validate_alloc_dtype(dtype) + resolved_scope = DSASemantic.validate_alloc_scope(scope) + + elem_ir_ty = elem_dtype.to_ir(_semantic.builder) + + if not hasattr(_dsa_ir, "create_dsa_alloc"): + raise RuntimeError("builder missing create_dsa_alloc for DSA alloc") + + alloc_value = _dsa_ir.create_dsa_alloc(_semantic.builder, list(unwrapped_shape), elem_ir_ty) + buf_ty = tle.buffered_tensor_type(unwrapped_shape, elem_dtype, resolved_scope) + buf_ty._ir_type = alloc_value.get_type() + return tle.buffered_tensor(alloc_value, buf_ty) + + +class CopyDirection(Enum): + """Copy direction enum for data transfer operations""" + GM_TO_LOCAL = "GMTOLOCAL" # Global memory to local memory + LOCAL_TO_GM = "LOCALTOGM" # Local memory to global memory + + +@tl.builtin +def copy(src, dst, shape, offsets: Sequence[constexpr | tensor] = None, + _semantic=None) -> None: + """Copy an entire contiguous buffer, or an equally shaped tensor of pointers. + + DSA CopyOp only accepts memrefs. GM transfers are expressed as explicit + loads/stores through a local pointer view so their IR types remain valid. + """ + from .semantic import DSASemantic + + # Copy the complete local buffer; sub-buffer offsets are not implemented. + # A caller can offset the GM pointer before passing it to this function. + if offsets is not None: + raise NotImplementedError("DSA copy offsets are not supported; pass an adjusted pointer") + shape = DSASemantic.validate_alloc_shape(tl._unwrap_if_constexpr(shape)) + src_is_buf = isinstance(src, tle.buffered_tensor) + dst_is_buf = isinstance(dst, tle.buffered_tensor) + if not (src_is_buf or dst_is_buf): + raise ValueError("copy requires at least one buffered_tensor") + for value in (src, dst): + if isinstance(value, tle.buffered_tensor): + if value.type.shape != shape: + raise ValueError("copy shape must match the complete DSA buffer") + # Remote buffer handles need an explicit remote pointer view. + if hasattr(value.type, "_tle_remote_shard_id"): + raise NotImplementedError("Use local_ptr(remote(buffer, tile)) for NoC transfers") + + # Local -> local: both handles are memrefs, as required by dsa.copy. + # The binding takes the Triton 3.5 semantic object's builder explicitly. + if src_is_buf and dst_is_buf: + DSASemantic.validate_copy_dtype_compat(src.dtype, dst.dtype) + _dsa_ir.create_dsa_copy(_semantic.builder, src.handle, dst.handle) + return + + # GM <-> local: exactly one operand is a buffer. A Triton pointer cannot + # be passed to dsa.copy, so use pointer views and emit load/store IR below. + buffer = src if src_is_buf else dst + gm = dst if src_is_buf else src + if not isinstance(gm, tl.tensor) or not gm.dtype.is_ptr(): + raise ValueError("The global operand of DSA copy must be a pointer tensor") + DSASemantic.validate_copy_dtype_compat(buffer.dtype, gm.dtype.element_ty) + + # Broadcast each coordinate axis to the full buffer shape. local_ptr only + # builds an address view of the buffer; it does not copy any data itself. + indices = _make_full_indices(buffer, _semantic) + ptr = local_ptr(buffer, indices, _semantic=_semantic) + if not gm.shape: + # A scalar GM base pointer denotes contiguous row-major storage. + # Expand it to one pointer per element: shape (M, N) gives i * N + j. + # Pointer addition uses element offsets, so no byte-size factor is needed. + linear = tl.full(shape, 0, tl.int32, _semantic=_semantic) + stride = 1 + for axis in builtins.range(len(shape) - 1, -1, -1): + offset = indices[axis].__mul__(stride, _semantic=_semantic) + linear = linear.__add__(offset, _semantic=_semantic) + stride *= shape[axis] + gm = gm.__add__(linear, _semantic=_semantic) + elif tuple(gm.shape) != shape: + raise ValueError("Global pointer tensor shape must match the DSA buffer") + + # An already-shaped pointer tensor keeps its caller-supplied addressing. + # These transfers are unmasked: every supplied GM address must be valid. + # _semantic makes these tl calls generate IR during JIT compilation. + if src_is_buf: + # Local -> GM: read the local pointer view and write global pointers. + tl.store(gm, tl.load(ptr, _semantic=_semantic), _semantic=_semantic) + else: + # GM -> local: read global pointers and write the local pointer view. + tl.store(ptr, tl.load(gm, _semantic=_semantic), _semantic=_semantic) + + +def _expand_index_to_shape(index: tl.tensor, shape: Sequence[int], axis: int, _semantic) -> tl.tensor: + idx = index + for _ in builtins.range(axis): + idx = tl.expand_dims(idx, 0, _semantic=_semantic) + for _ in builtins.range(len(shape) - axis - 1): + idx = tl.expand_dims(idx, len(idx.shape), _semantic=_semantic) + return tl.broadcast_to(idx, *shape, _semantic=_semantic) + + +def _make_full_indices(buffer: tle.buffered_tensor, _semantic) -> tuple[tl.tensor, ...]: + shape = tuple(int(tl._unwrap_if_constexpr(dim)) for dim in buffer.type.shape) + indices = [] + for axis, dim in enumerate(shape): + idx = tl.arange(0, dim, _semantic=_semantic) + idx = _expand_index_to_shape(idx, shape, axis, _semantic) + indices.append(idx) + return tuple(indices) + + +@tl.builtin +def local_ptr( + buffer: tle.buffered_tensor, + indices: Optional[Sequence] = None, + _semantic=None, + _generator=None, +) -> tl.tensor: + """ + Materialize shared-memory pointers that cover the given buffered tensor. + + Args: + buffer: Local memory buffer tensor returned by ``tle.alloc``. + indices: Tuple of integer index tensors. The tuple length must equal + the rank of ``buffer`` and every tensor must have the same shape. + The output pointer tensor will have that same shape. + + Returns: + Tensor of pointers suitable for ``tl.load``/``tl.store``. + """ + if not isinstance(buffer, tle.buffered_tensor): + raise ValueError(f"Buffer parameter must be tle.buffered_tensor, but got {type(buffer)}") + + if _semantic is None: + raise ValueError("local_ptr must be used inside @triton.jit") + + # Preferred metadata source: buffered_tensor.type (survives JIT value + # reconstruction). Keep value attrs as backward-compatibility fallback. + remote_shard_id = getattr(buffer.type, "_tle_remote_shard_id", None) + remote_scope = getattr(buffer.type, "_tle_remote_scope", None) + if remote_shard_id is None: + remote_shard_id = getattr(buffer, "_tle_remote_shard_id", None) + remote_scope = getattr(buffer, "_tle_remote_scope", None) + remote_buffer_marker = remote_shard_id is not None + + indices = tl._unwrap_if_constexpr(indices) + if indices is None: + raise ValueError("local_ptr indices must be provided as a tuple of tensors") + if isinstance(indices, tl.tuple): + indices_tuple = tuple(indices.values) + elif isinstance(indices, (tuple, list)): + indices_tuple = tuple(indices) + else: + raise ValueError("local_ptr indices must be a tuple or list of tensors") + + buffer_shape = tuple(int(tl._unwrap_if_constexpr(dim)) for dim in buffer.type.shape) + if len(indices_tuple) != len(buffer_shape): + raise ValueError(f"local_ptr indices must provide {len(buffer_shape)} tensors, got {len(indices_tuple)}") + + idx_tensors: list[tensor] = [] + view_shape: Optional[tuple[int, ...]] = None + scalar_index_flags: list[bool] = [] + for idx in indices_tuple: + idx_tensor = idx if isinstance(idx, tensor) else _semantic.to_tensor(idx) + if not idx_tensor.dtype.is_int(): + raise ValueError("local_ptr indices must use integer dtypes") + is_scalar_index = not idx_tensor.type.is_block() + scalar_index_flags.append(is_scalar_index) + if is_scalar_index: + idx_tensors.append(idx_tensor) + continue + if view_shape is None: + view_shape = tuple(idx_tensor.shape) + elif tuple(idx_tensor.shape) != view_shape: + raise ValueError("local_ptr indices must have identical shapes") + idx_tensors.append(idx_tensor) + + if not idx_tensors: + raise ValueError("local_ptr indices cannot be empty") + all_scalar_indices = all(scalar_index_flags) + any_scalar_indices = any(scalar_index_flags) + if any_scalar_indices and not all_scalar_indices: + raise ValueError("local_ptr indices must be either all scalar or all tensors with identical shapes") + if not all_scalar_indices and view_shape is None: + view_shape = tuple() + + ptr_dtype = tl.pointer_type(buffer.type.element_ty) + insert_block = _semantic.builder.get_insertion_block() + if insert_block is None: + raise RuntimeError("TLE local_ptr called without an insertion block") + if all_scalar_indices: + result_ty = ptr_dtype + result_ir = ptr_dtype.to_ir(_semantic.builder) + else: + result_ty = tl.block_type(ptr_dtype, list(view_shape)) + result_ir = result_ty.to_ir(_semantic.builder) + handles = [idx.handle for idx in idx_tensors] + if not hasattr(_dsa_ir, "create_dsa_local_pointers"): + raise RuntimeError("builder missing create_dsa_local_pointers for DSA local_ptr") + local_ptr_op = _dsa_ir.create_dsa_local_pointers(_semantic.builder, result_ir, buffer.handle, *handles) + + result_tensor = tl.tensor(local_ptr_op.get_result(0), result_ty) + + if remote_buffer_marker: + if remote_scope is not None: + raise NotImplementedError("Wafer DSA remote pointers currently require physical tile IDs without a scope") + if all_scalar_indices: + raise ValueError("local_ptr does not yet support scalar indices on remote buffers") + if not hasattr(_dsa_ir, "create_dsa_remote_pointers"): + raise RuntimeError("builder missing create_dsa_remote_pointers for remote buffers") + shard_val = (remote_shard_id.handle if isinstance(remote_shard_id, tl.tensor) else _semantic.to_tensor(remote_shard_id).handle) + remote_op = _dsa_ir.create_dsa_remote_pointers(_semantic.builder, result_ir, result_tensor.handle, shard_val) + result_tensor = tl.tensor(remote_op.get_result(0), result_ty) + + return result_tensor + + +@tl.builtin +def to_tensor(buffer: tle.buffered_tensor, writable: bool = True, _semantic=None) -> tl.tensor: + """ + Convert a DSA ``buffered_tensor`` (on-chip buffer) into a ``tl.tensor`` view. + + This is a zero-copy view: the returned tensor aliases the on-chip buffer, so it + can participate in standard Triton tensor expressions without any data + movement. Lowered by ``--tle-to-mk`` into ``bufferization.to_tensor``. + + Args: + buffer: A ``buffered_tensor`` previously allocated with ``tle.language.dsa.alloc``. + writable: Mark the resulting tensor as writable (default ``True``). + + Returns: + ``tl.tensor`` aliasing the on-chip buffer contents. + """ + builder = _semantic.builder + if builder is None: + raise ValueError("to_tensor must be used inside @triton.jit") + if not isinstance(buffer, tle.buffered_tensor): + raise ValueError(f"to_tensor requires a buffered_tensor, got {type(buffer).__name__}") + if hasattr(buffer.type, "_tle_remote_shard_id"): + raise NotImplementedError("to_tensor requires a local buffer; use local_ptr for NoC") + shape = tuple(int(tl._unwrap_if_constexpr(dim)) for dim in buffer.type.shape) + result_ty = tl.block_type(buffer.type.element_ty, list(shape)) + result_ir = result_ty.to_ir(builder) + handle = _dsa_ir.create_dsa_to_tensor(builder, result_ir, buffer.handle, bool(writable)) + return tl.tensor(handle, result_ty) + + +@tl.builtin +def to_buffer(src: tl.tensor, space=None, _semantic=None) -> tle.buffered_tensor: + """ + Copy a ``tl.tensor`` into a newly-allocated DSA buffer and return it. + + This is the reverse bridge of :func:`to_tensor`: it materialises the tensor + value into a fresh on-chip buffer. Lowered by ``--tle-to-mk`` into + ``bufferization.to_buffer`` + ``memref.copy``. + + Args: + src: A ``tl.tensor`` value to store. + space: Storage scope for the new buffer. Defaults to the internal + scratchpad scope; may be a per-backend selector such as + ``tle.dsa.wafer.SPM``. + may be an address-space selector such as + ``tle.language.dsa.wafer.SPM``. + + Returns: + A new ``buffered_tensor`` containing a copy of ``src``. + """ + builder = _semantic.builder + if builder is None: + raise ValueError("to_buffer must be used inside @triton.jit") + if not isinstance(src, tl.tensor): + raise ValueError(f"to_buffer src must be a tl.tensor, got {type(src).__name__}") + shape = tuple(int(tl._unwrap_if_constexpr(dim)) for dim in src.shape) + if not shape: + raise ValueError("to_buffer src must be a non-scalar tensor") + buf = alloc(shape, src.dtype, scope=space, _semantic=_semantic) + _dsa_ir.create_dsa_to_buffer(builder, src.handle, buf.handle) + return buf + + +def _check_binary_operands(opname, input, other, result): + """Validate three-operand elementwise arithmetic operands. + + All three must be ``buffered_tensor`` with identical shape and element + dtype (tle.md: no implicit broadcast in this API layer). + """ + for name, val in (("input", input), ("other", other), ("result", result)): + if not isinstance(val, tle.buffered_tensor): + raise ValueError(f"{opname} {name} must be a buffered_tensor, " + f"got {type(val).__name__}") + input_shape = tuple(int(tl._unwrap_if_constexpr(dim)) for dim in input.type.shape) + other_shape = tuple(int(tl._unwrap_if_constexpr(dim)) for dim in other.type.shape) + result_shape = tuple(int(tl._unwrap_if_constexpr(dim)) for dim in result.type.shape) + if input_shape != other_shape or input_shape != result_shape: + raise ValueError(f"{opname} shape mismatch: input={input_shape}, " + f"other={other_shape}, result={result_shape}") + input_dtype = input.type.element_ty + other_dtype = other.type.element_ty + result_dtype = result.type.element_ty + if input_dtype != other_dtype or input_dtype != result_dtype: + raise ValueError(f"{opname} dtype mismatch: input={input_dtype}, " + f"other={other_dtype}, result={result_dtype}") + if not input_dtype.is_floating(): + raise NotImplementedError(f"{opname}: DSA buffer arithmetic currently requires floating point") + if any(hasattr(value.type, "_tle_remote_shard_id") for value in (input, other, result)): + raise NotImplementedError(f"{opname}: use local_ptr for NoC transfers before local arithmetic") + + +@tl.builtin +def add(input, other, result, _semantic=None): + """``result = input + other`` elementwise on on-chip buffers (three-operand).""" + builder = _semantic.builder + _check_binary_operands("add", input, other, result) + _dsa_ir.create_dsa_add(builder, input.handle, other.handle, result.handle) + + +@tl.builtin +def sub(input, other, result, _semantic=None): + """``result = input - other`` elementwise on on-chip buffers (three-operand).""" + builder = _semantic.builder + _check_binary_operands("sub", input, other, result) + _dsa_ir.create_dsa_sub(builder, input.handle, other.handle, result.handle) + + +@tl.builtin +def mul(input, other, result, _semantic=None): + """``result = input * other`` elementwise on on-chip buffers (three-operand).""" + builder = _semantic.builder + _check_binary_operands("mul", input, other, result) + _dsa_ir.create_dsa_mul(builder, input.handle, other.handle, result.handle) + + +@tl.builtin +def max(input, other, result, _semantic=None): + """``result = max(input, other)`` elementwise on on-chip buffers (three-operand).""" + builder = _semantic.builder + _check_binary_operands("max", input, other, result) + _dsa_ir.create_dsa_maximum(builder, input.handle, other.handle, result.handle) + + +@tl.builtin +def min(input, other, result, _semantic=None): + """``result = min(input, other)`` elementwise on on-chip buffers (three-operand).""" + builder = _semantic.builder + _check_binary_operands("min", input, other, result) + _dsa_ir.create_dsa_minimum(builder, input.handle, other.handle, result.handle) + + +@tl.builtin +def div(input, other, result, _semantic=None): + """``result = input / other`` elementwise on on-chip buffers (three-operand).""" + builder = _semantic.builder + _check_binary_operands("div", input, other, result) + _dsa_ir.create_dsa_div(builder, input.handle, other.handle, result.handle) + + +# --------------------------------------------------------------------------- +# dsa.extract_slice / dsa.insert_slice +# +# Element-level strided slicing (spec 3.3.2.4). The tile grid-coordinate +# convenience forms are provided by the generic tle-lite tier +# (``tle.extract_tile`` / ``tle.insert_tile``) and lower through the shared +# tle dialect with backend-specific conversion patterns. +# Offsets may mix static (int/constexpr) and dynamic (scalar tl.tensor) dims, +# encoded with ShapedType::kDynamic as the sentinel in `static_offsets`. +# --------------------------------------------------------------------------- + +_DYNAMIC = -(1 << 63) # int64 min == ShapedType::kDynamic sentinel + + +def _try_unwrap_int(val): + """Return ``val`` as a Python int if int/constexpr-like, else ``None``.""" + if isinstance(val, int): + return val + v = tl._unwrap_if_constexpr(val) + return v if isinstance(v, int) else None + + +def _static_dims(shape, fn): + """Unwrap a shape into compile-time ints, erroring on dynamic dims.""" + shape = tl._unwrap_if_constexpr(shape) + dims = [tl._unwrap_if_constexpr(d) for d in shape] + if any(not isinstance(d, int) for d in dims): + raise ValueError(f"{fn}: shape must be compile-time constants, got {shape}") + return dims + + +def _split_offsets(offsets, fn): + """Split per-dim offsets (int/constexpr or scalar tl.tensor) into + ``(static_offsets, dyn_handles)``; dynamic dims use ``_DYNAMIC`` sentinel.""" + offsets = tl._unwrap_if_constexpr(offsets) + static = [] + dyn = [] + for v in offsets: + if isinstance(v, tl.tensor): + if len(v.shape) != 0: + raise ValueError(f"{fn}: dynamic offsets must be scalar tl.tensor") + static.append(_DYNAMIC) + dyn.append(v.handle) + else: + iv = _try_unwrap_int(v) + if iv is None: + raise ValueError(f"{fn}: offsets must be int/constexpr or scalar tl.tensor") + static.append(iv) + return static, dyn + + +def _check_static_slice_bounds(fn, src_shape, static_offsets, sizes, strides): + """Validate the static (non-sentinel) offsets against src/sizes/strides.""" + for i, (src, off, size, stride) in enumerate(zip(src_shape, static_offsets, sizes, strides)): + if size <= 0: + raise ValueError(f"{fn}: size[{i}]={size} must be positive") + if stride <= 0: + raise ValueError(f"{fn}: stride[{i}]={stride} must be positive") + if off == _DYNAMIC: + continue + if off < 0: + raise ValueError(f"{fn}: offset[{i}]={off} must be non-negative") + end = off + (size - 1) * stride + 1 + if end > src: + raise ValueError(f"{fn}: slice [{off}:{off}+{size}*{stride}] exceeds source dim " + f"{i} ({src})") + + +@tl._tensor_member_fn +@tl.builtin +def extract_slice(x: tl.tensor, offsets, sizes, strides, _semantic=None) -> tl.tensor: + """Extract a strided slice from ``x``. + + Args: + x: Source ``tl.tensor``. + offsets: Per-dim offsets; each element is an int/constexpr (static) or + a scalar ``tl.tensor`` (dynamic). + sizes: Per-dim slice sizes (compile-time constants); this is the + result shape (``tensor.extract_slice`` semantics). + strides: Per-dim strides (compile-time constants). + + Returns a tensor whose shape is ``sizes``. + """ + if not isinstance(x, tl.tensor): + raise ValueError(f"extract_slice: source must be tl.tensor, got {type(x)}") + + builder = _semantic.builder + if builder is None: + raise ValueError("extract_slice must be used inside @triton.jit") + + src_shape = _static_dims(x.type.shape, "extract_slice") + sizes = _static_dims(sizes, "extract_slice") + strides = _static_dims(strides, "extract_slice") + offsets = tl._unwrap_if_constexpr(offsets) + if len(offsets) != len(src_shape) or len(sizes) != len(src_shape) \ + or len(strides) != len(src_shape): + raise ValueError("extract_slice: offsets/sizes/strides rank must match source rank") + + static_offsets, dyn_offsets = _split_offsets(offsets, "extract_slice") + _check_static_slice_bounds("extract_slice", src_shape, static_offsets, sizes, strides) + + # tensor.extract_slice semantics: sizes is the result shape. + result_ty = tl.block_type(x.type.element_ty, sizes) + result_ir = result_ty.to_ir(builder) + + handle = _dsa_ir.create_dsa_extract_slice(builder, result_ir, x.handle, static_offsets, dyn_offsets, sizes, strides) + return tl.tensor(handle, result_ty) + + +@tl._tensor_member_fn +@tl.builtin +def insert_slice(x: tl.tensor, tile: tl.tensor, offsets, sizes=None, strides=None, _semantic=None) -> tl.tensor: + """Insert ``tile`` into ``x`` at ``offsets``. + + ``sizes`` defaults to ``tile.shape``; ``strides`` defaults to all ones. + Returns a new tensor with the same shape/type as ``x``. + """ + if not isinstance(x, tl.tensor): + raise ValueError(f"insert_slice: source must be tl.tensor, got {type(x)}") + if not isinstance(tile, tl.tensor): + raise ValueError(f"insert_slice: tile must be tl.tensor, got {type(tile)}") + + builder = _semantic.builder + if builder is None: + raise ValueError("insert_slice must be used inside @triton.jit") + + src_shape = _static_dims(x.type.shape, "insert_slice") + tile_shape = _static_dims(tile.type.shape, "insert_slice") + offsets = tl._unwrap_if_constexpr(offsets) + if len(offsets) != len(src_shape): + raise ValueError("insert_slice: offsets rank must match source rank") + if sizes is None: + sizes = tile_shape + else: + sizes = _static_dims(sizes, "insert_slice") + if strides is None: + strides = [1] * len(src_shape) + else: + strides = _static_dims(strides, "insert_slice") + if len(sizes) != len(src_shape) or len(strides) != len(src_shape): + raise ValueError("insert_slice: sizes/strides rank must match source rank") + if tuple(sizes) != tuple(tile_shape): + raise ValueError(f"insert_slice: sizes {sizes} must match tile shape {tile_shape}") + if x.type.element_ty != tile.type.element_ty: + raise ValueError(f"insert_slice: element type mismatch source={x.type.element_ty}, " + f"tile={tile.type.element_ty}") + + static_offsets, dyn_offsets = _split_offsets(offsets, "insert_slice") + _check_static_slice_bounds("insert_slice", src_shape, static_offsets, sizes, strides) + + handle = _dsa_ir.create_dsa_insert_slice(builder, x.type.to_ir(builder), x.handle, tile.handle, static_offsets, dyn_offsets, + sizes, strides) + return tl.tensor(handle, x.type) diff --git a/third_party/wafer/experimental/tle/language/dsa/semantic.py b/third_party/wafer/experimental/tle/language/dsa/semantic.py new file mode 100644 index 00000000..e354db2d --- /dev/null +++ b/third_party/wafer/experimental/tle/language/dsa/semantic.py @@ -0,0 +1,179 @@ +""" +DSA Semantic Validation Layer +============================= + +Provides early, human-readable error messages for invalid TLE DSA operations +before they reach the MLIR lowering pipeline. Mirrors the role of +``flagtree_tle``'s ``TLESemantic`` class but adapted for the TsingMicro / +DSA backend. +""" + +from __future__ import annotations + +from typing import Optional, Sequence, Tuple + +import triton.language.core as tl +from . import types as tle + + +class DSASemanticError(Exception): + """Raised when a DSA operation fails semantic validation.""" + pass + + +# Data types supported by the TsingMicro DSA backend for buffer allocation. +_SUPPORTED_ALLOC_DTYPES = frozenset([ + tl.float32, + tl.float16, + tl.bfloat16, + tl.int8, + tl.int16, + tl.int32, + tl.int64, + tl.uint8, + tl.uint16, + tl.uint32, + tl.uint64, +]) + + +class DSASemantic: + """Semantic analyzer for DSA TLE operations. + + Each ``validate_*`` method raises :class:`DSASemanticError` with a + descriptive message if validation fails, and returns silently on + success. + """ + + # ------------------------------------------------------------------ + # alloc() validation + # ------------------------------------------------------------------ + + @staticmethod + def validate_alloc_shape(shape: Sequence) -> Tuple[int, ...]: + """Validate and normalise *shape* for ``alloc()``. + + Returns the unwrapped shape tuple on success. + """ + if not isinstance(shape, (tuple, list)): + if hasattr(shape, "__iter__"): + shape = tuple(shape) + else: + raise DSASemanticError(f"alloc: shape must be a tuple or list, got {type(shape).__name__}") + + unwrapped = [] + for i, dim in enumerate(shape): + dim = tl._unwrap_if_constexpr(dim) + if not isinstance(dim, int) or dim <= 0: + raise DSASemanticError(f"alloc: shape[{i}] must be a positive integer, got {dim!r}") + unwrapped.append(dim) + return tuple(unwrapped) + + @staticmethod + def validate_alloc_dtype(dtype: tl.dtype) -> tl.dtype: + """Validate *dtype* for ``alloc()``.""" + dtype = tl._unwrap_if_constexpr(dtype) + if not isinstance(dtype, tl.dtype): + raise DSASemanticError(f"alloc: dtype must be a tl.dtype instance, got {type(dtype).__name__}") + if dtype not in _SUPPORTED_ALLOC_DTYPES: + supported = ", ".join(str(d) for d in sorted(_SUPPORTED_ALLOC_DTYPES, key=str)) + raise DSASemanticError(f"alloc: unsupported dtype {dtype}. Supported types: {supported}") + return dtype + + @staticmethod + def validate_alloc_scope(scope) -> tle.scope: + """Validate *scope* for ``alloc()``.""" + if scope is None: + return tle.spm # default + if not isinstance(scope, tle.scope): + raise DSASemanticError(f"alloc: scope must be a tle.scope instance, got {type(scope).__name__}") + return scope + + # ------------------------------------------------------------------ + # copy() validation + # ------------------------------------------------------------------ + + @staticmethod + def validate_copy_operands(src, dst) -> str: + """Validate *src*/*dst* types for ``copy()`` and return a direction tag. + + Returns one of ``"SPM_TO_SPM"``, ``"GM_TO_SPM"``, ``"SPM_TO_GM"``. + + Raises :class:`DSASemanticError` if the combination is unsupported. + """ + src_is_buf = isinstance(src, tle.buffered_tensor) + dst_is_buf = isinstance(dst, tle.buffered_tensor) + + if src_is_buf and dst_is_buf: + return "SPM_TO_SPM" + if (not src_is_buf) and dst_is_buf: + if not isinstance(src, tl.tensor): + raise DSASemanticError(f"copy: src must be tl.tensor or buffered_tensor, got {type(src).__name__}") + return "GM_TO_SPM" + if src_is_buf and (not dst_is_buf): + if not isinstance(dst, tl.tensor): + raise DSASemanticError(f"copy: dst must be tl.tensor or buffered_tensor, got {type(dst).__name__}") + return "SPM_TO_GM" + raise DSASemanticError("copy: at least one operand must be a buffered_tensor. " + f"Got src={type(src).__name__}, dst={type(dst).__name__}") + + @staticmethod + def validate_copy_dtype_compat(src_dtype, dst_dtype) -> None: + """Check that element types of *src* and *dst* are compatible.""" + if src_dtype != dst_dtype: + raise DSASemanticError(f"copy: element type mismatch – src has {src_dtype}, dst has {dst_dtype}") + + # ------------------------------------------------------------------ + # local_ptr() validation + # ------------------------------------------------------------------ + + @staticmethod + def validate_local_ptr_buffer(buffer) -> None: + """Validate that *buffer* is a proper ``buffered_tensor``.""" + if not isinstance(buffer, tle.buffered_tensor): + raise DSASemanticError(f"local_ptr: buffer must be a buffered_tensor, got {type(buffer).__name__}") + if buffer.type.shape is None: + raise DSASemanticError("local_ptr: buffer shape is None (deferred shapes not yet supported)") + + @staticmethod + def validate_local_ptr_indices( + indices: Sequence, + buffer_rank: int, + ) -> None: + """Validate *indices* for ``local_ptr()``. + + Checks: + - indices length matches buffer rank + - all indices are integer-typed + - indices are either all scalar or all tensor with matching shapes + """ + if indices is None: + raise DSASemanticError("local_ptr: indices must be provided as a tuple of tensors") + if len(indices) != buffer_rank: + raise DSASemanticError(f"local_ptr: expected {buffer_rank} index tensors, got {len(indices)}") + + view_shape: Optional[tuple] = None + has_scalar = False + has_tensor = False + + for i, idx in enumerate(indices): + if not isinstance(idx, tl.tensor): + raise DSASemanticError(f"local_ptr: indices[{i}] must be a tl.tensor, " + f"got {type(idx).__name__}") + if not idx.dtype.is_int(): + raise DSASemanticError(f"local_ptr: indices[{i}] must have integer dtype, " + f"got {idx.dtype}") + is_scalar = not idx.type.is_block() + if is_scalar: + has_scalar = True + else: + has_tensor = True + if view_shape is None: + view_shape = tuple(idx.shape) + elif tuple(idx.shape) != view_shape: + raise DSASemanticError(f"local_ptr: index tensor shape mismatch at dim {i}: " + f"expected {view_shape}, got {tuple(idx.shape)}") + + if has_scalar and has_tensor: + raise DSASemanticError("local_ptr: indices must be either all scalar or all " + "tensor with identical shapes (mixed not allowed)") diff --git a/third_party/wafer/experimental/tle/language/dsa/types.py b/third_party/wafer/experimental/tle/language/dsa/types.py new file mode 100644 index 00000000..c510154b --- /dev/null +++ b/third_party/wafer/experimental/tle/language/dsa/types.py @@ -0,0 +1,126 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, List, Tuple + +from triton.language.core import base_type, base_value, dtype, _unwrap_if_constexpr + + +@dataclass(frozen=True) +class scope: + """ + Simple storage descriptor for DSA buffers. + + This is intentionally backend-agnostic. `name` / `value` / `memory_space` + are carried through as metadata only; concrete lowering is handled by the + DSA dialect and backend. + """ + + name: str + value: str + memory_space: str + + def __repr__(self) -> str: + return self.name + + +# DSA storage scopes. +local = scope("local", "local", "local") + +# Scratch Pad Memory – the on-chip SRAM exposed by TsingMicro Wafer. +# This is the primary storage scope for DSA kernels and serves the same +# conceptual role as NVIDIA shared memory (smem). +spm = scope("spm", "spm", "spm") + + +class buffered_tensor(base_value): + """ + Symbolic handle to a buffer allocated via DSA. + + This is a thin wrapper over an IR value plus a `buffered_tensor_type` + describing shape / element dtype / memory space. + """ + + def __init__(self, handle: Any, ty: "buffered_tensor_type"): + self.handle = handle + self.type = ty + self.shape = ty.shape + self.dtype = ty.element_ty + + def _flatten_ir(self, handles: List[Any]) -> None: + handles.append(self.handle) + + +class buffered_tensor_type(base_type): + """ + Frontend description of a DSA buffer. + + - `shape`: logical block shape (may be None for deferred shapes) + - `element_ty`: scalar dtype + - `storage`: abstract storage scope (currently `local`) + - `memory_space`: backend-visible memory space string, defaults to + `storage.memory_space` + """ + + def __init__( + self, + shape, + element_ty: dtype, + storage: scope | None = None, + memory_space: str = "", + ): + if shape is None: + self.shape = None + else: + shape = _unwrap_if_constexpr(shape) + self.shape = tuple(int(_unwrap_if_constexpr(x)) for x in shape) + self.element_ty = _unwrap_if_constexpr(element_ty) + self.storage = storage if storage is not None else local + self.memory_space = (str(_unwrap_if_constexpr(memory_space)) if memory_space else self.storage.memory_space) + + @property + def scalar(self) -> dtype: + return self.element_ty + + def __eq__(self, other: object) -> bool: + if not isinstance(other, buffered_tensor_type): + return False + return ( + self.shape, + self.element_ty, + self.storage, + self.memory_space, + ) == ( + other.shape, + other.element_ty, + other.storage, + other.memory_space, + ) + + def __repr__(self) -> str: + shape = "?" if self.shape is None else "x".join(map(str, self.shape)) + return f"buffered_tensor_type<{shape}, {self.element_ty}, {self.memory_space}>" + + def _unflatten_ir(self, handles: List[Any], cursor: int) -> Tuple[base_value, int]: + value = buffered_tensor(handles[cursor], self) + # Preserve remote metadata if present on the type. + if hasattr(self, "_tle_remote_shard_id"): + shard_id = getattr(self, "_tle_remote_shard_id") + scope = getattr(self, "_tle_remote_scope", None) + setattr(value, "_tle_remote_shard_id", shard_id) + setattr(value, "_tle_remote_scope", scope) + setattr(value.type, "_tle_remote_shard_id", shard_id) + setattr(value.type, "_tle_remote_scope", scope) + return value, cursor + 1 + + def mangle(self) -> str: + if hasattr(self, "_tle_remote_shard_id"): + raise NotImplementedError("Passing a remote-marked DSA buffer between JIT functions is not supported") + return "dsa_" + self.memory_space + "_" + "_".join(map(str, self.shape)) + "_" + self.element_ty.mangle() + + def _flatten_ir_types(self, builder, out: List[Any]) -> None: + if hasattr(self, "_tle_remote_shard_id"): + raise NotImplementedError("Passing a remote-marked DSA buffer between JIT functions is not supported") + if not hasattr(self, "_ir_type"): + raise ValueError("DSA buffer types must originate from tle.dsa.alloc") + out.append(self._ir_type) diff --git a/third_party/wafer/experimental/tle/language/dsa/wafer/__init__.py b/third_party/wafer/experimental/tle/language/dsa/wafer/__init__.py new file mode 100644 index 00000000..906bbe64 --- /dev/null +++ b/third_party/wafer/experimental/tle/language/dsa/wafer/__init__.py @@ -0,0 +1,4 @@ +"""Wafer scratchpad and hardware random-number primitives.""" +from .core import SPM, randgen, rand, randn + +__all__ = ["SPM", "randgen", "rand", "randn"] diff --git a/third_party/wafer/experimental/tle/language/dsa/wafer/core.py b/third_party/wafer/experimental/tle/language/dsa/wafer/core.py new file mode 100644 index 00000000..d52ff137 --- /dev/null +++ b/third_party/wafer/experimental/tle/language/dsa/wafer/core.py @@ -0,0 +1,194 @@ +# flagtree tle +"""TsingMicro Wafer vendor-specific DSA primitives. + +Exposed as ``triton.experimental.tle.language.dsa.wafer``, mirroring the +upstream per-backend namespace convention (see ``dsa.ascend``). The hardware +random-generation family and vendor address-space selectors live here; the +multi-vendor DSA API surface stays at the ``dsa`` root. +""" + +from triton._C.libtriton.wafer import tle as _dsa_ir +import triton.language.core as tl +from triton.language.core import PropagateNan +from triton.language import math as tlmath + +from ..types import scope + +# Address-space selector for the on-chip scratchpad memory (vendor naming). +SPM = scope("wafer.spm", "spm", "spm") + +# Fmt_INT64 in wafer wafer Data_Format enum (see instr_def / op_def.h). +_DSA_RANDGEN_FMT_INT64 = 11 +_DSA_RANDGEN_NUM_STREAMS = 16 # values per step (128 bytes) + + +@tl.builtin +def randgen(seed0, seed1, n_out: tl.constexpr, _semantic=None): + """ + Hardware random-number generator (xorshift128+ peri) on Wafer. + + Args: + seed0: block tensor ``[16]`` of ``int64`` / ``uint64`` seeds (stream a). + seed1: block tensor ``[16]`` of ``int64`` / ``uint64`` seeds (stream b). + n_out: number of ``int64`` random outputs; must be a multiple of 16 + (hardware emits 16 values / 128 bytes per step). + + Returns: + ``(out, seed0_out, seed1_out)`` with ``out`` shaped ``[n_out]``. + ``seed0_out`` / ``seed1_out`` mirror the inputs: the peri does not + expose state readback, so identical seeds yield identical blocks. + Vary the seeds across calls to obtain different data. + + Notes: + Output values are raw xorshift128+ ``uint64`` bit patterns stored as + ``int64``. Convert to Uniform(0,1) / Normal yourself (see ``rand`` / + ``randn`` helpers), or feed downstream kernels. + """ + n_out = int(tl._unwrap_if_constexpr(n_out)) + if n_out <= 0 or (n_out % _DSA_RANDGEN_NUM_STREAMS) != 0: + raise ValueError(f"tle.dsa.wafer.randgen n_out must be a positive multiple of " + f"{_DSA_RANDGEN_NUM_STREAMS}, got {n_out}") + + builder = _semantic.builder + + if not isinstance(seed0, tl.tensor): + seed0 = tl.to_tensor(seed0, _semantic=_semantic) + if not isinstance(seed1, tl.tensor): + seed1 = tl.to_tensor(seed1, _semantic=_semantic) + + if not seed0.dtype.is_int() or not seed1.dtype.is_int(): + raise ValueError("tle.dsa.wafer.randgen seeds must be integer tensors") + if seed0.dtype != tl.int64: + seed0 = seed0.to(tl.int64, _semantic=_semantic) + if seed1.dtype != tl.int64: + seed1 = seed1.to(tl.int64, _semantic=_semantic) + + seed0_ty = seed0.type + seed1_ty = seed1.type + if (not seed0_ty.is_block() or not seed1_ty.is_block() + or tuple(int(tl._unwrap_if_constexpr(d)) for d in seed0_ty.shape) != (_DSA_RANDGEN_NUM_STREAMS, ) + or tuple(int(tl._unwrap_if_constexpr(d)) for d in seed1_ty.shape) != (_DSA_RANDGEN_NUM_STREAMS, )): + raise ValueError("tle.dsa.wafer.randgen seeds must be block tensors of shape [16]") + + out_ty = tl.block_type(tl.int64, [n_out]) + byte_count = n_out * 8 + + rand_op = _dsa_ir.create_dsa_randgen(builder, + out_ty.to_ir(builder), + seed0_ty.to_ir(builder), + seed1_ty.to_ir(builder), + seed0.handle, + seed1.handle, + int(byte_count), + int(_DSA_RANDGEN_FMT_INT64), + ) + out = tl.tensor(rand_op.get_result(0), out_ty) + seed0_out = tl.tensor(rand_op.get_result(1), seed0_ty) + seed1_out = tl.tensor(rand_op.get_result(2), seed1_ty) + return out, seed0_out, seed1_out + + +def _uint32_bits_to_uniform(bits32, semantic): + """ + Map random int32 bits to Uniform(0, 1) via IEEE754 mantissa stuffing. + + u = bitcast((bits & 0x7FFFFF) | 0x3F800000, f32) - 1.0 + + Avoids the sitofp / where / cmp / sub integer chain that lowers to + per-element scf.for on Wafer (no int vector ALU). Uses only bitwise + and/or + bitcast + float sub, which can stay on peri / float paths. + + ``semantic`` is the SemanticAnalyzer (bound methods, not a raw builder). + """ + # 0x7FFFFF fits int32; only the bit pattern of 0x3F800000 matters for + # the bitcast. + mant_mask = semantic.to_tensor(0x7FFFFF) + one_bits = semantic.to_tensor(0x3F800000) + one_f = semantic.to_tensor(1.0) + mant = semantic.and_(bits32, mant_mask) + packed = semantic.or_(mant, one_bits) + f12 = semantic.bitcast(packed, tl.float32) # [1, 2) + return semantic.sub(f12, one_f, True) # [0, 1) + + +def _i64_as_i32_view(raw_i64, n_i32: int, builder): + """ + Zero-copy view of an i64 buffer as i32 (little-endian: lo32, hi32, ...). + + ``raw_i64`` must have shape ``[n_i32 // 2]``. Emits vendor-neutral + ``dsa.bitcast`` (backends alias the buffer; no elementwise ``trunci``). + """ + n_i64 = int(tl._unwrap_if_constexpr(raw_i64.shape[0])) + if n_i32 != n_i64 * 2: + raise ValueError(f"i64->i32 view expects n_i32 == 2 * n_i64, got {n_i32} vs 2*{n_i64}") + dst_ty = tl.block_type(tl.int32, [n_i32]) + handle = _dsa_ir.create_dsa_bitcast(builder, dst_ty.to_ir(builder), raw_i64.handle) + return tl.tensor(handle, dst_ty) + + +@tl.builtin +def rand(seed0, seed1, n_out: tl.constexpr, _semantic=None): + """ + Uniform(0, 1) floats via hardware ``randgen`` + float scaling. + + ``n_out`` must be a multiple of 32 (``randgen`` emits i64; each i64 + contributes two i32 samples via a zero-copy view). + + Returns ``(u, seed0_out, seed1_out)`` with ``u`` shaped ``[n_out]`` float32. + """ + n_out = int(tl._unwrap_if_constexpr(n_out)) + if n_out <= 0 or (n_out % 32) != 0: + raise ValueError(f"tle.dsa.wafer.rand n_out must be a positive multiple of 32, got {n_out}") + + builder = _semantic.builder + # Half as many i64 draws: lo/hi 32-bit halves become two Uniform samples. + raw64, seed0_out, seed1_out = randgen(seed0, seed1, n_out // 2, _semantic=_semantic) + bits32 = _i64_as_i32_view(raw64, n_out, builder) + u = _uint32_bits_to_uniform(bits32, _semantic) + return u, seed0_out, seed1_out + + +@tl.builtin +def randn(seed0, seed1, n_out: tl.constexpr, _semantic=None): + """ + Normal(0, 1) floats via hardware ``randgen`` + Box-Muller. + + ``n_out`` must be a multiple of 32 (two Uniform halves of size + ``n_out // 2``, each backed by ``n_out // 4`` i64 draws viewed as i32). + + Uses a single ``randgen`` of length ``n_out // 2`` and pairs consecutive + Uniform samples ``(u[2i], u[2i+1])`` for Box-Muller. Calling ``rand`` + twice is unsafe while the peri does not advance ``seed0_out`` / + ``seed1_out``. + + Returns ``(n, seed0_out, seed1_out)`` with ``n`` shaped ``[n_out]`` float32 + (concatenation of the two Box-Muller outputs). + """ + n_out = int(tl._unwrap_if_constexpr(n_out)) + if n_out <= 0 or (n_out % 32) != 0: + raise ValueError(f"tle.dsa.wafer.randn n_out must be a positive multiple of 32, got {n_out}") + + half = n_out // 2 + + # Box-Muller over consecutive Uniform pairs. Scalar f32 constants are + # broadcast from u_half (``u*0+c``): direct ``semantic.to_tensor(float)`` + # scalars bias the result on Wafer. + uv, seed0_out, seed1_out = rand(seed0, seed1, n_out, _semantic=_semantic) + + pairs = _semantic.reshape(uv, [half, 2], False) + u_half, v_half = _semantic.split(pairs) + + # Force f32-typed scalars via broadcast from u_half (avoids f64 pitfalls). + zero = _semantic.mul(u_half, _semantic.to_tensor(0.0), True) + eps = _semantic.add(zero, _semantic.to_tensor(1.0e-7), True) + two_pi = _semantic.add(zero, _semantic.to_tensor(6.283185307179586), True) + neg_two = _semantic.add(zero, _semantic.to_tensor(-2.0), True) + + u1 = _semantic.maximum(u_half, eps, PropagateNan.NONE) + theta = _semantic.mul(two_pi, v_half, True) + log_u1 = tlmath.log(u1, _semantic=_semantic) + r = tlmath.sqrt(_semantic.mul(neg_two, log_u1, True), _semantic=_semantic) + n0 = _semantic.mul(r, tlmath.cos(theta, _semantic=_semantic), True) + n1 = _semantic.mul(r, tlmath.sin(theta, _semantic=_semantic), True) + out = _semantic.reshape(_semantic.join(n0, n1), [n_out], False) + return out, seed0_out, seed1_out diff --git a/third_party/wafer/include/Address/CMakeLists.txt b/third_party/wafer/include/Address/CMakeLists.txt new file mode 100755 index 00000000..557daa84 --- /dev/null +++ b/third_party/wafer/include/Address/CMakeLists.txt @@ -0,0 +1,2 @@ +add_subdirectory(Dialect) +add_subdirectory(Transforms) diff --git a/third_party/wafer/include/Address/Dialect/CMakeLists.txt b/third_party/wafer/include/Address/Dialect/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/include/Address/Dialect/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/include/Address/Dialect/IR/AddressDialect.h b/third_party/wafer/include/Address/Dialect/IR/AddressDialect.h new file mode 100755 index 00000000..c499db0c --- /dev/null +++ b/third_party/wafer/include/Address/Dialect/IR/AddressDialect.h @@ -0,0 +1,36 @@ +//===- AddressDialect.h - Address dialect -----------------------*- C++ -*-===// +// +// This file is licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file defines the Address dialect. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_DIALECT_ADDRESS_IR_ADDRESSDIALECT_H +#define MLIR_DIALECT_ADDRESS_IR_ADDRESSDIALECT_H + +#include "mlir/Bytecode/BytecodeOpInterface.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/Interfaces/CastInterfaces.h" +#include "mlir/Interfaces/InferTypeOpInterface.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" + +#include "Address/Dialect/IR/AddressOpsDialect.h.inc" + +namespace mlir { +class PatternRewriter; +} + +#define GET_TYPEDEF_CLASSES +#include "Address/Dialect/IR/AddressOpsTypes.h.inc" + +#define GET_OP_CLASSES +#include "Address/Dialect/IR/AddressOps.h.inc" + +#endif // MLIR_DIALECT_ADDRESS_IR_ADDRESSDIALECT_H diff --git a/third_party/wafer/include/Address/Dialect/IR/AddressDialect.td b/third_party/wafer/include/Address/Dialect/IR/AddressDialect.td new file mode 100755 index 00000000..f6dec91f --- /dev/null +++ b/third_party/wafer/include/Address/Dialect/IR/AddressDialect.td @@ -0,0 +1,77 @@ +//===- AddressDialect.td - Address dialect -----------*- tablegen -*-===// +// +// This file is licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#ifndef ADDRESS_DIALECT +#define ADDRESS_DIALECT + +include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/BuiltinTypeInterfaces.td" +include "mlir/IR/OpBase.td" + +//===----------------------------------------------------------------------===// +// Address dialect definition. +//===----------------------------------------------------------------------===// + +def Address_Dialect : Dialect { + let name = "addr"; + let summary = "address dialect"; + let cppNamespace = "::mlir::addr"; + let useDefaultTypePrinterParser = 1; + let extraClassDeclaration = [{ + void registerTypes(); + }]; +} + +//===----------------------------------------------------------------------===// +// Address type definitions +//===----------------------------------------------------------------------===// + +class Address_Type traits = []> + : TypeDef { + let mnemonic = typeMnemonic; +} + +def AddressType : Address_Type<"Address", "address", [ + MemRefElementTypeInterface + ]> { + let summary = "Type for holding addresses"; + let description = [{ + Syntax: + + ```mlir + address ::= `address` `<` (address-space)? `>` + address-space ::= attribute-value + ``` + `address` is a type for representing memory addresses, including its address + space. Its size is target and address space dependent, and it is only known + once it is lowered. + }]; + let parameters = (ins OptionalParameter<"Attribute">:$addressSpace); + let builders = [ + TypeBuilder<(ins + CArg<"Attribute", "{}">:$addressSpace), [{ + return Base::get($_ctxt, addressSpace); + }]>, + TypeBuilderWithInferredContext<(ins + "MemRefType":$type), [{ + assert(type && "expected a valid memref type"); + return Base::get(type.getContext(), type.getMemorySpace()); + }]> + ]; + let assemblyFormat = "(`<` $addressSpace^ `>`)?"; + let skipDefaultBuilders = 1; +} + +//===----------------------------------------------------------------------===// +// Base address operation definition. +//===----------------------------------------------------------------------===// + +class Address_Op traits = []> : + Op; + +#endif // ADDRESS_DIALECT diff --git a/third_party/wafer/include/Address/Dialect/IR/AddressOps.td b/third_party/wafer/include/Address/Dialect/IR/AddressOps.td new file mode 100755 index 00000000..9c486988 --- /dev/null +++ b/third_party/wafer/include/Address/Dialect/IR/AddressOps.td @@ -0,0 +1,247 @@ +//===- AddressOps.td - Address dialect ops -----------------*- tablegen -*-===// +// +// This file is licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#ifndef ADDRESS_OPS +#define ADDRESS_OPS + +include "Address/Dialect/IR/AddressDialect.td" +include "mlir/Interfaces/CastInterfaces.td" +include "mlir/Interfaces/InferTypeOpInterface.td" +include "mlir/Interfaces/SideEffectInterfaces.td" +include "mlir/IR/OpAsmInterface.td" + +//===----------------------------------------------------------------------===// +// ConstantOp +//===----------------------------------------------------------------------===// + +def Address_ConstantOp : Address_Op<"constant", [ + ConstantLike, Pure, + DeclareOpInterfaceMethods + ]> { + let summary = "Creates an address constant."; + let description = [{ + The `addr.constant` operation produces an address-typed SSA value equal to + some index constant. + + Example: + + ```mlir + %addr0 = address.constant 0 + %addr1 = address.constant 1 : !address<3 : i32> + ``` + }]; + let arguments = (ins IndexAttr:$value); + let results = (outs AddressType:$result); + let builders = [ + OpBuilder<(ins "int64_t":$value, CArg<"Attribute", "nullptr">:$addressSpace)> + ]; + let assemblyFormat = "attr-dict $value custom(type($result))"; + let hasFolder = 1; +} + +//===----------------------------------------------------------------------===// +// TypeOffsetOp +//===----------------------------------------------------------------------===// + +def Address_TypeOffsetOp : Address_Op<"type_offset", [ConstantLike, Pure]> { + let summary = "Creates a type offset constant."; + let description = [{ + The `addr.type_offset` operation produces an int or index-typed SSA value + equal to a target-specific constant representing the offset of a single + element of the given type. The default return type is `index`. + Example: + + ```mlir + %0 = addr.type_offset f32 + %1 = addr.type_offset memref<12 x f64> : i32 + ``` + }]; + + let arguments = (ins TypeAttr:$baseType); + let results = (outs AnySignlessIntegerOrIndex:$result); + let builders = [ + OpBuilder<(ins "TypeAttr":$baseType, CArg<"Type", "nullptr">:$resultTy)> + ]; + let assemblyFormat = [{ + attr-dict $baseType custom(type($result)) + }]; + let hasFolder = 1; +} + +//===----------------------------------------------------------------------===// +// CastOp +//===----------------------------------------------------------------------===// + +def Address_CastOp : Address_Op<"cast", [Pure]> { + let summary = "Creates an address space cast"; + let description = [{ + The `addr.cast` operation casts addresses between address spaces. + Example: + + ```mlir + %addr = addr.cast %addr : !address to !address<1 : i32> + ``` + }]; + let arguments = (ins AddressType:$input); + let results = (outs AddressType:$result); + let builders = [ + OpBuilder<(ins "Attribute":$addressSpace, "Value":$input)> + ]; + let assemblyFormat = "$input attr-dict `:` type($input) `to` type($result)"; + let hasCanonicalizeMethod = 1; +} + +//===----------------------------------------------------------------------===// +// CastIntOp +//===----------------------------------------------------------------------===// + +def Address_CastIntOp : Address_Op<"cast_int", [ + Pure, DeclareOpInterfaceMethods]> { + let summary = "Creates an int <-> address cast."; + let description = [{ + The `addr.cast_int` operation casts an int or index value to an address and + vice-versa. + Example: + + ```mlir + %addr = addr.cast_int %int : i32 to !address<1 : i32> + %index = addr.cast_int %addr : !address<1 : i32> to index + ``` + }]; + let arguments = (ins AnyTypeOf<[AnySignlessIntegerOrIndex, AddressType]>:$input); + let results = (outs AnyTypeOf<[AnySignlessIntegerOrIndex, AddressType]>:$result); + let builders = [ + OpBuilder<(ins "Value":$input, CArg<"Type", "nullptr">:$resultTy)> + ]; + let assemblyFormat = "$input attr-dict `:` type($input) `to` type($result)"; +} + +//===----------------------------------------------------------------------===// +// FromMemrefOp +//===----------------------------------------------------------------------===// + +def Address_FromMemRefOp : Address_Op<"from_memref", [Pure]> { + let summary = "Converts a memref to an address."; + let description = [{ + The `addr.from_memref` operation extracts the aligned or allocated address + from a memref. + Example: + + ```mlir + %addr = addr.from_memref [%memref : memref<2x4x?xf32>] + %baseAddr = addr.from_memref extract_base [%memref : memref<1xf32, 3>] + ``` + }]; + let arguments = (ins AnyMemRef:$input, UnitAttr:$extract_base); + let results = (outs AddressType:$result); + let assemblyFormat = [{ + (`extract_base` $extract_base^)? ` ` `[` $input + custom(type($input), type($result)) `]` attr-dict + }]; + let hasVerifier = 1; + let hasCanonicalizeMethod = 1; +} + +def Address_FromUnrankedMemRefOp : Address_Op<"from_unranked_memref", [Pure]> { + let summary = "Converts a unranked memref to an address."; + let description = [{ + The `addr.from_unranked_memref` operation extracts the aligned or allocated address + from a memref. + Example: + + ```mlir + %addr = addr.from_unranked_memref [%memref : memref<*xf32>] + %baseAddr = addr.from_unranked_memref extract_base [%memref : memref<*xf32, 3>] + ``` + }]; + let arguments = (ins AnyUnrankedMemRef:$input, UnitAttr:$extract_base); + let results = (outs AddressType:$result); + let assemblyFormat = [{ + (`extract_base` $extract_base^)? ` ` `[` $input + custom(type($input), type($result)) `]` attr-dict + }]; +} + +//===----------------------------------------------------------------------===// +// ToMemRefOp +//===----------------------------------------------------------------------===// + +def Address_ToMemRefOp : Address_Op<"to_memref", [Pure]> { + let summary = "Converts an address to a memref."; + let description = [{ + The `addr.to_memref` operation converts an `address` to a `memref` with + static shape. + Example: + + ```mlir + %memref = addr.to_memref %addr : memref<2x4xf32> + %memrefFull = addr.to_memref %addr base %baseAddr : memref<2xf32> + ``` + }]; + let arguments = (ins AddressType:$address, Optional:$base); + let results = (outs AnyStaticShapeMemRef:$result); + let assemblyFormat = [{ + $address attr-dict + custom($base, type($base), type($address), type($result)) + }]; + let builders = [ + OpBuilder<(ins "MemRefType":$type, "Value":$address)> + ]; + let hasVerifier = 1; + let hasCanonicalizeMethod = 1; +} + +def Address_ToUnrankedMemRefOp : Address_Op<"to_unranked_memref", [Pure]> { + let summary = "Converts an address to an unranked memref."; + let description = [{ + The `addr.to_unranked_memref` operation converts an `address` to an unranked `memref`. + Example: + + ```mlir + %memref = addr.to_unranked_memref %addr : memref<*xf32> + %memrefFull = addr.to_unranked_memref %addr base %baseAddr : memref<*xf32> + ``` + }]; + let arguments = (ins AddressType:$address, Optional:$base); + let results = (outs AnyUnrankedMemRef:$result); + let assemblyFormat = [{ + $address attr-dict + custom($base, type($base), type($address), type($result)) + }]; + let builders = [ + OpBuilder<(ins "UnrankedMemRefType":$type, "Value":$address)> + ]; +} + +//===----------------------------------------------------------------------===// +// PtrAddOp +//===----------------------------------------------------------------------===// + +def Address_PtrAddOp : Address_Op<"ptradd", [ + Pure, AllTypesMatch<["base", "result"]> + ]> { + let summary = "Adds an int or index to a addres."; + let description = [{ + The `addr.ptradd` operation adds an `address` and an integer or index to + produce a new address. + Example: + ```mlir + %addr = addr.ptradd %addr : !addr.address<3 : i32>, %c10 : i32 + ``` + }]; + + let arguments = (ins AddressType:$base, AnySignlessIntegerOrIndex:$offset); + let results = (outs AddressType:$result); + + let assemblyFormat = [{ + $base custom(type($base)) `,` $offset + custom(type($offset)) attr-dict + }]; +} + +#endif // ADDRESS_OPS diff --git a/third_party/wafer/include/Address/Dialect/IR/CMakeLists.txt b/third_party/wafer/include/Address/Dialect/IR/CMakeLists.txt new file mode 100755 index 00000000..7e42f4e7 --- /dev/null +++ b/third_party/wafer/include/Address/Dialect/IR/CMakeLists.txt @@ -0,0 +1,20 @@ +add_mlir_dialect(AddressOps addr) + +set(LLVM_TARGET_DEFINITIONS AddressBase.td) +mlir_tablegen(AddressOpInterfaces.h.inc -gen-op-interface-decls) +mlir_tablegen(AddressOpInterfaces.cpp.inc -gen-op-interface-defs) +add_public_tablegen_target(MLIRAddressOpInterfacesIncGen) + +set(LLVM_TARGET_DEFINITIONS AddressOps.td) +mlir_tablegen(AddressOpsEnums.h.inc -gen-enum-decls) +mlir_tablegen(AddressOpsEnums.cpp.inc -gen-enum-defs) +add_public_tablegen_target(MLIRAddressOpsEnumsGen) + +set(LLVM_TARGET_DEFINITIONS AddressOps.td) +mlir_tablegen( + AddressOpsAttributes.h.inc -gen-attrdef-decls -attrdefs-dialect=addr +) +mlir_tablegen( + AddressOpsAttributes.cpp.inc -gen-attrdef-defs -attrdefs-dialect=addr +) +add_public_tablegen_target(MLIRAddressOpsAttributesIncGen) diff --git a/third_party/wafer/include/Address/Transforms/CMakeLists.txt b/third_party/wafer/include/Address/Transforms/CMakeLists.txt new file mode 100755 index 00000000..f235eb5b --- /dev/null +++ b/third_party/wafer/include/Address/Transforms/CMakeLists.txt @@ -0,0 +1,5 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls -name Address) +mlir_tablegen(Passes.capi.h.inc -gen-pass-capi-header --prefix Address) +mlir_tablegen(Passes.capi.cpp.inc -gen-pass-capi-impl --prefix Address) +add_public_tablegen_target(MLIRAddressPassIncGen) diff --git a/third_party/wafer/include/Address/Transforms/Passes.h b/third_party/wafer/include/Address/Transforms/Passes.h new file mode 100755 index 00000000..00ee97dd --- /dev/null +++ b/third_party/wafer/include/Address/Transforms/Passes.h @@ -0,0 +1,33 @@ +//===- Passes.h - Address passes -------------------------------*- C++ -*-===// +// +// This file is licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This header file defines prototypes that expose pass constructors. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_DIALECT_ADDRESS_TRANSFORMS_PASSES_H +#define MLIR_DIALECT_ADDRESS_TRANSFORMS_PASSES_H + +#include "Address/Dialect/IR/AddressDialect.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Pass/Pass.h" +#include + +namespace mlir { +namespace addr { +void populateBarePtrConvetion(RewritePatternSet &patterns); + +#define GEN_PASS_DECL +#include "Address/Transforms/Passes.h.inc" + +#define GEN_PASS_REGISTRATION +#include "Address/Transforms/Passes.h.inc" +} // namespace addr +} // namespace mlir + +#endif // MLIR_DIALECT_ADDRESS_TRANSFORMS_PASSES_H diff --git a/third_party/wafer/include/Address/Transforms/Passes.td b/third_party/wafer/include/Address/Transforms/Passes.td new file mode 100755 index 00000000..d0d37c68 --- /dev/null +++ b/third_party/wafer/include/Address/Transforms/Passes.td @@ -0,0 +1,44 @@ +//===- StandalonePsss.td - Standalone dialect passes -------*- tablegen -*-===// +// +// This file is licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_DIALECT_ADDRESS_PASSES +#define MLIR_DIALECT_ADDRESS_PASSES + +include "mlir/Pass/PassBase.td" + +def AddrBarePtrConvention : Pass<"addr-apply-bare-ptr"> { + let summary = "Applies the bare pointer convention"; + let description = [{ + This pass applies the bare pointer convention, transforming `memref`s with + static shape in function parameters, results, and in call arguments to + `address` values, reconstructing the `memref` appropriately. + ```mlir + func.func @bar(%arg : memref) -> memref { + %memref = func.call @foo(%arg) : (memref) -> memref + return %memref : memref + } + // Gets transformed to: + func.func @bar(%arg0: !addr.address) -> !addr.address { + %0 = addr.to_memref %arg0 : memref + %1 = addr.from_memref [%0 : memref] + %2 = call @foo(%1) : (!addr.address) -> !addr.address + %3 = addr.to_memref %2 : memref + %4 = addr.from_memref [%3 : memref] + return %4 : !addr.address + } + + ``` + }]; +} + +def AddrToLLVM : Pass<"addr-to-llvm", "::mlir::ModuleOp"> { + let summary = "Convert from the Address dialect to the LLVM dialect"; + let dependentDialects = ["LLVM::LLVMDialect"]; +} + +#endif // MLIR_DIALECT_ADDRESS_PASSES diff --git a/third_party/wafer/include/Analysis/Alias.h b/third_party/wafer/include/Analysis/Alias.h new file mode 100755 index 00000000..b5efc70f --- /dev/null +++ b/third_party/wafer/include/Analysis/Alias.h @@ -0,0 +1,96 @@ +#ifndef TRITON_ANALYSIS_ALIAS_H +#define TRITON_ANALYSIS_ALIAS_H + +#include "mlir/Analysis/AliasAnalysis.h" +#include "mlir/Analysis/DataFlow/SparseAnalysis.h" +#include "llvm/ADT/DenseSet.h" + +namespace mlir::triton::alias { + +class AliasInfo { +public: + AliasInfo() = default; + AliasInfo(Value value) { insert(value); } + + void insert(Value value) { allocs.insert(value); } + + const DenseSet &getAllocs() const { return allocs; } + + bool operator==(const AliasInfo &other) const { + return allocs == other.allocs; + } + + /// The pessimistic value state of a value without alias + static AliasInfo getPessimisticValueState(MLIRContext *context = nullptr) { + return AliasInfo(); + } + static AliasInfo getPessimisticValueState(Value value) { return AliasInfo(); } + + /// The union of both arguments + static AliasInfo join(const AliasInfo &lhs, const AliasInfo &rhs); + + void print(raw_ostream &os) const { + llvm::interleaveComma(allocs, os, [&](Value alloc) { alloc.print(os); }); + } + +private: + /// The set of allocated values that are aliased by this lattice. + /// For now, we only consider aliased value produced by the following + /// situations: + /// 1. values returned by scf.yield + /// 2. block arguments in scf.for + /// Example: + /// alloc v1 alloc v2 + /// | | + /// |--------------| |------------| + /// scf.for v3 scf.for v4 scf.for v5 + /// | + /// scf.yield v6 + /// + /// v1's alloc [v1] + /// v2's alloc [v2] + /// v3's alloc [v1] + /// v4's alloc [v1, v2] + /// v5's alloc [v2] + /// v6's alloc [v1] + /// + /// Therefore, v1's liveness range is the union of v3, v4, and v6 + /// v2's liveness range is the union of v4 and v5. + DenseSet allocs; +}; + +//===----------------------------------------------------------------------===// +// Shared Memory Alias Analysis +//===----------------------------------------------------------------------===// +class SharedMemoryAliasAnalysis + : public dataflow::SparseForwardDataFlowAnalysis< + dataflow::Lattice> { +public: + using dataflow::SparseForwardDataFlowAnalysis< + dataflow::Lattice>::SparseForwardDataFlowAnalysis; + using dataflow::SparseForwardDataFlowAnalysis< + dataflow::Lattice>::getLatticeElement; + + /// XXX(Keren): Compatible interface with MLIR AliasAnalysis for future use. + /// Given two values, returns their aliasing behavior. + AliasResult alias(Value lhs, Value rhs); + + /// Returns the modify-reference behavior of `op` on `location`. + ModRefResult getModRef(Operation *op, Value location); + + void setToEntryState(dataflow::Lattice *lattice) override { + propagateIfChanged(lattice, + lattice->join(AliasInfo::getPessimisticValueState( + lattice->getAnchor()))); + } + + /// Computes if the alloc set of the results are changed. + LogicalResult + visitOperation(Operation *op, + ArrayRef *> operands, + ArrayRef *> results) override; +}; + +} // namespace mlir::triton::alias + +#endif // TRITON_ANALYSIS_ALIAS_H diff --git a/third_party/wafer/include/Analysis/Allocation.h b/third_party/wafer/include/Analysis/Allocation.h new file mode 100755 index 00000000..e6b05fee --- /dev/null +++ b/third_party/wafer/include/Analysis/Allocation.h @@ -0,0 +1,219 @@ +#ifndef TRITON_ANALYSIS_ALLOCATION_H +#define TRITON_ANALYSIS_ALLOCATION_H + +#include "Analysis/Utility.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir::triton::alloc { + +class AllocationAnalysis; + +/// Modified from llvm-15.0: llvm/ADT/AddressRanges.h +/// A class that represents an interval, specified using a start and an end +/// values: [Start, End). +template class Interval { +public: + Interval() {} + Interval(T S, T E) : Start(S), End(E) { assert(Start <= End); } + T start() const { return Start; } + T end() const { return End; } + T size() const { return End - Start; } + bool contains(T Addr) const { return Start <= Addr && Addr < End; } + bool intersects(const Interval &R) const { + return Start < R.End && R.Start < End; + } + bool operator==(const Interval &R) const { + return Start == R.Start && End == R.End; + } + bool operator!=(const Interval &R) const { return !(*this == R); } + bool operator<(const Interval &R) const { + return std::make_pair(Start, End) < std::make_pair(R.Start, R.End); + } + +private: + T Start = std::numeric_limits::min(); + T End = std::numeric_limits::max(); +}; + +template Interval(T, T) -> Interval; + +class Allocation { +public: + /// A unique identifier for shared memory buffers + using BufferId = size_t; + using BufferIdSetT = DenseSet; + using FuncAllocMapT = CallGraph::FuncDataMapT; + + static constexpr BufferId InvalidBufferId = + std::numeric_limits::max(); + + Allocation() = default; + /// Creates a new Allocation analysis that computes the shared memory + /// information for all associated shared memory values. + explicit Allocation(Operation *operation) : operation(operation) {} + + /// Runs allocation analysis on the given top-level operation. + void run(FuncAllocMapT &funcAllocMap); + + /// Returns the operation this analysis was constructed from. + Operation *getOperation() const { return operation; } + + /// Returns the offset of the given buffer in the shared memory. + size_t getOffset(BufferId bufferId) const { + return bufferSet.at(bufferId).offset; + } + + /// Returns the size of the given buffer in the shared memory. + size_t getAllocatedSize(BufferId bufferId) const { + return bufferSet.at(bufferId).size; + } + + /// Returns the allocated interval of the given buffer. + Interval getAllocatedInterval(BufferId bufferId) const { + auto &buffer = bufferSet.at(bufferId); + return Interval(buffer.offset, buffer.offset + buffer.size); + } + + /// Returns the buffer id of the given value. + /// This interface only returns the allocated buffer id. + /// If you want to get all the buffer ids that are associated with the given + /// value, including alias buffers, use getBufferIds. + BufferId getBufferId(Value value) const { + if (valueBuffer.count(value)) { + return valueBuffer.lookup(value)->id; + } else { + return InvalidBufferId; + } + } + + /// Returns all the buffer ids of the given value, including alias buffers. + BufferIdSetT getBufferIds(Value value) const { + BufferIdSetT bufferIds; + auto allocBufferId = getBufferId(value); + if (allocBufferId != InvalidBufferId) + bufferIds.insert(allocBufferId); + for (auto *buffer : aliasBuffer.lookup(value)) { + if (buffer->id != InvalidBufferId) + bufferIds.insert(buffer->id); + } + return bufferIds; + } + + /// Returns the size of total shared memory allocated + size_t getSharedMemorySize() const { return sharedMemorySize; } + + /// Returns mapping from operation to list of live LDS buffers + std::map> getLiveBuffers(); + +private: + /// A class that represents a shared memory buffer + struct BufferT { + /// Explicit: memref.alloc + /// TODO: Others + enum class BufferKind { Explicit }; + + BufferKind kind; + BufferId id; + Operation *owner; + size_t size; + size_t alignment; + size_t offset; + + bool operator==(const BufferT &other) const { return id == other.id; } + bool operator<(const BufferT &other) const { return id < other.id; } + + BufferT(BufferKind kind, BufferId id, Operation *owner, size_t size, + size_t alignment = 4, size_t offset = 0) + : kind(kind), id(id), owner(owner), size(size), alignment(alignment), + offset(offset) {} + + size_t setOffsetAligned(size_t newOffset) { + return offset = llvm::alignTo(newOffset, alignment); + } + }; + + /// Value -> Explicit Buffer + using ValueBufferMapT = llvm::MapVector; + /// Value -> Alias Buffer + using AliasBufferMapT = llvm::MapVector>; + /// BufferId -> Buffer + using BufferSetT = std::map; + +private: + template + void addBuffer(KeyType &key, Args &&...args) { + BufferId nextId = bufferIdCounter++; + auto [it, inserted] = bufferSet.insert_or_assign( + nextId, BufferT(Kind, nextId, key, std::forward(args)...)); + BufferT *buffer = &it->second; + static_assert(Kind == BufferT::BufferKind::Explicit, + "Expected explicit buffer"); + valueBuffer[key] = buffer; + } + + void addAlias(Value value, Value alloc) { + aliasBuffer[value].insert(valueBuffer[alloc]); + } + +private: + Operation *operation = nullptr; + ValueBufferMapT valueBuffer; + AliasBufferMapT aliasBuffer; + BufferSetT bufferSet; + size_t sharedMemorySize = 0; + + size_t bufferIdCounter = 0; + + friend class triton::alloc::AllocationAnalysis; +}; + +/// Static analysis that computes the allocation of shared memory buffers +/// of the entire call graph. +/// The allocation is performed in a post-order walk of the call graph. +/// Each call op is treated like convert_layout that allocates a scratch buffer. +/// At each call, we compute the start offset of the scratch buffer and pass it +/// as an argument to the callee. +class ModuleAllocation : public CallGraph { +public: + using FuncOffsetMapT = DenseMap; + + ModuleAllocation(ModuleOp moduleOp) : CallGraph(moduleOp) { + walk( + // Pre-order edge walk callback + [](CallOpInterface callOp, FunctionOpInterface funcOp) {}, + // Post-order node walk callback + [&](FunctionOpInterface funcOp) { + auto [iter, inserted] = funcMap.try_emplace(funcOp, funcOp); + if (inserted) + iter->second.run(funcMap); + }); + } + + size_t getSharedMemorySize() { + size_t size = 0; + for (auto funcOp : getRoots()) { + auto *alloc = getFuncData(funcOp); + size = std::max(size, alloc->getSharedMemorySize()); + } + return size; + } + + size_t getSharedMemorySize(FunctionOpInterface funcOp) { + return getFuncData(funcOp)->getSharedMemorySize(); + } + + void setFunctionSharedMemoryValue(FunctionOpInterface funcOp, Value value) { + sharedMemoryValue[funcOp] = value; + } + + Value getFunctionSharedMemoryBase(FunctionOpInterface funcOp) { + return sharedMemoryValue[funcOp]; + } + +private: + FuncOffsetMapT sharedMemoryValue; +}; + +} // namespace mlir::triton::alloc + +#endif // TRITON_ANALYSIS_ALLOCATION_H diff --git a/third_party/wafer/include/Analysis/Membar.h b/third_party/wafer/include/Analysis/Membar.h new file mode 100644 index 00000000..893f0ea2 --- /dev/null +++ b/third_party/wafer/include/Analysis/Membar.h @@ -0,0 +1,121 @@ +#ifndef WAFER_MEMBAR_H +#define WAFER_MEMBAR_H + +#include "Analysis/Allocation.h" +#include "Analysis/Utility.h" + +#include +#include +#include + +namespace mlir { +class OpBuilder; +} + +namespace mlir::triton::membar { + +/// Which dependency shape is being checked between two intersecting ops. +enum class MembarHazardKind { + WriteRead, ///< writer op (lhs set) vs reader op (rhs set) + ReadWrite, ///< reader op (lhs set) vs writer op (rhs set) + WriteWrite ///< two writers +}; + +/// Return true to suppress a barrier between two ops even if their intervals +/// intersect. Operand roles follow `MembarHazardKind`. +using MembarFilterFn = + std::function; + +/// Ops that only compute addresses / metadata (memref shape, arith indices, +/// scf/cf control flow) and do not constitute a real SPM/DDR data access for +/// barrier insertion. `memref.load` / `memref.store` / `memref.copy` are not +/// included here. +bool isPureAddressOp(Operation *op); + +struct BlockInfo { + using IntervalT = triton::alloc::Interval; + using IntervalMapT = std::map>; + + IntervalMapT syncReadIntervals; + IntervalMapT syncWriteIntervals; + + BlockInfo &join(const BlockInfo &other) { + for (auto &interval : other.syncReadIntervals) + syncReadIntervals[interval.first].insert(interval.second.begin(), + interval.second.end()); + for (auto &interval : other.syncWriteIntervals) + syncWriteIntervals[interval.first].insert(interval.second.begin(), + interval.second.end()); + return *this; + } + + void sync() { + syncReadIntervals.clear(); + syncWriteIntervals.clear(); + } + + bool isIntersected(const BlockInfo &other, MembarFilterFn filter) const; + + bool operator==(const BlockInfo &other) const { + return syncReadIntervals == other.syncReadIntervals && + syncWriteIntervals == other.syncWriteIntervals; + } + bool operator!=(const BlockInfo &other) const { return !(*this == other); } +}; + +/// Membar-like analysis for Wafer IR using `mlir::triton::alloc::Allocation`. +/// Inserts `wafer::BarrierOp` as needed. +class MembarAnalysis { + using VirtualBlock = std::pair; + +public: + using FuncBlockInfoMapT = CallGraph::FuncDataMapT; + + MembarAnalysis() = default; + explicit MembarAnalysis(triton::alloc::Allocation *allocation, + MembarFilterFn filter = nullptr) + : allocation(allocation), filter(std::move(filter)) {} + + void run(FuncBlockInfoMapT &funcBlockInfoMap); + +private: + void resolve(FunctionOpInterface funcOp, FuncBlockInfoMapT *funcBlockInfoMap, + OpBuilder *builder); + void update(Operation *op, BlockInfo *blockInfo, + FuncBlockInfoMapT *funcBlockInfoMap, OpBuilder *builder); + void visitTerminator(Operation *op, SmallVector &successors); + void insertBarrier(Operation *op, OpBuilder *builder); + +private: + triton::alloc::Allocation *allocation = nullptr; + MembarFilterFn filter = nullptr; +}; + +class ModuleMembarAnalysis : public CallGraph { +public: + explicit ModuleMembarAnalysis(triton::alloc::ModuleAllocation *moduleAlloc, + MembarFilterFn filter = nullptr) + : CallGraph(moduleAlloc->getModuleOp()), + moduleAlloc(moduleAlloc), filter(std::move(filter)) {} + + void run() { + walk( + [](CallOpInterface callOp, FunctionOpInterface funcOp) {}, + [&](FunctionOpInterface funcOp) { + auto *alloc = moduleAlloc->getFuncData(funcOp); + auto [it, inserted] = funcMap.try_emplace(funcOp, BlockInfo()); + if (inserted && alloc) { + MembarAnalysis analysis(alloc, filter); + analysis.run(funcMap); + } + }); + } + +private: + triton::alloc::ModuleAllocation *moduleAlloc; + MembarFilterFn filter; +}; + +} // namespace mlir::triton::membar + +#endif // WAFER_MEMBAR_H diff --git a/third_party/wafer/include/Analysis/Utility.h b/third_party/wafer/include/Analysis/Utility.h new file mode 100755 index 00000000..ebc58a3b --- /dev/null +++ b/third_party/wafer/include/Analysis/Utility.h @@ -0,0 +1,154 @@ +#ifndef TRITON_ANALYSIS_UTILITY_H +#define TRITON_ANALYSIS_UTILITY_H + +#include "mlir/Analysis/DataFlowFramework.h" +#include "mlir/Support/LLVM.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { + +/// Create a basic DataFlowSolver with constant and dead code analysis included. +std::unique_ptr createDataFlowSolver(); + +/// This class represents a call graph for a given ModuleOp and holds +/// data of type T associated with each FunctionOpInterface. +template class CallGraph { +public: + using FuncDataMapT = DenseMap; + + /// Constructor that builds the call graph for the given moduleOp. + explicit CallGraph(ModuleOp moduleOp) : moduleOp(moduleOp) { build(); } + + /// Walks the call graph and applies the provided update functions + /// to the edges and nodes. + template + void walk(UpdateEdgeFn updateEdgeFn, UpdateNodeFn updateNodeFn) { + DenseSet visited; + for (auto root : roots) { + doWalk(root, visited, updateEdgeFn, + updateNodeFn); + } + } + + /// Retrieves the data associated with a function + T *getFuncData(FunctionOpInterface funcOp) { + if (funcMap.count(funcOp)) { + return &funcMap[funcOp]; + } + return nullptr; + } + + /// Getters + ModuleOp getModuleOp() const { return moduleOp; } + SmallVector getRoots() const { return roots; } + size_t getNumFunctions() const { return funcMap.size(); } + + /// Returns true if the given function is a root. + bool isRoot(FunctionOpInterface funcOp) const { + return llvm::is_contained(roots, funcOp); + } + + /// Maps the data and the graph nodes associated with a funcOp to a + /// targetFuncOp. + template + void mapFuncOp(FROM funcOp, TO targetFuncOp) { + // Iterate over graph and replace + for (auto &kv : graph) { + for (auto &edge : kv.second) { + if (edge.second == funcOp) { + edge.second = targetFuncOp; + } + } + } + graph[targetFuncOp] = graph[funcOp]; + // Replace in roots + for (auto it = roots.begin(); it != roots.end(); ++it) { + if (*it == funcOp) { + *it = targetFuncOp; + break; + } + } + // Replace in funcMap + funcMap[targetFuncOp] = funcMap[funcOp]; + } + + /// Maps the graph edges associated with a callOp to a targetCallOp. + template + void mapCallOp(FROM callOp, TO targetCallOp) { + // Iterate over graph and replace + for (auto &kv : graph) { + for (auto &edge : kv.second) { + if (edge.first == callOp) { + edge.first = targetCallOp; + } + } + } + } + +private: + void build() { + SymbolTableCollection symbolTable; + DenseSet visited; + // Build graph + moduleOp.walk([&](Operation *op) { + auto caller = op->getParentOfType(); + if (auto callOp = dyn_cast(op)) { + auto *callee = callOp.resolveCallableInTable(&symbolTable); + auto funcOp = dyn_cast_or_null(callee); + if (funcOp) { + graph[caller].emplace_back( + std::pair(callOp, funcOp)); + visited.insert(funcOp); + } + } + }); + // Find roots + moduleOp.walk([&](FunctionOpInterface funcOp) { + if (!visited.count(funcOp)) { + roots.push_back(funcOp); + } + }); + } + + template + void doWalk(FunctionOpInterface funcOp, + DenseSet &visited, UpdateEdgeFn updateEdgeFn, + UpdateNodeFn updateNodeFn) { + if (visited.count(funcOp)) { + llvm::report_fatal_error("Cycle detected in call graph"); + } + if constexpr (UpdateNodeOrder == WalkOrder::PreOrder) { + updateNodeFn(funcOp); + } + for (auto [callOp, callee] : graph[funcOp]) { + if constexpr (UpdateEdgeOrder == WalkOrder::PreOrder) { + updateEdgeFn(callOp, callee); + } + doWalk(callee, visited, updateEdgeFn, + updateNodeFn); + if constexpr (UpdateEdgeOrder == WalkOrder::PostOrder) { + updateEdgeFn(callOp, callee); + } + } + if constexpr (UpdateNodeOrder == WalkOrder::PostOrder) { + updateNodeFn(funcOp); + } + visited.erase(funcOp); + } + +protected: + ModuleOp moduleOp; + DenseMap>> + graph; + FuncDataMapT funcMap; + SmallVector roots; +}; + +} // namespace mlir + +#endif // TRITON_ANALYSIS_UTILITY_H diff --git a/third_party/wafer/include/CMakeLists.txt b/third_party/wafer/include/CMakeLists.txt new file mode 100755 index 00000000..fd15195d --- /dev/null +++ b/third_party/wafer/include/CMakeLists.txt @@ -0,0 +1,7 @@ +add_subdirectory(triton-shared) +add_subdirectory(Address) +add_subdirectory(magic-kernel) +add_subdirectory(wafer) +# The following 2 dialects are currently unused. +#add_subdirectory(magic-kernel-func) +#add_subdirectory(magic-kernel-instr) diff --git a/third_party/wafer/include/ExecutionEngine/CRunnerUtils.cpp b/third_party/wafer/include/ExecutionEngine/CRunnerUtils.cpp new file mode 100755 index 00000000..87e47027 --- /dev/null +++ b/third_party/wafer/include/ExecutionEngine/CRunnerUtils.cpp @@ -0,0 +1,189 @@ +//===- CRunnerUtils.cpp - Utils for MLIR execution ------------------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file implements basic functions to manipulate structured MLIR types at +// runtime. Entities in this file are meant to be retargetable, including on +// targets without a C++ runtime, and must be kept C compatible. +// +//===----------------------------------------------------------------------===// + +#include "CRunnerUtils.h" +#include "Msan.h" + +#ifndef _WIN32 +#if defined(__FreeBSD__) || defined(__NetBSD__) || defined(__OpenBSD__) || \ + defined(__DragonFly__) +#include +#else +#include +#endif +#include +#else +#include "malloc.h" +#endif // _WIN32 + +#include +#include +#include +#include +#include +#include + +#ifdef MLIR_CRUNNERUTILS_DEFINE_FUNCTIONS + +namespace { +template void stdSort(uint64_t n, V *p) { std::sort(p, p + n); } + +} // namespace + +// Small runtime support "lib" for vector.print lowering. +// By providing elementary printing methods only, this +// library can remain fully unaware of low-level implementation +// details of our vectors. Also useful for direct LLVM IR output. +extern "C" void printI64(int64_t i) { fprintf(stdout, "%" PRId64, i); } +extern "C" void printU64(uint64_t u) { fprintf(stdout, "%" PRIu64, u); } +extern "C" void printF32(float f) { fprintf(stdout, "%g", f); } +extern "C" void printF64(double d) { fprintf(stdout, "%lg", d); } +extern "C" void printString(char const *s) { fputs(s, stdout); } +extern "C" void printOpen() { fputs("( ", stdout); } +extern "C" void printClose() { fputs(" )", stdout); } +extern "C" void printComma() { fputs(", ", stdout); } +extern "C" void printNewline() { fputc('\n', stdout); } + +extern "C" void memrefCopy(int64_t elemSize, UnrankedMemRefType *srcArg, + UnrankedMemRefType *dstArg) { + DynamicMemRefType src(*srcArg); + DynamicMemRefType dst(*dstArg); + + int64_t rank = src.rank; + MLIR_MSAN_MEMORY_IS_INITIALIZED(src.sizes, rank * sizeof(int64_t)); + + // Handle empty shapes -> nothing to copy. + for (int rankp = 0; rankp < rank; ++rankp) + if (src.sizes[rankp] == 0) + return; + + char *srcPtr = src.data + src.offset * elemSize; + char *dstPtr = dst.data + dst.offset * elemSize; + + if (rank == 0) { + memcpy(dstPtr, srcPtr, elemSize); + return; + } + + int64_t *indices = static_cast(alloca(sizeof(int64_t) * rank)); + int64_t *srcStrides = static_cast(alloca(sizeof(int64_t) * rank)); + int64_t *dstStrides = static_cast(alloca(sizeof(int64_t) * rank)); + + // Initialize index and scale strides. + for (int rankp = 0; rankp < rank; ++rankp) { + indices[rankp] = 0; + srcStrides[rankp] = src.strides[rankp] * elemSize; + dstStrides[rankp] = dst.strides[rankp] * elemSize; + } + + int64_t readIndex = 0, writeIndex = 0; + for (;;) { + // Copy over the element, byte by byte. + memcpy(dstPtr + writeIndex, srcPtr + readIndex, elemSize); + // Advance index and read position. + for (int64_t axis = rank - 1; axis >= 0; --axis) { + // Advance at current axis. + auto newIndex = ++indices[axis]; + readIndex += srcStrides[axis]; + writeIndex += dstStrides[axis]; + // If this is a valid index, we have our next index, so continue copying. + if (src.sizes[axis] != newIndex) + break; + // We reached the end of this axis. If this is axis 0, we are done. + if (axis == 0) + return; + // Else, reset to 0 and undo the advancement of the linear index that + // this axis had. Then continue with the axis one outer. + indices[axis] = 0; + readIndex -= src.sizes[axis] * srcStrides[axis]; + writeIndex -= dst.sizes[axis] * dstStrides[axis]; + } + } +} + +/// Prints GFLOPS rating. +extern "C" void printFlops(double flops) { + fprintf(stderr, "%lf GFLOPS\n", flops / 1.0E9); +} + +/// Returns the number of seconds since Epoch 1970-01-01 00:00:00 +0000 (UTC). +extern "C" double rtclock() { +#ifndef _WIN32 + struct timeval tp; + int stat = gettimeofday(&tp, nullptr); + if (stat != 0) + fprintf(stderr, "Error returning time from gettimeofday: %d\n", stat); + return (tp.tv_sec + tp.tv_usec * 1.0e-6); +#else + fprintf(stderr, "Timing utility not implemented on Windows\n"); + return 0.0; +#endif // _WIN32 +} + +extern "C" void *mlirAlloc(uint64_t size) { return malloc(size); } + +extern "C" void *mlirAlignedAlloc(uint64_t alignment, uint64_t size) { +#ifdef _WIN32 + return _aligned_malloc(size, alignment); +#elif defined(__APPLE__) + // aligned_alloc was added in MacOS 10.15. Fall back to posix_memalign to also + // support older versions. + void *result = nullptr; + (void)::posix_memalign(&result, alignment, size); + return result; +#else + return aligned_alloc(alignment, size); +#endif +} + +extern "C" void mlirFree(void *ptr) { free(ptr); } + +extern "C" void mlirAlignedFree(void *ptr) { +#ifdef _WIN32 + _aligned_free(ptr); +#else + free(ptr); +#endif +} + +extern "C" void *rtsrand(uint64_t s) { + // Standard mersenne_twister_engine seeded with s. + return new std::mt19937(s); +} + +extern "C" uint64_t rtrand(void *g, uint64_t m) { + std::mt19937 *generator = static_cast(g); + std::uniform_int_distribution distrib(0, m); + return distrib(*generator); +} + +extern "C" void rtdrand(void *g) { + std::mt19937 *generator = static_cast(g); + delete generator; +} + +#define IMPL_STDSORT(VNAME, V) \ + extern "C" void _mlir_ciface_stdSort##VNAME(uint64_t n, \ + StridedMemRefType *vref) { \ + assert(vref); \ + assert(vref->strides[0] == 1); \ + V *values = vref->data + vref->offset; \ + stdSort(n, values); \ + } +IMPL_STDSORT(I64, int64_t) +IMPL_STDSORT(F64, double) +IMPL_STDSORT(F32, float) +#undef IMPL_STDSORT + +#endif // MLIR_CRUNNERUTILS_DEFINE_FUNCTIONS diff --git a/third_party/wafer/include/ExecutionEngine/CRunnerUtils.h b/third_party/wafer/include/ExecutionEngine/CRunnerUtils.h new file mode 100755 index 00000000..1e55ca92 --- /dev/null +++ b/third_party/wafer/include/ExecutionEngine/CRunnerUtils.h @@ -0,0 +1,482 @@ +//===- CRunnerUtils.h - Utils for debugging MLIR execution ----------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file declares basic classes and functions to manipulate structured MLIR +// types at runtime. Entities in this file must be compliant with C++11 and be +// retargetable, including on targets without a C++ runtime. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_EXECUTIONENGINE_CRUNNERUTILS_H +#define MLIR_EXECUTIONENGINE_CRUNNERUTILS_H + +#ifdef _WIN32 +#ifndef MLIR_CRUNNERUTILS_EXPORT +#ifdef mlir_c_runner_utils_EXPORTS +// We are building this library +#define MLIR_CRUNNERUTILS_EXPORT __declspec(dllexport) +#define MLIR_CRUNNERUTILS_DEFINE_FUNCTIONS +#else +// We are using this library +#define MLIR_CRUNNERUTILS_EXPORT __declspec(dllimport) +#endif // mlir_c_runner_utils_EXPORTS +#endif // MLIR_CRUNNERUTILS_EXPORT +#else // _WIN32 +// Non-windows: use visibility attributes. +#define MLIR_CRUNNERUTILS_EXPORT __attribute__((visibility("default"))) +#define MLIR_CRUNNERUTILS_DEFINE_FUNCTIONS +#endif // _WIN32 + +#include +#include +#include +#include +#include + +//===----------------------------------------------------------------------===// +// Codegen-compatible structures for Vector type. +//===----------------------------------------------------------------------===// +namespace mlir { +namespace detail { + +constexpr bool isPowerOf2(int n) { return (!(n & (n - 1))); } + +constexpr unsigned nextPowerOf2(int n) { + return (n <= 1) ? 1 : (isPowerOf2(n) ? n : (2 * nextPowerOf2((n + 1) / 2))); +} + +template struct Vector1D; + +template struct Vector1D { + Vector1D() { + static_assert(detail::nextPowerOf2(sizeof(T[Dim])) == sizeof(T[Dim]), + "size error"); + } + inline T &operator[](unsigned i) { return vector[i]; } + inline const T &operator[](unsigned i) const { return vector[i]; } + +private: + T vector[Dim]; +}; + +// 1-D vector, padded to the next power of 2 allocation. +// Specialization occurs to avoid zero size arrays (which fail in -Werror). +template struct Vector1D { + Vector1D() { + static_assert(nextPowerOf2(sizeof(T[Dim])) > sizeof(T[Dim]), "size error"); + static_assert(nextPowerOf2(sizeof(T[Dim])) < 2 * sizeof(T[Dim]), + "size error"); + } + inline T &operator[](unsigned i) { return vector[i]; } + inline const T &operator[](unsigned i) const { return vector[i]; } + +private: + T vector[Dim]; + char padding[nextPowerOf2(sizeof(T[Dim])) - sizeof(T[Dim])]; +}; +} // namespace detail +} // namespace mlir + +// N-D vectors recurse down to 1-D. +template struct Vector { + inline Vector &operator[](unsigned i) { return vector[i]; } + inline const Vector &operator[](unsigned i) const { + return vector[i]; + } + +private: + Vector vector[Dim]; +}; + +// 1-D vectors in LLVM are automatically padded to the next power of 2. +// We insert explicit padding in to account for this. +template +struct Vector + : public mlir::detail::Vector1D { +}; + +template using Vector1D = Vector; +template using Vector2D = Vector; +template +using Vector3D = Vector; +template +using Vector4D = Vector; + +template void dropFront(int64_t arr[N], int64_t *res) { + for (unsigned i = 1; i < N; ++i) + *(res + i - 1) = arr[i]; +} + +//===----------------------------------------------------------------------===// +// Codegen-compatible structures for StridedMemRef type. +//===----------------------------------------------------------------------===// +template class StridedMemrefIterator; + +/// StridedMemRef descriptor type with static rank. +template struct StridedMemRefType { + T *basePtr; + T *data; + int64_t offset; + int64_t sizes[N]; + int64_t strides[N]; + + template ().begin())> + T &operator[](Range &&indices) { + assert(indices.size() == N && + "indices should match rank in memref subscript"); + int64_t curOffset = offset; + for (int dim = N - 1; dim >= 0; --dim) { + int64_t currentIndex = *(indices.begin() + dim); + assert(currentIndex < sizes[dim] && "Index overflow"); + curOffset += currentIndex * strides[dim]; + } + return data[curOffset]; + } + + StridedMemrefIterator begin() { return {*this, offset}; } + StridedMemrefIterator end() { return {*this, -1}; } + + // This operator[] is extremely slow and only for sugaring purposes. + StridedMemRefType operator[](int64_t idx) { + StridedMemRefType res; + res.basePtr = basePtr; + res.data = data; + res.offset = offset + idx * strides[0]; + dropFront(sizes, res.sizes); + dropFront(strides, res.strides); + return res; + } +}; + +/// StridedMemRef descriptor type specialized for rank 1. +template struct StridedMemRefType { + T *basePtr; + T *data; + int64_t offset; + int64_t sizes[1]; + int64_t strides[1]; + + template ().begin())> + T &operator[](Range indices) { + assert(indices.size() == 1 && + "indices should match rank in memref subscript"); + return (*this)[*indices.begin()]; + } + + StridedMemrefIterator begin() { return {*this, offset}; } + StridedMemrefIterator end() { return {*this, -1}; } + + T &operator[](int64_t idx) { return *(data + offset + idx * strides[0]); } +}; + +/// StridedMemRef descriptor type specialized for rank 0. +template struct StridedMemRefType { + T *basePtr; + T *data; + int64_t offset; + + template ().begin())> + T &operator[](Range indices) { + assert((indices.size() == 0) && + "Expect empty indices for 0-rank memref subscript"); + return data[offset]; + } + + StridedMemrefIterator begin() { return {*this, offset}; } + StridedMemrefIterator end() { return {*this, offset + 1}; } +}; + +/// Iterate over all elements in a strided memref. +template class StridedMemrefIterator { +public: + using iterator_category = std::forward_iterator_tag; + using value_type = T; + using difference_type = std::ptrdiff_t; + using pointer = T *; + using reference = T &; + + StridedMemrefIterator(StridedMemRefType &descriptor, + int64_t offset = 0) + : offset(offset), descriptor(&descriptor) {} + StridedMemrefIterator &operator++() { + int dim = Rank - 1; + while (dim >= 0 && indices[dim] == (descriptor->sizes[dim] - 1)) { + offset -= indices[dim] * descriptor->strides[dim]; + indices[dim] = 0; + --dim; + } + if (dim < 0) { + offset = -1; + return *this; + } + ++indices[dim]; + offset += descriptor->strides[dim]; + return *this; + } + + reference operator*() { return descriptor->data[offset]; } + pointer operator->() { return &descriptor->data[offset]; } + + const std::array &getIndices() { return indices; } + + bool operator==(const StridedMemrefIterator &other) const { + return other.offset == offset && other.descriptor == descriptor; + } + + bool operator!=(const StridedMemrefIterator &other) const { + return !(*this == other); + } + +private: + /// Offset in the buffer. This can be derived from the indices and the + /// descriptor. + int64_t offset = 0; + + /// Array of indices in the multi-dimensional memref. + std::array indices = {}; + + /// Descriptor for the strided memref. + StridedMemRefType *descriptor; +}; + +/// Iterate over all elements in a 0-ranked strided memref. +template class StridedMemrefIterator { +public: + using iterator_category = std::forward_iterator_tag; + using value_type = T; + using difference_type = std::ptrdiff_t; + using pointer = T *; + using reference = T &; + + StridedMemrefIterator(StridedMemRefType &descriptor, int64_t offset = 0) + : elt(descriptor.data + offset) {} + + StridedMemrefIterator &operator++() { + ++elt; + return *this; + } + + reference operator*() { return *elt; } + pointer operator->() { return elt; } + + // There are no indices for a 0-ranked memref, but this API is provided for + // consistency with the general case. + const std::array &getIndices() { + // Since this is a 0-array of indices we can keep a single global const + // copy. + static const std::array indices = {}; + return indices; + } + + bool operator==(const StridedMemrefIterator &other) const { + return other.elt == elt; + } + + bool operator!=(const StridedMemrefIterator &other) const { + return !(*this == other); + } + +private: + /// Pointer to the single element in the zero-ranked memref. + T *elt; +}; + +//===----------------------------------------------------------------------===// +// Codegen-compatible structure for UnrankedMemRef type. +//===----------------------------------------------------------------------===// +// Unranked MemRef +template struct UnrankedMemRefType { + int64_t rank; + void *descriptor; +}; + +//===----------------------------------------------------------------------===// +// DynamicMemRefType type. +//===----------------------------------------------------------------------===// +template class DynamicMemRefIterator; + +// A reference to one of the StridedMemRef types. +template class DynamicMemRefType { +public: + int64_t rank; + T *basePtr; + T *data; + int64_t offset; + const int64_t *sizes; + const int64_t *strides; + + explicit DynamicMemRefType(const StridedMemRefType &memRef) + : rank(0), basePtr(memRef.basePtr), data(memRef.data), + offset(memRef.offset), sizes(nullptr), strides(nullptr) {} + template + explicit DynamicMemRefType(const StridedMemRefType &memRef) + : rank(N), basePtr(memRef.basePtr), data(memRef.data), + offset(memRef.offset), sizes(memRef.sizes), strides(memRef.strides) {} + explicit DynamicMemRefType(const ::UnrankedMemRefType &memRef) + : rank(memRef.rank) { + auto *desc = static_cast *>(memRef.descriptor); + basePtr = desc->basePtr; + data = desc->data; + offset = desc->offset; + sizes = rank == 0 ? nullptr : desc->sizes; + strides = sizes + rank; + } + + template ().begin())> + T &operator[](Range &&indices) { + assert(indices.size() == rank && + "indices should match rank in memref subscript"); + if (rank == 0) + return data[offset]; + + int64_t curOffset = offset; + for (int dim = rank - 1; dim >= 0; --dim) { + int64_t currentIndex = *(indices.begin() + dim); + assert(currentIndex < sizes[dim] && "Index overflow"); + curOffset += currentIndex * strides[dim]; + } + return data[curOffset]; + } + + DynamicMemRefIterator begin() { return {*this, offset}; } + DynamicMemRefIterator end() { return {*this, -1}; } + + // This operator[] is extremely slow and only for sugaring purposes. + DynamicMemRefType operator[](int64_t idx) { + assert(rank > 0 && "can't make a subscript of a zero ranked array"); + + DynamicMemRefType res(*this); + --res.rank; + res.offset += idx * res.strides[0]; + ++res.sizes; + ++res.strides; + return res; + } + + // This operator* can be used in conjunction with the previous operator[] in + // order to access the underlying value in case of zero-ranked memref. + T &operator*() { + assert(rank == 0 && "not a zero-ranked memRef"); + return data[offset]; + } +}; + +/// Iterate over all elements in a dynamic memref. +template class DynamicMemRefIterator { +public: + using iterator_category = std::forward_iterator_tag; + using value_type = T; + using difference_type = std::ptrdiff_t; + using pointer = T *; + using reference = T &; + + DynamicMemRefIterator(DynamicMemRefType &descriptor, int64_t offset = 0) + : offset(offset), descriptor(&descriptor) { + indices.resize(descriptor.rank, 0); + } + + DynamicMemRefIterator &operator++() { + if (descriptor->rank == 0) { + offset = -1; + return *this; + } + + int dim = descriptor->rank - 1; + + while (dim >= 0 && indices[dim] == (descriptor->sizes[dim] - 1)) { + offset -= indices[dim] * descriptor->strides[dim]; + indices[dim] = 0; + --dim; + } + + if (dim < 0) { + offset = -1; + return *this; + } + + ++indices[dim]; + offset += descriptor->strides[dim]; + return *this; + } + + reference operator*() { return descriptor->data[offset]; } + pointer operator->() { return &descriptor->data[offset]; } + + const std::vector &getIndices() { return indices; } + + bool operator==(const DynamicMemRefIterator &other) const { + return other.offset == offset && other.descriptor == descriptor; + } + + bool operator!=(const DynamicMemRefIterator &other) const { + return !(*this == other); + } + +private: + /// Offset in the buffer. This can be derived from the indices and the + /// descriptor. + int64_t offset = 0; + + /// Array of indices in the multi-dimensional memref. + std::vector indices = {}; + + /// Descriptor for the dynamic memref. + DynamicMemRefType *descriptor; +}; + +//===----------------------------------------------------------------------===// +// Small runtime support library for memref.copy lowering during codegen. +//===----------------------------------------------------------------------===// +extern "C" MLIR_CRUNNERUTILS_EXPORT void +memrefCopy(int64_t elemSize, ::UnrankedMemRefType *src, + ::UnrankedMemRefType *dst); + +//===----------------------------------------------------------------------===// +// Small runtime support library for vector.print lowering during codegen. +//===----------------------------------------------------------------------===// +extern "C" MLIR_CRUNNERUTILS_EXPORT void printI64(int64_t i); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printU64(uint64_t u); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printF32(float f); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printF64(double d); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printString(char const *s); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printOpen(); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printClose(); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printComma(); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printNewline(); + +//===----------------------------------------------------------------------===// +// Small runtime support library for timing execution and printing GFLOPS +//===----------------------------------------------------------------------===// +extern "C" MLIR_CRUNNERUTILS_EXPORT void printFlops(double flops); +extern "C" MLIR_CRUNNERUTILS_EXPORT double rtclock(); + +//===----------------------------------------------------------------------===// +// Runtime support library for random number generation. +//===----------------------------------------------------------------------===// +// Uses a seed to initialize a random generator and returns the generator. +extern "C" MLIR_CRUNNERUTILS_EXPORT void *rtsrand(uint64_t s); +// Returns a random number in the range of [0, m). +extern "C" MLIR_CRUNNERUTILS_EXPORT uint64_t rtrand(void *, uint64_t m); +// Deletes the random number generator. +extern "C" MLIR_CRUNNERUTILS_EXPORT void rtdrand(void *); + +//===----------------------------------------------------------------------===// +// Runtime support library to allow the use of std::sort in MLIR program. +//===----------------------------------------------------------------------===// +extern "C" MLIR_CRUNNERUTILS_EXPORT void +_mlir_ciface_stdSortI64(uint64_t n, StridedMemRefType *vref); +extern "C" MLIR_CRUNNERUTILS_EXPORT void +_mlir_ciface_stdSortF64(uint64_t n, StridedMemRefType *vref); +extern "C" MLIR_CRUNNERUTILS_EXPORT void +_mlir_ciface_stdSortF32(uint64_t n, StridedMemRefType *vref); +#endif // MLIR_EXECUTIONENGINE_CRUNNERUTILS_H diff --git a/third_party/wafer/include/ExecutionEngine/Msan.h b/third_party/wafer/include/ExecutionEngine/Msan.h new file mode 100755 index 00000000..ee94660a --- /dev/null +++ b/third_party/wafer/include/ExecutionEngine/Msan.h @@ -0,0 +1,35 @@ +//===- Msan.h - Utils related to the memory sanitizer ---------------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file declares and defines macros related to msan. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_EXECUTIONENGINE_MSAN_H +#define MLIR_EXECUTIONENGINE_MSAN_H + +// Memory sanitizer currently can't be enabled for the jit-compiled code, and +// to suppress msan warnings we need to unpoison pointers and pointed-to +// datastructures before they can be accessed. + +#ifndef __has_feature +#define __has_feature(x) 0 +#endif + +#if __has_feature(memory_sanitizer) && !defined(MLIR_MEMORY_SANITIZER) +#define MLIR_MEMORY_SANITIZER +#endif + +#if defined(MLIR_MEMORY_SANITIZER) +#include +#define MLIR_MSAN_MEMORY_IS_INITIALIZED(p, s) __msan_unpoison((p), (s)) +#else // Memory sanitizer: OFF +#define MLIR_MSAN_MEMORY_IS_INITIALIZED(p, s) +#endif // MLIR_MEMORY_SANITIZER + +#endif // MLIR_EXECUTIONENGINE_MSAN_H diff --git a/third_party/wafer/include/ExecutionEngine/version.txt b/third_party/wafer/include/ExecutionEngine/version.txt new file mode 100755 index 00000000..c3f15e55 --- /dev/null +++ b/third_party/wafer/include/ExecutionEngine/version.txt @@ -0,0 +1 @@ +https://github.com/llvm/llvm-project/commit/3be3883e6d67bf908fd12b51219075293ebb3dff diff --git a/third_party/wafer/include/flagtree/Common/UnifiedHardware.h b/third_party/wafer/include/flagtree/Common/UnifiedHardware.h new file mode 100755 index 00000000..be2dc6f3 --- /dev/null +++ b/third_party/wafer/include/flagtree/Common/UnifiedHardware.h @@ -0,0 +1,31 @@ +#pragma once + +#include +#include +#include +#include +#include + +namespace mlir { +namespace flagtree { +// this is the unified hardware abstraction for hardware +// to determined if these abstraction is specified, using std::optional is +// needed using in passes: if(uh_flagtree->xxx()){...} + +class UnifiedHardware { + +public: + UnifiedHardware() = default; + virtual ~UnifiedHardware() = default; + virtual bool isRegistered() const; + virtual int getDMATag() const; + virtual int getSharedMemoryTag() const; + virtual bool getIncubatedTag() const; + virtual std::string getReduceStrategy() const; + virtual std::string getFlagTreeBackend() const; +}; + +std::unique_ptr createUnifiedHardwareManager(); + +} // namespace flagtree +} // namespace mlir diff --git a/third_party/wafer/include/magic-kernel-func/CMakeLists.txt b/third_party/wafer/include/magic-kernel-func/CMakeLists.txt new file mode 100755 index 00000000..629c08af --- /dev/null +++ b/third_party/wafer/include/magic-kernel-func/CMakeLists.txt @@ -0,0 +1,2 @@ +add_subdirectory(Conversion) +add_subdirectory(Dialect) diff --git a/third_party/wafer/include/magic-kernel-func/Dialect/CMakeLists.txt b/third_party/wafer/include/magic-kernel-func/Dialect/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/include/magic-kernel-func/Dialect/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/include/magic-kernel-func/Dialect/IR/MagicKernelFuncOps.td b/third_party/wafer/include/magic-kernel-func/Dialect/IR/MagicKernelFuncOps.td new file mode 100755 index 00000000..e930ab73 --- /dev/null +++ b/third_party/wafer/include/magic-kernel-func/Dialect/IR/MagicKernelFuncOps.td @@ -0,0 +1,19 @@ +//===------------------- MagicKernelFuncOps.td ----------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Common abstraction layer for non-instruction driven ML accelerator. +// +// The target glue layer that translates target independent kernel operations +// into NPU like APIs call (There are other jargons such as intrinsic or driver +// functions etc). +// +// The NPU APIs are categories by data type, like traditional compilers, integer +// and floating point function unit are separated, so for every MK(MagicKernel) +// op, it is lowered to 2 MKF(MagicKernelFunc) which are integer version and +// floating point version. +// +//===----------------------------------------------------------------------===// diff --git a/third_party/wafer/include/magic-kernel-instr/CMakeLists.txt b/third_party/wafer/include/magic-kernel-instr/CMakeLists.txt new file mode 100755 index 00000000..629c08af --- /dev/null +++ b/third_party/wafer/include/magic-kernel-instr/CMakeLists.txt @@ -0,0 +1,2 @@ +add_subdirectory(Conversion) +add_subdirectory(Dialect) diff --git a/third_party/wafer/include/magic-kernel-instr/Dialect/CMakeLists.txt b/third_party/wafer/include/magic-kernel-instr/Dialect/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/include/magic-kernel-instr/Dialect/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/include/magic-kernel-instr/Dialect/IR/MagicKernelInstrOps.td b/third_party/wafer/include/magic-kernel-instr/Dialect/IR/MagicKernelInstrOps.td new file mode 100755 index 00000000..2dbad73e --- /dev/null +++ b/third_party/wafer/include/magic-kernel-instr/Dialect/IR/MagicKernelInstrOps.td @@ -0,0 +1,13 @@ +//===------------------- MagicKernelInstrOps.td ---------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Common abstraction layer for instruction driven ML accelerator. +// +// The target glue layer that translates target independent kernel operations +// into intrinsics which fits LLVM dialect lowering path. +// +//===----------------------------------------------------------------------===// diff --git a/third_party/wafer/include/magic-kernel/CMakeLists.txt b/third_party/wafer/include/magic-kernel/CMakeLists.txt new file mode 100755 index 00000000..495310c6 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/CMakeLists.txt @@ -0,0 +1,3 @@ +add_subdirectory(Conversion) +add_subdirectory(Dialect) +add_subdirectory(Transforms) diff --git a/third_party/wafer/include/magic-kernel/Conversion/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Conversion/CMakeLists.txt new file mode 100755 index 00000000..eb67e697 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/CMakeLists.txt @@ -0,0 +1,6 @@ +add_subdirectory(LinalgToMK) +add_subdirectory(CoreDialectsToMK) +add_subdirectory(LegalizeTensorFormLoops) +add_subdirectory(TLEToMK) + +add_subdirectory(MKPipeline) diff --git a/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/CMakeLists.txt new file mode 100755 index 00000000..69690dec --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/CMakeLists.txt @@ -0,0 +1,10 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +# All rights reserved. +# +#===------------------------------------------------------------------------===# + +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name CoreDialectsMK) +add_public_tablegen_target(CoreDialectsToMKConversionPassIncGen) diff --git a/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/CoreDialectsToMK.h b/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/CoreDialectsToMK.h new file mode 100755 index 00000000..69750d40 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/CoreDialectsToMK.h @@ -0,0 +1,27 @@ +//===------------------- CoreDialectsToMK.h -------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This pass is the wrap all pass that populates all the conversion patterns +// from core dialects such as linalg, memref, buf etc to mk dialect. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_CORE_DIALECTS_TO_MK_H +#define TRITON_CONVERSION_CORE_DIALECTS_TO_MK_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createCoreDialectsToMKPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_CORE_DIALECTS_TO_MK_H diff --git a/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/Passes.h b/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/Passes.h new file mode 100755 index 00000000..7e1982c3 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/Passes.h @@ -0,0 +1,26 @@ +//===------------------- CoreDialectsToMK.h -------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Wrap all the conversion from core dialects to backend dialects(MK etc). +// +//===----------------------------------------------------------------------===// + +#ifndef CORE_DIALECTS_TO_MK_CONVERSION_PASSES_H +#define CORE_DIALECTS_TO_MK_CONVERSION_PASSES_H + +#include "magic-kernel/Conversion/CoreDialectsToMK/CoreDialectsToMK.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "magic-kernel/Conversion/CoreDialectsToMK/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // CORE_DIALECTS_TO_MK_CONVERSION_PASSES_H diff --git a/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/Passes.td b/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/Passes.td new file mode 100755 index 00000000..a2e137ed --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/Passes.td @@ -0,0 +1,24 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef CORE_DIALECTS_TO_MK_CONVERSION_PASSES +#define CORE_DIALECTS_TO_MK_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def CoreDialectsToMK : Pass<"core-dialects-to-mk", "mlir::ModuleOp"> { + let summary = "Convert core dialects including Linalg, Memref etc to MK"; + let constructor = "triton::createCoreDialectsToMKPass()"; + let options = [ + Option<"precisionPriority", "precision-priority", "bool", /*default*/"false", + "Enable precision priority mode, translating integer-related operations to RISC-V instructions (legacy alias for mode 2)">, + Option<"precisionMode", "precision-mode", "int", /*default*/"0", + "Integer precision: 0 permits fp32 lowering, 1 preserves i64, 2 preserves i32/i64 and division"> + ]; +} + +#endif diff --git a/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/CMakeLists.txt new file mode 100755 index 00000000..85245843 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name LegalizeTensorFormLoops) +add_public_tablegen_target(LegalizeTensorFormLoopsPassIncGen) diff --git a/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.h b/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.h new file mode 100755 index 00000000..f6017fa8 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.h @@ -0,0 +1,23 @@ +//===----------------------- Passes.h -------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef LEGALIZE_TENSOR_FORM_LOOPS_CONVERSION_PASSES_H +#define LEGALIZE_TENSOR_FORM_LOOPS_CONVERSION_PASSES_H + +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#define GEN_PASS_REGISTRATION +#include "magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // LEGALIZE_TENSOR_FORM_LOOPS_CONVERSION_PASSES_H diff --git a/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.td b/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.td new file mode 100755 index 00000000..ddceda79 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.td @@ -0,0 +1,17 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef LEGALIZE_TENSOR_FORM_LOOPS_CONVERSION_PASSES +#define LEGALIZE_TENSOR_FORM_LOOPS_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def LegalizeTensorFormLoops : Pass<"legalize-tensor-form-loops"> { + let summary = "Legalize tensor dependencies in the loop."; +} + +#endif diff --git a/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/CMakeLists.txt new file mode 100755 index 00000000..76b9d911 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name LinalgToMK) +add_public_tablegen_target(LinalgToMKConversionPassIncGen) diff --git a/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/LinalgToMK.h b/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/LinalgToMK.h new file mode 100755 index 00000000..c1786cd6 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/LinalgToMK.h @@ -0,0 +1,83 @@ +//===------------------- LinalgToMK.h -------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Lowering all linalg ops into mk ops. +// +//===----------------------------------------------------------------------===// + +#ifndef ZTC_CONVERSION_LINALG_TO_MK_H +#define ZTC_CONVERSION_LINALG_TO_MK_H + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton-shared/Utils/Utils.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "magic-kernel/Conversion/LinalgToMK/Passes.h.inc" + +// Fusion, identity reduction init, etc. Other ops which need to decompose into +// multiple integer type operators also are converted. +void populateLinalgToMKPreProcessPatterns(RewritePatternSet &patterns); + +// Type conversion: trans integer type to float type, etc. +void populateLinalgToMKTypeConversionPatterns(RewritePatternSet &patterns, + int precisionMode = 0); + +// Pattern rewrite to target dependent operation +void populateLinalgToMKCanonicalizationPatterns(RewritePatternSet &patterns, + int precisionMode = 0); + +// Reshape input shape to destination shape +void populateLinalgToMKShapeCanonicalizationPatterns( + RewritePatternSet &patterns, int precisionMode = 0); + +// Convertion patterns +void populateLinalgToMKConversionPatterns(RewritePatternSet &patterns); + +std::unique_ptr> createLinalgToMKPass(); +std::unique_ptr> +createLinalgToMKPass(LinalgToMKOptions &options); + +} // namespace triton +} // namespace mlir + +namespace { + +using namespace mlir; +using namespace triton; + +// Extract the operations from a linalg op region +template static bool checkGenericOp(linalg::GenericOp op) { + auto regionBlock = op.getBody(); + auto regionOps = llvm::map_to_vector(regionBlock->without_terminator(), + [](Operation &op) { return &op; }); + + return regionOps.size() == 1 && isa(regionOps[0]); +} + +static bool isConstantTensor(Value &v, double targetValue, + bool isApprox = false); + +// Check if the given value is a tensor filled with 0. +static bool isZeroTensor(Value &v); + +// Check if the given value is a tensor filled with 1. +static bool isOneTensor(Value &v); + +static bool isHalfTensor(Value &v); + +static bool isTwoTensor(Value &v); + +} // namespace + +#endif // ZTC_CONVERSION_MEMREF_TO_MAGICKERNEL_H diff --git a/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/Passes.h b/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/Passes.h new file mode 100755 index 00000000..7c45210e --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/Passes.h @@ -0,0 +1,22 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef LINALG_TO_MK_CONVERSION_PASSES_H +#define LINALG_TO_MK_CONVERSION_PASSES_H + +#include "magic-kernel/Conversion/LinalgToMK/LinalgToMK.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "magic-kernel/Conversion/LinalgToMK/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // LINALG_TO_MK_CONVERSION_PASSES_H diff --git a/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/Passes.td b/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/Passes.td new file mode 100755 index 00000000..2124025f --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/Passes.td @@ -0,0 +1,24 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef LINALG_TO_MK_CONVERSION_PASSES +#define LINALG_TO_MK_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def LinalgToMK : Pass<"linalg-to-mk", "mlir::ModuleOp"> { + let summary = "Convert linalg operations into magic kernel operations"; + + let options = [ + Option<"precisionPriority", "precision-priority", "bool", /*default*/"false", + "Enable precision priority mode, translating integer-related operations to RISC-V instructions (legacy alias for mode 2)">, + Option<"precisionMode", "precision-mode", "int", /*default*/"0", + "Integer precision: 0 permits fp32 lowering, 1 preserves i64, 2 preserves i32/i64 and division"> + ]; +} + +#endif diff --git a/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/CMakeLists.txt new file mode 100644 index 00000000..2dc17af8 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name MKPipeline) +add_public_tablegen_target(MKPipelinePassIncGen) diff --git a/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/Passes.h b/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/Passes.h new file mode 100644 index 00000000..c5b002ce --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/Passes.h @@ -0,0 +1,19 @@ +//===----------------------------------------------------------------------===// +// MKPipeline: software-pipeline scf.for loops with mk.dot / SPM buffers +//===----------------------------------------------------------------------===// +#ifndef MK_PIPELINE_PASSES_H +#define MK_PIPELINE_PASSES_H + +#include "mlir/Pass/Pass.h" + +namespace mlir::triton { + +#define GEN_PASS_DECL +#include "magic-kernel/Conversion/MKPipeline/Passes.h.inc" + +#define GEN_PASS_REGISTRATION +#include "magic-kernel/Conversion/MKPipeline/Passes.h.inc" + +} // namespace mlir::triton + +#endif // MK_PIPELINE_PASSES_H diff --git a/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/Passes.td b/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/Passes.td new file mode 100644 index 00000000..2c046ba5 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/MKPipeline/Passes.td @@ -0,0 +1,64 @@ +#ifndef MK_PIPELINE_PASSES +#define MK_PIPELINE_PASSES + +include "mlir/Pass/PassBase.td" + +def MKPipelinePass : Pass<"mk-pipeline", "mlir::ModuleOp"> { + let summary = "Software-pipeline loops with mk.dot using SPM multi-buffering"; + let description = [{ + Reads 'tt.num_stages' from scf.for loop attributes (written by the Triton + front-end via tl.range(..., num_stages=N)), or falls back to the global + num-stages option. + For each qualifying innermost scf.for that contains: + - at least one memref.copy (DDR -> SPM load) AND + - at least one mk.dot (compute) + the pass rewrites the loop into a software-pipelined form with + (num_stages - 1) prefetch buffers: + prologue: kick off first (num_stages-1) iterations of DDR->SPM copy + kernel: overlap next copy with current mk.dot + epilogue: drain remaining mk.dot after last copy + New memref.alloc ops for the extra SPM buffers are inserted BEFORE + spmd-allocate-shared-memory so that the existing allocator assigns + correct allocation.offset attributes automatically. + Barrier insertion: + The subsequent wafer-insert-barrier pass resolves SPM/DDR hazards + using the current SDK local completion barrier. + }]; + + let dependentDialects = ["mlir::mk::MagicKernelDialect", + "mlir::memref::MemRefDialect", + "mlir::scf::SCFDialect", + "mlir::arith::ArithDialect"]; + + let options = [ + Option<"numStages", "num-stages", "unsigned", /*default=*/"2u", + "Default software-pipeline stages when tt.num_stages loop attr is absent">, + Option<"maxStages", "max-stages", "unsigned", /*default=*/"2u", + "Upper clamp for effective stage count (loop attr and default)">, + ]; + +} + +def MKLoopBoundCanonicalizePass + : Pass<"mk-loop-bound-canonicalize", "mlir::ModuleOp"> { + let summary = "Canonicalize MK loop expressions using scf.for bounds"; + let description = [{ + Applies conservative loop-bound-aware rewrites after MKPipeline and + before MK-to-Wafer lowering. The initial pattern folds full-tile tail + size expressions of the form: + + size = max(min(iv + step, ub), iv) - iv + is_tail = size < step + + when the enclosing scf.for has constant bounds/step and every + iteration is known to satisfy iv + step <= ub. This exposes constant + copy sizes and false tail-fill guards to the following cse/canonicalize + passes. The pass is intentionally structured so more loop-bound + patterns can be added incrementally. + }]; + + let dependentDialects = ["mlir::scf::SCFDialect", + "mlir::arith::ArithDialect"]; +} + +#endif // MK_PIPELINE_PASSES diff --git a/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/CMakeLists.txt new file mode 100755 index 00000000..883e421e --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TLEToMK) +add_public_tablegen_target(TLEToMKConversionPassIncGen) diff --git a/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/Passes.h b/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/Passes.h new file mode 100755 index 00000000..3e95c8b1 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/Passes.h @@ -0,0 +1,22 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef TLE_TO_MK_CONVERSION_PASSES_H +#define TLE_TO_MK_CONVERSION_PASSES_H + +#include "magic-kernel/Conversion/TLEToMK/TLEToMK.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "magic-kernel/Conversion/TLEToMK/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // TLE_TO_MK_CONVERSION_PASSES_H diff --git a/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/Passes.td b/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/Passes.td new file mode 100755 index 00000000..d2b30228 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/Passes.td @@ -0,0 +1,19 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef TLE_TO_MK_CONVERSION_PASSES +#define TLE_TO_MK_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TLEToMK : Pass<"tle-to-mk", "mlir::ModuleOp"> { + let summary = "Convert TLE communication operations into magic kernel operations"; + + let options = []; +} + +#endif diff --git a/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/TLEToMK.h b/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/TLEToMK.h new file mode 100755 index 00000000..52185459 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Conversion/TLEToMK/TLEToMK.h @@ -0,0 +1,32 @@ +//===------------------- TLEToMK.h -------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Lowering TLE communication ops into mk ops. +// +//===----------------------------------------------------------------------===// + +#ifndef ZTC_CONVERSION_TLE_TO_MK_H +#define ZTC_CONVERSION_TLE_TO_MK_H + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "magic-kernel/Conversion/TLEToMK/Passes.h.inc" + +void populateTLEToMKConversionPatterns(RewritePatternSet &patterns); + +// std::unique_ptr> createTLEToMKPass(); + +} // namespace triton +} // namespace mlir + +#endif // ZTC_CONVERSION_TLE_TO_MK_H diff --git a/third_party/wafer/include/magic-kernel/Dialect/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Dialect/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Dialect/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/include/magic-kernel/Dialect/IR/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Dialect/IR/CMakeLists.txt new file mode 100755 index 00000000..437811f2 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Dialect/IR/CMakeLists.txt @@ -0,0 +1,11 @@ +set(LLVM_TARGET_DEFINITIONS MagicKernelOps.td) +mlir_tablegen(MagicKernelDialect.h.inc -gen-dialect-decls -dialect=mk) +mlir_tablegen(MagicKernelDialect.cpp.inc -gen-dialect-defs -dialect=mk) +mlir_tablegen(MagicKernelOps.h.inc -gen-op-decls) +mlir_tablegen(MagicKernelOps.cpp.inc -gen-op-defs) + +set(LLVM_TARGET_DEFINITIONS MagicKernelTypes.td) +mlir_tablegen(MagicKernelTypes.h.inc -gen-typedef-decls) +mlir_tablegen(MagicKernelTypes.cpp.inc -gen-typedef-defs) + +add_public_tablegen_target(MagicKernelTableGen) diff --git a/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelAttrDefs.td b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelAttrDefs.td new file mode 100755 index 00000000..666a7f41 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelAttrDefs.td @@ -0,0 +1,15 @@ +//===------------------- MagicKernelAttrDefs.td ---------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MAGIC_KERNEL_ATTR_DEFS +#define MAGIC_KERNEL_ATTR_DEFS + +include "mlir/IR/EnumAttr.td" + + + +#endif // MAGIC_KERNEL_ATTR_DEFS diff --git a/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelDialect.h b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelDialect.h new file mode 100755 index 00000000..06bd269a --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelDialect.h @@ -0,0 +1,32 @@ +//===------------------- MagicKernelDialect.h -----------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_DIALECT_MAGIC_KERNEL_IR_DIALECT_H_ +#define MLIR_DIALECT_MAGIC_KERNEL_IR_DIALECT_H_ + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/SymbolTable.h" +#include "mlir/IR/TypeSupport.h" +#include "mlir/IR/Types.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +//===----------------------------------------------------------------------===// +// MagicKernel Operations +//===----------------------------------------------------------------------===// +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h.inc" + +// Include the auto-generated header file containing the declarations of the +// TritonStructured operations. +#define GET_OP_CLASSES +#include "magic-kernel/Dialect/IR/MagicKernelOps.h.inc" + +#endif // MLIR_DIALECT_MAGIC_KERNEL_IR_DIALECT_H_ diff --git a/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelDialect.td b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelDialect.td new file mode 100755 index 00000000..4aee43bb --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelDialect.td @@ -0,0 +1,44 @@ +//===------------------- MagicKernelDialect.td ----------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MAGIC_KERNEL_DIALECT +#define MAGIC_KERNEL_DIALECT + +include "mlir/IR/OpBase.td" + +def MagicKernelDialect : Dialect { + let name = "mk"; + + let cppNamespace = "::mlir::mk"; + + let summary = "The Magic Kernel IR in MLIR"; + + let description = [{ + Magic Kernel Dialect. + + Dependent Dialects: + * Memref + * copy, alloc + * Bufferization + * to_tensor + }]; + + let dependentDialects = [ + ]; + + let extraClassDeclaration = [{ + void registerTypes(); + }]; + + // let hasConstantMaterializer = 1; + // let useDefaultTypePrinterParser = 1; + let usePropertiesForAttributes = 1; +} + +include "magic-kernel/Dialect/IR/MagicKernelTypes.td" + +#endif // MAGIC_KERNEL_DIALECT diff --git a/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelOps.td b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelOps.td new file mode 100755 index 00000000..9f38c3cc --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelOps.td @@ -0,0 +1,909 @@ +//===------------------- MagicKernelOps.td --------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// The abstract layer between MLIR core dialects and the lower target specific +// dialects of MagicKernelFunc and MagicKernelInstr. +// +// Compare to higher level MLIR dialects such as memref, arith, affine etc, the +// granularity of MK dialect is more suitable to map into ML accelerators. For +// example, tt.load is lowered to arith + memref.reinterpret_cast + memref.alloc +// + memref.copy + bufferization.to_tensor by decoding hidden high level info +// into detailed info carried in those core MLIR dialects. +// If we convert tt.load to mk.alloc + mk.load, we have to redo all the analysis +// and info constructions which triton-shared already does, so that we should +// generate mk.alloc + mk.load from the core dialects to avoid reconstructing +// the information. +// By doing so, we can lower arith + memref.reinterpret_cast + memref.copy + +// buf.to_tensor into mk.load, and lower arith + memref.alloc into mk.alloc. +// +//===----------------------------------------------------------------------===// + +#ifndef MAGIC_KERNEL_OPS +#define MAGIC_KERNEL_OPS + +include "magic-kernel/Dialect/IR/MagicKernelTypes.td" +include "triton/Dialect/Triton/IR/TritonAttrDefs.td" +include "mlir/Interfaces/SideEffectInterfaces.td" // Pure +include "mlir/Interfaces/InferTypeOpInterface.td" // SameOperandsAndResultType +include "mlir/Interfaces/DestinationStyleOpInterface.td" +include "mlir/IR/OpBase.td" +include "mlir/IR/EnumAttr.td" + +//===----------------------------------------------------------------------===// +// Bufferable type. +//===----------------------------------------------------------------------===// + +def TensorOrMemref : + AnyTypeOf<[AnyMemRef, AnyRankedTensor], "", "::mlir::ShapedType">; + +def I1TensorOrMemref : + AnyTypeOf<[MemRefOf<[I1]>, RankedTensorOf<[I1]>], "", "::mlir::ShapedType">; + +def IntTensorOrMemref : + AnyTypeOf<[MemRefOf<[AnyInteger]>, RankedTensorOf<[AnyInteger]>], "", "::mlir::ShapedType">; + +def FPTensorOrMemref : + AnyTypeOf<[MemRefOf<[AnyFloat]>, RankedTensorOf<[AnyFloat]>], "", "::mlir::ShapedType">; + +def MKCommAddrLike : + AnyTypeOf<[I64, TensorOrMemref], "comm-addr-like">; + +// +// Interfaces +// +def GlobalMemory : Resource<"::mlir::triton::GlobalMemory">; + + +class MKOp traits = []> : + Op { +} + +class MKUnElemWiseOp : MKOp { + let summary = "Element wise unary operation: $mnemonic"; + + let arguments = ( + ins + TensorOrMemref:$src, + // buffer for store result + Arg:$zeroes, + BoolAttr:$is_atomic + ); + + let results = (outs Variadic:$dst); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getZeroesMutable(); + } + }]; + +} + +class MKBinElemWiseOp : MKOp { + let summary = "Element wise binary operation: $mnemonic"; + + let arguments = ( + ins + AnyTensor:$src0, + AnyTensor:$src1, + BoolAttr:$is_atomic + ); + + let results = (outs AnyTensor:$dst); +} + +class MKTerElemWiseOp : MKOp { + let summary = "Element wise binary operation: $mnemonic"; + + let arguments = ( + ins + AnyTensor:$src0, + AnyTensor:$src1, + AnyTensor:$src2, + BoolAttr:$is_atomic + ); + + let results = (outs AnyTensor:$dst); +} + +// arith.bitcast doesn't support pointers +def BitcastOp : MKOp<"bitcast", [Pure]> { + let summary = "Cast between types of the same bitwidth"; + + let arguments = (ins AnyType:$src); + + let results = (outs AnyType:$result); + + let assemblyFormat = "$src attr-dict `:` type($src) `->` type($result)"; + + let hasFolder = 1; + + // TODO: Add verifier +} + +// ============================================================================= +// Memory allocation ops +// ============================================================================= + +def AllocOp : MKOp<"alloc", []> { + let summary = "Allocate a consecutive memory from given addressing space"; + + let description = [{ + It may or may not generate target intrinsic call or instruction, the + lowering from this operator to lower level operator is target specific. + }]; + + let arguments = ( + ins + I32Attr:$addr_space, // The addressing space + I64ArrayAttr:$dims // The size of memory to be allocated + ); + + // Return the pointer of the allocated memory + let results = (outs AnyRankedOrUnrankedMemRef:$ptr); +} + +// ============================================================================= +// Load/Store Ops +// ============================================================================= + +// Unit and strided memory load +def LoadOp : MKOp<"load", []> { + let summary = "Load from a memory with optional strides"; + + let description = [{ See RISC-V RVV unit/strided memory load for detail }]; + + let arguments = ( + ins + AnyRankedOrUnrankedMemRef:$ptr, // The base ptr + I32Attr:$addr_space, // The address space + I64ArrayAttr:$dims, // The shape to be loaded, can be dynamic + I64ArrayAttr:$strides, // The strides in each rank, can be dynamic + BoolAttr:$mask // element is not loaded if mask[i] == 0 + ); + + let results = (outs AnyTensor:$result); + + let assemblyFormat = [{ + $ptr `,` attr-dict `:` type($ptr) `->` type($result) + }]; +} + +// Index memory load +def IndexLoadOp : MKOp<"iload", [ +]> { + let summary = "Load from a memory with indexed offset"; + + let description = [{ See RISC-V RVV index memory load for detail }]; + + let arguments = ( + ins + AnyRankedOrUnrankedMemRef:$ptr, // The base ptr + I32Attr:$addr_space, // The address space + I64ArrayAttr:$dims, // The shape to be loaded + AnyTensor:$index, // The tensor contains memory offset for each element + BoolAttr:$mask // element is not loaded if mask[i] == 0 + ); + + let results = (outs MKType:$result); +} + +// Unit and strided memory store +def StoreOp : MKOp<"store", [MemoryEffects<[MemWrite]>]> { + let summary = "Store to a memory with optional strides"; + + let description = [{ See RISC-V RVV unit/strided memory store for detail }]; + + let arguments = ( + ins + AnyRankedOrUnrankedMemRef:$ptr, // The base ptr + I32Attr:$addr_space, // The address space + I64ArrayAttr:$dims, // The shape to be stored + I64ArrayAttr:$strides, // The strides in each rank + BoolAttr:$mask // element is not write to dest if mask[i] == 0 + ); + + let assemblyFormat = [{ + $ptr `,` attr-dict `:` type($ptr) + }]; +} + +// Index memory store +def IndexStoreOp : MKOp<"istore", [MemoryEffects<[MemWrite]>]> { + let summary = "Store to a memory with indexed offset"; + + let description = [{ See RISC-V RVV index memory store for detail }]; + + let arguments = ( + ins + AnyRankedOrUnrankedMemRef:$ptr, // The base ptr + I32Attr:$addr_space, // The address space + I64ArrayAttr:$dims, // The shape to be stored + AnyTensor:$index, // The tensor contains memory offset for each element + BoolAttr:$mask // element is not write to dest if mask[i] == 0 + ); +} + +// ============================================================================= +// DataMove +// ============================================================================= + +def MaskMoveOp : MKOp<"mask_move", [DestinationStyleOpInterface]> { + let summary = "Mask data move API"; + + let description = [{ When mask is 1, extract the data from src and write it to dst. +When mask=0, the corresponding elements of dst remain unchanged. + }]; + + let arguments = ( + ins + TensorOrMemref:$source, // The source address in SPM + TensorOrMemref:$mask, + Arg:$init // The init buffer + ); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitMutable(); + } + }]; + + // The dst address is not used, use init instead. + let results = (outs Variadic:$dst); +} + +def GatherScatter : MKOp<"gatherscatter", []> { + let summary = "Transfer data in strides and iterations"; + + let arguments = (ins + TensorOrMemref:$source, // The source + TensorOrMemref:$target, // The target + I32Attr:$bytes, // Inner loop data size in bytes + I32Attr:$src_strideN, + I32Attr:$src_strideH, + I32Attr:$src_strideW, + I32Attr:$src_iterN, + I32Attr:$src_iterH, + I32Attr:$src_iterW, + I32Attr:$dst_strideN, + I32Attr:$dst_strideH, + I32Attr:$dst_strideW, + I32Attr:$dst_iterN, + I32Attr:$dst_iterH, + I32Attr:$dst_iterW + ); + let results = (outs Variadic:$dst); +} + +// ============================================================================= +// Dot op +// ============================================================================= + +def DotOp : MKOp<"dot", [DestinationStyleOpInterface]> { + let summary = "Inner production of 2 vectors"; + + let description = [{ + TODO: It is currently one to one mapping from upper dialect tt.dot. + }]; + + let arguments = ( + ins + TensorOrMemref:$a, // Matrix A + TensorOrMemref:$b, // Matrix B + Arg:$inits, + BoolAttr:$en_psum // Enable psum. Used as accumulate buffer + ); + + let results = (outs Variadic:$d); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitsMutable(); + } + }]; + + // let hasVerifier = 1; +} + + +// ============================================================================= +// Dot Scaled op +// ============================================================================= + +// Type for ScaleDotElemType kind of floats. +def MK_ScaleDotElemTypeAttr : I32EnumAttr< + "ScaleDotElemType", "", + [ + I32EnumAttrCase<"E4M3", 0, "e4m3">, + I32EnumAttrCase<"E5M2", 1, "e5m2">, + I32EnumAttrCase<"E2M3", 2, "e2m3">, + I32EnumAttrCase<"E3M2", 3, "e3m2">, + I32EnumAttrCase<"E2M1", 4, "e2m1">, + I32EnumAttrCase<"BF16", 5, "bf16">, + I32EnumAttrCase<"FP16", 6, "fp16"> + ]>{ + let cppNamespace = "::mlir::triton"; +} + +def DotScaledOp : MKOp<"dot_scaled", [DestinationStyleOpInterface, AttrSizedOperandSegments]> { + let summary = "dot_scaled"; + + let description = [{ + $dst = matrix_multiply(scale($a, $a_scale), scale($b, $b_scale)). + Where scale(x, s) is a function that applies the scale per block following microscaling spec. + }]; + + let arguments = ( + ins + // inputs are floats if we have a type for them, otherwise (fp4), + // they are packed in pairs in an I8Tensor + TensorOrMemref:$a, + Optional:$a_scale, + TensorOrMemref:$b, + Optional:$b_scale, + Arg:$dst, + MK_ScaleDotElemTypeAttr:$a_elem_type, + MK_ScaleDotElemTypeAttr:$b_elem_type, + BoolAttr:$fastMath + ); + + let results = (outs Variadic:$res); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getDstMutable(); + } + }]; +} + +// ============================================================================= +// Reduction ops +// ============================================================================= +class ReduceOp : MKOp { + let summary = "Compute extremal values and their indices from source tensor."; + + let arguments = ( + ins + TensorOrMemref:$src, + Arg:$init, + I64ArrayAttr: $nhwcShape, + I32Attr:$axis + ); + + let results = (outs Variadic:$result); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitMutable(); + } + }]; +} + +// TODO: Better way to define reduce op +// A complementary operation for linalg.reduce since tsing-micro need channel-norm +def ReduceMaxOp : ReduceOp<"reduce_max"> {} +def ReduceMinOp : ReduceOp<"reduce_min"> {} +def ReduceSumOp : ReduceOp<"reduce_sum"> {} +def XorSumOp : MKOp<"xor_sum"> {} + +class ArgReduceOp : MKOp { + let summary = "Compute extremal values and their indices from source tensor."; + + let arguments = ( + ins + TensorOrMemref:$src, + Arg:$value, + Arg:$index, + I32Attr:$axis + ); + + let results = (outs Variadic:$result); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return MutableOperandRange(getOperation(), /*start=*/1, /*count=*/2); + } + }]; +} + +def ArgMaxOp : ArgReduceOp<"argmax"> {} +def ArgMinOp : ArgReduceOp<"argmin"> {} + + +// ============================================================================= +// Scan/Sort Ops +// ============================================================================= + +def SortOp : MKOp<"sort", [Pure]> {} +def GatherOp : MKOp<"gather", [DestinationStyleOpInterface]> { + let summary = "Gather from a tensor along a given dimension."; + + let description = [{ + TODO: It is currently one to one mapping from upper dialect tt.gather. + }]; + + let arguments = ( + ins + TensorOrMemref:$src, // input + TensorOrMemref:$indices, // indices + Arg:$dst, // output + I32Attr:$axis + ); + + let results = (outs Variadic:$result); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getDstMutable(); + } + }]; +} + + +// ============================================================================= +// Unary/Binary/Ternary Element-wise Math Ops +// ============================================================================= + +def RandGenOp : MKOp<"randgen", [DestinationStyleOpInterface]> { + let summary = "hardware xorshift128+ PRNG lowered to tx.randgen / __RandGen"; + + let arguments = ( + ins + TensorOrMemref:$seed0, + TensorOrMemref:$seed1, + Arg:$out, + Arg:$seed0_out, + Arg:$seed1_out, + I32Attr:$byte_count, + I16Attr:$fmt + ); + + let results = (outs Variadic:$result); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return MutableOperandRange(getOperation(), /*start=*/2, /*count=*/3); + } + }]; +} + + +def AbsOp : MKUnElemWiseOp<"abs">; +def AddOp : MKBinElemWiseOp<"add">; +def AndOp : MKBinElemWiseOp<"and">; +def CDivOp : MKBinElemWiseOp<"cdiv">; +def CeilOp : MKUnElemWiseOp<"ceil">; +def ClampOp : MKUnElemWiseOp<"clamp">; +def CosOp : MKUnElemWiseOp<"cos">; +def DivOp : MKBinElemWiseOp<"div">; +def ErfOp : MKUnElemWiseOp<"erf">; +def ExpOp : MKUnElemWiseOp<"exp">; +def Exp2Op : MKUnElemWiseOp<"exp2">; +def FdivOp : MKBinElemWiseOp<"fdiv">; +def FloorOp : MKUnElemWiseOp<"floor">; +def FmaOp : MKTerElemWiseOp<"fma">; +def LogOp : MKUnElemWiseOp<"log">; +def Log2Op : MKUnElemWiseOp<"log2">; +def MaxOp : MKUnElemWiseOp<"max">; +def MinOp : MKUnElemWiseOp<"min">; +def OrOp : MKBinElemWiseOp<"or">; +def RsqrtOp : MKUnElemWiseOp<"rsqrt">; +def SigmoidOp : MKUnElemWiseOp<"sigmoid">; + +def GeluOp : MKOp<"gelu", [DestinationStyleOpInterface]> { + let summary = "Fusion gelu op"; + + let arguments = ( + ins + TensorOrMemref:$src, + Optional:$imm, // WORKAROUND for wafer backend + // buffer for store result + Arg:$zeroes, + BoolAttr:$is_atomic, + I16Attr:$gelu_mode + ); + + let results = (outs Variadic:$dst); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getZeroesMutable(); + } + }]; + +} +def SinOp : MKUnElemWiseOp<"sin">; +def SqrtOp : MKUnElemWiseOp<"sqrt">; +def SqrtRnOp : MKUnElemWiseOp<"sqrt_rn">; +def XorOp : MKBinElemWiseOp<"xor">; +// def UmulhiOp : MKOp<"umulhi", [Pure]> {} + +def BarrierOp : MKOp<"barrier"> { + let summary = "Synchronizes all work items"; + let description = [{ + The "barrier" op synchronizes all work items. + }]; + let assemblyFormat = "attr-dict"; +} + +// ============================================================================= +// Communication Ops (TLE) +// ============================================================================= + +def RemoteStoreOp : MKOp<"remote_store", [MemoryEffects<[MemRead, MemWrite]>]> { + let summary = "Store local tensor/memref to a remote tile address"; + + let description = [{ + Asynchronously store data from the current tile to a remote destination + address. The operation reads from the given source buffer (tensor/memref) + and writes it to the remote tile. + }]; + + let arguments = ( + ins + I64:$remote_chip_id_x, // X-coordinate of the remote chip ID + I64:$remote_chip_id_y, // Y-coordinate of the remote chip ID + I64:$remote_die_id, // ID of the remote die + I64:$remote_tile_id, // ID of the remote tile + MKCommAddrLike:$dst_addr, // Remote destination base address (or placeholder) + Arg:$src // Source buffer (read from here) + ); + + let results = (outs); +} + +def RemoteLoadOp : MKOp<"remote_load", [DestinationStyleOpInterface, MemoryEffects<[MemWrite]>]> { + let summary = "Load data from remote tile into a destination buffer"; + + let description = [{ + Destination-style remote load: receive data from a source tile and write it + into the given destination buffer. Returns a tensor view of the received + data (same shape and element type as the destination). + }]; + + let arguments = ( + ins + I64:$remote_chip_id_x, // X-coordinate of the remote chip ID + I64:$remote_chip_id_y, // Y-coordinate of the remote chip ID + I64:$remote_die_id, // ID of the remote die + I64:$remote_tile_id, // ID of the remote tile + Arg:$dst // Destination buffer (receive into here) + ); + + let results = (outs Variadic:$result); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getDstMutable(); + } + }]; +} + +def PrintOp : MKOp<"print", [MemoryEffects<[MemWrite]>, DestinationStyleOpInterface]> { + let summary = "Print at most a single scalar or 1D TensorOrMemref on each line"; + + let description = [{ + It only takes a single scalar or 1D TensorOrMemref element. + }]; + + let arguments = (ins + StrAttr:$prefix, + BoolAttr:$hex, + Variadic>:$val, + DenseI32ArrayAttr:$isSigned + ); + + let results = (outs Variadic:$dst); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getValMutable(); + } + }]; + + let hasVerifier = 1; +} + +def AssertOp : MKOp<"assert"> { + let summary = "Assert the condition at runtime from the device"; + + let arguments = (ins + StrAttr:$message + ); +} + + +// ============================================================================= +// Binary(scalr and tensor) Element-wise Math Ops +// ============================================================================= + +class ArithVSOp traits = [DestinationStyleOpInterface]> : + MKOp { + + let summary = "Vector-Scalar arith op which return output with same input vector type"; + + let arguments = (ins + FPTensorOrMemref:$input, // First input vector address + AnyFloat:$value, // Const value + Arg:$init // Out vector address + ); + let results = (outs Variadic:$dst); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitMutable(); + } + }]; +} + +def AddVS : ArithVSOp<"addvs">; +def SubVS : ArithVSOp<"subvs">; +def MulVS : ArithVSOp<"mulvs">; + +// ============================================================================= +// Relation Ops +// ============================================================================= +class RelationVVOp traits = [DestinationStyleOpInterface]> : + MKOp { + let summary = "Vector-Vector relation op which return output with same input type"; + + let arguments = (ins + FPTensorOrMemref:$input0, // First input vector address + FPTensorOrMemref:$input1, // Second vector address + Arg:$init // init vector address + ); + let results = (outs Variadic:$dst); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitMutable(); + } + }]; +} + +def EqualVV : RelationVVOp<"equalvv"> { + let summary = "compare two input value, if equal, return 1.0"; +} + +def UnEqualVV : RelationVVOp<"unequalvv"> { + let summary = "compare two input value, if unequal, return 1.0"; +} + +def GreaterEqualVV : RelationVVOp<"greatrequalvv"> { + let summary = "compare two input value, if src0 >= src1, return 1.0"; +} + +def GreaterVV : RelationVVOp<"greatervv"> { + let summary = "compare two input value, if src0 > src1, return 1.0"; +} + +def LessEqualVV : RelationVVOp<"lessequalvv"> { + let summary = "compare two input value, if src0 <= src1, return 1.0"; +} + +def LessThenVV : RelationVVOp<"lessthenvv"> { + let summary = "compare two input value, if src0 < src1, return 1.0"; +} + +class RelationVSOp traits = [DestinationStyleOpInterface]> : + MKOp { + + let summary = "Vector-Scalar relation op which return output with same input vector type"; + + let arguments = (ins + FPTensorOrMemref:$input, // First input vector address + AnyFloat:$value, // Const value + Arg:$init // Out vector address + ); + let results = (outs Variadic:$dst); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitMutable(); + } + }]; +} + +def BoolEqualVS : RelationVSOp<"boolequalvs"> { + let summary = "compare input value with ConstantOp, if equal, return true"; +} + +def BoolUnEqualVS : RelationVSOp<"boolunequalvs"> { + let summary = "compare input value with ConstantOp, if unequal, return true"; +} + +def BoolGreaterEqualVS : RelationVSOp<"boolgreatrequalvs"> { + let summary = "compare input value with ConstantOp, if src0 >= src1, return true"; +} + +def BoolGreaterVS : RelationVSOp<"boolgreatervs"> { + let summary = "compare input value with ConstantOp, if src0 > src1, return true"; +} + +def BoolLessEqualVS : RelationVSOp<"boollessequalvs"> { + let summary = "compare input value with ConstantOp, if src0 <= src1, return true"; +} + +def BoolLessThenVS : RelationVSOp<"boollessthenvs"> { + let summary = "compare input value with ConstantOp, if src0 < src1, return true"; +} + +def EqualVS : RelationVSOp<"equalvs"> { + let summary = "compare input value with ConstantOp, if equal, return 1.0"; +} + +def UnEqualVS : RelationVSOp<"unequalvs"> { + let summary = "compare input value with ConstantOp, if unequal, return 1.0"; +} + +def GreaterEqualVS : RelationVSOp<"greatrequalvs"> { + let summary = "compare input value with ConstantOp, if src0 >= src1, return 1.0"; +} + +def GreaterVS : RelationVSOp<"greatervs"> { + let summary = "compare input value with ConstantOp, if src0 > src1, return 1.0"; +} + +def LessEqualVS : RelationVSOp<"lessequalvs"> { + let summary = "compare input value with ConstantOp, if src0 <= src1, return 1.0"; +} + +def LessThenVS : RelationVSOp<"lessthenvs"> { + let summary = "compare input value with ConstantOp, if src0 < src1, return 1.0"; +} + +// ============================================================================= +// Atomic Ops +// ============================================================================= + +def AtomicRMWOp : MKOp<"atomic_rmw", [ + DestinationStyleOpInterface +]> { + let summary = "perform atomic read-modify-write operation on a pointer"; + + let description = [{ + load data at $ptr, do $rmw_op with $val, and store result to $ptr. + + return old value at $ptr + }]; + + let arguments = (ins Arg:$ptr, + AnyType:$val, + Arg:$dst, + TT_AtomicRMWAttr:$atomic_rmw_op, + TT_MemSemanticAttr:$sem, + TT_MemSyncScopeAttr:$scope); + + let results = (outs Variadic:$result); + + // Explicitly list $atomic_rmw_op, $sem, and $scope rather than relying on + // attr-dict so they're printed as strings rather than opaque integers. + let assemblyFormat = [{ + $atomic_rmw_op `,` $sem `,` $scope `,` $ptr `,` $val `,` $dst attr-dict `:` + functional-type(operands, $result) + }]; + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getDstMutable(); + } + }]; +} + +def AtomicCASOp : MKOp<"atomic_cas", [ + DestinationStyleOpInterface +]> { + let summary = "perform atomic compare-and-swap operation on a pointer"; + + let description = [{ + compare $cmp with data $old at location $ptr, + + if $old == $cmp, store $val to $ptr, + + else store $old to $ptr, + + return $old + }]; + + + let arguments = (ins Arg:$ptr, + AnyType:$cmp, AnyType:$val, + Arg:$dst, + TT_MemSemanticAttr:$sem, + TT_MemSyncScopeAttr:$scope); + + let results = (outs Variadic:$result); + + // Explicitly list $sem and $scope rather than relying on attr-dict so + // they're printed as strings rather than opaque integers. + let assemblyFormat = [{ + $sem `,` $scope `,` $ptr `,` $cmp `,` $val `,` $dst attr-dict `:` + functional-type(operands, $result) + }]; + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getDstMutable(); + } + }]; +} + +// ============================================================================= +// Sync ops +// ============================================================================= +def AtomicBarrierInOp : MKOp<"atomic_barrier_in", [MemoryEffects<[MemRead, MemAlloc]>]> { + let summary = "Synchronizes all work items before input"; + let description = [{ + The "barrier in" op synchronizes all work items. + }]; + let assemblyFormat = "attr-dict"; +} + +def AtomicBarrierOutOp : MKOp<"atomic_barrier_out", [MemoryEffects<[MemAlloc, MemWrite]>]> { + let summary = "Synchronizes all work items after output"; + let description = [{ + The "barrier out" op synchronizes all work items. + }]; + let assemblyFormat = "attr-dict"; +} + +// ============================================================================= +// Peripheral instructions +// ============================================================================= +def Bit2FpOp : MKOp<"bit2fp", [DestinationStyleOpInterface]> { + let summary = "Convert a vector of the bitwise into the fp vector"; + + let arguments = (ins + I1TensorOrMemref:$src, // Input tensor + Arg:$init + ); + let results = (outs Variadic:$result); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitMutable(); + } + }]; +} + +def CastOp : MKOp<"cast", [DestinationStyleOpInterface]> { + let summary = "Cast a tensor or memref to a different type"; + + let arguments = (ins + TensorOrMemref:$src, // Input tensor + Arg:$init + ); + let results = (outs Variadic:$result); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitMutable(); + } + }]; +} + +// TODO: May be replaced by quant dialect ops in the future +def DequantOp : MKOp<"dequant", [DestinationStyleOpInterface]> { + let summary = "output = (input - zero_point) * scale"; + + let arguments = (ins + TensorOrMemref:$src, // Input tensor + TensorOrMemref:$scale, // scale: float but stored as uint8 + Arg:$init + ); + let results = (outs Variadic:$result); + + let extraClassDeclaration = [{ + MutableOperandRange getDpsInitsMutable() { + return getInitMutable(); + } + }]; +} + +#endif // MAGIC_KERNEL_OPS diff --git a/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelTypes.td b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelTypes.td new file mode 100755 index 00000000..19fb9e1b --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelTypes.td @@ -0,0 +1,102 @@ +//===------------------- MagicKernelTypes.td ------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MAGIC_KERNEL_TYPES_TD +#define MAGIC_KERNEL_TYPES_TD + +include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/BuiltinTypeInterfaces.td" +include "magic-kernel/Dialect/IR/MagicKernelDialect.td" + +// +// Types +// +class MKTypeDef traits = []> + : TypeDef { + // Used by printer/parser + let mnemonic = _mnemonic; +} + +// Floating-point Type +def MKFloat : AnyTypeOf<[F8E4M3FN, F8E4M3FNUZ, F8E5M2, F8E5M2FNUZ, F16, BF16, F32, F64], "floating-point">; +def MKFloatTensor : RankedTensorOf<[MKFloat]>; +def MKFloatLike : AnyTypeOf<[MKFloat, MKFloatTensor]>; + +// Boolean Type +// TT_Bool -> I1 +def MKBoolTensor : RankedTensorOf<[I1]>; +def MKBoolLike : AnyTypeOf<[I1, MKBoolTensor]>; + +// Integer Type +def I4 : I<4>; +def MKInt : AnyTypeOf<[I1, I4, I8, I16, I32, I64], "integer">; +def MKIntTensor : RankedTensorOf<[MKInt]>; +def MKIntLike : AnyTypeOf<[MKInt, MKIntTensor]>; + +// I32 Type +// MKI32 -> I32 +// MKI32Tensor -> I32Tensor +def MKI32Like : AnyTypeOf<[I32, I32Tensor]>; + +// I64 Type +// MKI64 -> I64 +// MKI64Tensor -> I64Tensor +def MKI64Like : AnyTypeOf<[I64, I64Tensor]>; + +// Pointer Type in TableGen +class MKPtrOf pointeeTypes> : + DialectType($_self)">, + Concat<"[](::mlir::Type pointeeType) { return ", + SubstLeaves<"$_self", "pointeeType", AnyTypeOf.predicate>, + "; }(::mlir::cast<::mlir::triton::PointerType>($_self).getPointeeType())">]>, + "ptr", "::mlir::triton::PointerType">; + +// Pointer Type in C++ (corresponding to `MKPtrOf`) +def MKPtrType : MKTypeDef<"Pointer", "ptr"> { + let summary = "Pointer type (`::mlir::triton::PointerType`) in Triton IR type system"; + + let description = [{ + Pointer type in Triton IR type system, which could be pointing to scalars or tensors. + }]; + + let parameters = (ins "Type":$pointeeType, "int":$addressSpace); + + let builders = [ + TypeBuilderWithInferredContext<(ins + "Type":$pointeeType, + "int":$addressSpace + ), [{ + return $_get(pointeeType.getContext(), pointeeType, addressSpace); + }]> + ]; + + let hasCustomAssemblyFormat = 1; + + let skipDefaultBuilders = 1; +} + +// Scalar Pointer Type: `ptr<>` +def MKPtr : MKPtrOf<[AnyType]>; + +// Tensor of Pointer Type: `tensor>` +def MKPtrTensor : RankedTensorOf<[MKPtr]>; + +// Tensor of Pointer Type or Pointer type: `tensor>` or `ptr<>` +def MKPtrLike : AnyTypeOf<[MKPtr, MKPtrTensor]>; + +// Tensor Type +def MKFpIntTensor : RankedTensorOf<[MKFloat, MKInt]>; +def MKTensor : RankedTensorOf<[MKFloat, MKInt, MKPtr]>; + +// Pointer Type to Tensor Type: `ptr>` +def MKTensorPtr : MKPtrOf<[MKTensor]>; + +// Any Type in Magic Kernel IR +def MKType : AnyTypeOf<[MKFloatLike, MKIntLike, MKPtrLike, MKTensorPtr]>; + +#endif // MAGIC_KERNEL_TYPES_TD diff --git a/third_party/wafer/include/magic-kernel/Transforms/BufferizableOpInterfaceImpl.h b/third_party/wafer/include/magic-kernel/Transforms/BufferizableOpInterfaceImpl.h new file mode 100755 index 00000000..789a0c8a --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Transforms/BufferizableOpInterfaceImpl.h @@ -0,0 +1,26 @@ +//===- BufferizableOpInterfaceImpl.h - Impl. of BufferizableOpInterface ---===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM +// Exceptions. See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file declares the implementation of the BufferizableOpInterface. +// +//===----------------------------------------------------------------------===// + +#ifndef _MK_DIALECT_BUFFERIZABLEOPINTERFACEIMPL_H +#define _MK_DIALECT_BUFFERIZABLEOPINTERFACEIMPL_H + +namespace mlir { +class DialectRegistry; + +namespace mk { +void registerBufferizableOpInterfaceExternalModels(DialectRegistry ®istry); +} // namespace mk +} // namespace mlir + +#endif // _MK_DIALECT_BUFFERIZABLEOPINTERFACEIMPL_H diff --git a/third_party/wafer/include/magic-kernel/Transforms/CMakeLists.txt b/third_party/wafer/include/magic-kernel/Transforms/CMakeLists.txt new file mode 100644 index 00000000..7f4d2aa9 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Transforms/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name MKTransforms) +add_public_tablegen_target(MKTransformsPassIncGen) diff --git a/third_party/wafer/include/magic-kernel/Transforms/Passes.h b/third_party/wafer/include/magic-kernel/Transforms/Passes.h new file mode 100644 index 00000000..9777cb29 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Transforms/Passes.h @@ -0,0 +1,20 @@ +#ifndef MK_TRANSFORMS_PASSES_H +#define MK_TRANSFORMS_PASSES_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> +createMaterializeStridedLinalgInputsPass(); + +#define GEN_PASS_REGISTRATION +#define GEN_PASS_DECL +#include "magic-kernel/Transforms/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // MK_TRANSFORMS_PASSES_H diff --git a/third_party/wafer/include/magic-kernel/Transforms/Passes.td b/third_party/wafer/include/magic-kernel/Transforms/Passes.td new file mode 100644 index 00000000..dc3a4c62 --- /dev/null +++ b/third_party/wafer/include/magic-kernel/Transforms/Passes.td @@ -0,0 +1,49 @@ +#ifndef MK_TRANSFORMS_PASSES +#define MK_TRANSFORMS_PASSES + +include "mlir/Pass/PassBase.td" + +def MaterializeStridedLinalgInputs + : Pass<"materialize-strided-linalg-inputs", "mlir::ModuleOp"> { + let summary = "Materialize non-contiguous subview inputs of compute linalg.generic ops"; + + let description = [{ + This pass runs after OneShotBufferize. + + Some TX device compute ops do not support strided memref operands. After + bufferization, tensor.extract_slice may become memref.subview. If a + linalg.generic compute op consumes such a non-contiguous subview directly, + later lowering may drop or ignore the subview strides and read incorrect + values. + + This pass finds compute linalg.generic ops whose input operands are + non-contiguous memref.subview values. For each such input, it creates a + contiguous memref.alloc, inserts memref.copy from the subview into the + allocation, and rewrites the linalg.generic input to use the contiguous + buffer. + + Example: + + %sv = memref.subview %arg0[...] : + memref<1x4x2xf32> to memref<1x4xf32, strided<[8, 2], offset: 1>> + + linalg.generic ins(%sv : memref<1x4xf32, strided<[8, 2], offset: 1>>) ... + + becomes: + + %tmp = memref.alloc() : memref<1x4xf32> + memref.copy %sv, %tmp : + memref<1x4xf32, strided<[8, 2], offset: 1>> to memref<1x4xf32> + + linalg.generic ins(%tmp : memref<1x4xf32>) ... + }]; + + let dependentDialects = [ + "mlir::linalg::LinalgDialect", + "mlir::memref::MemRefDialect" + ]; + + let constructor = "triton::createMaterializeStridedLinalgInputsPass()"; +} + +#endif // MK_TRANSFORMS_PASSES diff --git a/third_party/wafer/include/triton-shared/CMakeLists.txt b/third_party/wafer/include/triton-shared/CMakeLists.txt new file mode 100755 index 00000000..bd3c0c6c --- /dev/null +++ b/third_party/wafer/include/triton-shared/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(Conversion) diff --git a/third_party/wafer/include/triton-shared/Conversion/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/CMakeLists.txt new file mode 100755 index 00000000..55f02553 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/CMakeLists.txt @@ -0,0 +1,12 @@ +add_subdirectory(TritonArithToLinalg) +add_subdirectory(StructuredToMemref) +add_subdirectory(ConvertTritonPtr) +add_subdirectory(TritonToCoreDialects) + +add_subdirectory(ReconcilePtrCasts) +add_subdirectory(UnstructuredToMK) + +# flir +# add_subdirectory(TritonPtrToMemref) +# add_subdirectory(UnstructuredToMemref) +# add_subdirectory(TritonToUnstructured) diff --git a/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/CMakeLists.txt new file mode 100755 index 00000000..06e58e0e --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/CMakeLists.txt @@ -0,0 +1,9 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Triton Project Contributors. +# +#===------------------------------------------------------------------------===# + +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name ConvertTritonPtr) +add_public_tablegen_target(ConvertTritonPtrPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/Passes.h b/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/Passes.h new file mode 100755 index 00000000..b10a1705 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/Passes.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef CONVERT_TRITON_PTR_CONVERSION_PASSES_H +#define CONVERT_TRITON_PTR_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/ConvertTritonPtr/TritonPtrToAddress.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/ConvertTritonPtr/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/Passes.td b/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/Passes.td new file mode 100755 index 00000000..581ddef3 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/Passes.td @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_PTR_TO_MEMREF_CONVERSION_PASSES +#define TRITON_PTR_TO_MEMREF_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonPtrToAddress : Pass<"triton-ptr-to-address", "mlir::ModuleOp"> { + let summary = "Convert Triton ops on pointers to the address dialect"; + let constructor = "triton::createTritonPtrToAddressPass()"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/TritonPtrToAddress.h b/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/TritonPtrToAddress.h new file mode 100755 index 00000000..618cbd7f --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/TritonPtrToAddress.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_TRITON_PTR_TO_ADDRESS_H +#define TRITON_CONVERSION_TRITON_PTR_TO_ADDRESS_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonPtrToAddressPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITON_PTR_TO_ADDRESS_H diff --git a/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/CMakeLists.txt new file mode 100755 index 00000000..278b4906 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name ReconcilePtrCasts) +add_public_tablegen_target(ReconcilePtrCastsPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.h b/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.h new file mode 100755 index 00000000..941d5b8b --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.h @@ -0,0 +1,15 @@ +#ifndef RECONCILE_PTR_CASTS_CONVERSION_PASSES_H +#define RECONCILE_PTR_CASTS_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/ReconcilePtrCasts/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.td b/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.td new file mode 100755 index 00000000..d19c5e8a --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.td @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef RECONCILE_PTR_CASTS_CONVERSION_PASSES +#define RECONCILE_PTR_CASTS_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def ReconcilePtrCasts : Pass<"reconcile-ptr-casts", "mlir::ModuleOp"> { + let summary = "Convert unrealized_cast op between tt.ptr or ptr.ptr to memref to to_memref or from_memref"; + let constructor = "triton::createReconcilePtrCastsPass()"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h b/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h new file mode 100755 index 00000000..bea24e8f --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_TRITONTOLINALG_ReconcilePtrCasts_H +#define TRITON_CONVERSION_TRITONTOLINALG_ReconcilePtrCasts_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createReconcilePtrCastsPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONTOLINALG_ReconcilePtrCasts_H diff --git a/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/CMakeLists.txt new file mode 100755 index 00000000..9c99c0a1 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name StructuredToMK) +add_public_tablegen_target(StructuredToMKConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/Passes.h b/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/Passes.h new file mode 100755 index 00000000..39ab2b7a --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_STRUCTURED_TO_MK_CONVERSION_PASSES_H +#define TRITON_STRUCTURED_TO_MK_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/StructuredToMK/StructuredToMK.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/StructuredToMK/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/Passes.td b/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/Passes.td new file mode 100755 index 00000000..8cef59a3 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/Passes.td @@ -0,0 +1,10 @@ +#ifndef TRITON_STRUCTURED_TO_MK_CONVERSION_PASSES +#define TRITON_STRUCTURED_TO_MK_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def StructuredToMK : Pass<"structured-to-mk", "mlir::ModuleOp"> { + let summary = "Convert triton structured pointer ops to mk"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/StructuredToMK.h b/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/StructuredToMK.h new file mode 100755 index 00000000..a252efdb --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/StructuredToMK/StructuredToMK.h @@ -0,0 +1,24 @@ +#ifndef TRITON_CONVERSION_STRUCTUREDTOMK_StructuredToMK_H +#define TRITON_CONVERSION_STRUCTUREDTOMK_StructuredToMK_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +class TypeConverter; +namespace triton { + +#define GEN_PASS_DECL +#include "triton-shared/Conversion/StructuredToMK/Passes.h.inc" + +void populateStructuredToMKConversionPatterns(RewritePatternSet &patterns, + TypeConverter &typeConverter); + +std::unique_ptr> createStructuredToMKPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_StructuredToMK_StructuredToMK_H diff --git a/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/CMakeLists.txt new file mode 100755 index 00000000..83ff64d3 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name StructuredToMemref) +add_public_tablegen_target(StructuredToMemrefConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/Passes.h b/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/Passes.h new file mode 100755 index 00000000..198675b1 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_STRUCTURED_TO_MEMREF_CONVERSION_PASSES_H +#define TRITON_STRUCTURED_TO_MEMREF_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/StructuredToMemref/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/Passes.td b/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/Passes.td new file mode 100755 index 00000000..03d0bc38 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/Passes.td @@ -0,0 +1,10 @@ +#ifndef LYF_STRUCTURED_TO_MEMREF_CONVERSION_PASSES +#define LYF_STRUCTURED_TO_MEMREF_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def StructuredToMemref : Pass<"structured-to-memref", "mlir::ModuleOp"> { + let summary = "Convert triton structured pointer ops to memref"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h b/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h new file mode 100755 index 00000000..8c67c9ec --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h @@ -0,0 +1,24 @@ +#ifndef TRITON_CONVERSION_STRUCTUREDTOMEMREF_STRUCTUREDTOMEMREF_H +#define TRITON_CONVERSION_STRUCTUREDTOMEMREF_STRUCTUREDTOMEMREF_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +class TypeConverter; +namespace triton { + +#define GEN_PASS_DECL +#include "triton-shared/Conversion/StructuredToMemref/Passes.h.inc" + +void populateStructuredToMemrefConversionPatterns(RewritePatternSet &patterns, + TypeConverter &typeConverter); + +std::unique_ptr> createStructuredToMemrefPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_STRUCTUREDTOMEMREF_STRUCTUREDTOMEMREF_H diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/CMakeLists.txt new file mode 100755 index 00000000..85076bd1 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonArithToLinalg) +add_public_tablegen_target(TritonArithToLinalgConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.h b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.h new file mode 100755 index 00000000..208459cf --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.h @@ -0,0 +1,2572 @@ +#ifndef TRITON_CONVERSION_PATTERNS +#define TRITON_CONVERSION_PATTERNS + +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "triton-shared/Analysis/MaskAnalysis.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/Analysis/PtrAnalysis.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "triton-shared/Utils/FusionHelper.h" +#include "triton-shared/Utils/ReduceScanCommon.h" +#include "triton-shared/Utils/Utils.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/GPU/IR/GPUDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include +#include +#include +#include +#include + +using namespace mlir; +using namespace triton; + +//===----------------------------------------------------------------------===// +// Utilities +//===----------------------------------------------------------------------===// + +// Extract a scalar value from v. +// If v is a scalar, return that directly. Otherwise, parse through operations +// (currently only support splat, sitofp, and truncf) that produce it to +// extract the underlying scalar value. We then reconstruct the chain of +// operations that can produce this constant with the original type. If no +// scalar value can be extracted, a nullptr is returned. +static Value getScalarValue(Value operand, Location loc, + ConversionPatternRewriter &rewriter) { + SmallVector ops; + + auto reconstructScalarValue = [&](Value src) { + for (auto op = ops.rbegin(); op != ops.rend(); ++op) { + src = TypeSwitch(*op) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return rewriter.create(loc, resType, src); + }) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return rewriter.create(loc, resType, src); + }) + .Default([](Operation *op) { + llvm_unreachable("unsupported op in generating "); + return nullptr; + }); + } + return src; + }; + + while (true) { + if (!dyn_cast(operand.getType())) { + return reconstructScalarValue(operand); + } else if (auto op = operand.getDefiningOp()) { + if (auto attr = dyn_cast(op.getValue())) { + if (!attr.isSplat()) { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load " + "produced by unsupported instruction"; + return nullptr; + } + auto elemValue = attr.getSplatValue(); + auto constOp = arith::ConstantOp::materialize( + rewriter, elemValue, attr.getElementType(), op.getLoc()); + return reconstructScalarValue(constOp.getResult()); + } + } else if (auto op = operand.getDefiningOp()) { + operand = op.getSrc(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load produced " + "by unsupported instruction"; + return nullptr; + } + } + return nullptr; +} + +static SmallVector getNParallelLoopsAttrs(unsigned n) { + return SmallVector(n, utils::IteratorType::parallel); +} + +static Value getTransposedValue(Value source, const Location loc, + ConversionPatternRewriter &rewriter) { + + auto sourceType = cast(source.getType()); + auto sourceRank = sourceType.getRank(); + + SmallVector perm(sourceRank); + std::iota(std::begin(perm), std::end(perm), 0); + std::swap(perm[sourceRank - 1], perm[sourceRank - 2]); + + SmallVector transposedShape(sourceType.getShape()); + std::swap(transposedShape[sourceRank - 1], transposedShape[sourceRank - 2]); + + Value transposeInit = rewriter.create( + loc, transposedShape, sourceType.getElementType()); + + Value transpose = + rewriter.create(loc, source, transposeInit, perm) + .getResults()[0]; + + return transpose; +} + +// for IntLike and FloatLike types +static std::optional getBitWidth(Type a) { + if (auto type = dyn_cast(a)) { + auto elementType = type.getElementType(); + if (elementType.isIntOrFloat()) { + return type.getElementType().getIntOrFloatBitWidth(); + } + return std::nullopt; + } + + if (a.isIntOrFloat()) + return a.getIntOrFloatBitWidth(); + + return std::nullopt; +} + +//===----------------------------------------------------------------------===// +// Op Lowering Patterns +//===----------------------------------------------------------------------===// + +namespace { + +//----------------------------- +// Begin of monolithic only +//----------------------------- +struct AdvanceConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::AdvanceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + llvm::SmallDenseMap knownPtrs; + PtrState pointerState; + PtrAnalysis::rewriteAdvanceOp(op, rewriter, knownPtrs); + return success(); + } +}; + +struct MakeTensorPtrConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + void populateVectorAsIndex(SmallVector &vec, + Operation::operand_range ops, + ConversionPatternRewriter &rewriter, + Location loc) const { + for (auto opnd : ops) { + if (isa(opnd.getType())) { + auto castOp = rewriter.create( + loc, rewriter.getIndexType(), opnd); + vec.push_back(castOp.getResult()); + } else { + assert(isa(opnd.getType())); + vec.push_back(opnd); + } + } + } + + LogicalResult + matchAndRewrite(triton::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + PtrState pointerState; + + auto orderSize = op.getOrder().size(); + if (orderSize > 1) { + for (auto [first, second] : + llvm::zip(op.getOrder().slice(0, orderSize - 2), + op.getOrder().slice(1, orderSize - 1))) { + assert(first == second + 1 && + "Currently only support default order on block pointers"); + } + } + + pointerState.source = rewriter.getRemappedValue(op.getBase()); + populateVectorAsIndex(pointerState.offsets, op.getOffsets(), rewriter, loc); + populateVectorAsIndex(pointerState.strides, op.getStrides(), rewriter, loc); + + SmallVector newOffsets; + for (auto [offset, stride] : + llvm::zip(pointerState.offsets, pointerState.strides)) { + auto mulOp = rewriter.create(loc, cast(offset), + cast(stride)); + newOffsets.push_back(mulOp.getResult()); + } + + pointerState.offsets.clear(); + + for (auto offset : newOffsets) { + pointerState.offsets.push_back(offset); + } + + ArrayRef resultShape; + auto pointerType = + cast(op.getResult().getType()); + if (auto shapedType = dyn_cast(pointerType.getPointeeType())) { + resultShape = shapedType.getShape(); + for (auto dim_size : resultShape) { + pointerState.sizes.push_back( + IntegerAttr::get(IntegerType::get(op.getContext(), 64), dim_size)); + } + } else { + // scalar pointer, should produce a one dimensional memref + SmallVector scalarShape(1, 1); + resultShape = scalarShape; + assert(pointerState.getRank() == 1); + } + + auto castOp = pointerState.createCastOp(resultShape, loc, rewriter); + rewriter.replaceOp(op, castOp.getResult()); + return success(); + } +}; + +struct LegacyAddPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::AddPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + llvm::SmallDenseMap knownPtrs; + PtrAnalysis::rewriteAddptrOp(op, rewriter, knownPtrs); + return success(); + } +}; + +struct LoadConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + void createSideBySideCopies(Value block1, Value block2, Value dst, + Location loc, + ConversionPatternRewriter &rewriter) const { + + auto zero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + + auto one = + rewriter.create(loc, rewriter.getIndexAttr(1)); + + Value block1Row = rewriter.create(loc, block1, 0); + Value block1Col = rewriter.create(loc, block1, 1); + + Value block2Row = rewriter.create(loc, block2, 0); + Value block2Col = rewriter.create(loc, block2, 1); + + auto block1Dst = + rewriter.create(loc, dst, /* offsets */ + ValueRange{zero, zero}, + /* sizes */ + ValueRange{block1Row, block1Col}, + /* strides */ + ValueRange{one, one}); + + auto block2Dst = + rewriter.create(loc, dst, + /* offsets */ + ValueRange{zero, block1Col}, + /* sizes */ + ValueRange{block2Row, block2Col}, + /* strides */ + ValueRange{one, one}); + + rewriter.create(loc, block1, block1Dst); + rewriter.create(loc, block2, block2Dst); + } + + void createStackedCopies(Value block1, Value block2, Value dst, Location loc, + ConversionPatternRewriter &rewriter) const { + + auto zero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + auto one = + rewriter.create(loc, rewriter.getIndexAttr(1)); + + Value block1Row = rewriter.create(loc, block1, 0); + Value block1Col = rewriter.create(loc, block1, 1); + + Value block2Row = rewriter.create(loc, block2, 0); + Value block2Col = rewriter.create(loc, block2, 1); + + auto block1Dst = + rewriter.create(loc, dst, /* offsets */ + ValueRange{zero, zero}, + /* sizes */ + ValueRange{block1Row, block1Col}, + /* strides */ + ValueRange{one, one}); + + auto block2Dst = + rewriter.create(loc, dst, + /* offsets */ + ValueRange{block1Row, zero}, + /* sizes */ + ValueRange{block2Row, block2Col}, + /* strides */ + ValueRange{one, one}); + + rewriter.create(loc, block1, block1Dst); + rewriter.create(loc, block2, block2Dst); + } + +public: + LogicalResult + matchAndRewrite(triton::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto ptr = adaptor.getPtr(); + auto mask = op.getMask(); + auto other = op.getOther(); + auto loc = op.getLoc(); + + // 0. Shortcut for scalar loads + if (!isa(op.getResult().getType())) { + auto sMemRef = PtrAnalysis::getScalarMemRef(op.getPtr(), adaptor.getPtr(), + loc, rewriter); + auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext()); + auto loadOp = rewriter.create( + op.getLoc(), sMemRef, zeroMap, ValueRange{}); + rewriter.replaceOp(op, loadOp.getResult()); + return success(); + } + + // 1. Simple case where no mask is used. + auto type = dyn_cast(ptr.getType()); + if (!type) { + // Seen when implicit broadcasting is done late in a chain of operations. + // The workaround is to broadcast the pointers early in the address + // calculation. A proper fix is complicated, but at least we can provide a + // better error message. + return rewriter.notifyMatchFailure( + op, "LoadOp expects a memref, not a memref of pointers"); + } + + auto tensorType = + RankedTensorType::get(type.getShape(), type.getElementType()); + auto alloc = rewriter.create( + loc, MemRefType::get(type.getShape(), type.getElementType())); + + if (!mask) { + assert(!other && "other value used in non-masked load"); + if (auto unrealizedCast = + ptr.getDefiningOp()) { + if (auto wrapType = unrealizedCast->getAttrOfType( + ModuloState::WraparoundAttr)) { + + auto memrefs = unrealizedCast.getOperands(); + auto block1 = memrefs[0]; + auto block2 = memrefs[1]; + + if (wrapType.getValue() == ModuloState::WraparoundSideBySide) { + createSideBySideCopies(block1, block2, alloc, loc, rewriter); + } else if (wrapType.getValue() == ModuloState::WraparoundStacked) { + createStackedCopies(block1, block2, alloc, loc, rewriter); + } else { + llvm_unreachable("unexpected wraparound type"); + } + } else { + llvm_unreachable("unexpected unrealized cast op"); + } + + } else { + rewriter.create(loc, ptr, alloc); + } + + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + + return success(); + } + + // 2. Continuous masked loads. + // Analyze the mask operand to determine at runtime the size of the data we + // are moving. + MaskState mstate; + auto isContMask = mstate.parse(mask, loc, rewriter); + + if (isContMask.failed()) { + return rewriter.notifyMatchFailure( + op, "Cannot lower continuous masked loads"); + } + + // fill load destination with other value + if (other) { + auto scalarOther = getScalarValue(other, loc, rewriter); + assert(scalarOther && "other value used in masked load produced by " + "unsupported instruction"); + + // For each dimension check if mstate.dims[i] < shape[i], or-accumulate + // the result + auto shape = type.getShape(); + auto accBase = + rewriter.create(loc, rewriter.getBoolAttr(false)) + .getResult(); + for (size_t i = 0; i < type.getShape().size(); i++) { + auto shapei = rewriter.create( + loc, rewriter.getIndexAttr(shape[i])); + + Value dimi = dyn_cast(mstate.dims[i]); + if (!dimi) { + dimi = rewriter.create( + loc, cast(cast(mstate.dims[i]))); + } + + auto cmpOp = rewriter.create( + loc, arith::CmpIPredicate::slt, dimi, shapei); + accBase = rewriter.create(loc, accBase, cmpOp.getResult()) + .getResult(); + } + + // condition the memset on the or-accumulation + // initialize with padding prior to CopyOp + rewriter.create( + loc, accBase, [&](OpBuilder &builder, Location loc) { + builder.create(loc, ValueRange{scalarOther}, + ValueRange{alloc}); + builder.create(loc); + }); + } + + if (auto unrealizedCast = ptr.getDefiningOp()) { + if (auto wrapType = unrealizedCast->getAttrOfType( + ModuloState::WraparoundAttr)) { + + auto memrefs = unrealizedCast.getOperands(); + auto block1 = memrefs[0]; + auto block2 = memrefs[1]; + + if (wrapType.getValue() == ModuloState::WraparoundSideBySide) { + auto [subview1, subview2] = + mstate.getSideBySideSubviews(block1, block2, loc, rewriter); + + createSideBySideCopies(subview1, subview2, alloc, loc, rewriter); + } else if (wrapType.getValue() == ModuloState::WraparoundStacked) { + auto [subview1, subview2] = + mstate.getStackedSubviews(block1, block2, loc, rewriter); + + createStackedCopies(subview1, subview2, alloc, loc, rewriter); + } else { + llvm_unreachable("unexpected wraparound type"); + } + + } else { + llvm_unreachable("unexpected unrealized cast op"); + } + + } else { + memref::SubViewOp srcSubview = mstate.getSubview(ptr, loc, rewriter); + memref::SubViewOp dstSubview = mstate.getSubview(alloc, loc, rewriter); + rewriter.create(loc, srcSubview, dstSubview); + } + + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + + return success(); + } +}; + +struct StoreConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto ptr = adaptor.getPtr(); + auto val = adaptor.getValue(); + auto mask = op.getMask(); + auto loc = op.getLoc(); + + // 0. Shortcut for scalar stores + if (!isa(val.getType())) { + auto sMemRef = + PtrAnalysis::getScalarMemRef(op.getPtr(), ptr, loc, rewriter); + auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext()); + rewriter.create(loc, val, sMemRef, zeroMap, + ValueRange{}); + rewriter.eraseOp(op); + return success(); + } + + // 1. Simple case where no mask is used. + if (!mask) { + auto storeOp = rewriter.create( + loc, val, ptr); + storeOp.setWritable(true); + rewriter.eraseOp(op); + return success(); + } + + // 2. Continuous masked stores. + // Analyze the mask operand to determine at runtime the size of the data we + // are moving. + MaskState mstate; + auto isContMask = mstate.parse(mask, loc, rewriter); + + if (isContMask.failed()) + return failure(); + + auto srcSlice = mstate.getExtractSlice(val, loc, rewriter); + auto dstSubview = mstate.getSubview(ptr, loc, rewriter); + + auto storeOp = rewriter.create( + loc, srcSlice, dstSubview); + storeOp.setWritable(true); + rewriter.eraseOp(op); + + return success(); + } +}; + +struct LoopConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(scf::ForOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + llvm::SmallDenseMap knownPtrs; + PtrAnalysis::IndexMapSet + levelToBlockArgIndex; // level -> set of block arg index to be replaced + + PtrAnalysis::rewriteForOp(op, rewriter, levelToBlockArgIndex, 0, knownPtrs); + return success(); + } +}; + +struct YieldConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(scf::YieldOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, adaptor.getOperands()); + return success(); + } +}; + +// Remove all Meta ops except for AddPtr which is handled by AddPtrConverter. +// Use benefit == 10 to ensure that this pattern always takes precedence over +// other patterns. +struct MetaOpConverter : public RewritePattern { +private: + // UseAnalysis will tag operations whose results are used only as meta-data + // with "MetaUse" tag. + bool isMetaUse(Operation *op) const { return op->hasAttr("MetaUse"); } + +public: + MetaOpConverter(MLIRContext *context) + : RewritePattern(MatchAnyOpTypeTag(), /*benefit=*/10, context) {} + + LogicalResult matchAndRewrite(Operation *op, + PatternRewriter &rewriter) const final { + + if (isa(op)) { + return rewriter.notifyMatchFailure(op, + "AddPtrOp will be handled separately"); + } + + if (isMetaUse(op)) { + rewriter.eraseOp(op); + return success(); + } + + return rewriter.notifyMatchFailure(op, "requires meta ops"); + } +}; + +struct UnrealizedCastConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(UnrealizedConversionCastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.eraseOp(op); + return success(); + } +}; + +//----------------------------- +// End of monolithic only +//----------------------------- + +struct SplatConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::SplatOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto opType = cast(op.getType()); + auto loc = op.getLoc(); + + auto init = rewriter.create(loc, opType.getShape(), + opType.getElementType()); + + auto filledTensor = + rewriter + .create(loc, ValueRange{adaptor.getSrc()}, + ValueRange{init}) + .result(); + + rewriter.replaceOp(op, filledTensor); + return success(); + } +}; + +struct BroadcastConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + SmallVector getBroadcastDims(RankedTensorType src, + RankedTensorType dst) const { + SmallVector broadcastDims; + auto srcShape = src.getShape(); + auto dstShape = dst.getShape(); + + for (size_t i = 0; i < srcShape.size(); i++) { + if (dstShape[i] != srcShape[i]) { + assert(srcShape[i] == 1); + broadcastDims.push_back(i); + } + } + assert(!broadcastDims.empty() && "cannot identify broadcast dimension"); + return broadcastDims; + } + + // Broadcasts input tensor based on TosaToLinalg's broadcastToShape + AffineMap getBroadcastAffineMap(MLIRContext *context, + ArrayRef inputShape, + ArrayRef broadcastToShape) const { + + assert(broadcastToShape.size() >= inputShape.size()); + + // Create affine map and shapes for tensor initialization. + SmallVector outExpr; + + size_t diff = broadcastToShape.size() - inputShape.size(); + for (size_t i = 0; i < broadcastToShape.size(); i++) { + if (i < diff) { + continue; + } + size_t j = i - diff; + if (inputShape[j] == 1) { + // Broadcast singleton dimension + outExpr.push_back(mlir::getAffineConstantExpr(0, context)); + continue; + } + // Non-broadcast case + outExpr.push_back(mlir::getAffineDimExpr(i, context)); + } + return AffineMap::get(broadcastToShape.size(), 0, outExpr, context); + } + +public: + LogicalResult + matchAndRewrite(triton::BroadcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + + assert(op->getNumResults() == 1 && "code assumes single result!"); + RankedTensorType sourceType = + cast(adaptor.getSrc().getType()); + RankedTensorType resultType = cast(op.getType()); + auto elementType = resultType.getElementType(); + size_t resultRank = resultType.getRank(); + + SmallVector indexingMaps; + indexingMaps.reserve(op->getNumOperands() + op->getNumResults()); + + indexingMaps.push_back(getBroadcastAffineMap( + op->getContext(), sourceType.getShape(), resultType.getShape())); + indexingMaps.append(op->getNumResults(), + rewriter.getMultiDimIdentityMap(resultRank)); + + assert(op->getNumResults() == 1 && "code assumes single result!"); + auto init = rewriter.create(loc, resultType.getShape(), + elementType); + + auto linalgOp = rewriter.create( + loc, op->getResultTypes(), ValueRange{adaptor.getSrc()}, + ValueRange{init}, indexingMaps, getNParallelLoopsAttrs(resultRank), + [&](OpBuilder &nestedBuilder, Location nestedLoc, + ValueRange blockArgs) { + Value opResult = blockArgs[0]; + nestedBuilder.create(loc, opResult); + }); + + linalgOp->setAttr("broadcastDims", + rewriter.getDenseI64ArrayAttr( + getBroadcastDims(sourceType, resultType))); + + rewriter.replaceOp(op, linalgOp->getResults()); + return success(); + } +}; + +struct ExpandDimsConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ExpandDimsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto src = adaptor.getSrc(); + auto srcRank = cast(src.getType()).getRank(); + auto resType = cast(op->getResultTypes()[0]); + SmallVector reassoc; + int64_t c = 0; + for (int64_t i = 0; i < srcRank; i++) { + ReassociationIndices g; + g.push_back(c++); + if (op.getAxis() == i) { + g.push_back(c++); + } else if (op.getAxis() == i + 1 && i == srcRank - 1) { + g.push_back(c++); + } + reassoc.push_back(g); + } + + auto expandShapeOp = rewriter.create( + op.getLoc(), resType, src, reassoc); + + rewriter.replaceOp(op, expandShapeOp.getResult()); + return success(); + } +}; + +struct TransposeConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::TransOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto source = adaptor.getSrc(); + + auto sourceType = cast(source.getType()); + auto sourceShape = sourceType.getShape(); + auto sourceRank = sourceType.getRank(); + + auto order = op.getOrder(); + SmallVector perm(order.begin(), order.end()); + + SmallVector transposedShape(sourceType.getShape()); + for (int i = 0; i < sourceRank; i++) + transposedShape[i] = sourceShape[perm[i]]; + + Value transposeInit = rewriter.create( + op->getLoc(), transposedShape, sourceType.getElementType()); + + Value transpose = rewriter + .create(op->getLoc(), source, + transposeInit, perm) + .getResults()[0]; + + rewriter.replaceOp(op, transpose); + return success(); + } +}; + +struct MakeRangeConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::MakeRangeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto type = cast(op.getResult().getType()); + auto shape = type.getShape(); + auto elementType = type.getElementType(); + auto context = rewriter.getContext(); + + assert(type.getShape().size() == 1 && + type.getElementType().getIntOrFloatBitWidth() == 32 && + "make range can only return 1D int32 tensor"); + + SmallVector indexingMaps{AffineMap::get( + /* dimCount */ 1, /* symbolCount */ 0, + SmallVector{mlir::getAffineDimExpr(0, context)}, context)}; + + auto init = rewriter.create(loc, shape, elementType); + auto linalgOp = rewriter.create( + loc, op->getResultTypes(), /* operands */ ValueRange{}, + ValueRange{init}, indexingMaps, getNParallelLoopsAttrs(1), + [&](OpBuilder &nestedBuilder, Location nestedLoc, + ValueRange blockArgs) { + Value start = rewriter.create( + loc, op.getStart()); // start + Value index = nestedBuilder.create(loc, 0); + Value res = nestedBuilder.create( + loc, type.getElementType(), + rewriter.create(loc, index, start)); + nestedBuilder.create(loc, res); + }); + + rewriter.replaceOp(op, linalgOp->getResults()); + return success(); + } +}; + +// FIXME: There is no triton::BarrierOp currently. +struct BarrierConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mlir::gpu::BarrierOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + + rewriter.create(loc); + rewriter.eraseOp(op); + return success(); + } +}; + +// Similar with triton-cpu. +struct PrintOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::PrintOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op->getLoc(); + // If the op has no operands, we can just print the prefix. + if (op.getNumOperands() == 0) { + rewriter.create(loc, TypeRange{}, op.getPrefix(), + op.getHex(), ValueRange{}, + llvm::SmallVector{}); + rewriter.eraseOp(op); + return success(); + } + + for (size_t i = 0; i < op.getNumOperands(); i++) { + Value operand = op.getOperands()[i]; + auto isSigned = {op.getIsSigned()[i]}; + // If the operand is not a ranked tensor, we should create a new tensor. + // See mlir/lib/Interfaces/DestinationStyleOpInterface.cpp#L39 + if (!isa(operand.getType())) { + // NOTE: Use tensor.from_elements, the arith.constant will not translate + // to linalg.fill + auto emptyTensor = rewriter.create( + loc, SmallVector{}, operand.getType()); + + auto operandTensor = rewriter.create( + loc, operand, emptyTensor, ValueRange{}); + rewriter.create(loc, operandTensor.getType(), + op.getPrefix(), op.getHex(), + operandTensor.getResult(), isSigned); + continue; + } + + auto operandType = cast(operand.getType()); + auto flattenTensor = operand; + if (operandType.getRank() != 1) { + SmallVector flatten_shape = {operandType.getNumElements()}; + auto targetType = + RankedTensorType::get(flatten_shape, operandType.getElementType()); + // NOTE: Avoid to create global constant tensors + SmallVector reassociation(1); + for (unsigned i = 0; i < operandType.getRank(); ++i) { + reassociation.front().push_back(i); + } + flattenTensor = rewriter.create( + loc, targetType, operand, reassociation); + } + + rewriter.create(loc, flattenTensor.getType(), op.getPrefix(), + op.getHex(), flattenTensor, isSigned); + } + + rewriter.eraseOp(op); + return success(); + } +}; + +struct BitcastConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::BitcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto inputType = op.getOperand().getType(); + auto resultType = op.getType(); + if (isa(inputType) || isa(resultType)) { + return success(); + } + auto arithBitcast = rewriter.create( + op.getLoc(), op.getType(), op.getOperand()); + + rewriter.replaceOp(op, arithBitcast.getResult()); + return success(); + } +}; + +struct CallConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::CallOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + SmallVector args = adaptor.getOperands(); + + // We need to pass extra arguments added by addProgramInfo which are + // num_programs and program_ids + if (FuncOp parentFunc = op->getParentOfType()) { + SymbolRefAttr calleeAttr = op.getCalleeAttr(); + StringRef calleeName = calleeAttr.getRootReference(); + + if (ModuleOp module = op->getParentOfType()) { + if (FuncOp calleeFunc = module.lookupSymbol(calleeName)) { + size_t argsNeed = calleeFunc.getFunctionType().getInputs().size(); + Block &entryBlock = parentFunc.front(); + auto parentInputs = entryBlock.getArguments(); + size_t argsParent = parentInputs.size(); + + if (argsNeed > args.size()) { + int missing = argsNeed - args.size(); + for (int i = 0; i < missing; i++) { + args.push_back(parentInputs[args.size()]); + } + } + } + } + } + + auto call = rewriter.create(op.getLoc(), op.getCallee(), + op.getResultTypes(), args); + + if (!call) { + op.emitError("Failed to create func::CallOp"); + return failure(); + } + + rewriter.replaceOp(op, call); + return success(); + } +}; + +struct FpToFpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::FpToFpOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto roundingMode = triton::RoundingMode::RTNE; // default + + auto roundingModeAttr = op.getRounding(); + if (roundingModeAttr.has_value()) { + roundingMode = roundingModeAttr.value(); + } + + assert(roundingMode != triton::RoundingMode::RTZ && + "Rounding Towards Zero is not supported"); + + Type resultType = op.getResult().getType(); + + auto operandWidth = getBitWidth(op.getOperand().getType()); + auto resultWidth = getBitWidth(resultType); + + assert(operandWidth.has_value() && resultWidth.has_value() && + "Not a float-like operand or result"); + + if (operandWidth.value() > resultWidth.value()) { + Value truncatedValue = rewriter.create( + op.getLoc(), resultType, op.getOperand()); + rewriter.replaceOp(op, truncatedValue); + return success(); + } + + Value extendedValue = rewriter.create( + op.getLoc(), resultType, op.getOperand()); + rewriter.replaceOp(op, extendedValue); + + return success(); + } +}; + +struct ClampConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ClampFOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + bool propagateNan = op.getPropagateNan() == triton::PropagateNan::ALL; + + Location loc = op.getLoc(); + Value x = adaptor.getOperands()[0]; + Value min = adaptor.getOperands()[1]; + Value max = adaptor.getOperands()[2]; + + Value clamp = x; + auto maxMin = min; + + if (propagateNan) { + // Handle NaN propagation + maxMin = rewriter.create(loc, x, min); + clamp = rewriter.create(loc, maxMin, max); + } else { + // No NaN propagation. + maxMin = rewriter.create(loc, x, min); + clamp = rewriter.create(loc, maxMin, max); + } + + rewriter.replaceOp(op, clamp); + + return success(); + } +}; + +struct PreciseSqrtConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::PreciseSqrtOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto replacement = + rewriter.create(op.getLoc(), adaptor.getOperands()); + + rewriter.replaceOp(op, replacement); + return success(); + } +}; + +struct PreciseDivConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::PreciseDivFOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto replacement = + rewriter.create(op.getLoc(), adaptor.getOperands()); + + rewriter.replaceOp(op, replacement); + return success(); + } +}; + +struct CatConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::CatOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto replacement = rewriter.create( + op.getLoc(), 0 /* concat dimension */, adaptor.getOperands()); + + rewriter.replaceOp(op, replacement); + + return success(); + } +}; + +struct SplitConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::SplitOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + Value input = op.getOperand(); + auto inputType = cast(input.getType()); + + Type resultType = op.getResults().front().getType(); + auto resultTensor = cast(resultType); + auto shape = inputType.getShape(); + + SmallVector offsets(shape.size(), rewriter.getIndexAttr(0)); + SmallVector strides(shape.size(), rewriter.getIndexAttr(1)); + SmallVector sizes = llvm::to_vector( + llvm::map_range(shape, [&](int64_t dim) -> OpFoldResult { + return rewriter.getIndexAttr(dim); + })); + + SmallVector results; + + for (int i = 0; i < 2; ++i) { + offsets.pop_back(); + sizes.pop_back(); + + offsets.push_back(rewriter.getIndexAttr(i)); + sizes.push_back(rewriter.getIndexAttr(1)); + Value slice = rewriter.create( + loc, resultTensor, input, offsets, sizes, strides); + results.push_back(slice); + } + + rewriter.replaceOp(op, results); + return success(); + } +}; + +struct JoinConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::JoinOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + ValueRange inputs = op.getOperands(); + + auto resultType = cast(op.getResult().getType()); + + auto loc = op.getLoc(); + Value result = rewriter.create( + loc, resultType.getShape(), resultType.getElementType()); + + auto shape = resultType.getShape(); + + SmallVector offsets(shape.size(), rewriter.getIndexAttr(0)); + SmallVector strides(shape.size(), rewriter.getIndexAttr(1)); + SmallVector sizes = llvm::to_vector( + llvm::map_range(shape, [&](int64_t dim) -> OpFoldResult { + return rewriter.getIndexAttr(dim); + })); + + for (int i = 0; i < 2; ++i) { + offsets.pop_back(); + sizes.pop_back(); + + offsets.push_back(rewriter.getIndexAttr(i)); + sizes.push_back(rewriter.getIndexAttr(1)); + result = rewriter.create(loc, inputs[i], result, + offsets, sizes, strides); + } + + rewriter.replaceOp(op, result); + + return success(); + } +}; + +struct MulHiUIOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::MulhiUIOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + + auto mulResult = + rewriter.create(loc, adaptor.getOperands()); + rewriter.replaceOp(op, mulResult.getHigh()); + + return success(); + } +}; + +// Check if tensor is all zeros +// Returns true if all tensor elements are zero +// Returns false if tensor is not all zeros or cannot be determined +bool isZeroTensor(Value &v, bool integers) { + // Check SplatOp case + if (auto splatOp = v.getDefiningOp()) { + if (auto constOp = splatOp.getSrc().getDefiningOp()) { + // Check floating point constant + if (auto val = dyn_cast(constOp.getValue())) { + return val.getValueAsDouble() == 0.; + } + // Check integer constant + if (auto val = dyn_cast(constOp.getValue())) { + return val.getValue() == 0; + } + } + return false; + } + + // Check ConstantOp case + if (auto constOp = v.getDefiningOp()) { + if (auto denseAttr = dyn_cast(constOp.getValue())) { + if (denseAttr.isSplat()) { + // Check zero value based on type + if (integers) + return denseAttr.getSplatValue().isZero(); + return denseAttr.getSplatValue().isZero(); + } + } + } + + return false; +} + +struct MatmulConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::DotOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp( + op, ValueRange{adaptor.getA(), adaptor.getB()}, + ValueRange{adaptor.getC()}); + return success(); + } +}; + +struct DotScaledConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::DotScaledOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + + // Get operands + auto a = op.getA(); + auto b = op.getB(); + auto c = op.getC(); + Value aScale = op.getAScale(); + Value bScale = op.getBScale(); + auto aElemType = op.getAElemTypeAttr(); + auto bElemType = op.getBElemTypeAttr(); + auto fastMath = op.getFastMathAttr(); + + // Get type information + auto aType = a.getType(); + auto bType = b.getType(); + auto dstType = cast(op.getType()); + auto elementType = dstType.getElementType(); + + // Create initial zero tensor + auto init = + rewriter.create(loc, dstType.getShape(), elementType); + TypedAttr constantAttr = + static_cast(rewriter.getFloatAttr(elementType, 0)); + auto zero = rewriter.create( + op.getLoc(), elementType, constantAttr); + auto zeroes = + rewriter.create(loc, ValueRange{zero}, ValueRange{init}) + .result(); + + // Perform scaled dot product operation + Value res = rewriter + .create(loc, TypeRange{op.getType()}, a, + aScale, b, bScale, zeroes, + aElemType, bElemType, fastMath) + .getResult(0); + + // Check if C needs to be added + bool skipC = isZeroTensor(c, false); + if (!skipC) { + res = rewriter.create(loc, c, res); + } + + rewriter.replaceOp(op, res); + return success(); + } +}; + +struct ReduceConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +private: + llvm::SmallVector getRedOps(triton::ReduceOp redOp) const { + auto reduceBlock = redOp.getBody(); + return llvm::map_to_vector(reduceBlock->without_terminator(), + [](Operation &op) { return &op; }); + } + + Value getRedElement(Value lhs, Value rhs, const Location loc, + Operation *redOp, OpBuilder &b) const { + return llvm::TypeSwitch(redOp) + .Case([&](auto redOp) { + return b.create(loc, lhs, rhs); + }) + .Case([&](auto redOp) { + return b.create(loc, lhs, rhs); + }) + .Default([](Operation *op) { + op->dump(); + llvm_unreachable("Reduction op not yet supported"); + return nullptr; + }); + } + + LogicalResult + convertMultiOpsReduction(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const { + + auto loc = op->getLoc(); + int64_t axis = op.getAxis(); + + SmallVector initTensors; + for (auto result : op->getResultTypes()) { + SmallVector shape = + isa(result) + ? SmallVector( + cast(result).getShape().begin(), + cast(result).getShape().end()) + : SmallVector{}; + Type type = isa(result) + ? cast(result).getElementType() + : result; + initTensors.push_back(rewriter.create(loc, shape, type)); + } + auto reduceOp = rewriter.create( + loc, op.getSrcs(), initTensors, SmallVector{axis}, + [&](OpBuilder &opBuilder, Location loc, ValueRange inputs) { + auto reduceBlock = op.getBody(); + IRMapping mapping; + mapping.map(reduceBlock->getArguments(), inputs); + for (auto &innerOp : reduceBlock->without_terminator()) { + opBuilder.clone(innerOp, mapping); + } + auto yield = reduceBlock->getTerminator(); + auto results = + llvm::map_to_vector(yield->getOperands(), [&](Value val) { + return mapping.lookup(val); + }); + opBuilder.create(loc, results); + }); + SmallVector results; + for (int i = 0; i < reduceOp->getNumResults(); i++) { + Value finalResult = + (isa(op->getResultTypes()[i])) + ? reduceOp->getResults()[i] + : rewriter + .create(loc, op->getResultTypes()[i], + reduceOp->getResults()[i]) + ->getResults()[0]; + results.push_back(finalResult); + } + rewriter.replaceOp(op, results); + return success(); + } + + LogicalResult + convertToLinalgReduce(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto source = adaptor.getOperands().front(); + auto sourceType = cast(source.getType()); + auto elemType = sourceType.getElementType(); + auto resType = op.getResult().front().getType(); + auto loc = op.getLoc(); + auto reductionOps = getRedOps(op); + + if (reductionOps.size() != 1) + return convertMultiOpsReduction(op, adaptor, rewriter); + + // Reduction of arbitrary operations isn't supported because using the first + // element across the reduction dimension requires us to iterate over a + // subview that skips over each first element. + if (!isTritonAllowedReductionOp(reductionOps.front())) { + return rewriter.notifyMatchFailure( + op, "Only support lowering reduction with body " + "containing 1 max(i/f) or addf."); + } + + auto rop = reductionOps.front(); + auto axis = op.getAxis(); + auto isVectorReduce = sourceType.getRank() == 1; + + auto accBaseConstOp = getRedBaseConstOp(rewriter, rop, elemType); + Value initTensor; + + if (isVectorReduce) { + // The affine vectorizer cannot vectorize affine loops generated from + // linalg.reduce for the vector reduce case, so we must rewrite the + // linalg.reduce to affine loops manually. Here we lower to AllocTensor + // directly instead of EmptyOp so that the subsequent pass can recognize + // the patterns (EmptyOp is susceptible to being CSE'd away, making it + // harder to match the patterns correctly). + initTensor = rewriter.create( + loc, RankedTensorType::get({}, elemType), ValueRange{}); + initTensor = rewriter.create(loc, accBaseConstOp, + initTensor, ValueRange{}); + } else { + Value init = rewriter.create( + loc, cast(resType).getShape(), elemType); + initTensor = rewriter + .create(loc, ValueRange{accBaseConstOp}, + ValueRange{init}) + .result(); + } + + Value finalResult = + rewriter + .create( + loc, ValueRange{source}, ValueRange{initTensor}, + SmallVector{axis}, + [&](OpBuilder &opBuilder, Location loc, ValueRange inputs) { + assert(inputs.size() == 2); + Value result = + getRedElement(inputs[0], inputs[1], loc, rop, opBuilder); + opBuilder.create(loc, result); + }) + .getResult(0); + + if (sourceType.getRank() == 1) { + finalResult = + rewriter.create(loc, elemType, finalResult); + } + + rewriter.replaceOp(op, finalResult); + return success(); + } + +public: + LogicalResult + matchAndRewrite(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto sourceType = + cast(adaptor.getOperands().front().getType()); + assert(sourceType.hasRank() && "Expected input is " + "ranked"); + + int64_t axis = op.getAxis(); + assert(axis >= 0 && axis < sourceType.getRank() && + "Expected reduction " + "axis is within " + "operand's rank"); + + return convertToLinalgReduce(op, adaptor, rewriter); + } +}; + +template +class ArgMinMaxBaseConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + Value getInitTensor(ConversionPatternRewriter &rewriter, + ArrayRef shape, Value fillValue, + Location loc) const { + Value initTensor = + rewriter.create(loc, shape, fillValue.getType()); + return rewriter + .create(loc, ValueRange{fillValue}, + ValueRange{initTensor}) + .result(); + } + +public: + ArgMinMaxBaseConverter(MLIRContext *context) : OpConversionPattern(context) {} + bool isArgMin; + + LogicalResult matchReduction(ReduceOp op) const { + if (op.getBody()->getNumArguments() != 4) { + return failure(); + } + + auto block = op.getBody(); + auto ops = block->without_terminator(); + + Value currValue = block->getArgument(0); + Value currIndex = block->getArgument(1); + Value reduceValue = block->getArgument(2); + Value reduceIndex = block->getArgument(3); + + auto opsIt = ops.begin(); + Value indexSelectOp, valueSelectOp; + if (failed(matchArgMinMax(currValue, currIndex, reduceValue, reduceIndex, + opsIt, indexSelectOp, valueSelectOp, isArgMin))) { + return failure(); + } + + // matching: tt.reduce.return %16, %17 : f32, i32 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *opsIt << "\n"); + auto termOp = dyn_cast(*opsIt++); + if (termOp && termOp == block->getTerminator()) { + auto opnds = termOp.getOperands(); + if (opnds != ArrayRef{valueSelectOp, indexSelectOp}) { + return failure(); + } + } else { + return failure(); + } + + return success(); + } + + LogicalResult + matchAndRewrite(ReduceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override final { + if (failed(matchReduction(op))) + return failure(); + + auto loc = op.getLoc(); + + auto elemTypes = op.getElementTypes(); + + // Set the initial value of the rank-0 tensor containing + // the result value to either -inf or +inf depending on + // whether we're dealing with argmax or argmin + auto valueType = elemTypes[0]; + + auto valuesAccBaseVal = + valueType.isInteger() + ? rewriter.create( + loc, valueType, + rewriter.getIntegerAttr( + valueType, T::getBaseReductionIntValue( + elemTypes[0].getIntOrFloatBitWidth()))) + : rewriter.create( + loc, valueType, + rewriter.getFloatAttr(valueType, + T::getBaseReductionFloatValue())); + + // Set the initial value of the rank-0 tensor containing the index of the + // min or max value to -1 + auto indexType = elemTypes[1]; + auto indicesAccBaseVal = rewriter.create( + loc, indexType, rewriter.getIntegerAttr(indexType, -1)); + + // Get the shape of the resulting tensors (both for values and indices). If + // we are reducing to a single scalar, then the result's type is a tensor of + // rank-0, otherwise we can reuse the original result shape + auto valueResultType = dyn_cast(op.getType(0)); + const auto isScalarReduce = valueResultType == nullptr; + SmallVector reductionResultShape{ + isScalarReduce ? SmallVector{} + : SmallVector(valueResultType.getShape())}; + + SmallVector outputs{ + getInitTensor(rewriter, reductionResultShape, valuesAccBaseVal, loc), + getInitTensor(rewriter, reductionResultShape, indicesAccBaseVal, loc)}; + + auto linalgOp = rewriter.create( + loc, adaptor.getOperands(), outputs, + SmallVector{adaptor.getAxis()}, + [&](OpBuilder &b, Location loc, ValueRange inputs) { + assert(inputs.size() == 4); + + auto tritonReduceBlock = op.getBody(); + IRMapping mapping; + mapping.map(tritonReduceBlock->getArguments(), inputs); + + for (auto &op : tritonReduceBlock->without_terminator()) { + b.clone(op, mapping); + } + + auto tritonYield = tritonReduceBlock->getTerminator(); + auto results = + llvm::map_to_vector(tritonYield->getOperands(), [&](Value val) { + return mapping.lookup(val); + }); + b.create(loc, results); + }); + + if (isScalarReduce) { + SmallVector reduceResults{ + rewriter.create( + loc, valueType, linalgOp.getResults()[0], ValueRange{}), + rewriter.create( + loc, indexType, linalgOp.getResults()[1], ValueRange{})}; + rewriter.replaceOp(op, reduceResults); + } else { + rewriter.replaceOp(op, linalgOp); + } + return success(); + } +}; + +struct ArgMaxConverter : public ArgMinMaxBaseConverter { + // TODO: min value according the bitwidth? Now this is not used in wafer + static int64_t getBaseReductionIntValue(int bitwidth) { + return llvm::minIntN(bitwidth); + } + static float getBaseReductionFloatValue() { + return -std::numeric_limits::infinity(); + } + + ArgMaxConverter(MLIRContext *context) : ArgMinMaxBaseConverter(context) { + isArgMin = false; + } +}; + +struct ArgMinConverter : public ArgMinMaxBaseConverter { + static int64_t getBaseReductionIntValue(int bitwidth) { + return llvm::maxIntN(bitwidth); + } + static float getBaseReductionFloatValue() { + return std::numeric_limits::infinity(); + } + + ArgMinConverter(MLIRContext *context) : ArgMinMaxBaseConverter(context) { + isArgMin = true; + } +}; + +// get_program_id and get_num_programs: +// When launching triton kernels, we pass 6 additional arguments to indicate +// num_programs and program_id. Amongst those six, we have 3 arguments +// correspond to each axis for num_programs followed by 3 additional arguments +// for program_id. +// +// For instance, with triton kernel example_kernel(a, b, c), we have: +// example_kernel( +// a, b, c, +// num_programs_axis_0, +// num_programs_axis_1, +// num_programs_axis_2, +// program_id_axis_0, +// program_id_axis_1, +// program_id_axis_2, +// ) +// +struct GetProgramIDConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + static uint32_t constexpr LAUNCH_GRID_RANK = + getMaxEnumValForProgramIDDim() + 1; + +public: + LogicalResult + matchAndRewrite(triton::GetProgramIdOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto axis = (uint32_t)op.getAxis(); + assert(axis < LAUNCH_GRID_RANK && "program_id expects " + "axis to be either 0, " + "1, or 2"); + + auto func = op->getParentOfType(); + auto numArgs = func.getNumArguments(); + auto id = func.getArgument(numArgs - LAUNCH_GRID_RANK + axis); + + rewriter.replaceOp(op, id); + return success(); + } +}; + +struct GetNumProgramsConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +private: + static uint32_t constexpr LAUNCH_GRID_RANK = + getMaxEnumValForProgramIDDim() + 1; + +public: + GetNumProgramsConverter(MLIRContext *context) + : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(triton::GetNumProgramsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto axis = (uint32_t)op.getAxis(); + assert(axis < LAUNCH_GRID_RANK && "program_id expects " + "axis to be either 0, " + "1, or 2"); + + auto func = op->getParentOfType(); + auto numArgs = func.getNumArguments(); + auto id = func.getArgument(numArgs - LAUNCH_GRID_RANK * 2 + axis); + + rewriter.replaceOp(op, id); + return success(); + } +}; + +// Convert a pair of cmpf and select to either min or max. +// Leave the pattern as simple as possible because triton has plans to emit +// min and max directly. +template +struct MinMaxConverter : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + MinMaxConverter(MLIRContext *context) + : OpRewritePattern(context, /*benefit=*/10) {} + + LogicalResult matchAndRewrite(CmpOp cmpOp, + PatternRewriter &rewriter) const final { + if (!cmpOp.getResult().hasOneUse()) { + return failure(); + } + auto selectOp = + dyn_cast(*cmpOp.getResult().getUsers().begin()); + if (!selectOp) { + return failure(); + } + + if (!(cmpOp.getResult() == selectOp.getCondition() && + cmpOp.getLhs() == selectOp.getTrueValue() && + cmpOp.getRhs() == selectOp.getFalseValue())) { + return failure(); + } + + rewriteOpWithMinMax(rewriter, cmpOp, selectOp, cmpOp.getPredicate()); + rewriter.eraseOp(cmpOp); + + return success(); + } + + void rewriteOpWithMinMax(PatternRewriter &rewriter, arith::CmpFOp cmpOp, + arith::SelectOp selectOp, + arith::CmpFPredicate pred) const { + switch (pred) { + case arith::CmpFPredicate::OGT: + case arith::CmpFPredicate::OGE: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + case arith::CmpFPredicate::OLT: + case arith::CmpFPredicate::OLE: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + default: + llvm_unreachable("Unhandled predicate"); + } + } + + void rewriteOpWithMinMax(PatternRewriter &rewriter, arith::CmpIOp cmpOp, + arith::SelectOp selectOp, + arith::CmpIPredicate pred) const { + switch (pred) { + case arith::CmpIPredicate::sgt: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + case arith::CmpIPredicate::ugt: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + case arith::CmpIPredicate::slt: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + case arith::CmpIPredicate::ult: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + default: + llvm_unreachable("Unhandled predicate"); + } + } +}; + +struct DenseConstantConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(arith::ConstantOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto attr = cast(op.getValue()); + auto loc = op.getLoc(); + + auto splatConst = arith::ConstantOp::materialize( + rewriter, attr.getSplatValue(), attr.getElementType(), loc); + + auto init = rewriter.create( + loc, cast(op.getResult().getType()).getShape(), + attr.getElementType()); + + rewriter.replaceOpWithNewOp(op, ValueRange{splatConst}, + ValueRange{init}); + + return success(); + } +}; + +class CumSumConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + // CumSum is a specific instance of Scan that looks like the following: + // %1 = "tt.scan"(%0) <{axis = 1 : i32}> ({ + // ^bb0(%arg0: f32, %arg1: f32): + // %2 = arith.addf %arg0, %arg1 : f32 + // tt.scan.return %2 : f32 + // }) : (tensor<4x4xf32>) -> tensor<4x4xf32> + bool isCumSum(triton::ScanOp op) const { + auto scanBlock = op.getBody(); + auto ops = llvm::map_to_vector(scanBlock->without_terminator(), + [](Operation &op) { return &op; }); + + if (ops.size() != 1) { + return false; + } + + auto addOp = ops.front(); + if (isa(addOp)) { + if (addOp->getResult(0) != scanBlock->getTerminator()->getOperand(0)) { + return false; + } + + auto blockArgs = + llvm::map_range(scanBlock->getArguments(), [](BlockArgument arg) { + return dyn_cast(arg); + }); + + auto addArgs = addOp->getOperands(); + + return DenseSet(blockArgs.begin(), blockArgs.end()) == + DenseSet(addArgs.begin(), addArgs.end()); + } + + return false; + } + +public: + LogicalResult + matchAndRewrite(triton::ScanOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!isCumSum(op)) { + return rewriter.notifyMatchFailure( + op, "Only support cumsum variant of scan op"); + } + + auto input = op.getOperand(0); + auto axis = op.getAxis(); + auto type = dyn_cast(input.getType()); + + if (type.getRank() != 1 && type.getRank() != 2 && + axis != type.getRank() - 1) { + return rewriter.notifyMatchFailure( + op, "Only support lowering scan op to cumsum with rank " + "= {1, 2} and axis = rank - 1"); + } + + Value init = rewriter.create(op.getLoc(), type.getShape(), + type.getElementType()); + + rewriter.replaceOpWithNewOp( + op, input, rewriter.getUI32IntegerAttr(axis), init); + + return success(); + } +}; + +class AddPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::AddPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto resType = op.getResult().getType(); + assert(isa(resType)); + auto rank = cast(resType).getRank(); + SmallVector indexingMaps( + /*numResult + numOperands*/ 3, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + SmallVector outputs = {op.getPtr()}; + rewriter.replaceOpWithNewOp( + op, op->getResultTypes(), op->getOperands(), outputs, indexingMaps, + iteratorTypes, + [&](OpBuilder &builder, Location loc, ValueRange regionArgs) { + auto resultTypes = llvm::to_vector<6>( + llvm::map_range(op->getResultTypes(), [](Type type) { + return cast(type).getElementType(); + })); + auto *scalarOp = + builder.create(loc, op->getName().getIdentifier(), + regionArgs.take_front(op->getNumOperands()), + resultTypes, op->getAttrs()); + builder.create(loc, scalarOp->getResults()); + }); + return success(); + } +}; + +template +class TensorOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(OpType op, typename OpType::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto resultTensorType = dyn_cast(op.getResult().getType()); + if (!resultTensorType) { + return failure(); + } + auto rank = resultTensorType.getRank(); + SmallVector indexingMaps( + op->getNumResults() + op->getNumOperands(), + rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + SmallVector outputs = {rewriter.create( + op->getLoc(), resultTensorType.getShape(), + resultTensorType.getElementType())}; + rewriter.replaceOpWithNewOp( + op, op->getResultTypes(), op->getOperands(), outputs, indexingMaps, + iteratorTypes, + [&](OpBuilder &builder, Location loc, ValueRange regionArgs) { + auto resultTypes = llvm::map_to_vector( + op->getResultTypes(), [](Type type) { + return cast(type).getElementType(); + }); + auto *scalarOp = builder.create( + loc, op->getName().getIdentifier(), + regionArgs.take_front(op->getNumOperands()), resultTypes, + op->getAttrs()); + builder.create(loc, scalarOp->getResults()); + }); + return success(); + } +}; + +class StorePtrToLinalgConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto storeTensorType = dyn_cast(op.getValue().getType()); + if (!storeTensorType) { + return failure(); + } + auto rank = storeTensorType.getRank(); + SmallVector indexingMaps( + op->getNumResults() + op.getNumOperands(), + rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + rewriter.replaceOpWithNewOp( + op, op->getResultTypes(), op->getOperands(), ValueRange{}, indexingMaps, + iteratorTypes, + [&](OpBuilder &builder, Location loc, ValueRange regionArgs) { + auto resultTypes = llvm::map_to_vector( + op->getResultTypes(), [](Type type) { + return cast(type).getElementType(); + }); + auto *scalarOp = builder.create( + loc, op->getName().getIdentifier(), + regionArgs.take_front(op->getNumOperands()), resultTypes, + op->getAttrs()); + builder.create(loc, scalarOp->getResults()); + }); + return success(); + } +}; + +class ReshapeConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::ReshapeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto input = op.getSrc(); + auto output = op.getResult(); + + auto inputType = input.getType(); + auto outputType = output.getType(); + if (!outputType.hasStaticShape()) { + return failure(); + } + + if (auto maybeReassociationMap = + getReassociationIndicesForReshape(inputType, outputType)) { + auto reassociationMap = *maybeReassociationMap; + if (outputType.getRank() < inputType.getRank()) { + rewriter.replaceOpWithNewOp( + op, outputType, input, reassociationMap); + } else { + rewriter.replaceOpWithNewOp( + op, outputType, input, reassociationMap); + } + return success(); + } + + ArrayRef outputShape = outputType.getShape(); + + auto shape = rewriter.create( + loc, rewriter.getI64TensorAttr(outputShape)); + rewriter.replaceOpWithNewOp(op, outputType, input, + shape); + + return success(); + } +}; + +struct GatherConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::GatherOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + auto resultType = cast(op.getResult().getType()); + Value dstInit = rewriter.create( + op.getLoc(), resultType.getShape(), resultType.getElementType()); + + auto gatherOp = + rewriter.create(op.getLoc(), op.getType(), op.getSrc(), + op.getIndices(), dstInit, op.getAxis()); + + rewriter.replaceOp(op, gatherOp.getResult()); + return success(); + } +}; + +class ExternElementwiseFiniteOpConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::ExternElementwiseOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + if (!op.getPure() || op.getSrcs().size() != 1 || + op.getSymbol() != "__nv_finitef") + return failure(); + + // get input value + Value x = op.getSrcs()[0]; + auto shapedType = cast(x.getType()); + auto elemType = shapedType.getElementType(); + + if (!elemType.isF32()) { + return failure(); + } + + // step 1: check x == x (not NaN) + Value notNaN = + rewriter.create(loc, arith::CmpFPredicate::OEQ, x, x); + + // step 2: check x + 1 != x (not Infinity) + Value one = rewriter.create( + loc, rewriter.getFloatAttr(elemType, 1.0f)); + + // create constant tensor + auto constTensor = + rewriter.create(loc, shapedType.getShape(), elemType); + auto filledTensor = rewriter + .create(loc, ValueRange{one}, + ValueRange{constTensor}) + .getResult(0); + + // x + 1 + Value xPlusOne = rewriter.create(loc, x, filledTensor); + // x + 1 != x + Value notInf = rewriter.create( + loc, arith::CmpFPredicate::ONE, xPlusOne, x); + + // step 3: combine two conditions (notNaN && notInf) + Value isFinite = rewriter.create(loc, notNaN, notInf); + + rewriter.replaceOp(op, isFinite); + return success(); + } +}; + +class ExternElementwiseFmodOpConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::ExternElementwiseOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (op.getSrcs().size() != 2 || + (op.getSymbol() != "__nv_fmodf" && op.getSymbol() != "__nv_fmod")) { + return failure(); + } + + // get input values + Value input1 = op.getSrcs()[0]; + Value input2 = op.getSrcs()[1]; + auto elemType = cast(input1.getType()).getElementType(); + + if (!elemType.isF32()) { + return failure(); + } + rewriter.replaceOpWithNewOp(op, input1, input2); + return success(); + } +}; + +class ExternElementwiseBinaryOpConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::ExternElementwiseOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + if (!op.getPure() || op.getSrcs().size() != 2) + return failure(); +#define POPULATE_BINARY_OP(FUNC_NAME, DST_OP) \ + if (!op.getSymbol().compare(FUNC_NAME)) { \ + rewriter.replaceOpWithNewOp(op, op.getSrcs()[0], op.getSrcs()[1]); \ + return success(); \ + } + + POPULATE_BINARY_OP("__nv_atan2f", math::Atan2Op); + POPULATE_BINARY_OP("__nv_atan2", math::Atan2Op); + POPULATE_BINARY_OP("__nv_powf", math::PowFOp); + POPULATE_BINARY_OP("__nv_pow", math::PowFOp); + POPULATE_BINARY_OP("__nv_powif", math::FPowIOp); + POPULATE_BINARY_OP("__nv_powi", math::FPowIOp); + +#undef POPULATE_BINARY_OP + return failure(); + } +}; + +class ExternElementwiseUnaryOpConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::ExternElementwiseOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + if (!op.getPure() || op.getSrcs().size() != 1) + return failure(); +#define POPULATE_UNARY_OP(FUNC_NAME, DST_OP) \ + if (!op.getSymbol().compare(FUNC_NAME)) { \ + rewriter.replaceOpWithNewOp(op, op.getSrcs()[0]); \ + return success(); \ + } + + POPULATE_UNARY_OP("__nv_fabsf", math::AbsFOp); + POPULATE_UNARY_OP("__nv_fabs", math::AbsFOp); + POPULATE_UNARY_OP("__nv_sinf", math::SinOp); + POPULATE_UNARY_OP("__nv_sin", math::SinOp); + POPULATE_UNARY_OP("__nv_cosf", math::CosOp); + POPULATE_UNARY_OP("__nv_cos", math::CosOp); + POPULATE_UNARY_OP("__nv_tanf", math::TanOp); + POPULATE_UNARY_OP("__nv_tan", math::TanOp); + POPULATE_UNARY_OP("__nv_asinf", math::AsinOp); + POPULATE_UNARY_OP("__nv_asin", math::AsinOp); + POPULATE_UNARY_OP("__nv_acosf", math::AcosOp); + POPULATE_UNARY_OP("__nv_acos", math::AcosOp); + POPULATE_UNARY_OP("__nv_atanf", math::AtanOp); + POPULATE_UNARY_OP("__nv_atan", math::AtanOp); + POPULATE_UNARY_OP("__nv_sinhf", math::SinhOp); + POPULATE_UNARY_OP("__nv_sinh", math::SinhOp); + POPULATE_UNARY_OP("__nv_coshf", math::CoshOp); + POPULATE_UNARY_OP("__nv_cosh", math::CoshOp); + POPULATE_UNARY_OP("__nv_tanhf", math::TanhOp); + POPULATE_UNARY_OP("__nv_tanh", math::TanhOp); + POPULATE_UNARY_OP("__nv_acoshf", math::AcoshOp); + POPULATE_UNARY_OP("__nv_acosh", math::AcoshOp); + POPULATE_UNARY_OP("__nv_asinhf", math::AsinhOp); + POPULATE_UNARY_OP("__nv_asinh", math::AsinhOp); + POPULATE_UNARY_OP("__nv_atanhf", math::AtanhOp); + POPULATE_UNARY_OP("__nv_atanhf", math::AtanhOp); + POPULATE_UNARY_OP("__nv_logf", math::LogOp); + POPULATE_UNARY_OP("__nv_log", math::LogOp); + POPULATE_UNARY_OP("__nv_log10f", math::Log10Op); + POPULATE_UNARY_OP("__nv_log10", math::Log10Op); + POPULATE_UNARY_OP("__nv_log1pf", math::Log1pOp); + POPULATE_UNARY_OP("__nv_log1p", math::Log1pOp); + POPULATE_UNARY_OP("__nv_expf", math::ExpOp); + POPULATE_UNARY_OP("__nv_exp", math::ExpOp); + POPULATE_UNARY_OP("__nv_exp2f", math::Exp2Op); + POPULATE_UNARY_OP("__nv_exp2", math::Exp2Op); + POPULATE_UNARY_OP("__nv_erff", math::ErfOp); + POPULATE_UNARY_OP("__nv_erf", math::ErfOp); + POPULATE_UNARY_OP("__nv_sqrtf", math::SqrtOp); + POPULATE_UNARY_OP("__nv_sqrt", math::SqrtOp); + POPULATE_UNARY_OP("__nv_rsqrtf", math::RsqrtOp); + POPULATE_UNARY_OP("__nv_rsqrt", math::RsqrtOp); + POPULATE_UNARY_OP("__nv_ceilf", math::CeilOp); + POPULATE_UNARY_OP("__nv_ceil", math::CeilOp); + POPULATE_UNARY_OP("__nv_floorf", math::FloorOp); + POPULATE_UNARY_OP("__nv_floor", math::FloorOp); + POPULATE_UNARY_OP("__nv_truncf", math::TruncOp); + POPULATE_UNARY_OP("__nv_trunc", math::TruncOp); + POPULATE_UNARY_OP("__nv_isnanf", math::IsNaNOp); + POPULATE_UNARY_OP("__nv_isnand", math::IsNaNOp); + POPULATE_UNARY_OP("__nv_isinff", math::IsInfOp); + POPULATE_UNARY_OP("__nv_isinfd", math::IsInfOp); + POPULATE_UNARY_OP("__nv_rintf", math::RoundOp); + POPULATE_UNARY_OP("__nv_rint", math::RoundOp); + +#undef POPULATE_UNARY_OP + return failure(); + } +}; + +static void populateExternElementwiseOpToMLIROps(RewritePatternSet &patterns) { + patterns + .add(patterns.getContext()); +} + +struct HistogramOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::HistogramOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto src = rewriter.getRemappedValue(op.getSrc()); + + auto srcTy = dyn_cast(src.getType()); + auto resTy = dyn_cast(op.getType()); + + auto flattenTensor = src; + // NOTE: Triton only support rank1. But flatten the tensor to 1D can always + // implement the histogram operation. + if (srcTy.getRank() != 1) { + SmallVector flatten_shape = {srcTy.getNumElements()}; + auto targetType = + RankedTensorType::get(flatten_shape, srcTy.getElementType()); + + auto shapeAttr = rewriter.getI64TensorAttr(flatten_shape); + auto shapeConst = rewriter.create(loc, shapeAttr); + flattenTensor = rewriter.create( + loc, targetType, flattenTensor, shapeConst); + } + + Value zero = rewriter.create( + loc, resTy, rewriter.getZeroAttr(resTy)); + Value one = rewriter.create(loc, resTy, + rewriter.getOneAttr(resTy)); + RankedTensorType cmpVecTy = + RankedTensorType::get(resTy.getShape(), srcTy.getElementType()); + + Value allocatedRangeVec = rewriter.create( + loc, resTy.getShape(), cmpVecTy.getElementType()); + SmallVector indexingMaps{AffineMap::get( + /* dimCount */ 1, /* symbolCount */ 0, + SmallVector{ + mlir::getAffineDimExpr(0, rewriter.getContext())}, + rewriter.getContext())}; + auto linalgOp = rewriter.create( + loc, cmpVecTy, /* operands */ ValueRange{}, + ValueRange{allocatedRangeVec}, indexingMaps, getNParallelLoopsAttrs(1), + [&](OpBuilder &nestedBuilder, Location nestedLoc, + ValueRange blockArgs) { + Value index = nestedBuilder.create(loc, 0); + Value res = nestedBuilder.create( + loc, resTy.getElementType(), index); + nestedBuilder.create(loc, res); + }); + + Value res = zero; + + // Create loop bounds + Value lowerBound = rewriter.create(loc, 0); + Value upperBound = + rewriter.create(loc, srcTy.getNumElements()); + Value step = rewriter.create(loc, 1); + + auto forOp = rewriter.create( + loc, lowerBound, upperBound, step, ValueRange{res}, + [&](OpBuilder &builder, Location loc, Value iv, ValueRange iterArgs) { + Value currentRes = iterArgs[0]; + // Extract element at current index + Value elem = + builder.create(loc, flattenTensor, iv); + SmallVector elems(resTy.getNumElements(), elem); + // Create a splat of the element + Value elemVec = builder.create(loc, elems); + + // Compare with range vector + Value mask = builder.create( + loc, arith::CmpIPredicate::eq, elemVec, linalgOp.getResult(0)); + // Select based on mask + Value delta = + builder.create(loc, resTy, mask, one, zero); + // Add to running result + Value newRes = builder.create(loc, currentRes, delta); + // Yield the updated result + builder.create(loc, newRes); + }); + + // Replace the original op with the final histogram result + rewriter.replaceOp(op, forOp.getResults()[0]); + return success(); + } + + TypedAttr makeRangeAttr(RankedTensorType resTy, + ConversionPatternRewriter &rewriter) const { + Type elemTy = resTy.getElementType(); + if (elemTy.isInteger(32)) { + SmallVector range(resTy.getShape()[0]); + std::iota(range.begin(), range.end(), 0); + return rewriter.getI32TensorAttr(range); + } else if (elemTy.isInteger(64)) { + SmallVector range(resTy.getShape()[0]); + std::iota(range.begin(), range.end(), 0); + return rewriter.getI64TensorAttr(range); + } else { + llvm_unreachable( + "unsupported src elem type for histogram (expected i32 or i64)"); + } + } +}; + +struct ScanOpConverter + : public triton::ReduceScanOpConversionBase { +private: + using ReduceScanOpConversionBase::ReduceScanOpConversionBase; + + SmallVector + lower1DInput(ValueRange inputs, ScanOp op, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + Region &combineOp = op.getRegion(); + bool reverse = op.getReverse(); + + auto shape = cast(inputs[0].getType()).getShape(); + SmallVector res(inputs.size()); + std::transform(inputs.begin(), inputs.end(), res.begin(), [&](auto val) { + auto inputType = cast(val.getType()); + assert(inputType.getShape() == shape && + "All 1D input tensors must have the same shape"); + return rewriter.create(loc, shape, + inputType.getElementType()); + }); + + SmallVector acc(inputs.size()); + std::transform(inputs.begin(), inputs.end(), acc.begin(), [&](auto val) { + auto inputType = cast(val.getType()); + return rewriter.create( + loc, rewriter.getZeroAttr(inputType.getElementType())); + }); + + // scf.for loop bounds and step + // NOTE: Scf only support positive step, so we need to handle reverse + Value lowerBound = rewriter.create(loc, 0); + Value upperBound = rewriter.create(loc, shape[0]); + Value step = rewriter.create(loc, 1); + // Use tuple to pass acc and res as loop-carried variables + SmallVector initVals; + for (auto v : res) + initVals.push_back(v); + for (auto v : acc) + initVals.push_back(v); + + auto forOp = rewriter.create( + loc, lowerBound, upperBound, step, initVals, + [&](OpBuilder &b, Location loc, Value iv, ValueRange iterArgs) { + // iterArgs: [res..., acc...] + SmallVector curRes(iterArgs.begin(), + iterArgs.begin() + inputs.size()); + SmallVector currAcc(iterArgs.begin() + inputs.size(), + iterArgs.end()); + + SmallVector idxIndex; + if (reverse) { + auto actualIdx = b.create(loc, upperBound, iv); + actualIdx = b.create(loc, actualIdx, step); + idxIndex.push_back(actualIdx); + } else { + idxIndex.push_back(iv); + } + SmallVector inputsElem(inputs.size()); + std::transform( + inputs.begin(), inputs.end(), inputsElem.begin(), [&](auto val) { + return rewriter.create(loc, val, idxIndex); + }); + + // Check if this is the first iteration + Value isFirstValue = b.create( + loc, arith::CmpIPredicate::eq, iv, lowerBound); + // FIXME: Bufferize will generate dynamic stride memref type, which + // cause memref::copy conversion failure + scf::IfOp ifOp = b.create( + loc, isFirstValue, + [&](OpBuilder &b, Location loc) { + b.create(loc, inputsElem); + }, + [&](OpBuilder &b, Location loc) { + b.create( + loc, accumulate(inputsElem, currAcc, combineOp, b)); + }); + currAcc = ifOp.getResults(); + + assert(acc.size() == inputs.size() && + "accumulate should return the same number of results as " + "inputs"); + for (int i = 0; i < res.size(); ++i) { + curRes[i] = rewriter.create(loc, currAcc[i], + curRes[i], idxIndex); + } + + SmallVector yieldVals; + for (auto v : curRes) + yieldVals.push_back(v); + for (auto v : currAcc) + yieldVals.push_back(v); + + b.create(loc, yieldVals); + }); + // Extract result tensors from forOp + SmallVector results; + for (size_t i = 0; i < inputs.size(); ++i) { + results.push_back(forOp.getResult(i)); + } + return results; + } + + SmallVector + lowerLeadingDimension(ValueRange inputs, ScanOp op, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + Region &combineOp = op.getRegion(); + bool reverse = op.getReverse(); + + auto shape = cast(inputs[0].getType()).getShape(); + + SmallVector resTypes; + for (const auto &resTy : op.getResultTypes()) { + resTypes.push_back(RankedTensorType::get( + shape, cast(resTy).getElementType())); + } + + // Initialize result tensors + SmallVector res(inputs.size()); + std::transform(inputs.begin(), inputs.end(), res.begin(), [&](auto val) { + auto inputType = cast(val.getType()); + auto valShape = inputType.getShape(); + assert(shape[0] == valShape[0] && + "All input tensors must have the same leading dimension"); + return rewriter.create(loc, valShape, + inputType.getElementType()); + }); + + // Initialize accumulators as empty tensors of shape [1, ...] + SmallVector acc(inputs.size()); + std::transform(inputs.begin(), inputs.end(), acc.begin(), [&](auto val) { + auto inputType = cast(val.getType()); + auto valShape = inputType.getShape(); + SmallVector accShape({1}); + accShape.insert(accShape.end(), shape.begin() + 1, shape.end()); + return rewriter.create(loc, accShape, + inputType.getElementType()); + }); + + // scf.for loop bounds and step + // NOTE: Scf only support positive step, so we need to handle reverse + Value lowerBound = rewriter.create(loc, 0); + Value upperBound = rewriter.create(loc, shape[0]); + Value step = rewriter.create(loc, 1); + + // Use tuple to pass acc and res as loop-carried variables + SmallVector initVals; + for (auto v : res) + initVals.push_back(v); + for (auto v : acc) + initVals.push_back(v); + + auto forOp = rewriter.create( + loc, lowerBound, upperBound, step, initVals, + [&](OpBuilder &b, Location loc, Value iv, ValueRange iterArgs) { + // iterArgs: [res..., acc...] + SmallVector curRes(iterArgs.begin(), + iterArgs.begin() + inputs.size()); + SmallVector currAcc(iterArgs.begin() + inputs.size(), + iterArgs.end()); + + // Build offsets, sizes, and strides for ExtractSlice + SmallVector dynOffsets; + if (reverse) { + auto actualIdx = b.create(loc, upperBound, iv); + actualIdx = b.create(loc, actualIdx, step); + dynOffsets.push_back(actualIdx); + } else { + dynOffsets.push_back(iv); + } + + for (size_t j = 1; j < shape.size(); ++j) { + dynOffsets.push_back(b.create(loc, 0)); + } + SmallVector sizeVal({1}); + sizeVal.insert(sizeVal.end(), shape.begin() + 1, shape.end()); + SmallVector strides(shape.size(), 1); + + // iv is a Value, so build dynamic offsets for ExtractSlice + SmallVector subInputs(inputs.size()); + std::transform( + inputs.begin(), inputs.end(), subInputs.begin(), [&](auto val) { + return b.create( + loc, + RankedTensorType::get( + sizeVal, + cast(val.getType()).getElementType()), + val, dynOffsets, /*sizes*/ ValueRange(), + /*strides*/ ValueRange(), + /*static_offsets*/ + SmallVector(shape.size(), ShapedType::kDynamic), + sizeVal, strides); + }); + + // Check if this is the first iteration + Value isFirstValue = b.create( + loc, arith::CmpIPredicate::eq, iv, lowerBound); + // FIXME: Bufferize will generate dynamic stride memref type, which + // cause memref::copy conversion failure + scf::IfOp ifOp = b.create( + loc, isFirstValue, + [&](OpBuilder &b, Location loc) { + b.create(loc, subInputs); + }, + [&](OpBuilder &b, Location loc) { + b.create( + loc, accumulate(subInputs, currAcc, combineOp, b)); + }); + currAcc = ifOp.getResults(); + + // Insert current accumulator into result tensor + for (size_t i = 0; i < res.size(); ++i) { + curRes[i] = b.create( + loc, resTypes[i], currAcc[i], curRes[i], dynOffsets, + /*sizes*/ ValueRange(), /*strides*/ ValueRange(), + /*static_offsets*/ + SmallVector(shape.size(), ShapedType::kDynamic), + sizeVal, strides); + } + + SmallVector yieldVals; + for (auto v : curRes) + yieldVals.push_back(v); + for (auto v : currAcc) + yieldVals.push_back(v); + + b.create(loc, yieldVals); + }); + + // Extract result tensors from forOp + SmallVector results; + for (size_t i = 0; i < inputs.size(); ++i) { + results.push_back(forOp.getResult(i)); + } + return results; + } + + uint32_t getAxis(triton::ScanOp op) const override { return op.getAxis(); } + SmallVector getInputs(triton::ScanOp op) const override { + return op->getOperands(); + } +}; + +} // namespace + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/Passes.h b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/Passes.h new file mode 100755 index 00000000..b95cbde7 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_ARITH_TO_LINALG_CONVERSION_PASSES_H +#define TRITON_ARITH_TO_LINALG_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/Passes.td b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/Passes.td new file mode 100755 index 00000000..8678cca9 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/Passes.td @@ -0,0 +1,22 @@ +#ifndef TRITON_ARITH_TO_LINALG_CONVERSION_PASSES +#define TRITON_ARITH_TO_LINALG_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonArithToLinalg : Pass<"triton-arith-to-linalg", "mlir::ModuleOp"> { + let summary = "Convert Triton arithmetic operations into linalg"; + let options = [ + Option<"pidsToFuncArgs", "pids-to-func-args", "bool", /*default*/"true", + "Convert tt.get_program_id and tt.get_num_programs to reference to function arguments">, + Option<"ttToFuncFunc", "tt-to-func-func", "bool", /*default*/"true", + "Convert tt.func to func.func">, + Option<"addptrToLinalg", "addptr-to-linalg", "bool", /*default*/"true", + "Convert tt.addptr on tensors to linalg">, + Option<"assertToCf", "assert-to-cf", "bool", /*default*/"true", + "Convert tt.assert to cf.assert">, + Option<"tensorPtrToLinalg", "tensor-ptr-to-linalg", "bool", /*default*/"false", + "Convert triton ops on tensor of pointers to linalg.generic">, + ]; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h new file mode 100755 index 00000000..b1a7c549 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h @@ -0,0 +1,31 @@ +#ifndef TRITON_CONVERSION_TRITONARITHTOLINALG_TRITONARITHTOLINALG_H +#define TRITON_CONVERSION_TRITONARITHTOLINALG_TRITONARITHTOLINALG_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +void populateTritonArithToLinalgCanonicalizationPatterns( + RewritePatternSet &patterns); + +void populateTritonArithToLinalgConversionPatterns(bool pidsToFuncArgs, + bool addptrToLinalg, + bool assertToCf, + RewritePatternSet &patterns); + +void populateTritonTensorPtrConversionPatterns(RewritePatternSet &patterns); + +std::unique_ptr> +createTritonArithToLinalgPass(bool tensorPtrToLinalg = false); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONARITHTOLINALG_TRITONARITHTOLINALG_H diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/CMakeLists.txt new file mode 100755 index 00000000..07f9ad33 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonPtrToMemref) +add_public_tablegen_target(TritonPtrToMemrefConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/Passes.h b/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/Passes.h new file mode 100755 index 00000000..e1f6f33b --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_PTR_TO_MEMREF_CONVERSION_PASSES_H +#define TRITON_PTR_TO_MEMREF_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonPtrToMemref/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/Passes.td b/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/Passes.td new file mode 100755 index 00000000..c027b098 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/Passes.td @@ -0,0 +1,11 @@ +#ifndef TRITON_PTR_TO_MEMREF_CONVERSION_PASSES +#define TRITON_PTR_TO_MEMREF_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonPtrToMemref : Pass<"triton-ptr-to-memref", "mlir::ModuleOp"> { + let summary = "Convert triton pointer to unranked memref"; + let constructor = "triton::createTritonPtrToMemrefPass()"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h b/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h new file mode 100755 index 00000000..4476f7d6 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h @@ -0,0 +1,17 @@ +#ifndef TRITON_CONVERSION_TRITON_PTR_TO_MEMREF_TRITON_PTR_TO_MEMREF_H +#define TRITON_CONVERSION_TRITON_PTR_TO_MEMREF_TRITON_PTR_TO_MEMREF_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonPtrToMemrefPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITON_PTR_TO_MEMREF_TRITON_PTR_TO_MEMREF_H diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/CMakeLists.txt new file mode 100755 index 00000000..3cc51fcb --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/CMakeLists.txt @@ -0,0 +1,9 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Triton Project Contributors. +# +#===------------------------------------------------------------------------===# + +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToCoreDialects) +add_public_tablegen_target(TritonToCoreDialectsConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/Passes.h b/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/Passes.h new file mode 100755 index 00000000..32fc0104 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/Passes.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_TO_CORE_DIALECTS_CONVERSION_PASSES_H +#define TRITON_TO_CORE_DIALECTS_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonToCoreDialects/TritonToCoreDialects.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonToCoreDialects/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/Passes.td b/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/Passes.td new file mode 100755 index 00000000..6d10cfb6 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/Passes.td @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_TO_CORE_DIALECTS_CONVERSION_PASSES +#define TRITON_TO_CORE_DIALECTS_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToCoreDialects : Pass<"triton-to-core-dialects", "mlir::ModuleOp"> { + let summary = "Convert Triton to core dialects including Linalg, Memref etc"; + let constructor = "triton::createTritonToCoreDialectsPass()"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/TritonToCoreDialects.h b/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/TritonToCoreDialects.h new file mode 100755 index 00000000..d968cc05 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/TritonToCoreDialects.h @@ -0,0 +1,27 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// +// +// This pass is the wrapall pass that populates all the conversion patterns from +// triton to core dialects such as linalg, memref, buf etc. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_TRITON_TO_CORE_DIALECTS_H +#define TRITON_CONVERSION_TRITON_TO_CORE_DIALECTS_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonToCoreDialectsPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITON_TO_CORE_DIALECTS_H diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToMK/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/TritonToMK/CMakeLists.txt new file mode 100755 index 00000000..f90ffc19 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToMK/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToMK) +add_public_tablegen_target(TritonToMKConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToMK/Passes.h b/third_party/wafer/include/triton-shared/Conversion/TritonToMK/Passes.h new file mode 100755 index 00000000..8d96e357 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToMK/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_STRUCTURED_TO_MEMREF_CONVERSION_PASSES_H +#define TRITON_STRUCTURED_TO_MEMREF_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonToMK/TritonToMKPatterns.hpp" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonToMK/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToMK/Passes.td b/third_party/wafer/include/triton-shared/Conversion/TritonToMK/Passes.td new file mode 100755 index 00000000..415448c1 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToMK/Passes.td @@ -0,0 +1,10 @@ +#ifndef TRITON_TO_MK_CONVERSION_PASSES +#define TRITON_TO_MK_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToMK : Pass<"triton-to-mk", "mlir::ModuleOp"> { + let summary = "Convert triton triton pointer ops to mk"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToMK/TritonToMKPatterns.hpp b/third_party/wafer/include/triton-shared/Conversion/TritonToMK/TritonToMKPatterns.hpp new file mode 100755 index 00000000..6e86fb5c --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToMK/TritonToMKPatterns.hpp @@ -0,0 +1,194 @@ +#ifndef TRITON_CONVERSION_PATTERNS +#define TRITON_CONVERSION_PATTERNS + +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "flagtree/Common/UnifiedHardware.h" + +#include "mlir-ext/Dialect/MathExt/IR/MathExt.h" +#include "triton-shared/Analysis/MaskAnalysis.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/Analysis/PtrAnalysis.h" +#include "triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns_FlagTree.hpp" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" +#include "triton-shared/Utils/Utils.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include +#include +#include + +using namespace mlir; +using namespace triton; + +namespace { + +// FIXME: There is no triton::BarrierOp currently. +struct BarrierConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mlir::gpu::BarrierOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + + rewriter.create(loc); + rewriter.eraseOp(op); + return success(); + } +}; + +// Similar with triton-cpu. +struct PrintOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::PrintOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op->getLoc(); + // If the op has no operands, we can just print the prefix. + if (op.getNumOperands() == 0) { + rewriter.create(loc, TypeRange{}, op.getPrefix(), + op.getHex(), ValueRange{}, + llvm::SmallVector{}); + rewriter.eraseOp(op); + return success(); + } + + for (size_t i = 0; i < op.getNumOperands(); i++) { + Value operand = op.getOperands()[i]; + auto isSigned = {op.getIsSigned()[i]}; + // If the operand is not a ranked tensor, we should create a new tensor. + // See mlir/lib/Interfaces/DestinationStyleOpInterface.cpp#L39 + if (!isa(operand.getType())) { + // NOTE: Use tensor.from_elements, the arith.constant will not translate + // to linalg.fill + auto emptyTensor = rewriter.create( + loc, SmallVector{}, operand.getType()); + + auto operandTensor = rewriter.create( + loc, operand, emptyTensor, ValueRange{}); + rewriter.create(loc, operandTensor.getType(), + op.getPrefix(), op.getHex(), + operandTensor.getResult(), isSigned); + continue; + } + + auto operandType = cast(operand.getType()); + auto flattenTensor = operand; + if (operandType.getRank() != 1) { + SmallVector flatten_shape = {operandType.getNumElements()}; + auto targetType = + RankedTensorType::get(flatten_shape, operandType.getElementType()); + // NOTE: Avoid to create global constant tensors + SmallVector reassociation(1); + for (unsigned i = 0; i < operandType.getRank(); ++i) { + reassociation.front().push_back(i); + } + flattenTensor = rewriter.create( + loc, targetType, operand, reassociation); + } + + rewriter.create(loc, flattenTensor.getType(), op.getPrefix(), + op.getHex(), flattenTensor, isSigned); + } + + rewriter.eraseOp(op); + return success(); + } +}; + +struct DotScaledConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::DotScaledOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + + // Get operands + auto a = op.getA(); + auto b = op.getB(); + auto c = op.getC(); + Value aScale = op.getAScale(); + Value bScale = op.getBScale(); + auto aElemType = op.getAElemTypeAttr(); + auto bElemType = op.getBElemTypeAttr(); + auto fastMath = op.getFastMathAttr(); + + // Get type information + auto aType = a.getType(); + auto bType = b.getType(); + auto dstType = cast(op.getType()); + auto elementType = dstType.getElementType(); + + // Create initial zero tensor + auto init = + rewriter.create(loc, dstType.getShape(), elementType); + TypedAttr constantAttr = + static_cast(rewriter.getFloatAttr(elementType, 0)); + auto zero = rewriter.create( + op.getLoc(), elementType, constantAttr); + auto zeroes = + rewriter.create(loc, ValueRange{zero}, ValueRange{init}) + .result(); + + // Perform scaled dot product operation + Value res = rewriter + .create(loc, TypeRange{op.getType()}, a, + aScale, b, bScale, zeroes, + aElemType, bElemType, fastMath) + .getResult(0); + + // Check if C needs to be added + bool skipC = isZeroTensor(c, false); + if (!skipC) { + res = rewriter.create(loc, c, res); + } + + rewriter.replaceOp(op, res); + return success(); + } +}; + +struct GatherConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::GatherOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + auto resultType = cast(op.getResult().getType()); + Value dstInit = rewriter.create( + op.getLoc(), resultType.getShape(), resultType.getElementType()); + + auto gatherOp = + rewriter.create(op.getLoc(), op.getType(), op.getSrc(), + op.getIndices(), dstInit, op.getAxis()); + + rewriter.replaceOp(op, gatherOp.getResult()); + return success(); + } +}; + +} // namespace + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/CMakeLists.txt new file mode 100755 index 00000000..116a3e3f --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToUnstructured) +add_public_tablegen_target(TritonToUnstructuredConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/Passes.h b/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/Passes.h new file mode 100755 index 00000000..a2016c7a --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_TO_UNSTRUCTURED_CONVERSION_PASSES_H +#define TRITON_TO_UNSTRUCTURED_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonToUnstructured/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/Passes.td b/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/Passes.td new file mode 100755 index 00000000..542d087c --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/Passes.td @@ -0,0 +1,15 @@ +#ifndef TRITON_TO_UNSTRUCTURED_CONVERSION_PASSES +#define TRITON_TO_UNSTRUCTURED_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToUnstructured : Pass<"triton-to-unstructured", "mlir::ModuleOp"> { + let summary = "Transforms tt.addptr ops into offset accumulation ops"; + let constructor = "triton::createTritonToUnstructuredPass()"; + let options = [ + Option<"offsetBitWidth", "offset-bit-width", "size_t", /*default*/"32", + "Bitwidth used for the starting offset of each pointer"> + ]; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h b/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h new file mode 100755 index 00000000..03ccdcd1 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h @@ -0,0 +1,17 @@ +#ifndef TRITON_CONVERSION_TRITON_TO_UNSTRUCTURED_TRITON_TO_UNSTRUCTURED_H +#define TRITON_CONVERSION_TRITON_TO_UNSTRUCTURED_TRITON_TO_UNSTRUCTURED_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonToUnstructuredPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITON_TO_UNSTRUCTURED_TRITON_TO_UNSTRUCTURED_H diff --git a/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/CMakeLists.txt new file mode 100755 index 00000000..d31c6aee --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/CMakeLists.txt @@ -0,0 +1,10 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. +# +#===------------------------------------------------------------------------===# + +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name UnstructuredToMK) +add_public_tablegen_target(UnstructuredToMKConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/Passes.h b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/Passes.h new file mode 100755 index 00000000..09649d3f --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/Passes.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES_H +#define UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/UnstructuredToMK/UnstructuredToMK.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/UnstructuredToMK/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/Passes.td b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/Passes.td new file mode 100755 index 00000000..679d203a --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/Passes.td @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef UNSTRUCTURED_TO_MK_CONVERSION_PASSES +#define UNSTRUCTURED_TO_MK_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def UnstructuredToMK : Pass<"unstructured-to-mk", "mlir::ModuleOp"> { + let summary = "Convert unstructured triton ptr (gather / scatter) to mk"; + let constructor = "triton::createUnstructuredToMKPass()"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/UnstructuredToMK.h b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/UnstructuredToMK.h new file mode 100755 index 00000000..ff457c3a --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/UnstructuredToMK.h @@ -0,0 +1,21 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_UNSTRUCTUREDTOMK_UNSTRUCTUREDTOMK_H +#define TRITON_CONVERSION_UNSTRUCTUREDTOMK_UNSTRUCTUREDTOMK_H + +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createUnstructuredToMKPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_UNSTRUCTUREDTOMK_UNSTRUCTUREDTOMK_H diff --git a/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/CMakeLists.txt b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/CMakeLists.txt new file mode 100755 index 00000000..f988ac9f --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/CMakeLists.txt @@ -0,0 +1,10 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. +# +#===------------------------------------------------------------------------===# + +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name UnstructuredToMemref) +add_public_tablegen_target(UnstructuredToMemrefConversionPassIncGen) diff --git a/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/Passes.h b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/Passes.h new file mode 100755 index 00000000..f2d71174 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/Passes.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES_H +#define UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/UnstructuredToMemref/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/Passes.td b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/Passes.td new file mode 100755 index 00000000..a0bf316d --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/Passes.td @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES +#define UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def UnstructuredToMemref : Pass<"unstructured-to-memref", "mlir::ModuleOp"> { + let summary = "Convert unstructured triton ptr (gather / scatter) to memref"; + let constructor = "triton::createUnstructuredToMemrefPass()"; +} + +#endif diff --git a/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h new file mode 100755 index 00000000..ad0f5c46 --- /dev/null +++ b/third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h @@ -0,0 +1,21 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_UNSTRUCTUREDTOMEMREF_UNSTRUCTUREDTOMEMREF_H +#define TRITON_CONVERSION_UNSTRUCTUREDTOMEMREF_UNSTRUCTUREDTOMEMREF_H + +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createUnstructuredToMemrefPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_UNSTRUCTUREDTOMEMREF_UNSTRUCTUREDTOMEMREF_H diff --git a/third_party/wafer/include/utils/LinalgOpBuilderHelper.h b/third_party/wafer/include/utils/LinalgOpBuilderHelper.h new file mode 100755 index 00000000..795470cb --- /dev/null +++ b/third_party/wafer/include/utils/LinalgOpBuilderHelper.h @@ -0,0 +1,50 @@ +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" + +using namespace mlir; + +// buildLinalgElementwise with OneResult and custom result type +template +static Value buildLinalgElementwise(OpBuilder &rewriter, Location loc, + RankedTensorType resultType, + ValueRange inputs) { + static_assert(OpT::template hasTrait() && + OpT::template hasTrait(), + "OpT must have OneResult and Elementwise traits"); + auto inputType = cast(inputs[0].getType()); + auto rank = inputType.getRank(); + auto identityMap = + AffineMap::getMultiDimIdentityMap(rank, rewriter.getContext()); + SmallVector indexingMaps(inputs.size() + 1, identityMap); + SmallVector iteratorTypes(rank, + utils::IteratorType::parallel); + + auto output = rewriter.create(loc, resultType.getShape(), + resultType.getElementType()); + + auto linalgBuilder = [&](OpBuilder &nestedBuilder, Location nestedloc, + ValueRange iterArgs) { + Value opRes = nestedBuilder.create( + nestedloc, iterArgs.back().getType(), iterArgs.drop_back()); + nestedBuilder.create(nestedloc, opRes); + }; + + return rewriter + .create(loc, resultType, inputs, ValueRange{output}, + indexingMaps, iteratorTypes, linalgBuilder) + .getResult(0); +} + +// buildLinalgElementwise with OneResult and SameOperandsAndResultType +template +static Value buildLinalgElementwise(OpBuilder &rewriter, Location loc, + ValueRange inputs) { + static_assert( + OpT::template hasTrait() && + OpT::template hasTrait() && + OpT::template hasTrait(), + "OpT must have OneResult, Elementwise and SameOperandsAndResultType " + "traits"); + auto inputType = cast(inputs[0].getType()); + return buildLinalgElementwise(rewriter, loc, inputType, inputs); +} diff --git a/third_party/wafer/include/utils/TypeConvertor.h b/third_party/wafer/include/utils/TypeConvertor.h new file mode 100755 index 00000000..38af8856 --- /dev/null +++ b/third_party/wafer/include/utils/TypeConvertor.h @@ -0,0 +1,30 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "mlir/IR/TypeUtilities.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir::triton { + +class PtrToUnrankedMemrefConverter : public TypeConverter { +public: + PtrToUnrankedMemrefConverter() { + addConversion([](Type type) { return type; }); + addConversion([](triton::PointerType ptrType) { + return UnrankedMemRefType::get(ptrType.getPointeeType(), 0); + }); + addTargetMaterialization([&](OpBuilder &builder, + UnrankedMemRefType resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + } +}; + +} // namespace mlir::triton diff --git a/third_party/wafer/include/wafer/CMakeLists.txt b/third_party/wafer/include/wafer/CMakeLists.txt new file mode 100755 index 00000000..495310c6 --- /dev/null +++ b/third_party/wafer/include/wafer/CMakeLists.txt @@ -0,0 +1,3 @@ +add_subdirectory(Conversion) +add_subdirectory(Dialect) +add_subdirectory(Transforms) diff --git a/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/CMakeLists.txt b/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/CMakeLists.txt new file mode 100755 index 00000000..27effe00 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name AllocateSharedMemory) +add_public_tablegen_target(AllocateSharedMemoryPassIncGen) diff --git a/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/Passes.h b/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/Passes.h new file mode 100755 index 00000000..0c8a46f1 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/Passes.h @@ -0,0 +1,25 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef ALLOCATE_SHARED_MEMORY_PASSES_H +#define ALLOCATE_SHARED_MEMORY_PASSES_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir::triton::alloc { + +#define GEN_PASS_DECL +#include "wafer/Conversion/AllocateSharedMemory/Passes.h.inc" + +#define GEN_PASS_REGISTRATION +#include "wafer/Conversion/AllocateSharedMemory/Passes.h.inc" + +} // namespace mlir::triton::alloc + +#endif // ALLOCATE_SHARED_MEMORY_PASSES_H diff --git a/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/Passes.td b/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/Passes.td new file mode 100755 index 00000000..61b2a21d --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/Passes.td @@ -0,0 +1,24 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef ALLOCATE_SHARED_MEMORY_PASSES +#define ALLOCATE_SHARED_MEMORY_PASSES + +include "mlir/Pass/PassBase.td" + +def AllocateSharedMemory : Pass<"spmd-allocate-shared-memory", "mlir::ModuleOp"> { + let summary = "SPMD(single program multi-data) mode: add metadata for shared memory allocation"; + + let description = [{ + This pass uses the `ModuleAllocation` analysis to: + - Annotate modules with an attribute with the amount of shared/local + memory used. + - Annotate operations with an offset into the total shared/local memory. + }]; +} + +#endif diff --git a/third_party/wafer/include/wafer/Conversion/CMakeLists.txt b/third_party/wafer/include/wafer/Conversion/CMakeLists.txt new file mode 100755 index 00000000..98ea8006 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/CMakeLists.txt @@ -0,0 +1,7 @@ +add_subdirectory(LinalgTiling) +add_subdirectory(LinalgFusion) +add_subdirectory(MKToWafer) +add_subdirectory(WaferToLLVM) +add_subdirectory(WaferMemrefToLLVM) +add_subdirectory(AllocateSharedMemory) +add_subdirectory(ExportKernelSymbols) diff --git a/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/CMakeLists.txt b/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/CMakeLists.txt new file mode 100755 index 00000000..6f6c9b63 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name ExportKernelSymbols) +add_public_tablegen_target(ExportKernelSymbolsConversionPassIncGen) diff --git a/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/ExportKernelSymbols.h b/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/ExportKernelSymbols.h new file mode 100755 index 00000000..919a14dd --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/ExportKernelSymbols.h @@ -0,0 +1,26 @@ +//===------------------- ExportKernelSymbols.h -------------------------*- C++ +//-*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Ludt) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef EXPORT_KERNEL_SYMBOLS_CONVERSION_H +#define EXPORT_KERNEL_SYMBOLS_CONVERSION_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "wafer/Conversion/ExportKernelSymbols/Passes.h.inc" + +std::unique_ptr> createExportKernelSymbolsPass(); + +} // namespace triton +} // namespace mlir + +#endif // EXPORT_KERNEL_SYMBOLS_CONVERSION_H diff --git a/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/Passes.h b/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/Passes.h new file mode 100755 index 00000000..f97dbed4 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/Passes.h @@ -0,0 +1,22 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Ludt) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef EXPORT_KERNEL_SYMBOLS_CONVERSION_PASSES_H +#define EXPORT_KERNEL_SYMBOLS_CONVERSION_PASSES_H + +#include "wafer/Conversion/ExportKernelSymbols/ExportKernelSymbols.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "wafer/Conversion/ExportKernelSymbols/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // EXPORT_KERNEL_SYMBOLS_CONVERSION_PASSES_H diff --git a/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/Passes.td b/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/Passes.td new file mode 100755 index 00000000..ed062461 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/Passes.td @@ -0,0 +1,20 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Ludt) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef EXPORT_KERNEL_SYMBOLS_CONVERSION_PASSES +#define EXPORT_KERNEL_SYMBOLS_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def ExportKernelSymbols : Pass<"export-kernel-symbols", "mlir::ModuleOp"> { + let summary = "Export kernel function to a dynamic symbol table that can be accessed at runtime."; + let constructor = "triton::createExportKernelSymbolsPass()"; + let options = []; + let dependentDialects = ["LLVM::LLVMDialect"]; +} + +#endif diff --git a/third_party/wafer/include/wafer/Conversion/LinalgFusion/CMakeLists.txt b/third_party/wafer/include/wafer/Conversion/LinalgFusion/CMakeLists.txt new file mode 100755 index 00000000..907d8137 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/LinalgFusion/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name LinalgFusion) +add_public_tablegen_target(LinalgFusionConversionPassIncGen) diff --git a/third_party/wafer/include/wafer/Conversion/LinalgFusion/LinalgFusion.h b/third_party/wafer/include/wafer/Conversion/LinalgFusion/LinalgFusion.h new file mode 100755 index 00000000..478320db --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/LinalgFusion/LinalgFusion.h @@ -0,0 +1,27 @@ +#ifndef TRITON_CONVERSION_LINALG_FUSION_H +#define TRITON_CONVERSION_LINALG_FUSION_H + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "wafer/Conversion/LinalgFusion/Passes.h.inc" + +void populateLinalgBinaryOpFusionPatterns(RewritePatternSet &patterns); + +// TODO: Support linalg elementwise op fusion. +#if 0 +void populateLinalgFusionPatterns(RewritePatternSet &patterns); +#endif + +std::unique_ptr> createLinalgFusionPass(); + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/include/wafer/Conversion/LinalgFusion/Passes.h b/third_party/wafer/include/wafer/Conversion/LinalgFusion/Passes.h new file mode 100755 index 00000000..96b6ac07 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/LinalgFusion/Passes.h @@ -0,0 +1,22 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_LINALG_FUSION_PASSES_H +#define TRITON_CONVERSION_LINALG_FUSION_PASSES_H + +#include "wafer/Conversion/LinalgFusion/LinalgFusion.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "wafer/Conversion/LinalgFusion/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_LINALG_FUSION_PASSES_H diff --git a/third_party/wafer/include/wafer/Conversion/LinalgFusion/Passes.td b/third_party/wafer/include/wafer/Conversion/LinalgFusion/Passes.td new file mode 100755 index 00000000..97562200 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/LinalgFusion/Passes.td @@ -0,0 +1,23 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_LINALG_FUSION_PASSES +#define TRITON_CONVERSION_LINALG_FUSION_PASSES + +include "mlir/Pass/PassBase.td" + +def LinalgFusion : Pass<"linalg-fusion", "mlir::ModuleOp"> { + let summary = "Apply fusion transformation to linalg operations"; + let description = [{ + This pass applies scalar fusion transformation to linalg operations. + It fuses scalar input linalg ops to reduce redundant memory read and write operations. + Elementwise fusion needs to be implemented later. + }]; + let constructor = "triton::createLinalgFusionPass()"; +} + +#endif // TRITON_CONVERSION_LINALG_FUSION_PASSES diff --git a/third_party/wafer/include/wafer/Conversion/LinalgTiling/CMakeLists.txt b/third_party/wafer/include/wafer/Conversion/LinalgTiling/CMakeLists.txt new file mode 100755 index 00000000..9b8130ce --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/LinalgTiling/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name LinalgTiling) +add_public_tablegen_target(LinalgTilingConversionPassIncGen) diff --git a/third_party/wafer/include/wafer/Conversion/LinalgTiling/LinalgTiling.h b/third_party/wafer/include/wafer/Conversion/LinalgTiling/LinalgTiling.h new file mode 100755 index 00000000..8bbe4530 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/LinalgTiling/LinalgTiling.h @@ -0,0 +1,34 @@ +//===------------------- LinalgTiling.h ------------------------*- C++-*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file implements the patterns to tile linalg operations for better +// performance. It applies tiling transformations to improve data locality +// and parallelism. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_LINALG_TILING_H +#define TRITON_CONVERSION_LINALG_TILING_H + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "wafer/Conversion/LinalgTiling/Passes.h.inc" + +void populateLinalgTilingPatterns(RewritePatternSet &patterns); + +std::unique_ptr> createLinalgTilingPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_LINALG_TILING_H diff --git a/third_party/wafer/include/wafer/Conversion/LinalgTiling/Passes.h b/third_party/wafer/include/wafer/Conversion/LinalgTiling/Passes.h new file mode 100755 index 00000000..dbea5178 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/LinalgTiling/Passes.h @@ -0,0 +1,22 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_LINALG_TILING_PASSES_H +#define TRITON_CONVERSION_LINALG_TILING_PASSES_H + +#include "wafer/Conversion/LinalgTiling/LinalgTiling.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "wafer/Conversion/LinalgTiling/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_LINALG_TILING_PASSES_H diff --git a/third_party/wafer/include/wafer/Conversion/LinalgTiling/Passes.td b/third_party/wafer/include/wafer/Conversion/LinalgTiling/Passes.td new file mode 100755 index 00000000..67496b1c --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/LinalgTiling/Passes.td @@ -0,0 +1,22 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_LINALG_TILING_PASSES +#define TRITON_CONVERSION_LINALG_TILING_PASSES + +include "mlir/Pass/PassBase.td" + +def LinalgTiling : Pass<"linalg-tiling", "mlir::ModuleOp"> { + let summary = "Apply tiling transformation to linalg operations"; + let description = [{ + This pass applies tiling transformation to linalg operations. + It tiles and fuses linalg operations to improve data locality and parallelism. + }]; + let constructor = "triton::createLinalgTilingPass()"; +} + +#endif // TRITON_CONVERSION_LINALG_TILING_PASSES diff --git a/third_party/wafer/include/wafer/Conversion/MKToWafer/CMakeLists.txt b/third_party/wafer/include/wafer/Conversion/MKToWafer/CMakeLists.txt new file mode 100755 index 00000000..4c42f960 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/MKToWafer/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name MKToWafer) +add_public_tablegen_target(MKToWaferConversionPassIncGen) diff --git a/third_party/wafer/include/wafer/Conversion/MKToWafer/MKToWafer.h b/third_party/wafer/include/wafer/Conversion/MKToWafer/MKToWafer.h new file mode 100755 index 00000000..fd0e00d2 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/MKToWafer/MKToWafer.h @@ -0,0 +1,36 @@ +//===------------------- MKToWafer.h ---------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Lowering magic kernel ops to Wafer Wafer target. +// +//===----------------------------------------------------------------------===// + +#ifndef ZTC_CONVERSION_MK_TO_WAFER_H +#define ZTC_CONVERSION_MK_TO_WAFER_H + +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "wafer/Conversion/MKToWafer/Passes.h.inc" + +void populateMKToWaferCanonicalizationPatterns(RewritePatternSet &patterns); + +void populateMKToWaferConversionPatterns(RewritePatternSet &patterns); + +std::unique_ptr> createMKToWaferPass(); + +} // namespace triton +} // namespace mlir + +#endif // ZTC_CONVERSION_MK_TO_WAFER_H diff --git a/third_party/wafer/include/wafer/Conversion/MKToWafer/Passes.h b/third_party/wafer/include/wafer/Conversion/MKToWafer/Passes.h new file mode 100755 index 00000000..04ed9ed9 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/MKToWafer/Passes.h @@ -0,0 +1,22 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MK_TO_WAFER_CONVERSION_PASSES_H +#define MK_TO_WAFER_CONVERSION_PASSES_H + +#include "wafer/Conversion/MKToWafer/MKToWafer.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "wafer/Conversion/MKToWafer/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // MK_TO_WAFER_CONVERSION_PASSES_H diff --git a/third_party/wafer/include/wafer/Conversion/MKToWafer/Passes.td b/third_party/wafer/include/wafer/Conversion/MKToWafer/Passes.td new file mode 100755 index 00000000..d9539f23 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/MKToWafer/Passes.td @@ -0,0 +1,18 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MK_TO_WAFER_CONVERSION_PASSES +#define MK_TO_WAFER_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def MKToWafer : Pass<"mk-to-wafer", "mlir::ModuleOp"> { + let summary = "Convert magic kernel operations into Wafer Wafer operations"; + let constructor = "triton::createMKToWaferPass()"; +} + +#endif diff --git a/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/CMakeLists.txt b/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/CMakeLists.txt new file mode 100755 index 00000000..e17b4f8b --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name WaferMemrefToLLVM) +add_public_tablegen_target(WaferMemrefToLLVMConversionPassIncGen) diff --git a/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/Passes.h b/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/Passes.h new file mode 100755 index 00000000..8d333084 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/Passes.h @@ -0,0 +1,22 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MEMREF_TO_MK_CONVERSION_PASSES_H +#define MEMREF_TO_MK_CONVERSION_PASSES_H + +#include "wafer/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "wafer/Conversion/WaferMemrefToLLVM/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // MEMREF_TO_MK_CONVERSION_PASSES_H diff --git a/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/Passes.td b/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/Passes.td new file mode 100755 index 00000000..a2c99d95 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/Passes.td @@ -0,0 +1,19 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MEMREF_TO_MK_CONVERSION_PASSES +#define MEMREF_TO_MK_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def WaferMemrefToLLVM : Pass<"wafer-memref-to-llvm", "mlir::ModuleOp"> { + let summary = "Convert memref and bufferization operations into custom llvm function call."; + let constructor = "triton::createWaferMemrefToLLVMPass()"; + let options = []; +} + +#endif diff --git a/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.h b/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.h new file mode 100755 index 00000000..9d6e7a0d --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.h @@ -0,0 +1,40 @@ +//===------------------- WaferMemrefToLLVM.h -------------------------*- C++ +//-*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Lowering memref.copy, memref.alloc to mk.load, mk.alloc etc. +// +//===----------------------------------------------------------------------===// + +#ifndef ZTC_CONVERSION_MEMREF_TO_MK_H +#define ZTC_CONVERSION_MEMREF_TO_MK_H + +#include "mlir/Conversion/LLVMCommon/TypeConverter.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "wafer/Conversion/WaferMemrefToLLVM/Passes.h.inc" + +void populateWaferMemrefToLLVMCanonicalizationPatterns( + RewritePatternSet &patterns); + +void populateWaferMemrefToLLVMConversionPatterns(RewritePatternSet &patterns, + LLVMTypeConverter &converter); + +std::unique_ptr> createWaferMemrefToLLVMPass(); + +} // namespace triton +} // namespace mlir + +#endif // ZTC_CONVERSION_MEMREF_TO_MAGICKERNEL_H diff --git a/third_party/wafer/include/wafer/Conversion/WaferToLLVM/CMakeLists.txt b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/CMakeLists.txt new file mode 100755 index 00000000..b0f5849b --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/CMakeLists.txt @@ -0,0 +1,7 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name WaferToLLVM) +add_public_tablegen_target(WaferToLLVMConversionPassIncGen) + +set(LLVM_TARGET_DEFINITIONS KernelArgBufferPass.td) +mlir_tablegen(KernelArgBufferPass.h.inc -gen-pass-decls --name KernelArgBufferPass) +add_public_tablegen_target(KernelArgBufferPassIncGen) diff --git a/third_party/wafer/include/wafer/Conversion/WaferToLLVM/KernelArgBufferPass.h b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/KernelArgBufferPass.h new file mode 100755 index 00000000..352b82c1 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/KernelArgBufferPass.h @@ -0,0 +1,35 @@ +//===- KernelArgBufferPass.h ----------------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This pass transforms kernel function signatures by converting multiple +// arguments into a single void* buffer containing all the arguments. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_KERNEL_ARG_BUFFER_PASS_H +#define MLIR_KERNEL_ARG_BUFFER_PASS_H + +#include "mlir/Pass/Pass.h" +#include + +namespace mlir { +class ModuleOp; +class Pass; + +namespace triton { +/// Creates a pass that transforms kernel functions by replacing multiple +/// arguments with a single void* buffer argument. +std::unique_ptr createKernelArgBufferPass(); + +#define GEN_PASS_REGISTRATION +#define GEN_PASS_DECL +#include "wafer/Conversion/WaferToLLVM/KernelArgBufferPass.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // MLIR_KERNEL_ARG_BUFFER_PASS_H diff --git a/third_party/wafer/include/wafer/Conversion/WaferToLLVM/KernelArgBufferPass.td b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/KernelArgBufferPass.td new file mode 100755 index 00000000..a47c45d0 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/KernelArgBufferPass.td @@ -0,0 +1,32 @@ +//===- KernelArgBufferPass.td ---------------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef KERNEL_ARG_BUFFER_PASS +#define KERNEL_ARG_BUFFER_PASS + +include "mlir/Pass/PassBase.td" + +def KernelArgBufferPass : Pass<"kernel-arg-buffer", "ModuleOp"> { + let summary = "Convert kernel arguments to a single buffer argument"; + let description = [{ + This pass transforms kernel function signatures by converting multiple + arguments into a single void* buffer containing all the arguments. + + For example, a function like: + add_kernel(uint64_t* arg1, uint64_t* arg2, int64_t size, int gridX, int x) + + Will be converted to: + add_kernel(void* args) + + Where the args buffer contains pointers to arg1 and arg2, followed by the scalar + values size, gridX, and x. Each scalar value occupies 8 bytes in the buffer. + }]; + let constructor = "mlir::triton::createKernelArgBufferPass()"; + let dependentDialects = ["mlir::LLVM::LLVMDialect", "mlir::func::FuncDialect"]; +} + +#endif // KERNEL_ARG_BUFFER_PASS diff --git a/third_party/wafer/include/wafer/Conversion/WaferToLLVM/Passes.h b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/Passes.h new file mode 100755 index 00000000..99b0e3f8 --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/Passes.h @@ -0,0 +1,22 @@ +//===------------------- Passes.h -----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef WAFER_TO_LLVM_CONVERSION_PASSES_H +#define WAFER_TO_LLVM_CONVERSION_PASSES_H + +#include "wafer/Conversion/WaferToLLVM/WaferToLLVM.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "wafer/Conversion/WaferToLLVM/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // WAFER_TO_LLVM_CONVERSION_PASSES_H diff --git a/third_party/wafer/include/wafer/Conversion/WaferToLLVM/Passes.td b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/Passes.td new file mode 100755 index 00000000..22dd645d --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/Passes.td @@ -0,0 +1,38 @@ +//===------------------- Passes.td ----------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef WAFER_TO_LLVM_CONVERSION_PASSES +#define WAFER_TO_LLVM_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + + +def WaferToLLVM : Pass<"wafer-to-llvm", "ModuleOp"> { + let summary = "Convert Wafer dialect to LLVM dialect"; + let description = [{ + This pass converts operations in the Wafer dialect to the LLVM IR dialect. + + It handles the conversion of Wafer-specific operations like wafer.rdma, wafer.wdma, + wafer.gemm etc to appropriate LLVM calls to the Wafer runtime library. + + The pass also relies on existing conversion patterns for standard dialects + like arith, func, memref, etc. + }]; + + let constructor = "triton::createWaferToLLVMPass()"; + + let dependentDialects = [ + "mlir::LLVM::LLVMDialect", + "mlir::arith::ArithDialect", + "mlir::func::FuncDialect", + "mlir::memref::MemRefDialect", + "mlir::scf::SCFDialect", + "wafer::WaferDialect" + ]; +} + +#endif diff --git a/third_party/wafer/include/wafer/Conversion/WaferToLLVM/WaferToLLVM.h b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/WaferToLLVM.h new file mode 100755 index 00000000..c4af5e4d --- /dev/null +++ b/third_party/wafer/include/wafer/Conversion/WaferToLLVM/WaferToLLVM.h @@ -0,0 +1,33 @@ +//===------------------- WaferToLLVM.h -------------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_WAFER_TO_LLVM_H +#define TRITON_CONVERSION_WAFER_TO_LLVM_H + +#include "mlir/Conversion/LLVMCommon/Pattern.h" +#include "mlir/Conversion/LLVMCommon/TypeConverter.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "wafer/Conversion/WaferToLLVM/Passes.h.inc" + +void populateWaferToLLVMConversionPatterns(RewritePatternSet &patterns, + ConversionTarget &target, + LLVMTypeConverter &converter); + +std::unique_ptr> createWaferToLLVMPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_WAFER_TO_LLVM_H diff --git a/third_party/wafer/include/wafer/Dialect/CMakeLists.txt b/third_party/wafer/include/wafer/Dialect/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/include/wafer/Dialect/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/include/wafer/Dialect/IR/CMakeLists.txt b/third_party/wafer/include/wafer/Dialect/IR/CMakeLists.txt new file mode 100755 index 00000000..fb193137 --- /dev/null +++ b/third_party/wafer/include/wafer/Dialect/IR/CMakeLists.txt @@ -0,0 +1,14 @@ +set(LLVM_TARGET_DEFINITIONS WaferOps.td) +mlir_tablegen(WaferDialect.h.inc -gen-dialect-decls -dialect=wafer) +mlir_tablegen(WaferDialect.cpp.inc -gen-dialect-defs -dialect=wafer) +mlir_tablegen(WaferOps.h.inc -gen-op-decls) +mlir_tablegen(WaferOps.cpp.inc -gen-op-defs) + +mlir_tablegen(WaferEnums.h.inc -gen-enum-decls) +mlir_tablegen(WaferEnums.cpp.inc -gen-enum-defs) + +set(LLVM_TARGET_DEFINITIONS WaferTypes.td) +mlir_tablegen(WaferTypes.h.inc -gen-typedef-decls) +mlir_tablegen(WaferTypes.cpp.inc -gen-typedef-defs) + +add_public_tablegen_target(WaferTableGen) diff --git a/third_party/wafer/include/wafer/Dialect/IR/WaferAttrDefs.td b/third_party/wafer/include/wafer/Dialect/IR/WaferAttrDefs.td new file mode 100755 index 00000000..f4e6923e --- /dev/null +++ b/third_party/wafer/include/wafer/Dialect/IR/WaferAttrDefs.td @@ -0,0 +1,24 @@ +//===---------------------- WaferAttrDefs.td -------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef WAFER_ATTR_DEFS +#define WAFER_ATTR_DEFS + +include "mlir/IR/EnumAttr.td" + +// Round mode, aligned with RND_MODE in instr_def.h +def RoundModeAttr : I32EnumAttr<"RoundMode", "Round mode", [ + I32EnumAttrCase<"RND_NEAREST_EVEN", 0, "nearest">, + I32EnumAttrCase<"RND_ZERO", 1, "zero">, + I32EnumAttrCase<"RND_POS_INF", 2, "pos">, + I32EnumAttrCase<"RND_NEG_INF", 3, "neg">, + I32EnumAttrCase<"RND_STOCHASTIC", 4, "stochastic"> +]> { + let cppNamespace = "::mlir::wafer"; +} + +#endif // WAFER_ATTR_DEFS diff --git a/third_party/wafer/include/wafer/Dialect/IR/WaferDialect.h b/third_party/wafer/include/wafer/Dialect/IR/WaferDialect.h new file mode 100755 index 00000000..5fe82fc2 --- /dev/null +++ b/third_party/wafer/include/wafer/Dialect/IR/WaferDialect.h @@ -0,0 +1,33 @@ +//===-------------------------- WaferDialect.h -----------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_DIALECT_WAFER_IR_DIALECT_H +#define MLIR_DIALECT_WAFER_IR_DIALECT_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/SymbolTable.h" +#include "mlir/IR/TypeSupport.h" +#include "mlir/IR/Types.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +//===----------------------------------------------------------------------===// +// Wafer Wafer Operations +//===----------------------------------------------------------------------===// +#include "wafer/Dialect/IR/WaferDialect.h.inc" + +// Include the auto-generated header file containing the declarations of the +// TritonStructured operations. +#define GET_OP_CLASSES +#include "wafer/Dialect/IR/WaferEnums.h.inc" +#include "wafer/Dialect/IR/WaferOps.h.inc" + +#endif // MLIR_DIALECT_WAFER_IR_DIALECT_H diff --git a/third_party/wafer/include/wafer/Dialect/IR/WaferDialect.td b/third_party/wafer/include/wafer/Dialect/IR/WaferDialect.td new file mode 100755 index 00000000..6264da28 --- /dev/null +++ b/third_party/wafer/include/wafer/Dialect/IR/WaferDialect.td @@ -0,0 +1,43 @@ +//===----------------------- WaferDialect.td -------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef WAFER_DIALECT +#define WAFER_DIALECT + +include "mlir/IR/OpBase.td" + +def WaferDialect : Dialect { + let name = "wafer"; + + let cppNamespace = "::mlir::wafer"; + + let summary = "The Wafer Wafer IR in MLIR"; + + let description = [{ + Wafer Wafer Dialect. + + Dependent Dialects: + * MK + * Memref + * Bufferization + }]; + + let dependentDialects = [ + ]; + + let extraClassDeclaration = [{ + void registerTypes(); + }]; + + // let hasConstantMaterializer = 1; + // let useDefaultTypePrinterParser = 1; + let usePropertiesForAttributes = 1; +} + +include "wafer/Dialect/IR/WaferTypes.td" + +#endif // WAFER_DIALECT diff --git a/third_party/wafer/include/wafer/Dialect/IR/WaferOps.h b/third_party/wafer/include/wafer/Dialect/IR/WaferOps.h new file mode 100755 index 00000000..39a11590 --- /dev/null +++ b/third_party/wafer/include/wafer/Dialect/IR/WaferOps.h @@ -0,0 +1,26 @@ +//===-------------------------- WaferOps.h ---------------------*- C++ -*---===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_DIALECT_WAFER_IR_OPS_H +#define MLIR_DIALECT_WAFER_IR_OPS_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/SymbolTable.h" +#include "mlir/IR/TypeSupport.h" +#include "mlir/IR/Types.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#define GET_OP_CLASSES +#include "wafer/Dialect/IR/WaferEnums.h.inc" +#include "wafer/Dialect/IR/WaferOps.h.inc" + +#endif // MLIR_DIALECT_WAFER_IR_DIALECT_H diff --git a/third_party/wafer/include/wafer/Dialect/IR/WaferOps.td b/third_party/wafer/include/wafer/Dialect/IR/WaferOps.td new file mode 100755 index 00000000..a610af18 --- /dev/null +++ b/third_party/wafer/include/wafer/Dialect/IR/WaferOps.td @@ -0,0 +1,1278 @@ + +//===---------------------- WaferOps.td ------------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Definition of Wafer's Wafer ML accelerator operations. +// +// Data format supported by Wafer ML accelerator are: +// f16,fp16,tf32,fp32 +// +// For Wafer accelerator unsupported data type, we can either convert it by +// using `TsmConvert`, or lower the operations to run on RISC-V controller +// instead. +// +// NOTE: CHANGING THE ARGUMENTS AND RETURNS OF ANY OPS RESULT IN THE CHANGE OF +// THEIR RUNTIME INTERFACE AND IMPLEMENTATION IN crt/Target/Wafer. +// +//===----------------------------------------------------------------------===// + +#ifndef WAFER_OPS +#define WAFER_OPS + +include "wafer/Dialect/IR/WaferAttrDefs.td" +include "wafer/Dialect/IR/WaferTypes.td" +include "mlir/Interfaces/SideEffectInterfaces.td" // Pure +include "mlir/Interfaces/InferTypeOpInterface.td" // SameOperandsAndResultType +include "mlir/IR/OpBase.td" + +// +// Interfaces +// +def GlobalMemory : Resource<"::mlir::triton::GlobalMemory">; + +class WaferOp traits = []> : + Op { +} + +def MemRefOrInt + : AnyTypeOf<[AnyMemRef, AnySignlessIntegerOrIndex], + "MemRef or Int as address type.", "::mlir::Type">; + +// ============================================================================= +// 4.8/4.9 DDR and SPM transfer ops +// ============================================================================= + +def RdmaOp : WaferOp<"rdma", [ + AttrSizedOperandSegments + ]> { + + let summary = "Copy data from global memory DDR(dram) to per thread local SPM(sram)"; + + let description = [{ + Copy data from global memory DDR(dram) to per thread local SPM(sram). + }]; + + let arguments = ( + ins + MemRefOrInt:$source, // The source address in DDR + MemRefOrInt:$target, // The target address in SPM + Variadic:$src_shape, // src shape + Variadic:$src_strides, // src strides + Variadic:$dst_shape, // dst shape + Variadic:$dst_strides, // dst strides + I32Attr:$rank, // rank + I32Attr:$elem_bytes, // elem bytes + I32Attr:$fmt // elem fmt + ); + + let results = (outs I64:$dst); // The dest address in SPM +} + +def WdmaOp : WaferOp<"wdma", [ + AttrSizedOperandSegments + ]> { + let summary = "Copy data from per thread local SPM(sram) to global memory DDR(dram)"; + + let description = [{ + Copy data from per thread local SPM(sram) to global memory DDR(dram). + }]; + + let arguments = ( + ins + MemRefOrInt:$source, // The source address in DDR + MemRefOrInt:$target, // The target address in SPM + Variadic:$src_shape, // src shape + Variadic:$src_strides, // src strides + Variadic:$dst_shape, // dst shape + Variadic:$dst_strides, // dst strides + I32Attr:$rank, // rank + I32Attr:$elem_bytes, // elem bytes + I32Attr:$fmt // elem fmt + ); + + let results = (outs I64:$dst); // The dest address in DDR +} + +def Rdma4dOp : WaferOp<"rdma4d"> { + let arguments = ( + ins + MemRefOrInt:$target, // The target address in SPM + MemRefOrInt:$source, // The source address in DDR + I32:$elem_count, + I32:$stride0, + I32:$iteration0, + I32:$stride1, + I32:$iteration1, + I32:$stride2, + I32:$iteration2, + I32Attr:$fmt // elem fmt + ); + + let results = (outs); +} + +def Wdma4dOp : WaferOp<"wdma4d"> { + let arguments = ( + ins + MemRefOrInt:$target, // The target address in DDR + MemRefOrInt:$source, // The source address in SPM + I32:$elem_count, + I32:$stride0, + I32:$iteration0, + I32:$stride1, + I32:$iteration1, + I32:$stride2, + I32:$iteration2, + I32Attr:$fmt // elem fmt + ); + + let results = (outs); +} + +def Rdma1dOp : WaferOp<"rdma1d"> { + let arguments = ( + ins + MemRefOrInt:$target, // The target address in SPM + MemRefOrInt:$source, // The source address in DDR + I32:$elem_count, + I32Attr:$fmt // elem fmt + ); + + let results = (outs); +} + +def Wdma1dOp : WaferOp<"wdma1d"> { + let arguments = ( + ins + MemRefOrInt:$target, // The target address in DDR + MemRefOrInt:$source, // The source address in SPM + I32:$elem_count, + I32Attr:$fmt // elem fmt + ); + + let results = (outs); +} + +def MemCopyOp : WaferOp<"memcpy", [ + MemoryEffects<[MemRead, MemWrite]> +]> { + + let summary = "Copy data from local SPM(sram) to local SPM(sram)"; + + let description = [{ + Copy data from local SPM(sram) to local SPM(sram). + }]; + + let arguments = ( + ins + MemRefOrInt:$input, // The source address in DDR + MemRefOrInt:$out, // Out vector address + MemRefOrInt:$elem_count, // elem count + I32Attr:$fmt // elem fmt + ); + + let results = (outs I64:$dst); // The dest address in SPM +} + +// ============================================================================= +// 4.4~6 TsmConv, TsmDepthwiseConv, TsmBackwardConv +// ============================================================================= + +def ConvOp : WaferOp<"conv", [Pure]> { + let summary = "Convolution engine intrinsic runtime API"; + + let description = [{ + A common convolution op for TsmConv, TsmDepthwiseConv, TsmBackwardConv. + This TsmConv is not a 1 to 1 map to Wafer's TsmConv intrinsic, it is + the wrap of all APIs related to TsmConv. This op wraps the following APIs: + TsmNewConv, TsmDeleteConv, AddInput, AddWeight, AddBias, AddOutput, + SetOpType, SetNegativeAxisScale, SetPositiveAxisScale, SetSparse, SetPsum, + SetPads, SetUnPads, SetKernelStrides, SetDilations, EnableRelu, + EnableLeakyRelu, DisableRelu, DisableLeakyRelu, SetQuant. + }]; + + let arguments = ( + ins + I64Attr:$op_type, // 0: conv, 1: depthwise conv, 2: backward conv, + // 3: gemm + MemRefOrInt:$src_activation, // Input activation addr in SPM + I32ArrayAttr:$src_dims, // dims of src activation in NHWC format + MemRefOrInt:$weight, // Input weight addr in SPM + I16Attr:$weight_dims, // dims of weight(conv kernel) in Kx, Ky, Sx, Sy + // Where K and S is short for size(K) and step(S) + BoolAttr:$en_bias, // Enable bias add + MemRefOrInt:$src_bias, // The address of bias in SPM + BoolAttr:$en_neg_scale, // Enable negative axis scale + MemRefOrInt:$src_neg_scale, // The address of negative scale data in SPM + BoolAttr:$en_pos_scale, // Enable positive axis scale + MemRefOrInt:$src_pos_scale, // The address of positive scale data in SPM + BoolAttr:$en_sparse, // Enable sparse + MemRefOrInt:$src_sparse, // The sparse matrix addr in SPM + BoolAttr:$en_psum, // Enable psum? TODO: Production sum? + MemRefOrInt:$src_psum, // psum addr in SPM? + I32ArrayAttr:$pads, // Pad in top, bottom, left, right order + I32ArrayAttr:$unpads, // Unpad in top, bottom, left, right order + I32ArrayAttr:$strides, // Kernel strids in Kx, Ky, Sx, Sy + I32ArrayAttr:$dilations, // dialation d0, d1 for conv/backwardconv + BoolAttr:$en_leaky_relu, // Enable LeakyRelu or normal Relu + I32ArrayAttr:$out_dims, // dims of output in NHWC format + I64Attr:$src_fmt, // Data format of src activation + I64Attr:$weight_fmt, // Data format of weight + I64Attr:$out_fmt // Data format of output + // The param of SetQuant() is unused + ); + + // Output matrix C addr in SPM + let results = (outs I64:$dst); +} + +// ============================================================================= +// 4.7. TsmGemm +// ============================================================================= + +def GemmOp : WaferOp<"gemm", []> { + let summary = "Gemm engine intrinsic runtime API"; + + let description = [{ + This TsmGemm is not a 1 to 1 map to Wafer's TsmGemm intrinsic, it is + the wrap of all APIs related to TsmGemm. This op wraps the following APIs: + TsmNewGemm, TsmDeleteGemm, AddInput, ConfigMKN, AddOutput, SetPsum, + SetTransflag, SetQuant, ConfigBatch, EnableRelu, EnableLeakyRelu, + DisableRelu, DisableLeakyRelu, AddBias, SetNegativeAxisScale, + SetPositiveAxisScale. + }]; + + let arguments = ( + ins + MemRefOrInt:$src_a, // Input matrix A addr in SPM + MemRefOrInt:$src_b, // Input matrix B addr in SPM + MemRefOrInt:$src_bias, // The address of bias in SPM + // Output and initial zeroes buffer + // FIXME: Whether need add side effect to source operands? + Arg:$dst, + I32ArrayAttr:$dims, // The dimensions of M, K, N + BoolAttr:$en_psum, // Enable psum. Used as accumulate buffer + MemRefOrInt:$psum_addr, // The address of psum in SPM, Always same to output + BoolAttr:$trans_src_a, // Should matrix A be transposed + BoolAttr:$trans_src_b, // Should matrix B be transposed + I32Attr:$batch_src_a, // The batch of matrix A + I32Attr:$batch_src_b, // The batch of matrix B + I32Attr:$relu_mode, // Enable LeakyRelu or normal Relu or none + BoolAttr:$en_bias, // Enable bias add. Only support per channel(C dim), and int8 type + BoolAttr:$en_neg_scale, // Enable negative axis scale + MemRefOrInt:$src_neg_scale, // The address of negative scale data in SPM + BoolAttr:$en_pos_scale, // Enable positive axis scale + MemRefOrInt:$src_pos_scale, // The address of positive scale data in SPM + I32Attr:$src_fmt, // Input matrix data format + I32Attr:$dst_fmt // Output matrix data format + // The param of SetQuant() is unused + ); + + // Output matrix C addr in SPM + let results = (outs Variadic:$output); +} + + +// ============================================================================= +// Tsm crt ChannelNorm/Dechannelnorm +// ============================================================================= + +def ChannelNormOp : WaferOp<"channel_norm", []> { + let summary = "Align channel dim."; + + let description = [{ + Align (N,H,W,C) to (N,cx,H,W,64) + (N,cx,H,W,c0), + which align_base = 64, cx = C/align_base + c0 = C%align_base + c0_align = get_c0_align(c0) + }]; + + let arguments = ( + ins + MemRefOrInt:$src, // Input tensor address in SPM + Arg:$dst, // Output tensor address in SPM + DenseI64ArrayAttr:$shape, // The shape info of src + // I16Attr:$cx, + I16Attr:$c0_align, + I16Attr:$dtype_size + ); + + // Output matrix C addr in SPM + let results = (outs Variadic:$output); +} + +def DechannelNormOp : WaferOp<"dechannel_norm", []> { + let summary = "Inverse operation of channelnorm."; + + let description = [{ + Trans (N,cx,H,W,64) + (N,cx,H,W,c0) back to (N,H,W,C). + }]; + + let arguments = ( + ins + MemRefOrInt:$src, // Input tensor address in SPM + Arg:$dst, // Output tensor address in SPM + DenseI64ArrayAttr:$shape, // The shape info of src + // I16Attr:$cx, + I16Attr:$c0_align, + I16Attr:$dtype_size + ); + + // Output matrix C addr in SPM + let results = (outs Variadic:$output); +} + +// ============================================================================= +// 4.10. TsmArith +// ============================================================================= + +class UnaryOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input, // Input vector address + Arg:$out, // Out vector address + MemRefOrInt:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs Variadic:$dst); +} + +def AbsVVOp : UnaryOp<"absvv"> { + let summary = "Absolute value of input vector"; +} +def SqrtVVOp : UnaryOp<"sqrtvv", [Pure, Elementwise]> {} +def RsqrtVVOp : UnaryOp<"rsqrtvv", [Pure, Elementwise]> {} +def NegVVOp : UnaryOp<"negvv", [Pure, Elementwise]> {} +def RecipVVOp : UnaryOp<"recipvv", [Pure, Elementwise]> {} +def SquareVVOp : UnaryOp<"squarevv", [Pure, Elementwise]> {} + +class BinaryVVOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input0, // First input vector address + MemRefOrInt:$input1, // Second vector address + Arg:$out, // Out vector address + MemRefOrInt:$elem_count, // Number of input elements + I16Attr:$rnd_mode, // round mode + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs Variadic:$dst); +} + +def AddVVOp : BinaryVVOp<"addvv"> { + let summary = "Add two vectors element-wise"; +} +def SubVVOp : BinaryVVOp<"subvv">; +def MulVVOp : BinaryVVOp<"mulvv">; +def DivVVOp : BinaryVVOp<"divvv">; +def MaxVVOp : BinaryVVOp<"maxvv">; +def MinVVOp : BinaryVVOp<"minvv">; + +class BinaryVSOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input0, // First input vector address + I32:$value, // Const value + Arg:$out, // Out vector address + MemRefOrInt:$elem_count, // Number of input elements + I16Attr:$rnd_mode, // round mode + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs Variadic:$dst); +} + +def AddVSOp : BinaryVSOp<"addvs"> { + let summary = "Add input vector and constant value"; +} +def SubVSOp : BinaryVSOp<"subvs">; +def MulVSOp : BinaryVSOp<"mulvs">; +def DivVSOp : BinaryVSOp<"divvs">; + +// ... + +// ============================================================================= +// 4.11. TsmRelation +// ============================================================================= + +class RelationVVOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input0, // First input vector address + MemRefOrInt:$input1, // Second vector address + Arg:$out, // Out vector address + MemRefOrInt:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs Variadic:$dst); +} + +def BoolEqualVV : RelationVVOp<"boolequalvv"> { + let summary = "compare two input value, if equal, return true"; +} + +def BoolUnEqualVV : RelationVVOp<"boolunequalvv"> { + let summary = "compare two input value, if unequal, return true"; +} + +def BoolGreaterEqualVV : RelationVVOp<"boolgreatrequalvv"> { + let summary = "compare two input value, if src0 >= src1, return true"; +} + +def BoolGreaterVV : RelationVVOp<"boolgreatervv"> { + let summary = "compare two input value, if src0 > src1, return true"; +} + +def BoolLessEqualVV : RelationVVOp<"boollessequalvv"> { + let summary = "compare two input value, if src0 <= src1, return true"; +} + +def BoolLessThenVV : RelationVVOp<"boollessthenvv"> { + let summary = "compare two input value, if src0 < src1, return true"; +} + +def EqualVV : RelationVVOp<"equalvv"> { + let summary = "compare two input value, if equal, return 1.0"; +} + +def UnEqualVV : RelationVVOp<"unequalvv"> { + let summary = "compare two input value, if unequal, return 1.0"; +} + +def GreaterEqualVV : RelationVVOp<"greatrequalvv"> { + let summary = "compare two input value, if src0 >= src1, return 1.0"; +} + +def GreaterVV : RelationVVOp<"greatervv"> { + let summary = "compare two input value, if src0 > src1, return 1.0"; +} + +def LessEqualVV : RelationVVOp<"lessequalvv"> { + let summary = "compare two input value, if src0 <= src1, return 1.0"; +} + +def LessThenVV : RelationVVOp<"lessthenvv"> { + let summary = "compare two input value, if src0 < src1, return 1.0"; +} + +class RelationVSOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input0, // First input vector address + I32:$value, // Const value + Arg:$out, // Out vector address + MemRefOrInt:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs Variadic:$dst); +} + +def BoolEqualVS : RelationVSOp<"boolequalvs"> { + let summary = "compare input value with ConstantOp, if equal, return true"; +} + +def BoolUnEqualVS : RelationVSOp<"boolunequalvs"> { + let summary = "compare input value with ConstantOp, if unequal, return true"; +} + +def BoolGreaterEqualVS : RelationVSOp<"boolgreatrequalvs"> { + let summary = "compare input value with ConstantOp, if src0 >= src1, return true"; +} + +def BoolGreaterVS : RelationVSOp<"boolgreatervs"> { + let summary = "compare input value with ConstantOp, if src0 > src1, return true"; +} + +def BoolLessEqualVS : RelationVSOp<"boollessequalvs"> { + let summary = "compare input value with ConstantOp, if src0 <= src1, return true"; +} + +def BoolLessThenVS : RelationVSOp<"boollessthenvs"> { + let summary = "compare input value with ConstantOp, if src0 < src1, return true"; +} + +def EqualVS : RelationVSOp<"equalvs"> { + let summary = "compare input value with ConstantOp, if equal, return 1.0"; +} + +def UnEqualVS : RelationVSOp<"unequalvs"> { + let summary = "compare input value with ConstantOp, if unequal, return 1.0"; +} + +def GreaterEqualVS : RelationVSOp<"greatrequalvs"> { + let summary = "compare input value with ConstantOp, if src0 >= src1, return 1.0"; +} + +def GreaterVS : RelationVSOp<"greatervs"> { + let summary = "compare input value with ConstantOp, if src0 > src1, return 1.0"; +} + +def LessEqualVS : RelationVSOp<"lessequalvs"> { + let summary = "compare input value with ConstantOp, if src0 <= src1, return 1.0"; +} + +def LessThenVS : RelationVSOp<"lessthenvs"> { + let summary = "compare input value with ConstantOp, if src0 < src1, return 1.0"; +} + + +// ... +// ============================================================================= +// 4.12. TsmLogic +// ============================================================================= + +class BinaryLogicVVOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input0, // First input vector address + MemRefOrInt:$input1, // Second vector address + Arg:$out, // Out vector address + MemRefOrInt:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs Variadic:$dst); +} + +def AndVV : BinaryLogicVVOp<"andvv"> { + let summary = "And operation on elements at the same position. If the element is not 0, it is represented as 1."; +} + +def OrVV : BinaryLogicVVOp<"orvv"> { + let summary = "OR operation on elements at the same position. If the element is not 0, it is represented as 1."; +} + +def XorVV : BinaryLogicVVOp<"xorvv"> { + let summary = "XOR operation on elements at the same position. If the element is not 0, it is represented as 1."; +} + +class BoolUnaryLogicVVOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input, // input vector address + Arg:$out, // Out vector address + MemRefOrInt:$elem_count // Number of input elements + ); + let results = (outs Variadic:$dst); +} + +def BoolNotV : BoolUnaryLogicVVOp<"boolnotvv"> { + let summary = "Not operation on elements at each bit position. src_elem_num must be an integer multiple of 8"; +} + +class BoolBinaryLogicVVOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input0, // First input vector address + MemRefOrInt:$input1, // Second vector address + Arg:$out, // Out vector address + MemRefOrInt:$elem_count // Number of input elements + ); + let results = (outs Variadic:$dst); +} + +def BoolAndV : BoolBinaryLogicVVOp<"boolandvv"> { + let summary = "And operation on elements at each bit position. src_elem_num must be an integer multiple of 8"; +} + +def BoolXorV : BoolBinaryLogicVVOp<"boolxorvv"> { + let summary = "Xor operation on elements at each bit position. src_elem_num must be an integer multiple of 8"; +} + +def BoolOrV : BoolBinaryLogicVVOp<"boolorvv"> { + let summary = "Or operation on elements at each bit position. src_elem_num must be an integer multiple of 8"; +} + + +// ============================================================================= +// 4.13. TsmTranscendental +// ============================================================================= + +def Log2Op : UnaryOp<"log2", []> { + let summary = "Logarithm based 2"; +} +def LnOp : UnaryOp<"ln", []> { + let summary = "Logarithm based e"; +} +def Pow2Op : UnaryOp<"pow2", []> { + let summary = "2 ** x"; +} +def ExpOp : UnaryOp<"exp", []> { + let summary = "Exponential with high precision"; +} +def ExplpOp : UnaryOp<"explp", []> { + let summary = "Exponential with low precision"; +} +def SinOp : UnaryOp<"sin", []> { + let summary = "Sine"; +} +def CosOp : UnaryOp<"cos", []> { + let summary = "Cosine"; +} + +// ============================================================================= +// 4.13. TsmActivation +// ============================================================================= + +class ActivationOp traits> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input, // Input vector address + Arg:$out, // Out vector address + AnySignlessIntegerOrIndex:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); + + let results = (outs Variadic:$dst); +} + +def Tanh : ActivationOp<"tanh", []> { + let summary = "Hyperbolic tangent"; +} +def Sigmoid : ActivationOp<"sigmoid", []> { + let summary = "Logistic sigmoid"; +} +def GeluNone : ActivationOp<"gelu_none", []> { + let summary = "Gaussian Error Linear Unit by None"; + let arguments = (ins + MemRefOrInt:$input, // Input vector address + Arg:$out, // Out vector address + AnySignlessIntegerOrIndex:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); +} +def GeluTanh : ActivationOp<"gelu_tanh", []> { + let summary = "Gaussian Error Linear Unit by Tanh"; + let arguments = (ins + MemRefOrInt:$input, // Input vector address + Arg:$buffer,// Buffer vector address + Arg:$out, // Out vector address + AnySignlessIntegerOrIndex:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); +} +def Relu : ActivationOp<"relu", []> { + let summary = "Rectified linear unit"; +} +def Satrelu : ActivationOp<"satrelu", []> { + let summary = "Saturated ReLU"; +} +def Leakyrelu : ActivationOp<"leakyrelu", []> { + let summary = "Leaky rectified linear unit"; +} +def Softplus : ActivationOp<"softplus", []> { + let summary = "Smooth approximation of ReLU"; +} + +// ============================================================================= +// 4.15. TsmReduce +// ============================================================================= + +class Reduce : WaferOp { + let summary = "Reduction engine intrinsic runtime API"; + + let description = [{ + Includes ReduceSum, ReduceAvg, ReduceMin and ReduceMax interfaces. + Mapping between `dim` and NCHW: + Reduction on C: dim=0 + Reduction on W: dim=1 + Reduction on H: dim=2 + Reduction on HW: dim=4 + }]; + + let arguments = ( + ins + MemRefOrInt:$src, // Input tensor address in SPM + Arg:$dst, // Output tensor address in SPM + UI32Attr:$dim, // Which dimension to be reduced + I64ArrayAttr:$shape, // The shape info of src + I16Attr:$fmt // The data format of src & dst + ); + + // Output tensor address in SPM + let results = (outs Variadic); +} + +def ReduceSumOp : Reduce<"reduce_sum">; +def ReduceAvgOp : Reduce<"reduce_avg">; +def ReduceMaxOp : Reduce<"reduce_max">; +def ReduceMinOp : Reduce<"reduce_min">; +def ReduceMulOp : Reduce<"reduce_mul">; +// ============================================================================= +// 4.15. TsmMaskDataMove +// ============================================================================= + +def MaskMoveOp : WaferOp<"mask_move", []> { + let summary = "Mask data move engine intrinsic runtime API"; + + let description = [{ When mask is 1, extract the data from src and write it to dst. +When mask=0, the corresponding elements of dst remain unchanged. + }]; + + let arguments = ( + ins + MemRefOrInt:$source, // The source address in SPM + // The target address in SPM + Arg:$target, + AnySignlessIntegerOrIndex:$elem_count, // Number of elements to be copied + MemRefOrInt:$mask, + I32Attr:$fmt + ); + + // The dst address is not used, use target in arguments instead. + let results = (outs Variadic:$dst); +} + +// ============================================================================= +// 4.19. TsmConvert instructions +// ============================================================================= + +class ZeroPointConvertOp traits> : + WaferOp { + let arguments = (ins + MemRefOrInt:$src, + Arg:$dst, + UI32Attr:$zero_point, + UI32Attr:$elem_count + ); +} + +def INT8ToFP16Op : ZeroPointConvertOp<"int8_fp16", []> { + let summary = "Data format from int8 to fp16"; +} +def INT8ToBF16Op : ZeroPointConvertOp<"int8_bf16", []> { + let summary = "Data format from int8 to bf16"; +} +def INT8ToFP32Op : ZeroPointConvertOp<"int8_fp32", []> { + let summary = "Data format from int8 to fp32"; +} +def INT8ToTF32Op : ZeroPointConvertOp<"int8_tf32", []> { + let summary = "Data format from int8 to tf32"; +} + +class RoundConvertOp traits> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input, + Arg:$output, + AnySignlessIntegerOrIndex:$elem_count, + I16Attr:$rnd_mode + ); + let results = (outs I64:$dst); +} + +def INT16ToBF16Op : RoundConvertOp<"int16_bf16", []> { + let summary = "Data format from int16 to bf16"; +} +def INT16ToFP32Op : RoundConvertOp<"int16_fp32", []> { + let summary = "Data format from int16 to fp32"; +} +def INT16ToTF32Op : RoundConvertOp<"int16_tf32", []> { + let summary = "Data format from int16 to tf32"; +} +def INT32ToFP16Op : RoundConvertOp<"int32_fp16", []> { + let summary = "Data format from int32 to fp16"; +} +def INT32ToBF16Op : RoundConvertOp<"int32_bf16", []> { + let summary = "Data format from int32 to bf16"; +} +def INT32ToFP32Op : RoundConvertOp<"int32_fp32", []> { + let summary = "Data format from int32 to fp32"; +} +def INT32ToTF32Op : RoundConvertOp<"int32_tf32", []> { + let summary = "Data format from int32 to tf32"; +} +def BF16ToINT16Op : RoundConvertOp<"bf16_int16", []> { + let summary = "Data format from bf16 to int16"; +} +def BF16ToINT32Op : RoundConvertOp<"bf16_int32", []> { + let summary = "Data format from bf16 to int32"; +} +def FP16ToINT8Op : RoundConvertOp<"fp16_int8", []> { + let summary = "Data format from fp16 to int8"; +} +def FP16ToINT16Op : RoundConvertOp<"fp16_int16", []> { + let summary = "Data format from fp16 to int16"; +} +def FP16ToINT32Op : RoundConvertOp<"fp16_int32", []> { + let summary = "Data format from fp16 to int32"; +} +def FP16ToBF16Op : RoundConvertOp<"fp16_bf16", []> { + let summary = "Data format from fp16 to bf16"; +} +def FP32ToINT8Op : RoundConvertOp<"fp32_int8", []> { + let summary = "Data format from fp32 to int8"; +} +def FP32ToINT16Op : RoundConvertOp<"fp32_int16", []> { + let summary = "Data format from fp32 to int16"; +} +def FP32ToINT32Op : RoundConvertOp<"fp32_int32", []> { + let summary = "Data format from fp32 to int32"; +} +def FP32ToFP16Op : RoundConvertOp<"fp32_fp16", []> { + let summary = "Data format from fp32 to fp16"; +} +def FP32ToBF16Op : RoundConvertOp<"fp32_bf16", []> { + let summary = "Data format from fp32 to bf16"; +} +def FP32ToTF32Op : RoundConvertOp<"fp32_tf32", []> { + let summary = "Data format from fp32 to tf32"; +} +def TF32ToINT8Op : RoundConvertOp<"tf32_int8", []> { + let summary = "Data format from tf32 to int8"; +} +def TF32ToINT16Op : RoundConvertOp<"tf32_int16", []> { + let summary = "Data format from tf32 to int16"; +} +def TF32ToINT32Op : RoundConvertOp<"tf32_int32", []> { + let summary = "Data format from tf32 to int32"; +} +def TF32ToBF16Op : RoundConvertOp<"tf32_bf16", []> { + let summary = "Data format from tf32 to bf16"; +} + +class NormalConvertOp traits> : + WaferOp { + let arguments = (ins + MemRefOrInt:$input, + Arg:$output, + AnySignlessIntegerOrIndex:$elem_count + ); + let results = (outs I64:$dst); +} + +def INT16ToFP16Op : NormalConvertOp<"int16_fp16", []> { + let summary = "Data format from int16 to fp16"; +} +def BF16ToINT8Op : NormalConvertOp<"bf16_int8", []> { + let summary = "Data format from bf16 to int8"; +} +def BF16ToFP16Op : NormalConvertOp<"bf16_fp16", []> { + let summary = "Data format from bf16 to fp16"; +} +def BF16ToFP32Op : NormalConvertOp<"bf16_fp32", []> { + let summary = "Data format from bf16 to fp32"; +} +def BF16ToTF32Op : NormalConvertOp<"bf16_tf32", []> { + let summary = "Data format from bf16 to tf32"; +} +def FP16ToFP32Op : NormalConvertOp<"fp16_fp32", []> { + let summary = "Data format from fp16 to fp32"; +} +def FP16ToTF32Op : NormalConvertOp<"fp16_tf32", []> { + let summary = "Data format from fp16 to tf32"; +} +def TF32ToFP16Op : NormalConvertOp<"tf32_fp16", []> { + let summary = "Data format from tf32 to fp16"; +} +def TF32ToFP32Op : NormalConvertOp<"tf32_fp32", []> { + let summary = "Data format from tf32 to fp16"; +} + +// Ref: mlir/include/mlir/IR/BuiltinTypes.td +// Micro scaling format conversion operations +def FP8E4M3ToBF16Op : NormalConvertOp<"fp8E4M3_bf16", []> { + let summary = "Convert data format from FP8 (4 exponent bits, 3 mantissa bits) to BF16"; +} +def FP8E4M3FNToBF16Op : NormalConvertOp<"fp8E4M3FN_bf16", []> { + let summary = "Convert data format from FP8 (4 exponent bits, 3 mantissa bits) to BF16, with NAN no inf"; +} + +def FP8E5M2ToBF16Op : NormalConvertOp<"fp8E5M2_bf16", []> { + let summary = "Convert data format from FP8 (5 exponent bits, 2 mantissa bits) to BF16"; +} +def FP4E2M1ToBF16Op : NormalConvertOp<"fp4E2M1_bf16", []> { + let summary = "Convert data format from FP4 (2 exponent bits, 1 mantissa bit) to BF16"; +} + +def FP8E4M3ToFP16Op : NormalConvertOp<"fp8E4M3_fp16", []> { + let summary = "Convert data format from FP8 (4 exponent bits, 3 mantissa bits) to FP16"; +} +def FP8E4M3FNToFP16Op : NormalConvertOp<"fp8E4M3FN_fp16", []> { + let summary = "Convert data format from FP8 (4 exponent bits, 3 mantissa bits) to FP16, with NAN no inf"; +} + +def FP8E5M2ToFP16Op : NormalConvertOp<"fp8E5M2_fp16", []> { + let summary = "Convert data format from FP8 (5 exponent bits, 2 mantissa bits) to FP16"; +} +def FP4E2M1ToFP16Op : NormalConvertOp<"fp4E2M1_fp16", []> { + let summary = "Convert data format from FP4 (2 exponent bits, 1 mantissa bit) to FP16"; +} + +class MXFPScaleOp : WaferOp { + let summary = "Scaling up bf16/fp16 tensor with e8m0 scale."; + + let description = [{ + Performs conversion of MXFP numbers to BF16/FP16 format according to the + Open Compute Project (OCP) microscaling formats specification v1.0. + See: https://www.opencompute.org/documents/ocp-microscaling-formats-mx-v1-0-spec-final-pdf + }]; + + let arguments = ( + ins + MemRefOrInt:$src, // Source tensor address in SPM + MemRefOrInt:$scale, // Scale factor tensor address in SPM + MemRefOrInt:$dst, // Destination tensor address in SPM + I32Attr:$elemCount // Number of elements + ); + + let results = (outs Variadic:$output); // Output tensor address in SPM +} + +def MXFPScaleBF16Op : MXFPScaleOp<"MXFP_scale_bf16"> {} +def MXFPScaleFP16Op : MXFPScaleOp<"MXFP_scale_fp16"> {} + +// ============================================================================= +// 4.20. TsmPeripheral instructions +// ============================================================================= + +def CountOp : WaferOp<"count", [Pure]> { + let summary = "Count the non-zero elements from given tensor"; + + let arguments = ( + ins + MemRefOrInt:$src, // Input tensor address in SPM + I32Attr:$elem_count, // TODO: Ask Wafer for explain. + //I64Attr:$p_wb_data0, // TODO: Ask Wafer for explain. + //I64Attr:$p_wb_data1, // TODO: Ask Wafer for explain. + I16Attr:$fmt + ); + + // The output tensor address in SPM + let results = (outs MemRefOrInt:$dst); +} + +def MemsetOp : WaferOp<"memset", [ + AttrSizedOperandSegments +]> { + let summary = "Write given `value` to range of address on SPM(sram)"; + + let arguments = ( + ins + MemRefOrInt:$target, // SPM address to be memset + I32:$value, // Value to be written + Variadic:$dst_shape, // src shape + Variadic:$dst_strides, // src strides + I32Attr:$rank, // rank + I16Attr:$fmt + ); + + // The address updated by memset in SPM + let results = (outs MemRefOrInt:$dst); +} + +def Bit2FpOp : WaferOp<"bit2fp", []> { + let summary = "Convert a vector of the bitwise into the fp vector"; + + let arguments = (ins + MemRefOrInt:$src, // Input tensor + Arg:$target, + AnySignlessIntegerOrIndex:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs I64:$dst); +} + +def ArgMaxOp : WaferOp<"argmax", []> { + let summary = "Return a max value inner a vector and its corresponding index"; + + let arguments = (ins + MemRefOrInt:$src, // First input vector address + Arg:$value, // Address + Arg:$index, // Address + I32Attr:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); + + let results = (outs Variadic:$dst); +} + +def ArgMinOp : WaferOp<"argmin", []> { + let summary = "Return a min value inner a vector and its corresponding index"; + + let arguments = (ins + MemRefOrInt:$src, // First input vector address + Arg:$value, // Address + Arg:$index, // Address + I32Attr:$elem_count, // Number of input elements + I16Attr:$fmt // The data format of src & dst + ); + + let results = (outs Variadic:$dst); +} + +def BilinearOp : WaferOp<"bilinear", []> { + let summary = "Bilinear interpolation"; + + let arguments = (ins + UI64:$src, // Input tensor with the NHWC format + I32ArrayAttr:$src_shape, // Input tensor shape + I32ArrayAttr:$dst_shape, // Output tensor shape + F32:$scale_w, // Input tensor "w" divided by output tensor "w" + F32:$scale_h, // Input tensor "h" divided by output tensor "h" + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs UI64:$dst); +} + +def Lut16Op : WaferOp<"lut16", []> { + let summary = "16-bit lookup table"; + + let arguments = (ins + MemRefOrInt:$src, // Vector offset with respect to LUT + UI64:$lut16, + I32Attr:$src_elem_count, // Number of elements in vector offset + I32Attr:$lut_elem_count // Number of elements in LUT + ); + + let results = (outs UI64:$dst); +} + +def Lut32Op : WaferOp<"lut32", []> { + let summary = "32-bit lookup table"; + + let arguments = (ins + MemRefOrInt:$src, // Vector offset with respect to LUT + UI64:$lut32, + I32Attr:$src_elem_count, // Number of elements in vector offset + I32Attr:$lut_elem_count // Number of elements in LUT + ); + + let results = (outs UI64:$dst); +} + +def RandGenOp : WaferOp<"randgen", []> { + let summary = "Generate random numbers using two 64-bit seeds"; + + let arguments = (ins + MemRefOrInt:$src0, // First seed vector address (16 x u64) + MemRefOrInt:$src1, // Second seed vector address (16 x u64) + Arg:$dst0, + Arg:$dst1, + Arg:$dst2, + I32Attr:$elem_num, // SDK byte count, a multiple of 128 + I16Attr:$fmt // The date format of random value + ); +} + + +// +// 4.21. TsmDataMove +// + +class TransformOp traits = []> : + WaferOp { + let arguments = (ins + MemRefOrInt:$source, // Input tensor + MemRefOrInt:$target, // Output tensor + DenseI32ArrayAttr:$src_shape, // Input shape + DenseI32ArrayAttr:$dst_shape, // Output shape + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs I64:$dst); +} + +def Mirror : TransformOp<"mirror", []> { + let summary = "Horizontal mirror to a matrix"; +} +def Transpose : TransformOp<"transpose", []> { + let summary = "Transpose a matrix"; +} +def Rotate90 : TransformOp<"rotate90", []> { + let summary = "Rotate a matrix 90 degree clockwise"; +} +def Rotate180 : TransformOp<"rotate180", []> { + let summary = "Rotate a matrix 180 degree clockwise"; +} +def Rotate270 : TransformOp<"rotate270", []> { + let summary = "Rotate a matrix 270 degree clockwise"; +} +def Nchw2nhwc : TransformOp<"nchw2nhwc", []> { + let summary = "Tranform a tensor from nchw to nhwc"; +} +def Nhwc2nchw : TransformOp<"nhwc2nchw", []> { + let summary = "Tranform a tensor from nhwc to nchw"; +} +def TensorNorm : TransformOp<"tensornorm", []> { + let summary = "Make continuous tensor align std format in ch direction"; +} + +def Concat : WaferOp<"concat", []> { + let summary = "Concatenation based on the dim"; + + let arguments = (ins + UI64:$src1, // The first input tensor + I32ArrayAttr:$src1_shape, // The first input tensor shape + UI64:$src2, // The second input tensor + I32ArrayAttr:$src2_shape, // The second input tensor shape + I32ArrayAttr:$dst_shape, // Ouput tensor shape + I16Attr:$dim, // Represent the concat direction, such as: + // 0 is channel, 1 is width, and 2 is height + I16Attr:$fmt // The data format of input & output tensor + ); + let results = (outs UI64:$dst); +} + +def Pad : WaferOp<"pad", []> { + let summary = "Tensor padding"; + + let arguments = (ins + UI64:$src, // Input tensor + I32ArrayAttr:$src_shape, // Input tensor shape + I32ArrayAttr:$dst_shape, // Output tensor shape + I16Attr:$pad, // Padding mode: top, bottom, left, and right + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs UI64:$dst); +} + +def Img2col : WaferOp<"img2col", []> { + let summary = "Transform a feature-map tensor into a matrix"; + + let arguments = (ins + UI64:$src, // Input tensor + I32ArrayAttr:$src_shape, // Input tensor shape + I32ArrayAttr:$dst_shape, // Output tensor shape + I32Attr:$src_elem_num, // Number of elements in input tensor + I32Attr:$dst_elem_num, // Number of elements in output tensor + I32ArrayAttr:$swr, // Horizontal stride of convolution + I32ArrayAttr:$pdr, // Vertical stride of convolution + I16Attr:$fmt // The data format of src & dst + ); + let results = (outs UI64:$dst); +} + +def GatherScatter : WaferOp<"gatherscatter", []> { + let summary = "Transfer data in strides and iterations"; + + let arguments = (ins + MemRefOrInt:$source, // The source + MemRefOrInt:$target, // The target + I32Attr:$bytes, // Inner loop data size in bytes + I32Attr:$src_strideN, + I32Attr:$src_strideH, + I32Attr:$src_strideW, + I32Attr:$src_iterN, + I32Attr:$src_iterH, + I32Attr:$src_iterW, + I32Attr:$dst_strideN, + I32Attr:$dst_strideH, + I32Attr:$dst_strideW, + I32Attr:$dst_iterN, + I32Attr:$dst_iterH, + I32Attr:$dst_iterW + ); + let results = (outs I64:$dst); +} + +def BarrierOp : WaferOp<"barrier"> { + let summary = "Synchronizes all work items"; + let description = [{ + The "barrier" op synchronizes all work items. + }]; + let assemblyFormat = "attr-dict"; +} + +def AtomicBarrierInOp : WaferOp<"atomic_barrier_in", [MemoryEffects<[MemRead, MemAlloc]>]> { + let summary = "Synchronizes all work items before input"; + let description = [{ + The "barrier in" op synchronizes all work items. + }]; + let assemblyFormat = "attr-dict"; +} + +def AtomicBarrierOutOp : WaferOp<"atomic_barrier_out", [MemoryEffects<[MemAlloc, MemWrite]>]> { + let summary = "Synchronizes all work items after output"; + let description = [{ + The "barrier out" op synchronizes all work items. + }]; + let assemblyFormat = "attr-dict"; +} + +// ============================================================================= +// 4.22. TsmTile Communication (Recv/Send) +// ============================================================================= + +def RemoteBufferOp : WaferOp<"remote_buffer", [Pure]> { + let summary = "Build a remote buffer descriptor"; + + let description = [{ + Build a lightweight remote buffer descriptor from remote coordinates and a + destination base address. This op is used to carry `remote(buffer)` semantics + through Wafer lowering before it is consumed by remote_store. + }]; + + let arguments = ( + ins + I64:$remote_chip_id_x, + I64:$remote_chip_id_y, + I64:$remote_die_id, + I64:$remote_tile_id, + MemRefOrInt:$dst + ); + + let results = (outs I64:$remote_dst); + + let assemblyFormat = [{ + $remote_chip_id_x `,` $remote_chip_id_y `,` $remote_die_id `,` $remote_tile_id `,` $dst attr-dict `:` type($remote_chip_id_x) `,` type($remote_chip_id_y) `,` type($remote_die_id) `,` type($remote_tile_id) `,` type($dst) + }]; +} + +def RemoteLoadOp : WaferOp<"remote_load", [MemoryEffects<[MemWrite]>]> { + let summary = "Load data from a source tile into a destination buffer"; + + let description = [{ + Receive data from a source tile and write it into the given destination + buffer. + }]; + + let arguments = ( + ins + I64:$remote_chip_id_x, // X-coordinate of the remote chip ID + I64:$remote_chip_id_y, // Y-coordinate of the remote chip ID + I64:$remote_die_id, // ID of the remote die + I64:$remote_tile_id, // ID of the remote tile + MemRefOrInt:$dst, // Destination buffer address (receive into here) + I32:$elem_bytes, // Element size in bytes + I64:$data_size // Total data size in bytes + ); + + let results = (outs); + + let assemblyFormat = [{ + $remote_chip_id_x `,` $remote_chip_id_y `,` $remote_die_id `,` $remote_tile_id `,` $dst `,` $elem_bytes `,` $data_size attr-dict `:` type($remote_chip_id_x) `,` type($remote_chip_id_y) `,` type($remote_die_id) `,` type($remote_tile_id) `,` type($dst) `,` type($elem_bytes) `,` type($data_size) + }]; +} + +def RemoteStoreOp : WaferOp<"remote_store", [MemoryEffects<[MemRead, MemWrite]>]> { + let summary = "Store data to a destination tile"; + + let description = [{ + Store data from the current tile to a destination tile. The operation reads + from the given source buffer and stores it into the remote destination + address. + }]; + + let arguments = ( + ins + I64:$remote_chip_id_x, // X-coordinate of the remote chip ID + I64:$remote_chip_id_y, // Y-coordinate of the remote chip ID + I64:$remote_die_id, // ID of the remote die + I64:$remote_tile_id, // ID of the remote tile + MemRefOrInt:$dst, // Remote destination base address + MemRefOrInt:$src, // Source buffer address (read from here) + I32:$elem_bytes, // Element size in bytes + I64:$data_size // Total data size in bytes + ); + + let results = (outs); + + let assemblyFormat = [{ + $remote_chip_id_x `,` $remote_chip_id_y `,` $remote_die_id `,` $remote_tile_id `,` $dst `,` $src `,` $elem_bytes `,` $data_size attr-dict `:` type($remote_chip_id_x) `,` type($remote_chip_id_y) `,` type($remote_die_id) `,` type($remote_tile_id) `,` type($dst) `,` type($src) `,` type($elem_bytes) `,` type($data_size) + }]; +} + +#endif // WAFER_OPS diff --git a/third_party/wafer/include/wafer/Dialect/IR/WaferTypes.td b/third_party/wafer/include/wafer/Dialect/IR/WaferTypes.td new file mode 100755 index 00000000..b27e42b5 --- /dev/null +++ b/third_party/wafer/include/wafer/Dialect/IR/WaferTypes.td @@ -0,0 +1,107 @@ +//===-------------------------- WaferTypes.td ------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// TODO: Update this file to define the customized type used by Wafer dialect, +// it is now copy-and-pasted from MagicKernelTypes.td. +// +//===----------------------------------------------------------------------===// + +#ifndef WAFER_TYPES_TD +#define WAFER_TYPES_TD + +include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/BuiltinTypeInterfaces.td" +include "wafer/Dialect/IR/WaferDialect.td" + +// +// Types +// +class MKTypeDef traits = []> + : TypeDef { + // Used by printer/parser + let mnemonic = _mnemonic; +} + +// Floating-point Type +def MKFloat : AnyTypeOf<[F8E4M3FN, F8E4M3FNUZ, F8E5M2, F8E5M2FNUZ, F16, BF16, F32, F64], "floating-point">; +def MKFloatTensor : RankedTensorOf<[MKFloat]>; +def MKFloatLike : AnyTypeOf<[MKFloat, MKFloatTensor]>; + +// Boolean Type +// TT_Bool -> I1 +def MKBoolTensor : RankedTensorOf<[I1]>; +def MKBoolLike : AnyTypeOf<[I1, MKBoolTensor]>; + +// Integer Type +def I4 : I<4>; +def MKInt : AnyTypeOf<[I1, I4, I8, I16, I32, I64], "integer">; +def MKIntTensor : RankedTensorOf<[MKInt]>; +def MKIntLike : AnyTypeOf<[MKInt, MKIntTensor]>; + +// I32 Type +// MKI32 -> I32 +// MKI32Tensor -> I32Tensor +def MKI32Like : AnyTypeOf<[I32, I32Tensor]>; + +// I64 Type +// MKI64 -> I64 +// MKI64Tensor -> I64Tensor +def MKI64Like : AnyTypeOf<[I64, I64Tensor]>; + +// Pointer Type in TableGen +class MKPtrOf pointeeTypes> : + DialectType($_self)">, + Concat<"[](::mlir::Type pointeeType) { return ", + SubstLeaves<"$_self", "pointeeType", AnyTypeOf.predicate>, + "; }(::mlir::cast<::mlir::triton::PointerType>($_self).getPointeeType())">]>, + "ptr", "::mlir::triton::PointerType">; + +// Pointer Type in C++ (corresponding to `MKPtrOf`) +def MKPtrType : MKTypeDef<"Pointer", "ptr"> { + let summary = "Pointer type (`::mlir::triton::PointerType`) in Triton IR type system"; + + let description = [{ + Pointer type in Triton IR type system, which could be pointing to scalars or tensors. + }]; + + let parameters = (ins "Type":$pointeeType, "int":$addressSpace); + + let builders = [ + TypeBuilderWithInferredContext<(ins + "Type":$pointeeType, + "int":$addressSpace + ), [{ + return $_get(pointeeType.getContext(), pointeeType, addressSpace); + }]> + ]; + + let hasCustomAssemblyFormat = 1; + + let skipDefaultBuilders = 1; +} + +// Scalar Pointer Type: `ptr<>` +def MKPtr : MKPtrOf<[AnyType]>; + +// Tensor of Pointer Type: `tensor>` +def MKPtrTensor : RankedTensorOf<[MKPtr]>; + +// Tensor of Pointer Type or Pointer type: `tensor>` or `ptr<>` +def MKPtrLike : AnyTypeOf<[MKPtr, MKPtrTensor]>; + +// Tensor Type +def MKFpIntTensor : RankedTensorOf<[MKFloat, MKInt]>; +def MKTensor : RankedTensorOf<[MKFloat, MKInt, MKPtr]>; + +// Pointer Type to Tensor Type: `ptr>` +def MKTensorPtr : MKPtrOf<[MKTensor]>; + +// Any Type in Magic Kernel IR +def MKType : AnyTypeOf<[MKFloatLike, MKIntLike, MKPtrLike, MKTensorPtr]>; + +#endif // WAFER_TYPES_TD diff --git a/third_party/wafer/include/wafer/Transforms/CMakeLists.txt b/third_party/wafer/include/wafer/Transforms/CMakeLists.txt new file mode 100644 index 00000000..eb57629a --- /dev/null +++ b/third_party/wafer/include/wafer/Transforms/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name WaferTransforms) +add_public_tablegen_target(WaferTransformsPassIncGen) diff --git a/third_party/wafer/include/wafer/Transforms/Passes.h b/third_party/wafer/include/wafer/Transforms/Passes.h new file mode 100644 index 00000000..09eadfaf --- /dev/null +++ b/third_party/wafer/include/wafer/Transforms/Passes.h @@ -0,0 +1,10 @@ +#ifndef WAFER_TRANSFORMS_PASSES_H +#define WAFER_TRANSFORMS_PASSES_H +#include "mlir/Pass/Pass.h" +namespace mlir::triton { +std::unique_ptr> createInsertBarrierPass(); +#define GEN_PASS_DECL +#define GEN_PASS_REGISTRATION +#include "wafer/Transforms/Passes.h.inc" +} +#endif diff --git a/third_party/wafer/include/wafer/Transforms/Passes.td b/third_party/wafer/include/wafer/Transforms/Passes.td new file mode 100644 index 00000000..b445b84f --- /dev/null +++ b/third_party/wafer/include/wafer/Transforms/Passes.td @@ -0,0 +1,5 @@ +include "mlir/Pass/PassBase.td" +def InsertBarrier : Pass<"wafer-insert-barrier", "mlir::ModuleOp"> { + let summary = "Insert Wafer SPM/DDR producer-consumer barriers"; + let constructor = "mlir::triton::createInsertBarrierPass()"; +} diff --git a/third_party/wafer/language/cpu/__init__.py b/third_party/wafer/language/cpu/__init__.py new file mode 100755 index 00000000..229b57d8 --- /dev/null +++ b/third_party/wafer/language/cpu/__init__.py @@ -0,0 +1,3 @@ +from . import libdevice + +__all__ = ["libdevice"] diff --git a/third_party/wafer/language/cpu/libdevice.py b/third_party/wafer/language/cpu/libdevice.py new file mode 100755 index 00000000..f448d802 --- /dev/null +++ b/third_party/wafer/language/cpu/libdevice.py @@ -0,0 +1,1496 @@ +from triton.language import core +from triton.language.math import * +from triton.language.math import _check_dtype + + +@core.extern +def clz(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("int32"), ): ("__nv_clz", core.dtype("int32")), + (core.dtype("int64"), ): ("__nv_clzll", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def popc(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("int32"), ): ("__nv_popc", core.dtype("int32")), + (core.dtype("int64"), ): ("__nv_popcll", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def byte_perm(arg0, arg1, arg2, _builder=None): + return core.extern_elementwise("", "", [arg0, arg1, arg2], { + (core.dtype("int32"), core.dtype("int32"), core.dtype("int32")): ("__nv_byte_perm", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def mulhi(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_mulhi", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32")): ("__nv_umulhi", core.dtype("uint32")), + (core.dtype("int64"), core.dtype("int64")): ("__nv_mul64hi", core.dtype("int64")), + (core.dtype("uint64"), core.dtype("uint64")): ("__nv_umul64hi", core.dtype("uint64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def mul24(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_mul24", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32")): ("__nv_umul24", core.dtype("uint32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def brev(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("int32"), ): ("__nv_brev", core.dtype("int32")), + (core.dtype("int64"), ): ("__nv_brevll", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sad(arg0, arg1, arg2, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("int32"), core.dtype("int32"), core.dtype("uint32")): ("__nv_sad", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32"), core.dtype("uint32")): ("__nv_usad", core.dtype("uint32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rcp64h(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_rcp64h", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def trunc(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_trunc", core.dtype("fp64")), + (core.dtype("fp32"), ): ("__nv_truncf", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def saturatef(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_saturatef", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fma_rn(arg0, arg1, arg2, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmaf_rn", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_fma_rn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fma_rz(arg0, arg1, arg2, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmaf_rz", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_fma_rz", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fma_rd(arg0, arg1, arg2, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmaf_rd", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_fma_rd", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fma_ru(arg0, arg1, arg2, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmaf_ru", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_fma_ru", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_dividef(arg0, arg1, _builder=None): + return core.extern_elementwise("", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fast_fdividef", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def div_rz(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fdiv_rz", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_ddiv_rz", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def div_rd(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fdiv_rd", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_ddiv_rd", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def div_ru(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fdiv_ru", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_ddiv_ru", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rcp_rn(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_frcp_rn", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_drcp_rn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rcp_rz(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_frcp_rz", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_drcp_rz", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rcp_rd(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_frcp_rd", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_drcp_rd", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rcp_ru(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_frcp_ru", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_drcp_ru", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sqrt_rz(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fsqrt_rz", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_dsqrt_rz", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sqrt_rd(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fsqrt_rd", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_dsqrt_rd", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sqrt_ru(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fsqrt_ru", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_dsqrt_ru", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def add_rn(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dadd_rn", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fadd_rn", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def add_rz(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dadd_rz", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fadd_rz", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def add_rd(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dadd_rd", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fadd_rd", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def add_ru(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dadd_ru", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fadd_ru", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def mul_rn(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dmul_rn", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmul_rn", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def mul_rz(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dmul_rz", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmul_rz", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def mul_rd(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dmul_rd", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmul_rd", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def mul_ru(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [ + arg0, + arg1, + ], { + ( + core.dtype("fp64"), + core.dtype("fp64"), + ): ("__nv_dmul_ru", core.dtype("fp64")), + ( + core.dtype("fp32"), + core.dtype("fp32"), + ): ("__nv_fmul_ru", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2float_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2float_rn", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2float_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2float_rz", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2float_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2float_rd", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2float_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2float_ru", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2int_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2int_rn", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2int_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2int_rz", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2int_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2int_rd", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2int_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2int_ru", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2uint_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2uint_rn", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2uint_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2uint_rz", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2uint_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2uint_rd", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2uint_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2uint_ru", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def int2double_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2double_rn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def uint2double_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2double_rn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2int_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2int_rn", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2int_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2int_rz", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2int_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2int_rd", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2int_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2int_ru", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2uint_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2uint_rn", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2uint_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2uint_rz", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2uint_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2uint_rd", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2uint_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2uint_ru", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def int2float_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2float_rn", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def int2float_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2float_rz", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def int2float_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2float_rd", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def int2float_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2float_ru", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def uint2float_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2float_rn", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def uint2float_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2float_rz", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def uint2float_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2float_rd", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def uint2float_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2float_ru", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def hiloint2double(arg0, arg1, _builder=None): + return core.extern_elementwise("", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_hiloint2double", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2loint(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2loint", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2hiint(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2hiint", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2ll_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ll_rn", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2ll_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ll_rz", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2ll_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ll_rd", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2ll_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ll_ru", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2ull_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ull_rn", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2ull_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ull_rz", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2ull_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ull_rd", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float2ull_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ull_ru", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2ll_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ll_rn", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2ll_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ll_rz", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2ll_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ll_rd", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2ll_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ll_ru", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2ull_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ull_rn", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2ull_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ull_rz", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2ull_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ull_rd", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double2ull_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ull_ru", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ll2float_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2float_rn", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ll2float_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2float_rz", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ll2float_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2float_rd", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ll2float_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2float_ru", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ull2float_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2float_rn", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ull2float_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2float_rz", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ull2float_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2float_rd", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ull2float_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2float_ru", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ll2double_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2double_rn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ll2double_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2double_rz", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ll2double_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2double_rd", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ll2double_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2double_ru", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ull2double_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2double_rn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ull2double_rz(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2double_rz", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ull2double_rd(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2double_rd", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ull2double_ru(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2double_ru", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def int_as_float(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int_as_float", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float_as_int(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float_as_int", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def uint_as_float(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint_as_float", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def float_as_uint(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float_as_uint", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def longlong_as_double(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_longlong_as_double", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def double_as_longlong(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double_as_longlong", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_sinf(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_sinf", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_cosf(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_cosf", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_logf(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_logf", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_expf(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_expf", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_tanf(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_tanf", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_exp10f(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_exp10f", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_log10f(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_log10f", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fast_powf(arg0, arg1, _builder=None): + return core.extern_elementwise("", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fast_powf", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def hadd(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_hadd", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32")): ("__nv_uhadd", core.dtype("uint32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rhadd(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_rhadd", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32")): ("__nv_urhadd", core.dtype("uint32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sub_rn(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fsub_rn", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dsub_rn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sub_rz(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fsub_rz", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dsub_rz", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sub_rd(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fsub_rd", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dsub_rd", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sub_ru(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fsub_ru", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dsub_ru", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rsqrt_rn(arg0, _builder=None): + return core.extern_elementwise("", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_frsqrt_rn", core.dtype("fp32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ffs(arg0, _builder=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("int32"), ): ("__nv_ffs", core.dtype("int32")), + (core.dtype("int64"), ): ("__nv_ffsll", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rint(arg0, _builder=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_rintf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_rint", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def llrint(arg0, _builder=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_llrintf", core.dtype("int64")), + (core.dtype("fp64"), ): ("__nv_llrint", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def nearbyint(arg0, _builder=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_nearbyintf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_nearbyint", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def isnan(arg0, _builder=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_isnanf", core.dtype("int1")), + (core.dtype("fp64"), ): ("__nv_isnand", core.dtype("int1")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def signbit(arg0, _builder=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_signbitf", core.dtype("int32")), + (core.dtype("fp64"), ): ("__nv_signbitd", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def copysign(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_copysignf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_copysign", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def finitef(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_finitef", core.dtype("int1")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def isinf(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_isinff", core.dtype("int1")), + (core.dtype("fp64"), ): ("__nv_isinfd", core.dtype("int1")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def nextafter(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_nextafterf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_nextafter", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sinpi(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_sinpif", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_sinpi", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def cospi(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_cospif", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cospi", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def tan(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_tanf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_tan", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def exp10(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_exp10f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_exp10", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def cosh(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_coshf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cosh", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def sinh(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_sinhf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_sinh", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def tanh(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_tanhf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_tanh", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def atan2(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_atan2f", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_atan2", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def atan(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_atanf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_atan", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def asin(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_asinf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_asin", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def acos(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_acosf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_acos", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def log10(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_log10f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_log10", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def log1p(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_log1pf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_log1p", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def acosh(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_acoshf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_acosh", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def asinh(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_asinhf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_asinh", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def atanh(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_atanhf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_atanh", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def expm1(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_expm1f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_expm1", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def hypot(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_hypotf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_hypot", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rhypot(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_rhypotf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_rhypot", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def norm3d(arg0, arg1, arg2, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_norm3df", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_norm3d", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rnorm3d(arg0, arg1, arg2, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_rnorm3df", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_rnorm3d", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def norm4d(arg0, arg1, arg2, arg3, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2, arg3], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): + ("__nv_norm4df", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): + ("__nv_norm4d", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rnorm4d(arg0, arg1, arg2, arg3, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2, arg3], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): + ("__nv_rnorm4df", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): + ("__nv_rnorm4d", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def cbrt(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_cbrtf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cbrt", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def rcbrt(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_rcbrtf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_rcbrt", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def j0(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_j0f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_j0", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def j1(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_j1f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_j1", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def y0(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_y0f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_y0", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def y1(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_y1f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_y1", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def yn(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("fp32")): ("__nv_ynf", core.dtype("fp32")), + (core.dtype("int32"), core.dtype("fp64")): ("__nv_yn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def jn(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("fp32")): ("__nv_jnf", core.dtype("fp32")), + (core.dtype("int32"), core.dtype("fp64")): ("__nv_jn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def cyl_bessel_i0(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_cyl_bessel_i0f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cyl_bessel_i0", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def cyl_bessel_i1(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_cyl_bessel_i1f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cyl_bessel_i1", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +# Rewrite math.erf to support fp16/bf16 +@core.builtin +@_check_dtype(dtypes=["fp16", "fp32", "bf16"]) +@core._tensor_member_fn +def erf(x, _builder=None): + x = semantic.to_tensor(x, _builder) + return core.tensor(_builder.create_erf(x.handle), x.type) + + +@core.extern +def erfinv(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_erfinvf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_erfinv", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def erfc(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_erfcf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_erfc", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def erfcx(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_erfcxf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_erfcx", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def erfcinv(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_erfcinvf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_erfcinv", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def normcdfinv(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_normcdfinvf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_normcdfinv", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def normcdf(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_normcdff", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_normcdf", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def lgamma(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_lgammaf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_lgamma", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ldexp(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("int32")): ("__nv_ldexpf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("int32")): ("__nv_ldexp", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def scalbn(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("int32")): ("__nv_scalbnf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("int32")): ("__nv_scalbn", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fmod(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmodf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_fmod", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def remainder(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_remainderf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_remainder", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def pow(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("int32")): ("__nv_powif", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("int32")): ("__nv_powi", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_powf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_pow", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def tgamma(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_tgammaf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_tgamma", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def round(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_roundf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_round", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def llround(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_llroundf", core.dtype("int64")), + (core.dtype("fp64"), ): ("__nv_llround", core.dtype("int64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def fdim(arg0, arg1, _builder=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fdimf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_fdim", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def ilogb(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_ilogbf", core.dtype("int32")), + (core.dtype("fp64"), ): ("__nv_ilogb", core.dtype("int32")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def logb(arg0, _builder=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_logbf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_logb", core.dtype("fp64")), + }, is_pure=True, _builder=_builder) + + +@core.extern +def isfinited(arg0, _builder=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_isfinited", core.dtype("int32")), + }, is_pure=True, _builder=_builder).to(core.int1, _builder=_builder) diff --git a/third_party/wafer/language/txda/__init__.py b/third_party/wafer/language/txda/__init__.py new file mode 100644 index 00000000..85414620 --- /dev/null +++ b/third_party/wafer/language/txda/__init__.py @@ -0,0 +1,4 @@ +"""Compatibility import for the Wafer language extension.""" +from ..wafer import libdevice + +__all__ = ["libdevice"] diff --git a/third_party/wafer/language/txda/libdevice.py b/third_party/wafer/language/txda/libdevice.py new file mode 100644 index 00000000..73974c73 --- /dev/null +++ b/third_party/wafer/language/txda/libdevice.py @@ -0,0 +1,2 @@ +"""Compatibility module; new code imports triton.language.extra.wafer.""" +from ..wafer.libdevice import * # noqa: F401,F403 diff --git a/third_party/wafer/language/wafer/__init__.py b/third_party/wafer/language/wafer/__init__.py new file mode 100755 index 00000000..229b57d8 --- /dev/null +++ b/third_party/wafer/language/wafer/__init__.py @@ -0,0 +1,3 @@ +from . import libdevice + +__all__ = ["libdevice"] diff --git a/third_party/wafer/language/wafer/libdevice.py b/third_party/wafer/language/wafer/libdevice.py new file mode 100755 index 00000000..0506995d --- /dev/null +++ b/third_party/wafer/language/wafer/libdevice.py @@ -0,0 +1,1497 @@ +from triton.language import core +from triton.language.math import * + + +def _float_unary(arg, fp32_symbol, fp64_symbol, semantic, predicate=False): + # The device math ABI takes FP32/FP64. Promote low-precision inputs before + # extern dispatch, then restore their dtype (predicates always return bool). + arg = semantic.to_tensor(arg) + dtype = arg.dtype + if dtype in (core.float16, core.bfloat16): + arg = semantic.cast(arg, core.float32) + result = core.extern_elementwise("", "", [arg], { + (core.float32,): (fp32_symbol, core.int1 if predicate else core.float32), + (core.float64,): (fp64_symbol, core.int1 if predicate else core.float64), + }, is_pure=True, _semantic=semantic) + return result if predicate else semantic.cast(result, dtype) + + +@core.extern +def clz(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("int32"), ): ("__nv_clz", core.dtype("int32")), + (core.dtype("int64"), ): ("__nv_clzll", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def popc(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("int32"), ): ("__nv_popc", core.dtype("int32")), + (core.dtype("int64"), ): ("__nv_popcll", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def byte_perm(arg0, arg1, arg2, _semantic=None): + return core.extern_elementwise("", "", [arg0, arg1, arg2], { + (core.dtype("int32"), core.dtype("int32"), core.dtype("int32")): ("__nv_byte_perm", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def mulhi(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_mulhi", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32")): ("__nv_umulhi", core.dtype("uint32")), + (core.dtype("int64"), core.dtype("int64")): ("__nv_mul64hi", core.dtype("int64")), + (core.dtype("uint64"), core.dtype("uint64")): ("__nv_umul64hi", core.dtype("uint64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def mul24(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_mul24", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32")): ("__nv_umul24", core.dtype("uint32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def brev(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("int32"), ): ("__nv_brev", core.dtype("int32")), + (core.dtype("int64"), ): ("__nv_brevll", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sad(arg0, arg1, arg2, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("int32"), core.dtype("int32"), core.dtype("uint32")): ("__nv_sad", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32"), core.dtype("uint32")): ("__nv_usad", core.dtype("uint32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rcp64h(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_rcp64h", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def trunc(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_trunc", core.dtype("fp64")), + (core.dtype("fp32"), ): ("__nv_truncf", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def saturatef(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_saturatef", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fma_rn(arg0, arg1, arg2, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmaf_rn", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_fma_rn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fma_rz(arg0, arg1, arg2, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmaf_rz", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_fma_rz", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fma_rd(arg0, arg1, arg2, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmaf_rd", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_fma_rd", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fma_ru(arg0, arg1, arg2, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmaf_ru", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_fma_ru", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_dividef(arg0, arg1, _semantic=None): + return core.extern_elementwise("", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fast_fdividef", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def div_rz(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fdiv_rz", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_ddiv_rz", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def div_rd(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fdiv_rd", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_ddiv_rd", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def div_ru(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fdiv_ru", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_ddiv_ru", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rcp_rn(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_frcp_rn", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_drcp_rn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rcp_rz(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_frcp_rz", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_drcp_rz", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rcp_rd(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_frcp_rd", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_drcp_rd", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rcp_ru(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_frcp_ru", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_drcp_ru", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sqrt_rz(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fsqrt_rz", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_dsqrt_rz", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sqrt_rd(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fsqrt_rd", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_dsqrt_rd", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sqrt_ru(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fsqrt_ru", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_dsqrt_ru", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def add_rn(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dadd_rn", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fadd_rn", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def add_rz(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dadd_rz", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fadd_rz", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def add_rd(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dadd_rd", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fadd_rd", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def add_ru(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dadd_ru", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fadd_ru", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def mul_rn(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dmul_rn", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmul_rn", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def mul_rz(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dmul_rz", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmul_rz", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def mul_rd(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dmul_rd", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmul_rd", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def mul_ru(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [ + arg0, + arg1, + ], { + ( + core.dtype("fp64"), + core.dtype("fp64"), + ): ("__nv_dmul_ru", core.dtype("fp64")), + ( + core.dtype("fp32"), + core.dtype("fp32"), + ): ("__nv_fmul_ru", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2float_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2float_rn", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2float_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2float_rz", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2float_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2float_rd", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2float_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2float_ru", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2int_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2int_rn", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2int_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2int_rz", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2int_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2int_rd", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2int_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2int_ru", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2uint_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2uint_rn", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2uint_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2uint_rz", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2uint_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2uint_rd", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2uint_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2uint_ru", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def int2double_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2double_rn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def uint2double_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2double_rn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2int_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2int_rn", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2int_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2int_rz", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2int_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2int_rd", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2int_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2int_ru", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2uint_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2uint_rn", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2uint_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2uint_rz", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2uint_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2uint_rd", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2uint_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2uint_ru", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def int2float_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2float_rn", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def int2float_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2float_rz", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def int2float_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2float_rd", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def int2float_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int2float_ru", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def uint2float_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2float_rn", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def uint2float_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2float_rz", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def uint2float_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2float_rd", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def uint2float_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint2float_ru", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def hiloint2double(arg0, arg1, _semantic=None): + return core.extern_elementwise("", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_hiloint2double", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2loint(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2loint", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2hiint(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2hiint", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2ll_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ll_rn", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2ll_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ll_rz", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2ll_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ll_rd", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2ll_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ll_ru", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2ull_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ull_rn", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2ull_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ull_rz", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2ull_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ull_rd", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float2ull_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float2ull_ru", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2ll_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ll_rn", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2ll_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ll_rz", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2ll_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ll_rd", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2ll_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ll_ru", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2ull_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ull_rn", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2ull_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ull_rz", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2ull_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ull_rd", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double2ull_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double2ull_ru", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ll2float_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2float_rn", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ll2float_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2float_rz", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ll2float_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2float_rd", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ll2float_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2float_ru", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ull2float_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2float_rn", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ull2float_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2float_rz", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ull2float_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2float_rd", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ull2float_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2float_ru", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ll2double_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2double_rn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ll2double_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2double_rz", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ll2double_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2double_rd", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ll2double_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_ll2double_ru", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ull2double_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2double_rn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ull2double_rz(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2double_rz", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ull2double_rd(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2double_rd", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ull2double_ru(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint64"), ): ("__nv_ull2double_ru", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def int_as_float(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int32"), ): ("__nv_int_as_float", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float_as_int(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float_as_int", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def uint_as_float(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("uint32"), ): ("__nv_uint_as_float", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def float_as_uint(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_float_as_uint", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def longlong_as_double(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("int64"), ): ("__nv_longlong_as_double", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def double_as_longlong(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_double_as_longlong", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_sinf(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_sinf", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_cosf(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_cosf", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_logf(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_logf", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_expf(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_expf", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_tanf(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_tanf", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_exp10f(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_exp10f", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_log10f(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_fast_log10f", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fast_powf(arg0, arg1, _semantic=None): + return core.extern_elementwise("", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fast_powf", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def hadd(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_hadd", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32")): ("__nv_uhadd", core.dtype("uint32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rhadd(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("int32")): ("__nv_rhadd", core.dtype("int32")), + (core.dtype("uint32"), core.dtype("uint32")): ("__nv_urhadd", core.dtype("uint32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sub_rn(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fsub_rn", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dsub_rn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sub_rz(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fsub_rz", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dsub_rz", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sub_rd(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fsub_rd", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dsub_rd", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sub_ru(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fsub_ru", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_dsub_ru", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rsqrt_rn(arg0, _semantic=None): + return core.extern_elementwise("", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_frsqrt_rn", core.dtype("fp32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ffs(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("int32"), ): ("__nv_ffs", core.dtype("int32")), + (core.dtype("int64"), ): ("__nv_ffsll", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rint(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_rintf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_rint", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def llrint(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_llrintf", core.dtype("int64")), + (core.dtype("fp64"), ): ("__nv_llrint", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def nearbyint(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_nearbyintf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_nearbyint", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def isnan(arg0, _semantic=None): + return _float_unary(arg0, "__nv_isnanf", "__nv_isnand", _semantic, predicate=True) + + +@core.extern +def signbit(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [ + arg0, + ], { + (core.dtype("fp32"), ): ("__nv_signbitf", core.dtype("int32")), + (core.dtype("fp64"), ): ("__nv_signbitd", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def copysign(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_copysignf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_copysign", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def finitef(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_finitef", core.dtype("int1")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def isinf(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_isinff", core.dtype("int1")), + (core.dtype("fp64"), ): ("__nv_isinfd", core.dtype("int1")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def nextafter(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_nextafterf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_nextafter", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sinpi(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_sinpif", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_sinpi", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def cospi(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_cospif", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cospi", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def tan(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_tanf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_tan", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def exp10(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_exp10f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_exp10", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def cosh(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_coshf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cosh", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def sinh(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_sinhf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_sinh", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def tanh(arg0, _semantic=None): + arg0 = _semantic.to_tensor(arg0) + result = _float_unary(arg0, "__nv_tanhf", "__nv_tanh", _semantic) + # Wafer loses -0 and returns NaN for +inf. Preserve zeros and exact limits + # with device comparisons/selects while keeping its finite vector path. + one = core.full((), 1, arg0.dtype, _semantic=_semantic) + result = core.where(arg0.__eq__(float("inf"), _semantic=_semantic), one, result, _semantic=_semantic) + result = core.where(arg0.__eq__(-float("inf"), _semantic=_semantic), + one.__neg__(_semantic=_semantic), result, _semantic=_semantic) + return core.where(arg0.__eq__(0, _semantic=_semantic), arg0, result, _semantic=_semantic) + + +@core.extern +def atan2(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_atan2f", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_atan2", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def atan(arg0, _semantic=None): + return _float_unary(arg0, "__nv_atanf", "__nv_atan", _semantic) + + +@core.extern +def asin(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_asinf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_asin", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def acos(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_acosf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_acos", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def log10(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_log10f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_log10", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def log1p(arg0, _semantic=None): + return _float_unary(arg0, "__nv_log1pf", "__nv_log1p", _semantic) + + +@core.extern +def acosh(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_acoshf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_acosh", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def asinh(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_asinhf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_asinh", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def atanh(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_atanhf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_atanh", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def expm1(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_expm1f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_expm1", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def hypot(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_hypotf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_hypot", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rhypot(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_rhypotf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_rhypot", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def norm3d(arg0, arg1, arg2, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_norm3df", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_norm3d", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rnorm3d(arg0, arg1, arg2, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): ("__nv_rnorm3df", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): ("__nv_rnorm3d", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def norm4d(arg0, arg1, arg2, arg3, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2, arg3], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): + ("__nv_norm4df", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): + ("__nv_norm4d", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rnorm4d(arg0, arg1, arg2, arg3, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1, arg2, arg3], { + (core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32"), core.dtype("fp32")): + ("__nv_rnorm4df", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64"), core.dtype("fp64")): + ("__nv_rnorm4d", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def cbrt(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_cbrtf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cbrt", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def rcbrt(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_rcbrtf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_rcbrt", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def j0(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_j0f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_j0", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def j1(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_j1f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_j1", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def y0(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_y0f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_y0", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def y1(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_y1f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_y1", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def yn(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("fp32")): ("__nv_ynf", core.dtype("fp32")), + (core.dtype("int32"), core.dtype("fp64")): ("__nv_yn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def jn(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("int32"), core.dtype("fp32")): ("__nv_jnf", core.dtype("fp32")), + (core.dtype("int32"), core.dtype("fp64")): ("__nv_jn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def cyl_bessel_i0(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_cyl_bessel_i0f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cyl_bessel_i0", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def cyl_bessel_i1(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_cyl_bessel_i1f", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_cyl_bessel_i1", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def erf(x, _semantic=None): + # Triton 3.5 semantic converts values; IR operations live on its builder. + x = core.to_tensor(x, _semantic=_semantic) + return core.tensor(_semantic.builder.create_erf(x.handle), x.type) + + +@core.extern +def erfinv(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_erfinvf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_erfinv", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def erfc(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_erfcf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_erfc", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def erfcx(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_erfcxf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_erfcx", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def erfcinv(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_erfcinvf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_erfcinv", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def normcdfinv(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_normcdfinvf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_normcdfinv", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def normcdf(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_normcdff", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_normcdf", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def lgamma(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_lgammaf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_lgamma", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ldexp(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("int32")): ("__nv_ldexpf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("int32")): ("__nv_ldexp", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def scalbn(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("int32")): ("__nv_scalbnf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("int32")): ("__nv_scalbn", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fmod(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fmodf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_fmod", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def remainder(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_remainderf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_remainder", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def pow(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("int32")): ("__nv_powif", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("int32")): ("__nv_powi", core.dtype("fp64")), + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_powf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_pow", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def tgamma(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_tgammaf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_tgamma", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def round(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_roundf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_round", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def llround(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_llroundf", core.dtype("int64")), + (core.dtype("fp64"), ): ("__nv_llround", core.dtype("int64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def fdim(arg0, arg1, _semantic=None): + return core.extern_elementwise( + "", "", [arg0, arg1], { + (core.dtype("fp32"), core.dtype("fp32")): ("__nv_fdimf", core.dtype("fp32")), + (core.dtype("fp64"), core.dtype("fp64")): ("__nv_fdim", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def ilogb(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_ilogbf", core.dtype("int32")), + (core.dtype("fp64"), ): ("__nv_ilogb", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def logb(arg0, _semantic=None): + return core.extern_elementwise( + "", "", [arg0], { + (core.dtype("fp32"), ): ("__nv_logbf", core.dtype("fp32")), + (core.dtype("fp64"), ): ("__nv_logb", core.dtype("fp64")), + }, is_pure=True, _semantic=_semantic) + + +@core.extern +def isfinited(arg0, _semantic=None): + return core.extern_elementwise("", "", [arg0], { + (core.dtype("fp64"), ): ("__nv_isfinited", core.dtype("int32")), + }, is_pure=True, _semantic=_semantic).to(core.int1, _semantic=_semantic) diff --git a/third_party/wafer/language/wafer/slicing.py b/third_party/wafer/language/wafer/slicing.py new file mode 100644 index 00000000..7c591319 --- /dev/null +++ b/third_party/wafer/language/wafer/slicing.py @@ -0,0 +1,42 @@ +"""Opt-in static bounded tensor slices using Wafer's TLE slice operation.""" + +import builtins + +import triton.language.core as tl +from triton.experimental.tle.language.dsa import extract_slice + + +_upstream_getitem = tl.tensor.__getitem__ + + +@tl._tensor_member_fn +@tl.builtin +def __getitem__(self, slices, _semantic=None): + if isinstance(slices, tl.tuple): + indices = list(slices.values) + elif isinstance(slices, (builtins.tuple, list)): + indices = list(slices) + else: + indices = [slices] + slice_types = (builtins.slice, tl.slice) + bounded = any(isinstance(s, slice_types) and any( + tl._unwrap_if_constexpr(v) is not None for v in (s.start, s.stop, s.step) + ) for s in indices) + if not bounded: + # Preserve upstream full-slice and new-axis behavior. + return _upstream_getitem(self, slices, _semantic=_semantic) + if len(indices) > len(self.shape) or not all(isinstance(s, slice_types) for s in indices): + raise ValueError("Wafer bounded indexing accepts slices only; use expand_dims separately") + indices += [builtins.slice(None)] * (len(self.shape) - len(indices)) + offsets, sizes, strides = [], [], [] + for index, dim in zip(indices, self.type.shape): + values = [tl._unwrap_if_constexpr(v) for v in (index.start, index.stop, index.step)] + start, stop, step = [default if v is None else v for v, default in zip(values, (0, dim, 1))] + if not all(isinstance(v, int) for v in (start, stop, step)): + raise ValueError("Wafer slice bounds must be compile-time integers") + if not 0 <= start < stop <= dim or step <= 0: + raise ValueError("Wafer slices require 0 <= start < stop <= dimension and positive stride") + offsets.append(start) + sizes.append((stop - start + step - 1) // step) + strides.append(step) + return extract_slice(self, offsets, sizes, strides, _semantic=_semantic) diff --git a/third_party/wafer/lib/Analysis/Alias.cpp b/third_party/wafer/lib/Analysis/Alias.cpp new file mode 100755 index 00000000..cf2603fd --- /dev/null +++ b/third_party/wafer/lib/Analysis/Alias.cpp @@ -0,0 +1,79 @@ +#include "Analysis/Alias.h" +#include "Address/Dialect/IR/AddressDialect.h" +#include "Analysis/Utility.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/BuiltinTypes.h" + +namespace mlir::triton::alias { + +AliasInfo AliasInfo::join(const AliasInfo &lhs, const AliasInfo &rhs) { + if (lhs == rhs) + return lhs; + AliasInfo ret; + for (auto value : lhs.allocs) { + ret.insert(value); + } + for (auto value : rhs.allocs) { + ret.insert(value); + } + return ret; +} + +LogicalResult SharedMemoryAliasAnalysis::visitOperation( + Operation *op, ArrayRef *> operands, + ArrayRef *> results) { + AliasInfo aliasInfo; + bool pessimistic = true; + auto result = op->getResult(0); + // skip ops that return memdesc in a different memory space. + // TODO: Check if the memory space is shared memory + if (isa(op)) + return success(); + + // Only LocalAllocOp creates a new buffer. + if (isa(op)) { + aliasInfo.insert(result); + pessimistic = false; + } else if (isa(op)) { + // FIXME: memref::viewOp and memref::transposeOp + // FIXME: A common trait or interface to handle all memref op + // FIXME: memref::SubViewOp may need handled before this analysis. + aliasInfo = AliasInfo(operands[0]->getValue()); + pessimistic = false; + } else { + if (isa(result.getType())) { + op->dump(); + fflush(stdout); + } + assert(!isa(result.getType()) && + "unknown operation creating memory descriptor"); + } + + if (pessimistic) { + setAllToEntryStates(results); + return success(); + } + // Join all lattice elements + for (auto *result : results) + propagateIfChanged(result, result->join(aliasInfo)); + + return success(); +} + +AliasResult SharedMemoryAliasAnalysis::alias(Value lhs, Value rhs) { + // TODO: implement + return AliasResult::MayAlias; +} + +ModRefResult SharedMemoryAliasAnalysis::getModRef(Operation *op, + Value location) { + // TODO: implement + return ModRefResult::getModAndRef(); +} + +} // namespace mlir::triton::alias diff --git a/third_party/wafer/lib/Analysis/Allocation.cpp b/third_party/wafer/lib/Analysis/Allocation.cpp new file mode 100755 index 00000000..1f29f036 --- /dev/null +++ b/third_party/wafer/lib/Analysis/Allocation.cpp @@ -0,0 +1,756 @@ +#include "Analysis/Allocation.h" +#include "Analysis/Alias.h" +#include "mlir/Analysis/Liveness.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include +#include + +#define DEBUG_TYPE "allocation-shared-memory" +#define DBGS() (llvm::dbgs() << "[" DEBUG_TYPE "]: ") +#define LDBG(X) LLVM_DEBUG(DBGS() << X << "\n") + +//===----------------------------------------------------------------------===// +// Shared Memory Allocation Analysis +//===----------------------------------------------------------------------===// +namespace mlir::triton::alloc { + +class AllocationAnalysis { +public: + enum class BufferAccessMode { READ, WRITE, READ_WRITE, UnSupported }; + +public: + AllocationAnalysis(Operation *operation, + Allocation::FuncAllocMapT *funcAllocMap, + Allocation *allocation) + : operation(operation), funcAllocMap(funcAllocMap), + allocation(allocation) { + run(); + } + +private: + using BufferT = Allocation::BufferT; + + /// Value -> Liveness Range + /// Use MapVector to ensure determinism. + using BufferRangeMapT = llvm::MapVector>; + /// Nodes -> Nodes + using GraphT = DenseMap>; + + void run() { + getValuesAndSizes(); + resolveLiveness(); + computeOffsets(); + } + + /// Initializes explicitly defined shared memory values for a given operation. + void getExplicitValueSize(Operation *op) { + // FIXME: Support memory hierarchy (Multi-memory allocation. eg: Shared + // memory && scratch memory) + auto alloc = dyn_cast(op); + if (!alloc) + return; + // Bytes could be a different value once we support padding or other + // allocation policies. + auto allocType = alloc.getType(); + // FIXME: padding or other alignment + auto bitWidth = allocType.getElementTypeBitWidth(); + auto elemByte = (bitWidth + 7) / 8; + int64_t bytes = allocType.getNumElements() * elemByte; + + auto alignment = alloc.getAlignment().value_or(1); + + // WORKAROUND: Reduce op output will write to alignment 256 bytes + // FIXME: Handle tensors that require more memory than their shape suggests, + // for example, due to padding or alignment requirements. + bytes = (bytes + alignment - 1) / alignment * alignment; + + allocation->addBuffer(alloc, bytes, + alignment, 0); + } + + void getValueAlias(Value value, + triton::alias::SharedMemoryAliasAnalysis &analysis) { + dataflow::Lattice *latticeElement = + analysis.getLatticeElement(value); + if (!latticeElement) + return; + + triton::alias::AliasInfo &info = latticeElement->getValue(); + if (info.getAllocs().empty()) { + LLVM_DEBUG({ + llvm::dbgs() << "\tNo allocs found for value: "; + value.dump(); + }); + return; + } + + for (auto alloc : info.getAllocs()) { + // FIXME: Why this happens? DPS? + if (value == alloc) + continue; + LLVM_DEBUG({ + llvm::dbgs() << "\tAdd alias value: "; + value.dump(); + llvm::dbgs() << "\t to alloc: "; + alloc.dump(); + }); + allocation->addAlias(value, alloc); + } + } + + /// Extract all shared memory values and their sizes + void getValuesAndSizes() { + // Get the alloc values + operation->walk( + [&](Operation *op) { getExplicitValueSize(op); }); + + LDBG("\nGet buffer and size --"); + for (auto valueBufferIter : allocation->valueBuffer) { + auto *buffer = valueBufferIter.second; + LLVM_DEBUG(llvm::dbgs() + << "-- buffer " << buffer->id << " size: " << buffer->size + << " offset: " << buffer->offset << "\n\t"; + buffer->owner->dump();); + } + LLVM_DEBUG({ llvm::dbgs() << "\n\n"; }); + + // Get the alias values + std::unique_ptr solver = createDataFlowSolver(); + triton::alias::SharedMemoryAliasAnalysis *aliasAnalysis = + solver->load(); + // Run the analysis rooted at every isolated from above operation, including + // the top-level function but also any nested regions. + operation->walk([&](Operation *op) { + if (op->hasTrait() && + failed(solver->initializeAndRun(op))) { + // TODO: return error instead of bailing out.. + llvm_unreachable("failed to run SharedMemoryAliasAnalysis"); + } + }); + + LDBG("\n==== Value alias ============="); + operation->walk([&](Operation *op) { + LLVM_DEBUG({ + llvm::dbgs() << "\nValue Alias for op: "; + if (!op->hasTrait()) + op->dump(); + else + op->getName(); + }); + + for (auto operand : op->getOperands()) { + getValueAlias(operand, *aliasAnalysis); + } + for (auto value : op->getResults()) { + getValueAlias(value, *aliasAnalysis); + } + }); + LLVM_DEBUG({ llvm::dbgs() << "\n\n"; }); + } + + /// Traverse the use-def chain to find out the earliest memory access pattern + /// operation + + /// Check whether we need to update the memory access pattern + bool needUpdateMemAccessPattern(Operation *last, Operation *current, + DenseMap &operationId) { + assert(current && "current op is null"); + if (!last) + return true; + + auto lastOpParent = last->getParentOp(); + auto currentOpParent = current->getParentOp(); + assert(lastOpParent && currentOpParent && "parent op is null"); + // Same region, compare operation id + if (lastOpParent == currentOpParent) + return operationId[last] < operationId[current]; + + // Sub-region, always update + if (lastOpParent->isProperAncestor(currentOpParent)) + return true; + + return false; + } + + /// Set the access pattern according the live operation + void setAccessPattern( + DenseMap> &BufferAccessMap, + DenseMap &operationId, BufferT *buffer, + Operation *liveOp, BufferAccessMode mode) { + auto minAccessIDOp = BufferAccessMap[buffer][static_cast(mode)]; + + BufferAccessMap[buffer][static_cast(mode)] = + needUpdateMemAccessPattern(minAccessIDOp, liveOp, operationId) + ? liveOp + : minAccessIDOp; + } + + bool isPartialWrite(Value operand, BufferT *buffer) { + auto memrefType = cast(operand.getType()); + assert(isa(buffer->owner)); + auto bufferType = cast(buffer->owner->getResultTypes().front()); + return memrefType.getShape() != bufferType.getShape(); + } + + /// Handle MemoryEffectOpInterface access pattern + void memoryEffectOpInterfaceAccessPattern( + DenseMap> &BufferAccessMap, + DenseMap &operationId, BufferT *buffer, + Operation *liveOp, OpOperand &opOperand) { + auto memEffectOp = cast(liveOp); + SmallVector effects; + memEffectOp.getEffects(effects); + + auto operand = opOperand.get(); + for (const auto &effect : effects) { + if (effect.getValue() != operand) + continue; + if (isa(effect.getEffect())) { + setAccessPattern(BufferAccessMap, operationId, buffer, liveOp, + BufferAccessMode::READ); + LLVM_DEBUG(llvm::dbgs() + << " -> MemoryEffectOpInterface READ access for buffer " + << buffer->id << "\n"); + } else if (isa(effect.getEffect())) { + if (isPartialWrite(operand, buffer)) + continue; + + setAccessPattern(BufferAccessMap, operationId, buffer, liveOp, + BufferAccessMode::WRITE); + LLVM_DEBUG(llvm::dbgs() + << " -> MemoryEffectOpInterface WRITE access for buffer " + << buffer->id << "\n"); + } else { + assert(false && "unknown memory effect"); + } + } + } + + /// Resolve the buffer access pattern for all buffers + DenseMap> + resolveBufferAccessPattern(Liveness &liveness, + DenseMap &operationId) { + // Each access pattern will store the earliest operation that performs the + // access. + DenseMap> + InLoopBufferAccessPattern; + for (auto [K, V] : allocation->valueBuffer) { + if (V->owner == findOutmostParentLoopOp(V->owner)) + continue; + InLoopBufferAccessPattern[V] = llvm::SmallVector( + static_cast(BufferAccessMode::UnSupported), nullptr); + } + + // TODO: More complicated access pattern analysis + auto getMemoryAccessPattern = [&](Value value, BufferT *buffer) { + LLVM_DEBUG(llvm::dbgs() << "\nAnalyzing memory access pattern for buffer " + << buffer->id << " buffer: "; + buffer->owner->dump(); llvm::dbgs() << "\n";); + + auto liveOperations = liveness.resolveLiveness(value); + std::for_each( + liveOperations.begin(), liveOperations.end(), [&](Operation *liveOp) { + for (auto &opOperand : liveOp->getOpOperands()) { + auto operand = opOperand.get(); + auto bufferIds = allocation->getBufferIds(operand); + if (bufferIds.empty()) + continue; + + // scf::if may has multiple buffers associated with one + // operand + if (!bufferIds.contains(buffer->id)) + continue; + + LLVM_DEBUG(llvm::dbgs() << "Live Operation: \n"; liveOp->dump(); + llvm::dbgs() << " Operand: \n\t"; operand.dump(); + llvm::dbgs() << "\n"; fflush(stderr);); + + if (liveOp->mightHaveTrait()) { + // FIXME: Handle terminators properly. scf::if + InLoopBufferAccessPattern[buffer] = + llvm::SmallVector( + static_cast(BufferAccessMode::UnSupported), + nullptr); + + LLVM_DEBUG(llvm::dbgs() << " -> terminator for buffer " + << buffer->id << "\n"); + } else if (isa(liveOp)) { + // Memory effect ops + LLVM_DEBUG(llvm::dbgs() << " -> MemoryEffectOpInterface " + "buffer access pattern!\n";); + memoryEffectOpInterfaceAccessPattern(InLoopBufferAccessPattern, + operationId, buffer, + liveOp, opOperand); + + continue; + } else { + LLVM_DEBUG(llvm::dbgs() + << " -> unknown buffer access pattern!\n";); + assert(allocation->valueBuffer.contains(operand)); + assert(false && "unknown buffer access pattern"); + } + } + }); + }; + + // Compute access pattern for explicitly defined buffers + for (auto valueBufferIter : allocation->valueBuffer) { + auto value = valueBufferIter.first; + auto *buffer = valueBufferIter.second; + if (!InLoopBufferAccessPattern.contains(buffer)) + continue; + getMemoryAccessPattern(value, buffer); + } + + // Compute access pattern for alias buffers + for (const auto &[value, buffers] : allocation->aliasBuffer) { + for (auto *buffer : buffers) { + if (!InLoopBufferAccessPattern.contains(buffer)) + continue; + getMemoryAccessPattern(value, buffer); + } + } + return InLoopBufferAccessPattern; + } + + void updateToLoopLiveness(BufferT *buffer, + DenseMap &operationId) { + auto parentOp = findOutmostParentLoopOp(buffer->owner); + assert(bufferRange.count(buffer)); + + assert(parentOp->hasTrait()); + auto &entryBlock = parentOp->getRegion(0).getBlocks().front(); + auto firstOp = &entryBlock.front(); + + auto minId = std::min(operationId[firstOp], bufferRange[buffer].start()); + auto maxId = std::max(operationId[parentOp] + 1, bufferRange[buffer].end()); + bufferRange[buffer] = Interval(minId, maxId); + } + + void updateInLoopBufferLiveness( + DenseMap> &BufferAccessMap, + DenseMap &operationId) { + + LLVM_DEBUG( + llvm::dbgs() << "\n====== InLoopBuffer Access Pattern : ==========\n"; + for (auto [buffer, access] : BufferAccessMap) { + auto minWOp = access[static_cast(BufferAccessMode::WRITE)]; + auto minROp = access[static_cast(BufferAccessMode::READ)]; + auto minRWOp = + access[static_cast(BufferAccessMode::READ_WRITE)]; + llvm::dbgs() << "Buffer " << buffer->id << " "; + buffer->owner->dump(); + llvm::dbgs() << " READ="; + minROp == nullptr + ? llvm::dbgs() << std::numeric_limits::max() << "\n" + : llvm::dbgs() << operationId[minROp] << "\n"; + llvm::dbgs() << ", WRITE="; + minWOp == nullptr + ? llvm::dbgs() << std::numeric_limits::max() << "\n" + : llvm::dbgs() << operationId[minWOp] << "\n"; + llvm::dbgs() << ", READ_WRITE="; + minRWOp == nullptr + ? llvm::dbgs() << std::numeric_limits::max() << "\n" + : llvm::dbgs() << operationId[minRWOp] << "\n"; + llvm::dbgs() << ", \n\n"; + fflush(stderr); + };); + + for (auto [buffer, access] : BufferAccessMap) { + auto minWOp = access[static_cast(BufferAccessMode::WRITE)]; + auto minROp = access[static_cast(BufferAccessMode::READ)]; + auto minRWOp = access[static_cast(BufferAccessMode::READ_WRITE)]; + + assert(minRWOp == nullptr && + "READ_WRITE access pattern is not supported yet"); + if (!minWOp || (minROp && operationId[minROp] < operationId[minWOp])) + updateToLoopLiveness(buffer, operationId); + } + } + + /// Computes the liveness range of the allocated value. + /// Each buffer is allocated only once. + void resolveExplicitBufferLiveness( + function_ref(Value value, BufferT *buffer)> + getLiveness) { + for (auto valueBufferIter : allocation->valueBuffer) { + auto value = valueBufferIter.first; + auto *buffer = valueBufferIter.second; + bufferRange[buffer] = getLiveness(value, buffer); + LLVM_DEBUG({ + llvm::dbgs() << "-- buffer " << buffer->id << "; value: "; + value.dump(); + }); + } + } + + /// Extends the liveness range by unionizing the liveness range of the aliased + /// values because each allocated buffer could be an alias of others, if block + /// arguments are involved. + void resolveAliasBufferLiveness( + function_ref(Value value, BufferT *buffer)> + getLiveness) { + for (const auto &[value, buffers] : allocation->aliasBuffer) { + auto range = getLiveness(value, buffers.front()); + for (auto *buffer : buffers) { + auto minId = range.start(); + auto maxId = range.end(); + if (bufferRange.count(buffer)) { + // Extend the allocated buffer's range + minId = std::min(minId, bufferRange[buffer].start()); + maxId = std::max(maxId, bufferRange[buffer].end()); + } + bufferRange[buffer] = Interval(minId, maxId); + } + } + } + + Operation *findOutmostParentLoopOp(Operation *op) { + if (!op) { + return nullptr; + } + auto parentOp = op->getParentOp(); + if (!parentOp) { + return op; + } + if (!isa(parentOp) && !isa(parentOp) && + !isa(parentOp)) { + return op; + } + return findOutmostParentLoopOp(parentOp); + } + + DenseMap idToOperation; + + /// Resolves liveness of all values involved under the root operation. + void resolveLiveness() { + // Assign an ID to each operation using post-order traversal. + // To achieve the correct liveness range, the parent operation's ID + // should be greater than each of its child operation's ID . + // Example: + // ... + // %5 = triton.convert_layout %4 + // %6 = scf.for ... iter_args(%arg0 = %0) -> (i32) { + // %2 = triton.convert_layout %5 + // ... + // scf.yield %arg0 + // } + // For example, %5 is defined in the parent region and used in + // the child region, and is not passed as a block argument. + // %6 should should have an ID greater than its child operations, + // otherwise %5 liveness range ends before the child operation's liveness + // range ends. + DenseMap operationId; + LLVM_DEBUG( + llvm::dbgs() + << "\n=== Assigning operation IDs using post-order traversal ===\n"); + operation->walk([&](Operation *op) { + LLVM_DEBUG(llvm::dbgs() << "Assigning ID " << operationId.size() + << " to operation: "; + if (!op->hasTrait()) op->dump(); + else op->getName();); + operationId[op] = operationId.size(); + }); + LLVM_DEBUG(llvm::dbgs() << "\n\n"); + + for (auto [K, V] : operationId) + idToOperation[V] = K; + + // Analyze liveness of explicit buffers + Liveness liveness(operation); + auto getValueLivenessRange = [&](Value value, BufferT *buffer) { + auto liveOperations = liveness.resolveLiveness(value); + // TODO: Support async + // Update regions for buffer. + + auto minId = std::numeric_limits::max(); + auto maxId = std::numeric_limits::min(); + std::for_each(liveOperations.begin(), liveOperations.end(), + [&](Operation *liveOp) { + minId = std::min(minId, operationId[liveOp]); + // FIXME: Optimize. Since buffer in loop will always has + // same address, so assumed they have the same liveness + // range with the parent loop operation. + auto parentOp = liveOp; + maxId = std::max(maxId, operationId[parentOp] + 1); + }); + return Interval(minId, maxId); + }; + + resolveExplicitBufferLiveness(getValueLivenessRange); + resolveAliasBufferLiveness(getValueLivenessRange); + + // Process in-loop buffer access pattern to extend liveness range + auto BufferAccessMap = resolveBufferAccessPattern(liveness, operationId); + updateInLoopBufferLiveness(BufferAccessMap, operationId); + } + + void dumpBuffers() { + LDBG("\nDump bufferRange: id size offset ---------"); + for (auto bufferIter : bufferRange) { + LLVM_DEBUG({ + bufferIter.first->owner->dump(); + llvm::dbgs() << "-- " << bufferIter.first->id << " " + << bufferIter.first->size << " " + << bufferIter.first->offset << " " + << "interval " << bufferIter.second.start() << " " + << bufferIter.second.end() << "\n"; + llvm::dbgs() << "\t start: "; + idToOperation.at(bufferIter.second.start())->dump(); + llvm::dbgs() << "\t end: "; + idToOperation.at(bufferIter.second.end())->dump(); + }); + } + llvm::dbgs() << "\n\n"; + } + + void dumpAllocationSize() const { + LDBG("\nDump shared memory allocation size -----------"); + auto liveBuffers = allocation->getLiveBuffers(); + auto analyzedSize = 0; + for (auto [op, bufferIds] : liveBuffers) { + auto size = 0; + for (auto bufferId : bufferIds) { + auto bufferSize = allocation->getAllocatedSize(bufferId); + size += bufferSize; + } + analyzedSize = std::max(analyzedSize, size); + } + llvm::dbgs() << "Allocated: " << allocation->sharedMemorySize + << ", analyzed: " << analyzedSize << "\n"; + llvm::dbgs() << "\n\n"; + } + + void dumpInterferenceGraph(const GraphT &interference) const { + LDBG("\nDump interference graph: \n"); + for (auto edges : interference) { + llvm::dbgs() << "-- from " << edges.first->id << " to "; + for (auto node : edges.second) { + llvm::dbgs() << node->id << "; "; + } + llvm::dbgs() << "\n"; + } + llvm::dbgs() << "\n\n"; + } + + /// Computes the shared memory offsets for all related values. + /// Paper: Algorithms for Compile-Time Memory Optimization + /// (https://dl.acm.org/doi/pdf/10.5555/314500.315082) + void computeOffsets() { + SmallVector buffers; + for (auto bufferIter : bufferRange) { + buffers.emplace_back(bufferIter.first); + } + + // Sort buffers by size in descending order to reduce the fragmentation + // on big buffers caused by smaller buffers. Big buffers have a higher + // chance to overlap with multiple other buffers, and allocating them first + // (by calculateStarts) ensures a higher chance that they will occupy a + // standalone smem slot. + std::sort(buffers.begin(), buffers.end(), + [&](BufferT *A, BufferT *B) { return A->size > B->size; }); + + calculateStarts(buffers); + LLVM_DEBUG(dumpBuffers()); + + // NOTE: The original paper doesn't consider interference between + // the bumped ranges. Buffers that previously do not interfere with + // could interfere after offset bumping if their liveness ranges overlap. + // Therefore, we rerun the interference graph algorithm after bumping so + // that we regroup the buffers and color them again. Since we always + // increase the buffer offset and keep reducing conflicts, we will + // eventually reach a fixed point. + GraphT interference; + buildInterferenceGraph(buffers, interference); + do { + allocate(buffers, interference); + buildInterferenceGraph(buffers, interference); + } while (!interference.empty()); + + LLVM_DEBUG(dumpAllocationSize()); + // TODO: What is sharingGroup + // Update allocation for sharingGroup. + LLVM_DEBUG(dumpBuffers()); + } + + /// Computes the initial shared memory offsets. + void calculateStarts(const SmallVector &buffers) { + // v = values in shared memory + // t = triplet of (size, start, end) + // shared memory space + // - + // | *******t4 + // | /|\ v2 inserts t4, t5, and t6 + // | | + // | ******t5 ************t6 + // | ^^^^^v2^^^^^^ + // | | *********************t2 + // | \|/ v2 erases t1 + // | ******t1 ^^^^^^^^^v1^^^^^^^^^ ************t3 + // |---------------------------------------------| liveness range + // 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 ... + // If the available triple's range is less than a given buffer range, + // we won't know if there has been an overlap without using graph coloring. + // Start -> Liveness Range + using TripleMapT = std::multimap>; + TripleMapT tripleMap; + tripleMap.insert(std::make_pair(0, Interval())); + SmallVector xBuffers = buffers; + while (!xBuffers.empty()) { + auto tripleIt = tripleMap.begin(); + auto offset = tripleIt->first; + auto range = tripleIt->second; + tripleMap.erase(tripleIt); + auto bufferIt = + std::find_if(xBuffers.begin(), xBuffers.end(), [&](auto *buffer) { + auto xRange = bufferRange[buffer]; + bool res = xRange.intersects(range); + for (const auto &val : tripleMap) + res = res && + !val.second.intersects(xRange); // only one buffer intersect + return res; + }); + if (bufferIt != xBuffers.end()) { + auto buffer = *bufferIt; + auto xSize = buffer->size; + auto xRange = bufferRange.lookup(buffer); + // TODO(Keren): A buffer's size shouldn't be determined here, have to + // clean it up + size_t alignOffset = buffer->setOffsetAligned(offset); + tripleMap.insert({alignOffset + xSize, + Interval{std::max(range.start(), xRange.start()), + std::min(range.end(), xRange.end())}}); + // We could either insert (range.start, xRange.start) or (range.start, + // xRange.end), both are correct and determine the potential buffer + // offset, and the graph coloring algorithm will solve the interference, + // if any + if (range.start() < xRange.start()) + tripleMap.insert({offset, Interval{range.start(), xRange.end()}}); + if (xRange.end() < range.end()) + tripleMap.insert({offset, Interval{xRange.start(), range.end()}}); + xBuffers.erase(bufferIt); + } + } + } + + /// Builds a graph of all shared memory values. Edges are created between + /// shared memory values that are overlapping. + void buildInterferenceGraph(const SmallVector &buffers, + GraphT &interference) { + // Reset interference graph + interference.clear(); + for (auto x : buffers) { + for (auto y : buffers) { + if (x == y) + continue; + auto xStart = x->offset; + auto yStart = y->offset; + auto xSize = x->size; + auto ySize = y->size; + Interval xSizeRange = {xStart, xStart + xSize}; + Interval ySizeRange = {yStart, yStart + ySize}; + auto xOpRange = bufferRange.lookup(x); + auto yOpRange = bufferRange.lookup(y); + + // Buffers interfere if their allocation offsets overlap and they are + // live at the same time. + if (xOpRange.intersects(yOpRange) && + xSizeRange.intersects(ySizeRange)) { + interference[x].insert(y); + } + + // TODO: Async + // Buffers also interfere if their allocation offsets overlap and they + // exist within regions that may execute simultaneously with respect to + // each other. + // if x and y belong to different regions (ignore producer region). + } + } + + LLVM_DEBUG(dumpInterferenceGraph(interference)); + } + + /// Finalizes shared memory offsets considering interference. + void allocate(const SmallVector &buffers, + const GraphT &interference) { + LDBG("\n------------ graph coloring ------------"); + + // Reset shared memory size + allocation->sharedMemorySize = 0; + // First-fit graph coloring + // Neighbors are nodes that interfere with each other. + // We color a node by finding the index of the first available + // non-neighboring node or the first neighboring node without any color. + // Nodes with the same color do not interfere with each other. + DenseMap colors; + for (auto value : buffers) { + colors[value] = (value == buffers[0]) ? 0 : -1; + } + SmallVector available(buffers.size()); + for (auto x : buffers) { + std::fill(available.begin(), available.end(), true); + for (auto y : interference.lookup(x)) { + int color = colors[y]; + if (color >= 0) { + available[color] = false; + } + } + auto it = std::find(available.begin(), available.end(), true); + colors[x] = std::distance(available.begin(), it); + LLVM_DEBUG({ + llvm::dbgs() << "-- color " << x->id << " " << colors[x] << "\n"; + }); + } + LLVM_DEBUG({ llvm::dbgs() << "\n\n"; }); + + // Finalize allocation + // color0: [0, 7), [0, 8), [0, 15) -> [0, 7), [0, 8), [0, 15) + // color1: [7, 9) -> [0 + 1 * 15, 9 + 1 * 15) -> [15, 24) + // color2: [8, 12) -> [8 + 2 * 15, 12 + 2 * 15) -> [38, 42) + // TODO(Keren): We are wasting memory here. + // Nodes with color2 can actually start with 24. + for (auto x : buffers) { + size_t newOffset = 0; + for (auto y : interference.lookup(x)) { + newOffset = std::max(newOffset, y->offset + y->size); + } + if (colors.lookup(x) != 0) + x->setOffsetAligned(newOffset); + allocation->sharedMemorySize = + std::max(allocation->sharedMemorySize, x->offset + x->size); + } + LLVM_DEBUG(dumpBuffers()); + } + +private: + Operation *operation; + Allocation::FuncAllocMapT *funcAllocMap; + Allocation *allocation; + BufferRangeMapT bufferRange; +}; + +void Allocation::run(FuncAllocMapT &funcAllocMap) { + triton::alloc::AllocationAnalysis(getOperation(), &funcAllocMap, this); +} + +std::map> +Allocation::getLiveBuffers() { + std::map> liveBuffers; + + Operation *rootOperation = getOperation(); + Liveness liveness(rootOperation); + auto analyzeOperation = [&](Operation *op) -> void { + for (auto result : op->getOpResults()) { + auto bufferId = getBufferId(result); + if (bufferId == Allocation::InvalidBufferId) + continue; + auto liveOperations = liveness.resolveLiveness(result); + for (auto depOp : liveOperations) + liveBuffers[depOp].push_back(bufferId); + } + }; + rootOperation->walk(analyzeOperation); + return liveBuffers; +} + +} // namespace mlir::triton::alloc diff --git a/third_party/wafer/lib/Analysis/CMakeLists.txt b/third_party/wafer/lib/Analysis/CMakeLists.txt new file mode 100755 index 00000000..7dd83206 --- /dev/null +++ b/third_party/wafer/lib/Analysis/CMakeLists.txt @@ -0,0 +1,13 @@ +add_triton_library(ZTCAnalysis + Allocation.cpp + Alias.cpp + Membar.cpp + + DEPENDS + TritonTableGen + WaferTableGen + TritonGPUAttrDefsIncGen + + LINK_LIBS PUBLIC + MLIRAnalysis +) diff --git a/third_party/wafer/lib/Analysis/Membar.cpp b/third_party/wafer/lib/Analysis/Membar.cpp new file mode 100644 index 00000000..f4ef24e9 --- /dev/null +++ b/third_party/wafer/lib/Analysis/Membar.cpp @@ -0,0 +1,411 @@ +#include "Analysis/Membar.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/Interfaces/CallInterfaces.h" +#include "mlir/Interfaces/ControlFlowInterfaces.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "wafer/Dialect/IR/WaferOps.h" + +#include +#include +#include + +namespace mlir::triton::membar { +using namespace mlir; + +static bool isIntersectedMap(const BlockInfo::IntervalMapT &lhsIntervalSet, + const BlockInfo::IntervalMapT &rhsIntervalSet, + MembarFilterFn filter, MembarHazardKind kind) { + for (auto &lhs : lhsIntervalSet) + for (auto &rhs : rhsIntervalSet) + if (lhs.first.intersects(rhs.first)) + for (auto lhsOp : lhs.second) + for (auto rhsOp : rhs.second) + if (!filter || !filter(lhsOp, rhsOp, kind)) + return true; + return false; +} + +bool BlockInfo::isIntersected(const BlockInfo &other, + MembarFilterFn filter) const { + return isIntersectedMap(syncWriteIntervals, other.syncReadIntervals, filter, + MembarHazardKind::WriteRead) || + isIntersectedMap(syncReadIntervals, other.syncWriteIntervals, filter, + MembarHazardKind::ReadWrite) || + isIntersectedMap(syncWriteIntervals, other.syncWriteIntervals, filter, + MembarHazardKind::WriteWrite); +} + +// ----------------------------------------------------------------------------- +// Helpers: resolve i64/index address computations back to a memref Value +// (avoid touching Alias.cpp; keep local). +// ----------------------------------------------------------------------------- + +static Value getBaseBuffer(Value v) { + while (auto *defOp = v.getDefiningOp()) { + if (auto op = dyn_cast(defOp)) + v = op.getSource(); + else if (auto op = dyn_cast(defOp)) + v = op.getSource(); + else if (auto op = dyn_cast(defOp)) + v = op.getSource(); + else if (auto op = dyn_cast(defOp)) + v = op.getSource(); + else if (auto op = dyn_cast(defOp)) { + if (v == op->getResult(0)) + v = op.getSource(); + else + break; + } else + break; + } + return v; +} + +static Value traceToOriginMemRef(Value v, unsigned maxDepth = 16) { + if (maxDepth == 0) + return {}; + if (isa(v.getType())) + return getBaseBuffer(v); + Operation *defOp = v.getDefiningOp(); + if (!defOp) + return {}; + + if (auto op = dyn_cast(defOp)) + return traceToOriginMemRef(op.getIn(), maxDepth - 1); + if (auto op = dyn_cast(defOp)) + return traceToOriginMemRef(op.getIn(), maxDepth - 1); + if (auto op = dyn_cast(defOp)) + return traceToOriginMemRef(op.getIn(), maxDepth - 1); + if (auto op = dyn_cast(defOp)) + return traceToOriginMemRef(op.getIn(), maxDepth - 1); + + if (auto op = dyn_cast(defOp)) + return getBaseBuffer(op.getSource()); + + if (isa(defOp)) { + if (Value r = traceToOriginMemRef(defOp->getOperand(0), maxDepth - 1)) + return r; + return traceToOriginMemRef(defOp->getOperand(1), maxDepth - 1); + } + return {}; +} + +static Value resolveForBufferLookup(Value v) { + if (!v) + return {}; + if (isa(v.getType())) + return getBaseBuffer(v); + if (Value origin = traceToOriginMemRef(v)) + return origin; + return {}; +} + +static bool getAllocationOffsetInterval(Value v, + BlockInfo::IntervalT &interval) { + Value base = resolveForBufferLookup(v); + if (!base) + return false; + + auto alloc = base.getDefiningOp(); + if (!alloc) + return false; + + auto offsetAttr = alloc->getAttrOfType("allocation.offset"); + if (!offsetAttr) + return false; + + int64_t signedOffset = offsetAttr.getInt(); + if (signedOffset < 0) + return false; + + MemRefType allocType = alloc.getType(); + if (!allocType.hasStaticShape()) + return false; + + int64_t numElements = allocType.getNumElements(); + unsigned bitWidth = allocType.getElementTypeBitWidth(); + uint64_t elemBytes = (bitWidth + 7) / 8; + if (numElements < 0 || elemBytes == 0) + return false; + if (static_cast(numElements) > + std::numeric_limits::max() / elemBytes) + return false; + + uint64_t bytes = static_cast(numElements) * elemBytes; + uint64_t offset = static_cast(signedOffset); + uint64_t maxSize = std::numeric_limits::max(); + if (offset > maxSize || bytes > maxSize - offset) + return false; + + interval = BlockInfo::IntervalT(static_cast(offset), + static_cast(offset + bytes)); + return true; +} + +static bool isTxDialect(Operation *op) { + auto *d = op->getDialect(); + return d && d->getNamespace() == "wafer"; +} + +bool isPureAddressOp(Operation *op) { + if (!op || isTxDialect(op)) + return false; + // Real memref data movement / visibility — must participate in hazards. + if (isa(op)) + return false; + + auto *dialect = op->getDialect(); + if (!dialect) + return false; + StringRef ns = dialect->getNamespace(); + // Addressing, layout, and control flow only (no SPM/DDR data dependence). + if (ns == "arith" || ns == "memref" || ns == "scf" || ns == "cf" || + ns == "builtin" || ns == "affine" || ns == "index") + return true; + return false; +} + +static void collectWaferAccesses(Operation *op, SmallVector &reads, + SmallVector &writes) { + auto hasAccess = [&](Value v) { + for (Value read : reads) + if (read == v) + return true; + for (Value write : writes) + if (write == v) + return true; + return false; + }; + + // Prefer MemoryEffectOpInterface if present. + if (auto iface = dyn_cast(op)) { + SmallVector> effects; + iface.getEffects(effects); + for (auto &e : effects) { + Value v = e.getValue(); + if (!v) + continue; + if (isa(e.getEffect())) + reads.push_back(v); + else if (isa(e.getEffect())) + writes.push_back(v); + } + // Many Wafer ops annotate destination operands with MemWrite but leave + // source address operands unannotated. Treat remaining address-like + // operands as reads. + for (Value v : op->getOperands()) + if (!hasAccess(v) && resolveForBufferLookup(v)) + reads.push_back(v); + return; + } + + for (Value v : op->getOperands()) + reads.push_back(v); +} + +static void collectCpuAccesses(Operation *op, SmallVector &reads, + SmallVector &writes) { + if (auto load = dyn_cast(op)) { + reads.push_back(load.getMemRef()); + return; + } + if (auto store = dyn_cast(op)) { + writes.push_back(store.getMemRef()); + return; + } + if (auto copy = dyn_cast(op)) { + reads.push_back(copy.getSource()); + writes.push_back(copy.getTarget()); + return; + } + + // CPU / unknown: operands may be real reads of shared buffers. + for (Value v : op->getOperands()) + reads.push_back(v); +} + +void MembarAnalysis::run(FuncBlockInfoMapT &funcBlockInfoMap) { + FunctionOpInterface funcOp = + dyn_cast(allocation->getOperation()); + OpBuilder builder(funcOp.getContext()); + resolve(funcOp, &funcBlockInfoMap, &builder); +} + +void MembarAnalysis::resolve(FunctionOpInterface funcOp, + FuncBlockInfoMapT *funcBlockInfoMap, + OpBuilder *builder) { + DenseMap inputBlockInfoMap; + DenseMap outputBlockInfoMap; + std::deque blockList; + + funcOp.walk([&](Block *block) { + if (block->isEntryBlock() && + !isa(block->getParentOp())) + blockList.emplace_back(block, Block::iterator()); + }); + + while (!blockList.empty()) { + VirtualBlock block = blockList.front(); + blockList.pop_front(); + auto inputBlockInfo = inputBlockInfoMap[block]; + SmallVector successors; + Block::iterator startIt = + block.second.isValid() ? std::next(block.second) : block.first->begin(); + + for (Operation &op : llvm::make_range(startIt, block.first->end())) { + if (op.hasTrait() || + isa(op)) { + visitTerminator(&op, successors); + break; + } + update(&op, &inputBlockInfo, funcBlockInfoMap, builder); + } + + if (outputBlockInfoMap.count(block) && + inputBlockInfo == outputBlockInfoMap[block]) + continue; + + outputBlockInfoMap[block] = inputBlockInfo; + for (VirtualBlock successor : successors) { + inputBlockInfoMap[successor].join(outputBlockInfoMap[block]); + blockList.emplace_back(successor); + } + } + + // Join dangling buffers at return sites. + BlockInfo &funcBlockInfo = (*funcBlockInfoMap)[funcOp]; + funcOp.walk([&](Operation *retLike) { + if (!retLike->hasTrait()) + return; + SmallVector> virtualBlocks; + for (auto &[vb, blockInfo] : outputBlockInfoMap) + if (vb.first == retLike->getBlock()) + virtualBlocks.emplace_back(vb, blockInfo); + if (virtualBlocks.empty()) + return; + auto maxIt = llvm::max_element(virtualBlocks, [&](auto &lhs, auto &rhs) { + Block::iterator lhsIt = lhs.first.second, rhsIt = rhs.first.second; + return !lhsIt.isValid() || + (rhsIt.isValid() && lhsIt->isBeforeInBlock(&*rhsIt)); + }); + funcBlockInfo.join(maxIt->second); + }); +} + +void MembarAnalysis::visitTerminator(Operation *op, + SmallVector &successors) { + if (isa(op)) { + for (Block *successor : op->getSuccessors()) + successors.emplace_back(successor, Block::iterator()); + return; + } + + if (auto br = dyn_cast(op)) { + SmallVector regions; + br.getSuccessorRegions(RegionBranchPoint::parent(), regions); + for (RegionSuccessor ®ion : regions) { + if (region.isParent()) + successors.emplace_back(br->getBlock(), br->getIterator()); + else + successors.emplace_back(®ion.getSuccessor()->front(), + Block::iterator()); + } + return; + } + + auto br = dyn_cast(op); + if (br && isa(br->getParentOp())) { + SmallVector operands(br->getNumOperands()); + SmallVector regions; + br.getSuccessorRegions(operands, regions); + for (RegionSuccessor ®ion : regions) { + if (region.isParent()) { + Operation *parent = br->getParentOp(); + successors.emplace_back(parent->getBlock(), parent->getIterator()); + } else { + successors.emplace_back(®ion.getSuccessor()->front(), + Block::iterator()); + } + } + return; + } + + if (op->hasTrait()) + return; + llvm_unreachable("Unknown terminator encountered in Wafer membar analysis"); +} + +void MembarAnalysis::insertBarrier(Operation *op, OpBuilder *builder) { + OpBuilder::InsertionGuard g(*builder); + builder->create(op->getLoc()); +} + +void MembarAnalysis::update(Operation *op, BlockInfo *blockInfo, + FuncBlockInfoMapT *funcBlockInfoMap, + OpBuilder *builder) { + if (isa(op)) { + blockInfo->sync(); + return; + } + + BlockInfo curBlockInfo; + + // Inter-procedural: treat calls as their callee's block info. + if (isa(op)) { + auto callOpInterface = dyn_cast(op); + if (auto callee = + dyn_cast(callOpInterface.resolveCallable())) + curBlockInfo = funcBlockInfoMap->lookup(callee); + } else { + SmallVector reads, writes; + + if (isTxDialect(op)) { + collectWaferAccesses(op, reads, writes); + } else if (isPureAddressOp(op)) { + // memref/arith/scf scaffolding between tx ops — not a CPU data touch. + } else { + collectCpuAccesses(op, reads, writes); + } + + auto addIntervals = [&](ArrayRef vals, bool isWrite) { + for (Value v : vals) { + Value lookup = resolveForBufferLookup(v); + if (lookup) { + for (auto bufferId : allocation->getBufferIds(lookup)) { + if (bufferId == triton::alloc::Allocation::InvalidBufferId) + continue; + auto interval = allocation->getAllocatedInterval(bufferId); + if (isWrite) + curBlockInfo.syncWriteIntervals[interval].insert(op); + else + curBlockInfo.syncReadIntervals[interval].insert(op); + } + } + + BlockInfo::IntervalT physicalInterval; + if (getAllocationOffsetInterval(v, physicalInterval)) { + if (isWrite) + curBlockInfo.syncWriteIntervals[physicalInterval].insert(op); + else + curBlockInfo.syncReadIntervals[physicalInterval].insert(op); + } + } + }; + + addIntervals(reads, /*isWrite=*/false); + addIntervals(writes, /*isWrite=*/true); + } + + if (blockInfo->isIntersected(curBlockInfo, filter)) { + builder->setInsertionPoint(op); + insertBarrier(op, builder); + blockInfo->sync(); + } + blockInfo->join(curBlockInfo); +} + +} // namespace mlir::triton::membar diff --git a/third_party/wafer/lib/CMakeLists.txt b/third_party/wafer/lib/CMakeLists.txt new file mode 100755 index 00000000..36c2b915 --- /dev/null +++ b/third_party/wafer/lib/CMakeLists.txt @@ -0,0 +1,9 @@ +# Common (UnifiedHardware) is already built by FlagTree as target "Common" +# when wafer is loaded as a plugin via TRITON_PLUGIN_DIRS. +if(NOT WAFER_USE_EXTERNAL_TLE) + add_subdirectory(Common) +endif() +add_subdirectory(Analysis) +add_subdirectory(Conversion) +add_subdirectory(Dialect) +add_subdirectory(Registrar) diff --git a/third_party/wafer/lib/Common/CMakeLists.txt b/third_party/wafer/lib/Common/CMakeLists.txt new file mode 100755 index 00000000..61a9f7dd --- /dev/null +++ b/third_party/wafer/lib/Common/CMakeLists.txt @@ -0,0 +1 @@ +add_triton_library(Common UnifiedHardware.cc) diff --git a/third_party/wafer/lib/Common/UnifiedHardware.cc b/third_party/wafer/lib/Common/UnifiedHardware.cc new file mode 100755 index 00000000..1953ef5a --- /dev/null +++ b/third_party/wafer/lib/Common/UnifiedHardware.cc @@ -0,0 +1,32 @@ +#include "flagtree/Common/UnifiedHardware.h" +#include +namespace mlir { +namespace flagtree { + +bool UnifiedHardware::isRegistered() const { +#ifdef FLAGTREE_BACKEND + return true; +#else + return false; +#endif +} + +int UnifiedHardware::getDMATag() const { return 0; } + +int UnifiedHardware::getSharedMemoryTag() const { return 0; } + +bool UnifiedHardware::getIncubatedTag() const { return false; } + +std::string UnifiedHardware::getReduceStrategy() const { + return "linalg_reduce"; +} + +std::string UnifiedHardware::getFlagTreeBackend() const { return "default"; } + +__attribute__((weak)) std::unique_ptr +createUnifiedHardwareManager() { + return std::make_unique(); +} + +} // namespace flagtree +} // namespace mlir diff --git a/third_party/wafer/lib/Conversion/AllocateSharedMemory/AllocateSharedMemoryPass.cpp b/third_party/wafer/lib/Conversion/AllocateSharedMemory/AllocateSharedMemoryPass.cpp new file mode 100755 index 00000000..fb3c401c --- /dev/null +++ b/third_party/wafer/lib/Conversion/AllocateSharedMemory/AllocateSharedMemoryPass.cpp @@ -0,0 +1,135 @@ +//===------------------- AllocateSharedMemoryPass.cpp ---------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "Analysis/Allocation.h" +#include "Analysis/Utility.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "triton/Analysis/Allocation.h" +#include "wafer/Conversion/AllocateSharedMemory/Passes.h" + +#define DEBUG_TYPE "allocate-shared-memory" + +using namespace mlir; + +namespace mlir::triton::alloc { +#define GEN_PASS_DEF_ALLOCATESHAREDMEMORY +#include "wafer/Conversion/AllocateSharedMemory/Passes.h.inc" + +} // namespace mlir::triton::alloc + +namespace { +struct AllocateSharedMemory + : public mlir::triton::alloc::impl::AllocateSharedMemoryBase< + AllocateSharedMemory> { + + Operation *findAlignmentRestrictOpOperandBuffer(Value v) { + auto op = v.getDefiningOp(); + assert(op && "Value has no defining op"); + if (isa(op)) + return op; + // Memref op which has result: ViewLikeOpInterface. Eg: + // memref::ExpandShapeOp + assert(isa(op)); + return findAlignmentRestrictOpOperandBuffer(op->getOperand(0)); + } + + void handleTargetDependentAlignmentRequirements(ModuleOp &mod) { + // NOTE: Wafer gemm instructions require 256-byte alignment for shared + // memory operands. + mod.walk([&](FunctionOpInterface funcOp) { + funcOp.walk([&](Operation *op) { + // FIXME: Abstract for other ops that may require alignment, e.g. + if (!isa( + op)) + return; + for (auto user : op->getOperands()) { + auto allocOp = findAlignmentRestrictOpOperandBuffer(user); + assert(isa(allocOp)); + cast(allocOp).setAlignment(256); + } + }); + return WalkResult::skip(); + }); + } + + // Move the buffer allocations to just before their first user. (buffer and + // user need to be in the same block) + void relocateAllocationsToFirstUser(ModuleOp &mod) { + DenseMap operationId; + mod->walk( + [&](Operation *op) { operationId[op] = operationId.size(); }); + + OpBuilder builder(mod->getContext()); + DenseMap bufferCanMove; + mod->walk([&](Operation *op) { + if (!isa(op)) + return; + auto allocOp = cast(op); + auto users = allocOp->getUsers(); + assert(!users.empty() && "tensor.empty should have no uses here"); + + Operation *minIDUser = nullptr; + size_t minID = std::numeric_limits::max(); + for (auto user : users) { + if (operationId[user] > minID) + continue; + minID = operationId[user]; + minIDUser = user; + } + assert(minIDUser && "There should be at least one user"); + if (!minIDUser->getParentOp()->isAncestor(allocOp)) + return; + + bufferCanMove.insert({allocOp, minIDUser}); + }); + + for (auto [bufferOp, userOp] : bufferCanMove) { + bufferOp->moveBefore(userOp); + } + } + + void setAllocationOffsetAttrs(ModuleOp &mod, MLIRContext *ctx, + triton::alloc::ModuleAllocation &allocation) { + mod.walk([&](FunctionOpInterface funcOp) { + auto *funcAllocation = allocation.getFuncData(funcOp); + funcOp.walk([&](Operation *op) { + // Only handle memref::AllocOp + if (!isa(op)) + return; + int offset = -1; + Value value = op->getResult(0); + auto vBufferId = funcAllocation->getBufferId(value); + if (vBufferId != triton::alloc::Allocation::InvalidBufferId) + offset = funcAllocation->getOffset(vBufferId); + + if (offset == -1) + return; + if (op->hasAttr("allocation.offset")) + return; + op->setAttr("allocation.offset", + IntegerAttr::get(IntegerType::get(ctx, 32), offset)); + }); + return WalkResult::skip(); + }); + mod->setAttr("triton_tsm.spm_use", + mlir::IntegerAttr::get(mlir::IntegerType::get(ctx, 32), + allocation.getSharedMemorySize())); + } + + void runOnOperation() override { + ModuleOp mod = getOperation(); + MLIRContext *ctx = &getContext(); + + handleTargetDependentAlignmentRequirements(mod); + relocateAllocationsToFirstUser(mod); + triton::alloc::ModuleAllocation allocation(mod); + setAllocationOffsetAttrs(mod, ctx, allocation); + } +}; + +} // namespace diff --git a/third_party/wafer/lib/Conversion/AllocateSharedMemory/CMakeLists.txt b/third_party/wafer/lib/Conversion/AllocateSharedMemory/CMakeLists.txt new file mode 100755 index 00000000..a0a21677 --- /dev/null +++ b/third_party/wafer/lib/Conversion/AllocateSharedMemory/CMakeLists.txt @@ -0,0 +1,20 @@ +add_triton_library(AllocateSharedMemory + AllocateSharedMemoryPass.cpp + + DEPENDS + AllocateSharedMemoryPassIncGen + TritonSharedUtils + TritonSharedAnalysis + ZTCAnalysis + + LINK_LIBS PUBLIC + TritonSharedUtils + MLIRDialectUtils + MLIRIR + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonSharedAnalysis + ZTCAnalysis +) diff --git a/third_party/wafer/lib/Conversion/CMakeLists.txt b/third_party/wafer/lib/Conversion/CMakeLists.txt new file mode 100755 index 00000000..e6fbc906 --- /dev/null +++ b/third_party/wafer/lib/Conversion/CMakeLists.txt @@ -0,0 +1,26 @@ +add_subdirectory(TritonArithToLinalg) +add_subdirectory(StructuredToMemref) +add_subdirectory(ConvertTritonPtr) +# TritonPtrToMemref is provided by FLIR +add_subdirectory(WaferMemrefToLLVM) +add_subdirectory(LinalgToMK) +add_subdirectory(MKToWafer) +add_subdirectory(WaferToLLVM) +add_subdirectory(TritonToCoreDialects) +add_subdirectory(CoreDialectsToMK) +add_subdirectory(LegalizeTensorFormLoops) +add_subdirectory(LinalgTiling) +add_subdirectory(LinalgFusion) +add_subdirectory(AllocateSharedMemory) +add_subdirectory(ExportKernelSymbols) +add_subdirectory(TLEToMK) +add_subdirectory(UnstructuredToMK) +add_subdirectory(ReconcilePtrCasts) + +# FLIR conversions are built from third_party/flir/: +# - TritonToLinalg +# - TritonToStructured +# - TritonToUnstructured +# - UnstructuredToMemref + +add_subdirectory(MKPipeline) diff --git a/third_party/wafer/lib/Conversion/ConvertTritonPtr/CMakeLists.txt b/third_party/wafer/lib/Conversion/ConvertTritonPtr/CMakeLists.txt new file mode 100755 index 00000000..90f0b71a --- /dev/null +++ b/third_party/wafer/lib/Conversion/ConvertTritonPtr/CMakeLists.txt @@ -0,0 +1,25 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Triton Project Contributors. +# +#===------------------------------------------------------------------------===# + +add_triton_library(ConvertTritonPtr + TritonPtrToAddressPass.cpp + + DEPENDS + ConvertTritonPtrPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + MLIRReconcileUnrealizedCasts + TritonIR + MLIRAddress +) diff --git a/third_party/wafer/lib/Conversion/ConvertTritonPtr/TritonPtrToAddressPass.cpp b/third_party/wafer/lib/Conversion/ConvertTritonPtr/TritonPtrToAddressPass.cpp new file mode 100755 index 00000000..2882bd6e --- /dev/null +++ b/third_party/wafer/lib/Conversion/ConvertTritonPtr/TritonPtrToAddressPass.cpp @@ -0,0 +1,181 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// +// This pass lowers all triton ops on pointer to their equivalent form in the +// proposed Pointer Dialect: +// https://discourse.llvm.org/t/rfc-ptr-dialect-modularizing-ptr-ops-in-the-llvm-dialect/75142 +// +// This pass is intended to be used after all running +// triton-arith-to-linalg="tensor-ptr-to-linalg=true". +// All triton ops on tensors of pointers are expected to have been lowered to +// linalg ops, and that only triton ops on single pointers remain. +// +// Implementation notes: +// Because triton pointers are typed whereas the !ptr.ptr type isn't. The +// lowering for addptr will have to manually scale the offsets by pointee type. +// As a result, bitcasts are no-op after this pass. +//===----------------------------------------------------------------------===// + +#include "Address/Dialect/IR/AddressDialect.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton-shared/Conversion/ConvertTritonPtr/TritonPtrToAddress.h" +#include "triton-shared/Utils/Utils.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#define DEBUG_TYPE "triton-to-ptr" + +using namespace mlir; + +namespace { + +#define GEN_PASS_DEF_TRITONPTRTOADDRESS +#include "triton-shared/Conversion/ConvertTritonPtr/Passes.h.inc" + +// arith.select could operate on triton pointers. Convert to use !ptr.ptr +struct SelectOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + SelectOpConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(arith::SelectOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp( + op, getTypeConverter()->convertType(op.getType()), + adaptor.getCondition(), adaptor.getTrueValue(), + adaptor.getFalseValue()); + return success(); + } +}; + +// Convert bitcast which is a no-op because !ptr.ptr is opaque with no pointtee +// type. +struct BitCastConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + BitCastConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::BitcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getType())) { + return failure(); + } + + // If the source is a triton pointer, we can convert it to an address + // type. + rewriter.replaceOpWithNewOp( + op, getTypeConverter()->convertType(op.getType()), adaptor.getSrc()); + return success(); + } +}; + +// Convert tt.ptr_to_int to ptr.ptrtoint +struct PtrToIntConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + PtrToIntConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::PtrToIntOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getType())) { + return failure(); + } + rewriter.replaceOpWithNewOp(op, op.getType(), + adaptor.getSrc()); + return success(); + } +}; + +// Convert tt.int_to_ptr to ptr.ptrtoint +struct IntToPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + IntToPtrConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::IntToPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getType())) { + return failure(); + } + rewriter.replaceOpWithNewOp( + op, addr::AddressType::get(rewriter.getContext()), adaptor.getSrc()); + return success(); + } +}; + +class TritonPtrTypeConverter : public TypeConverter { +public: + TritonPtrTypeConverter(MLIRContext *context) { + addConversion([](Type type) { return type; }); + addConversion([context](triton::PointerType ptrType) { + return addr::AddressType::get(context); + }); + addConversion([context](RankedTensorType tensorType) { + if (isa(tensorType.getElementType())) { + return RankedTensorType::get(tensorType.getShape(), + addr::AddressType::get(context)); + } + return tensorType; + }); + auto createCast = [&](OpBuilder &builder, Type resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }; + addTargetMaterialization(createCast); + addSourceMaterialization(createCast); + } +}; + +class TritonPtrToAddressPass + : public impl::TritonPtrToAddressBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + TritonPtrTypeConverter typeConverter(&getContext()); + target.addLegalDialect(); + + target.addIllegalOp(); + target.addDynamicallyLegalOp([](auto op) { + return llvm::all_of( + llvm::concat(op->getOperands(), op->getResults()), + [&](Value v) { return !mlir::triton::isPtrTypeLike(v.getType()); }); + }); + + patterns.add(typeConverter, patterns.getContext()); + + mlir::scf::populateSCFStructuralTypeConversionsAndLegality( + typeConverter, patterns, target); + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> +triton::createTritonPtrToAddressPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/CoreDialectsToMK/CMakeLists.txt b/third_party/wafer/lib/Conversion/CoreDialectsToMK/CMakeLists.txt new file mode 100755 index 00000000..4f092c10 --- /dev/null +++ b/third_party/wafer/lib/Conversion/CoreDialectsToMK/CMakeLists.txt @@ -0,0 +1,25 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +# All rights reserved. +# +#===------------------------------------------------------------------------===# + +add_triton_library(CoreDialectsToMK + CoreDialectsToMKPass.cpp + + DEPENDS + CoreDialectsToMKConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + + LinalgToMagicKernel +) diff --git a/third_party/wafer/lib/Conversion/CoreDialectsToMK/CoreDialectsToMKPass.cpp b/third_party/wafer/lib/Conversion/CoreDialectsToMK/CoreDialectsToMKPass.cpp new file mode 100755 index 00000000..6db66d01 --- /dev/null +++ b/third_party/wafer/lib/Conversion/CoreDialectsToMK/CoreDialectsToMKPass.cpp @@ -0,0 +1,63 @@ +//===------------------- CoreDialectsToMKPass.cpp -------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Lowering core dialects to backend dialects +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Conversion/CoreDialectsToMK/CoreDialectsToMK.h" +#include "magic-kernel/Conversion/LinalgToMK/LinalgToMK.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" + +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/Passes.h" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "magic-kernel/Conversion/CoreDialectsToMK/Passes.h.inc" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" + +namespace { + +class CoreDialectsToMKPass : public CoreDialectsToMKBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + PassManager pm(&getContext(), moduleOp.getOperationName()); + + LinalgToMKOptions options; + options.precisionMode = precisionPriority ? 2 : precisionMode; + pm.addPass(createLinalgToMKPass(options)); + + // Erase dead code and fold constants created during lowering + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> triton::createCoreDialectsToMKPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/ExportKernelSymbols/CMakeLists.txt b/third_party/wafer/lib/Conversion/ExportKernelSymbols/CMakeLists.txt new file mode 100755 index 00000000..547a5deb --- /dev/null +++ b/third_party/wafer/lib/Conversion/ExportKernelSymbols/CMakeLists.txt @@ -0,0 +1,14 @@ +add_triton_library(ExportKernelSymbols + ExportKernelSymbols.cpp + + DEPENDS + ExportKernelSymbolsConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRIR + MLIRPass + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms +) diff --git a/third_party/wafer/lib/Conversion/ExportKernelSymbols/ExportKernelSymbols.cpp b/third_party/wafer/lib/Conversion/ExportKernelSymbols/ExportKernelSymbols.cpp new file mode 100755 index 00000000..0144d223 --- /dev/null +++ b/third_party/wafer/lib/Conversion/ExportKernelSymbols/ExportKernelSymbols.cpp @@ -0,0 +1,168 @@ +//===--------------------- ExportKernelSymbolsPass.cpp +//-----------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Ludt) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Conversion/ExportKernelSymbols/ExportKernelSymbols.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Support/LLVM.h" +#include "wafer/Dialect/IR/WaferDialect.h" +#include "llvm/Support/Debug.h" +#include +#include +#include + +#define DEBUG_TYPE "export-kernel-symbols" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "wafer/Conversion/ExportKernelSymbols/Passes.h.inc" + +namespace { + +class ExportKernelSymbolsPass + : public ExportKernelSymbolsBase { +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + ModuleOp module = getOperation(); + OpBuilder builder(&getContext()); + MLIRContext *ctx = &getContext(); + + Type ptrType = LLVM::LLVMPointerType::get(ctx); + Type symtabType = LLVM::LLVMStructType::getLiteral(ctx, {ptrType, ptrType}); + Type i8Type = IntegerType::get(ctx, 8); + + bool changed = false; + LLVM_DEBUG(llvm::dbgs() << "ExportKernelSymbols: Processing module\n"); + + LLVM::GlobalOp lastGlobal = nullptr; + for (auto global : module.getOps()) { + lastGlobal = global; + } + + auto addFunctionToSections = [&](LLVM::LLVMFuncOp funcOp, + StringRef funcName) { + std::string symbolName = funcName.str(); + Type arrayType = + LLVM::LLVMArrayType::get(ctx, i8Type, symbolName.size() + 1); + + SmallVector bytes; + bytes.append(symbolName.begin(), symbolName.end()); + bytes.push_back(0); + + RankedTensorType tensorType = + RankedTensorType::get({static_cast(bytes.size())}, i8Type); + DenseElementsAttr nameAttr = + DenseElementsAttr::get(tensorType, ArrayRef(bytes)); + + std::string nameVarName = ("_dynsym_" + funcName + "_name").str(); + if (lastGlobal) { + builder.setInsertionPointAfter(lastGlobal); + } else { + builder.setInsertionPointToStart(module.getBody()); + } + + LLVM::GlobalOp nameVar = builder.create( + module.getLoc(), arrayType, /*isConstant=*/true, LLVM::Linkage::Weak, + nameVarName, nameAttr); + nameVar.setSection(".rodata.name"); + nameVar.setAlignment(1); + lastGlobal = nameVar; + + std::string symtabVarName = ("_dynsym_" + funcName).str(); + builder.setInsertionPointAfter(nameVar); + + LLVM::GlobalOp symtabVar = builder.create( + module.getLoc(), symtabType, /*isConstant=*/true, + LLVM::Linkage::External, symtabVarName, Attribute()); + symtabVar.setSection("ExportedDYNSYMTab"); + symtabVar.setAlignment(8); + lastGlobal = symtabVar; + + Region &initRegion = symtabVar.getInitializerRegion(); + Block *initBlock = builder.createBlock(&initRegion); + builder.setInsertionPointToStart(initBlock); + + Value funcAddr = builder.create( + module.getLoc(), ptrType, funcOp.getSymNameAttr()); + Value nameAddr = builder.create( + module.getLoc(), ptrType, nameVar.getSymNameAttr()); + + Value structVal = + builder.create(module.getLoc(), symtabType); + structVal = builder.create( + module.getLoc(), symtabType, structVal, funcAddr, + builder.getDenseI64ArrayAttr({0})); + structVal = builder.create( + module.getLoc(), symtabType, structVal, nameAddr, + builder.getDenseI64ArrayAttr({1})); + + builder.create(module.getLoc(), structVal); + }; + + for (auto funcOp : module.getOps()) { + if (funcOp.isDeclaration()) + continue; + + StringRef funcName = funcOp.getSymName(); + changed = true; + addFunctionToSections(funcOp, funcName); + } + + Type voidType = LLVM::LLVMVoidType::get(ctx); + Type voidPtrType = LLVM::LLVMPointerType::get(ctx); + LLVM::LLVMFunctionType funcType = LLVM::LLVMFunctionType::get( + voidType, {voidPtrType}, /*isVarArg=*/false); + + auto getOrCreateFunc = [&](StringRef name) -> LLVM::LLVMFuncOp { + if (auto existing = module.lookupSymbol(name)) + return existing; + + if (lastGlobal) { + builder.setInsertionPointAfter(lastGlobal); + } else { + builder.setInsertionPointToStart(module.getBody()); + } + + auto func = builder.create( + module.getLoc(), name, funcType, LLVM::Linkage::External, + /*dsoLocal=*/false, /*cconv=*/LLVM::CConv::C); + changed = true; + return func; + }; + + LLVM::LLVMFuncOp moduleInitFunc = getOrCreateFunc("module_init"); + builder.setInsertionPointAfter(moduleInitFunc); + LLVM::LLVMFuncOp moduleCleanupFunc = getOrCreateFunc("module_cleanup"); + + addFunctionToSections(moduleInitFunc, "module_init"); + addFunctionToSections(moduleCleanupFunc, "module_cleanup"); + + LLVM_DEBUG(llvm::dbgs() + << "ExportKernelSymbols: " + << (changed ? "Modified" : "No changes") << " module\n"); + + if (!changed) + markAllAnalysesPreserved(); + } +}; + +} // namespace + +std::unique_ptr> +triton::createExportKernelSymbolsPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/LegalizeTensorFormLoops/CMakeLists.txt b/third_party/wafer/lib/Conversion/LegalizeTensorFormLoops/CMakeLists.txt new file mode 100755 index 00000000..b2b39659 --- /dev/null +++ b/third_party/wafer/lib/Conversion/LegalizeTensorFormLoops/CMakeLists.txt @@ -0,0 +1,23 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +# All rights reserved. +# +#===------------------------------------------------------------------------===# + +add_triton_library(LegalizeTensorFormLoops + LegalizeTensorFormLoops.cpp + + DEPENDS + LegalizeTensorFormLoopsPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport +) diff --git a/third_party/wafer/lib/Conversion/LegalizeTensorFormLoops/LegalizeTensorFormLoops.cpp b/third_party/wafer/lib/Conversion/LegalizeTensorFormLoops/LegalizeTensorFormLoops.cpp new file mode 100755 index 00000000..aaca9bf5 --- /dev/null +++ b/third_party/wafer/lib/Conversion/LegalizeTensorFormLoops/LegalizeTensorFormLoops.cpp @@ -0,0 +1,79 @@ +//===----------------- LegalizeTensorFormLoops.cpp ------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" + +#define DEBUG_TYPE "legalize-tensor-form-loops" + +using namespace mlir; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_LEGALIZETENSORFORMLOOPS +#include "magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { +struct ForOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(scf::ForOp forOp, + PatternRewriter &rewriter) const override { + auto result = failure(); + auto yieldOp = cast(forOp.getBody()->getTerminator()); + rewriter.setInsertionPoint(yieldOp); + for (auto op : llvm::enumerate(yieldOp->getOperands())) { + auto val = op.value(); + auto itArg = forOp.getRegionIterArgs()[op.index()]; + if (!isa(val.getType()) || val == itArg) + continue; + + bool insertCopy = false; + auto defOp = val.getDefiningOp(); + if (defOp) { + // TODO: Use BufferizableOpInterface to analyze whether the operand is + // equivalent to the corresponding iter bbArg. + auto copyOp = dyn_cast(defOp); + insertCopy = !copyOp || copyOp.getOutputs()[0] != itArg; + } else { + // BlockArgument && val != itArg + insertCopy = true; + } + + if (insertCopy) { + auto reduceVal = + rewriter.create(forOp.getLoc(), val, itArg); + yieldOp->setOperand(op.index(), reduceVal->getResult(0)); + result = success(); + } + } + + return result; + } +}; + +class LegalizeTensorFormLoopsPass + : public triton::impl::LegalizeTensorFormLoopsBase< + LegalizeTensorFormLoopsPass> { + using LegalizeTensorFormLoopsBase< + LegalizeTensorFormLoopsPass>::LegalizeTensorFormLoopsBase; + +public: + void runOnOperation() override { + RewritePatternSet patterns(&getContext()); + patterns.add(&getContext()); + if (failed(applyPatternsGreedily(getOperation(), std::move(patterns)))) { + signalPassFailure(); + } + } +}; + +} // namespace diff --git a/third_party/wafer/lib/Conversion/LinalgFusion/CMakeLists.txt b/third_party/wafer/lib/Conversion/LinalgFusion/CMakeLists.txt new file mode 100755 index 00000000..ab1ac6cd --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgFusion/CMakeLists.txt @@ -0,0 +1,17 @@ +add_triton_library(LinalgFusion + LinalgFusion.cpp + LinalgFusionPass.cpp + + DEPENDS + LinalgFusionConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRDialectUtils + MLIRIR + MLIRPass + MLIRLinalgDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms +) diff --git a/third_party/wafer/lib/Conversion/LinalgFusion/LinalgFusion.cpp b/third_party/wafer/lib/Conversion/LinalgFusion/LinalgFusion.cpp new file mode 100755 index 00000000..e8d3eb22 --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgFusion/LinalgFusion.cpp @@ -0,0 +1,275 @@ +//===------------------- LinalgFusion.cpp --------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file implements the patterns to fuse scalar input linalg operations for +// better performance. It applies scalar fusion transformations to reduce +// redundant memory read and write operations. Elementwise fusion needs to +// be implemented later. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Conversion/LinalgFusion/LinalgFusion.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Transforms/Transforms.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/PatternMatch.h" +#include "triton-shared/Utils/Utils.h" +#include "utils/LinalgOpBuilderHelper.h" + +#define DEBUG_TYPE "linalg-fusion" + +using namespace mlir; + +namespace { +template +struct BinaryScalarAndTensorOpFusion + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + bool isValidUnaryElementWiseLinalgOp(linalg::LinalgOp op) const { + if (!op.isSingleInputOutput() || !op.isAllParallelLoops() || + op.getNumLoops() < 1) { + return false; + } + auto regionOps = mlir::triton::getRegionOps(op); + return regionOps.size() == 1 || isa(op); + } + + // Check if the input can be traced back to a fill op through a chain of + // unary elementwise linalg ops + bool + isScalarPropagationInput(Value input, Value &fillInput, + SmallVector &transformOps) const { + // Check if the input is from a fill op or a chain of + auto fillOp = input.getDefiningOp(); + if (fillOp) { + fillInput = fillOp->getOperand(0); + return true; + } + // Check if the input is from a valid unary elementwise linalg op + auto linalgOp = input.getDefiningOp(); + if (linalgOp && isValidUnaryElementWiseLinalgOp(linalgOp)) { + transformOps.push_back(linalgOp); + return isScalarPropagationInput(linalgOp->getOperand(0), fillInput, + transformOps); + } + return false; + } + + // Create the scalar input by applying the chain of unary elementwise ops on + // the fill input + Value createScalarInput(Value input, PatternRewriter &rewriter, + Location loc) const { + Value fillInput; + SmallVector transformOps; + if (!isScalarPropagationInput(input, fillInput, transformOps)) { + return nullptr; + } + Value scalarInput = fillInput; + for (int i = transformOps.size() - 1; i >= 0; i--) { + auto op = transformOps[i]; + SmallVector arithInputs; + auto regionOp = mlir::triton::getRegionOps(op).back(); + auto type = regionOp->getResultTypes()[0]; + if (isa(op)) + arithInputs.push_back( + rewriter.create(loc, rewriter.getOneAttr(type))); + arithInputs.push_back(scalarInput); + scalarInput = rewriter + .create(loc, + rewriter.getStringAttr( + regionOp->getName().getStringRef()), + arithInputs, type) + ->getResult(0); + } + return scalarInput; + } + + // Both inputs are constant propagation, create scalar operation and fill + LogicalResult handleBothScalarInputs(linalg::GenericOp op, + PatternRewriter &rewriter, + Value scalarInput0, + Value scalarInput1) const { + auto regionOps = mlir::triton::getRegionOps(op); + auto resultType = cast(op.getResultTypes()[0]); + auto loc = op->getLoc(); + auto scalarResult = + rewriter + .create(loc, + rewriter.getStringAttr( + regionOps.front()->getName().getStringRef()), + ValueRange{scalarInput0, scalarInput1}, + resultType.getElementType()) + ->getResult(0); + auto fillRes = rewriter + .create(loc, ValueRange{scalarResult}, + op.getOutputs()[0]) + ->getResult(0); + rewriter.replaceAllOpUsesWith(op, fillRes); + return success(); + } + + LogicalResult handleScalarAndTensorInput(linalg::GenericOp op, + PatternRewriter &rewriter, + Value scalarInput, + Value tensorInput) const { + auto resultType = cast(op.getResultTypes()[0]); + auto loc = op->getLoc(); + auto newRes = + rewriter + .create(op->getLoc(), op->getResultTypes()[0], tensorInput, + scalarInput, op.getOutputs()[0]) + .getResult(0); + + // Handle special case: subtraction operation with tensor input in the + // second position. + if (isa(op.getRegion().front().front()) && + tensorInput == op.getInputs()[1]) { + newRes = buildLinalgElementwise( + rewriter, op->getLoc(), resultType, ValueRange{newRes}); + } + rewriter.replaceAllOpUsesWith(op, newRes); + return success(); + } + + LogicalResult convertToIntegerDivision(linalg::GenericOp op, + PatternRewriter &rewriter, + Value dividend, Value divisor) const { + auto divsi = + rewriter.create(op->getLoc(), dividend, divisor); + auto siToFp = rewriter.create( + op->getLoc(), + cast(op.getResultTypes()[0]).getElementType(), divsi); + + auto fillRes = rewriter + .create(op->getLoc(), ValueRange{siToFp}, + op.getOutputs()[0]) + ->getResult(0); + + rewriter.replaceAllOpUsesWith(op, fillRes); + return success(); + } + + // Linalg op DivF will convert to reciprocal + mul, scalar reciprocal + // will convert to scalar divf which has low precision. + LogicalResult handleDivisionCase(linalg::GenericOp op, + PatternRewriter &rewriter, + linalg::ReciprocalOp reciprocalOp) const { + auto reciprocal = reciprocalOp->getResult(0); + auto dividend = + reciprocal == op.getInputs()[0] ? op.getInputs()[1] : op.getInputs()[0]; + + Value dividendScalar; + SmallVector dividendTransforms; + // tensor + scalar/tensor recip: no conversion + if (!isScalarPropagationInput(dividend, dividendScalar, dividendTransforms)) + return failure(); + + Value recipScalar; + SmallVector recipTransforms; + // scalar + tensor recip -> mulvs + if (!isScalarPropagationInput(reciprocalOp->getOperand(0), recipScalar, + recipTransforms)) { + return handleScalarAndTensorInput(op, rewriter, dividendScalar, + reciprocal); + } + + const bool isIntegerInput = isa(dividendScalar.getType()) && + isa(recipScalar.getType()); + const bool hasSingleTransformOp = + dividendTransforms.size() == 1 && recipTransforms.size() == 1; + // Convert scalar + scalar divf + sitofp to divsi. + if (isIntegerInput && hasSingleTransformOp && + isa(mlir::triton::getRegionOps( + dividendTransforms.front()) + .front()) && + isa(mlir::triton::getRegionOps( + recipTransforms.front()) + .front())) { + return convertToIntegerDivision(op, rewriter, dividendScalar, + recipScalar); + } + + // scalar + scalar recip : no conversion + return failure(); + } + + linalg::ReciprocalOp matchDivAndGetRecip(linalg::GenericOp op) const { + auto lhsRecip = op->getOperand(0).getDefiningOp(); + auto rhsRecip = op->getOperand(1).getDefiningOp(); + assert(!(lhsRecip && rhsRecip) && + "Currently, we only handle cases where one input of the mul op is a " + "reciprocal op."); + if (auto reciprocalOp = lhsRecip ? lhsRecip : rhsRecip) + return reciprocalOp; + return nullptr; + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + if (!isaElemwiseSingleBinaryOpInterface(op)) { + return failure(); + } + auto regionOps = mlir::triton::getRegionOps(op); + if (!isa(regionOps.front())) { + return failure(); + } + + assert(regionOps.size() == 1 && + "Expected a single operation in the linalg.generic region"); + + // Handle multiplication special case with reciprocal: + // DivF/DivSI convert to reciprocal + mul. + if (std::is_same_v) { + if (auto reciprocalOp = matchDivAndGetRecip(op)) { + return handleDivisionCase(op, rewriter, reciprocalOp); + } + } + + auto input0 = op.getInputs()[0]; + auto input1 = op.getInputs()[1]; + auto loc = op->getLoc(); + // Try to create scalar input for both inputs + Value scalarInput0 = createScalarInput(input0, rewriter, loc); + Value scalarInput1 = createScalarInput(input1, rewriter, loc); + if (!scalarInput1 && !scalarInput0) { + return failure(); + } + + if (scalarInput0 && scalarInput1) { + return handleBothScalarInputs(op, rewriter, scalarInput0, scalarInput1); + } + + Value scalarInput = scalarInput1 ? scalarInput1 : scalarInput0; + Value tensorInput = scalarInput1 ? input0 : input1; + + return handleScalarAndTensorInput(op, rewriter, scalarInput, tensorInput); + } +}; + +} // namespace + +void mlir::triton::populateLinalgBinaryOpFusionPatterns( + RewritePatternSet &patterns) { + patterns.add>( + patterns.getContext()); + patterns.add>( + patterns.getContext()); + patterns.add>( + patterns.getContext()); +} + +// TODO: Support linalg elementwise op fusion. +#if 0 +void mlir::triton::populateLinalgFusionPatterns(RewritePatternSet &patterns) { + // Add folding with reshape by expansion patterns. + linalg::ControlFusionFn defaultControlFn = [](OpOperand *fusedOperand) { + return false; + }; + linalg::populateElementwiseOpsFusionPatterns(patterns, defaultControlFn); +} +#endif diff --git a/third_party/wafer/lib/Conversion/LinalgFusion/LinalgFusionPass.cpp b/third_party/wafer/lib/Conversion/LinalgFusion/LinalgFusionPass.cpp new file mode 100755 index 00000000..35765a6b --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgFusion/LinalgFusionPass.cpp @@ -0,0 +1,66 @@ +//===------------------- LinalgFusionPass.cpp -----------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file implements the pass infrastructure for linalg fusion +// transformations. The pass applies fusion patterns to improve performance of +// linalg operations. +// +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Transforms/DialectConversion.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "wafer/Conversion/LinalgFusion/LinalgFusion.h" +#include +#include +#include + +#define DEBUG_TYPE "linalg-fusion" + +using namespace mlir; + +namespace mlir { +namespace triton { + +#define GEN_PASS_DEF_LINALGFUSION +#include "wafer/Conversion/LinalgFusion/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class LinalgFusionPass + : public triton::impl::LinalgFusionBase { +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + registry.insert(); + } + + void runOnOperation() override { + ModuleOp module = getOperation(); + MLIRContext *context = &getContext(); + { + RewritePatternSet selfPatterns(context); + + mlir::triton::populateLinalgBinaryOpFusionPatterns(selfPatterns); + + if (failed(applyPatternsGreedily(module, std::move(selfPatterns)))) { + signalPassFailure(); + } + } + } +}; +} // namespace + +std::unique_ptr> +mlir::triton::createLinalgFusionPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/LinalgTiling/CMakeLists.txt b/third_party/wafer/lib/Conversion/LinalgTiling/CMakeLists.txt new file mode 100755 index 00000000..5f90ae6e --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgTiling/CMakeLists.txt @@ -0,0 +1,17 @@ +add_triton_library(LinalgTiling + LinalgTiling.cpp + LinalgTilingPass.cpp + + DEPENDS + LinalgTilingConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRDialectUtils + MLIRIR + MLIRPass + MLIRLinalgDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms +) diff --git a/third_party/wafer/lib/Conversion/LinalgTiling/LinalgTiling.cpp b/third_party/wafer/lib/Conversion/LinalgTiling/LinalgTiling.cpp new file mode 100755 index 00000000..66af3e5f --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgTiling/LinalgTiling.cpp @@ -0,0 +1,87 @@ +//===------------------- LinalgTiling.cpp --------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file implements the patterns to tile linalg operations for better +// performance. It applies tiling transformations to improve data locality +// and parallelism, focusing on operations like linalg.reduce and +// linalg.generic. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Conversion/LinalgTiling/LinalgTiling.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Transforms/Transforms.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "triton-shared/Utils/Utils.h" +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "linalg-tiling" + +using namespace mlir; + +namespace { + +// Extract the operations from a linalg op region +template llvm::SmallVector getRegionOps(T linalgOp) { + auto regionBlock = linalgOp.getBody(); + return llvm::map_to_vector(regionBlock->without_terminator(), + [](Operation &op) { return &op; }); +} + +struct TilingReduceRewrite : public OpRewritePattern { + TilingReduceRewrite(MLIRContext *context) + : OpRewritePattern(context, /*benefit=*/1) {} + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(linalg::ReduceOp op, + PatternRewriter &rewriter) const override { + auto dims = op.getDimensions(); + if (dims.size() != 1) { + op->emitError() << "Only support one dim reduce."; + return rewriter.notifyMatchFailure(op, "Only support one dim reduce."); + } + + auto dim = dims[0]; + auto inputType = cast(op.getInputs()[0].getType()); + auto inputShape = inputType.getShape(); + + // Tiling if shape[dim]>32768 + if (inputShape[dim] < 32768) { + return failure(); + } + + auto regionOps = getRegionOps(op); + // FIXME: Move after type conversion pass + if (regionOps.size() != 1 || + !triton::isTargetSupportedReductionOp(regionOps.front())) + return failure(); + + linalg::LinalgTilingOptions tilingOptions; + auto tileSizes = SmallVector(inputShape); + assert(dim == 0 && "Expected tiling on the first dimension"); + assert(llvm::isPowerOf2_64(inputShape[dim]) && + "Expected power of 2 for tiling size"); + tileSizes[dim] = 16384; + tilingOptions.setTileSizes(tileSizes); + + auto tiled = linalg::tileLinalgOp(rewriter, op, tilingOptions); + if (failed(tiled)) + return rewriter.notifyMatchFailure(op, "operation not supported yet."); + + rewriter.replaceOp(op, tiled.value().tensorResults); + return success(); + } +}; + +} // namespace + +void mlir::triton::populateLinalgTilingPatterns(RewritePatternSet &patterns) { + patterns.add(patterns.getContext()); +} diff --git a/third_party/wafer/lib/Conversion/LinalgTiling/LinalgTilingPass.cpp b/third_party/wafer/lib/Conversion/LinalgTiling/LinalgTilingPass.cpp new file mode 100755 index 00000000..f7ae91ba --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgTiling/LinalgTilingPass.cpp @@ -0,0 +1,64 @@ +//===------------------- LinalgTilingPass.cpp -----------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file implements the pass infrastructure for linalg tiling +// transformations. The pass applies tiling patterns to improve performance of +// linalg operations. +// +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Transforms/DialectConversion.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "wafer/Conversion/LinalgTiling/LinalgTiling.h" +#include +#include +#include + +#define DEBUG_TYPE "linalg-tiling" + +using namespace mlir; + +namespace mlir { +namespace triton { + +#define GEN_PASS_DEF_LINALGTILING +#include "wafer/Conversion/LinalgTiling/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class LinalgTilingPass + : public triton::impl::LinalgTilingBase { +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + ModuleOp module = getOperation(); + MLIRContext *context = &getContext(); + + RewritePatternSet patterns(context); + + mlir::triton::populateLinalgTilingPatterns(patterns); + + if (failed(applyPatternsGreedily(module, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> +mlir::triton::createLinalgTilingPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/LinalgToMK/CMakeLists.txt b/third_party/wafer/lib/Conversion/LinalgToMK/CMakeLists.txt new file mode 100755 index 00000000..dfb84692 --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgToMK/CMakeLists.txt @@ -0,0 +1,21 @@ +add_triton_library(LinalgToMagicKernel + LinalgToMK.cpp + LinalgToMKPass.cpp + + DEPENDS + MagicKernelTableGen + LinalgToMKConversionPassIncGen + TritonSharedUtils + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonSharedUtils +) diff --git a/third_party/wafer/lib/Conversion/LinalgToMK/LinalgToMK.cpp b/third_party/wafer/lib/Conversion/LinalgToMK/LinalgToMK.cpp new file mode 100755 index 00000000..5a740155 --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgToMK/LinalgToMK.cpp @@ -0,0 +1,4244 @@ +//===------------------- LinalgToMK.cpp -----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Conversion/LinalgToMK/LinalgToMK.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "triton-shared/Utils/FusionHelper.h" +#include "triton-shared/Utils/ReduceScanCommon.h" +#include "triton-shared/Utils/Utils.h" +#include "utils/LinalgOpBuilderHelper.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/LogicalResult.h" + +#define DEBUG_TYPE "linalg-to-mk" + +using namespace mlir; +using namespace mk; + +#define GEN_PASS_CLASSES +#include "magic-kernel/Conversion/LinalgToMK/Passes.h.inc" + +namespace { + +static int normalizePrecisionMode(int precisionMode) { + if (precisionMode <= 0) + return 0; + if (precisionMode == 1) + return 1; + return 2; +} + +static bool preservesIntegerPrecision(Type elementType, int precisionMode) { + if (!isa(elementType)) + return false; + + int bitWidth = elementType.getIntOrFloatBitWidth(); + switch (normalizePrecisionMode(precisionMode)) { + case 0: + return false; + case 1: + return bitWidth >= 64; + default: + return bitWidth >= 32; + } +} + +bool isConstantValue(Value &v, double targetValue, bool isApprox = false) { + auto constOp = v.getDefiningOp(); + if (!constOp) { + return false; + } + if (auto val = dyn_cast(constOp.getValue())) { + return isApprox ? (std::abs(val.getValueAsDouble() - targetValue) < 1e-5) + : (val.getValueAsDouble() == targetValue); + } + if (auto val = dyn_cast(constOp.getValue())) { + return val.getValue() == static_cast(targetValue); + } + return false; +} + +bool isConstantTensor(Value &v, double targetValue, bool isApprox) { + auto *defOp = v.getDefiningOp(); + if (!defOp) { + return false; + } + auto fillOp = dyn_cast(defOp); + if (!fillOp) { + return false; + } + + auto fillValue = fillOp.getInputs()[0]; + return isConstantValue(fillValue, targetValue, isApprox); +} + +// Check if the given value is a tensor filled with 0. +bool isZeroTensor(Value &v) { return isConstantTensor(v, 0.0); } + +// Check if the given value is a tensor filled with 1. +bool isOneTensor(Value &v) { return isConstantTensor(v, 1.0); } + +bool isHalfTensor(Value &v) { + const float halfValue = 0.5f; + return isConstantTensor(v, halfValue); +} + +bool isTwoTensor(Value &v) { + // Check if the value is a constant tensor with the value of 2.0. + const float twoValue = 2.0f; + return isConstantTensor(v, twoValue); +}; + +bool checkReductionBaseAttr(linalg::ReduceOp op, OpBuilder &builder, + TypedAttr &attr) { + auto out = op.getInits().front(); + + auto outDef = out.getDefiningOp(); + if (!outDef) { + return false; + } + Value value; + if (isa(outDef)) { + value = cast(outDef).value(); + } else if (isa(outDef)) { + value = cast(outDef).getScalar(); + } else { + // If the output is not a fill or insert op, we cannot determine the + // reduction base attribute. + return false; + } + auto valueDef = value.getDefiningOp(); + if (!valueDef) { + return false; + } + + return valueDef.getValueAttr() == attr; +} + +static linalg::ReduceOp createReduceOp(OpBuilder &rewriter, linalg::ReduceOp op, + Location loc, ValueRange sources, + SmallVector dims, + SmallVector shape, + Type elementType, TypedAttr attr) { + Value init = rewriter.create(loc, shape, elementType); + + auto accBaseConstOp = + rewriter.create(loc, elementType, attr); + Value initTensor = rewriter + .create( + loc, ValueRange{accBaseConstOp}, ValueRange{init}) + .result(); + + return rewriter.create( + loc, sources, ValueRange{initTensor}, dims, + [&](OpBuilder &opBuilder, Location loc, ValueRange inputs) { + assert(inputs.size() == 2); + + auto reduceBlock = op.getBody(); + IRMapping mapping; + mapping.map(reduceBlock->getArguments(), inputs); + + for (auto &op : reduceBlock->without_terminator()) { + opBuilder.clone(op, mapping); + } + + auto yield = reduceBlock->getTerminator(); + auto results = + llvm::map_to_vector(yield->getOperands(), + [&](Value val) { return mapping.lookup(val); }); + + opBuilder.create(loc, results); + }); +}; + +// Convert linalg.matmul to mk.dot +struct LinalgMatmulOpRewrite : public OpRewritePattern { +private: + using OpRewritePattern::OpRewritePattern; + + Value channelNorm(PatternRewriter &rewriter, Location loc, Value src) const { + auto type = cast(src.getType()); + auto elementType = type.getElementType(); + int64_t alignBase = elementType.getIntOrFloatBitWidth() == 8 ? 128 : 64; + int64_t M = type.getShape()[0]; + int64_t N = type.getShape()[1]; + + assert(N >= 4 && llvm::isPowerOf2_64(N) && + "N must be at least 4 and a power of 2"); + + int64_t N1 = std::max(1L, N / alignBase); + int64_t N2 = std::min(alignBase, N); + + if (N1 == 1) + return rewriter.create( + loc, type.clone({1, M, N}), src, + ArrayRef{{0, 1}, {2}}); + + Value reshpe = rewriter.create( + loc, type.clone({M, N1, N2}), src, + ArrayRef{{0}, {1, 2}}); + auto permsTensor = rewriter.create( + loc, ArrayRef{N1, M, N2}, elementType); + return rewriter + .create(loc, reshpe, permsTensor, + ArrayRef{1, 0, 2}) + ->getResult(0); + } + + Value dechannelNorm(PatternRewriter &rewriter, Location loc, + Value src) const { + auto type = cast(src.getType()); + auto elementType = type.getElementType(); + int64_t N1 = type.getShape()[0]; + int64_t M = type.getShape()[1]; + int64_t N2 = type.getShape()[2]; + + if (N1 == 1) + return rewriter.create( + loc, src, ArrayRef{{0, 1}, {2}}); + + auto permsTensor = rewriter.create( + loc, ArrayRef{M, N1, N2}, elementType); + Value transpose = rewriter + .create( + loc, src, permsTensor, ArrayRef{1, 0, 2}) + ->getResult(0); + return rewriter.create( + loc, transpose, ArrayRef{{0}, {1, 2}}); + } + + // Find the `scf.for` loop in which the current `linalg.matmul` is used as the + // iteration variable. + std::optional> + findForOpAndIdx(linalg::MatmulOp op) const { + if (!op.getResult(0).hasOneUse()) + return std::nullopt; + + auto U = op.getResult(0).use_begin(); + auto idx = U->getOperandNumber(); + auto yieldOp = dyn_cast(U->getOwner()); + + if (!yieldOp) + return std::nullopt; + + auto forOp = dyn_cast(yieldOp->getParentOp()); + + if (!forOp) + return std::nullopt; + + auto iterArg = forOp.getRegionIterArgs()[idx]; + + if (!iterArg.hasOneUse() || op.getOutputs()[0] != iterArg) + return std::nullopt; + + return std::make_pair(forOp, idx); + } + +public: + LogicalResult matchAndRewrite(linalg::MatmulOp op, + PatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + auto a = channelNorm(rewriter, loc, op.getInputs()[0]); + auto b = channelNorm(rewriter, loc, op.getInputs()[1]); + auto fillOp = op.getOutputs()[0].getDefiningOp(); + + if (fillOp && matchPattern(fillOp.getInputs()[0], m_AnyZeroFloat())) { + Value res = rewriter.create( + loc, channelNorm(rewriter, loc, op.getOutputs()[0]).getType(), + ValueRange{}); + res = rewriter + .create(loc, res.getType(), a, b, res, + false /* en_psum */) + ->getResult(0); + rewriter.replaceOp(op, dechannelNorm(rewriter, loc, res)); + } else if (auto optFromLoopInfo = findForOpAndIdx(op)) { + // If we find that the current `linalg.matmul` is the loop's iteration + // variable, we can hoist the accumulation matrix "channelNorm" outside + // the loop and sink the "dechannelNorm" outside as well. + auto [forOp, idx] = *optFromLoopInfo; + SmallVector inits = forOp.getInits(); + rewriter.setInsertionPoint(forOp); + inits[idx] = channelNorm(rewriter, loc, inits[idx]); + auto newForOp = rewriter.create( + forOp->getLoc(), forOp.getLowerBound(), forOp.getUpperBound(), + forOp.getStep(), inits); + auto body = newForOp.getBody(); + body->getOperations().splice(body->begin(), + forOp.getBody()->getOperations()); + forOp.getInductionVar().replaceAllUsesWith(newForOp.getInductionVar()); + for (unsigned i = 0; i < forOp.getNumRegionIterArgs(); ++i) { + if (i == idx) + continue; + forOp.getRegionIterArg(i).replaceAllUsesWith( + newForOp.getRegionIterArg(i)); + forOp->getResult(i).replaceAllUsesWith(newForOp->getResult(i)); + } + + rewriter.setInsertionPoint(op); + auto dot = rewriter.create( + loc, newForOp->getResultTypes()[idx], a, b, + newForOp.getRegionIterArg(idx), true /* en_psum */); + body->getTerminator()->setOperand(idx, dot->getResult(0)); + rewriter.eraseOp(op); + + rewriter.setInsertionPointAfter(newForOp); + Value res = + dechannelNorm(rewriter, forOp->getLoc(), newForOp->getResult(idx)); + forOp->getResult(idx).replaceAllUsesWith(res); + rewriter.eraseOp(forOp); + } else { + Value output = channelNorm(rewriter, loc, op.getOutputs()[0]); + output = rewriter + .create(loc, output.getType(), a, b, output, + true /* en_psum */) + ->getResult(0); + rewriter.replaceOp(op, dechannelNorm(rewriter, loc, output)); + } + + return success(); + } +}; + +struct MKDotScaleOpRewrite : public OpRewritePattern { +private: + using OpRewritePattern::OpRewritePattern; + + Value buildExtF(OpBuilder &rewriter, Location loc, Value input, + Type targetType) const { + auto inputType = cast(input.getType()); + auto empty = + rewriter.create(loc, inputType.getShape(), targetType); + + auto rank = inputType.getRank(); + + auto identityMap = + AffineMap::getMultiDimIdentityMap(rank, rewriter.getContext()); + SmallVector indexingMaps(2, identityMap); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + + return rewriter + .create( + loc, TypeRange{empty.getType()}, ValueRange{input}, + ValueRange{empty}, indexingMaps, iteratorTypes, + [&](OpBuilder &nestedBuilder, Location nestedloc, + ValueRange iterArgs) { + auto extf = nestedBuilder.create( + nestedloc, targetType, iterArgs[0]); + nestedBuilder.create(nestedloc, + ValueRange{extf}); + }) + ->getResult(0); + } + + Type dotTypeFromAttr(triton::ScaleDotElemType type, MLIRContext *ctx) const { + switch (type) { + case triton::ScaleDotElemType::E4M3: + return Float8E4M3FNType::get(ctx); + case triton::ScaleDotElemType::E5M2: + return Float8E5M2Type::get(ctx); + case triton::ScaleDotElemType::E2M1: + return Float4E2M1FNType::get(ctx); + case triton::ScaleDotElemType::BF16: + return BFloat16Type::get(ctx); + case triton::ScaleDotElemType::FP16: + return Float16Type::get(ctx); + // unsupported + // case triton::ScaleDotElemType::E2M3: + // return Float6E2M3FNType::get(ctx); + // case triton::ScaleDotElemType::E3M2: + // return Float6E3M2FNType::get(ctx); + default: + llvm_unreachable("unsupported type!"); + } + } + + Value upcast(OpBuilder &rewriter, Location loc, Value v, Value vScale, + Type vElemType, Type compType, bool transposed) const { + + if (vElemType == compType) + return v; + + if (!vScale) + return buildExtF(rewriter, loc, v, compType); + + Value scaleInput = v; + auto tensorType = cast(v.getType()); + // Since b matrix scale is stored in column major, need transpose + if (transposed) { + + auto originShape = tensorType.getShape(); + SmallVector perm{1, 0}; + SmallVector transposeShape = {originShape[1], originShape[0]}; + auto empty = rewriter.create( + loc, transposeShape, tensorType.getElementType()); + scaleInput = rewriter.create(loc, v, empty, perm) + ->getResult(0); + } + + // exsample: + // tt.dot_scaled %26, %22 scale %25, %cst_0 lhs = e4m3 rhs = e2m1 + // : tensor<4x32xf8E4M3FN> * tensor<16x4xi8>, tensor<4x1xi8> -> + // tensor<4x4xf32> + if (tensorType.getElementType().isInteger(8) && + isa(vElemType)) { + auto shape = cast(scaleInput.getType()).getShape(); + SmallVector newShape(shape); + newShape.back() *= 2; + scaleInput = rewriter.create( + loc, + RankedTensorType::get(newShape, + Float4E2M1FNType::get(vElemType.getContext())), + scaleInput); + } + + assert(cast(scaleInput.getType()).getElementType() == + vElemType); + + scaleInput = buildExtF(rewriter, loc, scaleInput, compType); + + // dstBuffer inplace + if (vScale) { + auto shape = cast(scaleInput.getType()).getShape(); + auto empty = rewriter.create(loc, shape, compType); + scaleInput = rewriter + .create( + loc, RankedTensorType::get(shape, compType), + scaleInput, vScale, empty) + ->getResult(0); + } + + // read: dstBuffer, write: transposeMid + if (transposed) { + // e2m1 need use shape after bitcast + auto shape = cast(scaleInput.getType()).getShape(); + assert(shape.size() == 2); + SmallVector newShape{shape[1], shape[0]}; + auto empty = rewriter.create(loc, newShape, compType); + SmallVector perm{1, 0}; + scaleInput = + rewriter.create(loc, scaleInput, empty, perm) + ->getResult(0); + } + + return scaleInput; + } + +public: + LogicalResult matchAndRewrite(mk::DotScaledOp op, + PatternRewriter &rewriter) const override { + + auto loc = op.getLoc(); + Value a = op.getA(); + Value b = op.getB(); + Value dst = op.getDst(); + Value aScale = op.getAScale(); + Value bScale = op.getBScale(); + auto aElemAttr = op.getAElemType(); + auto bElemAttr = op.getBElemType(); + + // Cast input to compType + auto aElemType = dotTypeFromAttr(aElemAttr, rewriter.getContext()); + auto bElemType = dotTypeFromAttr(bElemAttr, rewriter.getContext()); + + assert(cast(a.getType()).getRank() == 2 && + "support rank is 2 only"); + + Type compType = (aElemType.isF16() || bElemType.isF16()) + ? rewriter.getF16Type() + : rewriter.getBF16Type(); + + // If has scale, do quantization + + a = upcast(rewriter, loc, a, aScale, aElemType, compType, false); + b = upcast(rewriter, loc, b, bScale, bElemType, compType, true); + + // Do standard matmul + auto matmulOp = rewriter.create( + loc, op->getResultTypes(), ValueRange{a, b}, ValueRange{dst}); + + rewriter.replaceOp(op, matmulOp->getResults()); + + return success(); + } +}; + +struct NormalizeReduceInitToIdentityPattern + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + Operation *accumulateInit(OpBuilder &builder, linalg::ReduceOp op, + Value reduceVal, Location loc) const { + Value init = op.getInits()[0]; + auto outputType = cast(init.getType()); + auto rank = outputType.getRank(); + SmallVector idMaps(3, builder.getMultiDimIdentityMap(rank)); + SmallVector iterators( + rank, mlir::utils::IteratorType::parallel); + auto genericOp = builder.create( + loc, op->getResultTypes(), ValueRange{reduceVal, init}, + ValueRange{init}, idMaps, iterators); + genericOp.getRegion().takeBody(op.getRegion()); + genericOp.getRegion().front().addArgument(outputType.getElementType(), loc); + + return genericOp; + } + + LogicalResult matchAndRewrite(linalg::ReduceOp op, + PatternRewriter &rewriter) const override { + + auto reduceOps = getRegionOps(op); + + auto *reduceOp = reduceOps.front(); + if (reduceOps.size() != 1) + return failure(); + + // If the init value is reduction op base(reduction operation identity + // value), don't need to accumulate it + auto resType = cast(op.getInits()[0].getType()); + auto constantType = resType.getElementType(); + auto inputType = cast(op.getInputs()[0].getType()); + + // TODO: Config according backend + // NOTE: Assume has done integer to float + if (!(isReductionOpAndTypeSupportedByTarget(reduceOp, constantType) || + isReduceToElementWiseOpAndTypeSupportedByTarget( + reduceOp, constantType, inputType.getNumElements(), + inputType.getRank()))) { + return failure(); + } + + auto attr = getRedBaseAttr(rewriter, reduceOp, constantType); + if (checkReductionBaseAttr(op, rewriter, attr)) + return failure(); + + auto loc = op.getLoc(); + SmallVector dims(op.getDimensions().begin(), + op.getDimensions().end()); + SmallVector shape(resType.getShape().begin(), + resType.getShape().end()); + Value finalResult = createReduceOp(rewriter, op, loc, op.getInputs(), dims, + shape, constantType, attr) + ->getResult(0); + + auto newOp = accumulateInit(rewriter, op, finalResult, loc); + + rewriter.replaceOp(op, newOp->getResults()); + return success(); + } +}; + +// todo : Determine whether precision promotion is needed based on hardware +// characteristics, eg, fp16->fp32. +// Tsingmicro does not require precision promotion since it is handled in +// hardware computation. +struct ReducePrecisionPromotionRewrite + : public OpRewritePattern { + + bool requiresF32Conversion(const Type elemType, Operation *redOp) const { + return isa(elemType) && + elemType.getIntOrFloatBitWidth() < + cast(Float32Type::get(elemType.getContext())) + .getWidth() && + (isa(redOp) || isa(redOp)); + } +}; + +struct LinalgReduceToMKReduceConversion + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + SmallVector reshapeReduceShapeTo4d(ArrayRef inputShape, + int64_t dim) const { + + auto rank = inputShape.size(); + assert(dim < rank && "Dim out of range"); + + SmallVector newShape; + int64_t leftDimsElement = 1; + int64_t rightDimsElement = 1; + + for (int i = 0; i < dim; i++) + leftDimsElement *= inputShape[i]; + + if (dim == inputShape.size() - 1) + return {1, 1, leftDimsElement, inputShape[dim]}; + + for (int i = dim + 1; i < rank; i++) + rightDimsElement *= inputShape[i]; + + newShape = {1, leftDimsElement, inputShape[dim], rightDimsElement}; // NHWC + return newShape; + } + + Value createReshapeOp(OpBuilder &rewriter, Location loc, Value src, + ArrayRef outputShape, Type elementType) const { + + int32_t rank = outputShape.size(); + + auto shapeTensorShape = SmallVector{rank}; + Value shapeTensor = rewriter.create(loc, shapeTensorShape, + rewriter.getI32Type()); + + for (int32_t i = 0; i < rank; i++) { + auto index = rewriter.create(loc, i); + auto dim = rewriter.create( + loc, rewriter.getI32Type(), outputShape[i]); + shapeTensor = rewriter.create(loc, dim, shapeTensor, + ValueRange{index}); + } + return rewriter.create( + loc, RankedTensorType::get(outputShape, elementType), src, shapeTensor); + } + + Value convertToMKReduce(OpBuilder &rewriter, Location loc, + Operation *reduceOp, Value channelNormedInput, + SmallVector &inputShape4D, Type elementType, + bool lastDimReduce) const { + + auto newReduceDim = lastDimReduce ? 3 : 2; + + auto channelNormedShape = + cast(channelNormedInput.getType()).getShape(); + + // Reduce output + SmallVector channelNormedReduceOutput4D = + lastDimReduce ? SmallVector{inputShape4D[0], inputShape4D[1], + inputShape4D[2], 4} + : SmallVector{1, channelNormedShape[0], + channelNormedShape[1], + channelNormedShape[3]}; + + auto empty = rewriter.create( + loc, channelNormedReduceOutput4D, elementType); + auto outputType = + RankedTensorType::get(channelNormedReduceOutput4D, elementType); + ArrayAttr nhwcShape = rewriter.getI64ArrayAttr(inputShape4D); + + if (isa(reduceOp)) { + return rewriter + .create(loc, outputType, channelNormedInput, empty, + nhwcShape, newReduceDim) + ->getResult(0); + } + if (isa(reduceOp)) { + return rewriter + .create(loc, outputType, channelNormedInput, empty, + nhwcShape, newReduceDim) + ->getResult(0); + } + if (isa(reduceOp)) { + return rewriter + .create(loc, outputType, channelNormedInput, empty, + nhwcShape, newReduceDim) + ->getResult(0); + } + llvm_unreachable("Unsupported reduction operation"); + return nullptr; + } + + Value channelNorm(PatternRewriter &rewriter, Location loc, Value src, + SmallVector &inputShape4D) const { + + auto type = cast(src.getType()); + auto elementType = type.getElementType(); + + int64_t alignBase = elementType.getIntOrFloatBitWidth() == 8 ? 128 : 64; + + int lastDim = inputShape4D.back(); + // Triton always assume shape is power of 2 + assert(llvm::isPowerOf2_64(lastDim) && "LastDim must be power of 2"); + + if (lastDim > alignBase) { + + // {1, H, W, C} -> {H, W, C} + SmallVector collapseShape(inputShape4D.begin() + 1, + inputShape4D.end()); + Value collapse = rewriter.create( + loc, type.clone(collapseShape), src, + ArrayRef{{0, 1}, {2}, {3}}); + + // {H, W, C} -> {H, W, CX, C} + int64_t cx = lastDim / alignBase; + SmallVector expandShape{collapseShape[0], collapseShape[1], + cx, alignBase}; + + Value expand = rewriter.create( + loc, type.clone(expandShape), collapse, + ArrayRef{{0}, {1}, {2, 3}}); + + // {H, W, CX, C} -> {CX, H, W, C} + SmallVector channelnormShape = { + expandShape[2], expandShape[0], expandShape[1], expandShape[3]}; + auto permutateTensor = + rewriter.create(loc, channelnormShape, elementType); + auto channelNorm = + rewriter + .create(loc, expand, permutateTensor, + ArrayRef{2, 0, 1, 3}) + ->getResult(0); + + return channelNorm; + } else if (lastDim < 4) { + // {1, H, W, C} -> {1, H, W, 4} + SmallVector channelnormShape = inputShape4D; + channelnormShape.back() = 4; + + auto empty = + rewriter.create(loc, channelnormShape, elementType); + + auto insert = rewriter.create( + loc, RankedTensorType::get(channelnormShape, elementType), src, empty, + ValueRange(), ValueRange(), ValueRange(), + ArrayRef{0, 0, 0, 0}, inputShape4D, + ArrayRef{1, 1, 1, 1}); + return insert; + } + + return src; + } + + Value dechannelNorm(PatternRewriter &rewriter, Location loc, Value src, + SmallVector &inputShape4D, + bool lastDimReduce) const { + + auto type = cast(src.getType()); + auto elementType = type.getElementType(); + int64_t alignBase = elementType.getIntOrFloatBitWidth() == 8 ? 128 : 64; + + int lastDim = inputShape4D.back(); + + SmallVector outputShape = + lastDimReduce + ? SmallVector{1, 1, inputShape4D[2], 1} + : SmallVector{1, 1, inputShape4D[1], inputShape4D.back()}; + + if (4 <= lastDim && lastDim <= alignBase && !lastDimReduce) + return src; + + if (!lastDimReduce && lastDim > alignBase) { + // {1, cx, left, c0} -> {1, left cx, c0} + + SmallVector permutationShape = {1, inputShape4D[1], + lastDim / alignBase, alignBase}; + auto permutateTensor = + rewriter.create(loc, permutationShape, elementType); + return rewriter + .create(loc, src, permutateTensor, + ArrayRef{0, 2, 1, 3}) + ->getResult(0); + } + + return rewriter.create( + loc, RankedTensorType::get(outputShape, elementType), src, ValueRange(), + ValueRange(), ValueRange(), ArrayRef{0, 0, 0, 0}, outputShape, + ArrayRef{1, 1, 1, 1}); + } + + LogicalResult matchAndRewrite(linalg::ReduceOp op, + PatternRewriter &rewriter) const override { + + auto reduceOps = getRegionOps(op); + + auto *reduceOp = reduceOps.front(); + if (reduceOps.size() != 1) + return failure(); + // TODO: Config according backend + // Assume has done integer to float conversion + auto inputType = dyn_cast(op.getInputs()[0].getType()); + auto elementType = inputType.getElementType(); + if (!isReductionOpAndTypeSupportedByTarget(reduceOp, elementType)) { + return rewriter.notifyMatchFailure( + op, "Unsupported reduction operation or type."); + } + + auto dims = op.getDimensions(); + if (dims.size() != 1) + return rewriter.notifyMatchFailure(op, "Only support one dim reduce."); + + auto dim = dims[0]; + auto loc = op->getLoc(); + + // Don't do channel norm for output since we always canonicalize init to + // identity value + auto attr = getRedBaseAttr(rewriter, reduceOp, elementType); + // WORKAROUND: si to fp will generate uninitialized init value, also thought + // as base identity value + if (!(checkReductionBaseAttr(op, rewriter, attr) || + op.getInits().back().getDefiningOp())) + return rewriter.notifyMatchFailure( + op, "Init is not reduction op base value."); + + auto inputShape = inputType.getShape(); + bool lastDimReduce = (dim == inputShape.size() - 1); + // Reshape input to 4D + SmallVector inputShape4D = reshapeReduceShapeTo4d(inputShape, dim); + Value input4D = createReshapeOp(rewriter, loc, op.getInputs()[0], + inputShape4D, elementType); + Value channelNormedInput = + channelNorm(rewriter, loc, input4D, inputShape4D); + + Value reduce = + convertToMKReduce(rewriter, loc, reduceOp, channelNormedInput, + inputShape4D, elementType, lastDimReduce); + + auto dechannelNormOutput = + dechannelNorm(rewriter, loc, reduce, inputShape4D, lastDimReduce); + + auto resultType = cast(op->getResultTypes()[0]); + if (resultType.getRank() == 0) { + rewriter.replaceOpWithNewOp( + op, resultType, dechannelNormOutput, + ArrayRef{}); + return success(); + } + + Value result = createReshapeOp(rewriter, loc, dechannelNormOutput, + resultType.getShape(), elementType); + + rewriter.replaceOp(op, result); + // Implement layout transformation for ReduceOp here + return success(); + } +}; + +template +struct ScalarGlobalLoadRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(ScalarGlobalLoadOp op, + PatternRewriter &rewriter) const override { + // TODO: Implement memory space check + if (triton::isOperandMemorySpaceSPM(op.getMemref())) + return rewriter.notifyMatchFailure(op, "Not global memory load"); + + auto indices = op.getIndices(); + if (indices.size() > 1) { + return rewriter.notifyMatchFailure(op, "Load has multiple indices"); + } + + auto zero = rewriter.create(op.getLoc(), 0); + auto index = indices.empty() ? zero : indices[0]; + + auto loc = op.getLoc(); + auto type = cast(op.getMemref().getType()); + auto src = rewriter.create( + loc, op.getMemref(), SmallVector{index}, + SmallVector{rewriter.getIndexAttr(1)}, + SmallVector{rewriter.getIndexAttr(1)}); + + auto tensorType = RankedTensorType::get({1}, type.getElementType()); + Value tensor = rewriter.create( + loc, tensorType, src, true /* restrict */, true /* writable */); + auto empty = rewriter.create(loc, tensorType.getShape(), + tensorType.getElementType()); + // FIXME: Bufferization pass insert copy op according address space? + auto copyOp = rewriter.create( + loc, TypeRange{tensorType}, ValueRange{tensor}, ValueRange{empty}); + + rewriter.replaceOpWithNewOp(op, copyOp.getResult(0), + ValueRange{zero}); + + return success(); + } +}; + +template +struct ScalarGlobalStoreRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(ScalarGlobalStoreOp op, + PatternRewriter &rewriter) const override { + // TODO: Implement memory space check + if (triton::isOperandMemorySpaceSPM(op.getMemref())) + return failure(); + + auto indices = op.getIndices(); + if (indices.size() > 1) { + return rewriter.notifyMatchFailure(op, "StoreOp has multiple indices"); + } + + auto zero = rewriter.create(op.getLoc(), 0); + auto index = indices.empty() ? zero : indices[0]; + + auto dst = rewriter.create( + op.getLoc(), op.getMemref(), SmallVector{index}, + SmallVector{rewriter.getIndexAttr(1)}, + SmallVector{rewriter.getIndexAttr(1)}); + + auto val = op.getValue(); + if (!isa(val.getType())) { + // NOTE: tensor::FromElementsOp will optimize to arith::ConstantOp which + // has dense constant attribute. + auto empty = rewriter.create( + op.getLoc(), SmallVector{1}, val.getType()); + + val = rewriter.create(op.getLoc(), val, empty, + ValueRange{zero}); + } + auto storeOp = + rewriter.replaceOpWithNewOp( + op, val, dst); + storeOp.setWritable(true); + + return success(); + } +}; + +Operation *findOutmostLoopOp(scf::ForOp forOp) { + + if (!forOp->hasOneUse()) + return forOp; + auto user = *forOp->getUsers().begin(); + + if (!isa(user)) + return forOp; + auto parentOp = user->getParentOp(); + + if (isa(parentOp)) + return findOutmostLoopOp(cast(parentOp)); + else + return forOp; +} + +Operation *findLoopInsertPoint(Operation *atomicOp) { + if (!atomicOp->hasOneUse()) + return atomicOp; + + auto user = *atomicOp->getUsers().begin(); + + if (!isa(user) && !user->hasOneUse()) + return atomicOp; + user = *user->getUsers().begin(); + + // Whether has mask + if (isa(user)) { + auto parentOp = user->getParentOp(); + if (!isa(parentOp) || !parentOp->hasOneUse()) + return atomicOp; + user = *parentOp->getUsers().begin(); + } + + if (!isa(user) && !user->hasOneUse()) + return atomicOp; + user = *user->getUsers().begin(); + + if (!isa(user)) + return atomicOp; + auto parentOp = user->getParentOp(); + + if (!isa(parentOp)) { + return atomicOp; + } + + return findOutmostLoopOp(cast(parentOp)); +} + +struct AtomicRMWOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + Value createAtomicArithOp(OpBuilder &rewriter, Location loc, Value oldData, + Value val, RMWOp rmwOp) const { + + auto inputType = cast(val.getType()); + auto elementType = inputType.getElementType(); + + switch (rmwOp) { + case RMWOp::AND: + return buildLinalgElementwise(rewriter, loc, + ValueRange{val, oldData}); + case RMWOp::OR: + return buildLinalgElementwise(rewriter, loc, + ValueRange{val, oldData}); + case RMWOp::XOR: + return buildLinalgElementwise(rewriter, loc, + ValueRange{val, oldData}); + case RMWOp::ADD: + return buildLinalgElementwise(rewriter, loc, + ValueRange{val, oldData}); + case RMWOp::FADD: + return buildLinalgElementwise(rewriter, loc, + ValueRange{val, oldData}); + case RMWOp::MAX: + return elementType.isIntOrIndex() + ? buildLinalgElementwise( + rewriter, loc, ValueRange{val, oldData}) + : buildLinalgElementwise( + rewriter, loc, ValueRange{val, oldData}); + case RMWOp::UMAX: + return buildLinalgElementwise(rewriter, loc, + ValueRange{val, oldData}); + case RMWOp::MIN: + return elementType.isIntOrIndex() + ? buildLinalgElementwise( + rewriter, loc, ValueRange{val, oldData}) + : buildLinalgElementwise( + rewriter, loc, ValueRange{val, oldData}); + + case RMWOp::UMIN: + return buildLinalgElementwise(rewriter, loc, + ValueRange{val, oldData}); + case RMWOp::XCHG: + return val; + default: + llvm_unreachable("Unexpected atomic op"); + } + } + + LogicalResult matchAndRewrite(mk::AtomicRMWOp op, + PatternRewriter &rewriter) const override { + auto val = op.getVal(); + auto inputType = dyn_cast(val.getType()); + if (!inputType) + return rewriter.notifyMatchFailure(op, "expected ranked tensor type"); + + auto loc = op.getLoc(); + auto ptr = op.getPtr(); + + auto loopInsertPoint = findLoopInsertPoint(op); + auto insertionPoint = rewriter.saveInsertionPoint(); + rewriter.setInsertionPoint(loopInsertPoint); + rewriter.create(loc); + rewriter.restoreInsertionPoint(insertionPoint); + auto toTensorOp = rewriter.create( + loc, inputType, ptr, true /* restrict */, true /* writable */); + // Read oldData + auto oldData = rewriter.create(loc, inputType.getShape(), + inputType.getElementType()); + auto ddrToSpm = + rewriter + .create(loc, TypeRange{inputType}, + ValueRange{toTensorOp}, ValueRange{oldData}) + ->getResult(0); + + auto newData = + createAtomicArithOp(rewriter, loc, ddrToSpm, val, op.getAtomicRmwOp()); + + auto spmToDDR = rewriter.create( + loc, newData, ptr); + spmToDDR.setWritable(true); + + rewriter.setInsertionPointAfter(loopInsertPoint); + rewriter.create(loc); + rewriter.restoreInsertionPoint(insertionPoint); + + rewriter.replaceOp(op, ddrToSpm); + + return success(); + } +}; + +struct AtomicCASOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(mk::AtomicCASOp op, + PatternRewriter &rewriter) const override { + + auto val = op.getVal(); + auto inputType = dyn_cast(val.getType()); + if (!inputType) + return rewriter.notifyMatchFailure(op, "expected ranked tensor type"); + + auto loc = op.getLoc(); + auto ptr = op.getPtr(); + + auto loopInsertPoint = findLoopInsertPoint(op); + auto insertionPoint = rewriter.saveInsertionPoint(); + rewriter.setInsertionPoint(loopInsertPoint); + rewriter.create(loc); + rewriter.restoreInsertionPoint(insertionPoint); + auto toTensorOp = rewriter.create( + loc, inputType, ptr, true /* restrict */, true /* writable */); + auto empty = rewriter.create(loc, inputType.getShape(), + inputType.getElementType()); + + // Read oldData + auto ddrToSpm = + rewriter + .create(loc, TypeRange{inputType}, + ValueRange{toTensorOp}, ValueRange{empty}) + ->getResult(0); + + // compare $cmp with data $old at location $ptr, + Value cmp = op.getCmp(); + + int rank = inputType.getRank(); + auto conditionType = + RankedTensorType::get(inputType.getShape(), rewriter.getI1Type()); + auto i1Empty = rewriter.create( + loc, conditionType.getShape(), conditionType.getElementType()); + SmallVector binaryIndexingMaps( + 3, rewriter.getMultiDimIdentityMap(rank)); + + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + Value condition = + inputType.getElementType().isIntOrIndex() + ? rewriter + .create( + loc, TypeRange{conditionType}, ValueRange{ddrToSpm, cmp}, + ValueRange{i1Empty}, binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value result = b.create( + loc, arith::CmpIPredicate::eq, args[0], args[1]); + b.create(loc, result); + }) + ->getResult(0) + : rewriter + .create( + loc, TypeRange{conditionType}, ValueRange{ddrToSpm, cmp}, + ValueRange{i1Empty}, binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value val = b.create( + loc, arith::CmpFPredicate::OEQ, args[0], args[1]); + b.create(loc, val); + }) + ->getResult(0); + + SmallVector selectOpIndexingMaps( + 4, rewriter.getMultiDimIdentityMap(rank)); + + auto newData = + rewriter + .create( + loc, TypeRange{inputType}, ValueRange{condition, val, ddrToSpm}, + ValueRange{empty}, selectOpIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value result = + b.create(loc, args[0], args[1], args[2]); + b.create(loc, result); + }) + ->getResult(0); + + auto spmToDDR = rewriter.create( + loc, newData, ptr); + spmToDDR.setWritable(true); + + rewriter.setInsertionPointAfter(loopInsertPoint); + rewriter.create(loc); + rewriter.restoreInsertionPoint(insertionPoint); + + rewriter.replaceOp(op, ddrToSpm); + return success(); + } +}; + +struct BroadcastOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult handleLastDimensionBroadcast(linalg::GenericOp op, + PatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto inputs = op.getInputs(); + auto inputType = cast(inputs[0].getType()); + auto outputType = cast(op.getOutputs()[0].getType()); + int rank = inputType.getRank(); + + // Build loop bounds for all dimensions except the last one + SmallVector shapeWithoutLast(inputType.getShape()); + shapeWithoutLast.pop_back(); + auto zero = rewriter.create(loc, 0); + auto one = rewriter.create(loc, 1); + SmallVector lbs, ubs, steps; + for (int64_t size : shapeWithoutLast) { + lbs.push_back(zero); + ubs.push_back(rewriter.create(loc, size)); + steps.push_back(one); + } + + auto loopNest = scf::buildLoopNest( + rewriter, loc, lbs, ubs, steps, op.getOutputs(), + [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange indices, + ValueRange iterArgs) { + SmallVector regionInputs(op.getNumDpsInputs()); + SmallVector inputIndices = indices; + inputIndices.push_back(zero); + std::transform(inputs.begin(), inputs.end(), regionInputs.begin(), + [&](auto val) { + return rewriter.create( + loc, val, inputIndices); + }); + + SmallVector outputOffsets = indices; + outputOffsets.push_back(rewriter.getIndexAttr(0)); + SmallVector outputSizes(rank - 1, + rewriter.getIndexAttr(1)); + outputSizes.push_back( + rewriter.getIndexAttr(outputType.getDimSize(rank - 1))); + SmallVector outputStrides(rank, + rewriter.getIndexAttr(1)); + + Value outSlice = nestedBuilder.create( + nestedLoc, iterArgs[0], outputOffsets, outputSizes, + outputStrides); + + Value filled = + nestedBuilder + .create(nestedLoc, outSlice.getType(), + regionInputs, ValueRange{outSlice}) + ->getResults()[0]; + + auto outValTensor = nestedBuilder.create( + loc, filled, iterArgs[0], outputOffsets, outputSizes, + outputStrides); + return SmallVector{outValTensor}; + }); + + rewriter.replaceOp(op, loopNest.results); + return success(); + } + + std::tuple, SmallVector, SmallVector> + initializeSliceParams(RankedTensorType outputType, int rank) const { + SmallVector sliceOffsets(rank, 0); + SmallVector sliceSizes(outputType.getShape().begin(), + outputType.getShape().end()); + SmallVector sliceStrides(rank, 1); + return {sliceOffsets, sliceSizes, sliceStrides}; + } + + memref::AllocOp createOutputMemrefAndInitCopy(linalg::GenericOp op, + PatternRewriter &rewriter, + int64_t broadcastDim) const { + Location loc = op->getLoc(); + auto inputs = op.getInputs(); + auto inputType = cast(inputs[0].getType()); + auto outputType = cast(op.getOutputs()[0].getType()); + int rank = outputType.getRank(); + + auto [sliceOffsets, sliceSizes, sliceStrides] = + initializeSliceParams(outputType, rank); + sliceSizes[broadcastDim] = + 1; // Single element slice along broadcast dimension + + // WORKAROUND: For broadcast operations with multiple users, we need to + // explicitly allocate a new memref to avoid read-after-write conflicts in + // the bufferized representation. The one-shot-bufferize pass cannot + // automatically handle these conflicts for memref types in broadcast + // patterns, so we do it manually here. + // TODO : In the case where there are multiple users or where there is only + // one user, we can optimize memory usage by using buffer reuse analysis to + // reuse memory from output buffers. + auto outputMemref = rewriter.create( + loc, + MemRefType::get(outputType.getShape(), outputType.getElementType())); + auto inputMemref = rewriter.create( + loc, MemRefType::get(inputType.getShape(), inputType.getElementType()), + inputs[0]); + auto initSlice = + rewriter.create(loc, outputMemref, sliceOffsets, + /*sizes=*/sliceSizes, + /*strides=*/sliceStrides); + rewriter.create(loc, inputMemref, initSlice); + return outputMemref; + } + + Value copyWithTilingStrategy(linalg::GenericOp op, PatternRewriter &rewriter, + memref::AllocOp outputMemref, + int64_t broadcastDim) const { + Location loc = op->getLoc(); + auto outputType = cast(op.getOutputs()[0].getType()); + int rank = outputType.getRank(); + auto [sliceOffsets, sliceSizes, sliceStrides] = + initializeSliceParams(outputType, rank); + + int64_t broadcastDimSize = outputType.getShape()[broadcastDim]; + sliceSizes[broadcastDim] = 1; + + // WORKAROUND: Since one-shot-bufferize will insert extra copy for + // extract_slice and insert_slice even they are different slices. Here we + // directly use memref copy to avoid extra copy. + // Copy (1, 2, 4, 8, 16, 32) input slices to output + constexpr int64_t kTileSizes[] = {1, 2, 4, 8, 16, 32}; + + SmallVector currentOffsets = sliceOffsets; + SmallVector currentSizes = sliceSizes; + + for (int64_t tileSize : kTileSizes) { + if (tileSize >= broadcastDimSize) + break; + + currentSizes[broadcastDim] = tileSize; + auto sourceSlice = rewriter.create( + loc, outputMemref, sliceOffsets, currentSizes, sliceStrides); + + currentOffsets[broadcastDim] = tileSize; + auto destSlice = rewriter.create( + loc, outputMemref, currentOffsets, currentSizes, sliceStrides); + + rewriter.create(loc, sourceSlice, destSlice); + } + return rewriter.create( + loc, outputType, outputMemref, /*allow_memref_to_tensor=*/true, + /*allow_tensor_to_memref=*/true); + } + + LogicalResult handleLargeBroadcastCase(linalg::GenericOp op, + PatternRewriter &rewriter, + Value sourceTensor, + int64_t broadcastDim) const { + Location loc = op->getLoc(); + auto outputType = cast(op.getOutputs()[0].getType()); + int rank = outputType.getRank(); + int64_t broadcastDimSize = outputType.getShape()[broadcastDim]; + auto [sliceOffsets, sliceSizes, sliceStrides] = + initializeSliceParams(outputType, rank); + + SmallVector largeSliceShape = sliceSizes; + largeSliceShape[broadcastDim] = kMaxSliceSize; + auto largeSlice = rewriter.create( + loc, + RankedTensorType::get(largeSliceShape, outputType.getElementType()), + sourceTensor, ValueRange(), ValueRange(), ValueRange(), sliceOffsets, + largeSliceShape, sliceStrides); + + Value lowerBound = + rewriter.create(loc, kMaxSliceSize); + Value upperBound = + rewriter.create(loc, broadcastDimSize); + Value step = rewriter.create(loc, kMaxSliceSize); + + auto forOp = rewriter.create( + loc, lowerBound, upperBound, step, ValueRange{sourceTensor}, + [&](OpBuilder &nestedBuilder, Location nestedLoc, Value iv, + ValueRange iterArgs) { + SmallVector outputOffsets(rank, + rewriter.getIndexAttr(0)); + outputOffsets[broadcastDim] = iv; + SmallVector outputSizes; + for (auto s : outputType.getShape()) + outputSizes.push_back(rewriter.getIndexAttr(s)); + outputSizes[broadcastDim] = rewriter.getIndexAttr(kMaxSliceSize); + SmallVector outputStrides(rank, + rewriter.getIndexAttr(1)); + auto outputSlice = nestedBuilder.create( + nestedLoc, largeSlice, iterArgs[0], outputOffsets, outputSizes, + outputStrides); + nestedBuilder.create(nestedLoc, + outputSlice.getResult()); + }); + + rewriter.replaceOp(op, forOp.getResult(0)); + return success(); + } + + LogicalResult handleOtherDimensionBroadcast(linalg::GenericOp op, + PatternRewriter &rewriter, + int64_t broadcastDim) const { + Location loc = op->getLoc(); + auto inputs = op.getInputs(); + auto inputType = cast(inputs[0].getType()); + auto outputType = cast(op.getOutputs()[0].getType()); + int rank = inputType.getRank(); + int64_t broadcastDimSize = outputType.getShape()[broadcastDim]; + + auto outputMemref = + createOutputMemrefAndInitCopy(op, rewriter, broadcastDim); + auto resultTensor = + copyWithTilingStrategy(op, rewriter, outputMemref, broadcastDim); + + if (broadcastDimSize <= kMaxSliceSize) { + rewriter.replaceAllOpUsesWith(op, resultTensor); + return success(); + } + + return handleLargeBroadcastCase(op, rewriter, resultTensor, broadcastDim); + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (!regionOps.empty() || !op->hasAttr("broadcastDims")) + return failure(); + + auto dims = cast(op->getAttr("broadcastDims")); + auto inputs = op.getInputs(); + auto inputType = cast(inputs[0].getType()); + auto rank = inputType.getRank(); + if (dims.size() != 1) + return failure(); + + if (dims[0] == rank - 1) { + return handleLastDimensionBroadcast(op, rewriter); + } else { + return handleOtherDimensionBroadcast(op, rewriter, dims[0]); + } + } + +private: + constexpr static int64_t kMaxSliceSize = 64; +}; + +struct DivFloatOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult convertDivFloatOp(linalg::GenericOp op, + PatternRewriter &rewriter) const { + + Location loc = op->getLoc(); + + // Read rnd_mode attribute from the original DivFOp + auto regionOps = getRegionOps(op); + auto divOp = cast(regionOps[0]); + auto rndModeAttr = divOp->getAttr("rnd_mode"); + + auto inputTensorType = + dyn_cast(op.getInputs()[0].getType()); + auto outputTensorType = + dyn_cast(op.getOutputs()[0].getType()); + + // Regular (tensor) path: out = lhs / rhs -> recip(rhs) * lhs. + if (inputTensorType && outputTensorType) { + auto rank = inputTensorType.getRank(); + auto empty = rewriter.create( + loc, inputTensorType.getShape(), inputTensorType.getElementType()); + + Value recip = rewriter + .create( + loc, inputTensorType, ValueRange{op.getInputs()[1]}, + ValueRange{empty}) + ->getResult(0); + + SmallVector binaryIndexingMaps( + 3, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + rewriter.replaceOpWithNewOp( + op, outputTensorType, ValueRange{op.getInputs()[0], recip}, + ValueRange{empty}, binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + auto mulOp = b.create(loc, args[0], args[1]); + if (rndModeAttr) + mulOp->setAttr("rnd_mode", rndModeAttr); + b.create(loc, mulOp.getResult()); + }); + + return success(); + } + + // DSA memref path (mode 0/1): out = lhs / rhs + // -> scratch = recip(rhs); out = mul(lhs, scratch). + // A scratch buffer keeps the result correct even when out aliases an + // input. It is allocated here (this pass runs before + // spmd-allocate-shared-memory) so it receives an allocation.offset attr. + auto lhsTy = dyn_cast(op.getInputs()[0].getType()); + auto rhsTy = dyn_cast(op.getInputs()[1].getType()); + auto outTy = dyn_cast(op.getOutputs()[0].getType()); + if (!lhsTy || !rhsTy || !outTy) + return failure(); + + if (lhsTy.getShape() != rhsTy.getShape() || + lhsTy.getShape() != outTy.getShape()) + return op->emitRemark("dsa binary op shape mismatch between lhs/rhs/out"); + if (lhsTy.getElementType() != rhsTy.getElementType() || + lhsTy.getElementType() != outTy.getElementType()) + return op->emitRemark( + "dsa binary op element type mismatch between lhs/rhs/out"); + + auto scratch = rewriter.create(loc, outTy); + auto scratchMemref = scratch.getResult(); + + rewriter.create(loc, TypeRange{}, + ValueRange{op.getInputs()[1]}, + ValueRange{scratchMemref}); + + auto rank = static_cast(lhsTy.getShape().size()); + auto identityMap = rewriter.getMultiDimIdentityMap(rank); + SmallVector indexingMaps = {identityMap, identityMap, + identityMap}; + SmallVector iteratorTypes( + rank, mlir::utils::IteratorType::parallel); + + rewriter.create( + loc, + /*resultTensorTypes=*/TypeRange{}, + ValueRange{op.getInputs()[0], scratchMemref}, + ValueRange{op.getOutputs()[0]}, indexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + auto mulOp = b.create(loc, args[0], args[1]); + if (rndModeAttr) + mulOp->setAttr("rnd_mode", rndModeAttr); + b.create(loc, mulOp.getResult()); + }); + + rewriter.eraseOp(op); + return success(); + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1 || !isa(regionOps[0])) + return failure(); + + return convertDivFloatOp(op, rewriter); + } +}; + +struct DivIntOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1 || + !isa(regionOps[0])) + return failure(); + + // FIXME: Canonicalize non-precision mode divint in linalg-to-mk, others + // default to scf.for + // DSA memref-operand generics have no results and non-tensor operands; + // only handle tensor-form integer division here. + if (op.getNumResults() != 1) + return failure(); + auto resultTensorType = + dyn_cast(op.getResult(0).getType()); + if (!resultTensorType) + return failure(); + SmallVector inputs(op.getInputs().begin(), op.getInputs().end()); + + SmallVector outputs = {rewriter.create( + op->getLoc(), resultTensorType.getShape(), + resultTensorType.getElementType())}; + assert(op->getResultTypes().size() == 1); + + auto scalarResultType = resultTensorType.getElementType(); + + // NOTE: linalgOpToloop function only support memref type + auto shape = resultTensorType.getShape(); + auto loc = op->getLoc(); + auto zero = rewriter.create(loc, 0); + auto one = rewriter.create(loc, 1); + SmallVector lbs, ubs, steps; + for (auto [i, size] : enumerate(shape)) { + auto sizeValue = rewriter.create(loc, size); + lbs.push_back(zero); + ubs.push_back(sizeValue); + steps.push_back(one); + } + auto loopNest = scf::buildLoopNest( + rewriter, loc, lbs, ubs, steps, outputs, + [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange indices, + ValueRange iterArgs) { + SmallVector regionInputs(op.getNumDpsInputs()); + std::transform(inputs.begin(), inputs.end(), regionInputs.begin(), + [&](auto val) { + return rewriter.create(loc, val, + indices); + }); + auto *scalarOp = nestedBuilder.create( + loc, regionOps[0]->getName().getIdentifier(), regionInputs, + scalarResultType, regionOps[0]->getAttrs()); + + auto outValTensor = nestedBuilder.create( + loc, scalarOp->getResult(0), iterArgs[0], indices); + return SmallVector{outValTensor}; + }); + + rewriter.replaceOp(op, loopNest.results); + return success(); + } +}; + +struct SelectOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult convertI1Select(linalg::GenericOp op, + PatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto mask = op.getInputs()[0]; + auto input1 = op.getInputs()[1]; + auto input2 = op.getInputs()[2]; + + auto inputType = dyn_cast(input1.getType()); + auto rank = inputType.getRank(); + + auto empty = rewriter.create(loc, inputType.getShape(), + inputType.getElementType()); + + SmallVector binaryIndexingMaps( + 3, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + // result[i] = (mask[i] AND A[i]) OR (NOT mask[i] AND B[i]) + // mask[i] AND A[i] + auto lhs = rewriter.create( + loc, inputType, ValueRange{mask, input1}, ValueRange{empty}, + binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value result = b.create(loc, args[0], args[1]); + b.create(loc, result); + }); + + // NOT mask + // NOTE: Also can define a mk.not + Value allTrue = rewriter.create( + loc, inputType.getShape(), inputType.getElementType()); + allTrue = + rewriter + .create(loc, + ValueRange{ + rewriter.create( + loc, inputType.getElementType(), + rewriter.getBoolAttr(true)), + }, + ValueRange{empty}) + ->getResult(0); + auto notMask = rewriter.create( + loc, inputType, ValueRange{mask, allTrue}, ValueRange{empty}, + binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value result = b.create(loc, args[0], args[1]); + b.create(loc, result); + }); + + // NOT mask[i] AND B[i] + auto rhs = rewriter.create( + loc, inputType, ValueRange{notMask.getResult(0), input2}, + ValueRange{empty}, binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value result = b.create(loc, args[0], args[1]); + b.create(loc, result); + }); + + // result + rewriter.replaceOpWithNewOp( + op, inputType, ValueRange{lhs.getResult(0), rhs.getResult(0)}, + ValueRange{empty}, binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value result = b.create(loc, args[0], args[1]); + b.create(loc, result); + }); + + return success(); + } + + LogicalResult convertI8Select(linalg::GenericOp op, + PatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto mask = op.getInputs()[0]; + auto input1 = op.getInputs()[1]; + auto input2 = op.getInputs()[2]; + auto inputType = dyn_cast(input1.getType()); + auto midType = + RankedTensorType::get(inputType.getShape(), rewriter.getF32Type()); + auto fpinput1 = + buildLinalgElementwise(rewriter, loc, midType, input1); + auto fpinput2 = + buildLinalgElementwise(rewriter, loc, midType, input2); + auto fpOutput = buildLinalgElementwise( + rewriter, loc, midType, {mask, fpinput1, fpinput2}); + auto i8Output = buildLinalgElementwise( + rewriter, loc, inputType, fpOutput); + rewriter.replaceAllUsesWith(op->getResults(), {i8Output}); + rewriter.eraseOp(op); + return success(); + } + + Value createBitcastOp(OpBuilder &rewriter, Location loc, Value input, + RankedTensorType targetType) const { + auto empty = rewriter.create(loc, targetType.getShape(), + targetType.getElementType()); + int rank = targetType.getRank(); + + SmallVector binaryIndexingMaps( + 2, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + + return rewriter + .create( + loc, targetType, ValueRange{input}, ValueRange{empty}, + binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value result = b.create( + loc, targetType.getElementType(), args[0]); + b.create(loc, result); + }) + ->getResult(0); + } + + LogicalResult SelectConvertOp(linalg::GenericOp op, + PatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto mask = op.getInputs()[0]; + auto input1 = op.getInputs()[1]; + auto input2 = op.getInputs()[2]; + + auto inputType = dyn_cast(input1.getType()); + auto rank = inputType.getRank(); + + auto castType = inputType; + if (inputType.getElementType().isIntOrIndex()) { + auto bitWidth = inputType.getElementTypeBitWidth(); + assert(bitWidth == 16 || bitWidth == 32); + FloatType floatType = + bitWidth == 16 ? rewriter.getF16Type() : rewriter.getF32Type(); + castType = RankedTensorType::get(inputType.getShape(), floatType); + // TODO: Bitcast integer to float + input1 = createBitcastOp(rewriter, op.getLoc(), input1, castType); + input2 = createBitcastOp(rewriter, op.getLoc(), input2, castType); + } + + auto empty = rewriter.create(loc, castType.getShape(), + castType.getElementType()); + + // Maskmove mask only support int8/fp, here mask is memref + auto maskCast = + rewriter.create(op.getLoc(), castType, mask, empty); + + // Res = input2; + Value res = rewriter + .create(loc, castType, ValueRange{input2}, + ValueRange{empty}) + ->getResult(0); + + // if input0 = 1, Res = input1; + // if input0 = 0, Res = input2; + res = rewriter + .create(loc, castType, input1, + maskCast->getResult(0), res) + ->getResult(0); + + if (inputType != castType) { + res = createBitcastOp(rewriter, op.getLoc(), res, inputType); + } + rewriter.replaceOp(op, res); + return success(); + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1 || !isa(regionOps[0])) + return failure(); + + auto inputType = dyn_cast(op.getInputs()[1].getType()); + assert(inputType && "Only support ranked tensor type"); + auto elemType = inputType.getElementType(); + if (elemType.getIntOrFloatBitWidth() == 64) + return failure(); + if (elemType.isInteger(1)) + return convertI1Select(op, rewriter); + // maskmove does not support int8, so convert to fp + if (elemType.isInteger(8)) + return convertI8Select(op, rewriter); + + return SelectConvertOp(op, rewriter); + } +}; + +struct IsInfOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult convertIsInfOp(linalg::GenericOp op, + PatternRewriter &rewriter) const { + Location loc = op->getLoc(); + + auto inputTensorType = + dyn_cast(op.getInputs()[0].getType()); + auto outputTensorType = + dyn_cast(op.getOutputs()[0].getType()); + auto empty = rewriter.create( + loc, inputTensorType.getShape(), inputTensorType.getElementType()); + // 1 / inf == 0, Use recip and boolequalvs to calculate isinf. + Value recip = + rewriter + .create(loc, inputTensorType, op.getInputs(), + ValueRange{empty}) + ->getResult(0); + + auto zero = rewriter.create( + op.getLoc(), rewriter.getZeroAttr(inputTensorType.getElementType())); + rewriter.replaceOpWithNewOp(op, outputTensorType, recip, + zero, op.getOutputs()[0]); + return success(); + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1 || !isa(regionOps[0])) + return failure(); + + return convertIsInfOp(op, rewriter); + } +}; + +struct PowFOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + Value computeOddFloatIndicator(PatternRewriter &rewriter, Value exponent, + Location loc) const { + auto resultType = cast(exponent.getType()); + auto emptyTensor = rewriter.create( + loc, resultType.getShape(), resultType.getElementType()); + + auto zeroVal = rewriter.create( + loc, rewriter.getZeroAttr(resultType.getElementType())); + auto zeroTensor = rewriter + .create(loc, ValueRange{zeroVal}, + ValueRange{emptyTensor}) + ->getResult(0); + + auto oneVal = rewriter.create( + loc, rewriter.getOneAttr(resultType.getElementType())); + auto oneTensor = rewriter + .create(loc, ValueRange{oneVal}, + ValueRange{emptyTensor}) + ->getResult(0); + + auto twoVal = rewriter.create( + loc, rewriter.getFloatAttr(resultType.getElementType(), 2.0)); + auto twoTensor = rewriter + .create(loc, ValueRange{twoVal}, + ValueRange{emptyTensor}) + ->getResult(0); + + // Calculate exponent % 2: returns 1 if odd, 0 if even + auto remainder = buildLinalgElementwise( + rewriter, loc, resultType, ValueRange{exponent, twoTensor}); + + return rewriter + .create(loc, resultType, oneTensor, remainder, + zeroTensor) + ->getResult(0); + } + + Value computeNegativeResultIndicator(PatternRewriter &rewriter, Location loc, + Value base, Value exponent, + Value &absoluteBase) const { + auto resultType = cast(base.getType()); + auto intResultType = + RankedTensorType::get(resultType.getShape(), rewriter.getI32Type()); + + auto zeroVal = rewriter.create( + loc, rewriter.getZeroAttr(resultType.getElementType())); + auto emptyTensor = rewriter.create( + loc, resultType.getShape(), resultType.getElementType()); + + // Check if base is negative (a < 0) + auto isBaseNegative = + rewriter + .create(loc, resultType, base, zeroVal, emptyTensor) + ->getResult(0); + + // Check if exponent is integer-like (truncated value equals original) + auto truncatedExponent = buildLinalgElementwise( + rewriter, loc, resultType, ValueRange{exponent}); + auto isIntegerLikeExponent = + rewriter + .create(loc, resultType, exponent, truncatedExponent, + emptyTensor) + ->getResult(0); + + // Compute absolute base for integer-like exponents : a = |a| + auto absoluteBaseValue = buildLinalgElementwise( + rewriter, loc, resultType, ValueRange{base}); + absoluteBase = + rewriter + .create(loc, resultType, absoluteBaseValue, + isIntegerLikeExponent, base) + ->getResult(0); + + // Convert boolean masks to integer for bitwise operations + auto isBaseNegativeInt = rewriter.create( + loc, intResultType, ValueRange{isBaseNegative}); + auto isIntegerLikeExponentInt = rewriter.create( + loc, intResultType, ValueRange{isIntegerLikeExponent}); + auto canTakeAbsolute = buildLinalgElementwise( + rewriter, loc, intResultType, + ValueRange{isBaseNegativeInt, isIntegerLikeExponentInt}); + + // Check if integer-like exponent is odd : b % 2 != 0 + auto isOddIndicator = + computeOddFloatIndicator(rewriter, truncatedExponent, loc); + auto isOddIndicatorInt = rewriter.create( + loc, intResultType, ValueRange{isOddIndicator}); + + // Final condition: a < 0 & b is IntergerLike & b % 2 != 0 + auto isResultNegative = buildLinalgElementwise( + rewriter, loc, intResultType, + ValueRange{canTakeAbsolute, isOddIndicatorInt}); + + return rewriter.create(loc, resultType, + ValueRange{isResultNegative}); + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1 || !isa(regionOps.front())) + return failure(); + + // Skip DSA memref-operand generics (no results / non-tensor operands). + if (op->getResultTypes().empty() || + !isa(op->getResultTypes()[0])) + return failure(); + + auto base = op.getInputs()[0]; + auto exponent = op.getInputs()[1]; + auto loc = op->getLoc(); + auto resultType = cast(op->getResultTypes()[0]); + + auto zeroVal = rewriter.create( + loc, rewriter.getZeroAttr(resultType.getElementType())); + auto emptyTensor = rewriter.create( + loc, resultType.getShape(), resultType.getElementType()); + + Value absoluteBase; + // Determine if result should be negative: + // (a <0 && b isIntegerLike && b % 2!= 0) + auto isResultNegative = computeNegativeResultIndicator( + rewriter, loc, base, exponent, absoluteBase); + + // Compute power using identity: a^b = 2^(b * log2(|a|)) + auto log2Base = buildLinalgElementwise( + rewriter, loc, resultType, ValueRange{absoluteBase}); + auto exponentTimesLog = buildLinalgElementwise( + rewriter, loc, resultType, ValueRange{log2Base, exponent}); + auto powerResult = buildLinalgElementwise( + rewriter, loc, resultType, ValueRange{exponentTimesLog}); + + // Apply negative sign if isResultNegative: a ^ b => -2 ^ (b * log2(|a|)) + auto negativePowerResult = buildLinalgElementwise( + rewriter, loc, resultType, ValueRange{powerResult}); + auto signedResult = + rewriter + .create(loc, resultType, negativePowerResult, + isResultNegative, powerResult) + ->getResult(0); + + // Handle special case: if exponent = 0 , a ^ b = 1 + auto isZeroExponent = rewriter + .create(loc, resultType, exponent, + zeroVal, emptyTensor) + ->getResult(0); + auto oneVal = rewriter.create( + loc, rewriter.getOneAttr(resultType.getElementType())); + auto oneTensor = rewriter + .create(loc, ValueRange{oneVal}, + ValueRange{emptyTensor}) + ->getResult(0); + + auto finalResult = rewriter.create( + loc, resultType, oneTensor, isZeroExponent, signedResult); + + rewriter.replaceOp(op, finalResult); + return success(); + } +}; + +struct MinMaxOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + template + LogicalResult convertMinMaxOp(linalg::GenericOp op, + PatternRewriter &rewriter) const { + Location loc = op->getLoc(); + + auto lhs = op.getInputs()[0]; + auto rhs = op.getInputs()[1]; + + auto inputType = dyn_cast(lhs.getType()); + if (!inputType) + return failure(); + auto rank = inputType.getRank(); + + auto inputTypeEmpty = rewriter.create( + loc, inputType.getShape(), inputType.getElementType()); + + SmallVector binaryIndexingMaps( + 3, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + + // auto isANan = UnEqualVV(lhs, lhs) + auto isANan = rewriter.create(loc, inputType, lhs, lhs, + inputTypeEmpty); + + // auto result = lhs + Value res = rewriter + .create(loc, inputType, ValueRange{lhs}, + ValueRange{inputTypeEmpty}) + ->getResult(0); + + // result = maskmove(isANan, rhs) + res = rewriter + .create(loc, inputType, rhs, isANan->getResult(0), + res) + ->getResult(0); + + // auto isBNan = UnEqualVV(rhs, rhs) + auto isBNan = + rewriter + .create(loc, inputType, rhs, rhs, inputTypeEmpty) + ->getResult(0); + + // auto shouldApplyResult = EqualVS(isBNan, 0) + auto constValue = rewriter.create( + op.getLoc(), rewriter.getZeroAttr(inputType.getElementType())); + auto shouldApplyResult = rewriter.create( + loc, inputType, isBNan, constValue, inputTypeEmpty); + + // auto minMaxValue = maxvv/minvv(result, rhs) + auto minMaxValue = rewriter.create( + loc, inputType, ValueRange{res, rhs}, ValueRange{inputTypeEmpty}, + binaryIndexingMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value result = b.create(loc, args[0], args[1]); + b.create(loc, result); + }); + + // result = maskmove(shouldApplyResult, minMaxValue) + res = rewriter + .create(loc, inputType, minMaxValue->getResult(0), + shouldApplyResult->getResult(0), res) + .getResult(0); + + rewriter.replaceOp(op, res); + + return success(); + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1 || + !isa(regionOps[0])) + return failure(); + + if (isa(regionOps[0])) + return convertMinMaxOp(op, rewriter); + else + return convertMinMaxOp(op, rewriter); + } +}; + +template +Value createElemwiseNaryOp(OpBuilder &builder, Location loc, ValueRange inputs, + Value output) { + auto outputTy = cast(output.getType()); + auto rank = outputTy.getRank(); + if (rank == 0) { + SmallVector loadVals; + llvm::transform(inputs, std::back_inserter(loadVals), [&](Value input) { + return builder.create(loc, input, ValueRange{}); + }); + auto val = builder.create(loc, outputTy.getElementType(), loadVals); + return builder.create(loc, outputTy, + val.getResult()); + } else { + SmallVector idMaps(2, builder.getMultiDimIdentityMap(rank)); + SmallVector iterators( + rank, mlir::utils::IteratorType::parallel); + return builder + .create( + loc, outputTy, inputs, ValueRange{output}, idMaps, iterators, + [](OpBuilder &b, Location loc, ValueRange args) { + Value val = b.create(loc, args.back().getType(), + args.drop_back()); + b.create(loc, val); + }) + .getResult(0); + } +} + +static LogicalResult convertSIOpToF32Op( + Operation *srcOp, PatternRewriter &rewriter, ValueRange inputs, + ValueRange outputs, + std::function + fpOpBuildFn, + bool convertOutputs = false) { + Location loc = srcOp->getLoc(); + SmallVector fpInputs, fpOutputs, intResults; + // Convert integer input + for (auto input : inputs) { + auto inputTy = cast(input.getType()); + Value fpInput = rewriter.create(loc, inputTy.getShape(), + rewriter.getF32Type()); + fpInputs.push_back( + createElemwiseNaryOp(rewriter, loc, input, fpInput)); + } + + for (auto output : outputs) { + + auto outputTy = cast(output.getType()); + Value fpOutput = rewriter.create(loc, outputTy.getShape(), + rewriter.getF32Type()); + + if (convertOutputs) { + // Reduce path: convert the init value from int to fp32 instead of + // discarding it. Elementwise callers (default convertOutputs=false) + // still pass EmptyOp since their outputs are pure output buffers. + fpOutputs.push_back(createElemwiseNaryOp( + rewriter, loc, output, fpOutput)); + } else { + fpOutputs.push_back(fpOutput); + } + } + + auto fpResults = fpOpBuildFn(srcOp, rewriter, fpInputs, fpOutputs); + auto resultTy = cast(srcOp->getResultTypes()[0]); + for (auto fpResult : fpResults) { + + Value intResult = rewriter.create( + loc, resultTy.getShape(), resultTy.getElementType()); + intResults.push_back(createElemwiseNaryOp( + rewriter, loc, fpResult, intResult)); + } + rewriter.replaceOp(srcOp, intResults); + return success(); +} + +// Unsigned counterpart of convertSIOpToF32Op: uses UIToFP/FPToUI so that +// unsigned integers round-trip through f32 without sign-extension artifacts. +static LogicalResult convertUIOpToF32Op( + Operation *srcOp, PatternRewriter &rewriter, ValueRange inputs, + ValueRange outputs, + std::function + fpOpBuildFn) { + Location loc = srcOp->getLoc(); + SmallVector fpInputs, fpOutputs, intResults; + for (auto input : inputs) { + auto inputTy = cast(input.getType()); + Value fpInput = rewriter.create(loc, inputTy.getShape(), + rewriter.getF32Type()); + fpInputs.push_back( + createElemwiseNaryOp(rewriter, loc, input, fpInput)); + } + + for (auto output : outputs) { + auto outputTy = cast(output.getType()); + Value fpOutput = rewriter.create(loc, outputTy.getShape(), + rewriter.getF32Type()); + fpOutputs.push_back(fpOutput); + } + + auto fpResults = fpOpBuildFn(srcOp, rewriter, fpInputs, fpOutputs); + auto resultTy = cast(srcOp->getResultTypes()[0]); + for (auto fpResult : fpResults) { + Value intResult = rewriter.create( + loc, resultTy.getShape(), resultTy.getElementType()); + intResults.push_back(createElemwiseNaryOp( + rewriter, loc, fpResult, intResult)); + } + rewriter.replaceOp(srcOp, intResults); + return success(); +} + +// Build a linalg.generic wrapping arith.cmpf with the given predicate. +// Result is an i1 tensor with the same shape as the inputs. +static Value buildLinalgCmpF(OpBuilder &rewriter, Location loc, + arith::CmpFPredicate pred, Value lhs, Value rhs) { + auto inputTy = cast(lhs.getType()); + auto i1Ty = RankedTensorType::get(inputTy.getShape(), rewriter.getI1Type()); + auto rank = inputTy.getRank(); + auto idMap = AffineMap::getMultiDimIdentityMap(rank, rewriter.getContext()); + SmallVector maps(3, idMap); + SmallVector iters(rank, utils::IteratorType::parallel); + auto out = rewriter.create(loc, i1Ty.getShape(), + i1Ty.getElementType()); + return rewriter + .create( + loc, i1Ty, ValueRange{lhs, rhs}, ValueRange{out}, maps, iters, + [&](OpBuilder &b, Location l, ValueRange args) { + Value c = b.create(l, pred, args[0], args[1]); + b.create(l, c); + }) + .getResult(0); +} + +// Build a linalg.generic wrapping arith.select (elementwise, 3 inputs). +static Value buildLinalgSelect(OpBuilder &rewriter, Location loc, Value cond, + Value trueV, Value falseV) { + auto resTy = cast(trueV.getType()); + auto rank = resTy.getRank(); + auto idMap = AffineMap::getMultiDimIdentityMap(rank, rewriter.getContext()); + SmallVector maps(4, idMap); + SmallVector iters(rank, utils::IteratorType::parallel); + auto out = rewriter.create(loc, resTy.getShape(), + resTy.getElementType()); + return rewriter + .create( + loc, resTy, ValueRange{cond, trueV, falseV}, ValueRange{out}, maps, + iters, + [&](OpBuilder &b, Location l, ValueRange args) { + Value s = b.create(l, args[0], args[1], args[2]); + b.create(l, s); + }) + .getResult(0); +} + +// Build the corrective integer-division quotient in the positive f32 domain. +// q = trunc(a / b) ; may be (true_q - 1) when RECIP rounds low +// r = a - q*b +// q_out = q + (r >= b ? 1 : 0) +// `aAbs`/`bAbs` must be non-negative f32 tensors. All intermediate values +// must stay < 2^24 (the FP32 exact-representation bound) for the remainder +// check to be precise. Callers must therefore restrict this path to operands +// known to fit (gated by precision mode at the use site). +// Returns the corrected (still non-negative) quotient as an f32 tensor. +static Value buildCorrectivePosDiv(OpBuilder &rewriter, Location loc, + Value aAbs, Value bAbs) { + auto f32Ty = cast(aAbs.getType()); + // Use explicit recip+mul instead of divf so the integer path is clearly + // separated from user float divisions (which go through NRM_DIV). + Value recipOut = rewriter.create(loc, f32Ty.getShape(), + f32Ty.getElementType()); + Value recip = rewriter + .create(loc, f32Ty, ValueRange{bAbs}, + ValueRange{recipOut}) + ->getResult(0); + Value qf = + buildLinalgElementwise(rewriter, loc, {aAbs, recip}); + Value qTrunc = buildLinalgElementwise(rewriter, loc, {qf}); + Value chk = + buildLinalgElementwise(rewriter, loc, {qTrunc, bAbs}); + Value r = buildLinalgElementwise(rewriter, loc, {aAbs, chk}); + Value needsCorr = + buildLinalgCmpF(rewriter, loc, arith::CmpFPredicate::OGE, r, bAbs); + // i1 -> f32 (1.0/0.0). Use createElemwiseNaryOp (no Elementwise-trait + // requirement) since arith cast ops aren't guaranteed to satisfy the + // buildLinalgElementwise static_assert. + Value corrEmpty = rewriter.create(loc, f32Ty.getShape(), + f32Ty.getElementType()); + Value corr = createElemwiseNaryOp(rewriter, loc, needsCorr, + corrEmpty); + return buildLinalgElementwise(rewriter, loc, {qTrunc, corr}); +} + +// Signed corrective integer division, using explicit recip+mul for the +// initial quotient estimate (not divf, which would conflate with the user +// float-division -> NRM_DIV path). trunc-toward-zero semantics, matching +// arith.divsi / C. Same FP32 exact-representation bound as above. +// qf = a * recip(b) ; |qf| may be (|true_q| - 1) +// q = trunc(qf) +// r = a - q*b +// q += (|r| >= |b|) ? sign(qf) : 0 +static Value buildCorrectiveDivSigned(OpBuilder &rewriter, Location loc, + Value aF, Value bF) { + auto f32Ty = cast(aF.getType()); + // Use explicit recip+mul instead of divf to keep the integer path + // independent from the user-float-divf -> NRM_DIV path. + Value recipOut = rewriter.create(loc, f32Ty.getShape(), + f32Ty.getElementType()); + Value recip = rewriter + .create(loc, f32Ty, ValueRange{bF}, + ValueRange{recipOut}) + ->getResult(0); + Value qf = buildLinalgElementwise(rewriter, loc, {aF, recip}); + Value qTrunc = buildLinalgElementwise(rewriter, loc, {qf}); + Value chk = + buildLinalgElementwise(rewriter, loc, {qTrunc, bF}); + Value r = buildLinalgElementwise(rewriter, loc, {aF, chk}); + Value rAbs = buildLinalgElementwise(rewriter, loc, {r}); + Value bAbs = buildLinalgElementwise(rewriter, loc, {bF}); + Value needsCorr = + buildLinalgCmpF(rewriter, loc, arith::CmpFPredicate::OGE, rAbs, bAbs); + + // Splat tensor constant via DenseElementsAttr (avoids a separate FillOp). + // The rest of the file uses scalar Constant + FillOp; DenseConstantToFill + // canonicalizes both forms to the same IR, so either is fine. + Value zero = rewriter.create( + loc, DenseElementsAttr::get(f32Ty, rewriter.getF32FloatAttr(0.0f))); + Value qfNeg = + buildLinalgCmpF(rewriter, loc, arith::CmpFPredicate::OLT, qf, zero); + Value corrEmpty = rewriter.create(loc, f32Ty.getShape(), + f32Ty.getElementType()); + Value mag = createElemwiseNaryOp(rewriter, loc, needsCorr, + corrEmpty); + Value magNeg = buildLinalgElementwise(rewriter, loc, {mag}); + Value dir = buildLinalgSelect(rewriter, loc, qfNeg, magNeg, mag); + return buildLinalgElementwise(rewriter, loc, {qTrunc, dir}); +} + +struct CannonicalizeRedudantTypeConversion + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + bool isValidSIToFPOp(linalg::GenericOp op, PatternRewriter &rewriter) const { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1) + return false; + if (!dyn_cast(regionOps.front())) + return false; + + auto inputType = cast(op->getOperandTypes()[0]); + auto outputType = cast(op->getResultTypes()[0]); + + return inputType.getElementType() == rewriter.getI64Type() && + outputType.getElementType() == rewriter.getF32Type(); + } + + bool isValidFPToSIOp(linalg::GenericOp op, PatternRewriter &rewriter) const { + auto regionOps = getRegionOps(op); + if (regionOps.empty()) + return false; + if (!dyn_cast(regionOps[0])) + return false; + + auto inputType = cast(op->getOperandTypes()[0]); + auto outputType = cast(op->getResultTypes()[0]); + + return inputType.getElementType() == rewriter.getF32Type() && + outputType.getElementType() == rewriter.getI64Type(); + } + + bool isBroadcastOp(linalg::GenericOp op) const { + auto regionOps = getRegionOps(op); + return regionOps.empty() && op->hasAttr("broadcastDims"); + } + + LogicalResult handleBroadcastCase(linalg::GenericOp op, + linalg::GenericOp broadcastOp, + PatternRewriter &rewriter) const { + auto broadcastInput = broadcastOp.getInputs()[0]; + auto broadcastInputType = + cast(broadcastOp->getOperandTypes()[0]); + auto outputType = cast(op->getResultTypes()[0]); + auto inputType = cast(op->getOperandTypes()[0]); + + auto newSITOFPType = RankedTensorType::get(broadcastInputType.getShape(), + outputType.getElementType()); + auto newSITOFP = + createNewSIToFpOp(op, broadcastInput, newSITOFPType, rewriter); + auto newBroadcast = + createNewBroadcast(broadcastOp, newSITOFP, outputType, rewriter); + rewriter.replaceAllOpUsesWith(op, newBroadcast); + return success(); + } + + Value createNewBroadcast(linalg::GenericOp broadcastOp, Value input, + RankedTensorType outputType, + PatternRewriter &rewriter) const { + + auto empty = rewriter.create(broadcastOp->getLoc(), + outputType.getShape(), + outputType.getElementType()); + auto broadcast = rewriter.create( + broadcastOp->getLoc(), outputType, input, ValueRange{empty}, + broadcastOp.getIndexingMapsArray(), broadcastOp.getIteratorTypesArray(), + [](OpBuilder &b, Location loc, ValueRange args) { + b.create(loc, args.drop_back()); + }); + + broadcast->setAttr("broadcastDims", broadcastOp->getAttr("broadcastDims")); + return broadcast->getResult(0); + } + + Value createNewSIToFpOp(linalg::GenericOp originalOp, Value input, + RankedTensorType outputType, + PatternRewriter &rewriter) const { + auto empty = rewriter.create(originalOp->getLoc(), + outputType.getShape(), + outputType.getElementType()); + auto newOp = rewriter.create( + originalOp->getLoc(), outputType, ValueRange{input}, ValueRange{empty}, + originalOp.getIndexingMapsArray(), originalOp.getIteratorTypesArray(), + [](OpBuilder &b, Location loc, ValueRange args) { + auto sitofp = b.create(loc, args.back().getType(), + args.drop_back()); + b.create(loc, sitofp->getResult(0)); + }); + + return newOp->getResult(0); + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + // Todo : support eliminate other type fptosi and sitofp. + // As far, only support eliminate i64 type fptosi and sitofp(i64/f32). + if (!isValidSIToFPOp(op, rewriter)) { + return failure(); + } + + auto input = op.getInputs()[0]; + auto prevOp = input.getDefiningOp(); + if (!prevOp) { + return failure(); + } + if (isBroadcastOp(prevOp)) { + return handleBroadcastCase(op, prevOp, rewriter); + } else if (isValidFPToSIOp(prevOp, rewriter)) { + rewriter.replaceAllOpUsesWith(op, prevOp.getInputs()[0]); + return success(); + } + + return failure(); + } +}; + +struct CastElementwiseOpIOToFloatPattern + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + void initialize() { + // Register conversions from SIOp to FPOp + registerSIOpMapFPOp(); + registerSIOpMapFPOp(); + registerSIOpMapFPOp(); + registerSIOpMapFPOp(); + registerSIOpMapFPOp(); + registerSIOpMapFPOp(); + } + + template void registerSIOpMapFPOp() { + OperationName SIOpName(SIOp::getOperationName(), getContext()); + assert(!SIToFPOpBuildFnMap.contains(SIOpName) && + "SIOp already registered for conversion to FPOp"); + SIToFPOpBuildFnMap[SIOpName] = + [](Operation *srcOp, PatternRewriter &rewriter, ValueRange inputs, + ValueRange outputs) -> ValueRange { + auto genericOp = cast(srcOp); + return rewriter + .create( + srcOp->getLoc(), outputs.back().getType(), inputs, outputs, + genericOp.getIndexingMapsArray(), + genericOp.getIteratorTypesArray(), + [](OpBuilder &b, Location loc, ValueRange args) { + Value val = b.create(loc, args.back().getType(), + args.drop_back()); + b.create(loc, val); + }) + .getResults(); + }; + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1) + return failure(); + + Location loc = op->getLoc(); + auto elemWiseOp = regionOps[0]; + OperationName OpName = elemWiseOp->getName(); + auto inputs = op.getInputs(); + // NOTE: Output not always exist + auto outputs = op.getOutputs(); + + // Skip DSA memref-operand generics: this pattern is tensor-only and uses + // hard casts (would assert on memref operand types). + if (outputs.empty() || !isa(outputs[0].getType())) + return failure(); + + if (SIToFPOpBuildFnMap.contains(OpName) && + !preservesIntegerPrecision( + cast(outputs[0].getType()).getElementType(), + precisionMode)) { + assert(outputs.size() == 1 && + "Elementwise conversion only support single output"); + assert(cast(outputs[0].getType()) + .getElementType() + .isInteger() && + "Output type must be integer type"); + + return convertSIOpToF32Op(op, rewriter, op.getInputs(), op.getOutputs(), + SIToFPOpBuildFnMap.at(OpName)); + } + + if (auto cmpiOp = dyn_cast(elemWiseOp)) { + auto inputType = cast(inputs.front().getType()); + + if (inputType.getNumElements() < 8) + return failure(); + + if (preservesIntegerPrecision(inputType.getElementType(), precisionMode)) + return failure(); + + auto outputType = + cast(op.getOutputs().front().getType()); + + arith::CmpFPredicate fpPred; + switch (cmpiOp.getPredicate()) { + default: + return failure(); + case arith::CmpIPredicate::eq: + fpPred = arith::CmpFPredicate::OEQ; + break; + case arith::CmpIPredicate::ne: + fpPred = arith::CmpFPredicate::ONE; + break; + case arith::CmpIPredicate::sge: + fpPred = arith::CmpFPredicate::OGE; + break; + case arith::CmpIPredicate::sgt: + fpPred = arith::CmpFPredicate::OGT; + break; + case arith::CmpIPredicate::sle: + fpPred = arith::CmpFPredicate::OLE; + break; + case arith::CmpIPredicate::slt: + fpPred = arith::CmpFPredicate::OLT; + break; + } + + return convertSIOpToF32Op( + op, rewriter, op.getInputs(), ValueRange{}, + [&](Operation *srcOp, PatternRewriter &rewriter, ValueRange inputs, + ValueRange outputs) { + auto genericOp = cast(srcOp); + + return rewriter + .create( + srcOp->getLoc(), outputType, inputs, genericOp.getOutputs(), + genericOp.getIndexingMapsArray(), + genericOp.getIteratorTypesArray(), + [&](OpBuilder &b, Location loc, ValueRange args) { + Value val = b.create(loc, fpPred, args[0], + args[1]); + b.create(loc, val); + }) + ->getResults(); + }); + } + + if (auto divsiOp = dyn_cast(elemWiseOp)) { + // Route integer division/remainder by precision mode: + // mode 0: always cast to f32 and lower on Wafer (this path), any width. + // mode 1: widths < 64 use f32/Wafer; i64 falls back to exact RISC-V. + // mode 2: all widths fall back to exact RISC-V integer division. + // The `>= 2` term is what forces i8/i16 to RISC-V at mode 2 (which + // preservesIntegerPrecision alone would miss, as it gates on >= 32). + if (normalizePrecisionMode(precisionMode) >= 2 || + preservesIntegerPrecision( + cast(outputs[0].getType()).getElementType(), + precisionMode)) + return failure(); + // Integer division via f32 loses a unit when the hardware reciprocal + // rounds low (e.g. 7//7 -> 0). Detect it with a remainder check and + // nudge the quotient by sign(q_f) (trunc-toward-zero, like arith.divsi). + return convertSIOpToF32Op( + op, rewriter, op.getInputs(), op.getOutputs(), + [&](Operation *srcOp, PatternRewriter &rewriter, ValueRange inputs, + ValueRange outputs) -> ValueRange { + Value qOut = buildCorrectiveDivSigned(rewriter, srcOp->getLoc(), + inputs[0], inputs[1]); + return qOut.getDefiningOp()->getResults(); + }); + } + + if (auto divuiOp = dyn_cast(elemWiseOp)) { + // Route integer division/remainder by precision mode: + // mode 0: always cast to f32 and lower on Wafer (this path), any width. + // mode 1: widths < 64 use f32/Wafer; i64 falls back to exact RISC-V. + // mode 2: all widths fall back to exact RISC-V integer division. + // The `>= 2` term is what forces i8/i16 to RISC-V at mode 2 (which + // preservesIntegerPrecision alone would miss, as it gates on >= 32). + if (normalizePrecisionMode(precisionMode) >= 2 || + preservesIntegerPrecision( + cast(outputs[0].getType()).getElementType(), + precisionMode)) + return failure(); + // Unsigned: inputs are non-negative, so the positive-domain corrective + // division is sufficient (no sign handling needed). + return convertUIOpToF32Op( + op, rewriter, op.getInputs(), op.getOutputs(), + [&](Operation *srcOp, PatternRewriter &rewriter, ValueRange inputs, + ValueRange outputs) -> ValueRange { + Location loc = srcOp->getLoc(); + Value qOut = + buildCorrectivePosDiv(rewriter, loc, inputs[0], inputs[1]); + return qOut.getDefiningOp()->getResults(); + }); + } + + if (auto remsiOp = dyn_cast(elemWiseOp)) { + // Route integer division/remainder by precision mode: + // mode 0: always cast to f32 and lower on Wafer (this path), any width. + // mode 1: widths < 64 use f32/Wafer; i64 falls back to exact RISC-V. + // mode 2: all widths fall back to exact RISC-V integer division. + // The `>= 2` term is what forces i8/i16 to RISC-V at mode 2 (which + // preservesIntegerPrecision alone would miss, as it gates on >= 32). + if (normalizePrecisionMode(precisionMode) >= 2 || + preservesIntegerPrecision( + cast(outputs[0].getType()).getElementType(), + precisionMode)) + return failure(); + // r = a - q*b with q from the corrective signed division. + // q*b and the subtraction are exact in FP32 (all values < 2^24). + return convertSIOpToF32Op( + op, rewriter, op.getInputs(), op.getOutputs(), + [&](Operation *srcOp, PatternRewriter &rewriter, ValueRange inputs, + ValueRange outputs) -> ValueRange { + Location loc = srcOp->getLoc(); + Value q = + buildCorrectiveDivSigned(rewriter, loc, inputs[0], inputs[1]); + Value qb = buildLinalgElementwise(rewriter, loc, + {q, inputs[1]}); + Value r = buildLinalgElementwise(rewriter, loc, + {inputs[0], qb}); + return r.getDefiningOp()->getResults(); + }); + } + + if (auto remuiOp = dyn_cast(elemWiseOp)) { + // Route integer division/remainder by precision mode: + // mode 0: always cast to f32 and lower on Wafer (this path), any width. + // mode 1: widths < 64 use f32/Wafer; i64 falls back to exact RISC-V. + // mode 2: all widths fall back to exact RISC-V integer division. + // The `>= 2` term is what forces i8/i16 to RISC-V at mode 2 (which + // preservesIntegerPrecision alone would miss, as it gates on >= 32). + if (normalizePrecisionMode(precisionMode) >= 2 || + preservesIntegerPrecision( + cast(outputs[0].getType()).getElementType(), + precisionMode)) + return failure(); + // r = a - q*b with q from the corrective unsigned division. + return convertUIOpToF32Op( + op, rewriter, op.getInputs(), op.getOutputs(), + [&](Operation *srcOp, PatternRewriter &rewriter, ValueRange inputs, + ValueRange outputs) -> ValueRange { + Location loc = srcOp->getLoc(); + Value q = + buildCorrectivePosDiv(rewriter, loc, inputs[0], inputs[1]); + Value qb = buildLinalgElementwise(rewriter, loc, + {q, inputs[1]}); + Value r = buildLinalgElementwise(rewriter, loc, + {inputs[0], qb}); + return r.getDefiningOp()->getResults(); + }); + } + + return failure(); + } + + CastElementwiseOpIOToFloatPattern(MLIRContext *context, int precisionMode) + : OpRewritePattern(context), + precisionMode(precisionMode) {} + +private: + // Map from SIOp to FPOp conversion functions + llvm::DenseMap> + SIToFPOpBuildFnMap; + + int precisionMode = 0; +}; + +struct CastReduceOpIOToFloatPattern + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + void initialize() { + // Register conversions from SIOp to FPOp + registerSIOpMapFPOp(); + registerSIOpMapFPOp(); + registerSIOpMapFPOp(); + registerSIOpMapFPOp(); + } + + template void registerSIOpMapFPOp() { + OperationName SIOpName(SIOp::getOperationName(), getContext()); + assert(!SIToFPOpBuildFnMap.contains(SIOpName) && + "SIOp already registered for conversion to FPOp"); + SIToFPOpBuildFnMap[SIOpName] = [](Operation *op, PatternRewriter &rewriter, + ValueRange inputs, + ValueRange outputs) -> ValueRange { + auto reduceOp = cast(op); + return rewriter + .create( + reduceOp->getLoc(), inputs, outputs, reduceOp.getDimensions(), + [](OpBuilder &b, Location loc, ValueRange args) { + Value val = b.create(loc, args.back().getType(), args); + b.create(loc, val); + }) + .getResults(); + }; + } + + LogicalResult matchAndRewrite(linalg::ReduceOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1) + return failure(); + + auto reduceOp = regionOps[0]; + OperationName OpName = reduceOp->getName(); + + if (SIToFPOpBuildFnMap.contains(OpName) && + !preservesIntegerPrecision( + cast(op.getInits().front().getType()) + .getElementType(), + precisionMode)) { + + assert(op.getInits().size() == 1 && + "Reduce conversion only support single output"); + + auto constantType = + cast(op.getInits().front().getType()) + .getElementType(); + + auto attr = getRedBaseAttr(rewriter, reduceOp, constantType); + if (checkReductionBaseAttr(op, rewriter, attr)) + return rewriter.notifyMatchFailure( + op, "Reduction op has invalid init value"); + + return convertSIOpToF32Op(op, rewriter, op.getInputs(), op.getInits(), + SIToFPOpBuildFnMap.at(OpName), + /*convertOutputs=*/true); + } + + return failure(); + } + + CastReduceOpIOToFloatPattern(MLIRContext *context, int precisionMode) + : OpRewritePattern(context), + precisionMode(precisionMode) {} + +private: + // Map from SIOp to FPOp conversion functions + llvm::DenseMap> + SIToFPOpBuildFnMap; + + int precisionMode = 0; +}; + +template +struct CastArgMinMaxOpIOToFloatPattern : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + using OpAdaptor = typename MKOpT::Adaptor; + LogicalResult matchAndRewrite(MKOpT op, + PatternRewriter &rewriter) const override { + auto input = op.getSrc(); + auto inputTy = cast(input.getType()); + if (!inputTy.getElementType().isInteger()) { + return failure(); + } + + auto bitWidth = inputTy.getElementTypeBitWidth(); + if (bitWidth != 16 && bitWidth != 32 && bitWidth != 64) { + return failure(); + } + auto loc = op->getLoc(); + + Value inputEmpty = rewriter.create(loc, inputTy.getShape(), + rewriter.getF32Type()); + Value fpInput = createElemwiseNaryOp(rewriter, loc, + {input}, inputEmpty); + // FIXME: Since don't do tiling for argmin/argmax, we assume the init value + // is always empty + auto outValue = op.getValue(); + auto valueTy = cast(outValue.getType()); + auto fpValueTy = + RankedTensorType::get(valueTy.getShape(), rewriter.getF32Type()); + Value valueEmpty = rewriter.create(loc, valueTy.getShape(), + rewriter.getF32Type()); + + auto outIdx = op.getIndex(); + auto axis = op.getAxis(); + auto newOp = + rewriter.create(loc, TypeRange{fpValueTy, outIdx.getType()}, + fpInput, valueEmpty, outIdx, axis); + Value fpValue = createElemwiseNaryOp( + rewriter, loc, ValueRange{newOp.getResults()[0]}, outValue); + rewriter.replaceOp(op, ValueRange{fpValue, newOp.getResults()[1]}); + return success(); + } +}; + +struct BoolOpShapeCanonicalizePattern : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + bool linearizeShape(linalg::GenericOp op, PatternRewriter &rewriter) const { + assert(op.getOutputs().size() == 1 && "Only support single output"); + assert(llvm::all_of(op.getIndexingMapsArray(), + [](AffineMap &map) { return map.isIdentity(); }) && + "All affine maps must be identity affine map."); + + Location loc = op->getLoc(); + auto dstTensorType = cast(op.getOutputs()[0].getType()); + + if (dstTensorType.getRank() == 1) + return false; + + auto elemCount = dstTensorType.getNumElements(); + Value zero = rewriter.create(loc, 0); + Value elemCountVal = rewriter.create( + loc, rewriter.getI32Type(), elemCount); + + auto indices = llvm::seq(0, dstTensorType.getRank()); + SmallVector reassociation(1); + reassociation[0].insert(reassociation[0].end(), indices.begin(), + indices.end()); + + SmallVector inputs1D = llvm::map_to_vector( + llvm::concat(op.getInputs(), op.getOutputs()), + [&](Value val) -> Value { + return rewriter.create(loc, val, + reassociation); + }); + + Value output1D = inputs1D.pop_back_val(); + SmallVector idMaps(inputs1D.size() + 1, + rewriter.getMultiDimIdentityMap(1)); + SmallVector iters( + 1, mlir::utils::IteratorType::parallel); + auto newOp = rewriter.create( + loc, RankedTensorType::get({elemCount}, dstTensorType.getElementType()), + inputs1D, ValueRange{output1D}, idMaps, iters); + newOp.getRegion().takeBody(op.getRegion()); + + rewriter.replaceOpWithNewOp( + op, dstTensorType, newOp->getResult(0), reassociation); + + return true; + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1) + return failure(); + + Location loc = op->getLoc(); + auto elemWiseOp = regionOps[0]; + OperationName OpName = elemWiseOp->getName(); + auto inputs = op.getInputs(); + auto outputs = op.getOutputs(); + + // Check if the operation is a boolean operation + if (!(isa(elemWiseOp))) + return failure(); + + if (linearizeShape(op, rewriter)) + return success(); + + auto inputTensorType = cast(inputs[0].getType()); + auto dstTensorType = cast(outputs[0].getType()); + auto elemCount = dstTensorType.getNumElements(); + + assert(dstTensorType.getRank() == 1); + + if (!(elemCount & 0x7)) + return failure(); + + Value result = rewriter.create( + loc, dstTensorType.getShape(), dstTensorType.getElementType()); + // Legalize operations that are not multiples of 8 + unsigned mainCount = elemCount & ~0x7; + if (mainCount) { + + SmallVector ins = llvm::map_to_vector( + llvm::concat(inputs, outputs), [&](Value val) -> Value { + return rewriter.create( + loc, + RankedTensorType::get( + {mainCount}, + cast(val.getType()).getElementType()), + val, ValueRange(), ValueRange(), ValueRange(), + ArrayRef{0}, ArrayRef{mainCount}, + ArrayRef{1}); + }); + + Value out = ins.pop_back_val(); + SmallVector idMaps(inputs.size() + 1, + rewriter.getMultiDimIdentityMap(1)); + SmallVector iters( + 1, mlir::utils::IteratorType::parallel); + auto newOp = rewriter.create( + loc, + RankedTensorType::get({mainCount}, dstTensorType.getElementType()), + ins, ValueRange{out}, idMaps, iters); + newOp.getRegion().takeBody(op.getRegion()); + + result = rewriter.create( + loc, newOp.getResult(0), result, ValueRange(), ValueRange(), + ValueRange(), ArrayRef{0}, ArrayRef{mainCount}, + ArrayRef{1}); + } + + for (unsigned idx = mainCount; idx < elemCount; ++idx) { + auto idxVal = rewriter.create(loc, idx); + auto loadIns = llvm::map_to_vector(inputs, [&](Value source) { + return rewriter.create(loc, source, + ValueRange{idxVal}); + }); + IRMapping mapper; + mapper.map(elemWiseOp->getOperands(), loadIns); + auto newVal = rewriter.clone(*elemWiseOp, mapper); + + result = rewriter.create(loc, newVal->getResult(0), + result, ValueRange{idxVal}); + } + + rewriter.replaceOp(op, result); + + return success(); + } + + BoolOpShapeCanonicalizePattern(MLIRContext *context, int precisionMode) + : OpRewritePattern(context), + precisionMode(precisionMode) {} + +private: + int precisionMode = 0; +}; + +struct SigmoidFusionPattern : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + bool matchSigmoid(linalg::GenericOp op, Value &input) const { + // 1. sub (0 - x = -x) + // 2. exp (e^(-x)) + // 3. add (1 + e^(-x)) + // 4. div (1 / (1 + e(^-x))) + // We match the sigmoid pattern from down to up. + + // 1. Match div first. + if (!checkGenericOp(op)) { + return false; + } + + auto divLhs = op.getInputs()[0]; + if (!isOneTensor(divLhs)) { + return false; + } + + // 2. Match add. + auto addResult = op.getInputs()[1]; + auto addGenericOp = addResult.getDefiningOp(); + if (!addGenericOp || !checkGenericOp(addGenericOp)) { + return false; + } + + auto addLhs = addGenericOp.getInputs()[0]; + auto addRhs = addGenericOp.getInputs()[1]; + bool isAddLhsOne = isOneTensor(addLhs); + bool isAddRhsOne = isOneTensor(addRhs); + if (!isAddLhsOne && !isAddRhsOne) { + return false; + } + + // 3. Match exp. + auto expResult = isAddLhsOne ? addRhs : addLhs; + auto expGenericOp = expResult.getDefiningOp(); + if (!expGenericOp || !checkGenericOp(expGenericOp)) { + return false; + } + + // 4. Match sub. + auto subResult = expGenericOp.getInputs()[0]; + auto subGenericOp = subResult.getDefiningOp(); + if (!subGenericOp || !checkGenericOp(subGenericOp)) { + return false; + } + + auto subLhs = subGenericOp.getInputs()[0]; + if (!isZeroTensor(subLhs)) { + return false; + } + + // Set input of Sub operation to the input of the sigmoid op. + input = subGenericOp.getInputs()[1]; + + // Match sigmoid pattern successfully. + return true; + } + +public: + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + // Match sigmoid pattern + Location loc = op.getLoc(); + Value input; + if (!matchSigmoid(op, input)) { + return rewriter.notifyMatchFailure(op, "sigmoid pattern not matched"); + } + + auto dstType = cast(op.getType(0)); + auto elementType = dstType.getElementType(); + auto init = + rewriter.create(loc, dstType.getShape(), elementType); + + // Replace the div GenericOp with mk::SigmoidOp + // We can use CSE to erase other unused generic ops. + auto sigmoidOp = rewriter.replaceOpWithNewOp( + op, dstType, input, init, rewriter.getBoolAttr(false)); + + return success(); + } +}; + +struct GeluFusionPattern : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + bool isErfScaleTensor(Value &v) const { + const float kSqrtScaleOverPiF32 = 1.0f / std::sqrt(2.0f); + const float kSqrtScaleOverPiF16 = 0.70703125f; + const float kSqrtScaleOverPiBF16 = 0.70703125f; + auto elementType = cast(v.getType()).getElementType(); + float erfScale; + if (elementType.isF32()) { + erfScale = kSqrtScaleOverPiF32; + } else if (elementType.isF16()) { + erfScale = kSqrtScaleOverPiF16; + } else if (elementType.isBF16()) { + erfScale = kSqrtScaleOverPiBF16; + } else { + return false; + } + return isConstantTensor(v, erfScale); + } + + bool matchGeluErf(linalg::GenericOp op, Value input) const { + // Match the mul op: (x * scale) + auto mulErfResult = op.getInputs()[0]; + auto mulErfGenericOp = mulErfResult.getDefiningOp(); + if (!mulErfGenericOp || !checkGenericOp(mulErfGenericOp)) { + return false; + } + auto mulErfLhs = mulErfGenericOp.getInputs()[0]; + auto mulErfRhs = mulErfGenericOp.getInputs()[1]; + if (!isErfScaleTensor(mulErfRhs) && !isErfScaleTensor(mulErfLhs)) { + return false; + } + + Value input1 = isErfScaleTensor(mulErfRhs) ? mulErfLhs : mulErfRhs; + + return input1 == input; + } + + bool isTanhScaledTensor(Value &v) const { + // Check if the value is a constant tensor with the value of sqrt(2 / pi). + const float kSqrt2OverPiF32 = std::sqrt(2.0f / M_PI); + const float kSqrt2OverPiF16 = 0.7978515625f; + const float kSqrt2OverPiBF16 = 0.796875f; + auto elementType = cast(v.getType()).getElementType(); + float tanhScale; + if (elementType.isF32()) { + tanhScale = kSqrt2OverPiF32; + } else if (elementType.isF16()) { + tanhScale = kSqrt2OverPiF16; + } else if (elementType.isBF16()) { + tanhScale = kSqrt2OverPiBF16; + } else { + return false; // Unsupported element type + } + return isConstantTensor(v, tanhScale); + } + + bool isPowScaledTensor(Value &v) const { + // Check if the value is a constant tensor with the value of 0.044715. + const float powScale = 0.044715f; + return isConstantTensor(v, powScale, true); + } + + bool isAddAndMulOp(linalg::GenericOp op1, linalg::GenericOp op2) const { + // Check if the given linalg generic op is an add and mul op. + return (checkGenericOp(op1) && + checkGenericOp(op2)) || + (checkGenericOp(op2) && + checkGenericOp(op1)); + } + + bool isExtfAndAddOp(linalg::GenericOp op1, linalg::GenericOp op2) const { + // Check if the given linalg generic op is an extf and add op. + return (checkGenericOp(op1) && + checkGenericOp(op2)) || + (checkGenericOp(op2) && + checkGenericOp(op1)); + } + + // According LHS and RHS of the outer mul op, get the nested mul generic op + // If the element type is F16 or BF16, we need to extend the input. + // If the element type is F32, we can directly use the input. + linalg::GenericOp getNestedMulGenericOp(linalg::GenericOp lhsGenericOp, + linalg::GenericOp rhsGenericOp, + bool isBit16) const { + if (!isBit16) { // match case : mul (mul(lhs * rhs) * add (lhs1 *rhs1)) + if (!isAddAndMulOp(lhsGenericOp, rhsGenericOp)) { + return linalg::GenericOp(); + } + return checkGenericOp(lhsGenericOp) ? lhsGenericOp + : rhsGenericOp; + } else { // match case : mul (extf (mul (lhs * rhs)) * add (lhs1 *rhs1)) + if (!isExtfAndAddOp(lhsGenericOp, rhsGenericOp)) { + return linalg::GenericOp(); + } + auto extfGenericOp = checkGenericOp(lhsGenericOp) + ? lhsGenericOp + : rhsGenericOp; + auto nestedMulGenericOp = + extfGenericOp.getInputs()[0].getDefiningOp(); + return (!nestedMulGenericOp || + !checkGenericOp(nestedMulGenericOp)) + ? linalg::GenericOp() + : nestedMulGenericOp; + } + } + + bool matchGeluTanh(linalg::GenericOp op, Value input, bool isBit16) const { + // 1. pow (pow(x, 2))) + // 2. mul (0.044715 * pow(x, 2))) + // 3. add (1 + 0.044715 * pow(x, 2))) + // 4. mul (x * 0.79788456) + // 5. mul (x * 0.79788456 * (1 + 0.044715 * pow(x, 2))) + + // We match the gelu tanh pattern from down to up. + + // 1. Match the mul op + auto mulOfTanhResult = op.getInputs()[0]; + auto mulOfTanhGenericOp = + mulOfTanhResult.getDefiningOp(); + if (!mulOfTanhGenericOp || + !checkGenericOp(mulOfTanhGenericOp)) { + return false; + } + + auto mulOfTanhLhs = mulOfTanhGenericOp.getInputs()[0]; + auto mulOfTanhRhs = mulOfTanhGenericOp.getInputs()[1]; + auto mulOfTanhLhsGenericOp = + mulOfTanhLhs.getDefiningOp(); + auto mulOfTanhRhsGenericOp = + mulOfTanhRhs.getDefiningOp(); + if (!mulOfTanhLhsGenericOp || !mulOfTanhRhsGenericOp) { + return false; + } + + // Get the nested mul generic op. + // If the element type is F16 or BF16, match the extf op and add op : + // mul_result = mul (extf (mul (x, 0.79788456)), add (1 , operand)) + // If the element type is F32, match the mul op and add op : + // mul_result = mul (mul (x, 0.79788456), add (1 , operand)) + linalg::GenericOp nestedMulGenericOp = getNestedMulGenericOp( + mulOfTanhLhsGenericOp, mulOfTanhRhsGenericOp, isBit16); + if (!nestedMulGenericOp) { + return false; + } + + // 2. Match the mul op: mul (x * 0.79788456) + auto nestedMulInput1 = nestedMulGenericOp.getInputs()[0]; + auto nestedMulInput2 = nestedMulGenericOp.getInputs()[1]; + + if (!isTanhScaledTensor(nestedMulInput1) && + !isTanhScaledTensor(nestedMulInput2)) { + // If both operands are not scale tensors, we cannot match the gelu + // pattern. + return false; + } + Value input1 = + isTanhScaledTensor(nestedMulInput1) ? nestedMulInput2 : nestedMulInput1; + if (input1 != input) { + // If the inputs of the mul ops are not the same, we cannot match the gelu + // pattern. + return false; + } + + // 3. Match add (1 + 0.044715 * pow(x, 2))) + auto nestedAddGenericOp = + checkGenericOp(mulOfTanhLhsGenericOp) + ? mulOfTanhLhsGenericOp + : mulOfTanhRhsGenericOp; + auto nestedAddLhs = nestedAddGenericOp.getInputs()[0]; + auto nestedAddRhs = nestedAddGenericOp.getInputs()[1]; + if (!isOneTensor(nestedAddLhs) && !isOneTensor(nestedAddRhs)) { + return false; + } + + // 4. Match the mul op: mul (0.044715 * pow(x, 2))) + auto finalMulResult = + isOneTensor(nestedAddLhs) ? nestedAddRhs : nestedAddLhs; + auto finalMulGenericOp = finalMulResult.getDefiningOp(); + if (!finalMulGenericOp || + !checkGenericOp(finalMulGenericOp)) { + return false; + } + auto finalMulLhs = finalMulGenericOp.getInputs()[0]; + auto finalMulRhs = finalMulGenericOp.getInputs()[1]; + if (!isPowScaledTensor(finalMulLhs) && !isPowScaledTensor(finalMulRhs)) { + // If both operands are not half tensors, we cannot match the gelu + // pattern. + return false; + } + + // 5. Match pow (pow(x, 2))) + auto powResult = isPowScaledTensor(finalMulRhs) ? finalMulLhs : finalMulRhs; + auto powGenericOp = powResult.getDefiningOp(); + if (!powGenericOp || !checkGenericOp(powGenericOp)) { + return false; + } + + auto powLhs = powGenericOp.getInputs()[0]; + auto powRhs = powGenericOp.getInputs()[1]; + if (!isTwoTensor(powRhs)) { + // If the exponent is not a two tensor, we cannot match the gelu pattern. + return false; + } + + if (!isBit16) { + return powLhs == input; + } + // If the element type is F16 or BF16, match the extf op nested pow op : + // pow_result = pow (extf (x), 2) + auto finalExtfGenericOp = powLhs.getDefiningOp(); + if (!finalExtfGenericOp || + !checkGenericOp(finalExtfGenericOp)) { + return false; + } + + return finalExtfGenericOp.getInputs()[0] == input; + } + + bool matchGelu(linalg::GenericOp op, Value &input, GeluMode &geluMode) const { + // Match gelu none or gelu tanh pattern. + // match 0.5 * x * (1 + tanh/erf) + + // 1. match mul first. + bool isBit16 = false; + linalg::GenericOp mulGenericOp = op; + if (checkGenericOp(op) && + dyn_cast(op.getType(0)).getElementTypeBitWidth() == + 16) { + isBit16 = true; + mulGenericOp = op.getInputs()[0].getDefiningOp(); + } + + if (!mulGenericOp || !checkGenericOp(mulGenericOp)) { + return false; + } + + auto mulLhs = mulGenericOp.getInputs()[0]; + auto mulRhs = mulGenericOp.getInputs()[1]; + auto mulLhsGenericOp = mulLhs.getDefiningOp(); + auto mulRhsGenericOp = mulRhs.getDefiningOp(); + if (!mulLhsGenericOp || !mulRhsGenericOp) { + return false; + } + + // Get the nested mul generic op. + // If the element type is F16 or BF16, and is tanh op: + // mul_result = trunf( mul (extf (mul (x, 0.5)), add (1, tanh)) ) + // If the element type is F32, match the mul op and add op : + // mul_result = mul (mul (x, 0.5), add (1, tanh/erf)) + linalg::GenericOp nestedMulOp = + getNestedMulGenericOp(mulLhsGenericOp, mulRhsGenericOp, isBit16); + if (!nestedMulOp) { + return false; + } + + // 2. Match the mul op: mul (0.5 * x) + auto nestedMulLhs = nestedMulOp.getInputs()[0]; + auto nestedMulRhs = nestedMulOp.getInputs()[1]; + bool isNestedMulLhsHalf = isHalfTensor(nestedMulLhs); + bool isNestedMulRhsHalf = isHalfTensor(nestedMulRhs); + if (!isNestedMulLhsHalf && !isNestedMulRhsHalf) { + // If both operands are not half tensors, we cannot match the gelu + // pattern. + return false; + } + Value input1 = isNestedMulRhsHalf ? nestedMulLhs : nestedMulRhs; + + // 3. Match add (1 + tanh/erf). + auto addGenericOp = checkGenericOp(mulLhsGenericOp) + ? mulLhsGenericOp + : mulRhsGenericOp; + auto addLhs = addGenericOp.getInputs()[0]; + auto addRhs = addGenericOp.getInputs()[1]; + bool isAddLhsOne = isOneTensor(addLhs); + bool isAddRhsOne = isOneTensor(addRhs); + if (!isAddLhsOne && !isAddRhsOne) { + return false; + } + + // 4. Match tanh/erf. + auto tanhOrErfResult = isAddLhsOne ? addRhs : addLhs; + auto tanhOrErfGenericOp = + tanhOrErfResult.getDefiningOp(); + if (!tanhOrErfGenericOp) { + return false; + } + if (checkGenericOp(tanhOrErfGenericOp) && + matchGeluTanh(tanhOrErfGenericOp, input1, isBit16)) { + geluMode = GeluMode::Tanh; + input = input1; + return true; + } + if (checkGenericOp(tanhOrErfGenericOp) && + matchGeluErf(tanhOrErfGenericOp, input1)) { + geluMode = GeluMode::None; + input = input1; + return true; + } + + // If the tanh/erf op is not matched, we cannot match the gelu pattern. + return false; + } + +public: + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + // Match gelu pattern + Location loc = op.getLoc(); + Value input; + GeluMode geluMode = GeluMode::None; + if (!matchGelu(op, input, geluMode)) { + // If the gelu pattern is not matched, we cannot rewrite the op. + return rewriter.notifyMatchFailure(op, "gelu pattern not matched"); + } + + auto dstType = cast(op.getType(0)); + auto elementType = dstType.getElementType(); + auto init = + rewriter.create(loc, dstType.getShape(), elementType); + + // Replace the mul generic op with mk::GeluOp + switch (geluMode) { + case GeluMode::None: { + rewriter.replaceOpWithNewOp( + op, dstType, input, nullptr, init, rewriter.getBoolAttr(false), + rewriter.getI16IntegerAttr(static_cast(geluMode))); + break; + } + case GeluMode::Tanh: { + // Double elementcount F32 buffer for immediate variable + auto imm = rewriter.create( + loc, dstType.getNumElements() * 2, rewriter.getF32Type()); + rewriter.replaceOpWithNewOp( + op, dstType, input, imm, init, rewriter.getBoolAttr(false), + rewriter.getI16IntegerAttr(static_cast(geluMode))); + break; + } + default: { + llvm::report_fatal_error("Unsupported gelu mode!"); + } + } + + return success(); + } +}; + +// copy from newest llvm lib +mlir::tensor::CollapseShapeOp +dropGivenUnitDims(OpBuilder &b, Location loc, Value src, + const llvm::SmallBitVector &dropDims) { + auto srcType = cast(src.getType()); + int64_t rank = srcType.getRank(); + assert(rank == static_cast(dropDims.size()) && + "dropDims dimension does not match src tensor rank"); + assert(llvm::all_of( + dropDims.set_bits(), + [&](unsigned dim) { return srcType.getShape()[dim] == 1; }) && + "Dropping non unit dimension"); + // Computed reassociation map for the corresponding tensor.collapse_shape. + SmallVector reassocMaps; + // Current reassociation group to add dropped dimension to. + + int64_t nextDimToGroup = 0; + llvm::SmallBitVector keptDims(dropDims); + keptDims.flip(); + int64_t lastSetBit = keptDims.find_last(); + for (int64_t setBit : keptDims.set_bits()) { + // Group consecutive dropped dimension with the next non-dropped dimension. + // If this is the last set dimension, also group all subsequent dropped + // dimension, if any. + int64_t upTo = setBit == lastSetBit ? rank - 1 : setBit; + auto seq = llvm::seq_inclusive(nextDimToGroup, upTo); + reassocMaps.emplace_back(llvm::make_range(seq.begin(), seq.end())); + nextDimToGroup = setBit + 1; + } + return b.create(loc, src, reassocMaps); +} + +template +struct ArgMinMaxFusionPattern : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + bool checkReductionBaseAttr(Value outVal) const { + if (outVal.getDefiningOp()) + return true; + if (auto fillOp = outVal.getDefiningOp()) { + // TODO: check init is argmin/argmax identity value. + return fillOp.getInputs()[0].getDefiningOp(); + } + assert(false && "Unsupported init op"); + return false; + } + +public: + LogicalResult matchAndRewrite(linalg::ReduceOp op, + PatternRewriter &rewriter) const override { + if (op.getBody()->getNumArguments() != 4 || + op.getDimensions().size() != 1) { + return failure(); + } + + // Get input and output types + auto input = op.getInputs()[0]; + auto outVal = op.getInits()[0]; + auto outIdx = op.getInits()[1]; + + if (!checkReductionBaseAttr(outVal)) + return rewriter.notifyMatchFailure( + op.getLoc(), "mk.argmin/max not support non-identity init\n"); + + auto inputType = cast(input.getType()); + auto valueType = cast(outVal.getType()); + auto indexType = cast(outIdx.getType()); + auto inputShape = inputType.getShape(); + + // skip unsupport dtype + auto elementType = inputType.getElementType(); + auto bitWidth = elementType.getIntOrFloatBitWidth(); + if (bitWidth == 64 && elementType.isF64()) { + return failure(); + } + + assert(bitWidth == 16 || bitWidth == 32 || bitWidth == 64); + + // Get the reduction block and its operations + auto block = op.getBody(); + auto ops = block->without_terminator(); + + // Extract block arguments for current and reduced values/indices + Value currValue = block->getArgument(0); + Value currIndex = block->getArgument(1); + Value reduceValue = block->getArgument(2); + Value reduceIndex = block->getArgument(3); + + // Match the ArgMin/ArgMax pattern in the block + bool isArgMin = std::is_same::value; + auto opsIter = ops.begin(); + Value indexSelectOp, valueSelectOp; + if (failed(matchArgMinMax(currValue, currIndex, reduceValue, reduceIndex, + opsIter, indexSelectOp, valueSelectOp, + isArgMin))) { + return failure(); + } + + // Verify the terminator operation matches expected pattern + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *opsIter << "\n"); + auto termOp = dyn_cast(*opsIter++); + if (!termOp || termOp != block->getTerminator()) + return failure(); + if (termOp.getOperands() != ArrayRef{valueSelectOp, indexSelectOp}) { + return failure(); + } + + auto loc = op->getLoc(); + auto reduceDim = op.getDimensions()[0]; + int64_t reduceSize = inputShape[reduceDim]; + bool keepDim = inputType.getRank() == valueType.getRank(); + assert(valueType.getRank() == indexType.getRank()); + // Support reduce on last dim only temporary + // TODO: for other dim, we can create a new contiguous buffer and copy in + assert(reduceDim == inputType.getRank() - 1); + + auto zero = rewriter.create(loc, 0); + auto one = rewriter.create(loc, 1); + + SmallVector lbs, ubs, steps; + for (auto [i, size] : enumerate(inputShape)) { + if (i != reduceDim) { + auto sizeValue = rewriter.create(loc, size); + lbs.push_back(zero); + ubs.push_back(sizeValue); + steps.push_back(one); + } + } + + auto loopNest = scf::buildLoopNest( + rewriter, loc, lbs, ubs, steps, ValueRange{outVal, outIdx}, + [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange indices, + ValueRange iterArgs) { + SmallVector inputOffsets, inputSizes, inputStrides; + SmallVector outputOffsets, outputSizes, outputStrides; + llvm::SmallBitVector dropDims; + for (auto [i, size] : enumerate(inputShape)) { + if (i == reduceDim) { + inputOffsets.push_back(rewriter.getIndexAttr(0)); + inputSizes.push_back(rewriter.getIndexAttr(size)); + dropDims.push_back(false); + } else { + inputOffsets.push_back(indices[i < reduceDim ? i : i - 1]); + inputSizes.push_back(rewriter.getIndexAttr(1)); + dropDims.push_back(true); + } + inputStrides.push_back(rewriter.getIndexAttr(1)); + } + for (auto [i, size] : enumerate(valueType.getShape())) { + if (keepDim && i == reduceDim) + outputOffsets.push_back(rewriter.getIndexAttr(0)); + else + outputOffsets.push_back(indices[i < reduceDim ? i : i - keepDim]); + outputSizes.push_back(rewriter.getIndexAttr(1)); + outputStrides.push_back(rewriter.getIndexAttr(1)); + } + auto inputVec = nestedBuilder.create( + loc, input, inputOffsets, inputSizes, inputStrides); + auto inputCollapsed = + dropGivenUnitDims(rewriter, loc, inputVec, dropDims); + auto outValVec = nestedBuilder.create( + loc, iterArgs[0], outputOffsets, outputSizes, outputStrides); + auto outIdxVec = nestedBuilder.create( + loc, iterArgs[1], outputOffsets, outputSizes, outputStrides); + auto argOp = nestedBuilder.create( + loc, TypeRange{outValVec.getType(), outIdxVec.getType()}, + inputCollapsed, outValVec, outIdxVec, 0); + auto outValTensor = nestedBuilder.create( + loc, argOp.getResult()[0], iterArgs[0], outputOffsets, + outputSizes, outputStrides); + auto outIdxTensor = nestedBuilder.create( + loc, argOp.getResult()[1], iterArgs[1], outputOffsets, + outputSizes, outputStrides); + return SmallVector{outValTensor, outIdxTensor}; + }); + rewriter.replaceAllUsesWith(op->getResults(), loopNest.results); + rewriter.eraseOp(op); + return success(); + } +}; + +struct ReduceOpToElementwiseOpConverter + : public ReduceScanOpConversionBase { +private: + using ReduceScanOpConversionBase::ReduceScanOpConversionBase; + + // memref: Assume base is least 8 bit align. offset is calculated as + // byte. So we don't expected extract i1 in bytes. + SmallVector lowerBool1DInput(ConversionPatternRewriter &rewriter, + Location loc, Type elementType, + ValueRange inputs, + linalg::ReduceOp op) const { + + auto rop = getRegionOps(op).front(); + + auto attr = getRedBaseAttr(rewriter, rop, elementType); + + auto finalResult = + createReduceOp(rewriter, op, loc, inputs, SmallVector{0}, + SmallVector{}, elementType, attr) + .getResults(); + return finalResult; + } + + bool isInputsIncludeI1Type(ValueRange inputs) const { + return llvm::any_of(inputs, [](Value input) { + auto inputType = dyn_cast(input.getType()); + return inputType && inputType.getElementType().isInteger(1); + }); + } + + SmallVector + lower1DInput(ValueRange inputs, linalg::ReduceOp op, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + + auto leadingInputType = cast(inputs[0].getType()); + auto shape = leadingInputType.getShape(); + + int32_t tileSize = shape[0] > 64 ? 64 : shape[0]; + + SmallVector lastRes(inputs.size()); + + // NOTE: Use scf.while may exist dynamic shape problem. + // shape > 64: tiling n * 64, reduction n dim + // remain 64: tiling 2 * n, reduction half parts + if (shape[0] > 64) { + // Reshape to 2D tensor with shape [tile, N / tile] + // Call lowering leading dimension reduction + SmallVector tiledShape = {shape[0] >> 6, 64}; + SmallVector reshapedInputs; + for (auto input : inputs) { + auto inputType = cast(input.getType()); + Value reshape = rewriter.create( + loc, RankedTensorType::get(tiledShape, inputType.getElementType()), + input, ArrayRef{{0, 1}}); + reshapedInputs.push_back(reshape); + } + lastRes = lowerLeadingDimension(reshapedInputs, op, rewriter); + } else { + lastRes = inputs; + } + + if (inputs.size() == 1 && leadingInputType.getElementType().isInteger(1)) { + // TODO: Can optimized to only 8 elements + return lowerBool1DInput(rewriter, loc, leadingInputType.getElementType(), + lastRes, op); + } + + assert(!isInputsIncludeI1Type(inputs) && + "I1 type inputs not supported for multi-op reductions: " + "byte-unaligned element access requires special handling"); + // TODO: Implement i1 type support for reduction operations by handling + // byte-unaligned element access in address calculation(lowerBool1DInput). + + Region &combineOp = op.getRegion(); + auto createExtractSliceOp = [&](Value val, + SmallVector static_offsets, + SmallVector static_size, + SmallVector static_stride) { + auto inputType = cast(val.getType()); + return rewriter.create( + loc, RankedTensorType::get(static_size, inputType.getElementType()), + val, ValueRange(), /*sizes*/ ValueRange(), + /*strides*/ ValueRange(), static_offsets, static_size, static_stride); + }; + + for (int32_t i = tileSize >> 1; i >= 1; i >>= 1) { + auto idx = rewriter.create(loc, i); + + SmallVector binaryInputs, binaryAcc; + for (auto &val : lastRes) { + auto curRes = createExtractSliceOp(val, SmallVector{0}, + SmallVector{i}, + SmallVector{1}); + auto RHS = createExtractSliceOp(val, SmallVector{i}, + SmallVector{i}, + SmallVector{1}); + binaryInputs.push_back(RHS); + binaryAcc.push_back(curRes); + } + lastRes = accumulate(binaryInputs, binaryAcc, combineOp, rewriter); + } + // Collapse the shape of the last result to a scalar tensor + std::transform( + lastRes.begin(), lastRes.end(), lastRes.begin(), [&](auto val) { + auto inputType = cast(val.getType()); + return rewriter.create( + loc, RankedTensorType::get({}, inputType.getElementType()), val, + ArrayRef{}); + }); + + return lastRes; + } + + SmallVector + lowerLeadingDimension(ValueRange inputs, linalg::ReduceOp op, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + Region &combineOp = op.getRegion(); + + auto leadingInputType = cast(inputs[0].getType()); + auto shape = leadingInputType.getShape(); + + // Initialize accumulators as empty tensors of shape [shape[1], ..] + SmallVector accShape(shape.begin() + 1, shape.end()); + + auto results = op.getResults(); + SmallVector acc(results.size()); + + // Build offsets, sizes, and strides for ExtractSlice + SmallVector static_offsets(shape.size(), 0); + + SmallVector sizeVal({1}); + sizeVal.insert(sizeVal.end(), shape.begin() + 1, shape.end()); + + SmallVector strides(shape.size(), 1); + // {1,shape[axis+1],..shape[rank]} -> + // {shape[axis+1],..shape[rank]} + SmallVector reassociation(shape.size() - 1); + // The first group: [0, 1] + reassociation[0].resize(2); + std::iota(reassociation[0].begin(), reassociation[0].end(), 0); + // The remaining groups: [2, ..., shape.size()-1] + for (size_t i = 2; i < shape.size(); ++i) { + reassociation[i - 1].push_back(i); + } + + std::transform(inputs.begin(), inputs.end(), acc.begin(), [&](auto val) { + auto inputType = cast(val.getType()); + auto extract_tensor = rewriter.create( + loc, RankedTensorType::get(sizeVal, inputType.getElementType()), val, + /*offset*/ ValueRange(), /*sizes*/ ValueRange(), + /*strides*/ ValueRange(), + /*static_offsets*/ + static_offsets, sizeVal, strides); + + // {1,shape[axis+1],..shape[rank]} -> + // {shape[axis+1],..shape[rank]} + return rewriter.create(loc, extract_tensor, + reassociation); + }); + + // scf.for loop bounds and step + Value lowerBound = rewriter.create(loc, 1); + Value upperBound = rewriter.create(loc, shape[0]); + Value step = rewriter.create(loc, 1); + + auto forOp = rewriter.create( + loc, lowerBound, upperBound, step, acc, + [&](OpBuilder &b, Location loc, Value iv, ValueRange iterArgs) { + SmallVector currAcc = iterArgs; + + // Build offsets, sizes, and strides for ExtractSlice + SmallVector dynOffsets = {iv}; + for (size_t j = 1; j < shape.size(); ++j) { + dynOffsets.push_back(b.create(loc, 0)); + } + + // iv is a Value, so build dynamic offsets for ExtractSlice + SmallVector subInputs(inputs.size()); + + std::transform( + inputs.begin(), inputs.end(), subInputs.begin(), [&](auto val) { + auto inputType = cast(val.getType()); + auto extract_tensor = b.create( + loc, + RankedTensorType::get(sizeVal, inputType.getElementType()), + val, dynOffsets, /*sizes*/ ValueRange(), + /*strides*/ ValueRange(), + /*static_offsets*/ + SmallVector(shape.size(), ShapedType::kDynamic), + sizeVal, strides); + + // {1,shape[axis+1],..shape[rank]} -> + // {shape[axis+1],..shape[rank]} + return rewriter.create( + loc, extract_tensor, reassociation); + }); + currAcc = accumulate(subInputs, currAcc, combineOp, b); + b.create(loc, currAcc); + }); + + // Extract result tensors from forOp + return forOp.getResults(); + } + + uint32_t getAxis(linalg::ReduceOp op) const override { + // For linalg.reduce, the axis is always 0. + auto dims = op.getDimensions(); + assert(dims.size() == 1 && "Expected a single dimension"); + return dims[0]; + } + + SmallVector getInputs(linalg::ReduceOp op) const override { + // For linalg.reduce, we return the inputs directly. + return op.getInputs(); + } +}; + +struct ArithRemFRewrite : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + + LogicalResult rewriteRemFOp(linalg::GenericOp op, + PatternRewriter &rewriter) const { + auto loc = op->getLoc(); + auto input1 = op.getInputs()[0]; + auto input2 = op.getInputs()[1]; + + auto divResult = buildLinalgElementwise( + rewriter, loc, ValueRange{input1, input2}); + auto truncResult = buildLinalgElementwise( + rewriter, loc, ValueRange{divResult}); + auto mulResult = buildLinalgElementwise( + rewriter, loc, ValueRange{truncResult, input2}); + auto subResult = buildLinalgElementwise( + rewriter, loc, ValueRange{input1, mulResult}); + + rewriter.replaceOp(op, subResult); + return success(); + } + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + auto regionOps = getRegionOps(op); + if (regionOps.size() != 1) { + return failure(); + } + + auto bodyOp = regionOps[0]; + if (!isa(bodyOp)) + return failure(); + + return rewriteRemFOp(op, rewriter); + } +}; + +struct I1ExtUIOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + + auto regionOps = triton::getRegionOps(op); + + if (regionOps.size() != 1 || !isa(regionOps.front())) + return rewriter.notifyMatchFailure(op, "only rewrite i1 extension op\n"); + + auto extOp = cast(regionOps.front()); + + if (!extOp->getOperandTypes()[0].isInteger(1)) + return rewriter.notifyMatchFailure(op, "only rewrite i1 extension op\n"); + + Location loc = op.getLoc(); + + auto input = op.getInputs()[0]; + auto inputType = cast(input.getType()); + auto resultType = cast(op->getResultTypes()[0]); + + Type f16Type = rewriter.getF16Type(); + auto empty = + rewriter.create(loc, resultType.getShape(), f16Type); + + auto castType = RankedTensorType::get(inputType.getShape(), f16Type); + auto f16Reseult = + rewriter.create(loc, castType, input, empty) + ->getResult(0); + + auto rank = resultType.getRank(); + SmallVector indexingMaps( + 2, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iterators( + rank, mlir::utils::IteratorType::parallel); + auto dstEmpty = rewriter.create( + loc, resultType.getShape(), resultType.getElementType()); + auto result = rewriter.create( + loc, TypeRange{resultType}, ValueRange{f16Reseult}, + ValueRange{dstEmpty}, indexingMaps, iterators, + [&](OpBuilder &builder, Location loc, ValueRange args) { + auto src = args[0]; + auto fPToSIOp = builder.create( + loc, resultType.getElementType(), src); + builder.create(loc, ValueRange{fPToSIOp}); + }); + + rewriter.replaceOp(op, result->getResult(0)); + return success(); + } +}; + +struct I1ExtSIOpRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + + auto regionOps = triton::getRegionOps(op); + + if (regionOps.size() != 1 || !isa(regionOps.front())) + return rewriter.notifyMatchFailure(op, "only rewrite i1 extension op\n"); + + auto extOp = cast(regionOps.front()); + + if (!extOp->getOperandTypes()[0].isInteger(1)) + return rewriter.notifyMatchFailure(op, "only rewrite i1 extension op\n"); + + Location loc = op.getLoc(); + + auto input = op.getInputs()[0]; + auto inputType = cast(input.getType()); + auto resultType = cast(op->getResultTypes()[0]); + + Type f16Type = rewriter.getF16Type(); + auto empty = + rewriter.create(loc, resultType.getShape(), f16Type); + + Value zero = rewriter.create( + loc, f16Type, rewriter.getZeroAttr(f16Type)); + Value negOne = rewriter.create( + loc, f16Type, rewriter.getFloatAttr(f16Type, -1.0)); + + auto zeroTensor = + rewriter + .create(loc, ValueRange{zero}, ValueRange{empty}) + .result(); + auto negOneTensor = + rewriter + .create(loc, ValueRange{negOne}, ValueRange{empty}) + .result(); + + auto castType = RankedTensorType::get(inputType.getShape(), f16Type); + auto maskCast = rewriter.create(loc, castType, input, empty); + Value maskmoveResult = + rewriter + .create(loc, negOneTensor.getType(), negOneTensor, + maskCast->getResult(0), zeroTensor) + ->getResult(0); + + auto rank = resultType.getRank(); + SmallVector indexingMaps( + 2, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iterators( + rank, mlir::utils::IteratorType::parallel); + auto dstEmpty = rewriter.create( + loc, resultType.getShape(), resultType.getElementType()); + auto result = rewriter.create( + loc, TypeRange{resultType}, ValueRange{maskmoveResult}, + ValueRange{dstEmpty}, indexingMaps, iterators, + [&](OpBuilder &builder, Location loc, ValueRange args) { + auto src = args[0]; + auto fPToSIOp = builder.create( + loc, resultType.getElementType(), src); + builder.create(loc, ValueRange{fPToSIOp}); + }); + + rewriter.replaceOp(op, result->getResult(0)); + return success(); + } +}; + +struct I1ToF32Rewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + + auto regionOps = triton::getRegionOps(op); + + if (regionOps.size() != 1 || !isa(regionOps.front())) + return rewriter.notifyMatchFailure(op, "only rewrite i1 to f32 op\n"); + + auto siToFP = cast(regionOps.front()); + + if (!siToFP->getOperandTypes()[0].isInteger(1) || + !siToFP->getResultTypes()[0].isF32()) + return rewriter.notifyMatchFailure(op, "only rewrite i1 to f32 op\n"); + + Location loc = op.getLoc(); + + auto input = op.getInputs()[0]; + auto inputType = cast(input.getType()); + auto resultType = cast(op->getResultTypes()[0]); + + auto f32Type = + RankedTensorType::get(inputType.getShape(), rewriter.getF32Type()); + auto empty = rewriter.create(loc, resultType.getShape(), + rewriter.getF32Type()); + + auto f32Result = + rewriter.create(loc, f32Type, input, empty)->getResult(0); + + rewriter.replaceOp(op, f32Result); + return success(); + } +}; + +struct FP32ToI1Rewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(linalg::GenericOp op, + PatternRewriter &rewriter) const override { + + auto regionOps = triton::getRegionOps(op); + + if (regionOps.size() != 1 || !isa(regionOps.front())) + return rewriter.notifyMatchFailure(op, "only rewrite f32 to i1 op\n"); + + auto fpToSI = cast(regionOps.front()); + + if (!fpToSI->getOperandTypes()[0].isF32() || + !fpToSI->getResultTypes()[0].isInteger(1)) + return rewriter.notifyMatchFailure(op, "only rewrite f32 to i1 op\n"); + + Location loc = op.getLoc(); + + auto input = op.getInputs()[0]; + auto inputType = cast(input.getType()); + auto resultType = cast(op->getResultTypes()[0]); + + auto rank = inputType.getRank(); + SmallVector identityMaps( + 3, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iterators( + rank, mlir::utils::IteratorType::parallel); + + auto I1Empty = rewriter.create(loc, resultType.getShape(), + rewriter.getIntegerType(1)); + + Value zeroF32Const = rewriter.create( + loc, rewriter.getF32Type(), APFloat(0.0f)); + + auto zeroTensor = + rewriter + .create( + loc, ValueRange{zeroF32Const}, + ValueRange{rewriter.create( + loc, inputType.getShape(), rewriter.getF32Type())}) + .getResult(0); + + auto result = rewriter.create( + loc, TypeRange{resultType}, ValueRange{input, zeroTensor}, + ValueRange{I1Empty}, identityMaps, iterators, + [&](OpBuilder &builder, Location loc, ValueRange args) { + Value cmp = builder.create( + loc, arith::CmpFPredicate::ONE, args[0], args[1]); + builder.create(loc, cmp); + }); + + rewriter.replaceOp(op, result->getResult(0)); + return success(); + } +}; + +struct AssertOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::AssertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!assertToCf) { + return convertToMK(op, adaptor, rewriter); + } else { + return convertToCF(op, adaptor, rewriter); + } + } + +private: + static LogicalResult convertToCF(triton::AssertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) { + Value condVal = op.getCondition(); + + if (isa(condVal.getType())) { + auto scalarVal = getScalarValue(op.getCondition(), op.getLoc(), rewriter); + condVal = scalarVal ? scalarVal : condVal; + } + assert(condVal && isa(condVal.getType()) && + "Only asserts on scalars are currently supported"); + + if (!condVal.getType().isInteger(1)) { + auto zero = + rewriter.create(op.getLoc(), 0, 32); + auto newCond = rewriter.create( + op.getLoc(), arith::CmpIPredicate::ne, condVal, zero); + condVal = newCond.getResult(); + } + + auto assertMessage = + llvm::formatv("Assertion `{0}` failed", op.getMessage()); + rewriter.create(op.getLoc(), condVal, + assertMessage.str()); + + rewriter.eraseOp(op); + return success(); + } + + static LogicalResult convertToMK(triton::AssertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) { + auto loc = op->getLoc(); + auto condition = adaptor.getCondition(); + auto conditionType = condition.getType(); + Value isTrueCondition; + + if (conditionType.isInteger(1)) { + isTrueCondition = condition; + } else if (auto tensorType = cast(conditionType)) { + if (!tensorType.getElementType().isInteger(1)) { + return rewriter.notifyMatchFailure( + op, "Condition tensor must have i1 element type"); + } + auto trueVal = + rewriter.create(loc, rewriter.getBoolAttr(true)); + auto emptyInit = rewriter.create( + loc, ArrayRef{}, tensorType.getElementType()); + auto filledInit = rewriter + .create(loc, ValueRange{trueVal}, + ValueRange{emptyInit}) + .getResult(0); + int64_t rank = tensorType.getRank(); + SmallVector dimensions; + for (int64_t i = 0; i < rank; ++i) { + dimensions.push_back(i); + } + + auto reduceOp = rewriter.create( + loc, ValueRange{condition}, ValueRange{filledInit}, dimensions, + [&](OpBuilder &b, Location loc, ValueRange args) { + Value inputElem = args[0]; + Value accumulated = args[1]; + auto anyFalse = + b.create(loc, accumulated, inputElem); + b.create(loc, anyFalse->getResult(0)); + }); + isTrueCondition = + rewriter.create(loc, reduceOp.getResult(0)); + } else { + return rewriter.notifyMatchFailure( + op, "Condition must be i1 scalar or tensor with i1 element type"); + } + auto ifOp = rewriter.create( + loc, isTrueCondition, + [&](OpBuilder &builder, Location loc) { + builder.create(loc); + }, + [&](OpBuilder &builder, Location loc) { + builder.create(loc, TypeRange{}, adaptor.getMessage()); + builder.create(loc); + }); + + rewriter.eraseOp(op); + return success(); + } + + bool assertToCf = false; +}; + +/// Convert a dense tensor arith.constant to linalg.fill(scalar, tensor.empty). +/// This is the missing pattern referenced by the comment in LinalgToMKPass: +/// "Lower dense constant to linalg.fill" +struct DenseConstantToFillPattern + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(arith::ConstantOp op, OpAdaptor /*adaptor*/, + ConversionPatternRewriter &rewriter) const override { + auto resultType = dyn_cast(op.getResult().getType()); + if (!resultType) + return failure(); + auto denseAttr = dyn_cast(op.getValue()); + if (!denseAttr) + return failure(); + if (!isa(denseAttr.getElementType())) + return failure(); + if (!denseAttr.isSplat()) + return failure(); + + auto loc = op.getLoc(); + auto elemType = resultType.getElementType(); + auto splatValue = denseAttr.getSplatValue(); + Value scalar = rewriter.create( + loc, elemType, cast(splatValue)); + Value empty = + rewriter.create(loc, resultType.getShape(), elemType); + Value fill = + rewriter.create(loc, scalar, empty).getResult(0); + rewriter.replaceOp(op, fill); + return success(); + } +}; + +// Non-splat constants include reshape/view shape vectors. Preserve every +// element and its row-major index instead of treating the tensor as a fill. +struct DenseConstantToInsertPattern : OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(arith::ConstantOp op, OpAdaptor /*adaptor*/, + ConversionPatternRewriter &rewriter) const override { + auto tensorType = dyn_cast(op.getType()); + if (!tensorType) + return failure(); + + auto denseAttr = dyn_cast(op.getValue()); + if (!denseAttr) + return failure(); + + if (denseAttr.isSplat()) + return failure(); + + Type elemType = tensorType.getElementType(); + if (!isa(elemType)) + return failure(); + + if (!tensorType.hasStaticShape()) + return failure(); + + Location loc = op.getLoc(); + + Value result = + rewriter.create(loc, tensorType.getShape(), elemType); + + SmallVector shape(tensorType.getShape()); + int64_t rank = tensorType.getRank(); + + int64_t linear = 0; + for (Attribute attr : denseAttr.getValues()) { + SmallVector indices(rank); + + int64_t tmp = linear; + for (int64_t d = rank - 1; d >= 0; --d) { + int64_t idx = tmp % shape[d]; + tmp /= shape[d]; + indices[d] = rewriter.create(loc, idx); + } + + Value scalar; + if (isa(elemType)) { + auto intAttr = cast(attr); + scalar = rewriter.create(loc, intAttr.getInt()); + } else { + scalar = rewriter.create(loc, elemType, + cast(attr)); + } + + result = rewriter.create(loc, scalar, result, indices); + + ++linear; + } + + rewriter.replaceOp(op, result); + return success(); + } +}; + +} // namespace + +void mlir::triton::populateLinalgToMKPreProcessPatterns( + RewritePatternSet &patterns) { + // clang-format off + patterns.add, + ArgMinMaxFusionPattern, + AtomicRMWOpRewrite, + AtomicCASOpRewrite>( + patterns.getContext()); + // clang-format on +} + +void mlir::triton::populateLinalgToMKTypeConversionPatterns( + RewritePatternSet &patterns, int precisionMode) { + patterns.add( + patterns.getContext(), precisionMode /* precisionMode */); + patterns.add(patterns.getContext()); + patterns + .add( + patterns.getContext()); + // TODO: if need precision mode + patterns.add, + CastArgMinMaxOpIOToFloatPattern>( + patterns.getContext()); +} + +void mlir::triton::populateLinalgToMKCanonicalizationPatterns( + RewritePatternSet &patterns, int precisionMode) { + // clang-format off + patterns.add, + ScalarGlobalLoadRewrite, + ScalarGlobalStoreRewrite, + ScalarGlobalStoreRewrite>( + patterns.getContext()); + // clang-format on + + if (normalizePrecisionMode(precisionMode) <= 1) + patterns.add( + patterns.getContext()); + else + patterns.add(patterns.getContext()); +} + +void mlir::triton::populateLinalgToMKShapeCanonicalizationPatterns( + RewritePatternSet &patterns, int precisionMode) { + patterns.add( + patterns.getContext(), precisionMode /* precisionMode */); +} + +void mlir::triton::populateLinalgToMKConversionPatterns( + RewritePatternSet &patterns) { + patterns.add(patterns.getContext()); + // After NormalizeReduceInitToIdentityPattern and si-to-fp + patterns.add(patterns.getContext()); + patterns.add( + patterns.getContext()); +} diff --git a/third_party/wafer/lib/Conversion/LinalgToMK/LinalgToMKPass.cpp b/third_party/wafer/lib/Conversion/LinalgToMK/LinalgToMKPass.cpp new file mode 100755 index 00000000..7cc7fb5a --- /dev/null +++ b/third_party/wafer/lib/Conversion/LinalgToMK/LinalgToMKPass.cpp @@ -0,0 +1,156 @@ +//===------------------- LinalgToMKPass.cpp -------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Conversion/LinalgToMK/LinalgToMK.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/Transforms/Transforms.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" + +#define DEBUG_TYPE "linalg-to-mk" + +using namespace mlir; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_LINALGTOMK +#include "magic-kernel/Conversion/LinalgToMK/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class LinalgToMKPass : public triton::impl::LinalgToMKBase { + using LinalgToMKBase::LinalgToMKBase; + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + { + // Fusion, identity reduction init, etc. Other ops which need to decompose + // into multiple integer type operators also are converted. + RewritePatternSet fusionPatterns(&getContext()); + triton::populateLinalgToMKPreProcessPatterns(fusionPatterns); + if (failed(applyPatternsGreedily(moduleOp, std::move(fusionPatterns)))) { + signalPassFailure(); + } + } + + { + RewritePatternSet typePatterns(&getContext()); + triton::populateLinalgToMKTypeConversionPatterns(typePatterns, + precisionPriority ? 2 : precisionMode); + if (failed(applyPatternsGreedily(moduleOp, std::move(typePatterns)))) { + signalPassFailure(); + } + } + + { + // Layout transformation, and other canonicalization + RewritePatternSet canonicalizePatterns(&getContext()); + triton::populateLinalgToMKCanonicalizationPatterns(canonicalizePatterns, + precisionPriority ? 2 : precisionMode); + if (failed(applyPatternsGreedily(moduleOp, + std::move(canonicalizePatterns)))) { + signalPassFailure(); + } + } + + { + RewritePatternSet shapePatterns(&getContext()); + triton::populateLinalgToMKShapeCanonicalizationPatterns( + shapePatterns, precisionPriority ? 2 : precisionMode); + if (failed(applyPatternsGreedily(moduleOp, std::move(shapePatterns)))) { + signalPassFailure(); + } + } + + { + // Target dependent conversion patterns + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + tensor::TensorDialect, bufferization::BufferizationDialect, + memref::MemRefDialect, mk::MagicKernelDialect>(); + target.addDynamicallyLegalOp([&](linalg::ReduceOp op) { + auto regionBlock = op.getBody(); + auto reduceOps = llvm::map_to_vector(regionBlock->without_terminator(), + [](Operation &op) { return &op; }); + if (reduceOps.size() != 1) + return false; + // TODO: Config according backend + // TODO: Optimize for i1 reduction. i1 reduction is not supported + // because memref.subviews may cause the offset to be inside the byte. + auto inputType = + cast(op.getInputs().front().getType()); + + // NOTE: Assume has done integer to float conversion + return !isReduceToElementWiseOpAndTypeSupportedByTarget( + reduceOps.front(), inputType.getElementType(), + inputType.getNumElements(), inputType.getRank()); + }); + + // Reduce op conversion will generate arith/math tensor type op + target.addDynamicallyLegalDialect( + [](Operation *op) { + // Lower dense constants to fill/insert, including index shapes. + if (auto constOp = dyn_cast(op)) { + if (!isa(constOp.getResult().getType())) { + return true; + } + + if (auto denseAttr = + dyn_cast(constOp.getValue())) { + if (isa(denseAttr.getElementType())) { + return false; + } + } + return true; + } + + bool operateOnTensors = + llvm::all_of(op->getOperandTypes(), [](Type type) { + return isa(type); + }); + + return !operateOnTensors; + }); + + triton::populateLinalgToMKConversionPatterns(patterns); + // FIXME: Fixed pass pipeline order to avoid repeatedly adding + // ElementwiseToLinalg patterns + linalg::populateElementwiseToLinalgConversionPatterns(patterns); + if (failed( + applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + } + } +}; + +} // namespace + +std::unique_ptr> triton::createLinalgToMKPass() { + return std::make_unique(); +} + +std::unique_ptr> +triton::createLinalgToMKPass(LinalgToMKOptions &options) { + return std::make_unique(options); +} diff --git a/third_party/wafer/lib/Conversion/MKPipeline/CMakeLists.txt b/third_party/wafer/lib/Conversion/MKPipeline/CMakeLists.txt new file mode 100644 index 00000000..bbf5b56e --- /dev/null +++ b/third_party/wafer/lib/Conversion/MKPipeline/CMakeLists.txt @@ -0,0 +1,18 @@ +add_triton_library(MKPipeline + MKPipelinePass.cpp + MKLoopBoundCanonicalizePass.cpp + DEPENDS + MKPipelinePassIncGen + MagicKernelTableGen + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRPass + MLIRSCFDialect + MLIRMemRefDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonSharedUtils +) diff --git a/third_party/wafer/lib/Conversion/MKPipeline/MKLoopBoundCanonicalizePass.cpp b/third_party/wafer/lib/Conversion/MKPipeline/MKLoopBoundCanonicalizePass.cpp new file mode 100644 index 00000000..e16508f3 --- /dev/null +++ b/third_party/wafer/lib/Conversion/MKPipeline/MKLoopBoundCanonicalizePass.cpp @@ -0,0 +1,229 @@ +//===----------------------------------------------------------------------===// +// MKLoopBoundCanonicalizePass +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Conversion/MKPipeline/Passes.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_MKLOOPBOUNDCANONICALIZEPASS +#include "magic-kernel/Conversion/MKPipeline/Passes.h.inc" +} // namespace triton +} // namespace mlir + +using namespace mlir; + +namespace { + +std::optional getConstantInt(Value value) { + APInt intValue; + if (!matchPattern(value, m_ConstantInt(&intValue))) + return std::nullopt; + return intValue.getSExtValue(); +} + +Value getLoopInductionValue(Value value) { + if (auto cast = value.getDefiningOp()) + return cast.getIn(); + return value; +} + +bool sameValue(Value lhs, Value rhs) { + return getLoopInductionValue(lhs) == getLoopInductionValue(rhs); +} + +Value createIntegerConstant(PatternRewriter &rewriter, Location loc, Type type, + int64_t value) { + if (type.isIndex()) + return rewriter.create(loc, value); + return rewriter.create( + loc, type, rewriter.getIntegerAttr(type, value)); +} + +bool matchAddOfIvAndStep(arith::AddIOp addOp, Value ivLike, + int64_t &stepValue) { + Value lhs = addOp.getLhs(); + Value rhs = addOp.getRhs(); + + if (sameValue(lhs, ivLike)) { + if (auto cst = getConstantInt(rhs)) { + stepValue = *cst; + return true; + } + } + + if (sameValue(rhs, ivLike)) { + if (auto cst = getConstantInt(lhs)) { + stepValue = *cst; + return true; + } + } + + return false; +} + +bool matchMinOfAddAndBound(arith::MinSIOp minOp, arith::AddIOp &addOp, + int64_t &boundValue) { + if ((addOp = minOp.getLhs().getDefiningOp())) { + if (auto cst = getConstantInt(minOp.getRhs())) { + boundValue = *cst; + return true; + } + } + + if ((addOp = minOp.getRhs().getDefiningOp())) { + if (auto cst = getConstantInt(minOp.getLhs())) { + boundValue = *cst; + return true; + } + } + + return false; +} + +bool matchMaxOfMinAndIv(arith::MaxSIOp maxOp, Value ivLike, + arith::MinSIOp &minOp) { + if (sameValue(maxOp.getLhs(), ivLike)) { + minOp = maxOp.getRhs().getDefiningOp(); + return static_cast(minOp); + } + + if (sameValue(maxOp.getRhs(), ivLike)) { + minOp = maxOp.getLhs().getDefiningOp(); + return static_cast(minOp); + } + + return false; +} + +bool loopHasOnlyFullTiles(scf::ForOp forOp, int64_t clampUpper, + int64_t tileStep) { + auto lower = getConstantInt(forOp.getLowerBound()); + auto upper = getConstantInt(forOp.getUpperBound()); + auto loopStep = getConstantInt(forOp.getStep()); + if (!lower || !upper || !loopStep) + return false; + + if (*loopStep <= 0 || tileStep <= 0) + return false; + + // Keep the first pattern intentionally conservative: the tile size must be + // the loop step and the clamp bound must be the loop upper bound. + if (tileStep != *loopStep || clampUpper != *upper) + return false; + + if (*lower >= *upper) + return true; + + return ((*upper - *lower) % *loopStep) == 0; +} + +bool isTailCmp(arith::CmpIOp cmpOp, Value size, int64_t tileStep) { + if (cmpOp.getPredicate() != arith::CmpIPredicate::slt) + return false; + + Value lhs = cmpOp.getLhs(); + Value rhs = cmpOp.getRhs(); + auto matchConstStep = [&](Value value) { + auto cst = getConstantInt(value); + return cst && *cst == tileStep; + }; + + return (lhs == size && matchConstStep(rhs)) || + (rhs == size && matchConstStep(lhs)); +} + +// pattern-1 +// %end = arith.addi %iv, %step +// %clamped = arith.minsi %end, %ub +// %safe = arith.maxsi %clamped, %iv +// %size = arith.subi %safe, %iv +// %is_tail = arith.cmpi slt, %size, %step +// scf.if %is_tail { ... zero fill ... } +// +// 当 scf.for 满足常量 iv in [lb, ub) step step,并能证明 iv + step <= ub +// 对所有迭代成立,就把 %size 替换成 %step,%is_tail 替换成 false,交给后续 +// canonicalize/DCE 删除 scf.if。 +struct FullTileTailSizePattern : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(arith::SubIOp subOp, + PatternRewriter &rewriter) const override { + auto forOp = subOp->getParentOfType(); + if (!forOp) + return failure(); + + Value ivLike = subOp.getRhs(); + if (getLoopInductionValue(ivLike) != forOp.getInductionVar()) + return failure(); + + auto maxOp = subOp.getLhs().getDefiningOp(); + if (!maxOp) + return failure(); + + arith::MinSIOp minOp; + if (!matchMaxOfMinAndIv(maxOp, ivLike, minOp)) + return failure(); + + arith::AddIOp addOp; + int64_t clampUpper = 0; + if (!matchMinOfAddAndBound(minOp, addOp, clampUpper)) + return failure(); + + int64_t tileStep = 0; + if (!matchAddOfIvAndStep(addOp, ivLike, tileStep)) + return failure(); + + if (!loopHasOnlyFullTiles(forOp, clampUpper, tileStep)) + return failure(); + + SmallVector tailCmps; + for (Operation *user : llvm::make_early_inc_range(subOp->getUsers())) { + if (auto cmpOp = dyn_cast(user); + cmpOp && isTailCmp(cmpOp, subOp.getResult(), tileStep)) + tailCmps.push_back(cmpOp); + } + + for (arith::CmpIOp cmpOp : tailCmps) { + rewriter.setInsertionPoint(cmpOp); + Value falseValue = rewriter.create( + cmpOp.getLoc(), rewriter.getBoolAttr(false)); + rewriter.replaceOp(cmpOp, falseValue); + } + + rewriter.setInsertionPoint(subOp); + Value fullTileSize = createIntegerConstant(rewriter, subOp.getLoc(), + subOp.getType(), tileStep); + rewriter.replaceOp(subOp, fullTileSize); + return success(); + } +}; + +void populateMKLoopBoundCanonicalizePatterns(RewritePatternSet &patterns) { + patterns.add(patterns.getContext()); +} + +struct MKLoopBoundCanonicalizePass + : public triton::impl::MKLoopBoundCanonicalizePassBase< + MKLoopBoundCanonicalizePass> { + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + RewritePatternSet patterns(&getContext()); + populateMKLoopBoundCanonicalizePatterns(patterns); + + if (failed(applyPatternsGreedily(getOperation(), std::move(patterns)))) + signalPassFailure(); + } +}; + +} // namespace diff --git a/third_party/wafer/lib/Conversion/MKPipeline/MKPipelinePass.cpp b/third_party/wafer/lib/Conversion/MKPipeline/MKPipelinePass.cpp new file mode 100644 index 00000000..bdef2f77 --- /dev/null +++ b/third_party/wafer/lib/Conversion/MKPipeline/MKPipelinePass.cpp @@ -0,0 +1,1965 @@ +//===----------------------------------------------------------------------===// +// MKPipelinePass — software-pipeline scf.for loops with mk.dot +// +//===----------------------------------------------------------------------===/ + +#include "magic-kernel/Conversion/MKPipeline/Passes.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/SCF/Transforms/Transforms.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/IRMapping.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "mlir/Transforms/RegionUtils.h" +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/DenseSet.h" +#include "llvm/ADT/SmallPtrSet.h" +#include +#include + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_MKPIPELINEPASS +#include "magic-kernel/Conversion/MKPipeline/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { +using namespace mlir; +using namespace mlir::scf; + +// ───────────────────────────────────────────────────────────────────────────── +// 工具函数 +// ───────────────────────────────────────────────────────────────────────────── + +static bool sanitizeScfIfs(Operation *container) { + if (!container) + return true; + + bool ok = true; + container->walk([&](scf::IfOp ifOp) { + Location loc = ifOp.getLoc(); + for (Region ®ion : ifOp->getRegions()) { + if (region.empty()) { + OpBuilder b(ifOp.getContext()); + b.createBlock(®ion); + } + + Block &block = region.front(); + + // dedup yields:一次性删除"非最后一个" scf.yield。 + SmallVector yields; + for (auto y : block.getOps()) + yields.push_back(y); + for (size_t i = 0; i + 1 < yields.size(); ++i) + yields[i]->erase(); + + // 补齐 terminator。 + if (block.empty() || !isa(block.back())) { + OpBuilder b(&block, block.end()); + if (ifOp.getNumResults() == 0) { + b.create(loc); + } else { + // 带结果的 ifOp 缺 yield 是上游构造错误。无法凭空补齐合法 + // 默认值;标记失败,让 caller 触发 signalPassFailure。 + ifOp.emitError( + "scf.if with results is missing a terminating scf.yield; " + "MKPipelinePass produced malformed IR"); + ok = false; + } + } + } + }); + return ok; +} + +// Returns the effective pipeline depth for `forOp`. +// n <= 1 → caller should skip pipelining (return 1). +// n >= 2 → force 2 (ping-pong) — the only depth currently supported. +// +// Rationale: the kernel-body `mk.barrier` at the end of each pipelined +// iteration drains every outstanding DMA+compute, so additional prefetch +// buffers add latency without exposing more parallelism. Until the barrier +// placement is reworked, clamp to 2. +static int getEffectiveNumStages(scf::ForOp forOp, int globalDefault, + int maxStages) { + int n = globalDefault; // default 2 + if (forOp->hasAttr("tt.num_stages")) + n = mlir::cast(forOp->getAttr("tt.num_stages")).getInt(); + + if (n > maxStages) { + // forOp->emitWarning("tt.num_stages is greater than max-stages, clamping to + // 2"); + llvm::errs() << "warning: tt.num_stages is greater than max-stages, " + "clamping to 2\n"; + n = 2; + } else if (n <= 1) { + n = 1; + } + return n; +} + +static bool isAllocBacked(Value v); +static void remapMixedFoldResults(ArrayRef in, + SmallVectorImpl &out, + const IRMapping &rm); +static MemRefType inferSubviewResultType(memref::SubViewOp sv, + MemRefType srcType, + ArrayRef offsets, + ArrayRef sizes, + ArrayRef strides); + +// A loop is treated as "pipelineable" (effectively innermost) even if it has +// nested scf.for loops, provided those nested loops are "trivial" — they +// contain no mk::DotOp and no alloc-backed memref::CopyOp. Trivial inner +// loops (e.g., the scalar row-fill loops that broadcast softmax row-max/sum +// into a 2D buffer) are cloned intact into the compute phase by +// cloneComputeOps and do not interfere with multi-buffering. +static bool isInnermostForOp(scf::ForOp forOp) { + bool hasSignificantNestedFor = false; + forOp.walk([&](scf::ForOp inner) -> WalkResult { + if (inner == forOp) + return WalkResult::advance(); + bool significant = false; + inner.walk([&](Operation *op) -> WalkResult { + if (isa(op)) { + significant = true; + return WalkResult::interrupt(); + } + if (auto copyOp = dyn_cast(op)) { + if (isAllocBacked(copyOp.getTarget())) { + significant = true; + return WalkResult::interrupt(); + } + } + return WalkResult::advance(); + }); + if (significant) { + hasSignificantNestedFor = true; + return WalkResult::interrupt(); + } + return WalkResult::advance(); + }); + return !hasSignificantNestedFor; +} + +// ───────────────────────────────────────────────────────────────────────────── +// Fill-once-per-bank:把 OOB-zero 的 linalg.fill 提到 K-loop 之外 +// +// 动机 +// ---- +// 对未对齐的 GEMM,stage1 IR 形如: +// scf.for %k = ... { +// %a = memref.alloc() : memref<1024x64xf16> // base alloc +// %sv_in = subview %a[0,0][%M_tail, 64] // in-bound view +// scf.if %needFillA { linalg.fill ins(0.0) outs(%a) } // zero整 base +// memref.copy %ddr_sv, %sv_in // 仅写 in-bound +// mk.dot ... reads %a ... // 读整 slot +// } +// 多 buffer 化之后,base alloc 被替换成 `memref<2x1024x64xf16>` 的两个 slot; +// 原 IR 的 fill 被搬进 (a) prologue scf.if、(b) kernel-load scf.for 的每轮 +// iter,导致 K-loop 每轮都对整 slot 写 0(mm_kernel 1024×64 单 fill 就是 +// 1024×64×2B = 128KB SPM 写)。 +// +// 关键观察:%needFillA 是 per-PID 常量(M_tail 由 PID 决定,K-loop 内不 +// 变),fill 的写区域 = 整 base alloc,copy 的写区域 ⊆ in-bound 视图, +// **K-loop 期间无任何其它 op 写 OOB 区域**。所以"每个 bank 的 OOB 只需 +// 在循环外被零化一次,之后所有 K-iter 的 copy 都不会污染它"。 +// +// 实现 +// ---- +// 1) collectPipelinedLoads 给 LoadGroup 增加 3 条字段: +// matchedZeroFillIf —— 直接持有匹配到的 scf.if; +// fillCoversBase —— fill outs 直接是 baseAlloc 时为 true(结构 +// 合格性,是把 fill 提到外面的【必要】前提); +// preInitEmitted —— emitPerBankPreInits 真正成功发射 pre-init +// 后才置 true。clonePreludeIntoLoad 据此决定 +// 是否跳过原 fill-if(不能用 fillCoversBase, +// 否则 fallback 路径会漏发 fill —— 见 addmm +// K-tail 案例的修复说明)。 +// 2) PipelineRewriter::run 在 prologue 之前调用 emitPerBankPreInits: +// 对每个合格 multi-buffer 发射一条 scf.if (cond) { fill 全部 N 个 slot }, +// 把"每轮 K-iter 一次 fill"压缩为"K-loop 之外仅 N 次 fill"。 +// 3) clonePreludeIntoLoad 在 preInitEmitted 时跳过 matchedZeroFillIf, +// 使 prologue / kernel-load 路径都不再发射 fill;compute 路径仍由 +// skipInComputeSet 跳过同一条 if。 +// +// 安全性 +// ------ +// * fillCoversBase 只在 fill 写满整 base alloc 时为 true:保证"slot \ copy +// 的 OOB 区域 ⊆ fill 区域",预填充能完全覆盖。fill 写 partial subview +// 的 group 一律 fallback 到原 per-iter 行为。 +// * 条件值必须由 forOp 之外定义(改造点 A 的 LICM 应已实现)。若个别条件 +// 未被提走(典型:addmm 的 ori(M_tail, K_tail) 因 K_tail 依赖 K-IV +// 而仍在 forOp 内),emitPerBankPreInits 跳过该 group 不发射 pre-init, +// preInitEmitted 保持 false,clonePreludeIntoLoad 自然 fallback 仍发 +// 原 per-iter fill。这一双闸结构是改造点 E 的关键安全保证: +// - 第一闸(fillCoversBase):fill 写区域是否能整 slot 覆盖 OOB; +// - 第二闸(cond 是否循环不变 → preInitEmitted):fill 触发频度 +// 是否能被"一次性"等价替代。两闸都通过才允许 hoist。 +// * 多个 LoadGroup 共享同一 base alloc 时,按 multi-buffer 聚合所有 +// conditions 取 OR:因为每个原 fill 都写满整 slot,OR 等价于"任一 +// condition 真则该 slot 整个置零"。 +// +// 收益(mm_kernel 1024×64×64, K=8, numStages=2 为例): +// K-loop 内 fill 数:8 → 0 +// loop 外一次性 fill 数:0 → 2 (slot 0 + slot 1,guarded by %needFill) +// 总 fill 写量:8×128KB = 1MB → 2×128KB = 256KB(4× 减少) +// 且 OOB-fill 与所有 K-iter 的 RDMA/dot 解耦,可被流水掩盖。 +// ───────────────────────────────────────────────────────────────────────────── + +// ───────────────────────────────────────────────────────────────────────────── +// Pre-LICM:把 `forOp` body 中循环不变的纯 op 提到 forOp 之前 +// +// 动机 +// ---- +// MKPipeline 后端会把"边界 mask 的 condition 链"(典型形态: +// %a = arith.addi %M_base, %c1024 : index +// %b = arith.index_cast %arg_M : i32 to index +// %c = arith.minsi %a, %b : index +// %d = arith.maxsi %c, %M_base : index +// %e = arith.subi %d, %M_base : index +// %f = arith.cmpi slt, %e, %c1024 : index +// scf.if %f { linalg.fill ... } +// )连同 zero-fill scf.if 一起作为 prelude,独立 clone 进 +// (a) prologue scf.if +// (b) kernel-load 路径(clonePreludeIntoLoad) +// 而 cloneComputeOps 走原 forOp body 时,由于 subview 的动态尺寸操作数 +// (上链中的 sub 结果 %e)并非 load-only(它同时被 copy 的 subview 和 +// fill 的 scf.if 使用),整条链在 compute 路径上又会再被 clone 一次。 +// 三次发射的同一组 6 个 arith op 在 LL IR 阶段会真实变成多份 +// `add/index_cast/minsi/maxsi/sub/cmpi`,并各自驱动一份 alloca/store +// 描述符链(详见 stage2_ok/ll_0.mlir 主循环 bb8/bb10/bb12 重复段)。 +// +// 但这条链的全部输入(M_base/N_base 来自 PID、arg_M/arg_N 来自 func +// 形参、常数)在整个 K 维 scf.for 里都是循环不变量。把它们提到 forOp +// 之前一次性算完,prologue / kernel-load / kernel-compute 三处 clone +// 时直接引用同一个循环外 SSA value,副本就只剩一份。 +// +// 实现 +// ---- +// 迭代式 LICM:每轮挑出 body 内"纯 + 无 region + 所有 operand 均定义 +// 在 body 之外(含来自 forOp 之外的 BlockArgument / 函数 arg / 常量 +// op)"的 op,整体 moveBefore(forOp);下一轮再扫描,处理刚被解锁的 +// 上层依赖(比如 cmpi 在 minsi/maxsi 提走之后才被识别为不变)。 +// +// 与 PipelineRewriter 的协作 +// -------------------------- +// 提之后: +// * collectIfConditionDefChain:def->getBlock() != forOp.getBody() +// 立即 return,preludeOps 仅含 zero-fill scf.if 自身; +// * cloneDefChainInLoopBody:同样早返,subview 动态 size operand +// 直接复用循环外 SSA; +// * clonePreludeIntoLoad / cloneComputeOps 不需要任何改动。 +// +// 安全性 +// ------ +// 只搬 isMemoryEffectFree 且无 region 的 op;这覆盖了边界 mask 链涉及 +// 的 arith.addi/subi/minsi/maxsi/cmpi/index_cast,并自动排除: +// - memref.alloc / memref.copy / linalg.fill(有内存副作用) +// - scf.if / scf.while(有 region) +// - mk.dot(写 acc,有副作用) +// 这些 op 必须留在 body 内才能被 PipelineRewriter 正确流水化。 +// +// 提升不会越过 forOp 的支配关系(forOp 之前的位置严格支配 forOp body +// 内任何 use),所以 SSA 合法性自然保留。 +// ───────────────────────────────────────────────────────────────────────────── +static size_t hoistLoopInvariantPureOps(scf::ForOp forOp) { + Block *body = forOp.getBody(); + size_t moved = 0; + bool changed = true; + while (changed) { + changed = false; + SmallVector hoistList; + for (Operation &op : body->without_terminator()) { + // 仅搬纯 op(无内存副作用 + 不会 trap 的算术)。memref.alloc / + // memref.copy / linalg.fill / mk.dot 等都不满足。 + if (!isMemoryEffectFree(&op)) + continue; + // 含 region 的 op(scf.if / scf.while / scf.for / linalg.generic + // 等)不在这里搬:它们的 body 内可能引用循环 IV / iter_args, + // 整体上提会破坏 SSA。需要搬只能逐个分析其 region。 + if (op.getNumRegions() != 0) + continue; + // 没有 result 的纯 op 在 LICM 视角下没有"被使用 = 必须重算"的 + // 价值,跳过避免无谓搬移。 + if (op.getNumResults() == 0) + continue; + + bool invariant = true; + for (Value v : op.getOperands()) { + if (Operation *def = v.getDefiningOp()) { + if (def->getBlock() == body) { + invariant = false; + break; + } + } else if (auto ba = dyn_cast(v)) { + // forOp 自身的 BlockArgument(IV / iter_args)一定挂在 body + // 上,命中此分支即视为"循环依赖"。其它来源的 BlockArgument + // (函数参数、外层 region 的入参)owner != body,落到 else + // 之后被认为是不变的输入,符合预期。 + if (ba.getOwner() == body) { + invariant = false; + break; + } + } + } + if (invariant) + hoistList.push_back(&op); + } + for (Operation *op : hoistList) { + op->moveBefore(forOp); + ++moved; + changed = true; + } + } + return moved; +} + +// ───────────────────────────────────────────────────────────────────────────── +// Bank index 推进:把 `(x + 1) % numStages` 这条 +// index → i64 → addi → remsi → index +// round-trip 折成位运算 / 单条算子,且全程留在 index 域。 +// +// 现状(pre-fix):每个流水化 for 都会发射两组(写槽 + 读槽)形如 +// %a = arith.index_cast %x : index to i64 +// %b = arith.addi %a, %c1_i64 +// %c = arith.remsi %b, %c2_i64 +// %d = arith.index_cast %c : i64 to index +// 主循环每轮 8 个 op、epilogue 每个 stage 4 个 op,连同 epilogue 的 +// `%c1_i64 / %c2_i64 / nVal_i64` 函数级常量与 `arith.index_cast` +// 一起把 LL IR 描述符链拉长 + 把 srem 这种除法器单元留在主循环里。 +// +// 观察: +// 1) `getEffectiveNumStages` 注释里写明实际只可能返回 1 或 2, +// `n <= 1` 在 runOnOperation 里被早 return。所以走到 +// PipelineRewriter::run 时 `numStages == 2` 是不变量。 +// 2) 写/读 slot 索引 ∈ {0, 1},`(x + 1) % 2` 等价 `x ^ 1`。 +// 3) `arith.xori` / `arith.andi` 直接接受 index 操作数, +// 不必绕 i64。 +// +// 因此为 3 个 numStages 分支分别给最便宜的实现,未来要放宽 +// numStages 支持时也无需再回头改: +// * numStages == 2 → arith.xori %x, %c1_idx +// * numStages 为 2 的幂 (>2) → arith.andi (addi %x %c1) (N-1) +// * 其它任意 N → arith.remsi (addi %x %c1) %N +// (仍在 index 域,省掉 cast 来回) +// +// `oneIdx` 由 caller 传入(复用 PipelineRewriter::run 里已经构造好 +// 的 `c1 : index` 常量),避免重复发射常量 op。 +// ───────────────────────────────────────────────────────────────────────────── +static Value makeBankAdvance(IRRewriter &rewriter, Location loc, Value cur, + int numStages, Value oneIdx) { + assert(numStages >= 2 && "bank advance only meaningful for >=2 stages"); + if (numStages == 2) { + // 0 ↔ 1 toggle:单条 xori,无加法、无除法、无 cast。 + return rewriter.create(loc, cur, oneIdx); + } + Value next = rewriter.create(loc, cur, oneIdx); + if (llvm::isPowerOf2_64(static_cast(numStages))) { + Value mask = rewriter.create(loc, numStages - 1); + return rewriter.create(loc, next, mask); + } + Value modulo = rewriter.create(loc, numStages); + return rewriter.create(loc, next, modulo); +} + +// 判断一个 memref 值是否最终来自 memref::AllocOp(SPM 缓冲)。 +// 通过递归穿透所有 view op(不改数据、只改类型/形状/布局的操作)。 +// ────────────────────────────────────────────────────────────────────────── +// ⚠️ 风险:此函数必须覆盖 ALL view op 类型。如果遗漏某个 view op: +// - SPM→SPM 的 copy 会被误判为 DDR→SPM → 错误多缓冲化 → 精度错误 +// - DDR→SPM 的 copy 可能找不到 base alloc → pipeline 跳过 → 依赖顺序错误 +// 所有 memref 的 ViewLikeOpInterface 实现均在此覆盖。新增 view op +// 时务必同步更新。 +// ────────────────────────────────────────────────────────────────────────── +static bool isAllocBacked(Value v) { + if (!v) + return false; + if (v.getDefiningOp()) + return true; + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getSource()); + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getSource()); + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getSrc()); + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getSrc()); + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getSource()); + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getSource()); + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getSource()); + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getSource()); + if (auto op = v.getDefiningOp()) + return isAllocBacked(op.getIn()); + return false; +} + +// 与 isAllocBacked 配套:沿 view 链找到最终的 memref::AllocOp。 +// 覆盖范围必须与 isAllocBacked 保持同步。 +static memref::AllocOp findBaseAlloc(Value v) { + if (!v) + return {}; + if (auto a = v.getDefiningOp()) + return a; + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getSource()); + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getSource()); + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getSrc()); + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getSrc()); + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getSource()); + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getSource()); + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getSource()); + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getSource()); + if (auto op = v.getDefiningOp()) + return findBaseAlloc(op.getIn()); + return {}; +} + +// 一组与流水线化 copy 配对的"前置准备 op": +// - copy 本身 +// - 紧邻 copy 之前、对同一 base alloc 做"清零"的 scf.if (含其 fill), +// 以及该 if 的 condition 仅由该 copy 之前的 cmp/ori 链构成 +// 这些 op 在生成 prologue/kernel-load 时必须与 copy 一起迁移到 load 分支, +// 并把 fill 的 outs 重映射到 multi-buffer slot;compute 阶段则跳过它们, +// 否则会在 load 之后把已经载入的数据清零,破坏 OOB 区域的 0 语义。 +struct LoadGroup { + Operation *copy; // 可以是 memref::CopyOp 或 linalg::CopyOp + // 标记该copy的目标alloc是否在循环外分配(如cu_seqlens偏移缓冲)。 + // 循环外分配的缓冲不需要多缓冲,但copy本身仍需从compute移到load阶段。 + bool isOuterAlloc = false; + // prelude 中的所有 op(含 scf.if、其内部 fill、condition 计算的 cmp/ori 等) + // 全部位于 forOp body 顶层、copy 之前。collectPipelinedLoads 一并填充。 + // 在 load 阶段需要按顺序 clone 全部 preludeOps 到 multi-buffer slot + // 上下文中。 + SmallVector preludeOps; + // preludeOps 中"必须在 compute 阶段被跳过"的 op 子集。 + // 当前仅含 zero-fill scf.if 自身:它的副作用 (linalg.fill) 在 load 阶段已 + // 写入到 slot 的 OOB 区域,若 compute 阶段重新执行同一 if(指向原始 + // baseAlloc 的 view 在 compute 时已被 mapping 指向 extract slot),将 + // 把刚 load 的有效数据再次清零。 + // + // condition 计算链 (cmp/min/max/ori 等) 是纯 op,且经常在 compute 体内 + // 仍有非 prelude 的 use(例如作为 subview 的动态尺寸/偏移)。它们必须 + // 保留在 compute 阶段,否则原 forOp 被 erase 后那些 use 会指向已销毁的 + // SSA value,触发 "operation destroyed but still has uses" crash。 + llvm::SmallPtrSet skipInCompute; + + // 直接持有匹配到的 zero-fill scf.if(preludeOps 末尾那条), + // 用于在 PipelineRewriter::run 中识别"可一次性预填充"的 group: + // - matchedZeroFillIf 非空 = 找到了一组 fill+copy; + // - fillCoversBase = true = fill 的 outs 直接是 baseAlloc 本身(无 + // subview 链),意味着 fill 写满整个 base alloc。这是"一次性预填充" + // 的【结构性】前提(详见文件级注释 [改造点 E]),但仅此还不够。 + // - preInitEmitted = true = emitPerBankPreInits 真的为该 group 在 + // K-loop 之外发射了 pre-init scf.if。只有这一标志为 true, + // clonePreludeIntoLoad 才允许跳过原 fill-if;否则必须 fallback 回 + // per-iter fill 行为。 + // + // 解耦动机(addmm-495-5333-71 精度回归):fillCoversBase 只看 fill 写 + // 区域,不看 fill condition 是否循环不变。当 condition 依赖 K-tail + // (如 addmm 中的 ori(M_tail, K_tail),K_tail = arg8 - arg19*32 不变量 + // 不成立),emitPerBankPreInits 会安全跳过;但若 clonePreludeIntoLoad + // 仍按 fillCoversBase 跳过原 fill,则 prologue / kernel-load 两条 + // load 路径都不再发 fill,OOB 区永久未清零,mk.dot 读到 SPM 残留 + // 数据 → 精度错误。 + scf::IfOp matchedZeroFillIf; + bool fillCoversBase = false; + bool preInitEmitted = false; +}; + +// 判定 ifOp 是否仅做 "fill base 0",且填充的 base 与 copy 的 target 同源。 +// scf.if %58 { +// linalg.fill ins(%cst : f16) outs(%alloc_9 : memref<128x1024xf16>) +// } +// fill 的 input 是 constant,且值为 0。 +// fill 的 output 是 memref::AllocOp。 +// fill 的 output 与 copy 的 target 同源。 +static bool isZeroFillIfFor(scf::IfOp ifOp, memref::AllocOp baseAlloc) { + if (!ifOp || !baseAlloc) + return false; + if (ifOp.getNumResults() != 0) + return false; + // 只接受 else 区域为空,或 else 区域仅含 scf.yield 的 if。 + // 带副作用的 else 被迁移到 load 阶段会改变 OOB 之外的 slot 内容, + // 故直接放弃匹配此 if,保留原 IR 行为。 + if (!ifOp.getElseRegion().empty()) { + Block &elseBlock = ifOp.getElseRegion().front(); + for (Operation &op : elseBlock) { + if (!isa(op)) + return false; + } + } + // 仅允许 then 区域非空、else 区域为空或仅 yield。 + Region &thenRegion = ifOp.getThenRegion(); + if (thenRegion.empty()) + return false; + Block &thenBlock = thenRegion.front(); + // 找到唯一一个 linalg.fill;其它非 yield op 视为有副作用,拒绝。 + linalg::FillOp fillOp; + for (Operation &op : thenBlock) { + if (isa(op)) + continue; + auto f = dyn_cast(&op); + if (!f) + return false; + if (fillOp) + return false; + fillOp = f; + } + if (!fillOp) + return false; + if (fillOp.getOutputs().size() != 1) + return false; + + Value fillVal = fillOp.getInputs()[0]; + auto cstOp = fillVal.getDefiningOp(); + if (!cstOp) + return false; + + Attribute attr = cstOp.getValue(); + bool isZero = false; + if (auto fa = dyn_cast(attr)) + isZero = fa.getValue().isZero(); + else if (auto ia = dyn_cast(attr)) + isZero = ia.getValue().isZero(); + if (!isZero) + return false; + + Value out = fillOp.getOutputs()[0]; + // 允许 out 直接是 base,或经过 subview/reinterpret_cast 链回到 base。 + return findBaseAlloc(out) == baseAlloc; +} + +// 沿 condition 反向收集仅由 cmp/ori/and/xor/constant 等纯计算构成的 def-chain, +// 且所有 op 都位于 forOp body 顶层。失败时返回 false(不当作可迁移 prelude)。 +static bool +collectIfConditionDefChain(Value cond, scf::ForOp forOp, Operation *boundary, + SmallVectorImpl &out, + llvm::SmallPtrSetImpl &seen) { + Operation *def = cond.getDefiningOp(); + if (!def) + return true; // BlockArg / 常量外引用,认为无须迁移 + if (def->getBlock() != forOp.getBody()) + return true; // 来自循环外,无需克隆 + if (def->isBeforeInBlock(boundary) == false) + return false; + if (!seen.insert(def).second) + return true; + // 只允许这些"纯计算" op 作为 if condition 的来源。 + if (!isa(def)) + return false; + for (Value opnd : def->getOperands()) + if (!collectIfConditionDefChain(opnd, forOp, boundary, out, seen)) + return false; + out.push_back(def); + return true; +} + +static Value getCopyTarget(Operation *copyOp) { + if (auto mcpy = dyn_cast(copyOp)) + return mcpy.getTarget(); + if (auto lcpy = dyn_cast(copyOp)) + return lcpy.getOutputs()[0]; + return Value(); +} +static Value getCopySource(Operation *copyOp) { + if (auto mcpy = dyn_cast(copyOp)) + return mcpy.getSource(); + if (auto lcpy = dyn_cast(copyOp)) + return lcpy.getInputs()[0]; + return Value(); +} + +static SmallVector collectPipelinedLoads(scf::ForOp forOp) { + SmallVector groups; + // 标记被收集为 prelude 的 op,避免不同 copy 之间错误共享。 + llvm::SmallPtrSet claimed; + + for (Operation &op : forOp.getBody()->without_terminator()) { + // 识别 memref.copy 或 linalg.copy 作为 DDR→SPM 加载 + Operation *copyOp = nullptr; + Value target, source; + if (auto mcpy = dyn_cast(&op)) { + copyOp = mcpy.getOperation(); + target = mcpy.getTarget(); + source = mcpy.getSource(); + } else if (auto lcpy = dyn_cast(&op)) { + copyOp = lcpy.getOperation(); + target = lcpy.getOutputs()[0]; + source = lcpy.getInputs()[0]; + } + if (!copyOp) + continue; + if (!isAllocBacked(target)) + continue; + if (isAllocBacked(source)) + continue; // skip SPM→SPM copies; only DDR→SPM loads are pipelined + memref::AllocOp baseAlloc = findBaseAlloc(target); + if (!baseAlloc) + continue; + + LoadGroup grp; + grp.copy = copyOp; + grp.isOuterAlloc = (baseAlloc->getBlock() != forOp.getBody()); + + // 在 copy 之前向上扫描,找到第一个对 baseAlloc 做清零的 scf.if。 + // 一旦遇到任何对 baseAlloc 有写副作用的其他 op,立刻停止匹配(保守)。 + scf::IfOp matchedIf; + for (Operation *cur = copyOp->getPrevNode(); cur != nullptr; + cur = cur->getPrevNode()) { + if (claimed.contains(cur)) + continue; + if (auto ifOp = dyn_cast(cur)) { + if (isZeroFillIfFor(ifOp, baseAlloc)) { + matchedIf = ifOp; + break; + } + } + // 任何对 baseAlloc 的潜在写入(除了我们正在找的 fill-if)都断开匹配。 + bool writesBase = false; + cur->walk([&](Operation *inner) { + if (writesBase) + return; + + if (auto store = dyn_cast(inner)) + if (findBaseAlloc(store.getMemRef()) == baseAlloc) + writesBase = true; + if (auto cpy = dyn_cast(inner)) + if (cpy.getOperation() != copyOp && + findBaseAlloc(cpy.getTarget()) == baseAlloc) + writesBase = true; + if (auto lcpy = dyn_cast(inner)) + if (lcpy.getOperation() != copyOp && + findBaseAlloc(lcpy.getOutputs()[0]) == baseAlloc) + writesBase = true; + if (auto fill = dyn_cast(inner)) { + for (Value o : fill.getOutputs()) + if (findBaseAlloc(o) == baseAlloc) + writesBase = true; + } + }); + if (writesBase) + break; + } + + if (matchedIf) { + // 收集 if condition 的 def-chain(必须在 matchedIf 之前的同一 block + // 内)。 + SmallVector condChain; + llvm::SmallPtrSet seen; + if (collectIfConditionDefChain(matchedIf.getCondition(), forOp, matchedIf, + condChain, seen)) { + for (Operation *c : condChain) { + grp.preludeOps.push_back(c); + // 注意:condition 链是纯计算 op,不放入 claimed —— 同一组 cmp/ + // min/max 链可能被多个 copy 共用(例如 alloc_3 与 alloc_7 共用 + // 同一个 OOB mask 的 cmpi/maxsi/minsi)。允许重复 clone 到各 + // load 分支(pure op,无副作用)。同样不放入 skipInCompute: + // 它们必须在 compute 阶段保留,因为 compute 体内其它 op 可能 + // 仍引用它们的结果(典型:作为 dynamic subview size)。 + } + grp.preludeOps.push_back(matchedIf); + // 仅 scf.if (含 fill 的副作用) 才需要在 compute 阶段被跳过 + + // 在 copy 之间唯一归属。 + claimed.insert(matchedIf); + grp.skipInCompute.insert(matchedIf); + + // 记录 matchedIf + 判定 fillCoversBase。 + // fillCoversBase 仅在 fill 的 outs 直接是 baseAlloc 本身(无任何 + // subview / reinterpret_cast 包装)时为 true:因为后续我们要把 + // fill 一次性提到 K-loop 之外、按 bank 写满整个 slot —— 这要求 + // 原 fill 也是写满整个 base alloc。若 fill 只覆盖 base alloc 的 + // 一部分,则 "OOB(slot) ⊆ fill region" 不成立,按 bank 预填充 + // 会漏写 fill 之外、copy 之外的那部分,破坏 OOB-0 不变量。 + grp.matchedZeroFillIf = matchedIf; + for (Operation &innerOp : + matchedIf.getThenRegion().front().getOperations()) { + if (auto fillOp = dyn_cast(&innerOp)) { + if (fillOp.getOutputs().size() == 1 && + fillOp.getOutputs()[0] == baseAlloc.getResult()) { + grp.fillCoversBase = true; + } + break; + } + } + } + // 若 condition chain 不可识别,则放弃迁移此 fill:保留原 IR 行为 + // (在 OOB 边界 tile 下仍可能错,但至少不会引入更糟的 IR)。 + } + + groups.push_back(std::move(grp)); + } + return groups; +} + +// 兼容旧接口:返回所有 copies。 +static SmallVector collectDDRToSPMCopies(scf::ForOp forOp) { + SmallVector copies; + for (auto copyOp : forOp.getOps()) { + if (isAllocBacked(copyOp.getTarget())) + copies.push_back(copyOp); + } + return copies; +} + +static std::optional tryGetMemRefStaticBytes(MemRefType ty) { + if (!ty.hasStaticShape()) + return std::nullopt; + Type elemTy = ty.getElementType(); + if (!elemTy.isIntOrFloat()) + return std::nullopt; + unsigned bitWidth = elemTy.getIntOrFloatBitWidth(); + int64_t numElems = 1; + for (int64_t d : ty.getShape()) + numElems *= d; + int64_t bitTotal = numElems * static_cast(bitWidth); + return (bitTotal + 7) / 8; +} + +static void +cloneDefChainInLoopBody(Operation *root, scf::ForOp forOp, IRMapping &mapping, + OpBuilder &builder, + llvm::DenseSet &alreadyCloned, + llvm::DenseMap &cloneOf); + +static void dedupIfYields(scf::IfOp ifOp, OpBuilder &builder) { + Location ifLoc = ifOp.getLoc(); + for (Region ®ion : ifOp->getRegions()) { + if (region.empty()) + continue; + for (Block &block : region) { + // 一次性收集并删除所有"非最后一个" scf.yield。 + // 删除后剩余至多 1 个 yield,无需重复扫描。 + SmallVector yields; + for (Operation &op : block) + if (auto y = dyn_cast(&op)) + yields.push_back(y); + for (size_t i = 0; i + 1 < yields.size(); ++i) + yields[i]->erase(); + } + bool needTerminator = false; + for (Block &block : region) { + if (!block.mightHaveTerminator()) { + needTerminator = true; + break; + } + } + if (needTerminator) + scf::IfOp::ensureTerminator(region, builder, ifLoc); + } + for (Region ®ion : ifOp->getRegions()) { + if (region.empty()) + continue; + region.walk([&](scf::IfOp nested) { + if (nested != ifOp) + dedupIfYields(nested, builder); + }); + } +} + +static void +cloneDefChainInLoopBody(Operation *root, scf::ForOp forOp, IRMapping &mapping, + OpBuilder &builder, + llvm::DenseSet &alreadyCloned, + llvm::DenseMap &cloneOf); + +// 在 generic clone(`op`) 之前,把 op region 内部所引用的外层 forOp body +// 顶层 SSA value 提前 clone 到当前 builder/mapping 中。 +// +// 对 region-bearing op(scf.for / scf.while / scf.if / linalg.generic …): +// builder.clone(op, mapping) 会原样复制 region 中的子 op,并在子 op 的 +// operand 上做 mapping.lookupOrDefault。如果某外层 value 没在 mapping, +// clone 出来的子 op 会继续指向原始 SSA value——等到原 forOp 被 erase 时, +// "operation destroyed but still has uses" 立即触发。 +// +// 因此对所有 region-bearing op,clone 之前必须先把 region 中"defined above +// 且属于 forOp body 顶层"的 def 全部 cloneDefChainInLoopBody 进来。 +static void +ensureRegionExternalsCloned(Operation *op, scf::ForOp forOp, IRMapping &mapping, + OpBuilder &builder, + llvm::DenseSet &alreadyCloned, + llvm::DenseMap &cloneOf) { + if (!op || op->getNumRegions() == 0) + return; + llvm::SetVector outerValues; + for (Region ®ion : op->getRegions()) + mlir::getUsedValuesDefinedAbove(region, outerValues); + + SmallVector outerDefs; + llvm::SmallPtrSet outerDefsSeen; + for (Value v : outerValues) { + Operation *def = v.getDefiningOp(); + if (!def) + continue; + if (def->getBlock() != forOp.getBody()) + continue; + if (def == op) + continue; + if (mapping.contains(v)) + continue; + if (outerDefsSeen.insert(def).second) + outerDefs.push_back(def); + } + for (Operation *def : outerDefs) + cloneDefChainInLoopBody(def, forOp, mapping, builder, alreadyCloned, + cloneOf); +} + +static void +cloneDefChainInLoopBody(Operation *root, scf::ForOp forOp, IRMapping &mapping, + OpBuilder &builder, + llvm::DenseSet &alreadyCloned, + llvm::DenseMap &cloneOf) { + if (!root || root->getBlock() != forOp.getBody()) + return; + if (alreadyCloned.contains(root)) { + Operation *cached = cloneOf.lookup(root); + if (cached == root) + return; + assert(cached && "clone cache missing for already-cloned op"); + for (auto [from, to] : llvm::zip(root->getResults(), cached->getResults())) + mapping.map(from, to); + return; + } + if (auto allocOp = dyn_cast(root)) { + if (mapping.contains(allocOp.getResult())) { + alreadyCloned.insert(root); + cloneOf[root] = root; + return; + } + } + for (Value opnd : root->getOperands()) { + if (Operation *def = opnd.getDefiningOp()) { + cloneDefChainInLoopBody(def, forOp, mapping, builder, alreadyCloned, + cloneOf); + } else if (auto ba = dyn_cast(opnd)) { + // forOp body 的 BlockArg(induction var / iter_args)必须由 caller + // 在 mapping 里预先映射;否则 generic clone 会引用原 forOp 的 + // BlockArg,导致跨 region 的 SSA 非法。这里同时保留 assert + // (debug 构建快速发现)和运行时 emit error(release 构建可见)。 + if (ba.getOwner() == forOp.getBody() && !mapping.contains(opnd)) { + assert(false && "IRMapping must cover loop IV / iter_args before " + "cloning op into pipeline stage"); + root->emitError("MKPipelinePass: block-argument operand of ") + << root->getName() + << " is not present in IRMapping; pipeline " + "transform produced invalid IR"; + return; + } + } + } + + // 对 region-bearing root,clone 前先把 region 内引用的外层 forOp body + // 顶层 def 提前 clone 到 mapping 中,避免 generic clone 出来的子 op + // 仍然指向原 forOp 内即将被 erase 的 SSA value。详见 + // ensureRegionExternalsCloned 的注释。 + ensureRegionExternalsCloned(root, forOp, mapping, builder, alreadyCloned, + cloneOf); + + // 在 generic clone 之前,确保 mapping 中没有任何针对 root + // 自身 region 内 block-arg 的预映射。如果存在(例如其他路径误把这些 + // BlockArgument 当作"已映射"加入),Region::cloneInto 会跳过 + // addArgument,导致克隆后的 region 入口 block 没有 block args, + // 触发 "region control flow edge from parent operands to Region #N: + // source has K operands, but target successor needs 0" verifier 错误。 + for (Region ®ion : root->getRegions()) + for (Block &block : region) + for (BlockArgument arg : block.getArguments()) + if (mapping.contains(arg)) + mapping.erase(arg); + + // SubViewOp/ExpandShapeOp/CollapseShapeOp:当 source 的类型因 mapping 而改变 + // (例如 alloc → multi-buffer slot,引入 dynamic offset/stride),结果类型 + // 必须重新推导,否则 builder.clone 会原样保留旧 result type,产生 + // "mismatch of result layout" 验证错误。 + Operation *cloned = nullptr; + if (auto sv = dyn_cast(root)) { + SmallVector offs, sizes, strides; + remapMixedFoldResults(sv.getMixedOffsets(), offs, mapping); + remapMixedFoldResults(sv.getMixedSizes(), sizes, mapping); + remapMixedFoldResults(sv.getMixedStrides(), strides, mapping); + Value newSrc = mapping.lookupOrDefault(sv.getSource()); + auto newSrcTy = cast(newSrc.getType()); + MemRefType resultTy = + inferSubviewResultType(sv, newSrcTy, offs, sizes, strides); + cloned = builder.create(sv.getLoc(), resultTy, newSrc, + offs, sizes, strides); + } else if (auto expand = dyn_cast(root)) { + Value newSrc = mapping.lookupOrDefault(expand.getSrc()); + auto newSrcTy = cast(newSrc.getType()); + auto reassoc = expand.getReassociationIndices(); + FailureOr resultTy = memref::ExpandShapeOp::computeExpandedType( + newSrcTy, expand.getResultType().getShape(), reassoc); + if (failed(resultTy)) { + cloned = builder.clone(*root, mapping); + } else { + OpBuilder shapeBuilder(builder.getContext()); + SmallVector mixedOut = getMixedValues( + expand.getStaticOutputShape(), expand.getOutputShape(), shapeBuilder); + SmallVector remappedMixed; + remapMixedFoldResults(mixedOut, remappedMixed, mapping); + cloned = builder.create( + expand.getLoc(), *resultTy, newSrc, reassoc, remappedMixed); + } + } else if (auto collapse = dyn_cast(root)) { + Value newSrc = mapping.lookupOrDefault(collapse.getSrc()); + auto newSrcTy = cast(newSrc.getType()); + auto reassoc = collapse.getReassociationIndices(); + MemRefType resultTy = + memref::CollapseShapeOp::computeCollapsedType(newSrcTy, reassoc); + cloned = builder.create(collapse.getLoc(), + resultTy, newSrc, reassoc); + } else { + cloned = builder.clone(*root, mapping); + if (auto clonedIf = dyn_cast(cloned)) + dedupIfYields(clonedIf, builder); + } + + cloneOf[root] = cloned; + for (auto [from, to] : llvm::zip(root->getResults(), cloned->getResults())) + mapping.map(from, to); + alreadyCloned.insert(root); +} + +static void remapCopySourceProducersInLoop( + memref::CopyOp copyOp, scf::ForOp forOp, IRMapping &mapping, + OpBuilder &builder, llvm::DenseSet &alreadyCloned, + llvm::DenseMap &cloneOf) { + Value src = copyOp.getSource(); + if (!src) + return; + if (Operation *def = src.getDefiningOp()) + cloneDefChainInLoopBody(def, forOp, mapping, builder, alreadyCloned, + cloneOf); +} + +static MemRefType inferSubviewResultType(memref::SubViewOp sv, + MemRefType srcType, + ArrayRef offsets, + ArrayRef sizes, + ArrayRef strides) { + if (sv.getType().getRank() < srcType.getRank()) + return memref::SubViewOp::inferRankReducedResultType( + sv.getType().getShape(), srcType, offsets, sizes, strides); + return cast( + memref::SubViewOp::inferResultType(srcType, offsets, sizes, strides)); +} + +static OpFoldResult remapOfr(OpFoldResult ofr, const IRMapping &rm) { + if (Value v = llvm::dyn_cast_if_present(ofr)) + return rm.lookupOrDefault(v); + return ofr; +} + +static void remapMixedFoldResults(ArrayRef in, + SmallVectorImpl &out, + const IRMapping &rm) { + out.clear(); + for (OpFoldResult ofr : in) + out.push_back(remapOfr(ofr, rm)); +} + +static Value remapDestToSlot(OpBuilder &b, Location loc, + scf::ForOp pipelinedLoop, IRMapping &mapping, + llvm::DenseSet &producerDone, + llvm::DenseMap &cloneOf, + Value origDest, Value slot, + memref::AllocOp baseAlloc) { + if (!baseAlloc || origDest == baseAlloc.getResult()) + return slot; + + SmallVector path; + Value cur = origDest; + Value baseVal = baseAlloc.getResult(); + while (cur != baseVal) { + Operation *def = cur.getDefiningOp(); + if (!def) + return slot; + path.push_back(def); + if (auto sv = dyn_cast(def)) { + cur = sv.getSource(); + continue; + } + if (auto rc = dyn_cast(def)) { + cur = rc.getSource(); + continue; + } + return slot; + } + + Value mapped = slot; + for (Operation *op : llvm::reverse(path)) { + Value srcV; + if (auto sv = dyn_cast(op)) + srcV = sv.getSource(); + else if (auto rc = dyn_cast(op)) + srcV = rc.getSource(); + else + continue; + + for (Value oper : op->getOperands()) { + if (oper == srcV) + continue; + if (Operation *def = oper.getDefiningOp()) { + if (def->getBlock() == pipelinedLoop.getBody()) + cloneDefChainInLoopBody(def, pipelinedLoop, mapping, b, producerDone, + cloneOf); + } + } + + IRMapping rm(mapping); + rm.map(srcV, mapped); + + if (auto sv = dyn_cast(op)) { + SmallVector offs, sizes, strides; + remapMixedFoldResults(sv.getMixedOffsets(), offs, rm); + remapMixedFoldResults(sv.getMixedSizes(), sizes, rm); + remapMixedFoldResults(sv.getMixedStrides(), strides, rm); + Value newSrc = rm.lookupOrDefault(sv.getSource()); + auto newSrcTy = cast(newSrc.getType()); + MemRefType resultTy = + inferSubviewResultType(sv, newSrcTy, offs, sizes, strides); + mapped = b.create(loc, resultTy, newSrc, offs, sizes, + strides); + } else if (isa(op)) { + mapped = b.clone(*op, rm)->getResult(0); + } + } + return mapped; +} + +// ───────────────────────────────────────────────────────────────────────────── +// 流水线重写器 +// ───────────────────────────────────────────────────────────────────────────── + +struct PipelineRewriter { + scf::ForOp forOp; + SmallVector copies; // 可以是 memref::CopyOp 或 linalg::CopyOp + // 与 copies 一一对应,保存每个 copy 的 prelude(含 fill-if + 其 cmp/ori + // 链)。 + SmallVector loadGroups; + // compute 阶段必须跳过的 op 集合 —— 仅含 zero-fill 的 scf.if 自身。 + // 详见 LoadGroup::skipInCompute 的注释。 + llvm::SmallPtrSet skipInComputeSet; + // compute 阶段要跳过的"仅服务于 load"的 op 集合。详见 buildLoadOnlyOps + // 的注释。超集包含 copies、skipInComputeSet 以及它们上游的计算/ + // scf.while 等;核心目的是避免把 scf.while 内部的 memref.copy 在 compute + // 阶段重新执行一次(本已在 load 阶段写入 slot[insertIdx],compute 再写 + // slot[extractIdx] 是纯浪费 DMA)。 + llvm::SmallPtrSet loadOnlyOps; + int numStages; + IRRewriter &rewriter; + Location loc; + + PipelineRewriter(scf::ForOp forOp, SmallVector groups, + int numStages, IRRewriter &rewriter) + : forOp(forOp), loadGroups(std::move(groups)), numStages(numStages), + rewriter(rewriter), loc(forOp.getLoc()) { + for (auto &g : loadGroups) { + copies.push_back(g.copy); + for (Operation *op : g.skipInCompute) + skipInComputeSet.insert(op); + } + buildLoadOnlyOps(); + } + + // Classify forOp.getBody() top-level ops as "load-only": those whose + // every direct top-level user is also load-only. Seeds: + // - Pipelined memref.copy ops. + // - skipInComputeSet (zero-fill scf.if whose side effects on + // pipelined alloc slots are already done in the load phase). + // + // Then iterate until fixed point: an op with at least one use, whose + // users all end at load-only ops, is load-only too. + // + // Motivation: an inner memref.copy inside scf.while (e.g. for partial + // row-wraparound reads from DDR into SPM) would otherwise be cloned + // twice — once into load phase via the pipelined copy's source + // def-chain, and once into compute phase via cloneComputeOps's + // top-level walk. The second clone issues a redundant DMA into the + // extract slot that was already correctly populated by the earlier + // load phase. Skipping the whole load-only sub-graph in compute kills + // the duplicate DMA and the wasted arith that feeds it. + void buildLoadOnlyOps() { + for (Operation *c : copies) + loadOnlyOps.insert(c); + for (Operation *op : skipInComputeSet) + loadOnlyOps.insert(op); + + // Pipelined alloc bases are NOT load-only (compute reads them via + // the extract slot), but they must not block propagation either. + llvm::SmallPtrSet pipelinedAllocBases; + for (Operation *c : copies) + if (auto base = findBaseAlloc(getCopyTarget(c))) + pipelinedAllocBases.insert(base.getOperation()); + + Block *body = forOp.getBody(); + bool changed = true; + while (changed) { + changed = false; + for (Operation &op : body->without_terminator()) { + if (loadOnlyOps.contains(&op)) + continue; + if (pipelinedAllocBases.contains(&op)) + continue; + // Side-effect-only ops (no results) cannot be classified via + // user-propagation — they're seeded explicitly through + // skipInComputeSet when appropriate. + if (op.getNumResults() == 0) + continue; + bool hasUser = false; + bool allUsersLoadOnly = true; + for (Operation *user : op.getUsers()) { + Operation *topUser = user; + while (topUser && topUser->getBlock() != body) + topUser = topUser->getParentOp(); + if (!topUser || topUser->getBlock() != body) { + allUsersLoadOnly = false; + break; + } + hasUser = true; + if (!loadOnlyOps.contains(topUser)) { + allUsersLoadOnly = false; + break; + } + } + if (hasUser && allUsersLoadOnly) { + loadOnlyOps.insert(&op); + changed = true; + } + } + } + } + + // Returns true on success. On failure, `destToMultiBuf` may be partially + // populated but the caller MUST treat the rewrite as failed — downstream + // code looks up every `copyOp.getTarget()` and would dereference null if + // any entry is missing. + bool createMultiBuffers(DenseMap &destToMultiBuf) { + DenseMap baseMemToMulti; + rewriter.setInsertionPoint(forOp); + for (Operation *copyOp : copies) { + Value origDest = getCopyTarget(copyOp); + memref::AllocOp baseAlloc = findBaseAlloc(origDest); + if (!baseAlloc) { + mlir::emitError(loc, + "copy target has no memref.alloc base (isAllocBacked " + "mismatch?)"); + return false; + } + Value baseMem = baseAlloc.getResult(); + // 循环外分配的缓冲(如cu_seqlens偏移缓冲1xi32)不需要多缓冲, + // 只需把copy移到load阶段即可。后续compute/load阶段直接用原始缓冲。 + if (baseAlloc->getBlock() != forOp.getBody()) { + destToMultiBuf[origDest] = baseMem; + continue; + } + Value multi = baseMemToMulti.lookup(baseMem); + if (!multi) { + auto baseType = mlir::cast(baseMem.getType()); + SmallVector newShape; + newShape.push_back(numStages); + for (int64_t d : baseType.getShape()) + newShape.push_back(d); + + auto multiBufType = MemRefType::get(newShape, baseType.getElementType(), + MemRefLayoutAttrInterface{}, + baseType.getMemorySpace()); + + ValueRange dynamicOperands = baseAlloc.getDynamicSizes(); + IntegerAttr alignment = baseAlloc.getAlignmentAttr(); + auto newAlloc = rewriter.create( + loc, multiBufType, dynamicOperands, alignment); + multi = newAlloc.getResult(); + baseMemToMulti[baseMem] = multi; + } + destToMultiBuf[origDest] = multi; + } + return true; + } + + // 将某个 copy 的 prelude(fill-if + 其 condition 的 + // cmp/ori 链)按原顺序 clone 到当前 builder 的位置。调用方需要保证: + // - mapping 已经把 baseAlloc.getResult() 映射到当前 stage 的 slot; + // - mapping 已经覆盖 forOp 的 induction var / iter_args。 + // 注意:fill 的 outs 有可能是 baseAlloc 的派生 view;那些派生 view 在 + // 原循环体内由本身就属于 forOp body 顶层的 op 产生,会被外层 + // cloneDefChainInLoopBody 自动拉入。这里只负责按顺序 clone prelude 自身。 + // 注:参数取非 const ref 是因为 mlir 的 Op 包装类型(scf::IfOp 等) + // 的 accessor 没有 const 重载,访问 grp.matchedZeroFillIf.getOperation() + // 需要非 const 路径。grp 内部的 preludeOps / 其它字段不会被本函数修改。 + void clonePreludeIntoLoad(LoadGroup &grp, OpBuilder &b, IRMapping &mapping) { + Operation *zeroFillIfOp = + grp.matchedZeroFillIf ? grp.matchedZeroFillIf.getOperation() : nullptr; + for (Operation *p : grp.preludeOps) { + // 仅当 emitPerBankPreInits 已经为本 group 在 K-loop + // 之外发射了 pre-init scf.if(preInitEmitted == true)才跳过原 + // fill-if。注意不能改用 fillCoversBase:那只是结构合格性,与 + // "pre-init 是否真发射"无关;emitPerBankPreInits 还会用 condition + // 是否循环不变量做二次过滤(addmm K-tail 类 kernel 在那一步会被 + // 安全跳过),此时 prologue / kernel-load 必须保留原 per-iter + // fill,否则 OOB 区不会被清零,mk.dot 读到 SPM 残留垃圾。 + if (grp.preInitEmitted && p == zeroFillIfOp) + continue; + + // 多个 LoadGroup 可能共享同一段 condition 链(例如 alloc_3 与 alloc_7 + // 共用一个 OOB mask)。若该 op 的所有 result 已经在 mapping 中(即上 + // 一组 prelude 或同 stage 内的前置已经 clone 过),就不再重复 clone, + // 也不要覆盖 mapping —— 否则后续使用者指向新克隆的副本,旧引用 + // 仍然散落在已生成的 IR 中,行为虽然合法但会产生死代码。 + bool allMapped = !p->getResults().empty(); + for (Value r : p->getResults()) { + if (!mapping.contains(r)) { + allMapped = false; + break; + } + } + if (allMapped) + continue; + + // 任一 prelude op 的 operand 若来自 forOp body 的非 prelude/ + // 非 copies 部分,则已经被 cloneDefChainInLoopBody 处理; + // 这里只需顺序 clone。 + // [关键修复] 同样在 generic clone 前清理 mapping 中针对 op 自身 + // region block-arg 的旧映射,避免 Region::cloneInto 跳过 addArgument。 + for (Region ®ion : p->getRegions()) + for (Block &block : region) + for (BlockArgument arg : block.getArguments()) + if (mapping.contains(arg)) + mapping.erase(arg); + Operation *cloned = b.clone(*p, mapping); + // 对 scf.if 的 fill-then-yield 子块,clone 后必须确保 terminator 合法。 + if (auto clonedIf = dyn_cast(cloned)) + dedupIfYields(clonedIf, b); + } + } + + Value getSlot(Value multiBuf, Value slotIdx) { + auto bufType = mlir::cast(multiBuf.getType()); + // 非多缓冲(循环外分配的原始缓冲)直接返回本身 + if (bufType.getShape().empty() || bufType.getShape()[0] != numStages) + return multiBuf; + SmallVector subShape(bufType.getShape().drop_front()); + + SmallVector offsets, sizes, strides; + offsets.push_back(slotIdx); + sizes.push_back(rewriter.getIndexAttr(1)); + strides.push_back(rewriter.getIndexAttr(1)); + for (int64_t d : subShape) { + offsets.push_back(rewriter.getIndexAttr(0)); + sizes.push_back(rewriter.getIndexAttr(d)); + strides.push_back(rewriter.getIndexAttr(1)); + } + MemRefType subType = memref::SubViewOp::inferRankReducedResultType( + subShape, bufType, offsets, sizes, strides); + assert(subType && + "inferRankReducedResultType failed for multi-buffer slot"); + return rewriter.create(loc, subType, multiBuf, offsets, + sizes, strides); + } + + SmallVector cloneComputeOps(IRMapping &mapping) { + llvm::DenseSet producerDone; + llvm::DenseMap computeCloneOf; + for (Operation &op : forOp.getBody()->without_terminator()) { + if (llvm::is_contained(copies, &op)) + continue; + if (auto allocOp = dyn_cast(&op)) { + bool isPipelinedSpmBase = false; + for (Operation *c : copies) { + if (findBaseAlloc(getCopyTarget(c)) == allocOp) { + isPipelinedSpmBase = true; + break; + } + } + if (isPipelinedSpmBase) + continue; + } + // 跳过已在 load 阶段执行过的、带副作用的 prelude + // (zero-fill scf.if)。否则 compute 阶段会再次执行 fill,把刚 + // load 进 slot 的数据清零。 + // + // 注意:condition 链 (cmp/min/max/ori 等) 不放进 skipInComputeSet, + // 因为它们是纯计算 op,compute 体内(例如 scf.while 内的 subview + // 动态尺寸)可能仍有 use;只有整个子图确实"只服务于 load"时才能 + // 跳过它们。这由下面的 loadOnlyOps 全图传播来保证——其他使用者 + // 如果落在 compute 侧,子图里的任何 op 都不会被标成 load-only。 + if (skipInComputeSet.contains(&op)) + continue; + // [Perf] 跳过"仅服务于 load"的 op。这类 op 的所有 top-level 使用者 + // 要么是 pipelined copy,要么是 skipInCompute 的 zero-fill,要么是 + // 它们的上游链(例如 scf.while 计算 DDR 偏移给 copy)。在 load 阶 + // 段已被 cloneDefChainInLoopBody 克隆过,compute 阶段再克隆一次会 + // 让 scf.while 里的 memref.copy 发起一次重复 DMA 打到 extract slot。 + // 见 buildLoadOnlyOps 的注释。 + if (loadOnlyOps.contains(&op)) + continue; + for (Value opnd : op.getOperands()) { + if (Operation *def = opnd.getDefiningOp()) + cloneDefChainInLoopBody(def, forOp, mapping, rewriter, producerDone, + computeCloneOf); + } + // 若 op 自身带 region(scf.for / scf.while / scf.if 等),还要把 + // region 内部引用的外层 forOp body 顶层 def 也提前 clone,否则 + // 下面 rewriter.clone(op, mapping) 出来的 IR 仍会指向原 forOp 中 + // 即将被 erase 的 SSA value,触发 "operation destroyed but still has + // uses"。 + ensureRegionExternalsCloned(&op, forOp, mapping, rewriter, producerDone, + computeCloneOf); + Operation *cloned = nullptr; + if (auto sv = dyn_cast(&op)) { + SmallVector offs, sizes, strides; + remapMixedFoldResults(sv.getMixedOffsets(), offs, mapping); + remapMixedFoldResults(sv.getMixedSizes(), sizes, mapping); + remapMixedFoldResults(sv.getMixedStrides(), strides, mapping); + Value newSrc = mapping.lookupOrDefault(sv.getSource()); + auto newSrcTy = cast(newSrc.getType()); + MemRefType resultTy = + inferSubviewResultType(sv, newSrcTy, offs, sizes, strides); + cloned = rewriter.create( + sv.getLoc(), resultTy, newSrc, offs, sizes, strides); + } else if (auto expand = dyn_cast(&op)) { + Value newSrc = mapping.lookupOrDefault(expand.getSrc()); + auto newSrcTy = cast(newSrc.getType()); + auto reassoc = expand.getReassociationIndices(); + FailureOr resultTy = + memref::ExpandShapeOp::computeExpandedType( + newSrcTy, expand.getResultType().getShape(), reassoc); + if (failed(resultTy)) { + cloned = rewriter.clone(op, mapping); + } else { + OpBuilder shapeBuilder(rewriter.getContext()); + SmallVector mixedOut = + getMixedValues(expand.getStaticOutputShape(), + expand.getOutputShape(), shapeBuilder); + SmallVector remappedMixed; + remapMixedFoldResults(mixedOut, remappedMixed, mapping); + cloned = rewriter.create( + expand.getLoc(), *resultTy, newSrc, reassoc, remappedMixed); + } + } else if (auto collapse = dyn_cast(&op)) { + Value newSrc = mapping.lookupOrDefault(collapse.getSrc()); + auto newSrcTy = cast(newSrc.getType()); + auto reassoc = collapse.getReassociationIndices(); + MemRefType resultTy = + memref::CollapseShapeOp::computeCollapsedType(newSrcTy, reassoc); + cloned = rewriter.create( + collapse.getLoc(), resultTy, newSrc, reassoc); + } else { + // 与 cloneDefChainInLoopBody 中保持一致:在 generic + // clone 之前清掉 mapping 里针对 op 自身 region 内 block-arg 的 + // 旧映射;否则 Region::cloneInto 会跳过 addArgument,导致克隆 + // 后的入口 block 没有 block args,触发 verifier 错误 + // ("region control flow edge ... target successor needs 0")。 + for (Region ®ion : op.getRegions()) + for (Block &block : region) + for (BlockArgument arg : block.getArguments()) + if (mapping.contains(arg)) + mapping.erase(arg); + + cloned = rewriter.clone(op, mapping); + // cloneComputeOps 的 generic clone 路径:同样需要 dedupIfYields。 + // 注意:若此 IfOp 已被 cloneDefChainInLoopBody 提前克隆并写入 + // producerDone,则上面 cloneDefChainInLoopBody 调用会走 cached 路径 + // 不再 clone,此处的 rewriter.clone 不会执行到。 + // 若确实走到此处(顶层直接遇到 IfOp),仍需修复。 + if (auto clonedIf = dyn_cast(cloned)) + dedupIfYields(clonedIf, rewriter); + } + computeCloneOf[&op] = cloned; + for (auto [from, to] : llvm::zip(op.getResults(), cloned->getResults())) + mapping.map(from, to); + producerDone.insert(&op); + } + auto origYield = cast(forOp.getBody()->getTerminator()); + SmallVector yieldedVals; + for (Value v : origYield.getOperands()) + yieldedVals.push_back(mapping.lookupOrDefault(v)); + return yieldedVals; + } + + // 在 prologue 之前、按 (multi-buffer × condition) 一次性 + // 预填充 OOB 区域。返回值:被成功"提到 K-loop 之外"的 LoadGroup 数量 + // (主要用于日志/调试,目前未使用,保留 size_t 以备将来扩展)。 + // + // 安全前提(参见文件级注释 [改造点 E]): + // * grp.fillCoversBase == true:fill 写满整个 baseAlloc; + // * grp.matchedZeroFillIf.getCondition() 在 forOp 之外定义(改造点 A 的 + // LICM 应已把整条 cmp 链提走;若未提走则保守跳过该 group,保持 + // per-iter prelude 行为,正确性不受影响)。 + // + // 输出 IR 形态(每个 multi-buffer 一条 scf.if): + // scf.if (%cond_or_combined) { + // for s = 0 .. numStages-1: + // %slot_s = subview %multi[s, 0, 0][1, full_dims...] + // linalg.fill ins(%cst) outs(%slot_s) + // } + // + // 与同一 multi-buffer 关联的多个 group 共用一条 scf.if,conditions 之间 + // 取 OR:因为每个原 fill 都写满整 base alloc,等效于"任一 condition 为 + // 真时整 slot 都该被零初始化"。 + size_t emitPerBankPreInits(const DenseMap &destToMultiBuf) { + auto condDominatesForOp = [&](Value c) -> bool { + if (!c) + return false; + Operation *def = c.getDefiningOp(); + if (!def) + return true; // BlockArg / func arg / 常量外引用:默认认为支配 + return !forOp->isProperAncestor(def); + }; + + // 收集 (multi -> { (groupIdx, cond, fillConstantOp), ... }),按 multi + // 聚合,发射时合并 conditions(OR 语义),最后把 preInitEmitted=true + // 写回所有共享该 multi 的合格 group。 + struct PerMultiItem { + size_t groupIdx; + Value cond; + arith::ConstantOp cstOp; + }; + SmallVector orderedMulti; + DenseMap> perMulti; + + // 注:grp 取非 const ref 是因为 mlir 的 Op 包装类型 accessor 没有 + // const 重载(getOperation / getCondition / getThenRegion / getTarget + // 都不是 const-qualified)。 + for (size_t i = 0; i < loadGroups.size(); ++i) { + LoadGroup &grp = loadGroups[i]; + if (!grp.fillCoversBase || !grp.matchedZeroFillIf) + continue; + Value multi = destToMultiBuf.lookup(getCopyTarget(grp.copy)); + if (!multi) + continue; + // [关键安全检查] 必须确认 condition 是循环不变量。否则 fill 与 + // 其触发条件本就是 per-iter 行为(典型:addmm 的 K-tail 让 + // ori(M_tail, K_tail) 含 arg19),任何"提到外面做一次"的尝试都 + // 不能复刻原 per-iter 语义 —— 只能 fallback 保留原 fill。 + Value cond = grp.matchedZeroFillIf.getCondition(); + if (!condDominatesForOp(cond)) + continue; + + // 提取 then-region 内的 linalg.fill 常量。 + arith::ConstantOp cstOp; + for (Operation &innerOp : + grp.matchedZeroFillIf.getThenRegion().front().getOperations()) { + if (auto fillOp = dyn_cast(&innerOp)) { + cstOp = fillOp.getInputs()[0].getDefiningOp(); + break; + } + } + if (!cstOp) + continue; + + if (perMulti.find(multi) == perMulti.end()) + orderedMulti.push_back(multi); + perMulti[multi].push_back({i, cond, cstOp}); + } + + if (perMulti.empty()) + return 0; + + rewriter.setInsertionPoint(forOp); + size_t emitted = 0; + for (Value multi : orderedMulti) { + const auto &items = perMulti[multi]; + assert(!items.empty()); + + // 合并所有 condition:cond_0 OR cond_1 OR ... + Value mergedCond = items[0].cond; + for (size_t i = 1; i < items.size(); ++i) + mergedCond = + rewriter.create(loc, mergedCond, items[i].cond); + + // 任选一个常量作为 fill 值(同 multi 上的 fill 都是同一个零常量)。 + arith::ConstantOp anyCstOp = items[0].cstOp; + + auto bufType = mlir::cast(multi.getType()); + SmallVector subShape(bufType.getShape().drop_front()); + + auto preIf = rewriter.create( + loc, mergedCond, + [&](OpBuilder &b, Location l) { + // 在 then-block 中本地克隆常量,避免依赖 anyCstOp 的位置假设 + // (scf.if 的 then-region 与 anyCstOp 通常都在 func 顶层,但 + // 局部克隆使得后续 IR 重写不会被远程引用打扰)。 + Value localCst = b.clone(*anyCstOp)->getResult(0); + for (int s = 0; s < numStages; ++s) { + Value slotIdx = b.create(l, s); + SmallVector offsets, sizes, strides; + offsets.push_back(slotIdx); + sizes.push_back(b.getIndexAttr(1)); + strides.push_back(b.getIndexAttr(1)); + for (int64_t d : subShape) { + offsets.push_back(b.getIndexAttr(0)); + sizes.push_back(b.getIndexAttr(d)); + strides.push_back(b.getIndexAttr(1)); + } + MemRefType subType = + memref::SubViewOp::inferRankReducedResultType( + subShape, bufType, offsets, sizes, strides); + assert(subType && + "inferRankReducedResultType failed (pre-init slot)"); + Value slot = b.create(l, subType, multi, + offsets, sizes, strides); + b.create(l, ValueRange{localCst}, + ValueRange{slot}); + } + b.create(l); + }, + [&](OpBuilder &b, Location l) { b.create(l); }); + dedupIfYields(preIf, rewriter); + // 发射成功后回写到 group:clonePreludeIntoLoad 据此跳过原 fill-if。 + for (const PerMultiItem &it : items) + loadGroups[it.groupIdx].preInitEmitted = true; + ++emitted; + } + return emitted; + } + + LogicalResult run() { + DenseMap destToMultiBuf; + if (!createMultiBuffers(destToMultiBuf) || destToMultiBuf.empty()) { + mlir::emitError(loc, "create multi buffer failed for dest mem"); + return failure(); + } + + Value lb = forOp.getLowerBound(); + Value ub = forOp.getUpperBound(); + Value step = forOp.getStep(); + Type ivType = step.getType(); + + rewriter.setInsertionPoint(forOp); + auto makeIvConst = [&](int64_t v) -> Value { + if (ivType.isIndex()) + return rewriter.create(loc, v); + auto intTy = dyn_cast(ivType); + assert(intTy && "scf.for induction type must be index or integer"); + return rewriter.create(loc, v, intTy.getWidth()); + }; + Value zero = rewriter.create(loc, 0); + Value one = rewriter.create(loc, 1); + // 不再发射 `nVal / i64Type / nVal_i64`:bank rotation + // 的 `(x + 1) % numStages` 全部改走 makeBankAdvance(index 域 + + // 位运算优先),无 i64 round-trip,函数级也少 3 条常量/cast。 + + // 在 prologue 之前对每个合格 multi-buffer 一次性预填充 + // 所有 bank slot 的 OOB 区域。命中后,对应 LoadGroup 的 zero-fill + // scf.if 在 prologue / kernel-load 路径上都会被 clonePreludeIntoLoad + // 跳过;compute 路径仍由 skipInComputeSet 跳过。语义对等于"OOB 区 + // 一旦被零化、K-loop 内任何 op 都不会再触碰它"——参见 collectPipelined + // Loads 中的 fillCoversBase 注释 + run() 函数体顶部的整体说明。 + (void)emitPerBankPreInits(destToMultiBuf); + + // ══════════════════════════════════════════════════════════════════════════ + // 1. Prologue + // ══════════════════════════════════════════════════════════════════════════ + rewriter.setInsertionPoint(forOp); + + for (int i = 0; i < numStages - 1; ++i) { + Value prologueIv; + if (i == 0) { + prologueIv = lb; + } else { + Value iVal = makeIvConst(i); + prologueIv = rewriter.create( + loc, lb, rewriter.create(loc, iVal, step)); + } + Value slotIdx = rewriter.create(loc, i); + + Value inBound = rewriter.create( + loc, arith::CmpIPredicate::slt, prologueIv, ub); + + auto prologueIf = rewriter.create( + loc, inBound, + [&](OpBuilder &b, Location l) { + llvm::DenseSet copyProducersDone; + llvm::DenseMap copyCloneOf; + IRMapping copyMapping; + copyMapping.map(forOp.getInductionVar(), prologueIv); + for (unsigned j = 0; j < forOp.getNumRegionIterArgs(); ++j) + copyMapping.map(forOp.getRegionIterArgs()[j], + forOp.getInitArgs()[j]); + for (auto &grp : loadGroups) { + Operation *copyOp = grp.copy; + Value origTarget = getCopyTarget(copyOp); + Value multi = destToMultiBuf.lookup(origTarget); + Value slot; + memref::AllocOp baseAlloc = findBaseAlloc(origTarget); + if (grp.isOuterAlloc) { + // 循环外分配的缓冲直接用原缓冲 + slot = multi; + } else { + auto bufType = mlir::cast(multi.getType()); + SmallVector subShape(bufType.getShape().drop_front()); + SmallVector offsets, sizes, strides; + offsets.push_back(slotIdx); + sizes.push_back(b.getIndexAttr(1)); + strides.push_back(b.getIndexAttr(1)); + for (int64_t d : subShape) { + offsets.push_back(b.getIndexAttr(0)); + sizes.push_back(b.getIndexAttr(d)); + strides.push_back(b.getIndexAttr(1)); + } + MemRefType subType = + memref::SubViewOp::inferRankReducedResultType( + subShape, bufType, offsets, sizes, strides); + assert(subType && + "inferRankReducedResultType failed (prologue slot)"); + slot = b.create(l, subType, multi, offsets, + sizes, strides); + } + + if (baseAlloc) + copyMapping.map(baseAlloc.getResult(), slot); + clonePreludeIntoLoad(grp, b, copyMapping); + + Value mappedTarget = + remapDestToSlot(b, l, forOp, copyMapping, copyProducersDone, + copyCloneOf, origTarget, slot, baseAlloc); + copyMapping.map(origTarget, mappedTarget); + Value origSource = getCopySource(copyOp); + if (Operation *def = origSource.getDefiningOp()) + cloneDefChainInLoopBody(def, forOp, copyMapping, b, + copyProducersDone, copyCloneOf); + b.clone(*copyOp, copyMapping); + } + b.create(l); + }, + [&](OpBuilder &b, Location l) { b.create(l); }); + dedupIfYields(prologueIf, rewriter); + } + + rewriter.setInsertionPoint(forOp); + // [DISABLED] post-prologue barrier. + // + // `mk.barrier` 只有一个来源:`dsa.DistributedBarrierOp`,语义是跨 tile + // 的全局同步("Synchronizes all work items")。单 tile 的 matmul 调用它 + // 会等所有协作 tile 到达,但实际 grid 里根本没有其他参与者,于是 + // `__Barrier()` 永远不返回——栈跟踪 `txStreamSynchronize → + // Stream::finish → Event::awaitCompletion` 死循环即由此而来。 + // + // 正常(未流水)版本的 tx/LLVM IR 里一次 `tx.barrier` 都没有,完全靠 + // 硬件 stream 的 FIFO 顺序保证 `__Rdma`/`__Gemm`/`__Memset` 之间的 + // program-order 依赖。流水化改写只是 SSA 层面的多缓冲切分,同样不需要 + // 显式 barrier。 + // + // 在 mk 方言补齐 per-tile stream fence 之前,这里不再插入 barrier。 + // rewriter.create(loc); + + // ══════════════════════════════════════════════════════════════════════════ + // 2. Kernel Loop + // ══════════════════════════════════════════════════════════════════════════ + Value kernelLb = rewriter.create( + loc, lb, + rewriter.create(loc, makeIvConst(numStages - 1), step)); + + SmallVector initArgs; + initArgs.push_back( + rewriter.create(loc, numStages - 1)); + initArgs.push_back(zero); + for (Value v : forOp.getInitArgs()) + initArgs.push_back(v); + + auto newForOp = rewriter.create( + loc, kernelLb, ub, step, initArgs, + [&](OpBuilder &b, Location nestedLoc, Value, ValueRange iterArgs) { + b.create(nestedLoc, iterArgs); + }); + auto oldYieldInKernel = + cast(newForOp.getBody()->getTerminator()); + rewriter.setInsertionPoint(oldYieldInKernel); + + Value insertIdx = newForOp.getBody()->getArgument(1); + Value extractIdx = newForOp.getBody()->getArgument(2); + + Value kernelIv = newForOp.getInductionVar(); + Value computeIv = rewriter.create( + loc, kernelIv, + rewriter.create( + loc, makeIvConst(static_cast(numStages - 1)), step)); + IRMapping computeMapping; + computeMapping.map(forOp.getInductionVar(), computeIv); + for (int i = 0; i < forOp.getNumRegionIterArgs(); ++i) + computeMapping.map(forOp.getRegionIterArgs()[i], + newForOp.getBody()->getArgument(3 + i)); + + llvm::DenseSet kernelCopyProducersDone; + llvm::DenseMap kernelCopyCloneOf; + IRMapping copyMapping; + copyMapping.map(forOp.getInductionVar(), kernelIv); + auto origYield = cast(forOp.getBody()->getTerminator()); + for (unsigned j = 0; j < forOp.getNumRegionIterArgs(); ++j) { + Value nextState = origYield.getOperand(j); + if (Operation *def = nextState.getDefiningOp()) + cloneDefChainInLoopBody(def, forOp, computeMapping, rewriter, + kernelCopyProducersDone, kernelCopyCloneOf); + copyMapping.map(forOp.getRegionIterArgs()[j], + computeMapping.lookupOrDefault(nextState)); + } + for (auto &grp : loadGroups) { + Operation *copyOp = grp.copy; + Value origTarget = getCopyTarget(copyOp); + Value multi = destToMultiBuf.lookup(origTarget); + Value slot = getSlot(multi, insertIdx); + memref::AllocOp baseAlloc = findBaseAlloc(origTarget); + + if (baseAlloc) + copyMapping.map(baseAlloc.getResult(), slot); + clonePreludeIntoLoad(grp, rewriter, copyMapping); + + Value mappedTarget = remapDestToSlot( + rewriter, loc, forOp, copyMapping, kernelCopyProducersDone, + kernelCopyCloneOf, origTarget, slot, baseAlloc); + copyMapping.map(origTarget, mappedTarget); + Value origSource = getCopySource(copyOp); + if (Operation *def = origSource.getDefiningOp()) + cloneDefChainInLoopBody(def, forOp, copyMapping, rewriter, + kernelCopyProducersDone, kernelCopyCloneOf); + rewriter.clone(*copyOp, copyMapping); + } + + llvm::DenseSet computeDestProducersDone; + llvm::DenseMap computeDestCloneOf; + for (auto [origDest, multi] : destToMultiBuf) { + memref::AllocOp baseAlloc = findBaseAlloc(origDest); + Value extSlot = getSlot(multi, extractIdx); + if (baseAlloc) + computeMapping.map(baseAlloc.getResult(), extSlot); + Value remapped = remapDestToSlot( + rewriter, loc, forOp, computeMapping, computeDestProducersDone, + computeDestCloneOf, origDest, extSlot, baseAlloc); + if (baseAlloc && origDest != baseAlloc.getResult()) + computeMapping.map(origDest, remapped); + } + + SmallVector clonedYieldVals = cloneComputeOps(computeMapping); + + // [DISABLED] kernel-body barrier。`mk.barrier` 是 distributed barrier, + // 单 tile 调用会 hang。详见 run() 上方 post-prologue barrier 处的说明。 + // rewriter.create(loc); + + // index 域单 op rotation;numStages == 2 时落到一条 + // arith.xori,主循环每轮少 8 个 op(详见 makeBankAdvance 注释)。 + Value nextInsert = makeBankAdvance(rewriter, loc, insertIdx, numStages, + /*oneIdx=*/one); + Value nextExtract = makeBankAdvance(rewriter, loc, extractIdx, numStages, + /*oneIdx=*/one); + + SmallVector yieldVals = {nextInsert, nextExtract}; + for (Value v : clonedYieldVals) + yieldVals.push_back(v); + rewriter.create(loc, yieldVals); + rewriter.eraseOp(oldYieldInKernel); + + // 再收敛 kernel 内所有 scf.if(含嵌套):clone 顺序下仍可能残留双 yield。 + newForOp.walk([&](scf::IfOp op) { dedupIfYields(op, rewriter); }); + + // ══════════════════════════════════════════════════════════════════════════ + // 3. Epilogue + // ══════════════════════════════════════════════════════════════════════════ + rewriter.setInsertionPointAfter(newForOp); + Value hasAnyWork = + rewriter.create(loc, arith::CmpIPredicate::slt, lb, ub); + + auto guardedEpilogue = rewriter.create( + loc, forOp.getResultTypes(), hasAnyWork, /*withElseRegion=*/true); + + { + OpBuilder::InsertionGuard g(rewriter); + OpBuilder thenB = guardedEpilogue.getThenBodyBuilder(); + rewriter.setInsertionPoint(thenB.getInsertionBlock(), + thenB.getInsertionPoint()); + + Value curExtractIdx = newForOp.getResult(1); + + SmallVector epilogueAccs; + for (unsigned i = 2; i < newForOp.getNumResults(); ++i) + epilogueAccs.push_back(newForOp.getResult(i)); + + for (int e = 0; e < numStages - 1; ++e) { + Value eOffsetVal = makeIvConst(static_cast(numStages - 1 - e)); + Value epilogueIv = rewriter.create( + loc, ub, rewriter.create(loc, eOffsetVal, step)); + + // 当原循环实际迭代数 N_orig < numStages - 1 时, + // 部分 epilogue stage 对应的 epilogueIv 会 < lb(即原循环根本没有 + // 跑到这一轮)。此时不能执行该 stage 的 compute(slot 未被 + // load 过),也不能用其结果污染 acc。 + // 用 scf.if (epilogueIv >= lb) 包裹本轮:then 走真实 compute 并 + // yield 新 acc;else 直接透传当前 acc 与 curExtractIdx。 + Value isValid = rewriter.create( + loc, arith::CmpIPredicate::sge, epilogueIv, lb); + + SmallVector stageResultTypes; + for (Value v : epilogueAccs) + stageResultTypes.push_back(v.getType()); + // 同时在 if 里前进 extract 指针,避免后续轮使用了无效的 slot。 + stageResultTypes.push_back(rewriter.getIndexType()); + + auto stageIf = rewriter.create( + loc, stageResultTypes, isValid, /*withElseRegion=*/true); + + // ── then: 真实 compute ───────────────────────────────── + { + OpBuilder::InsertionGuard g(rewriter); + OpBuilder thenB2 = stageIf.getThenBodyBuilder(); + rewriter.setInsertionPoint(thenB2.getInsertionBlock(), + thenB2.getInsertionPoint()); + + IRMapping epilogueMapping; + epilogueMapping.map(forOp.getInductionVar(), epilogueIv); + for (int i = 0; i < forOp.getNumRegionIterArgs(); ++i) + epilogueMapping.map(forOp.getRegionIterArgs()[i], epilogueAccs[i]); + + llvm::DenseSet epilogueDestProducersDone; + llvm::DenseMap epilogueDestCloneOf; + for (auto [origDest, multi] : destToMultiBuf) { + memref::AllocOp baseAlloc = findBaseAlloc(origDest); + Value extSlot = getSlot(multi, curExtractIdx); + if (baseAlloc) + epilogueMapping.map(baseAlloc.getResult(), extSlot); + Value remapped = + remapDestToSlot(rewriter, loc, forOp, epilogueMapping, + epilogueDestProducersDone, epilogueDestCloneOf, + origDest, extSlot, baseAlloc); + if (baseAlloc && origDest != baseAlloc.getResult()) + epilogueMapping.map(origDest, remapped); + } + + SmallVector newAccs = cloneComputeOps(epilogueMapping); + // [DISABLED] epilogue barrier。`mk.barrier` 是 distributed barrier, + // 单 tile 调用会 hang。之前加这一处是想修 "epilogue 后 dealloc/ + // truncf/Wdma 早于 Gemm 完成" 的问题,但实际正常(未流水)版本 + // 根本没有 barrier,靠 stream FIFO 顺序就能保证最后的 + // `tx.fp32_fp16`/`tx.wdma` 在 `tx.gemm` 之后执行。详见 run() + // 上方 post-prologue barrier 处的完整说明。 + // rewriter.create(loc); + + // 同主循环:index 域单 op,numStages==2 → xori。 + Value advancedExtract = makeBankAdvance(rewriter, loc, curExtractIdx, + numStages, /*oneIdx=*/one); + + SmallVector thenYield(newAccs.begin(), newAccs.end()); + thenYield.push_back(advancedExtract); + rewriter.create(loc, thenYield); + } + + // ── else: 透传当前 acc 与 extract idx ───────────────── + { + OpBuilder elseB2 = stageIf.getElseBodyBuilder(); + SmallVector elseYield(epilogueAccs.begin(), + epilogueAccs.end()); + elseYield.push_back(curExtractIdx); + elseB2.create(loc, elseYield); + } + + // 更新 epilogueAccs / curExtractIdx 为 if 的结果,下一轮使用。 + epilogueAccs.clear(); + for (unsigned i = 0; i + 1 < stageIf.getNumResults(); ++i) + epilogueAccs.push_back(stageIf.getResult(i)); + curExtractIdx = stageIf.getResult(stageIf.getNumResults() - 1); + } + thenB.create(loc, epilogueAccs); + } + + { + OpBuilder elseB = guardedEpilogue.getElseBodyBuilder(); + elseB.create(loc, forOp.getInitArgs()); + } + + guardedEpilogue.walk([&](scf::IfOp op) { dedupIfYields(op, rewriter); }); + + // The buffer deallocation pass has been deprecated in favor of the + // ownership-based buffer deallocation pipeline. The deprecated pass has + // some limitations that may cause memory leaks in the resulting IR. + // llvm::DenseSet deallocatedMultis; + // for (auto &entry : destToMultiBuf) { + // if (deallocatedMultis.insert(entry.second).second) + // rewriter.create(loc, entry.second); + // } + + for (unsigned i = 0; i < forOp.getNumResults(); ++i) + forOp.getResult(i).replaceAllUsesWith(guardedEpilogue.getResult(i)); + + rewriter.eraseOp(forOp); + return success(); + } +}; + +// ───────────────────────────────────────────────────────────────────────────── +// Pass 入口 +// ───────────────────────────────────────────────────────────────────────────── + +struct MKPipelinePass + : public mlir::triton::impl::MKPipelinePassBase { + using Base::Base; + + void runOnOperation() override { + ModuleOp module = getOperation(); + IRRewriter rewriter(module.getContext()); + + SmallVector candidates; + module.walk([&](scf::ForOp forOp) { + if (!isInnermostForOp(forOp)) + return; + int n = getEffectiveNumStages(forOp, this->numStages, this->maxStages); + if (n <= 1) + return; + auto groups = collectPipelinedLoads(forOp); + if (groups.empty()) + return; + // Spec requires at least one mk.dot (compute) in a pipelined loop — + // without it, multi-buffering just adds overhead. + bool hasDot = false; + forOp.walk([&](mlir::mk::DotOp) { + hasDot = true; + return WalkResult::interrupt(); + }); + if (!hasDot) + return; + + candidates.push_back(forOp); + }); + + for (auto forOp : candidates) { + // 在拆 prologue / kernel / epilogue 之前先做一次本地 + // LICM。把循环内"边界 mask 的 condition 链"等纯不变 op 上提到 + // forOp 之前,使得 collectPipelinedLoads 收到的 condChain 为空, + // 后续 prologue / kernel-load / kernel-compute 三处 clone 都直接 + // 引用同一份循环外 SSA value,不再重复发射 6 个 arith op + 一串 + // 描述符 alloca/store。详细动机见 hoistLoopInvariantPureOps 注释。 + (void)hoistLoopInvariantPureOps(forOp); + + int n = getEffectiveNumStages(forOp, this->numStages, this->maxStages); + auto groups = collectPipelinedLoads(forOp); + if (failed(PipelineRewriter(forOp, groups, n, rewriter).run())) + signalPassFailure(); + } + + if (!sanitizeScfIfs(module)) + signalPassFailure(); + } +}; + +} // namespace diff --git a/third_party/wafer/lib/Conversion/MKToWafer/CMakeLists.txt b/third_party/wafer/lib/Conversion/MKToWafer/CMakeLists.txt new file mode 100755 index 00000000..bda76f2c --- /dev/null +++ b/third_party/wafer/lib/Conversion/MKToWafer/CMakeLists.txt @@ -0,0 +1,20 @@ +add_triton_library(MKToWafer + MKToWafer.cpp + MKToWaferPass.cpp + + DEPENDS + WaferTableGen + MKToWaferConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRLinalgDialect + MLIRDialectUtils + MLIRIR + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms +) diff --git a/third_party/wafer/lib/Conversion/MKToWafer/MKToWafer.cpp b/third_party/wafer/lib/Conversion/MKToWafer/MKToWafer.cpp new file mode 100755 index 00000000..5b023329 --- /dev/null +++ b/third_party/wafer/lib/Conversion/MKToWafer/MKToWafer.cpp @@ -0,0 +1,2438 @@ +//===--------------------- MKToWafer.cpp -----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file implements the patterns to convert operations from mk dialect to +// wafer dialect. It converts memory operations to RdmaOp/WdmaOp and converts +// mk.dot to wafer.gemm etc. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Conversion/MKToWafer/MKToWafer.h" +#include "instr_def.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Func/Transforms/FuncConversions.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Transforms/Transforms.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/MemRef/Utils/MemRefUtils.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/ValueRange.h" +#include "triton-shared/Utils/FusionHelper.h" +#include "triton-shared/Utils/Utils.h" +#include "triton/Conversion/TritonGPUToLLVM/Utility.h" +#include "wafer/Dialect/IR/WaferDialect.h" +#include "llvm/ADT/TypeSwitch.h" + +// FIXME: triton/Conversion/TritonGPUToLLVM/Utility.h which defined +// TritonLLVMOpBuilder and other utilities has defined DEBUG_TYPE. +#ifdef DEBUG_TYPE +#undef DEBUG_TYPE +#endif +#define DEBUG_TYPE "mk-to-wafer" + +using namespace mlir; +using namespace wafer; + +#define GEN_PASS_CLASSES +#include "wafer/Conversion/MKToWafer/Passes.h.inc" + +namespace { + +//===----------------------------------------------------------------------===// +// Type Conversion +//===----------------------------------------------------------------------===// + +class MKToWaferTypeConverter : public TypeConverter { +public: + MKToWaferTypeConverter() { + // Add conversions for MemRef types to UI64 (representing SPM addresses) + addConversion([](MemRefType type) -> Type { + return IntegerType::get(type.getContext(), 64, IntegerType::Unsigned); + }); + + // Add conversions for Tensor types to UI64 (representing SPM addresses) + addConversion([](TensorType type) -> Type { + return IntegerType::get(type.getContext(), 64, IntegerType::Unsigned); + }); + + // Keep other types as is + addConversion([](Type type) -> Type { return type; }); + } + +private: + MLIRContext *context; +}; + +//===----------------------------------------------------------------------===// +// Utilities +//===----------------------------------------------------------------------===// + +LogicalResult convertLinalgOpToLoops(linalg::LinalgOp op, + ConversionPatternRewriter &rewriter) { + if (failed(linalg::linalgOpToLoops(rewriter, op))) + return rewriter.notifyMatchFailure(op, "operation not supported yet."); + rewriter.eraseOp(op); + return success(); +} + +// Get format code for tensor element type +// This maps MLIR types to Wafer format codes +Data_Format getFormatCode(MemRefType type) { + auto elemType = type.getElementType(); + if (elemType.isF32()) { + return Fmt_FP32; + } else if (elemType.isF16()) { + return Fmt_FP16; + } else if (elemType.isBF16()) { + return Fmt_BF16; + } else if (elemType.isInteger(1)) { + return Fmt_BOOL; + } else if (elemType.isInteger(8)) { + return Fmt_INT8; + } else if (elemType.isInteger(16)) { + return Fmt_INT16; + } else if (elemType.isInteger(32)) { + return Fmt_INT32; + } else if (elemType.isInteger(64)) { + return Fmt_INT64; + } else { + llvm_unreachable("Wafer does not support the element type\n"); + } + // Default to F32 format + return Fmt_FP32; +} + +bool isSupportedType(MemRefType type) { + auto elemType = type.getElementType(); + return elemType.isF32() || elemType.isF16() || elemType.isBF16() || + elemType.isInteger(8); +} + +static uint64_t getElemByte(Type type) { + static DataLayout dataLayout; + auto typeSize = dataLayout.getTypeSize(type); + if (!typeSize.isFixed()) { + llvm::llvm_unreachable_internal("All element type should have fixed size."); + } + return typeSize.getFixedValue(); +} + +static std::tuple, SmallVector> +createMetadata(ConversionPatternRewriter &rewriter, Location loc, + Value operand) { + auto stridedMetadata = + rewriter.create(loc, operand); + Value indexBasePtr = rewriter.create( + loc, rewriter.getIndexType(), stridedMetadata.getBaseBuffer()); + auto elemType = dyn_cast(operand.getType()).getElementType(); + Value elemByte = + rewriter.create(loc, getElemByte(elemType)); + Value offset = stridedMetadata.getOffset(); + Value byteOffset = + rewriter.create(loc, offset.getType(), offset, elemByte); + + if (elemType.isInteger(1)) { + auto [stride, offset] = + cast(operand.getType()).getStridesAndOffset(); + // Expected 8 bit alignment + assert(offset % 8 == 0); + + byteOffset = offset != 0 + ? rewriter.create( + loc, byteOffset.getType(), byteOffset, + rewriter.create(loc, 3)) + : byteOffset; + } + + Value offsetPtr = rewriter.create(loc, indexBasePtr.getType(), + indexBasePtr, byteOffset); + Value i64SPMPtr = rewriter.create( + loc, rewriter.getI64Type(), offsetPtr); + + // FIXME: For multi-dimensional(rank > 2), strides need to be multiplied. + return {i64SPMPtr, stridedMetadata.getSizes(), stridedMetadata.getStrides()}; +} + +static Value createAddressFromMemref(ConversionPatternRewriter &rewriter, + Location loc, Value operand) { + auto [i64SPMPtr, sizes, strides] = createMetadata(rewriter, loc, operand); + return i64SPMPtr; +} + +static SmallVector padSizesToNHWC(ConversionPatternRewriter &rewriter, + Location loc, ValueRange sizes) { + Value one = rewriter.create(loc, 1); + int numPad = 4 - sizes.size(); + SmallVector nhwcShape; + while (numPad--) { + nhwcShape.push_back(one); + } + for (auto dim : sizes) { + nhwcShape.push_back(dim); + } + return nhwcShape; +} + +// The last stride is always 1, skip it, nhwcStrides.size() will be 3. +static SmallVector +padStridesToNHWC(ConversionPatternRewriter &rewriter, Location loc, + ValueRange strides) { + Value one = rewriter.create(loc, 1); + int numPad = 4 - strides.size(); + SmallVector nhwcStrides; + while (numPad--) { + nhwcStrides.push_back(one); + } + for (auto dim : strides) { + nhwcStrides.push_back(dim); + } + return nhwcStrides; +} + +static Value calculateElemCount(ConversionPatternRewriter &rewriter, + Location loc, ValueRange sizes) { + // If we get scalar data, sizes is empty, return 1 + if (sizes.empty()) { + return rewriter.create(loc, 1); + } + + Value elemCount = sizes[0]; + for (int i = 1; i < sizes.size(); i++) { + elemCount = rewriter.create(loc, elemCount.getType(), + elemCount, sizes[i]); + } + return elemCount; +} + +// Extract the operations from a linalg op region +template llvm::SmallVector getRegionOps(T linalgOp) { + auto regionBlock = linalgOp.getBody(); + return llvm::map_to_vector(regionBlock->without_terminator(), + [](Operation &op) { return &op; }); +} + +static Data_Format getFormatFromElemType(mlir::Type elemType) { + // Convert the integer type to float type by just convert fmt. + // So here elemType can be integer type. + auto bitWidth = elemType.getIntOrFloatBitWidth(); + switch (bitWidth) { + case 8: + return Fmt_INT8; + case 16: + return elemType.isBF16() ? Fmt_BF16 : Fmt_FP16; + case 32: + return elemType.isTF32() ? Fmt_TF32 : Fmt_FP32; + default: + llvm_unreachable("Unsupported bit width\n"); + } + return Fmt_FP32; +} + +static Data_Format getFormatFromValueType(MemRefType valueType) { + // Convert the integer type to float type by just convert fmt. + // So here elemType can be integer type. + auto elemType = valueType.getElementType(); + return getFormatFromElemType(elemType); +} + +// Convert integer type to float type for CGRA instruction +// Return the convert float type format code +// TODO: Directly convert memref type? +Data_Format insertConvertTypeOp(Value valuePtr, MemRefType valueType, + Value elemCount, + ConversionPatternRewriter &rewriter, + Location loc) { + + // TODO: Other integer type. May need realloc the memory + auto elemType = valueType.getElementType(); + + if (!isa(elemType)) + return getFormatCode(valueType); + + Data_Format fmt = Fmt_FP32; + // Get the bit width from the element type + auto bitWidth = elemType.getIntOrFloatBitWidth(); + switch (bitWidth) { + case 16: { // 16 bit integer + rewriter.create(loc, rewriter.getI64Type(), valuePtr, + valuePtr, elemCount); + fmt = Fmt_FP16; + break; + } + case 32: { // 32 bit integer + rewriter.create(loc, rewriter.getI64Type(), valuePtr, + valuePtr, elemCount, + rewriter.getI16IntegerAttr(0)); + break; + } + default: { + llvm_unreachable("Unsupported integer type\n"); + } + } + return fmt; +} + +// Restore float type to integer type to for CGRA instruction +Value insertRestoreTypeOp(Value valuePtr, MemRefType valueType, Value elemCount, + ConversionPatternRewriter &rewriter, Location loc, + int16_t roundMode = RND_MODE::RND_NEAREST_EVEN) { + // TODO: Other integer type. May need realloc the memory + auto elemType = valueType.getElementType(); + auto newValue = valuePtr; + if (!isa(elemType)) + return newValue; + + // Get the bit width from the element type + auto bitWidth = elemType.getIntOrFloatBitWidth(); + switch (bitWidth) { + case 16: { // 16 bit integer + newValue = rewriter.create( + loc, rewriter.getI64Type(), valuePtr, valuePtr, elemCount, + rewriter.getI16IntegerAttr(roundMode)); + break; + } + case 32: { // 32 bit integer + newValue = rewriter.create( + loc, rewriter.getI64Type(), valuePtr, valuePtr, elemCount, + rewriter.getI16IntegerAttr(roundMode)); + break; + } + default: { + llvm_unreachable("Unsupported integer type\n"); + } + } + return newValue; +} + +SmallVector reshapeReduceShapeTo4d(ArrayRef inputShape, + int64_t dim) { + + auto rank = inputShape.size(); + SmallVector newShape; + int64_t leftDimsElement = 1; + int64_t rightDimsElement = 1; + + for (int i = 0; i < dim; i++) + leftDimsElement *= inputShape[i]; + + if (dim == inputShape.size() - 1) + return {1, 1, leftDimsElement, inputShape[dim]}; + + for (int i = dim + 1; i < rank; i++) + rightDimsElement *= inputShape[i]; + + newShape = {1, leftDimsElement, inputShape[dim], rightDimsElement}; // NHWC + return newShape; +} + +uint64_t next_power_of_two_64(uint64_t x) { + if (x == 0) { + return 1; + } + x--; + x |= x >> 1; + x |= x >> 2; + x |= x >> 4; + x |= x >> 8; + x |= x >> 16; + x |= x >> 32; + return x + 1; +} + +LLVM::AtomicOrdering getOrdering(MemSemantic sem) { + switch (sem) { + case MemSemantic::RELAXED: + return LLVM::AtomicOrdering::monotonic; + case MemSemantic::ACQUIRE: + return LLVM::AtomicOrdering::acquire; + case MemSemantic::RELEASE: + return LLVM::AtomicOrdering::release; + case MemSemantic::ACQUIRE_RELEASE: + return LLVM::AtomicOrdering::acq_rel; + default: + llvm_unreachable("Unexpected atomic mem semantic"); + } +} + +class MemoryCopyConvertPattern : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(memref::CopyOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + assert(op->hasAttr("srcSpm") && op->hasAttr("dstSpm") && + "Can't get memory space attribute\n"); + bool isSrcSPM = op->getAttrOfType("srcSpm").getValue(); + bool isDstSPM = op->getAttrOfType("dstSpm").getValue(); + // DDR to DDR + if (!isSrcSPM && !isDstSPM) + return rewriter.notifyMatchFailure( + op, "Can not copy memory from DDR to DDR.\n"); + + // SPM to SPM + if (isSrcSPM && isDstSPM) { + int64_t rank = cast(op.getSource().getType()).getRank(); + // A scalar copy has no axes to transpose. In particular, i1 falls + // back to linalg loops, whose permutation map cannot be empty. + // Keep the scalar memory access so the later pass applies SPM mapping. + if (rank == 0) { + Location loc = op.getLoc(); + Value val = rewriter.create(loc, op.getSource()); + rewriter.create(loc, val, op.getTarget()); + rewriter.eraseOp(op); + return success(); + } + SmallVector perm(rank); + std::iota(perm.begin(), perm.end(), 0); + rewriter.replaceOpWithNewOp(op, op.getSource(), + op.getTarget(), perm); + return success(); + } + + Location loc = op.getLoc(); + auto [srcPtr, srcSizes, srcStrides] = + createMetadata(rewriter, loc, adaptor.getSource()); + auto [dstPtr, dstSizes, dstStrides] = + createMetadata(rewriter, loc, adaptor.getTarget()); + auto srcMemrefType = cast(op.getSource().getType()); + auto dstMemrefType = cast(op.getTarget().getType()); + int64_t rank = srcMemrefType.getRank(); + auto elemType = srcMemrefType.getElementType(); + + if (srcMemrefType.areTrailingDimsContiguous(rank) && + dstMemrefType.areTrailingDimsContiguous(rank)) { + Value elemCount = calculateElemCount(rewriter, loc, dstSizes); + if (elemType.getIntOrFloatBitWidth() == 64) { + elemType = rewriter.getF32Type(); + elemCount = rewriter.create(loc, elemCount, elemCount); + } + elemCount = rewriter.create( + loc, rewriter.getI32Type(), elemCount); + + Data_Format fmt = getFormatFromElemType(elemType); + + if (isDstSPM) + rewriter.replaceOpWithNewOp(op, dstPtr, srcPtr, elemCount, + fmt); + else + rewriter.replaceOpWithNewOp(op, dstPtr, srcPtr, elemCount, + fmt); + return success(); + } + + auto srcFmt = getFormatCode(cast(srcMemrefType.clone( + rewriter.getIntegerType(srcMemrefType.getElementTypeBitWidth())))); + + // Update rank to 4 if rank less than 4. + if (rank < 4) { + srcSizes = padSizesToNHWC(rewriter, op->getLoc(), srcSizes); + srcStrides = padStridesToNHWC(rewriter, op->getLoc(), srcStrides); + dstSizes = padSizesToNHWC(rewriter, op->getLoc(), dstSizes); + dstStrides = padStridesToNHWC(rewriter, op->getLoc(), dstStrides); + rank = 4; + } + int elemBytes = srcMemrefType.getElementTypeBitWidth() >> 3; + if (isDstSPM) { + auto rdmaOp = rewriter.create( + op.getLoc(), rewriter.getI64Type(), srcPtr, dstPtr, + srcSizes, // src shape + srcStrides, // src stride + dstSizes, // dst shape + dstStrides, // dst stride + rewriter.getI32IntegerAttr(rank), // rank + rewriter.getI32IntegerAttr(elemBytes), // elem bytes + rewriter.getI32IntegerAttr(srcFmt) // Format + ); + } else { + auto wdmaOp = rewriter.create( + op.getLoc(), rewriter.getI64Type(), srcPtr, dstPtr, + srcSizes, // src shape + srcStrides, // src stride + dstSizes, // dst shape + dstStrides, // dst stride + rewriter.getI32IntegerAttr(rank), // rank + rewriter.getI32IntegerAttr(elemBytes), // elem bytes + rewriter.getI32IntegerAttr(srcFmt) // Format + ); + } + rewriter.eraseOp(op); + return success(); + } +}; + +// Convert linalg.fill to MemsetOp +class LinalgFillOpConversion : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + void preprocessI1Fill(ConversionPatternRewriter &rewriter, linalg::FillOp op, + SmallVector &srcSizes, + SmallVector &srcStrides, int64_t &rank, + Type &inputType) const { + + auto elementCount = calculateElemCount(rewriter, op->getLoc(), srcSizes); + SmallVector memoryLinearizedSizes{rewriter.create( + op.getLoc(), elementCount.getType(), elementCount, + rewriter.create(op.getLoc(), 4))}; + + SmallVector memoryLinearizedStrides{ + rewriter.create(op.getLoc(), 1)}; + + srcSizes = memoryLinearizedSizes; + srcStrides = memoryLinearizedStrides; + rank = 1; + inputType = rewriter.getF16Type(); + } + + bool isMemoryContiguousType(MemRefType type) const { + auto shape = type.getShape(); + auto rank = type.getRank(); + auto firstNonOne = std::find_if_not(shape.begin(), shape.end(), + [](int64_t dim) { return dim == 1; }); + int leadingOnes = std::distance(shape.begin(), firstNonOne); + return type.areTrailingDimsContiguous(rank - leadingOnes); + } + + bool isSupportedBitWidthAndType(int bitWidth, MemRefType type) const { + assert(isMemoryContiguousType(type) && "Type's memory must be contiguous"); + return bitWidth == 16 || bitWidth == 32 || + (bitWidth == 1 && type.getNumElements() >= 16); + } + + LogicalResult + matchAndRewrite(linalg::FillOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the value to fill with + Value fillValue = op.getInputs()[0]; // adaptor.getValue(); + + if (op.getOutputs().size() != 1) + return rewriter.notifyMatchFailure(op, "Only support single output\n"); + + auto rank = cast(op.getOutputs()[0].getType()).getRank(); + if (rank == 0) { + rewriter.create(op.getLoc(), adaptor.getInputs()[0], + adaptor.getOutputs()[0]); + rewriter.eraseOp(op); + return success(); + } + + auto [srcPtr, srcSizes, srcStrides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + auto inputType = op.getInputs()[0].getType(); + auto outputType = cast(op.getOutputs()[0].getType()); + auto bitWidth = inputType.getIntOrFloatBitWidth(); + + if (!isSupportedBitWidthAndType(bitWidth, outputType)) { + return convertLinalgOpToLoops(op, rewriter); + } + + fillValue = rewriter.create( + op.getLoc(), rewriter.getIntegerType(bitWidth), fillValue); + fillValue = bitWidth != 32 + ? rewriter.create( + op.getLoc(), rewriter.getI32Type(), fillValue) + : fillValue; + + if (bitWidth == 1) { + preprocessI1Fill(rewriter, op, srcSizes, srcStrides, rank, inputType); + } + Data_Format fmt = getFormatFromElemType(inputType); + + // NOTE: When encounter NaN, use xor + addvs to simulate memset operation + // will get wrong result. + auto resultOp = rewriter.create( + op.getLoc(), rewriter.getI64Type(), srcPtr, fillValue, srcSizes, + srcStrides, rewriter.getI32IntegerAttr(rank), + rewriter.getI16IntegerAttr(fmt)); + + rewriter.eraseOp(op); + + return success(); + } +}; + +class TransposeOpConversion : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + bool convertToGatherScatter(linalg::TransposeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto perms = op.getPermutation(); + auto src = op.getInput(); + auto dst = op.getInit(); + auto srcType = cast(src.getType()); + auto dstType = cast(dst.getType()); + auto srcShape = srcType.getShape(); + auto srcStrides = srcType.getStridesAndOffset().first; + auto dstShape = dstType.getShape(); + auto dstStrides = dstType.getStridesAndOffset().first; + unsigned bitWidth = srcType.getElementTypeBitWidth(); + + if (!srcType.hasStaticShape() || !dstType.hasStaticShape() || + llvm::any_of(srcStrides, ShapedType::isDynamic) || + llvm::any_of(dstStrides, ShapedType::isDynamic)) { + LDBG("TransposeOpConversion: dynamic shape/strides not supported\n"); + return false; + } + + // Get inner bits + uint64_t bits = bitWidth; + int64_t contiguousStride = 1; + size_t rank = perms.size(); + for (; rank > 0 && srcStrides[perms[rank - 1]] == contiguousStride && + dstStrides[rank - 1] == contiguousStride; + --rank) { + bits *= dstShape[rank - 1]; + contiguousStride *= dstShape[rank - 1]; + } + + if (bits & 0x7 || rank > 3) { + LDBG("TransposeOpConversion: bits not byte aligned or rank not " + "supported\n"); + return false; + } + + unsigned bytes = bits >> 3; + SmallVector srcStrideArgs(3); + SmallVector srcIterArgs(3, 1); + SmallVector dstStrideArgs(3); + SmallVector dstIterArgs(3, 1); + + for (size_t i = 0; i < rank; ++i) { + uint64_t srcStrideBits = srcStrides[perms[i]] * bitWidth; + uint64_t dstStrideBits = dstStrides[i] * bitWidth; + if (srcStrideBits & 0x7 || dstStrideBits & 0x7) { + LDBG("TransposeOpConversion: stride not byte aligned\n"); + return false; + } + srcStrideArgs[i + 3 - rank] = srcStrideBits >> 3; + dstStrideArgs[i + 3 - rank] = dstStrideBits >> 3; + srcIterArgs[i + 3 - rank] = srcShape[perms[i]]; + dstIterArgs[i + 3 - rank] = dstShape[i]; + } + + auto srcPtr = createAddressFromMemref(rewriter, op->getLoc(), src); + auto dstPtr = createAddressFromMemref(rewriter, op->getLoc(), dst); + + rewriter.create( + op.getLoc(), rewriter.getI64Type(), srcPtr, dstPtr, bytes, + srcStrideArgs[0], srcStrideArgs[1], srcStrideArgs[2], srcIterArgs[0], + srcIterArgs[1], srcIterArgs[2], dstStrideArgs[0], dstStrideArgs[1], + dstStrideArgs[2], dstIterArgs[0], dstIterArgs[1], dstIterArgs[2]); + + rewriter.eraseOp(op); + return true; + } + + template + LogicalResult transposeChannel(linalg::TransposeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto src = op.getInput(); + auto dst = op.getInit(); + auto srcType = cast(src.getType()); + auto dstType = cast(dst.getType()); + SmallVector srcShape(srcType.getShape().begin(), + srcType.getShape().end()); + SmallVector dstShape(dstType.getShape().begin(), + dstType.getShape().end()); + + auto srcPtr = createAddressFromMemref(rewriter, op->getLoc(), src); + auto dstPtr = createAddressFromMemref(rewriter, op->getLoc(), dst); + + // TODO: Through fmt conversion to support more element types. + if (!isSupportedType(srcType)) { + return rewriter.notifyMatchFailure(op, "Unsupported element type\n"); + } + Data_Format fmt = getFormatCode(srcType); + + auto newOp = + rewriter.create(op->getLoc(), rewriter.getI64Type(), srcPtr, + dstPtr, srcShape, dstShape, fmt); + + rewriter.eraseOp(op); + return success(); + } + + LogicalResult + matchAndRewrite(linalg::TransposeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto perm = op.getPermutation(); + + if (perm == ArrayRef({0, 2, 3, 1})) { + return transposeChannel(op, adaptor, rewriter); + } + + if (perm == ArrayRef({0, 3, 1, 2})) { + return transposeChannel(op, adaptor, rewriter); + } + + if (convertToGatherScatter(op, adaptor, rewriter)) + return success(); + + // Default handling of remaining cases. + // TODO: Convert higher rank to wafer. + return convertLinalgOpToLoops(op, rewriter); + } +}; + +class ReciprocalOpConversionPattern + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(linalg::ReciprocalOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + + auto [inputPtr, sizes, strides] = + createMetadata(rewriter, loc, adaptor.getInputs()[0]); + auto outputPtr = + createAddressFromMemref(rewriter, loc, adaptor.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + // Tx neural engine not support fp32 for input + auto inputType = dyn_cast(op.getInputs()[0].getType()); + Data_Format srcFmt = getFormatCode(inputType); + + rewriter.create(loc, rewriter.getI64Type(), inputPtr, + outputPtr, elemCount, + rewriter.getI16IntegerAttr(srcFmt)); + rewriter.eraseOp(op); + + return success(); + } +}; + +//===----------------------------------------------------------------------===// +// mk.dot to wafer.gemm Conversion Pattern +//===----------------------------------------------------------------------===// + +class MKDotToWaferGemmOpConversion : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mk::DotOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op->getLoc(); + auto a = adaptor.getA(); + auto b = adaptor.getB(); + auto dst = adaptor.getInits(); + auto aType = cast(a.getType()); // (K/64)xMx64 + auto bType = cast(b.getType()); // (N/64)xKx64 + auto dstType = cast(dst.getType()); // (N/64)xMx64 + int32_t M = aType.getShape()[1]; + int32_t K = bType.getShape()[1]; + int32_t N = bType.getShape()[0] * bType.getShape()[2]; + auto dims = rewriter.getI32ArrayAttr({M, K, N}); + Data_Format srcFmt = getFormatCode(aType); + Data_Format dstFmt = getFormatCode(dstType); + + auto aPtr = createAddressFromMemref(rewriter, loc, a); + auto bPtr = createAddressFromMemref(rewriter, loc, b); + auto dstPtr = createAddressFromMemref(rewriter, loc, dst); + + // Assume input type is same. Tx neural engine not support fp32 for input + // FIXME: There are encoding differences between f32 and tf32 for certain + // special values (e.g., Inf and NaN). We assume that such special values do + // not occur. + if (aType.getElementType().isF32()) { + // Warning for neural engine that fp32 is not supported + LLVM_DEBUG(llvm::dbgs() << "Neural engine not support FP32. Convert FP32 " + "to TF32 for wafer.Gemm Op\n"); + srcFmt = Data_Format::Fmt_TF32; + } + + auto zero = + rewriter.create(loc, rewriter.getI64Type(), 0); + + // Create GemmOp + rewriter.create( + loc, rewriter.getI64Type(), + aPtr, // src_a (Matrix A in SPM) + bPtr, // src_b (Matrix B in SPM) + dstPtr, // src_bias. Unused for now. + dstPtr, // dst, + dims, // dimensions [M,K,N] + op.getEnPsumAttr(), // en_psum. Used as accumulate buffer + dstPtr, // The address of psum in SPM, Always same to output + rewriter.getBoolAttr(false), // trans_src_a + rewriter.getBoolAttr(true), // trans_src_b. + rewriter.getI32IntegerAttr(1), // batch_src_a + rewriter.getI32IntegerAttr(1), // batch_src_b + rewriter.getI32IntegerAttr(0), // relu_mode: no activation. + rewriter.getBoolAttr(false), // en_bias + rewriter.getBoolAttr(false), // en_neg_scale + zero, // src_neg_scale + rewriter.getBoolAttr(false), // en_pos_scale + zero, // src_pos_scale + rewriter.getI32IntegerAttr(srcFmt), // src_fmt + rewriter.getI32IntegerAttr(dstFmt) // dst_fmt + ); + + rewriter.eraseOp(op); + return success(); + } +}; + +class MKDequantOpConversionPattern : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mk::DequantOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + auto loc = op.getLoc(); + auto inputPtr = createAddressFromMemref(rewriter, loc, adaptor.getSrc()); + auto scalePtr = createAddressFromMemref(rewriter, loc, adaptor.getScale()); + auto outputPtr = createAddressFromMemref(rewriter, loc, adaptor.getInit()); + + auto outputType = cast(adaptor.getInit().getType()); + auto elemCount = outputType.getNumElements(); + auto elemType = outputType.getElementType(); + assert(elemType.isF16() || elemType.isBF16()); + if (elemType.isBF16()) + rewriter.create(loc, TypeRange{}, inputPtr, scalePtr, + outputPtr, elemCount); + else + rewriter.create(loc, TypeRange{}, inputPtr, scalePtr, + outputPtr, elemCount); + rewriter.eraseOp(op); + return success(); + } +}; + +class GatherConvertPattern : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mlir::mk::GatherOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + + auto indices = adaptor.getIndices(); + auto indicesType = cast(indices.getType()); + auto shape = indicesType.getShape(); + + auto axis = op.getAxis(); + + int64_t numElems = indicesType.getNumElements(); + auto strides = computeStrides(shape); + for (int64_t idx = 0; idx < numElems; idx += 1) { + auto tensorIdx = delinearize(idx, strides); + + SmallVector idxIndex(tensorIdx.size()); + std::transform(tensorIdx.begin(), tensorIdx.end(), idxIndex.begin(), + [&](auto val) { + return rewriter.create(loc, val); + }); + // Read the index value from indices tensor + Value indexValue = + rewriter.create(loc, indices, idxIndex); + + // Read value from source using computed indices + SmallVector inputIndex = idxIndex; + assert(axis < inputIndex.size() && axis >= 0 && + "Axis index out of bounds"); + inputIndex[axis] = rewriter.create( + loc, rewriter.getIndexType(), indexValue); + + Value gatheredValue = + rewriter.create(loc, adaptor.getSrc(), inputIndex); + + // Write value to destination + rewriter.create(loc, gatheredValue, adaptor.getDst(), + idxIndex); + } + + rewriter.eraseOp(op); + + return success(); + } +}; + +class MKSigmoidToWaferSigmoidOpConversion + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mlir::mk::SigmoidOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + auto [input, sizes, strides] = + createMetadata(rewriter, loc, adaptor.getSrc()); + auto [dst, dstSizes, dstStrides] = + createMetadata(rewriter, loc, adaptor.getZeroes()); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + // Tx neural engine not support fp32 for input + auto inputType = dyn_cast(op.getSrc().getType()); + assert(isSupportedType(inputType) && + "Unsupported element type for Sigmoid operation\n"); + Data_Format srcFmt = getFormatCode(inputType); + + rewriter.create(loc, rewriter.getI64Type(), input, dst, + elemCount, rewriter.getI16IntegerAttr(srcFmt)); + rewriter.eraseOp(op); + + return success(); + } +}; + +struct MKGeluToWaferGeluOpConversion + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(mlir::mk::GeluOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + auto [input, sizes, strides] = + createMetadata(rewriter, loc, adaptor.getSrc()); + auto [dst, dstSizes, dstStrides] = + createMetadata(rewriter, loc, adaptor.getZeroes()); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + // Check input type: only support f16, bf16, f32. + auto inputType = dyn_cast(op.getSrc().getType()); + assert(isSupportedType(inputType) && + "Unsupported element type for Gelu operation\n"); + Data_Format srcFmt = getFormatCode(inputType); + + auto geluMode = static_cast(op.getGeluMode()); + switch (geluMode) { + case GeluMode::None: { + rewriter.create(loc, rewriter.getI64Type(), input, dst, + elemCount, + rewriter.getI16IntegerAttr(srcFmt)); + break; + } + case GeluMode::Tanh: { + auto immAddr = createAddressFromMemref(rewriter, loc, adaptor.getImm()); + rewriter.create(loc, rewriter.getI64Type(), input, immAddr, + dst, elemCount, + rewriter.getI16IntegerAttr(srcFmt)); + break; + } + default: { + llvm::report_fatal_error("Unsupported gelu mode!"); + } + } + rewriter.eraseOp(op); + return success(); + } +}; + +class MKBit2FPOpConversionPattern : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mlir::mk::Bit2FpOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + auto [inputPtr, sizes, strides] = + createMetadata(rewriter, loc, adaptor.getSrc()); + auto outputPtr = createAddressFromMemref(rewriter, loc, adaptor.getInit()); + + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto outputType = dyn_cast(op.getInit().getType()); + Data_Format srcFmt = getFormatCode(outputType); + + rewriter.create(loc, rewriter.getI64Type(), inputPtr, + outputPtr, elemCount, + rewriter.getI16IntegerAttr(srcFmt)); + rewriter.eraseOp(op); + + return success(); + } +}; + +class MKMaskMoveOpConversionPattern + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mlir::mk::MaskMoveOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + auto [inputPtr, sizes, strides] = + createMetadata(rewriter, loc, adaptor.getSource()); + auto outputPtr = createAddressFromMemref(rewriter, loc, adaptor.getInit()); + auto maskPtr = createAddressFromMemref(rewriter, loc, adaptor.getMask()); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto inputType = dyn_cast(op.getSource().getType()); + Data_Format srcFmt = getFormatCode(inputType); + + rewriter.create(loc, rewriter.getI64Type(), inputPtr, + outputPtr, elemCount, maskPtr, + rewriter.getI32IntegerAttr(srcFmt)); + rewriter.eraseOp(op); + + return success(); + } +}; + +template +struct MKRelationVVOpConversionPattern : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename MKOpT::Adaptor; + + LogicalResult + matchAndRewrite(MKOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto input0 = adaptor.getInput0(); + auto input1 = adaptor.getInput1(); + auto output = adaptor.getInit(); + auto inputType = cast(input0.getType()); + + auto loc = op.getLoc(); + + auto [input0Ptr, sizes, strides] = createMetadata(rewriter, loc, input0); + auto input1Ptr = createAddressFromMemref(rewriter, loc, input1); + auto outputPtr = createAddressFromMemref(rewriter, loc, output); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto waferOp = rewriter.create( + loc, rewriter.getI64Type(), input0Ptr, input1Ptr, outputPtr, elemCount, + rewriter.getI16IntegerAttr(getFormatCode(inputType))); + + rewriter.eraseOp(op); + return success(); + } +}; + +template +struct MKArithVSOpConversionPattern : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename MKOpT::Adaptor; + + LogicalResult + matchAndRewrite(MKOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto input = adaptor.getInput(); + auto val = adaptor.getValue(); + auto output = adaptor.getInit(); + auto inputType = cast(input.getType()); + + auto loc = op.getLoc(); + + // Store float value on a integer type 32bit memory + auto bitWidth = inputType.getElementTypeBitWidth(); + assert(bitWidth == 16 || bitWidth == 32); + auto bitcastType = + bitWidth == 16 ? rewriter.getI16Type() : rewriter.getI32Type(); + Value i32Value = + rewriter.create(op.getLoc(), bitcastType, val); + // Extend bitcode to 32 bit + if (bitWidth == 16) { + i32Value = rewriter.create( + op.getLoc(), rewriter.getI32Type(), i32Value); + } + + auto elemCount = inputType.getNumElements(); + auto inputPtr = createAddressFromMemref(rewriter, loc, input); + auto outputPtr = createAddressFromMemref(rewriter, loc, output); + auto waferOp = rewriter.create( + loc, rewriter.getI64Type(), inputPtr, i32Value, outputPtr, + rewriter.create(loc, elemCount), + rewriter.getI16IntegerAttr(RND_MODE::RND_NEAREST_EVEN), // Round mode + rewriter.getI16IntegerAttr(getFormatCode(inputType))); + + rewriter.eraseOp(op); + return success(); + } +}; + +template +struct MKRelationVSOpConversionPattern : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename MKOpT::Adaptor; + + LogicalResult + matchAndRewrite(MKOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto input = adaptor.getInput(); + auto val = adaptor.getValue(); + auto output = adaptor.getInit(); + auto inputType = cast(input.getType()); + + auto loc = op.getLoc(); + + // Store float value on a integer type 32bit memory + auto bitWidth = inputType.getElementTypeBitWidth(); + assert(bitWidth == 16 || bitWidth == 32); + auto bitcastType = + bitWidth == 16 ? rewriter.getI16Type() : rewriter.getI32Type(); + Value i32Value = + rewriter.create(op.getLoc(), bitcastType, val); + // Extend bitcode to 32 bit + if (bitWidth == 16) { + i32Value = rewriter.create( + op.getLoc(), rewriter.getI32Type(), i32Value); + } + + // BoolRelationVSOp need 8 bit alignment + auto elemCount = inputType.getNumElements(); + auto outputType = cast(output.getType()); + if (outputType.getElementType().isInteger(1)) { + elemCount = ((elemCount + 7) / 8) * 8; + op->emitRemark() << "element count was expanded to a multiple of 8, may " + "access memory out of bounds!"; + } + + auto inputPtr = createAddressFromMemref(rewriter, loc, input); + auto outputPtr = createAddressFromMemref(rewriter, loc, output); + auto waferOp = rewriter.create( + loc, rewriter.getI64Type(), inputPtr, i32Value, outputPtr, + rewriter.create(loc, elemCount), + rewriter.getI16IntegerAttr(getFormatCode(inputType))); + + rewriter.eraseOp(op); + return success(); + } +}; + +template +struct MKArgMinMaxConversionPattern : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename MKOpT::Adaptor; + + LogicalResult + matchAndRewrite(MKOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto input = adaptor.getSrc(); + auto outVal = adaptor.getValue(); + auto outIdx = adaptor.getIndex(); + auto inputType = cast(input.getType()); + auto valueType = cast(outVal.getType()); + auto indexType = cast(outIdx.getType()); + auto inputShape = inputType.getShape(); + + assert(!inputType.getElementType().isInteger() && + "mk.argmax/argmin op's input type should not be integer."); + + auto loc = op.getLoc(); + auto reduceDim = op.getAxis(); + int64_t innerSize = inputShape.empty() ? 1 : inputShape.back(); + // NOTE: LinalgToMK Pass already sliced input to a vector + assert(inputType.getRank() == 1); + // TODO: support input's stride != 1 + auto [strides, offset] = inputType.getStridesAndOffset(); + assert(strides[0] == 1); + + auto inputPtr = createAddressFromMemref(rewriter, loc, input); + auto outValPtr = createAddressFromMemref(rewriter, loc, outVal); + auto outIdxPtr = createAddressFromMemref(rewriter, loc, outIdx); + + auto waferOp = rewriter.create( + loc, TypeRange{}, inputPtr, outValPtr, outIdxPtr, + rewriter.getI32IntegerAttr(innerSize), + rewriter.getI16IntegerAttr(getFormatCode(valueType))); + rewriter.eraseOp(op); + return success(); + } +}; + +struct ElementwiseConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult convertIsNaNOp(linalg::GenericOp op, OpAdaptor adapter, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input = createAddressFromMemref(rewriter, loc, adapter.getInputs()[0]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adapter.getOutputs()[0]); + auto inputType = dyn_cast(op.getInputs()[0].getType()); + auto elemCount = inputType.getNumElements(); + assert((elemCount % 8) == 0 && + "ElemCount must be a multiple of 8 due to ElementwiseRewrite pass!"); + + auto elemCountValue = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto fmt = getFormatCode(inputType); + rewriter.create(loc, // loc + rewriter.getI64Type(), // result type + input, // input0 + input, // input1 + output, // out + elemCountValue, // elem_count + rewriter.getI16IntegerAttr(fmt) // fmt + ); + rewriter.eraseOp(op); + return success(); + } + + template + LogicalResult convertUnaryOp(linalg::GenericOp op, OpAdaptor adapter, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input = createAddressFromMemref(rewriter, loc, adapter.getInputs()[0]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adapter.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto inputType = dyn_cast(op.getInputs()[0].getType()); + // Data format after conversion + Data_Format srcFmt = getFormatCode(inputType); + + // Create the unary operation + rewriter.create(loc, rewriter.getI64Type(), input, output, elemCount, + rewriter.getI16IntegerAttr(srcFmt)); + + rewriter.eraseOp(op); + return success(); + } + + template + LogicalResult convertBinaryOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input0 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto input1 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[1]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto inputType = dyn_cast(op.getInputs()[0].getType()); + // Data format after conversion + Data_Format srcFmt = getFormatCode(inputType); + + // Create the elementwise operation + // TODO: Fix attribute + rewriter.create(loc, rewriter.getI64Type(), input0, input1, output, + elemCount, + rewriter.getI16IntegerAttr(0), // Round mode + rewriter.getI16IntegerAttr(srcFmt)); + + rewriter.eraseOp(op); + return success(); + } + + template + LogicalResult + convertBoolBinaryLogicOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input0 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto input1 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[1]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + + auto inputType = dyn_cast(op.getInputs()[0].getType()); + auto bitWidth = inputType.getElementType().getIntOrFloatBitWidth(); + auto elemCount = inputType.getNumElements(); + + // If bit width is 1 and element count is not divisible by 8, expand + // the number of elements to a multiple of 8. + if (bitWidth == 1 && elemCount % 8) { + elemCount = ((elemCount + 7) / 8) * 8; + } + + elemCount *= bitWidth; + + // Creat new element count value. + Value elemCountValue = + rewriter.create(loc, elemCount); + rewriter.create(loc, rewriter.getI64Type(), input0, input1, output, + elemCountValue); + + rewriter.eraseOp(op); + return success(); + } + + template + LogicalResult ZeroPointConvertOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input = createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto output = createAddressFromMemref(rewriter, op->getLoc(), + adaptor.getOutputs()[0]); + auto elemCount = + cast(op->getOperandTypes()[0]).getNumElements(); + + rewriter.create(loc, input, output, 0, (uint32_t)elemCount); + rewriter.eraseOp(op); + return success(); + } + + template + LogicalResult NormalConvertOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input = createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + rewriter.create(loc, rewriter.getI64Type(), input, output, + elemCount); + rewriter.eraseOp(op); + return success(); + } + + template + LogicalResult + RoundConvertOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter, + RND_MODE roundMode = RND_MODE::RND_NEAREST_EVEN) const { + Location loc = op->getLoc(); + auto input = createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + // TODO: Fix attribute + auto result = rewriter.create( + loc, + rewriter.getI64Type(), // Result type + input, // Input + output, // Output + elemCount, // Element count + rewriter.getI16IntegerAttr(roundMode) // Round mode + ); + rewriter.eraseOp(op); + return success(); + } + + template + LogicalResult BoolRelationVVOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input0 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto input1 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[1]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto inputType = dyn_cast(op.getInputs()[0].getType()); + + assert(inputType.getNumElements() % 8 == 0); + Data_Format srcFmt = getFormatCode(inputType); + + // Create the elementwise operation + // TODO: Fix attribute + rewriter.create(loc, rewriter.getI64Type(), input0, input1, output, + elemCount, + rewriter.getI16IntegerAttr(srcFmt) // Format + ); + + rewriter.eraseOp(op); + return success(); + } + + LogicalResult FmaConvertOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input0 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto input1 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[1]); + auto input2 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[2]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto inputType = dyn_cast(op.getInputs()[0].getType()); + + auto mulResult = rewriter.create( + loc, rewriter.getI64Type(), input0, input1, output, elemCount, + rewriter.getI16IntegerAttr(0), // Round mode + rewriter.getI16IntegerAttr(getFormatCode(inputType))); + auto addResult = rewriter.create( + loc, rewriter.getI64Type(), output, input2, output, elemCount, + rewriter.getI16IntegerAttr(0), // Round mode + rewriter.getI16IntegerAttr(getFormatCode(inputType))); + rewriter.eraseOp(op); + return success(); + } + + LogicalResult convertDivIntOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Location loc = op->getLoc(); + auto input0 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto input1 = + createAddressFromMemref(rewriter, loc, adaptor.getInputs()[1]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto inputType = dyn_cast(op.getInputs()[0].getType()); + // Data format after conversion + Data_Format srcFmt = + insertConvertTypeOp(input0, inputType, elemCount, rewriter, loc); + if (adaptor.getInputs()[0] != adaptor.getInputs()[1]) { + // If input0 and input1 are not the same, we need to convert input1 type + insertConvertTypeOp(input1, inputType, elemCount, rewriter, loc); + } + + if (adaptor.getInputs()[0] != adaptor.getOutputs()[0] && + adaptor.getInputs()[1] != adaptor.getOutputs()[0]) + // If input and output are not the same, we need to convert output type + insertConvertTypeOp(output, inputType, elemCount, rewriter, loc); + + rewriter.create(loc, rewriter.getI64Type(), input0, input1, + output, elemCount, + rewriter.getI16IntegerAttr(0), // Round mode + rewriter.getI16IntegerAttr(srcFmt)); + + insertRestoreTypeOp(output, inputType, elemCount, rewriter, loc, + RND_MODE::RND_ZERO); + + if (adaptor.getInputs()[0] != adaptor.getOutputs()[0]) { + insertRestoreTypeOp(input0, inputType, elemCount, rewriter, loc); + } + if (adaptor.getInputs()[1] != adaptor.getOutputs()[0] && + adaptor.getInputs()[1] != adaptor.getInputs()[0]) { + insertRestoreTypeOp(input1, inputType, elemCount, rewriter, loc); + } + + rewriter.eraseOp(op); + return success(); + } + + LogicalResult + convertRoundOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter, + RND_MODE roundMode = RND_MODE::RND_NEAREST_EVEN) const { + Location loc = op->getLoc(); + auto input = createAddressFromMemref(rewriter, loc, adaptor.getInputs()[0]); + auto [output, sizes, strides] = + createMetadata(rewriter, op->getLoc(), adaptor.getOutputs()[0]); + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + // Use IEEE round to nearest mode + auto fpToInt = rewriter.create( + loc, rewriter.getI64Type(), input, output, elemCount, + rewriter.getI16IntegerAttr(roundMode)); // Round mode + auto intToFp = rewriter.create( + loc, rewriter.getI64Type(), output, output, elemCount, + rewriter.getI16IntegerAttr(0)); // Round mode + + rewriter.eraseOp(op); + return success(); + } + + LogicalResult convertUIToFPOp(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto inputType = + dyn_cast(op.getInputs()[0].getType()).getElementType(); + auto outputType = + dyn_cast(op.getOutputs()[0].getType()).getElementType(); + if (inputType.isInteger(1) && (outputType.isF32() || outputType.isF16())) { + Location loc = op.getLoc(); + auto [inputPtr, sizes, strides] = + createMetadata(rewriter, loc, adaptor.getInputs()[0]); + auto outputPtr = + createAddressFromMemref(rewriter, loc, adaptor.getOutputs()[0]); + + auto elemCount = calculateElemCount(rewriter, op->getLoc(), sizes); + + auto outputType = dyn_cast(op.getOutputs()[0].getType()); + Data_Format srcFmt = getFormatCode(outputType); + + rewriter.create(loc, rewriter.getI64Type(), inputPtr, + outputPtr, elemCount, + rewriter.getI16IntegerAttr(srcFmt)); + rewriter.eraseOp(op); + + return success(); + } else { + return rewriter.notifyMatchFailure( + op, "Unsupported input/output type combination for integer to " + "FP conversion"); + } + } + + LogicalResult + matchAndRewrite(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + auto regionOps = getRegionOps(op); + + if (!op.getOutputs().empty() && + cast(op.getOutputs()[0].getType()).getRank() == 0) + return convertLinalgOpToLoops(op, rewriter); + + // Check if the operation is elementwise + if (op.getIteratorTypesArray().front() != + mlir::utils::IteratorType::parallel) + return rewriter.notifyMatchFailure(op, "Only support elementwise op."); + + // WORKAROUND: Select op input0 is bool(i1), cmp op result is bool(i1) + // I64/F64 lowering to llvm + // NOTE: May exist scf.if which has not output + if (regionOps.size() != 1 || + (!op.getOutputs().empty() && + dyn_cast(op.getOutputs()[0].getType()) + .getElementType() + .getIntOrFloatBitWidth() == 64) || + (!op.getInputs().empty() && + dyn_cast(op.getInputs()[0].getType()) + .getElementType() + .getIntOrFloatBitWidth() == 64)) { + return convertLinalgOpToLoops(op, rewriter); + } + + auto elemWiseOp = regionOps[0]; + return llvm::TypeSwitch(elemWiseOp) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertIsNaNOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertBinaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertBinaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertBinaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertDivIntOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertBinaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertBinaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertBoolBinaryLogicOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertBoolBinaryLogicOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertBoolBinaryLogicOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertRoundOp(op, adaptor, rewriter, RND_MODE::RND_POS_INF); + }) + .Case([&](auto elemWiseOp) { + return convertRoundOp(op, adaptor, rewriter, RND_MODE::RND_NEG_INF); + }) + .Case([&](auto elemWiseOp) { + return convertRoundOp(op, adaptor, rewriter, RND_MODE::RND_ZERO); + }) + .Case([&](auto elemWiseOp) { + return convertRoundOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + return convertUnaryOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + auto inputType = elemWiseOp.getIn().getType(); + auto targetType = elemWiseOp.getOut().getType(); + if (inputType.isF16() && targetType.isF32()) + return NormalConvertOp(op, adaptor, rewriter); + else if (inputType.isBF16() && targetType.isF32()) + return NormalConvertOp(op, adaptor, rewriter); + + else if (isa(inputType) && targetType.isBF16()) + return NormalConvertOp(op, adaptor, rewriter); + else if (isa(inputType) && targetType.isBF16()) + return NormalConvertOp(op, adaptor, rewriter); + else if (isa(inputType) && targetType.isBF16()) + return NormalConvertOp(op, adaptor, + rewriter); + else if (isa(inputType) && targetType.isBF16()) + return NormalConvertOp(op, adaptor, rewriter); + + else if (isa(inputType) && targetType.isF16()) + return NormalConvertOp(op, adaptor, rewriter); + else if (isa(inputType) && targetType.isF16()) + return NormalConvertOp(op, adaptor, rewriter); + else if (isa(inputType) && targetType.isF16()) + return NormalConvertOp(op, adaptor, + rewriter); + else if (isa(inputType) && targetType.isF16()) + return NormalConvertOp(op, adaptor, rewriter); + else + return rewriter.notifyMatchFailure( + op, "Unsupported input/output type combination for ExtFOp " + "conversion"); + }) + .Case([&](auto elemWiseOp) { + return FmaConvertOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + // TODO: Need add more int to fp convert. + auto inputType = dyn_cast(op.getInputs()[0].getType()) + .getElementType(); + auto outputType = dyn_cast(op.getOutputs()[0].getType()) + .getElementType(); + + if (inputType.isInteger(8) && outputType.isF32()) { + return ZeroPointConvertOp(op, adaptor, rewriter); + } else if (inputType.isInteger(8) && outputType.isF16()) { + return ZeroPointConvertOp(op, adaptor, rewriter); + } else if (inputType.isInteger(16) && outputType.isF32()) { + return RoundConvertOp(op, adaptor, rewriter); + } else if (inputType.isInteger(16) && outputType.isF16()) { + return NormalConvertOp(op, adaptor, rewriter); + } else if (inputType.isInteger(32) && outputType.isF16()) { + return RoundConvertOp(op, adaptor, rewriter); + } else if (inputType.isInteger(32) && outputType.isF32()) { + return RoundConvertOp(op, adaptor, rewriter); + } else { + return rewriter.notifyMatchFailure( + op, "Unsupported input/output type combination for integer to " + "FP conversion"); + } + }) + .Case([&](auto elemWiseOp) { + return convertUIToFPOp(op, adaptor, rewriter); + }) + .Case([&](auto elemWiseOp) { + // TODO: Need add more int to fp convert. + auto inputType = dyn_cast(op.getInputs()[0].getType()) + .getElementType(); + auto outputType = dyn_cast(op.getOutputs()[0].getType()) + .getElementType(); + // arith.fptosi: Cast from a value interpreted as floating-point to + // the nearest (rounding towards zero) signed integer value. + if (inputType.isF16() && outputType.isInteger(8)) { + return RoundConvertOp(op, adaptor, rewriter, + RND_MODE::RND_ZERO); + } else if (inputType.isF16() && outputType.isInteger(16)) { + return RoundConvertOp(op, adaptor, rewriter, + RND_MODE::RND_ZERO); + } else if (inputType.isF16() && outputType.isInteger(32)) { + return RoundConvertOp(op, adaptor, rewriter, + RND_MODE::RND_ZERO); + } else if (inputType.isF32() && outputType.isInteger(8)) { + return RoundConvertOp(op, adaptor, rewriter, + RND_MODE::RND_ZERO); + } else if (inputType.isF32() && outputType.isInteger(16)) { + return RoundConvertOp(op, adaptor, rewriter, + RND_MODE::RND_ZERO); + } else if (inputType.isF32() && outputType.isInteger(32)) { + return RoundConvertOp(op, adaptor, rewriter, + RND_MODE::RND_ZERO); + } else { + return rewriter.notifyMatchFailure( + op, "Unsupported input/output type combination for fp to " + "integer conversion"); + } + }) + .Case([&](auto elemWiseOp) { + arith::CmpFPredicate predicate = elemWiseOp.getPredicate(); + switch (predicate) { + case arith::CmpFPredicate::OEQ: + case arith::CmpFPredicate::UEQ: + return BoolRelationVVOp(op, adaptor, rewriter); + case arith::CmpFPredicate::ONE: + case arith::CmpFPredicate::UNE: + return BoolRelationVVOp(op, adaptor, rewriter); + case arith::CmpFPredicate::OGE: + case arith::CmpFPredicate::UGE: + return BoolRelationVVOp(op, adaptor, + rewriter); + case arith::CmpFPredicate::OGT: + case arith::CmpFPredicate::UGT: + return BoolRelationVVOp(op, adaptor, rewriter); + case arith::CmpFPredicate::OLE: + case arith::CmpFPredicate::ULE: + return BoolRelationVVOp(op, adaptor, rewriter); + case arith::CmpFPredicate::OLT: + case arith::CmpFPredicate::ULT: + return BoolRelationVVOp(op, adaptor, rewriter); + default: + llvm_unreachable("Not yet supported"); + break; + } + }) + .Case([&](auto elemWiseOp) { + // May exist elemWiseOp has no result + auto resultType = elemWiseOp->getResult(0).getType(); + if (resultType.isF16()) + return RoundConvertOp(op, adaptor, rewriter); + else if (resultType.isBF16()) + return RoundConvertOp(op, adaptor, rewriter); + else + return rewriter.notifyMatchFailure( + op, "Unsupported input/output type combination for trunc " + "conversion"); + }) + .Default([&](auto elemWiseOp) { + // Affine dialect should handled before this pass. So here lower it + // to scf.for + return convertLinalgOpToLoops(op, rewriter); + }); + } +}; + +struct LinalgReduceConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(linalg::ReduceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto reductionOps = getRegionOps(op); + // If there is only one reduction operation, try to convert it to Wafer + // reduce op. + // TODO: Delete, here use to check linalg-to-mk pass has finished conversion + // for target supported reduction ops + if (reductionOps.size() == 1) { + auto redOp = reductionOps[0]; + + auto inputType = cast(op.getInputs()[0].getType()); + auto elementType = inputType.getElementType(); + // Check if linalg-to-mk pass has finished conversion for target supported + // reduction ops + assert(!(isReductionOpAndTypeSupportedByTarget(redOp, elementType) || + isReduceToElementWiseOpAndTypeSupportedByTarget( + redOp, elementType, inputType.getNumElements(), + inputType.getRank()))); + } + + return convertLinalgOpToLoops(op, rewriter); + } +}; + +template +struct MKReduceOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename MKOpT::Adaptor; + +public: + LogicalResult + matchAndRewrite(MKOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto input = op.getSrc(); + auto inputType = dyn_cast(input.getType()); + if (!isSupportedType(inputType)) { + return failure(); + } + // TODO: Check init buffer has no init value + + auto axis = op.getAxis(); + assert(axis == 3 || axis == 2); + auto loc = op->getLoc(); + + auto srcPtr = createAddressFromMemref(rewriter, loc, input); + auto outputPtr = createAddressFromMemref(rewriter, loc, op.getInit()); + + auto format = getFormatCode(inputType); + + rewriter.replaceOpWithNewOp( + op, TypeRange{}, srcPtr, outputPtr, + rewriter.getUI32IntegerAttr(axis == 3 ? 0 /*reduce C dim*/ + : 1 /*reduce W dim*/), + op.getNhwcShapeAttr(), rewriter.getI16IntegerAttr(format)); + return success(); + } +}; + +template +struct BarrierConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename MKOpT::Adaptor; + + LogicalResult + matchAndRewrite(MKOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + rewriter.create(loc); + rewriter.eraseOp(op); + + return success(); + } +}; + +struct RemoteLoadConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mk::RemoteLoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + + // mk.remote_load has 5 operands: + // 4 I64 coords (indices 0-3) + 1 dst (index 4). + Value dstVal = adaptor.getOperands()[4]; + + // Compute elem_bytes and data_size (in bytes) from the dst shaped type. + // Prefer compile-time constants for static shapes; fall back to runtime + // computation using extracted sizes if needed. + Type dstOrigTy = op.getDst().getType(); + ShapedType shapedTy = dyn_cast(dstOrigTy); + if (!shapedTy) + return rewriter.notifyMatchFailure( + op, "mk.remote_load dst must be shaped type"); + + int64_t elemBytesConst = + static_cast(getElemByte(shapedTy.getElementType())); + Value elemBytesI32 = rewriter.create( + loc, rewriter.getI32Type(), elemBytesConst); + + Value dataSizeI64; + if (shapedTy.hasStaticShape()) { + int64_t numElems = shapedTy.getNumElements(); + int64_t totalBytes = numElems * elemBytesConst; + dataSizeI64 = rewriter.create( + loc, rewriter.getI64Type(), totalBytes); + } else { + // Dynamic shape: compute element count from runtime sizes. + // Requires memref operand to extract metadata. + if (!isa(dstVal.getType())) + return rewriter.notifyMatchFailure( + op, "dynamic-shaped remote_load requires memref dst"); + auto [basePtr, sizes, strides] = createMetadata(rewriter, loc, dstVal); + (void)basePtr; + (void)strides; + Value elemCount = calculateElemCount(rewriter, loc, sizes); + Value elemCountI64 = rewriter.create( + loc, rewriter.getI64Type(), elemCount); + Value elemBytesI64 = rewriter.create( + loc, rewriter.getI64Type(), elemBytesI32); + dataSizeI64 = rewriter.create(loc, elemCountI64.getType(), + elemCountI64, elemBytesI64); + } + + // Convert dst memref to address (I64) + Value dstAddr = createAddressFromMemref(rewriter, loc, dstVal); + + // Create wafer.remote_load operation. + rewriter.create( + loc, + adaptor.getOperands()[0], // remote_chip_id_x + adaptor.getOperands()[1], // remote_chip_id_y + adaptor.getOperands()[2], // remote_die_id + adaptor.getOperands()[3], // remote_tile_id + dstAddr, // dst (I64 address) + elemBytesI32, // elem_bytes (I32) + dataSizeI64 // data_size (I64) + ); + + // mk.remote_load has results; wafer.remote_load is void. Replace results with + // dst. (converted) dst operand value, which represents the destination + // buffer. + if (op->getNumResults() > 0) { + SmallVector repl; + repl.reserve(op->getNumResults()); + for (unsigned i = 0; i < op->getNumResults(); ++i) + repl.push_back(dstVal); + rewriter.replaceOp(op, repl); + } else { + rewriter.eraseOp(op); + } + + return success(); + } +}; + +struct RemoteStoreConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mk::RemoteStoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + + // mk.remote_store has 6 operands: + // 4 I64 coords (indices 0-3) + 1 dst_addr (index 4) + 1 src (index 5) + Value dstAddrVal = adaptor.getOperands()[4]; + Value srcVal = adaptor.getOperands()[5]; + + // Compute elem_bytes and data_size (in bytes) from the src shaped type. + Type srcOrigTy = op.getSrc().getType(); + ShapedType shapedTy = dyn_cast(srcOrigTy); + if (!shapedTy) + return rewriter.notifyMatchFailure( + op, "mk.remote_store src must be shaped type"); + + int64_t elemBytesConst = + static_cast(getElemByte(shapedTy.getElementType())); + Value elemBytesI32 = rewriter.create( + loc, rewriter.getI32Type(), elemBytesConst); + + Value dataSizeI64; + if (shapedTy.hasStaticShape()) { + int64_t numElems = shapedTy.getNumElements(); + int64_t totalBytes = numElems * elemBytesConst; + dataSizeI64 = rewriter.create( + loc, rewriter.getI64Type(), totalBytes); + } else { + if (!isa(srcVal.getType())) + return rewriter.notifyMatchFailure( + op, "dynamic-shaped remote_store requires memref src"); + auto [basePtr, sizes, strides] = createMetadata(rewriter, loc, srcVal); + (void)basePtr; + (void)strides; + Value elemCount = calculateElemCount(rewriter, loc, sizes); + Value elemCountI64 = rewriter.create( + loc, rewriter.getI64Type(), elemCount); + Value elemBytesI64 = rewriter.create( + loc, rewriter.getI64Type(), elemBytesI32); + dataSizeI64 = rewriter.create(loc, elemCountI64.getType(), + elemCountI64, elemBytesI64); + } + + // Convert src memref to address (I64) + Value srcAddr = createAddressFromMemref(rewriter, loc, srcVal); + + // Convert dst "addr-like" to I64 address. + Value dstAddr; + if (dstAddrVal.getType().isInteger(64)) { + dstAddr = dstAddrVal; + } else if (isa(dstAddrVal.getType())) { + dstAddr = createAddressFromMemref(rewriter, loc, dstAddrVal); + } else { + return rewriter.notifyMatchFailure( + op, "mk.remote_store dst_addr must be i64 or memref at MKToWafer"); + } + + // Create wafer.remote_store operation directly with the destination address. + rewriter.create( + loc, + adaptor.getOperands()[0], // remote_chip_id_x + adaptor.getOperands()[1], // remote_chip_id_y + adaptor.getOperands()[2], // remote_die_id + adaptor.getOperands()[3], // remote_tile_id + dstAddr, // dst (I64 address) + srcAddr, // src (I64 address) + elemBytesI32, // elem_bytes (I32) + dataSizeI64 // data_size (I64) + ); + + // mk.remote_store has no results, just erase it + rewriter.eraseOp(op); + + return success(); + } +}; + +struct PrintConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(mk::PrintOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op->getLoc(); + + // printf scalar value. + if (printScalar(op)) { + if (op.getNumOperands() == 0) { + createRuntimePrintScalarCall(rewriter, op.getPrefix(), std::nullopt); + } else { + createRuntimePrintScalarCall(rewriter, op.getPrefix(), + adaptor.getOperands()[0], op.getHex(), + op.getIsSigned()[0]); + } + rewriter.eraseOp(op); + return success(); + } + + // print memref value. + createPrintMemrefCall(op, rewriter); + + rewriter.eraseOp(op); + return success(); + } + +private: + static std::string getFormatSubstr(Type type, bool hex = false, + std::optional width = std::nullopt, + bool isSigned = false) { + // If the `value` is a pointer, just return %p. + if (isa(type)) { + return "%p"; + } + // Hex is "0x%0nx" or "0x%0nllx", where n is the number of hex digits in + // the type (so 4 for fp16, 8 for int32, 16 for int64). + if (hex) { + // Ignore `width` for `hex` values, pad to typeWidth. + std::string ret = + "0x%0" + std::to_string(type.getIntOrFloatBitWidth() / 4); + if (type.getIntOrFloatBitWidth() > 32) { + ret += "ll"; + } + ret += "x"; + return ret; + } + + std::string prefix = "%"; + if (width.has_value()) { + prefix += std::to_string(*width); + } + + if (type.isBF16() || type.isF16() || type.isF32() || type.isF64()) { + return prefix + "f"; + } else if (type.isInteger()) { + if (type.getIntOrFloatBitWidth() == 64) + return prefix + (isSigned ? "lli" : "llu"); + else + return prefix + (isSigned ? "i" : "u"); + } + assert(false && "not supported type"); + return ""; + } + + // C varargs require integers of at least 32 bits and floating-point values + // promoted to double. Scalar and tensor printing must use the same ABI. + static Value printfPromoteValue(RewriterBase &rewriter, Value value, + bool hex, bool isSigned) { + auto *context = rewriter.getContext(); + auto type = value.getType(); + auto loc = UnknownLoc::get(context); + auto b = LLVM::TritonLLVMOpBuilder(loc, rewriter); + + if (hex && isa(type)) { + type = IntegerType::get(context, type.getIntOrFloatBitWidth()); + value = rewriter.create(loc, type, value); + } + if (type.isIntOrIndex() && type.getIntOrFloatBitWidth() < 32) { + if (hex || !isSigned) { + return b.zext(i32_ty, value); + } else { + return b.sext(i32_ty, value); + } + } else if (type.isBF16() || type.isF16() || type.isF32()) { + return b.fpext(f64_ty, value); + } + + return value; + } + + static LLVM::LLVMFuncOp + getOrAddPrintFuncDecl(ConversionPatternRewriter &rewriter, + StringRef funcName = "__Print") { + auto moduleOp = + rewriter.getBlock()->getParent()->getParentOfType(); + Operation *funcOp = moduleOp.lookupSymbol(funcName); + if (funcOp) + return cast(*funcOp); + + auto *ctx = rewriter.getContext(); + SmallVector argsType = {ptr_ty(ctx)}; + auto funcType = + LLVM::LLVMFunctionType::get(i32_ty, argsType, /*isVarArg*/ true); + + ConversionPatternRewriter::InsertionGuard guard(rewriter); + rewriter.setInsertionPointToStart(moduleOp.getBody()); + + return rewriter.create(UnknownLoc::get(ctx), funcName, + funcType); + } + + static bool printScalar(mk::PrintOp op) { + // Simply use printf if no operand or the operand is scalar. + if (op.getNumOperands() == 0) + return true; + + assert(op.getNumOperands() == 1); + Type oprType = op.getOperands()[0].getType(); + return (oprType.isIntOrIndexOrFloat() || isa(oprType)); + } + + static void createRuntimePrintScalarCall(ConversionPatternRewriter &rewriter, + StringRef prefix, + std::optional arg, + bool hex = false, + bool isSigned = false) { + assert(!prefix.empty() && "printf with empty string not supported"); + auto loc = UnknownLoc::get(rewriter.getContext()); + auto b = LLVM::TritonLLVMOpBuilder(loc, rewriter); + + std::string formatStr; + llvm::raw_string_ostream os(formatStr); + os << prefix; + if (arg.has_value()) + os << getFormatSubstr(arg.value().getType(), hex, std::nullopt, isSigned); + + llvm::SmallString<64> formatStrNewline(formatStr); + formatStrNewline.push_back('\n'); + formatStrNewline.push_back('\0'); + Value formatStrValue = LLVM::addStringToModule( + loc, rewriter, "printfFormat_", formatStrNewline); + + SmallVector allArgs{formatStrValue}; + if (arg.has_value()) + allArgs.push_back(printfPromoteValue(rewriter, arg.value(), hex, isSigned)); + b.call(getOrAddPrintFuncDecl(rewriter), allArgs); + } + + static LLVM::LLVMFunctionType getPrintfType(MLIRContext *context) { + auto llvmI32Ty = IntegerType::get(context, 32); + auto llvmPtr = LLVM::LLVMPointerType::get(context); + return LLVM::LLVMFunctionType::get(llvmI32Ty, llvmPtr, true); + } + + static FlatSymbolRefAttr getOrInsertPrintf(PatternRewriter &rewriter, + ModuleOp module, + StringRef funcName = "__Print") { + auto *context = module.getContext(); + if (module.lookupSymbol(funcName)) + return SymbolRefAttr::get(context, funcName); + + PatternRewriter::InsertionGuard insertGuard(rewriter); + rewriter.setInsertionPointToStart(module.getBody()); + rewriter.create(module.getLoc(), funcName, + getPrintfType(context)); + return SymbolRefAttr::get(context, funcName); + } + + static Value getOrCreateGlobalString(Location loc, OpBuilder &builder, + StringRef name, StringRef value, + ModuleOp module) { + LLVM::GlobalOp global; + if (!(global = module.lookupSymbol(name))) { + OpBuilder::InsertionGuard insertGuard(builder); + builder.setInsertionPointToStart(module.getBody()); + auto type = LLVM::LLVMArrayType::get( + IntegerType::get(builder.getContext(), 8), value.size()); + global = builder.create(loc, type, true, + LLVM::Linkage::Internal, name, + builder.getStringAttr(value), 0); + } + + Value globalPtr = builder.create(loc, global); + Value cst0 = builder.create(loc, builder.getI64Type(), + builder.getIndexAttr(0)); + return builder.create( + loc, LLVM::LLVMPointerType::get(builder.getContext()), global.getType(), + globalPtr, ArrayRef({cst0, cst0})); + } + + static void createPrintMemrefCall(mk::PrintOp op, + ConversionPatternRewriter &rewriter) { + auto loc = op->getLoc(); + auto context = rewriter.getContext(); + auto memRefType = llvm::cast(*op->operand_type_begin()); + auto memRefShape = memRefType.getShape(); + Type memElementType = memRefType.getElementType(); + ModuleOp parentModule = op->getParentOfType(); + + auto printfRef = getOrInsertPrintf(rewriter, parentModule); + std::string formatSpecifierStr = getFormatSubstr( + memElementType, op.getHex(), std::nullopt, op.getIsSigned()[0]); + formatSpecifierStr += ' '; + formatSpecifierStr.push_back('\0'); + auto prefix = op.getPrefix(); + std::string prefixNewline = "\n" + prefix.str(); + prefixNewline.push_back('\0'); + Value prefixValue = getOrCreateGlobalString( + loc, rewriter, "frmt_prefix" + prefix.str(), + StringRef(prefixNewline), parentModule); + Value formatSpecifierCst = getOrCreateGlobalString( + loc, rewriter, "frmt_spec" + StringRef(formatSpecifierStr).drop_back().str(), + StringRef(formatSpecifierStr), parentModule); + Value newLineCst = getOrCreateGlobalString( + loc, rewriter, "nl", StringRef("\n\0", 2), parentModule); + + // print prefix firstly. + rewriter.create(loc, getPrintfType(context), printfRef, + prefixValue); + + SmallVector loopIvs; + for (unsigned i = 0, e = memRefShape.size(); i != e; ++i) { + auto lowerBound = rewriter.create(loc, 0); + auto upperBound = + rewriter.create(loc, memRefShape[i]); + auto step = rewriter.create(loc, 1); + auto loop = + rewriter.create(loc, lowerBound, upperBound, step); + for (Operation &nested : *loop.getBody()) + rewriter.eraseOp(&nested); + loopIvs.push_back(loop.getInductionVar()); + + rewriter.setInsertionPointToEnd(loop.getBody()); + + if (i != e - 1) + rewriter.create(loc, getPrintfType(context), printfRef, + newLineCst); + rewriter.create(loc); + rewriter.setInsertionPointToStart(loop.getBody()); + } + + Value elementLoad = + rewriter.create(loc, op.getOperands()[0], loopIvs); + elementLoad = printfPromoteValue(rewriter, elementLoad, op.getHex(), + op.getIsSigned()[0]); + rewriter.create( + loc, getPrintfType(context), printfRef, + ArrayRef({formatSpecifierCst, elementLoad})); + } +}; +struct AtomicRMWOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LLVM::AtomicBinOp getAtomicBinOp(RMWOp op, Type type) const { + switch (op) { + case RMWOp::AND: + return LLVM::AtomicBinOp::_and; + case RMWOp::OR: + return LLVM::AtomicBinOp::_or; + case RMWOp::XOR: + return LLVM::AtomicBinOp::_xor; + case RMWOp::ADD: + return LLVM::AtomicBinOp::add; + case RMWOp::FADD: + return LLVM::AtomicBinOp::fadd; + case RMWOp::MAX: + return type.isIntOrIndex() ? LLVM::AtomicBinOp::max + : LLVM::AtomicBinOp::fmax; + case RMWOp::MIN: + return type.isIntOrIndex() ? LLVM::AtomicBinOp::min + : LLVM::AtomicBinOp::fmin; + case RMWOp::UMAX: + return LLVM::AtomicBinOp::umax; + case RMWOp::UMIN: + return LLVM::AtomicBinOp::umin; + case RMWOp::XCHG: + return LLVM::AtomicBinOp::xchg; + default: + llvm_unreachable("Unexpected atomic op"); + } + } + + LogicalResult + matchAndRewrite(mk::AtomicRMWOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto val = adaptor.getVal(); + + // TODO: Support sem and scope + if (isa(val.getType())) + return failure(); + + assert(0 && "Now, wafer backend only support memref type. Don't " + "support llvm atomic op conversion."); + auto ptr = + createAddressFromMemref(rewriter, op->getLoc(), adaptor.getPtr()); + + auto opKind = getAtomicBinOp(op.getAtomicRmwOp(), val.getType()); + auto ordering = getOrdering(op.getSem()); + ptr = rewriter.create( + loc, LLVM::LLVMPointerType::get(rewriter.getContext()), ptr); + + rewriter.replaceOpWithNewOp(op, opKind, ptr, val, + ordering); + + return success(); + } +}; + +struct AtomicCASOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mk::AtomicCASOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto cmp = adaptor.getCmp(); + auto val = adaptor.getVal(); + // TODO: Support sem and scope + if (isa(val.getType())) + return failure(); + + assert(0 && "Now, wafer backend only support memref type. Don't " + "support llvm atomic op conversion. It will be supported in " + "the future by using llvm.bitcast and llvm.atomic.cmpxchg."); + auto ptr = + createAddressFromMemref(rewriter, op->getLoc(), adaptor.getPtr()); + + ptr = rewriter.create( + loc, LLVM::LLVMPointerType::get(rewriter.getContext()), ptr); + + auto ordering = getOrdering(op.getSem()); + auto failureOrdering = ordering != LLVM::AtomicOrdering::monotonic + ? LLVM::AtomicOrdering::acquire + : ordering; + // TODO: Use llvm.bitcast to support other types: f32, etc. + Value cmpXchg = rewriter.create( + loc, ptr, cmp, val, ordering, failureOrdering); + Value oldVal = rewriter.create(loc, cmpXchg, 0); + rewriter.replaceOp(op, oldVal); + + return success(); + } +}; +} // namespace + +//===----------------------------------------------------------------------===// +// Legalize magic kernel operations to be convertible to Wafer operations +// patterns +//===----------------------------------------------------------------------===// +namespace { + +struct MKRandGenOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mk::RandGenOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto seed0Type = dyn_cast(op.getSeed0().getType()); + auto seed1Type = dyn_cast(op.getSeed1().getType()); + auto outType = dyn_cast(op.getOut().getType()); + auto seed0OutType = dyn_cast(op.getSeed0Out().getType()); + auto seed1OutType = dyn_cast(op.getSeed1Out().getType()); + if (!seed0Type || !seed1Type || !outType || !seed0OutType || !seed1OutType) + return rewriter.notifyMatchFailure( + op, "mk.randgen expects memref operands after bufferization"); + if (seed0Type.getShape() != ArrayRef({16}) || + seed1Type.getShape() != ArrayRef({16})) + return rewriter.notifyMatchFailure( + op, "mk.randgen seeds must have shape [16]"); + if (!seed0Type.getElementType().isInteger(64) || + !outType.getElementType().isInteger(64)) + return rewriter.notifyMatchFailure( + op, "mk.randgen currently supports only i64 element type"); + + int32_t byteCount = op.getByteCount(); + if (byteCount <= 0 || (byteCount % 128) != 0) + return rewriter.notifyMatchFailure( + op, "mk.randgen byte_count must be a positive multiple of 128"); + if (outType.getNumElements() * 8 != byteCount) + return rewriter.notifyMatchFailure( + op, "mk.randgen out numel * 8 must equal byte_count"); + + Location loc = op.getLoc(); + Value seed0Ptr = createAddressFromMemref(rewriter, loc, op.getSeed0()); + Value seed1Ptr = createAddressFromMemref(rewriter, loc, op.getSeed1()); + Value outPtr = createAddressFromMemref(rewriter, loc, op.getOut()); + Value seed0OutPtr = + createAddressFromMemref(rewriter, loc, op.getSeed0Out()); + Value seed1OutPtr = + createAddressFromMemref(rewriter, loc, op.getSeed1Out()); + + rewriter.replaceOpWithNewOp( + op, TypeRange{}, seed0Ptr, seed1Ptr, seed0OutPtr, seed1OutPtr, outPtr, + rewriter.getI32IntegerAttr(byteCount), + rewriter.getI16IntegerAttr(op.getFmt())); + return success(); + } +}; + + +struct LinalgCopyRewrite : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(linalg::CopyOp op, + PatternRewriter &rewriter) const override { + assert(op.getInputs().size() == 1 && op.getOutputs().size() == 1 && + "LinalgCopyRewrite only supports single input and output"); + rewriter.replaceOpWithNewOp(op, op.getInputs()[0], + op.getOutputs()[0]); + return success(); + } +}; + +} // namespace + +void mlir::triton::populateMKToWaferCanonicalizationPatterns( + RewritePatternSet &patterns) { + // Backend op canonicalization + patterns.add(patterns.getContext()); +} + +void mlir::triton::populateMKToWaferConversionPatterns( + RewritePatternSet &patterns) { + + MKToWaferTypeConverter typeConverter; + + // Add type conversion patterns + populateFunctionOpInterfaceTypeConversionPattern(patterns, + typeConverter); + populateReturnOpTypeConversionPattern(patterns, typeConverter); + populateCallOpTypeConversionPattern(patterns, typeConverter); + + // NOTE: Only convert ops that have been legalized to be convertible to Wafer + // ops. + // Convert only float type input/output ops to Wafer ops except some special + // ops that support any types. + // clang-format off + patterns.add(patterns.getContext()); + + patterns.add, + MKReduceOpConversion, + MKReduceOpConversion, + TransposeOpConversion, + ReciprocalOpConversionPattern, + LinalgFillOpConversion, + MKDotToWaferGemmOpConversion, + MKDequantOpConversionPattern, + MKSigmoidToWaferSigmoidOpConversion, + MKGeluToWaferGeluOpConversion, + MKArgMinMaxConversionPattern, + MKArgMinMaxConversionPattern, + MKMaskMoveOpConversionPattern, + MKBit2FPOpConversionPattern, + MKRelationVSOpConversionPattern, + MKRelationVSOpConversionPattern, + MKRelationVSOpConversionPattern, + MKRelationVVOpConversionPattern, + MKArithVSOpConversionPattern, + MKArithVSOpConversionPattern, + MKArithVSOpConversionPattern, + MKRelationVVOpConversionPattern, + GatherConvertPattern, + BarrierConversion, + BarrierConversion, + BarrierConversion, + PrintConversion, + RemoteStoreConversion, + RemoteLoadConversion, + AtomicRMWOpConversion, + AtomicCASOpConversion>( + patterns.getContext()); + // clang-format on +} diff --git a/third_party/wafer/lib/Conversion/MKToWafer/MKToWaferPass.cpp b/third_party/wafer/lib/Conversion/MKToWafer/MKToWaferPass.cpp new file mode 100755 index 00000000..023e7067 --- /dev/null +++ b/third_party/wafer/lib/Conversion/MKToWafer/MKToWaferPass.cpp @@ -0,0 +1,123 @@ +//===--------------------- MKToWaferPass.cpp -------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/DialectConversion.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "triton-shared/Utils/Utils.h" +#include "wafer/Conversion/MKToWafer/MKToWafer.h" +#include "wafer/Dialect/IR/WaferDialect.h" +#include "llvm/Support/Debug.h" +#include +#include +#include + +#define DEBUG_TYPE "mk-to-wafer" + +using namespace mlir; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_MKTOWAFER +#include "wafer/Conversion/MKToWafer/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class MKToWaferPass : public triton::impl::MKToWaferBase { + using MKToWaferBase::MKToWaferBase; + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + auto *ctx = &getContext(); + + // If disable precision priority mode, we need to legalize integer + // operations to float operations. + RewritePatternSet canonicalizePatterns(ctx); + triton::populateMKToWaferCanonicalizationPatterns(canonicalizePatterns); + + if (failed( + applyPatternsGreedily(moduleOp, std::move(canonicalizePatterns)))) { + signalPassFailure(); + } + + // Use to memory::CopyOp to wafer dialect op + moduleOp->walk([&](Operation *op) { + if (isa(op)) { + auto copyOp = cast(op); + op->setAttr("srcSpm", + BoolAttr::get(ctx, triton::isOperandMemorySpaceSPM( + copyOp.getSource()))); + op->setAttr("dstSpm", + BoolAttr::get(ctx, triton::isOperandMemorySpaceSPM( + copyOp.getTarget()))); + } + }); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + + // Register illegal ops for Dialect Conversion + target.addIllegalDialect(); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + affine::AffineDialect, scf::SCFDialect, memref::MemRefDialect, + cf::ControlFlowDialect, wafer::WaferDialect, LLVM::LLVMDialect>(); + + target.addIllegalOp(); + + target.addLegalOp(); + + target.addLegalOp(); + + triton::populateMKToWaferConversionPatterns(patterns); + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + + // linalg::linalgOpToLoops will generate memref::LoadOp/memref::StoreOp + // before and after the arith calculation. + // Use to check whether add spm mapping offset in + // memref::LoadOp/memref::StoreOp lowering + moduleOp->walk([&](Operation *op) { + if (isa(op)) { + bool isSpm = isa(op) + ? triton::isOperandMemorySpaceSPM(op->getOperand(0)) + : triton::isOperandMemorySpaceSPM(op->getOperand(1)); + + op->setAttr("isSpm", + IntegerAttr::get(IntegerType::get(op->getContext(), 32), + llvm::APInt(32, isSpm))); + } + }); + } +}; + +} // namespace + +std::unique_ptr> triton::createMKToWaferPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/ReconcilePtrCasts/CMakeLists.txt b/third_party/wafer/lib/Conversion/ReconcilePtrCasts/CMakeLists.txt new file mode 100755 index 00000000..e8b7cbd3 --- /dev/null +++ b/third_party/wafer/lib/Conversion/ReconcilePtrCasts/CMakeLists.txt @@ -0,0 +1,19 @@ +add_triton_library(ReconcilePtrCasts + ReconcilePtrCastsPass.cpp + + DEPENDS + ReconcilePtrCastsPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + MLIRReconcileUnrealizedCasts + TritonIR + MLIRAddress +) diff --git a/third_party/wafer/lib/Conversion/ReconcilePtrCasts/ReconcilePtrCastsPass.cpp b/third_party/wafer/lib/Conversion/ReconcilePtrCasts/ReconcilePtrCastsPass.cpp new file mode 100755 index 00000000..075bac77 --- /dev/null +++ b/third_party/wafer/lib/Conversion/ReconcilePtrCasts/ReconcilePtrCastsPass.cpp @@ -0,0 +1,161 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// +// Throughout the conversion process, we convert !tt.ptr -> {!ptr.ptr or +// memref<*>}. This process leaves around unrealized_conversion_cast ops between +// these types. We want to remove these unrealized casts and use the proper +// conversion ops in the PtrDialect: to_memref or from_memref. To do this, we +// use a pattern that simplifies the chain of conversions by removing +// intermediate conversion cast ops. At the end, we are left with just pointer +// to memref or vice versa. We then convert the unrealized cast to to_memref or +// from_memref accordingly. +//===----------------------------------------------------------------------===// + +#include "Address/Dialect/IR/AddressDialect.h" +#include "mlir/Conversion/ReconcileUnrealizedCasts/ReconcileUnrealizedCasts.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/BuiltinDialect.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h" +#include "triton/Dialect/Triton/IR/Types.h" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/ReconcilePtrCasts/Passes.h.inc" + +namespace { + +static bool isOneToOneCast(UnrealizedConversionCastOp op) { + return (op.getInputs().size() == 1 && op->getNumResults() == 1); +} + +struct SimplifyUnrealizedCast + : public OpRewritePattern { + SimplifyUnrealizedCast(MLIRContext *context, PatternBenefit benefit = 1) + : OpRewritePattern(context, benefit) {} + + LogicalResult matchAndRewrite(UnrealizedConversionCastOp op, + PatternRewriter &rewriter) const override { + if (!isOneToOneCast(op)) { + return failure(); + } + auto in = op.getInputs().front(); + + auto unrealizedCast = in.getDefiningOp(); + if (!unrealizedCast) + return failure(); + if (!isOneToOneCast(unrealizedCast)) { + return failure(); + } + + if (!isa(unrealizedCast.getType(0))) + return failure(); + auto prevInput = unrealizedCast.getInputs().front(); + auto newCast = rewriter.create( + op->getLoc(), op->getResultTypes(), ValueRange{prevInput}); + + rewriter.replaceOp(op, newCast); + return success(); + } +}; + +struct FromMemrefConverter + : public OpRewritePattern { + FromMemrefConverter(MLIRContext *context, PatternBenefit benefit = 1) + : OpRewritePattern(context, benefit) {} + + LogicalResult matchAndRewrite(UnrealizedConversionCastOp op, + PatternRewriter &rewriter) const override { + if (!isOneToOneCast(op)) { + return failure(); + } + + auto input = op.getInputs().front(); + auto unrankedInput = dyn_cast(input.getType()); + auto output = op.getResult(0); + auto outType = output.getType(); + + if (unrankedInput && isa(outType)) { + // from_memref only takes ranked memref, cast the unranked memref to + // ranked memref first. + auto rankedMemref = rewriter.create( + op.getLoc(), MemRefType::get({1}, unrankedInput.getElementType()), + input); + auto memrefToPtr = rewriter.create( + op->getLoc(), addr::AddressType::get(rewriter.getContext()), + rankedMemref); + + rewriter.replaceAllUsesWith(output, memrefToPtr); + rewriter.eraseOp(op); + + return success(); + } + + return failure(); + } +}; + +struct ToMemrefConverter : public OpRewritePattern { + ToMemrefConverter(MLIRContext *context, PatternBenefit benefit = 1) + : OpRewritePattern(context, benefit) {} + + LogicalResult matchAndRewrite(UnrealizedConversionCastOp op, + PatternRewriter &rewriter) const override { + if (!isOneToOneCast(op)) { + return failure(); + } + auto input = op.getInputs().front(); + auto inType = input.getType(); + auto output = op.getResult(0); + auto outUnrankedMemrefType = dyn_cast(output.getType()); + if (isa(inType) && outUnrankedMemrefType) { + // to_memref can only cast to ranked static shape memref, we have to cast + // the resulting memref back to unranked + auto elemType = outUnrankedMemrefType.getElementType(); + auto ptrToMemref = rewriter.create( + op->getLoc(), MemRefType::get({1}, elemType), input); + + auto newUnrankedMemref = rewriter.create( + op.getLoc(), MemRefType::get({ShapedType::kDynamic}, elemType), + ptrToMemref); + + rewriter.replaceAllUsesWith(output, newUnrankedMemref); + rewriter.eraseOp(op); + return success(); + } + + return failure(); + } +}; + +class ReconcilePtrCastsPass + : public ReconcilePtrCastsBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + RewritePatternSet patterns(&getContext()); + patterns + .add( + &getContext()); + if (failed(applyPatternsGreedily(moduleOp, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> triton::createReconcilePtrCastsPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/StructuredToMK/CMakeLists.txt b/third_party/wafer/lib/Conversion/StructuredToMK/CMakeLists.txt new file mode 100755 index 00000000..7269a60a --- /dev/null +++ b/third_party/wafer/lib/Conversion/StructuredToMK/CMakeLists.txt @@ -0,0 +1,23 @@ +add_triton_library(StructuredToMK + StructuredToMK.cpp + StructuredToMKPass.cpp + + DEPENDS + StructuredToMKConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRSCFTransforms + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonTilingExtIR + TritonStructuredIR + MLIRAddress +) diff --git a/third_party/wafer/lib/Conversion/StructuredToMK/StructuredToMK.cpp b/third_party/wafer/lib/Conversion/StructuredToMK/StructuredToMK.cpp new file mode 100755 index 00000000..313c2c22 --- /dev/null +++ b/third_party/wafer/lib/Conversion/StructuredToMK/StructuredToMK.cpp @@ -0,0 +1,151 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton/Dialect/Triton/IR/Types.h" + +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/Conversion/StructuredToMK/StructuredToMK.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" + +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/TypeUtilities.h" +#include "mlir/IR/Types.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR//MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" + +#include +#include +#include + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" + +#define DEBUG_TYPE "structured-to-memref" + +using namespace mlir; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/StructuredToMK/Passes.h.inc" + +static memref::SubViewOp getSubview(int rank, ArrayRef dims, + Value source, Location loc, OpBuilder &b) { + auto sourceType = cast(source.getType()); + SmallVector offsets(rank, b.getIndexAttr(0)); + SmallVector strides(rank, b.getIndexAttr(1)); + auto dstType = + memref::SubViewOp::inferResultType(sourceType, offsets, dims, strides); + + return b.create(loc, cast(dstType), source, + offsets, dims, strides); +} + +namespace { + +struct AtomicRMWOpConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + static tensor::ExtractSliceOp + getExtractSlice(int rank, ArrayRef dims, Value source, + const Location loc, OpBuilder &b) { + auto sourceType = cast(source.getType()); + SmallVector offsets(rank, b.getIndexAttr(0)); + SmallVector strides(rank, b.getIndexAttr(1)); + + auto dstType = tensor::ExtractSliceOp::inferResultType(sourceType, offsets, + dims, strides); + + return b.create(loc, dstType, source, offsets, dims, + strides); + } + +public: + LogicalResult + matchAndRewrite(tts::AtomicRMWOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto ptr = adaptor.getPtr(); + auto value = adaptor.getValue(); + + auto type = cast(value.getType()); + auto rank = type.getRank(); + + Value init = rewriter.create(loc, type.getShape(), + type.getElementType()); + + if (op.hasMask()) { + auto mixedDims = op.getMixedMaskDims(); + + auto valueSlice = getExtractSlice(rank, mixedDims, value, loc, rewriter); + auto ptrSubview = getSubview(rank, mixedDims, ptr, loc, rewriter); + + auto atomicRMWOp = rewriter.create( + loc, op.getType(), ptrSubview, valueSlice, init, + op.getAtomicRmwOpAttr(), op.getSemAttr(), op.getScopeAttr()); + rewriter.replaceOp(op, atomicRMWOp); + } else { + auto atomicRMWOp = rewriter.create( + loc, op.getType(), ptr, value, init, op.getAtomicRmwOpAttr(), + op.getSemAttr(), op.getScopeAttr()); + rewriter.replaceOp(op, atomicRMWOp); + } + return success(); + } +}; + +struct AtomicCASOpConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(tts::AtomicCASOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (op.getOffset()) + return failure(); + + auto loc = op.getLoc(); + auto ptr = adaptor.getPtr(); + auto cmp = adaptor.getCmp(); + auto value = adaptor.getValue(); + + auto type = cast(value.getType()); + + Value init = rewriter.create(loc, type.getShape(), + type.getElementType()); + + auto atomicCASOp = rewriter.create( + loc, op.getType(), ptr, cmp, value, init, op.getSemAttr(), + op.getScopeAttr()); + rewriter.replaceOp(op, atomicCASOp); + + return success(); + } +}; + +} // namespace + +void mlir::triton::populateStructuredToMKConversionPatterns( + RewritePatternSet &patterns, TypeConverter &typeConverter) { + patterns.add( + patterns.getContext()); +} diff --git a/third_party/wafer/lib/Conversion/StructuredToMK/StructuredToMKPass.cpp b/third_party/wafer/lib/Conversion/StructuredToMK/StructuredToMKPass.cpp new file mode 100755 index 00000000..82b230d1 --- /dev/null +++ b/third_party/wafer/lib/Conversion/StructuredToMK/StructuredToMKPass.cpp @@ -0,0 +1,151 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "Address/Dialect/IR/AddressDialect.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/Support/LogicalResult.h" +#include "triton-shared/Conversion/StructuredToMK/StructuredToMK.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "utils/TypeConvertor.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/Pass/PassManager.h" +#include "triton/Dialect/Triton/IR/Types.h" +#include "llvm/Support/Casting.h" + +#include + +#define DEBUG_TYPE "structured-to-memref" + +using namespace mlir; +using namespace triton; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_STRUCTUREDTOMK +#include "triton-shared/Conversion/StructuredToMK/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class LoopTypeConverter : public TypeConverter { +public: + LoopTypeConverter(MLIRContext *context) { + // The order of type conversion is important: later ones are tried earlier. + addConversion([](Type type) { return type; }); + + // A tensor of pointers can be passed in as scf.for's init-args, in such + // cases, we convert the type to a memref with dynamic offsets and + // strides. + addConversion( + [context](RankedTensorType tensorType) -> std::optional { + if (auto ptrType = llvm::dyn_cast( + tensorType.getElementType())) { + auto layout = StridedLayoutAttr::get( + context, ShapedType::kDynamic, + SmallVector(tensorType.getRank(), + ShapedType::kDynamic)); + Type elemType = ptrType.getPointeeType(); + return MemRefType::get(tensorType.getShape(), elemType, layout); + } + + return std::nullopt; + }); + + addSourceMaterialization([&](OpBuilder &builder, Type resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + + addArgumentMaterialization([&](OpBuilder &builder, Type resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + + // Convert the current memref type to a memref type with dynamic offsets and + // strides through another reinterpret_cast with the same offsets. + // Canonicalization will simplify this sequence by removing the inital + // reinterpret_cast. + addTargetMaterialization([&](OpBuilder &builder, MemRefType memrefType, + ValueRange inputs, Location loc) -> Value { + auto reinterpretCast = + inputs[0].getDefiningOp(); + if (!reinterpretCast) { + return builder + .create(loc, memrefType, inputs) + .getResult(0); + } + return builder.create( + loc, memrefType, inputs[0], reinterpretCast.getMixedOffsets()[0], + reinterpretCast.getMixedSizes(), reinterpretCast.getMixedStrides()); + }); + } +}; + +class StructuredToMKPass + : public triton::impl::StructuredToMKBase { + using StructuredToMKBase::StructuredToMKBase; + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, ttx::TritonTilingExtDialect, + memref::MemRefDialect, mk::MagicKernelDialect>(); + + target.addIllegalOp(); + + target.addLegalOp(); + + LoopTypeConverter loopTypeConverter(patterns.getContext()); + + mlir::scf::populateSCFStructuralTypeConversionsAndLegality( + loopTypeConverter, patterns, target); + + PtrToUnrankedMemrefConverter typeConverter; + triton::populateStructuredToMKConversionPatterns(patterns, typeConverter); + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> triton::createStructuredToMKPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/StructuredToMemref/CMakeLists.txt b/third_party/wafer/lib/Conversion/StructuredToMemref/CMakeLists.txt new file mode 100755 index 00000000..65f8580a --- /dev/null +++ b/third_party/wafer/lib/Conversion/StructuredToMemref/CMakeLists.txt @@ -0,0 +1,23 @@ +add_triton_library(StructuredToMemref + StructuredToMemref.cpp + StructuredToMemrefPass.cpp + + DEPENDS + StructuredToMemrefConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRSCFTransforms + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonTilingExtIR + TritonStructuredIR + MLIRAddress +) diff --git a/third_party/wafer/lib/Conversion/StructuredToMemref/StructuredToMemref.cpp b/third_party/wafer/lib/Conversion/StructuredToMemref/StructuredToMemref.cpp new file mode 100755 index 00000000..fb992b25 --- /dev/null +++ b/third_party/wafer/lib/Conversion/StructuredToMemref/StructuredToMemref.cpp @@ -0,0 +1,966 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h" +#include "Address/Dialect/IR/AddressDialect.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/TypeUtilities.h" +#include "mlir/IR/Types.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR//MemRef.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/Support/Debug.h" + +#include +#include +#include + +#define DEBUG_TYPE "structured-to-memref" + +using namespace mlir; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +static const std::string WRAP_SIDE_BY_SIDE = "wrap_side_by_side"; +static const std::string WRAP_STACKED = "wrap_stacked"; + +static memref::SubViewOp getSubview(int rank, ArrayRef dims, + Value source, Location loc, OpBuilder &b) { + auto sourceType = cast(source.getType()); + SmallVector offsets(rank, b.getIndexAttr(0)); + SmallVector strides(rank, b.getIndexAttr(1)); + auto dstType = + memref::SubViewOp::inferResultType(sourceType, offsets, dims, strides); + + return b.create(loc, cast(dstType), source, + offsets, dims, strides); +} + +static OpFoldResult accumulateTargetOffset(tts::MakeTensorPtrOp op, + OpBuilder &b) { + Location loc = op->getLoc(); + OpFoldResult targetOffset = b.getIndexAttr(0); + for (auto o : op.getMixedOffsets()) { + targetOffset = addOFRs(targetOffset, o, loc, b); + } + return targetOffset; +} + +namespace { + +struct MakeTensorPtrConverter + : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + static Type getElementTypeStructuredPtr(tts::MakeTensorPtrOp op) { + assert(!op.isBlockPtr()); + // tensor<1024x!tt.ptr> + auto ptrType = cast( + cast(op.getType()).getElementType()); + return ptrType.getPointeeType(); + } + + static Type getElementTypeBlockPtr(tts::MakeTensorPtrOp op) { + assert(op.isBlockPtr()); + // !tt.ptr, 1> + auto shapedType = cast( + cast(op.getType()).getPointeeType()); + return shapedType.getElementType(); + } + + static MemRefType getResultMemrefType(tts::MakeTensorPtrOp op, int64_t offset, + ArrayRef staticStrides, + ArrayRef resultShape) { + auto layout = + StridedLayoutAttr::get(op.getContext(), offset, staticStrides); + Type elemType; + if (op.isBlockPtr()) { + elemType = getElementTypeBlockPtr(op); + } else { + elemType = getElementTypeStructuredPtr(op); + } + return MemRefType::get(resultShape, elemType, layout); + } + + // If there are dimensions with size 1 and stride 0, replace 0 stride with + // the product of sizes of all lower dimensions. This avoids creating memref + // with zero stride. + static llvm::SmallVector + getMixedStridesForMemref(tts::MakeTensorPtrOp op, OpBuilder &b) { + llvm::SmallVector strides; + auto accumulate = 1; + for (auto [size, stride] : + llvm::reverse(llvm::zip(op.getSizes(), op.getMixedStrides()))) { + auto strideIntAttr = getIntAttr(stride); + if (size == 1 && strideIntAttr && strideIntAttr.value() == 0) { + strides.push_back(b.getIndexAttr(accumulate)); + } else if (auto v = llvm::dyn_cast_if_present(stride)) { + OpFoldResult result = getAsOpFoldResult(v); + strides.push_back(result); + } else { + strides.push_back(stride); + } + accumulate *= size; + } + std::reverse(strides.begin(), strides.end()); + return strides; + } + + LogicalResult rewritePtr(ArrayRef resultShape, bool isBlockPtr, + tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + + auto mixedStrides = getMixedStridesForMemref(op, rewriter); + SmallVector staticStrides; + SmallVector dynamicStrides; + dispatchIndexOpFoldResults(mixedStrides, dynamicStrides, staticStrides); + + auto targetOffset = accumulateTargetOffset(op, rewriter); + auto staticTargetOffset = getIntAttr(targetOffset); + auto resultType = getResultMemrefType( + op, staticTargetOffset.value_or(ShapedType::kDynamic), staticStrides, + resultShape); + + auto castOp = rewriter.create( + op.getLoc(), resultType, adaptor.getBase(), targetOffset, + op.getMixedSizes(), mixedStrides); + + rewriter.replaceOp(op, castOp); + + return success(); + } + + LogicalResult + rewriteStructuredPtr(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + ArrayRef resultShape = cast(op.getType()).getShape(); + return rewritePtr(resultShape, false, op, adaptor, rewriter); + } + + LogicalResult rewriteBlockPtr(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + // Block pointers are basically the same as structured pointers except that + // the return types are !tt.ptr> instead of + // tensor> + ArrayRef resultShape = + cast( + cast(op.getType()).getPointeeType()) + .getShape(); + return rewritePtr(resultShape, true, op, adaptor, rewriter); + } + +public: + MakeTensorPtrConverter(const TypeConverter &typeConverter, + MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // TODO: Order is a compiler hint. We can optimize data load/store according + // the order attribute. + if (op.isBlockPtr()) { + return rewriteBlockPtr(op, adaptor, rewriter); + } + + if (op.isStructuredPtr()) { + return rewriteStructuredPtr(op, adaptor, rewriter); + } + + if (op.isSplitPtr()) { + return success(); + } + + return failure(); + } +}; + +memref::SubViewOp createSubview(Value src, ArrayRef offsets, + ArrayRef sizes, + ArrayRef strides, Location loc, + ConversionPatternRewriter &rewriter) { + auto srcType = cast(src.getType()); + auto dstType = + memref::SubViewOp::inferResultType(srcType, offsets, sizes, strides); + return rewriter.create(loc, cast(dstType), src, + offsets, sizes, strides); +} + +Value createCastOps(tts::MakeTensorPtrOp op, + ConversionPatternRewriter &rewriter, Value start, + SmallVector sizesValues, + SmallVector strideVals) { + + Type elemType = + cast(op.getBase().getType()).getPointeeType(); + + auto unrankedMemrefType = UnrankedMemRefType::get(elemType, 0); + // WARNING: TypeConverter cannot automatically insert + // `UnrealizedConversionCastOp` through the materialization mechanism. + auto unrankedMemref = rewriter + .create( + op->getLoc(), unrankedMemrefType, op.getBase()) + ->getResults()[0]; + + auto layout = StridedLayoutAttr::get( + op.getContext(), ShapedType::kDynamic, + SmallVector(sizesValues.size(), ShapedType::kDynamic)); + MemRefType resultType = MemRefType::get( + SmallVector(sizesValues.size(), ShapedType::kDynamic), elemType, + layout); + auto block = rewriter.create( + op->getLoc(), resultType, unrankedMemref, start, sizesValues, strideVals); + return block; +} + +std::pair +getMemSubviews(SmallVector &dims, Value block, Location loc, + int64_t splitDim, ConversionPatternRewriter &rewriter) { + + auto rank = dims.size(); + OpFoldResult maskSize = + rewriter.create(loc, block, splitDim).getResult(); + + OpFoldResult subviewDimFull = dims[splitDim]; + OpFoldResult subviewDim = minOFRs(maskSize, subviewDimFull, loc, rewriter); + + SmallVector offsets(rank, rewriter.getIndexAttr(0)); + SmallVector strides(rank, rewriter.getIndexAttr(1)); + + SmallVector sizes(dims.begin(), dims.end()); + sizes[splitDim] = subviewDim; + + auto sv = createSubview(block, offsets, sizes, strides, loc, rewriter); + auto remainMask = rewriter.create( + loc, ofrToIndexValue(subviewDimFull, loc, rewriter), + ofrToIndexValue(subviewDim, loc, rewriter)); + + return {sv, remainMask}; +} + +void createMemCopies(Value block, Value dst, Location loc, + ConversionPatternRewriter &rewriter, Value &dstOffset, + int64_t splitDim, bool isLoadToDst) { + auto zero = rewriter.create(loc, rewriter.getIndexAttr(0)); + + auto one = rewriter.create(loc, rewriter.getIndexAttr(1)); + + auto rank = cast(dst.getType()).getRank(); + SmallVector blockShape; + for (int i = 0; i < rank; i++) { + blockShape.push_back(rewriter.create(loc, block, i)); + } + + SmallVector dstOffsets(rank, zero); + dstOffsets[splitDim] = dstOffset; + + auto blockDst = + rewriter.create(loc, dst, + /* offsets */ + dstOffsets, + /* sizes */ + blockShape, + /* strides */ + SmallVector(rank, one)); + dstOffset = + rewriter.create(loc, dstOffset, blockShape[splitDim]); + + if (isLoadToDst) { + rewriter.create(loc, block, blockDst); + } else { + rewriter.create(loc, blockDst, block); + } +} + +Value processMemSubviewCopies(Location loc, ConversionPatternRewriter &rewriter, + Value alloc, Value block, + SmallVector mixedDims, + Value &allocOffset, int64_t splitDim, + bool isLoadToDst) { + + Value subview; + Value remainMask; + if (mixedDims.empty()) { + subview = block; + } else { + auto res = getMemSubviews(mixedDims, block, loc, splitDim, rewriter); + subview = res.first; + remainMask = res.second; + } + + createMemCopies(subview, alloc, loc, rewriter, allocOffset, splitDim, + isLoadToDst); + return remainMask; +} + +void rewriteSideBySideMemAccess(tts::MakeTensorPtrOp makeTensorPtrOp, + ConversionPatternRewriter &rewriter, + Value alloc, + SmallVector mixedDims, + bool isLoadToDst) { + assert(makeTensorPtrOp.getStaticShape().size() == 1 || + makeTensorPtrOp.getStaticShape()[0] == 0); + auto loc = makeTensorPtrOp->getLoc(); + auto targetOffset = ofrToIndexValue( + accumulateTargetOffset(makeTensorPtrOp, rewriter), loc, rewriter); + + //////////////////////////////////////////////////////////////////////////// + // + // Handling side-by-side wraparound + // + // Same limitations apply to the stacked wraparound case. + // + //////////////////////////////////////////////////////////////////////////// + // + // nextOffset - targetOffset = colSize + // d1 + d2 = colSize + // N + // x clampedOffset + // --------------------------*----------------*-----* + // | | nextOffset (might + // | targetOffset | overflow) + // y *----- *----------------| + // | | | | + // M |----- -----------------| + // | d2 d1 | + // -------------------------------------------- + // + // x = targetOffset % N + // offset_dim_0 = scaled_offset_0 + // offset_dim_1 = scaled_offset_1 + // col_start = scaled_offset_1 % N + // remainSize = colSize + // size = N - col_start + // while (remainSize > size): + // reinterpret (col_start, size, stride ) + // remainSize = remainSize - size + // col_start = (scaled_offset_0 + size) %N + // size = N + // + // reinterpret (col_start, remainSize, stride ) + // + //////////////////////////////////////////////////////////////////////////// + auto rank = cast(makeTensorPtrOp.getType()).getRank(); + SmallVector scaledOffset(rank); + auto offsets = makeTensorPtrOp.getMixedOffsets(); + std::transform(offsets.begin(), offsets.end(), scaledOffset.begin(), + [&](auto val) { return ofrToIndexValue(val, loc, rewriter); }); + auto lastDim = rank - 1; + + assert(rank == makeTensorPtrOp.getSizes().size()); + // Data block shape to be read + auto sizesInt = makeTensorPtrOp.getSizes(); + SmallVector sizesValues(rank); + std::transform(sizesInt.begin(), sizesInt.end(), sizesValues.begin(), + [&](auto val) { + return rewriter.create( + loc, rewriter.getIndexAttr(val)); + }); + // Total side by side size + Value totalSize = sizesValues.back(); + + // NOTE: We use `scaledOffset[lastdim]` for modulo because the loop and the + // modulo dimension cannot be in the same dimension (otherwise ptrAnalysis + // cannot analyze `make_tptr` of `splitMemory`). + Value N = + ofrToIndexValue(makeTensorPtrOp.getMixedShape()[lastDim], loc, rewriter); + Value x = rewriter.create(loc, scaledOffset[lastDim], N); + Value y = + rewriter.create(loc, targetOffset, scaledOffset[lastDim]); + SmallVector strideVals = + ofrsToIndexValues(makeTensorPtrOp.getMixedStrides(), loc, rewriter); + + Value remainSize = totalSize; + Value size = rewriter.create(loc, N, x); + Value colStart = rewriter.create(loc, y, x); + Value allocOffset = rewriter.create(loc, 0); + SmallVector typeR{remainSize.getType(), size.getType(), + colStart.getType(), allocOffset.getType()}; + SmallVector valueR{remainSize, size, colStart, allocOffset}; + if (!mixedDims.empty()) { + Value mixedDim = ofrsToIndexValues(mixedDims[lastDim], loc, rewriter)[0]; + typeR.push_back(mixedDim.getType()); + valueR.push_back(mixedDim); + } + auto whileOp = rewriter.create( + loc, typeR, valueR, + /*beforeBuilder=*/ + [&](OpBuilder &b, Location loc, ValueRange args) { + Value cond = b.create(loc, arith::CmpIPredicate::sgt, + args[0], args[1]); + b.create(loc, cond, args); + }, + /*afterBuilder=*/ + [&](OpBuilder &b, Location loc, ValueRange args) { + Value remainSize = args[0]; + Value size = args[1]; + Value colStart = args[2]; + Value allocOffset = args[3]; + sizesValues[lastDim] = size; + if (!mixedDims.empty()) { + Value mixedDim = args[4]; + mixedDims[lastDim] = mixedDim; + } + Value block = createCastOps(makeTensorPtrOp, rewriter, colStart, + sizesValues, strideVals); + Value mixedDim = + processMemSubviewCopies(loc, rewriter, alloc, block, mixedDims, + allocOffset, lastDim, isLoadToDst); + remainSize = b.create(loc, remainSize, size); + colStart = b.create(loc, colStart, size); + colStart = b.create(loc, colStart, N); + colStart = b.create(loc, colStart, y); + size = N; + SmallVector newArgs{remainSize, size, colStart, allocOffset}; + if (!mixedDims.empty()) { + newArgs.push_back(mixedDim); + } + b.create(loc, newArgs); + }); + + remainSize = whileOp->getResult(0); + colStart = whileOp->getResult(2); + sizesValues[lastDim] = remainSize; + allocOffset = whileOp->getResult(3); + if (!mixedDims.empty()) { + mixedDims[lastDim] = whileOp->getResult(4); + } + + Value block = createCastOps(makeTensorPtrOp, rewriter, colStart, sizesValues, + strideVals); + processMemSubviewCopies(loc, rewriter, alloc, block, mixedDims, allocOffset, + lastDim, isLoadToDst); +} + +void rewriteStackedMemAccess(tts::MakeTensorPtrOp makeTensorPtrOp, + ConversionPatternRewriter &rewriter, Value alloc, + SmallVector mixedDims, + bool isLoadToDst) { + assert(makeTensorPtrOp.getStaticShape()[1] == 0); + + auto loc = makeTensorPtrOp->getLoc(); + auto resultShape = + cast(makeTensorPtrOp.getType()).getShape(); + + assert(resultShape.size() == 2); + auto rank = cast(makeTensorPtrOp.getType()).getRank(); + auto targetOffset = ofrToIndexValue( + accumulateTargetOffset(makeTensorPtrOp, rewriter), loc, rewriter); + + //////////////////////////////////////////////////////////////////////////// + // + // Handling stacked wraparound + // See side-by-side wraparound for details. + // + //////////////////////////////////////////////////////////////////////////// + // We're loading a tensor of dim (rowSize, colSize) + // d1 + d2 = rowSize + // d2 is the number of rows that overflow + // + // cols + // + // wrappedAroundOff + // --------------*------------*-------- + // | d2 | | | + // | |------------| | + // rows| | + // | | + // | targetOffset | + // | *------------| | + // | | | | + // | d1 | | | + // | | clampedOff | | + // --------------*--------------------- + // | overflow | + // *------------- + // nextOff + // + // wrappedAroundOff = targetOffset % cols + // clampedOff = (rows * strideRows) + wrappedAroundOff + // ~~~~~~~~~~~~~~~~~ + // ^ + // | + // We have already computed + // rows * strideRows = modRow = shape[1] + // in TritonToStructured + // + // clampedOff - targetOffset + // d1 = -------------------- + // strideRows + // + // N = stride[0] + // M = shape[0] % N + // row_start = targetOffset / N % M + // remainSize = rowsize + // size = M - row_start + // start = row_start * N + targetOffset % N + // while (remainSize > size): + // reinterpret (start, size, stride ) + // remainSize = remainSize - size + // start = scaled_offset_1 + // size = M + // reinterpret (start, remainSize, stride ) + //////////////////////////////////////////////////////////////////////////// + // cols + // + // wrappedAroundOff + // --------------*------------*-------- + // | | + // | targetOffset | + // | *------------| | + // | | | | + // | | | | + // rows| rowSize | | | + // | | | | + // | | | | + // | *------------| | + // | nextOff | + // | | + // | clampedOff | + // --------------*--------------------- + // + // d1 = rowSize + // + // d2 = 0 + Value modM = makeTensorPtrOp.getShape()[0]; + Value N = + ofrToIndexValue(makeTensorPtrOp.getMixedStrides()[0], loc, rewriter); + + auto sizesInt = makeTensorPtrOp.getSizes(); + SmallVector sizesValues(rank); + std::transform(sizesInt.begin(), sizesInt.end(), sizesValues.begin(), + [&](auto val) { + return rewriter.create( + loc, rewriter.getIndexAttr(val)); + }); + SmallVector strideVals = + ofrsToIndexValues(makeTensorPtrOp.getMixedStrides(), loc, rewriter); + + // NOTE: Here, we need to use `targetOffset` for integer division, instead of + // `scaledOffset[1]` as in `sidebyside`. This is because the offset of the + // column loop analyzed by ptrAnalysis is added to `scaledOffset[0]` (row), + // and we assume that the offset in the 1-dimensional dimension must be less + // than N. + Value M = rewriter.create(loc, modM, N); + Value rowStart = rewriter.create(loc, targetOffset, N); + rowStart = rewriter.create(loc, rowStart, M); + Value colStart = rewriter.create(loc, targetOffset, N); + + Value remainSize = rewriter.create( + loc, rewriter.getIndexAttr(makeTensorPtrOp.getSizes()[0])); + Value size = rewriter.create(loc, M, rowStart); + Value start = rewriter.create(loc, rowStart, N); + start = rewriter.create(loc, start, colStart); + Value allocOffset = rewriter.create(loc, 0); + SmallVector typeR{remainSize.getType(), size.getType(), start.getType(), + allocOffset.getType()}; + SmallVector valueR{remainSize, size, start, allocOffset}; + if (!mixedDims.empty()) { + Value mixedDim = ofrsToIndexValues(mixedDims[0], loc, rewriter)[0]; + typeR.push_back(mixedDim.getType()); + valueR.push_back(mixedDim); + } + auto whileOp = rewriter.create( + loc, typeR, valueR, + /*beforeBuilder=*/ + [&](OpBuilder &b, Location loc, ValueRange args) { + Value cond = b.create(loc, arith::CmpIPredicate::sgt, + args[0], args[1]); + b.create(loc, cond, args); + }, + /*afterBuilder=*/ + [&](OpBuilder &b, Location loc, ValueRange args) { + Value remainSize = args[0]; + Value size = args[1]; + Value start = args[2]; + Value allocOffset = args[3]; + sizesValues[0] = size; + if (!mixedDims.empty()) { + Value mixedDim = args[4]; + mixedDims[0] = mixedDim; + } + Value block = createCastOps(makeTensorPtrOp, rewriter, start, + sizesValues, strideVals); + Value mixedDim = processMemSubviewCopies( + loc, rewriter, alloc, block, mixedDims, allocOffset, + 0 /*dim of mod row*/, isLoadToDst); + remainSize = b.create(loc, remainSize, size); + Value addOffsets = b.create(loc, size, N); + start = b.create(loc, start, addOffsets); + start = b.create(loc, start, modM); + size = M; + SmallVector newArgs{remainSize, size, start, allocOffset}; + if (!mixedDims.empty()) { + newArgs.push_back(mixedDim); + } + b.create(loc, newArgs); + }); + remainSize = whileOp->getResult(0); + start = whileOp->getResult(2); + sizesValues[0] = remainSize; + allocOffset = whileOp->getResult(3); + if (!mixedDims.empty()) { + mixedDims[0] = whileOp->getResult(4); + } + Value block = + createCastOps(makeTensorPtrOp, rewriter, start, sizesValues, strideVals); + processMemSubviewCopies(loc, rewriter, alloc, block, mixedDims, allocOffset, + 0 /*dim of mod row*/, isLoadToDst); +} + +void rewriteMakeTensorPtrAndMemAccess(tts::MakeTensorPtrOp makeTensorPtrOp, + ConversionPatternRewriter &rewriter, + Value alloc, + SmallVector mixedDims, + bool isLoadToDst) { + auto parentShape = makeTensorPtrOp.getStaticShape(); + if (parentShape.size() > 1 && parentShape[0] == ShapedType::kDynamic) { + rewriteStackedMemAccess(makeTensorPtrOp, rewriter, alloc, mixedDims, + isLoadToDst); + } else { + rewriteSideBySideMemAccess(makeTensorPtrOp, rewriter, alloc, mixedDims, + isLoadToDst); + } +} + +tts::MakeTensorPtrOp isSplitMemoryAccess(Operation *op) { + auto makeTensorPtrOp = + op->getOperand(0).getDefiningOp(); + if (makeTensorPtrOp && makeTensorPtrOp.isSplitPtr()) { + return makeTensorPtrOp; + } + return nullptr; +} + +struct LoadConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + rewriteStructuredLoad(tts::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + assert(!op.hasMask()); + + auto loc = op->getLoc(); + auto ptr = adaptor.getPtr(); + auto other = op.getOther(); + + auto tensorType = cast(op.getType()); + auto elemType = tensorType.getElementType(); + + Value alloc; + MemRefType memrefType = MemRefType::get(tensorType.getShape(), elemType); + + // No mask + assert(!other && "other value used in non-masked load"); + + if (auto makeTensorPtrOp = isSplitMemoryAccess(op)) { + alloc = rewriter.create(loc, memrefType); + rewriteMakeTensorPtrAndMemAccess(makeTensorPtrOp, rewriter, alloc, + SmallVector{}, true); + } else { + alloc = rewriter.create(loc, memrefType, ptr); + } + + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + + return success(); + } + + LogicalResult rewriteMaskedLoad(tts::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + assert(op.hasMask()); + + auto loc = op->getLoc(); + auto ptr = adaptor.getPtr(); + + auto tensorType = cast(op.getType()); + auto elemType = tensorType.getElementType(); + + auto alloc = rewriter.create( + loc, MemRefType::get(tensorType.getShape(), elemType)); + + SmallVector mixedDims = op.getMixedMaskDims(); + + // Fill load destination with other value + auto other = op.getOther(); + if (!other) { + + LLVM_DEBUG(op->emitRemark( + "Masked load without other value, using zero padding instead\n")); + // FIXME: Different reduction op need different reduce base value + other = rewriter.create( + loc, elemType, rewriter.getZeroAttr(elemType)); + } + + // For each dimension check if dims[i] < shape[i], or-accumulate + // the result + auto shape = tensorType.getShape(); + auto accBase = + rewriter.create(loc, rewriter.getBoolAttr(false)) + .getResult(); + for (size_t i = 0; i < shape.size(); i++) { + auto shapei = rewriter.create( + loc, rewriter.getIndexAttr(shape[i])); + + Value dimi = dyn_cast(mixedDims[i]); + if (!dimi) { + dimi = rewriter.create( + loc, rewriter.getIndexAttr(op.getStaticMaskDims()[i])); + } + + Value cmp = rewriter.create(loc, arith::CmpIPredicate::slt, + dimi, shapei); + accBase = rewriter.create(loc, accBase, cmp); + } + + // condition the memset on the or-accumulation + // initialize with padding prior to CopyOp + rewriter.create(loc, accBase, [&](OpBuilder &b, Location loc) { + b.create(loc, ValueRange{other}, ValueRange{alloc}); + b.create(loc); + }); + + if (auto makeTensorPtrOp = isSplitMemoryAccess(op)) { + rewriteMakeTensorPtrAndMemAccess(makeTensorPtrOp, rewriter, alloc, + mixedDims, true); + } else { + memref::SubViewOp srcSubview = + getSubview(tensorType.getRank(), mixedDims, ptr, loc, rewriter); + memref::SubViewOp dstSubview = + getSubview(tensorType.getRank(), mixedDims, alloc, loc, rewriter); + rewriter.create(loc, srcSubview, dstSubview); + } + + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + + return success(); + } + +public: + LogicalResult + matchAndRewrite(tts::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (op.hasMask()) { + return rewriteMaskedLoad(op, adaptor, rewriter); + } else { + return rewriteStructuredLoad(op, adaptor, rewriter); + } + } +}; + +struct StoreConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + static tensor::ExtractSliceOp + getExtractSlice(int rank, ArrayRef dims, Value source, + const Location loc, OpBuilder &b) { + auto sourceType = cast(source.getType()); + SmallVector offsets(rank, b.getIndexAttr(0)); + SmallVector strides(rank, b.getIndexAttr(1)); + + auto dstType = tensor::ExtractSliceOp::inferResultType(sourceType, offsets, + dims, strides); + + return b.create(loc, dstType, source, offsets, dims, + strides); + } + + LogicalResult rewriteMaskedStore(tts::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + assert(op.hasMask()); + + auto loc = op.getLoc(); + auto ptr = adaptor.getPtr(); + auto storeValue = op.getValue(); + auto tensorType = cast(storeValue.getType()); + auto elemType = tensorType.getElementType(); + auto rank = cast(storeValue.getType()).getRank(); + + auto mixedDims = op.getMixedMaskDims(); + if (auto makeTensorPtrOp = isSplitMemoryAccess(op)) { + auto srcSlice = + getExtractSlice(rank, mixedDims, storeValue, loc, rewriter); + auto srcType = cast(srcSlice.getType()); + auto srcSliceMemRef = rewriter.create( + loc, MemRefType::get(srcType.getShape(), srcType.getElementType()), + srcSlice); + rewriteMakeTensorPtrAndMemAccess(makeTensorPtrOp, rewriter, + srcSliceMemRef, mixedDims, false); + } else { + auto srcSlice = + getExtractSlice(rank, mixedDims, storeValue, loc, rewriter); + auto dstSubview = getSubview(rank, mixedDims, ptr, loc, rewriter); + + auto storeOp = rewriter.create( + loc, srcSlice, dstSubview); + storeOp.setWritable(true); + } + rewriter.eraseOp(op); + return success(); + } + + LogicalResult + rewriteStructuredStore(tts::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + assert(!op.hasMask()); + + auto loc = op.getLoc(); + auto ptr = adaptor.getPtr(); + auto storeValue = op.getValue(); + + auto tensorType = cast(storeValue.getType()); + auto elemType = tensorType.getElementType(); + auto rank = cast(storeValue.getType()).getRank(); + MemRefType memrefType = MemRefType::get(tensorType.getShape(), elemType); + + if (auto makeTensorPtrOp = isSplitMemoryAccess(op)) { + auto srcMemRef = rewriter.create( + loc, memrefType, storeValue); + rewriteMakeTensorPtrAndMemAccess(makeTensorPtrOp, rewriter, srcMemRef, + SmallVector{}, false); + } else { + auto storeOp = rewriter.create( + loc, storeValue, ptr); + storeOp.setWritable(true); + } + + rewriter.eraseOp(op); + return success(); + } + +public: + LogicalResult + matchAndRewrite(tts::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (op.hasMask()) { + return rewriteMaskedStore(op, adaptor, rewriter); + } else { + return rewriteStructuredStore(op, adaptor, rewriter); + } + } +}; + +struct AtomicRMWOpConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + static tensor::ExtractSliceOp + getExtractSlice(int rank, ArrayRef dims, Value source, + const Location loc, OpBuilder &b) { + auto sourceType = cast(source.getType()); + SmallVector offsets(rank, b.getIndexAttr(0)); + SmallVector strides(rank, b.getIndexAttr(1)); + + auto dstType = tensor::ExtractSliceOp::inferResultType(sourceType, offsets, + dims, strides); + + return b.create(loc, dstType, source, offsets, dims, + strides); + } + +public: + LogicalResult + matchAndRewrite(tts::AtomicRMWOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto ptr = adaptor.getPtr(); + auto value = adaptor.getValue(); + + auto type = cast(value.getType()); + auto rank = type.getRank(); + + Value init = rewriter.create(loc, type.getShape(), + type.getElementType()); + + if (op.hasMask()) { + auto mixedDims = op.getMixedMaskDims(); + + auto valueSlice = getExtractSlice(rank, mixedDims, value, loc, rewriter); + auto ptrSubview = getSubview(rank, mixedDims, ptr, loc, rewriter); + + auto atomicRMWOp = rewriter.create( + loc, op.getType(), ptrSubview, valueSlice, init, + op.getAtomicRmwOpAttr(), op.getSemAttr(), op.getScopeAttr()); + rewriter.replaceOp(op, atomicRMWOp); + } else { + auto atomicRMWOp = rewriter.create( + loc, op.getType(), ptr, value, init, op.getAtomicRmwOpAttr(), + op.getSemAttr(), op.getScopeAttr()); + rewriter.replaceOp(op, atomicRMWOp); + } + return success(); + } +}; + +struct AtomicCASOpConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(tts::AtomicCASOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (op.getOffset()) + return failure(); + + auto loc = op.getLoc(); + auto ptr = adaptor.getPtr(); + auto cmp = adaptor.getCmp(); + auto value = adaptor.getValue(); + + auto type = cast(value.getType()); + + Value init = rewriter.create(loc, type.getShape(), + type.getElementType()); + + auto atomicCASOp = rewriter.create( + loc, op.getType(), ptr, cmp, value, init, op.getSemAttr(), + op.getScopeAttr()); + rewriter.replaceOp(op, atomicCASOp); + + return success(); + } +}; + +} // namespace + +void mlir::triton::populateStructuredToMemrefConversionPatterns( + RewritePatternSet &patterns, TypeConverter &typeConverter) { + patterns.add(typeConverter, patterns.getContext()); + patterns.add(patterns.getContext()); +} diff --git a/third_party/wafer/lib/Conversion/StructuredToMemref/StructuredToMemrefPass.cpp b/third_party/wafer/lib/Conversion/StructuredToMemref/StructuredToMemrefPass.cpp new file mode 100755 index 00000000..b3173b43 --- /dev/null +++ b/third_party/wafer/lib/Conversion/StructuredToMemref/StructuredToMemrefPass.cpp @@ -0,0 +1,153 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "Address/Dialect/IR/AddressDialect.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/Support/LogicalResult.h" +#include "triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "utils/TypeConvertor.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/Pass/PassManager.h" +#include "triton/Dialect/Triton/IR/Types.h" +#include "llvm/Support/Casting.h" + +#include + +#define DEBUG_TYPE "structured-to-memref" + +using namespace mlir; +using namespace triton; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_STRUCTUREDTOMEMREF +#include "triton-shared/Conversion/StructuredToMemref/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class LoopTypeConverter : public TypeConverter { +public: + LoopTypeConverter(MLIRContext *context) { + // The order of type conversion is important: later ones are tried earlier. + addConversion([](Type type) { return type; }); + + // A tensor of pointers can be passed in as scf.for's init-args, in such + // cases, we convert the type to a memref with dynamic offsets and + // strides. + addConversion( + [context](RankedTensorType tensorType) -> std::optional { + if (auto ptrType = llvm::dyn_cast( + tensorType.getElementType())) { + auto layout = StridedLayoutAttr::get( + context, ShapedType::kDynamic, + SmallVector(tensorType.getRank(), + ShapedType::kDynamic)); + Type elemType = ptrType.getPointeeType(); + return MemRefType::get(tensorType.getShape(), elemType, layout); + } + + return std::nullopt; + }); + + addSourceMaterialization([&](OpBuilder &builder, Type resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + + addTargetMaterialization([&](OpBuilder &builder, Type resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + + // Convert the current memref type to a memref type with dynamic offsets and + // strides through another reinterpret_cast with the same offsets. + // Canonicalization will simplify this sequence by removing the inital + // reinterpret_cast. + addTargetMaterialization([&](OpBuilder &builder, MemRefType memrefType, + ValueRange inputs, Location loc) -> Value { + auto reinterpretCast = + inputs[0].getDefiningOp(); + if (!reinterpretCast) { + return builder + .create(loc, memrefType, inputs) + .getResult(0); + } + return builder.create( + loc, memrefType, inputs[0], reinterpretCast.getMixedOffsets()[0], + reinterpretCast.getMixedSizes(), reinterpretCast.getMixedStrides()); + }); + } +}; + +class StructuredToMemrefPass + : public triton::impl::StructuredToMemrefBase { + using StructuredToMemrefBase::StructuredToMemrefBase; + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, ttx::TritonTilingExtDialect, + memref::MemRefDialect, mk::MagicKernelDialect>(); + + target.addIllegalOp(); + + target.addLegalOp(); + + LoopTypeConverter loopTypeConverter(patterns.getContext()); + + mlir::scf::populateSCFStructuralTypeConversionsAndLegality( + loopTypeConverter, patterns, target); + + PtrToUnrankedMemrefConverter typeConverter; + triton::populateStructuredToMemrefConversionPatterns(patterns, + typeConverter); + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> +triton::createStructuredToMemrefPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/TLEToMK/CMakeLists.txt b/third_party/wafer/lib/Conversion/TLEToMK/CMakeLists.txt new file mode 100755 index 00000000..23187c92 --- /dev/null +++ b/third_party/wafer/lib/Conversion/TLEToMK/CMakeLists.txt @@ -0,0 +1,21 @@ +add_triton_library(TLEToMagicKernel + TLEToMK.cpp + TLEToMKPass.cpp + + DEPENDS + MagicKernelTableGen + TLEToMKConversionPassIncGen + TleDsaDialectIncGen + TleDsaTypesIncGen + TleDsaOpsIncGen + + LINK_LIBS PUBLIC + MLIRIR + MLIRPass + MLIRTransforms + MLIRSupport + MLIRArithDialect + TritonIR + TritonStructuredIR + TleDsaIR +) diff --git a/third_party/wafer/lib/Conversion/TLEToMK/MKCommonBufferPlanningPass.cpp b/third_party/wafer/lib/Conversion/TLEToMK/MKCommonBufferPlanningPass.cpp new file mode 100755 index 00000000..76916a29 --- /dev/null +++ b/third_party/wafer/lib/Conversion/TLEToMK/MKCommonBufferPlanningPass.cpp @@ -0,0 +1,183 @@ +//===---------------- MKCommBufferPlanningPass.cpp ---------------------===// +// +// Insert local SPM buffers for paired mk.recv/mk.send that share the same +// placeholder base address. The placeholder base is represented by the i64 +// src_addr/dst_addr operands. +// +// Strategy: +// - Find pairs where recv.src_addr and send.dst_addr are the same SSA value. +// - Replace the shared placeholder "addr" with two distinct buffers: +// one buffer for send's remote dst, one buffer for recv's remote src. +// - The buffers are created as tensor.empty (or memref.alloc if already +// bufferized). +// +//===--------------------------------------------------------------------===// + +#include "magic-kernel/Conversion/TLEToMK/TLEToMK.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/Dominance.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Pass/Pass.h" + +#include + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "magic-kernel/Conversion/TLEToMK/Passes.h.inc" + +namespace { + +static int64_t alignUp(int64_t v, int64_t a) { return (v + a - 1) / a * a; } + +static std::optional getStaticBytes(ShapedType ty) { + if (!ty || !ty.hasStaticShape()) + return std::nullopt; + auto elemTy = ty.getElementType(); + if (!elemTy.isIntOrFloat()) + return std::nullopt; + int64_t elemBytes = elemTy.getIntOrFloatBitWidth() / 8; + if (elemBytes <= 0) + return std::nullopt; + return ty.getNumElements() * elemBytes; +} + +/// Return a canonical "placeholder root" for an addr-like value. +/// This lets us match send/recv pairs even if the i64 addr was computed by +/// distinct ptr_to_int/extract ops. +static Value getPlaceholderRoot(Value addrLike) { + Value v = addrLike; + // Peel trivial casts (best-effort). + if (auto cast = v.getDefiningOp()) { + if (!cast.getOperands().empty()) + v = cast.getOperands().front(); + } + + // If it's an integer address derived from a triton ptr, use the ptr source. + if (v.getType().isInteger(64)) { + if (auto p2i = v.getDefiningOp()) + v = p2i.getSrc(); + } + + // If it comes from extracting element [0,0,...] from a tensor of ptrs, use + // the tensor-of-ptrs as the root. + if (auto ex = v.getDefiningOp()) { + // Only treat it as placeholder root if indices are all constants. + // (We expect [0,0,...] here.) + v = ex.getTensor(); + } + + return v; +} + +static Value createEmptyLikeShaped(OpBuilder &b, Location loc, ShapedType ty) { + if (auto t = dyn_cast(ty)) { + return b.create(loc, t.getShape(), t.getElementType()); + } + if (auto m = dyn_cast(ty)) { + return b.create(loc, m); + } + return Value(); +} + +struct MKCommBufferPlanningPass + : public MKCommBufferPlanningBase { + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + ModuleOp module = getOperation(); + + module.walk([&](triton::FuncOp func) { + // Map base -> first send/recv. + DenseMap rootToSend; + DenseMap rootToRecv; + + func.walk([&](Operation *op) { + if (auto send = dyn_cast(op)) { + Value root = getPlaceholderRoot(send.getDstAddr()); + rootToSend.try_emplace(root, send); + } else if (auto recv = dyn_cast(op)) { + Value root = getPlaceholderRoot(recv.getDst()); + rootToRecv.try_emplace(root, recv); + } + }); + + OpBuilder b(func.getContext()); + for (auto &it : rootToRecv) { + Value root = it.first; + auto recv = it.second; + auto sendIt = rootToSend.find(root); + if (sendIt == rootToSend.end()) + continue; + auto send = sendIt->second; + + Location loc = recv.getLoc(); + + // Create a single shared buffer as close as possible while still + // dominating both send and recv. + DominanceInfo dom(func); + Block *sendBlock = send->getBlock(); + Block *recvBlock = recv->getBlock(); + Block *insBlock = dom.findNearestCommonDominator(sendBlock, recvBlock); + if (!insBlock) + insBlock = &func.getBody().front(); + + if (insBlock == sendBlock && insBlock == recvBlock) { + // Same block: insert before the earlier op. + Operation *insPt = send->isBeforeInBlock(recv) ? send.getOperation() + : recv.getOperation(); + b.setInsertionPoint(insPt); + } else { + // Different blocks: insert at end of common dominator block + // (before terminator if any). + b.setInsertionPointToEnd(insBlock); + if (!insBlock->empty() && + insBlock->back().hasTrait()) + b.setInsertionPoint(&insBlock->back()); + } + + auto sendSrcTy = dyn_cast(send.getSrc().getType()); + auto recvTy = recv.getNumResults() > 0 + ? dyn_cast(recv->getResult(0).getType()) + : dyn_cast(recv.getDst().getType()); + if (!sendSrcTy || !recvTy) + continue; + + // For scheme C we expect send/recv to communicate same shape/type. + if (sendSrcTy.getElementType() != recvTy.getElementType() || + sendSrcTy.getShape() != recvTy.getShape()) + continue; + + Value sharedBuf = createEmptyLikeShaped(b, loc, sendSrcTy); + if (!sharedBuf) + continue; + + // Replace the shared placeholder root with the shared buffer. + // Operand layout: + // mk.send: 4 coords + dst_addr + src + // mk.recv: 4 coords + dst_key + dst + send->setOperand(4, sharedBuf); + recv->setOperand(4, sharedBuf); // dst_key + } + }); + } +}; + +} // namespace + +std::unique_ptr triton::createMKCommBufferPlanningPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/TLEToMK/TLEToMK.cpp b/third_party/wafer/lib/Conversion/TLEToMK/TLEToMK.cpp new file mode 100755 index 00000000..1cba4388 --- /dev/null +++ b/third_party/wafer/lib/Conversion/TLEToMK/TLEToMK.cpp @@ -0,0 +1,717 @@ +//===------------------- TLEToMK.cpp -----------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Conversion/TLEToMK/TLEToMK.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "tle/include/tle-dsa/Dialect/IR/DsaDialect.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/IR/TypeUtilities.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "llvm/ADT/STLExtras.h" + +#define DEBUG_TYPE "tle-to-mk" + +using namespace mlir; +using namespace triton; +using namespace mk; +using namespace tts; + +namespace { + +static constexpr llvm::StringLiteral kRemoteShardCarrierAttr = + "tle.remote_shard_id_carrier"; + +static bool isConstantZeroIndex(Value v) { + if (auto cst = v.getDefiningOp()) + return cst.value() == 0; + if (auto cst = v.getDefiningOp()) { + if (!isa(cst.getType())) + return false; + if (auto intAttr = dyn_cast(cst.getValue())) + return intAttr.getValue().isZero(); + } + return false; +} + +static bool areAllZeroIndices(ValueRange indices) { + return llvm::all_of(indices, isConstantZeroIndex); +} + +static bool isBeforeOrAtInSameBlock(Operation *a, Operation *b) { + return a && b && a->getBlock() == b->getBlock() && + (a == b || a->isBeforeInBlock(b)); +} + +static Value getOrCreateScalarPtr(PatternRewriter &rewriter, Location loc, + Value ptrLike, Operation *useAnchor) { + if (!isa(ptrLike.getType())) + return ptrLike; + + for (Operation *user : ptrLike.getUsers()) { + auto ex = dyn_cast(user); + if (!ex) + continue; + if (ex.getTensor() != ptrLike) + continue; + if (!areAllZeroIndices(ex.getIndices())) + continue; + if (!useAnchor || isBeforeOrAtInSameBlock(ex.getOperation(), useAnchor)) + return ex.getResult(); + } + + auto ranked = cast(ptrLike.getType()); + SmallVector idxs; + idxs.reserve(ranked.getRank()); + for (int i = 0; i < ranked.getRank(); ++i) + idxs.push_back(rewriter.create(loc, 0)); + return rewriter.create(loc, ptrLike, idxs); +} + +static Value getOrCreatePtrToIntI64(PatternRewriter &rewriter, Location loc, + Value scalarPtr, Operation *useAnchor) { + for (Operation *user : scalarPtr.getUsers()) { + auto p2i = dyn_cast(user); + if (!p2i) + continue; + if (p2i.getSrc() != scalarPtr) + continue; + if (p2i.getType() != rewriter.getI64Type()) + continue; + if (!useAnchor || isBeforeOrAtInSameBlock(p2i.getOperation(), useAnchor)) + return p2i.getResult(); + } + + return rewriter.create(loc, rewriter.getI64Type(), + scalarPtr); +} + +/// Extract a flat i64 base-address from a pointer-like value. +/// +/// When \p ptrLike is the result of a \c dsa.local_pointers op we go straight +/// to the underlying memref, avoiding the creation of any \c !tt.ptr typed +/// intermediate values. +static Value getOrCreatePtrLikeAddrI64(PatternRewriter &rewriter, Location loc, + Value ptrLike, Operation *useAnchor) { + // --- Fast path: dsa.local_pointers → extract base from memref directly --- + if (auto localPtrOp = ptrLike.getDefiningOp()) { + OpBuilder::InsertionGuard g(rewriter); + // Insert right after the local_pointers op so that the new ops dominate + // all users. + if (localPtrOp->getNextNode()) + rewriter.setInsertionPoint(localPtrOp->getNextNode()); + else + rewriter.setInsertionPointAfter(localPtrOp); + auto idxTy = rewriter.getIndexType(); + auto i64Ty = rewriter.getI64Type(); + Value baseIndex = rewriter.create( + loc, idxTy, localPtrOp.getSrc()); + return rewriter.create(loc, i64Ty, baseIndex); + } + + // --- Original path: Triton pointer value --- + OpBuilder::InsertionGuard g(rewriter); + Block *block = rewriter.getInsertionBlock(); + if (auto def = ptrLike.getDefiningOp()) { + if (block && def->getBlock() == block) + rewriter.setInsertionPointAfter(def); + } else if (block) { + rewriter.setInsertionPointToStart(block); + } + + Value scalarPtr = getOrCreateScalarPtr(rewriter, loc, ptrLike, useAnchor); + return getOrCreatePtrToIntI64(rewriter, loc, scalarPtr, useAnchor); +} + +static Value castIntegerLikeToI64(PatternRewriter &rewriter, Location loc, + Value v) { + auto i64Ty = rewriter.getI64Type(); + Type ty = v.getType(); + if (ty == i64Ty) + return v; + if (isa(ty)) + return rewriter.create(loc, i64Ty, v); + if (auto intTy = dyn_cast(ty)) { + if (intTy.getWidth() < 64) + return rewriter.create(loc, i64Ty, v); + if (intTy.getWidth() > 64) + return rewriter.create(loc, i64Ty, v); + return v; + } + return Value(); +} + +static Value peelShardScalar(Value shardLike) { + if (auto splat = shardLike.getDefiningOp()) + return splat.getSrc(); + return shardLike; +} + +static LogicalResult getCoordsFromShardIdValue(PatternRewriter &rewriter, + Location loc, Value shardIdLike, + SmallVector &coords) { + Value shardId = peelShardScalar(shardIdLike); + Value tileId = castIntegerLikeToI64(rewriter, loc, shardId); + if (!tileId) + return failure(); + Value four = + rewriter.create(loc, rewriter.getI64IntegerAttr(4)); + Value zero = + rewriter.create(loc, rewriter.getI64IntegerAttr(0)); + Value chipX = rewriter.create(loc, tileId, four); + Value chipY = rewriter.create(loc, tileId, four); + coords = {chipX, chipY, zero, tileId}; + return success(); +} + +static LogicalResult extractRemoteInfoFromPtr(PatternRewriter &rewriter, + Location loc, Value ptrLike, + SmallVector &coords, + Value &basePtrLike) { + if (auto remotePtrOp = ptrLike.getDefiningOp()) { + if (failed(getCoordsFromShardIdValue(rewriter, loc, + remotePtrOp.getShardId(), coords))) + return failure(); + basePtrLike = remotePtrOp.getSrc(); + return success(); + } + if (auto addPtr = ptrLike.getDefiningOp(); + addPtr && addPtr->hasAttr(kRemoteShardCarrierAttr)) { + if (failed(getCoordsFromShardIdValue(rewriter, loc, addPtr.getOffset(), + coords))) + return failure(); + basePtrLike = addPtr.getPtr(); + return success(); + } + return failure(); +} + +// ===----------------------------------------------------------------------===// +// Barrier +// ===----------------------------------------------------------------------===// + +struct DsaDistributedBarrierToMkPattern + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(mlir::dsa::DistributedBarrierOp op, + PatternRewriter &rewriter) const override { + rewriter.create(op.getLoc()); + rewriter.eraseOp(op); + return success(); + } +}; + +// ===----------------------------------------------------------------------===// +// Remote load / store (dsa.remote_pointers → mk.remote_load/store) +// ===----------------------------------------------------------------------===// + +struct DsaRemoteLoadToMkPattern : public OpRewritePattern { + explicit DsaRemoteLoadToMkPattern(MLIRContext *ctx) + : OpRewritePattern(ctx, /*benefit=*/2) {} + + LogicalResult matchAndRewrite(triton::LoadOp loadOp, + PatternRewriter &rewriter) const override { + Location loc = loadOp.getLoc(); + SmallVector recvCoords; + Value basePtrLike = loadOp.getPtr(); + if (failed(extractRemoteInfoFromPtr(rewriter, loc, loadOp.getPtr(), + recvCoords, basePtrLike))) + return failure(); + + auto resultType = dyn_cast(loadOp.getResult().getType()); + if (!resultType) + return loadOp->emitRemark( + "remote load currently expects ranked tensor result"); + for (int64_t s : resultType.getShape()) { + if (ShapedType::isDynamic(s)) + return loadOp->emitRemark( + "remote load with dynamic shape not supported"); + } + + Value dstBuffer = rewriter.create( + loc, resultType.getShape(), resultType.getElementType()); + auto recvOp = rewriter.create( + loc, resultType, recvCoords[0], recvCoords[1], recvCoords[2], + recvCoords[3], dstBuffer); + rewriter.replaceOp(loadOp, recvOp.getResults().front()); + return success(); + } +}; + +struct DsaRemoteStoreToMkPattern : public OpRewritePattern { + explicit DsaRemoteStoreToMkPattern(MLIRContext *ctx) + : OpRewritePattern(ctx, /*benefit=*/2) {} + + LogicalResult matchAndRewrite(triton::StoreOp storeOp, + PatternRewriter &rewriter) const override { + Location loc = storeOp.getLoc(); + SmallVector sendCoords; + Value basePtrLike = storeOp.getPtr(); + if (failed(extractRemoteInfoFromPtr(rewriter, loc, storeOp.getPtr(), + sendCoords, basePtrLike))) + return failure(); + + if (storeOp.getMask()) + return storeOp->emitRemark("masked remote store not supported"); + + Value dstAddrI64 = getOrCreatePtrLikeAddrI64(rewriter, loc, basePtrLike, + storeOp.getOperation()); + rewriter.create(loc, sendCoords[0], sendCoords[1], + sendCoords[2], sendCoords[3], dstAddrI64, + storeOp.getValue()); + rewriter.eraseOp(storeOp); + return success(); + } +}; + +// ===----------------------------------------------------------------------===// +// Local load / store (dsa.local_pointers + tt.load/store → memref ops) +// +// Instead of lowering dsa.local_pointers to Triton pointer arithmetic +// (tt.splat/tt.addptr with tensor>), we directly convert the +// load/store users to memref-level operations. This avoids producing +// !tt.ptr element types that downstream triton-to-core-dialects cannot +// convert to valid memref types. +// ===----------------------------------------------------------------------===// + +/// tt.load whose pointer comes from dsa.local_pointers → +/// bufferization.to_tensor of the underlying memref. +struct DsaLocalLoadToMemrefPattern : public OpRewritePattern { + explicit DsaLocalLoadToMemrefPattern(MLIRContext *ctx) + : OpRewritePattern(ctx, /*benefit=*/3) {} + + LogicalResult matchAndRewrite(triton::LoadOp loadOp, + PatternRewriter &rewriter) const override { + // Only match loads whose pointer is produced by dsa.local_pointers. + auto localPtrOp = + loadOp.getPtr().getDefiningOp(); + if (!localPtrOp) + return failure(); + + auto memrefTy = dyn_cast(localPtrOp.getSrc().getType()); + if (!memrefTy) + return failure(); + + auto resultTy = dyn_cast(loadOp.getResult().getType()); + if (!resultTy) + return failure(); + + // Build a tensor type from the memref shape + element type. + auto tensorTy = + RankedTensorType::get(memrefTy.getShape(), memrefTy.getElementType()); + + // Shapes must agree (the common DSA pattern uses identity indices). + if (tensorTy.getShape() != resultTy.getShape()) + return loadOp->emitRemark( + "local load shape mismatch between memref and result tensor"); + + // Element type may differ if an implicit cast is present (e.g. f32→f16). + // For now we require them to match. + if (memrefTy.getElementType() != resultTy.getElementType()) + return loadOp->emitRemark( + "local load element type mismatch between memref and result tensor"); + + // Replace with: bufferization.to_tensor %memref + // writable=true because the SPM buffer is mutable. + auto toTensor = rewriter.create( + loadOp.getLoc(), resultTy, localPtrOp.getSrc(), + /*restrict=*/true, /*writable=*/true); + rewriter.replaceOp(loadOp, toTensor.getResult()); + return success(); + } +}; + +/// tt.store whose pointer comes from dsa.local_pointers → +/// bufferization.to_memref + memref.copy into the underlying SPM buffer. +struct DsaLocalStoreToMemrefPattern : public OpRewritePattern { + explicit DsaLocalStoreToMemrefPattern(MLIRContext *ctx) + : OpRewritePattern(ctx, /*benefit=*/3) {} + + LogicalResult matchAndRewrite(triton::StoreOp storeOp, + PatternRewriter &rewriter) const override { + auto localPtrOp = + storeOp.getPtr().getDefiningOp(); + if (!localPtrOp) + return failure(); + + auto destMemrefTy = dyn_cast(localPtrOp.getSrc().getType()); + if (!destMemrefTy) + return failure(); + + Value val = storeOp.getValue(); + auto valTy = dyn_cast(val.getType()); + if (!valTy) + return failure(); + + // Shapes must match. + if (valTy.getShape() != destMemrefTy.getShape()) + return storeOp->emitRemark( + "local store shape mismatch between value tensor and SPM memref"); + + // Element types must match (no implicit cast support yet). + if (valTy.getElementType() != destMemrefTy.getElementType()) + return storeOp->emitRemark( + "local store element type mismatch between value and SPM memref"); + + Location loc = storeOp.getLoc(); + + // Materialise the tensor value as a memref, then copy into the SPM buffer. + // Use a contiguous memref type for the intermediate to_memref result. + auto srcMemrefTy = + MemRefType::get(valTy.getShape(), valTy.getElementType()); + auto srcMemref = + rewriter.create(loc, srcMemrefTy, val); + rewriter.create(loc, srcMemref, localPtrOp.getSrc()); + rewriter.eraseOp(storeOp); + return success(); + } +}; + +// ===----------------------------------------------------------------------===// +// Remote pointers fallback (kept for edge cases) +// ===----------------------------------------------------------------------===// + +struct DsaRemotePointersToTritonPattern + : public OpRewritePattern { + explicit DsaRemotePointersToTritonPattern(MLIRContext *ctx) + : OpRewritePattern(ctx, /*benefit=*/1) {} + + LogicalResult matchAndRewrite(mlir::dsa::RemotePointersOp op, + PatternRewriter &rewriter) const override { + Value offset = op.getShardId(); + if (auto srcTy = dyn_cast(op.getSrc().getType())) { + auto shardTy = dyn_cast(offset.getType()); + if (!shardTy || shardTy.getShape() != srcTy.getShape()) { + auto offsetTy = + RankedTensorType::get(srcTy.getShape(), offset.getType()); + offset = + rewriter.create(op.getLoc(), offsetTy, offset); + } + } + auto addPtr = rewriter.create(op.getLoc(), op.getType(), + op.getSrc(), offset); + addPtr->setAttr(kRemoteShardCarrierAttr, rewriter.getUnitAttr()); + rewriter.replaceOp(op, addPtr.getResult()); + return success(); + } +}; + +struct DsaRandGenToMkPattern : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(mlir::dsa::RandGenOp op, + PatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + auto seed0Ty = cast(op.getSeed0().getType()); + auto seed1Ty = cast(op.getSeed1().getType()); + auto outTy = cast(op.getOut().getType()); + auto seed0OutTy = cast(op.getSeed0Out().getType()); + auto seed1OutTy = cast(op.getSeed1Out().getType()); + if (seed0Ty.getShape() != ArrayRef({16}) || + seed1Ty.getShape() != ArrayRef({16}) || + seed0OutTy.getShape() != ArrayRef({16}) || + seed1OutTy.getShape() != ArrayRef({16})) + return rewriter.notifyMatchFailure( + op, "dsa.randgen seeds must have shape [16]"); + + int32_t byteCount = op.getByteCount(); + if (byteCount <= 0 || (byteCount % 128) != 0) + return rewriter.notifyMatchFailure( + op, "dsa.randgen byte_count must be a positive multiple of 128"); + int64_t expectedOutElems = static_cast(byteCount) / 8; + if (outTy.getNumElements() != expectedOutElems) + return rewriter.notifyMatchFailure( + op, "dsa.randgen out numel must equal byte_count / 8"); + + auto outInit = rewriter.create(loc, outTy.getShape(), + outTy.getElementType()); + auto seed0Init = rewriter.create( + loc, seed0OutTy.getShape(), seed0OutTy.getElementType()); + auto seed1Init = rewriter.create( + loc, seed1OutTy.getShape(), seed1OutTy.getElementType()); + + auto mkOp = rewriter.create( + loc, TypeRange{outTy, seed0OutTy, seed1OutTy}, op.getSeed0(), + op.getSeed1(), outInit, seed0Init, seed1Init, + rewriter.getI32IntegerAttr(byteCount), op.getFmtAttr()); + + rewriter.replaceOp(op, mkOp->getResults()); + return success(); + } +}; + +// ===----------------------------------------------------------------------===// +// dsa.bitcast → mk.bitcast (zero-cost SPM buffer alias) +// ===----------------------------------------------------------------------===// + +struct DsaBitcastToMkPattern : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(mlir::dsa::BitcastOp op, + PatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, op.getResult().getType(), + op.getSrc()); + return success(); + } +}; + +// ===----------------------------------------------------------------------===// +// Remote pointers fallback (kept for edge cases) +// ===----------------------------------------------------------------------===// + + +static LogicalResult buildSliceOffsets(PatternRewriter &rewriter, Location loc, + ArrayRef staticOffsets, + ValueRange dynOffsets, + SmallVectorImpl &offsets) { + unsigned dynIdx = 0; + for (int64_t s : staticOffsets) { + if (ShapedType::isDynamic(s)) { + if (dynIdx >= dynOffsets.size()) + return failure(); + Value v = dynOffsets[dynIdx++]; + if (isa(v.getType())) + v = rewriter.create(loc, v, ValueRange{}); + if (!v.getType().isIndex()) + v = rewriter.create(loc, rewriter.getIndexType(), + v); + offsets.push_back(v); + } else { + offsets.push_back(rewriter.getI64IntegerAttr(s)); + } + } + return success(dynIdx == dynOffsets.size()); +} + +static void buildSliceSizesStrides(PatternRewriter &rewriter, + ArrayRef dims, + SmallVectorImpl &result) { + for (int64_t d : dims) + result.push_back(rewriter.getI64IntegerAttr(d)); +} + +struct DsaExtractSliceToTensorSlicePattern + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(mlir::dsa::ExtractSliceOp op, + PatternRewriter &rewriter) const override { + auto srcTy = dyn_cast(op.getSrc().getType()); + if (!srcTy) + return failure(); + auto resultTy = dyn_cast(op.getResult().getType()); + if (!resultTy) + return failure(); + + SmallVector offsets; + if (failed(buildSliceOffsets(rewriter, op.getLoc(), op.getStaticOffsets(), + op.getOffsets(), offsets))) + return failure(); + SmallVector sizes, strides; + buildSliceSizesStrides(rewriter, op.getSizes(), sizes); + buildSliceSizesStrides(rewriter, op.getStrides(), strides); + + rewriter.replaceOpWithNewOp( + op, resultTy, op.getSrc(), offsets, sizes, strides); + return success(); + } +}; + +struct DsaInsertSliceToTensorSlicePattern + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(mlir::dsa::InsertSliceOp op, + PatternRewriter &rewriter) const override { + auto srcTy = dyn_cast(op.getSrc().getType()); + if (!srcTy) + return failure(); + + SmallVector offsets; + if (failed(buildSliceOffsets(rewriter, op.getLoc(), op.getStaticOffsets(), + op.getOffsets(), offsets))) + return failure(); + SmallVector sizes, strides; + buildSliceSizesStrides(rewriter, op.getSizes(), sizes); + buildSliceSizesStrides(rewriter, op.getStrides(), strides); + + rewriter.replaceOpWithNewOp( + op, op.getTile(), op.getSrc(), offsets, sizes, strides); + return success(); + } +}; + +// Three-operand elementwise arithmetic on memrefs; benefit=4 fires before +// DsaLocalLoadToMemrefPattern (benefit=3). MKToWafer maps the arith op to +// tx.*VV. +template +struct DsaBinaryOpToLinalgPattern : public OpRewritePattern { + explicit DsaBinaryOpToLinalgPattern(MLIRContext *ctx) + : OpRewritePattern(ctx, /*benefit=*/4) {} + + LogicalResult matchAndRewrite(DsaOpT op, + PatternRewriter &rewriter) const override { + auto lhsTy = dyn_cast(op.getLhs().getType()); + auto rhsTy = dyn_cast(op.getRhs().getType()); + auto outTy = dyn_cast(op.getOut().getType()); + if (!lhsTy || !rhsTy || !outTy) + return failure(); + + if (lhsTy.getShape() != rhsTy.getShape() || + lhsTy.getShape() != outTy.getShape()) + return op->emitRemark("dsa binary op shape mismatch between lhs/rhs/out"); + if (lhsTy.getElementType() != rhsTy.getElementType() || + lhsTy.getElementType() != outTy.getElementType()) + return op->emitRemark( + "dsa binary op element type mismatch between lhs/rhs/out"); + + Location loc = op.getLoc(); + auto elemTy = lhsTy.getElementType(); + auto rank = static_cast(lhsTy.getShape().size()); + auto identityMap = rewriter.getMultiDimIdentityMap(rank); + SmallVector indexingMaps = {identityMap, identityMap, + identityMap}; + SmallVector iteratorTypes( + rank, mlir::utils::IteratorType::parallel); + + auto linalgOp = rewriter.create( + loc, + /*resultTensorTypes=*/TypeRange{}, ValueRange{op.getLhs(), op.getRhs()}, + ValueRange{op.getOut()}, indexingMaps, iteratorTypes); + + Block &block = linalgOp.getRegion().emplaceBlock(); + block.addArgument(elemTy, loc); + block.addArgument(elemTy, loc); + block.addArgument(elemTy, loc); + + { + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPointToStart(&block); + Value result = rewriter.create(loc, block.getArgument(0), + block.getArgument(1)); + rewriter.create(loc, result); + } + + rewriter.eraseOp(op); + return success(); + } +}; + +// dsa.to_tensor / dsa.to_buffer → bufferization ops. + +/// dsa.to_tensor %src {writable} : memref<...> -> tensor<...> +/// → bufferization.to_tensor %src {restrict, writable} +struct DsaToTensorToBufferizationPattern + : public OpRewritePattern { + explicit DsaToTensorToBufferizationPattern(MLIRContext *ctx) + : OpRewritePattern(ctx, /*benefit=*/3) {} + + LogicalResult matchAndRewrite(mlir::dsa::ToTensorOp op, + PatternRewriter &rewriter) const override { + auto memrefTy = dyn_cast(op.getSrc().getType()); + if (!memrefTy) + return failure(); + + auto resultTy = dyn_cast(op.getResult().getType()); + if (!resultTy) + return failure(); + + auto tensorTy = + RankedTensorType::get(memrefTy.getShape(), memrefTy.getElementType()); + + // Shapes must agree (identity view). + if (tensorTy.getShape() != resultTy.getShape()) + return op->emitRemark( + "dsa.to_tensor shape mismatch between memref and result tensor"); + + // Element types must match (no implicit cast support yet). + if (memrefTy.getElementType() != resultTy.getElementType()) + return op->emitRemark("dsa.to_tensor element type mismatch between " + "memref and result tensor"); + + auto toTensor = rewriter.create( + op.getLoc(), tensorTy, op.getSrc(), + /*restrict=*/true, /*writable=*/op.getWritable()); + rewriter.replaceOp(op, toTensor.getResult()); + return success(); + } +}; + +/// dsa.to_buffer %src, %dst : tensor<...>, memref<...> +/// → bufferization.to_buffer %src : memref<...>; memref.copy %tmp, %dst +struct DsaToBufferToBufferizationPattern + : public OpRewritePattern { + explicit DsaToBufferToBufferizationPattern(MLIRContext *ctx) + : OpRewritePattern(ctx, /*benefit=*/3) {} + + LogicalResult matchAndRewrite(mlir::dsa::ToBufferOp op, + PatternRewriter &rewriter) const override { + Value val = op.getSrc(); + auto valTy = dyn_cast(val.getType()); + if (!valTy) + return failure(); + + auto destMemrefTy = dyn_cast(op.getDst().getType()); + if (!destMemrefTy) + return failure(); + + // Shapes must match. + if (valTy.getShape() != destMemrefTy.getShape()) + return op->emitRemark( + "dsa.to_buffer shape mismatch between value tensor and SPM memref"); + + // Element types must match (no implicit cast support yet). + if (valTy.getElementType() != destMemrefTy.getElementType()) + return op->emitRemark( + "dsa.to_buffer element type mismatch between value and SPM memref"); + + Location loc = op.getLoc(); + + // Materialise the tensor value as a memref, then copy into the SPM buffer. + auto srcMemrefTy = + MemRefType::get(valTy.getShape(), valTy.getElementType()); + auto srcMemref = + rewriter.create(loc, srcMemrefTy, val); + rewriter.create(loc, srcMemref, op.getDst()); + rewriter.eraseOp(op); + return success(); + } +}; + + +} // namespace + +void mlir::triton::populateTLEToMKConversionPatterns( + RewritePatternSet &patterns) { + patterns.add(patterns.getContext()); + patterns.add, + DsaBinaryOpToLinalgPattern, + DsaBinaryOpToLinalgPattern, + DsaBinaryOpToLinalgPattern, + DsaBinaryOpToLinalgPattern, + DsaBinaryOpToLinalgPattern>(patterns.getContext()); + // Highest benefit (3): local load/store → memref ops. + // These MUST fire before any pattern that would produce !tt.ptr types. + patterns.add( + patterns.getContext()); + + // Benefit 2: remote load/store → mk ops. + patterns.add( + patterns.getContext()); + + // Benefit 1: remaining remote_pointers / barrier. + patterns + .add( + patterns.getContext()); +} diff --git a/third_party/wafer/lib/Conversion/TLEToMK/TLEToMKPass.cpp b/third_party/wafer/lib/Conversion/TLEToMK/TLEToMKPass.cpp new file mode 100755 index 00000000..5c0b071d --- /dev/null +++ b/third_party/wafer/lib/Conversion/TLEToMK/TLEToMKPass.cpp @@ -0,0 +1,59 @@ +//===------------------- TLEToMKPass.cpp -------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// Lowering TLE communication ops to backend dialects +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Conversion/TLEToMK/TLEToMK.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "mlir/Transforms/Passes.h" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "magic-kernel/Conversion/TLEToMK/Passes.h.inc" + +namespace { + +class TLEToMKPass : public TLEToMKBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + RewritePatternSet patterns(&getContext()); + populateTLEToMKConversionPatterns(patterns); + + if (failed(applyPatternsGreedily(moduleOp, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr triton::createTLEToMK() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/TritonArithToLinalg/CMakeLists.txt b/third_party/wafer/lib/Conversion/TritonArithToLinalg/CMakeLists.txt new file mode 100755 index 00000000..f20b2102 --- /dev/null +++ b/third_party/wafer/lib/Conversion/TritonArithToLinalg/CMakeLists.txt @@ -0,0 +1,22 @@ +add_triton_library(TritonArithToLinalg + TritonArithToLinalg.cpp + TritonArithToLinalgPass.cpp + + DEPENDS + TritonArithToLinalgConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRLinalgTransforms + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonTilingExtIR + TritonStructuredIR +) diff --git a/third_party/wafer/lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp b/third_party/wafer/lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp new file mode 100755 index 00000000..cd15620a --- /dev/null +++ b/third_party/wafer/lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp @@ -0,0 +1,107 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include +#include + +#define DEBUG_TYPE "triton-arith-to-linalg" +#include "triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.h" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +void mlir::triton::populateTritonArithToLinalgCanonicalizationPatterns( + RewritePatternSet &patterns) { + patterns.add, MinMaxConverter>( + patterns.getContext()); +} + + void mlir::triton::populateTritonTensorPtrConversionPatterns( + RewritePatternSet &patterns) { + patterns.add, + TensorOpConverter, + TensorOpConverter, + TensorOpConverter>(patterns.getContext()); + } + +void mlir::triton::populateTritonArithToLinalgConversionPatterns( + bool pidsToFuncArgs, bool addptrToLinalg, bool assertToCf, + RewritePatternSet &patterns) { + + if (pidsToFuncArgs) { + // Need use wafer interface to get pid. + patterns.add( + patterns.getContext()); + } + if (addptrToLinalg) { + patterns.add(patterns.getContext()); + } + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + + populateExternElementwiseOpToMLIROps(patterns); + + // Reduce converters + // Triton's reduce op is idential to linalg.reduce op, so we can clone + // `tt.reduce` body to `linalg.reduce`. Unfortunately, we still need to + // perform pattern matching to know what reduce ops we are dealing with + // so that we know how to initialize the initial reduce values correctly. + // + // We can do this in a generic way without pattern matching by always using + // the first elements along the reduction axis and perform the reduction on + // the remaining elements. However, this results in creatings sub-tensors that + // aren't always multiple of 2s, which are sub-optimal for certain hardwares. + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + + // Note: the ordering here matters! + // These patterns are added last to they will be tried last. + linalg::populateElementwiseToLinalgConversionPatterns(patterns); +} diff --git a/third_party/wafer/lib/Conversion/TritonArithToLinalg/TritonArithToLinalgPass.cpp b/third_party/wafer/lib/Conversion/TritonArithToLinalg/TritonArithToLinalgPass.cpp new file mode 100755 index 00000000..2eb160fd --- /dev/null +++ b/third_party/wafer/lib/Conversion/TritonArithToLinalg/TritonArithToLinalgPass.cpp @@ -0,0 +1,250 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" +#include "triton-shared/Utils/Utils.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Tensor/Transforms/Transforms.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" + +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "triton-arith-to-linalg" + +using namespace mlir; +using namespace triton; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_TRITONARITHTOLINALG +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class TritonArithToLinalgPass + : public triton::impl::TritonArithToLinalgBase { + using TritonArithToLinalgBase< + TritonArithToLinalgPass>::TritonArithToLinalgBase; + + static auto constexpr LAUNCH_GRID_RANK = getMaxEnumValForProgramIDDim() + 1; + static unsigned int constexpr TRITON_PROGRAM_INFO_ARG_COUNT = + LAUNCH_GRID_RANK * 2; + + // Add additional I32 arguments to represent: + // - num_programs, 3 in total, one for each axis of the launch grid + // - program_id, 3 in total, one for each axis of the launch grid + static void addProgramInfo(triton::FuncOp func) { + OpBuilder b(func); + + auto origFuncType = func.getFunctionType(); + auto origInputTypes = origFuncType.getInputs(); + SmallVector newInputTypes(origInputTypes); + newInputTypes.append(TRITON_PROGRAM_INFO_ARG_COUNT, b.getI32Type()); + + auto newFuncType = + b.getFunctionType(newInputTypes, origFuncType.getResults()); + + func.setFunctionType(newFuncType); + + // Add empty attributes for each new argument if needed + if (func.getAllArgAttrs()) { + SmallVector newArgAttrs; + func.getAllArgAttrs(newArgAttrs); + newArgAttrs.append(TRITON_PROGRAM_INFO_ARG_COUNT, DictionaryAttr()); + func.setAllArgAttrs(newArgAttrs); + } + + // Add the corresponding arguments to function body + for (unsigned int i = 0; i < TRITON_PROGRAM_INFO_ARG_COUNT; i++) { + func.getBody().front().addArgument(b.getI32Type(), func.getLoc()); + } + } + + LogicalResult applyTensorConcatDecomposition() { + auto moduleOp = getOperation(); + MLIRContext *context = &getContext(); + RewritePatternSet patterns(context); + + tensor::populateDecomposeTensorConcatPatterns(patterns); + + if (failed(applyPatternsGreedily(moduleOp, std::move(patterns)))) { + return failure(); + } + return success(); + } + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + { + RewritePatternSet patterns(&getContext()); + populateTritonArithToLinalgCanonicalizationPatterns(patterns); + if (failed(applyPatternsGreedily(moduleOp, std::move(patterns)))) { + signalPassFailure(); + } + } + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, ttx::TritonTilingExtDialect, + tts::TritonStructuredDialect, mk::MagicKernelDialect>(); + + target.addLegalOp(); + + target.addLegalOp(); + + target.addDynamicallyLegalDialect( + [](Operation *op) { + // Lower dense constant to linalg.fill + if (auto constOp = dyn_cast(op)) { + if (!isa(constOp.getResult().getType())) { + return true; + } + + if (auto denseAttr = + dyn_cast(constOp.getValue())) { + if (denseAttr.isSplat() && + isa(denseAttr.getElementType())) { + return false; + } + } + return true; + } + + bool operateOnTensors = + llvm::all_of(op->getOperandTypes(), [](Type type) { + return isa(type); + }); + + return !operateOnTensors; + }); + + if (pidsToFuncArgs) { + // Need use wafer interface to get pid. + target.addIllegalOp< + /* triton::GetProgramIdOp, */ triton::GetNumProgramsOp>(); + } + + if (addptrToLinalg) { + target.addDynamicallyLegalOp([](triton::AddPtrOp op) { + return !isa(op.getResult().getType()); + }); + } + + target.addDynamicallyLegalOp( + [this](triton::BitcastOp op) { + if (!tensorPtrToLinalg) { + return triton::isPtrTypeLike(op.getType()); + } + if (triton::isPtrTypeLike(op.getType())) { + return !isa(op.getType()); + } + return false; + }); + + if (tensorPtrToLinalg) { + target.addDynamicallyLegalOp( + [](auto op) { + return !isa(op->getOperands()[0].getType()); + }); + populateTritonTensorPtrConversionPatterns(patterns); + } + + target.addLegalOp(); + + triton::populateTritonArithToLinalgConversionPatterns( + pidsToFuncArgs, addptrToLinalg, assertToCf, patterns); + + if (pidsToFuncArgs) { + for (auto func : getOperation().getOps()) { + addProgramInfo(func); + } + } + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + + if (failed(applyTensorConcatDecomposition())) { + signalPassFailure(); + } + + // Convert tt.func and tt.return into func's counterparts + if (ttToFuncFunc) { + moduleOp.walk([&](triton::FuncOp func) { + OpBuilder builder(func); + + auto name = func.getName(); + auto type = func.getFunctionType(); + + SmallVector argAttrs, resAttrs; + func.getAllArgAttrs(argAttrs); + func.getAllResultAttrs(resAttrs); + + auto funcFunc = builder.create(func.getLoc(), name, type); + funcFunc.setAllArgAttrs(argAttrs); + funcFunc.setAllResultAttrs(resAttrs); + + auto &funcFuncBody = funcFunc.getBody(); + auto &funcBody = func.getBody(); + + IRMapping map; + funcBody.cloneInto(&funcFuncBody, map); + + for (Block &block : funcFuncBody.getBlocks()) { + auto term = block.getTerminator(); + // Only convert to func.return if the terminator is a tt.return. + // Otherwise, we will accidentally convert cf.br ops which are also + // considered terminators. + if (isa(term)) { + builder.setInsertionPoint(term); + builder.create(func.getLoc(), term->getOperands()); + term->erase(); + } + } + func.erase(); + }); + } + } +}; + +} // namespace + +std::unique_ptr> +triton::createTritonArithToLinalgPass(bool tensorPtrToLinalg) { + TritonArithToLinalgOptions options; + options.tensorPtrToLinalg = tensorPtrToLinalg; + return std::make_unique(options); +} diff --git a/third_party/wafer/lib/Conversion/TritonToCoreDialects/CMakeLists.txt b/third_party/wafer/lib/Conversion/TritonToCoreDialects/CMakeLists.txt new file mode 100755 index 00000000..9d77c79a --- /dev/null +++ b/third_party/wafer/lib/Conversion/TritonToCoreDialects/CMakeLists.txt @@ -0,0 +1,38 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Triton Project Contributors. +# +#===------------------------------------------------------------------------===# + +add_triton_library(TritonToCoreDialects + TritonToCoreDialectsPass.cpp + + DEPENDS + TritonToCoreDialectsConversionPassIncGen + TLEToMKConversionPassIncGen + + LINK_LIBS PUBLIC + TritonTilingExtIR + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + MLIRMathExtDialect + TritonIR + TritonTransforms + ZTCAnalysis + + TritonArithToLinalg + StructuredToMemref + UnstructuredToMemref + WaferTritonToStructured + TritonToUnstructured + TLEToMagicKernel + ConvertTritonPtr + ReconcilePtrCasts + TritonPtrToMemref +) diff --git a/third_party/wafer/lib/Conversion/TritonToCoreDialects/TritonToCoreDialectsPass.cpp b/third_party/wafer/lib/Conversion/TritonToCoreDialects/TritonToCoreDialectsPass.cpp new file mode 100755 index 00000000..5ab0641d --- /dev/null +++ b/third_party/wafer/lib/Conversion/TritonToCoreDialects/TritonToCoreDialectsPass.cpp @@ -0,0 +1,102 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "Address/Dialect/IR/AddressDialect.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "triton-shared/Conversion/ConvertTritonPtr/TritonPtrToAddress.h" +#include "triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h" +#include "triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h" +#include "triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h" +#include "triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h" +#include "triton-shared/Conversion/TritonToCoreDialects/TritonToCoreDialects.h" +#include "triton-shared/Conversion/TritonToStructured/TritonToStructured.h" +#include "triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h" +#include "triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "magic-kernel/Conversion/TLEToMK/TLEToMK.h" + +#include "triton-shared/Conversion/UnstructuredToMK/UnstructuredToMK.h" + +#include "mlir/Conversion/ReconcileUnrealizedCasts/ReconcileUnrealizedCasts.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" + +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/Passes.h" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonToCoreDialects/Passes.h.inc" + +namespace { + +class TritonToCoreDialectsPass + : public TritonToCoreDialectsBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createWaferTritonToStructuredPass()); // flir + + // Erase dead code and fold constants created during lowering + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + pm.addPass(createTritonToUnstructuredPass()); // flir + + pm.addPass(createTritonArithToLinalgPass()); // Tsingmicro + pm.addPass(createStructuredToMemrefPass()); // Tsingmicro + + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + + pm.addPass(createUnstructuredToMemrefPass()); // flir + pm.addPass(createUnstructuredToMKPass()); // Tsingmicro only + + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + + // Convert triton pointers to memref + address dialect + // TODO: Un-ranked memref will all converted to address dialect pointers + pm.addPass(createTritonPtrToMemrefPass()); + pm.addPass(createTritonPtrToAddressPass()); // Tsingmicro only + pm.addPass(createReconcileUnrealizedCastsPass()); // flir + pm.addPass(createReconcilePtrCastsPass()); // Tsingmicro + + // FIXME: RemoveDeadValuesPass is not working now + // pm.addPass(createRemoveDeadValuesPass()); + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + + pm.addPass(createTritonPtrToMemrefPass()); // flir + + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> +triton::createTritonToCoreDialectsPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/UnstructuredToMK/CMakeLists.txt b/third_party/wafer/lib/Conversion/UnstructuredToMK/CMakeLists.txt new file mode 100755 index 00000000..8360cf93 --- /dev/null +++ b/third_party/wafer/lib/Conversion/UnstructuredToMK/CMakeLists.txt @@ -0,0 +1,19 @@ +add_triton_library(UnstructuredToMK + UnstructuredToMKPass.cpp + + DEPENDS + UnstructuredToMKConversionPassIncGen + + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonTilingExtIR + MLIRAddress +) diff --git a/third_party/wafer/lib/Conversion/UnstructuredToMK/UnstructuredToMKPass.cpp b/third_party/wafer/lib/Conversion/UnstructuredToMK/UnstructuredToMKPass.cpp new file mode 100755 index 00000000..5643fc5e --- /dev/null +++ b/third_party/wafer/lib/Conversion/UnstructuredToMK/UnstructuredToMKPass.cpp @@ -0,0 +1,293 @@ +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "triton-shared/Conversion/UnstructuredToMK/UnstructuredToMK.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" +#include "utils/TypeConvertor.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/Support/ErrorHandling.h" +#include + +#define DEBUG_TYPE "unstructured-to-memref" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/UnstructuredToMK/Passes.h.inc" + +namespace { + +struct ScalarAtomicRMWOpConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(tts::IndexedAtomicRMWOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto value = adaptor.getValue(); + if (isa(value.getType())) { + return failure(); + } + + auto loc = op->getLoc(); + + // Calculate the ptr from the offset + auto ptr = adaptor.getPtr(); + auto offset = adaptor.getOffset(); + auto index = rewriter.create( + loc, rewriter.getIndexType(), offset); + auto rankedMemref = rewriter.create( + loc, ptr, getAsOpFoldResult(index), + ArrayRef{rewriter.getIndexAttr(1)} /*sizes*/, + ArrayRef{rewriter.getIndexAttr(1)} /*strides*/); + + auto inputTensorType = + RankedTensorType::get(SmallVector(1, 1), value.getType()); + + auto empty = rewriter.create( + loc, inputTensorType.getShape(), inputTensorType.getElementType()); + auto zero = rewriter.create(loc, 0); + auto valueTensor = rewriter.create( + loc, inputTensorType, value, empty, ValueRange{zero}); + + auto init = rewriter.create( + loc, inputTensorType.getShape(), inputTensorType.getElementType()); + if (op.getMask()) { + // If there is a mask, we need to check it before performing the + // atomic RMW operation. + auto mask = op.getMask(); + auto ifOp = rewriter.create( + loc, mask, + [&](OpBuilder &b, Location loc) { + // TODO: Support other types of inputs and outputs: f32. + auto atomic = rewriter + .create( + loc, inputTensorType, rankedMemref, + valueTensor, init, op.getAtomicRmwOpAttr(), + op.getSemAttr(), op.getScopeAttr()) + ->getResult(0); + + auto resultValue = rewriter.create( + loc, atomic, ValueRange{zero}); + + b.create(loc, resultValue.getResult()); + }, + [&](OpBuilder &b, Location loc) { + // else branch + Value zero = + b.create(loc, b.getZeroAttr(op.getType())); + b.create(loc, zero); + }); + + rewriter.replaceOp(op, ifOp); + } else { + + auto atomic = + rewriter + .create( + loc, inputTensorType, rankedMemref, valueTensor, init, + op.getAtomicRmwOpAttr(), op.getSemAttr(), op.getScopeAttr()) + ->getResult(0); + + auto resultValue = + rewriter.create(loc, atomic, ValueRange{zero}); + + rewriter.replaceOp(op, resultValue.getResult()); + } + + return success(); + } +}; + +struct ScalarAtomicCASOpConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(tts::AtomicCASOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!op.getOffset()) + return failure(); + + if (!op.getType().isIntOrIndexOrFloat()) { + return failure(); + } + + auto loc = op->getLoc(); + auto ptr = adaptor.getPtr(); + auto cmp = adaptor.getCmp(); + auto value = adaptor.getValue(); + + if (isa(value.getType())) { + return failure(); + } + + // Calculate the ptr from the offset + auto offset = adaptor.getOffset(); + auto index = rewriter.create( + loc, rewriter.getIndexType(), offset); + auto rankedMemref = rewriter.create( + loc, ptr, getAsOpFoldResult(index), + ArrayRef{rewriter.getIndexAttr(1)} /*sizes*/, + ArrayRef{rewriter.getIndexAttr(1)} /*strides*/); + + auto inputTensorType = + RankedTensorType::get(SmallVector(1, 1), value.getType()); + + auto empty = rewriter.create( + loc, inputTensorType.getShape(), inputTensorType.getElementType()); + auto zero = rewriter.create(loc, 0); + auto valueTensor = rewriter.create( + loc, inputTensorType, value, empty, ValueRange{zero}); + auto cmpTensor = rewriter.create( + loc, inputTensorType, cmp, empty, ValueRange{zero}); + + auto init = rewriter.create( + loc, inputTensorType.getShape(), inputTensorType.getElementType()); + + // TODO: Support other types of inputs and outputs: f32. + auto atomic = rewriter + .create( + loc, inputTensorType, rankedMemref, cmpTensor, + valueTensor, init, op.getSemAttr(), op.getScopeAttr()) + ->getResult(0); + + auto resultValue = + rewriter.create(loc, atomic, ValueRange{zero}); + + rewriter.replaceOp(op, resultValue.getResult()); + + return success(); + } +}; + +template +struct IndexedAtomicOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename TTS_AtomicOp::Adaptor; + + LogicalResult + matchAndRewrite(TTS_AtomicOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto resultTensorType = + dyn_cast(op.getResult().getType()); + if (!resultTensorType) { + return failure(); + } + + SmallVector inputs(op->getOperands().begin() + 1, + op->getOperands().end()); + + SmallVector outputs = {rewriter.create( + op->getLoc(), resultTensorType.getShape(), + resultTensorType.getElementType())}; + assert(op->getResultTypes().size() == 1); + + auto scalarResultType = + cast(op->getResultTypes().front()).getElementType(); + + // NOTE: linalg.generic cannot nested with linalg.generic (mk.atomic will + // generate linalg.generic inside), so we need to use scf.for to build the + // loop + auto shape = resultTensorType.getShape(); + auto loc = op->getLoc(); + auto zero = rewriter.create(loc, 0); + auto one = rewriter.create(loc, 1); + SmallVector lbs, ubs, steps; + for (auto [i, size] : enumerate(shape)) { + auto sizeValue = rewriter.create(loc, size); + lbs.push_back(zero); + ubs.push_back(sizeValue); + steps.push_back(one); + } + auto loopNest = scf::buildLoopNest( + rewriter, loc, lbs, ubs, steps, outputs, + [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange indices, + ValueRange iterArgs) { + SmallVector regionInputs(op.getNumOperands()); + regionInputs.front() = adaptor.getPtr(); + std::transform(inputs.begin(), inputs.end(), regionInputs.begin() + 1, + [&](auto val) { + return rewriter.create(loc, val, + indices); + }); + auto *scalarOp = nestedBuilder.create( + loc, op->getName().getIdentifier(), regionInputs, + scalarResultType, op->getAttrs()); + + auto outValTensor = nestedBuilder.create( + loc, scalarOp->getResult(0), iterArgs[0], indices); + return SmallVector{outValTensor}; + }); + + rewriter.replaceOp(op, loopNest.results); + + return success(); + } +}; + +class UnstructuredToMKPass : public UnstructuredToMKBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, memref::MemRefDialect, + ttx::TritonTilingExtDialect, mk::MagicKernelDialect>(); + + target.addIllegalOp(); + + PtrToUnrankedMemrefConverter typeConverter; + + patterns.add, + ScalarAtomicCASOpConverter, + IndexedAtomicOpConverter>( + typeConverter, patterns.getContext()); + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) + signalPassFailure(); + } +}; +} // namespace + +std::unique_ptr> triton::createUnstructuredToMKPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/CMakeLists.txt b/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/CMakeLists.txt new file mode 100755 index 00000000..4000bb43 --- /dev/null +++ b/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/CMakeLists.txt @@ -0,0 +1,19 @@ +add_triton_library(WaferMemrefToLLVM + WaferMemrefToLLVM.cpp + WaferMemrefToLLVMPass.cpp + + DEPENDS + WaferMemrefToLLVMConversionPassIncGen + TritonSharedUtils + + LINK_LIBS PUBLIC + TritonSharedUtils + MLIRDialectUtils + MLIRIR + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms +) diff --git a/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.cpp b/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.cpp new file mode 100755 index 00000000..b2ed703d --- /dev/null +++ b/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.cpp @@ -0,0 +1,523 @@ +//===------------------- WaferMemrefToLLVM.cpp------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.h" +#include "mlir/Conversion/LLVMCommon/Pattern.h" +#include "mlir/Conversion/MemRefToLLVM/MemRefToLLVM.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton-shared/Utils/Utils.h" +#include "wafer/Dialect/IR/WaferDialect.h" +#include +#include + +#define DEBUG_TYPE "wafer-memref-to-llvm" + +using namespace mlir; + +#define GEN_PASS_CLASSES +#include "wafer/Conversion/WaferMemrefToLLVM/Passes.h.inc" + +namespace { + +//===----------------------------------------------------------------------===// +// Wafer Custom MemRef Op Conversion Patterns +//===----------------------------------------------------------------------===// + +struct TsmMemRefAllocOpLowering + : public ConvertOpToLLVMPattern { + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + + std::tuple + allocateBufferFromSPM(ConversionPatternRewriter &rewriter, Location loc, + memref::AllocOp allocOp) const { + MemRefType memRefType = allocOp.getType(); + + assert(allocOp->hasAttr("allocation.offset") && + "Expected allocation.offset attribute"); + auto offsetAttr = cast(allocOp->getAttr("allocation.offset")); + // spm memory should start from 0x10000 + auto offset = offsetAttr.getInt() + 0x10000; + + // Align spm address. + if (allocOp.getAlignment().has_value()) { + assert(offset % allocOp.getAlignment().value() == 0 && + "allocation.offset should be aligned to alignment"); + } + + Value spmOffsetOp = + rewriter.create(loc, getIndexType(), offset); + auto elementPtrType = getElementPtrType(memRefType); + Value spmAddr = rewriter.create(loc, elementPtrType); + + spmAddr = rewriter.create(allocOp.getLoc(), + rewriter.getI64Type(), spmAddr); + spmAddr = rewriter.create( + allocOp.getLoc(), rewriter.getI64Type(), spmAddr, spmOffsetOp); + + spmAddr = rewriter.create(allocOp.getLoc(), + elementPtrType, spmAddr); + Value allocatedPtr = spmAddr; + if (!allocatedPtr) + return std::make_tuple(Value(), Value()); + Value alignedPtr = allocatedPtr; + + return std::make_tuple(allocatedPtr, alignedPtr); + } + + LogicalResult + matchAndRewrite(memref::AllocOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto memRefType = op.getType(); + if (!isConvertibleAndHasIdentityMaps(memRefType)) + return rewriter.notifyMatchFailure(op, "incompatible memref type"); + + SmallVector sizes; + SmallVector strides; + Value sizeBytes; + getMemRefDescriptorSizes(op.getLoc(), memRefType, adaptor.getOperands(), + rewriter, sizes, strides, sizeBytes); + + auto [allocatedPtr, alignedPtr] = + allocateBufferFromSPM(rewriter, op.getLoc(), op); + if (!allocatedPtr) + return failure(); + + auto descriptor = createMemRefDescriptor( + op.getLoc(), memRefType, allocatedPtr, alignedPtr, sizes, strides, + rewriter); + rewriter.replaceOp(op, ValueRange{descriptor}); + return success(); + } +}; + +template +struct MemrefLoadOrStoreOpLowering : public ConvertOpToLLVMPattern { + + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + using OpAdaptor = typename MemrefOp::Adaptor; + + LogicalResult + matchAndRewrite(MemrefOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + // Workaround: Should add memory space analysis pass. + Operation *opBase = op; + if (!opBase->hasAttr("isSpm")) { + return rewriter.notifyMatchFailure( + op, "Load/Store should have isSpm attribute."); + } + int isSpm = + cast(opBase->getAttr("isSpm")).getValue().getSExtValue(); + + assert(isSpm && + "Load/Store to global memory has been rewrite to memref::Copy " + "in linalg-to-mk pass."); + + Location loc = op->getLoc(); + auto type = op.getMemRefType(); + bool isMaskEle = type.getElementType().isInteger(1); + MemRefDescriptor memRefDescriptor(adaptor.getMemref()); + Type indexType = ConvertToLLVMPattern::getTypeConverter()->getIndexType(); + + Value dataPtr, Offset; + if (isMaskEle) { + dataPtr = memRefDescriptor.alignedPtr(rewriter, loc); + Value index = memRefDescriptor.offset(rewriter, loc); + auto indices = adaptor.getIndices(); + for (int i = 0, e = indices.size(); i < e; ++i) { + Value stride = memRefDescriptor.stride(rewriter, loc, i); + Value increment = rewriter.create(loc, indices[i], stride); + index = rewriter.create(loc, index, increment); + } + + Value MaskC = rewriter.create(loc, indexType, 7); + Value ShrAmtC = rewriter.create(loc, indexType, 3); + Offset = rewriter.create(loc, index, MaskC); + Offset = + rewriter.create(loc, rewriter.getI8Type(), Offset); + index = rewriter.create(loc, index, ShrAmtC); + dataPtr = rewriter.create( + loc, memRefDescriptor.getElementPtrType(), rewriter.getI8Type(), + dataPtr, index); + } else { + dataPtr = ConvertToLLVMPattern::getStridedElementPtr( + rewriter, op.getLoc(), type, adaptor.getMemref(), + adaptor.getIndices()); + } + + // TODO: Add spm offset according the memory space + auto intPtrType = ConvertToLLVMPattern::getIntPtrType( + memRefDescriptor.getElementPtrType().getAddressSpace()); + Value ptrValue = + rewriter.create(op.getLoc(), intPtrType, dataPtr); + + Value adjustedPtr = dataPtr; + + // Get the module for function declarations + auto module = op->template getParentOfType(); + // Types for function declaration + SmallVector argTypes = { + rewriter.getI64Type() // offset + }; + + auto i8PtrTy = LLVM::LLVMPointerType::get( + rewriter.getContext(), + *ConvertToLLVMPattern::getTypeConverter()->getMemRefAddressSpace(type)); + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction( + module, rewriter, op.getLoc(), "get_spm_memory_mapping_wrapper", + i8PtrTy, argTypes); + + // Create the call to __Rdma + auto spmMemoryAddrPtr = rewriter.create( + op.getLoc(), TypeRange{i8PtrTy}, + "get_spm_memory_mapping_wrapper", // funcPtr, + ValueRange{ptrValue}); + + adjustedPtr = spmMemoryAddrPtr.getResult(); + + // Whether need memoryspace cast + if constexpr (std::is_same()) { + if (isMaskEle) { + Value newVal = rewriter.create(loc, rewriter.getI8Type(), + adjustedPtr, 0, false, + op.getNontemporal()); + Value Zero = + rewriter.create(loc, rewriter.getI8Type(), 0); + Value MaskC = + rewriter.create(loc, rewriter.getI8Type(), 1); + MaskC = rewriter.create(loc, MaskC, Offset); + newVal = rewriter.create(loc, newVal, MaskC); + newVal = rewriter.create(loc, LLVM::ICmpPredicate::ne, + newVal, Zero); + rewriter.replaceOp(op, ValueRange{newVal}); + } else { + rewriter.replaceOpWithNewOp( + op, op.getType(), adjustedPtr, 0, false, op.getNontemporal()); + } + } else { + + Value StoreVal = adaptor.getValue(); + if (isMaskEle) { + Value srcVal = rewriter.create(loc, rewriter.getI8Type(), + adjustedPtr, 0, false, + op.getNontemporal()); + Value MaskC = + rewriter.create(loc, rewriter.getI8Type(), 1); + MaskC = rewriter.create(loc, MaskC, Offset); + Value TrueVal = rewriter.create(loc, srcVal, MaskC); + Value FalseVal = rewriter.create(loc, TrueVal, MaskC); + StoreVal = + rewriter.create(loc, StoreVal, TrueVal, FalseVal); + } + + rewriter.replaceOpWithNewOp(op, StoreVal, adjustedPtr, 0, + false, op.getNontemporal()); + } + + return success(); + } +}; + +struct MemRefReinterpretCastOpLowering + : public ConvertOpToLLVMPattern { + using ConvertOpToLLVMPattern< + memref::ReinterpretCastOp>::ConvertOpToLLVMPattern; + + LogicalResult + matchAndRewrite(memref::ReinterpretCastOp castOp, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type srcType = castOp.getSource().getType(); + + Value descriptor; + if (failed(convertSourceMemRefToDescriptor(rewriter, srcType, castOp, + adaptor, &descriptor))) + return failure(); + rewriter.replaceOp(castOp, {descriptor}); + return success(); + } + +private: + /// Extracts allocated, aligned pointers and offset from a ranked or unranked + /// memref type. In unranked case, the fields are extracted from the + /// underlying ranked descriptor. + void extractPointersAndOffset(Location loc, + ConversionPatternRewriter &rewriter, + const LLVMTypeConverter &typeConverter, + Value originalOperand, Value convertedOperand, + Value *allocatedPtr, Value *alignedPtr, + Value *offset = nullptr) const { + Type operandType = originalOperand.getType(); + if (isa(operandType)) { + MemRefDescriptor desc(convertedOperand); + *allocatedPtr = desc.allocatedPtr(rewriter, loc); + *alignedPtr = desc.alignedPtr(rewriter, loc); + if (offset != nullptr) + *offset = desc.offset(rewriter, loc); + return; + } + + // These will all cause assert()s on unconvertible types. + unsigned memorySpace = *typeConverter.getMemRefAddressSpace( + cast(operandType)); + auto elementPtrType = + LLVM::LLVMPointerType::get(rewriter.getContext(), memorySpace); + + // Extract pointer to the underlying ranked memref descriptor and cast it to + // ElemType**. + UnrankedMemRefDescriptor unrankedDesc(convertedOperand); + + // FIXME: workaround, take memRefDescPtr as naked ptr. + Value underlyingDescPtr = unrankedDesc.memRefDescPtr(rewriter, loc); + *allocatedPtr = underlyingDescPtr; + *alignedPtr = underlyingDescPtr; + + if (offset != nullptr) { + *offset = rewriter.create( + loc, getIndexType(), rewriter.getI32IntegerAttr(0)); + } + } + + LogicalResult convertSourceMemRefToDescriptor( + ConversionPatternRewriter &rewriter, Type srcType, + memref::ReinterpretCastOp castOp, + memref::ReinterpretCastOp::Adaptor adaptor, Value *descriptor) const { + MemRefType targetMemRefType = + cast(castOp.getResult().getType()); + auto llvmTargetDescriptorTy = dyn_cast_or_null( + typeConverter->convertType(targetMemRefType)); + if (!llvmTargetDescriptorTy) + return failure(); + + // Create descriptor. + Location loc = castOp.getLoc(); + auto desc = MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy); + + // Set allocated and aligned pointers. + Value allocatedPtr, alignedPtr; + extractPointersAndOffset(loc, rewriter, *getTypeConverter(), + castOp.getSource(), adaptor.getSource(), + &allocatedPtr, &alignedPtr); + desc.setAllocatedPtr(rewriter, loc, allocatedPtr); + desc.setAlignedPtr(rewriter, loc, alignedPtr); + + // Set offset. + if (castOp.isDynamicOffset(0)) + desc.setOffset(rewriter, loc, adaptor.getOffsets()[0]); + else + desc.setConstantOffset(rewriter, loc, castOp.getStaticOffset(0)); + + // Set sizes and strides. + unsigned dynSizeId = 0; + unsigned dynStrideId = 0; + for (unsigned i = 0, e = targetMemRefType.getRank(); i < e; ++i) { + if (castOp.isDynamicSize(i)) + desc.setSize(rewriter, loc, i, adaptor.getSizes()[dynSizeId++]); + else + desc.setConstantSize(rewriter, loc, i, castOp.getStaticSize(i)); + + if (castOp.isDynamicStride(i)) + desc.setStride(rewriter, loc, i, adaptor.getStrides()[dynStrideId++]); + else + desc.setConstantStride(rewriter, loc, i, castOp.getStaticStride(i)); + } + *descriptor = desc; + return success(); + } +}; + +/// Materialize the MemRef descriptor represented by the results of +/// ExtractStridedMetadataOp. +class ExtractStridedMetadataOpLowering + : public ConvertOpToLLVMPattern { +public: + using ConvertOpToLLVMPattern< + memref::ExtractStridedMetadataOp>::ConvertOpToLLVMPattern; + + LogicalResult + matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp, + OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + if (!LLVM::isCompatibleType(adaptor.getOperands().front().getType())) + return failure(); + + // Create the descriptor. + MemRefDescriptor sourceMemRef(adaptor.getSource()); + Location loc = extractStridedMetadataOp.getLoc(); + Value source = extractStridedMetadataOp.getSource(); + + auto sourceMemRefType = cast(source.getType()); + int64_t rank = sourceMemRefType.getRank(); + SmallVector results; + results.reserve(2 + rank * 2); + + // Base buffer. + Value baseBuffer = sourceMemRef.allocatedPtr(rewriter, loc); + Value alignedBuffer = sourceMemRef.alignedPtr(rewriter, loc); + MemRefDescriptor dstMemRef = MemRefDescriptor::fromStaticShape( + rewriter, loc, *getTypeConverter(), + cast(extractStridedMetadataOp.getBaseBuffer().getType()), + baseBuffer, alignedBuffer); + results.push_back((Value)dstMemRef); + + // Offset. + results.push_back(sourceMemRef.offset(rewriter, loc)); + + // Sizes. + for (unsigned i = 0; i < rank; ++i) + results.push_back(sourceMemRef.size(rewriter, loc, i)); + // Strides. + for (unsigned i = 0; i < rank; ++i) + results.push_back(sourceMemRef.stride(rewriter, loc, i)); + + rewriter.replaceOp(extractStridedMetadataOp, results); + return success(); + } +}; + +/// Unpack the pointer returned by a memref.extract_aligned_pointer_as_index. +class ConvertExtractAlignedPointerAsIndex + : public ConvertOpToLLVMPattern { +public: + using ConvertOpToLLVMPattern< + memref::ExtractAlignedPointerAsIndexOp>::ConvertOpToLLVMPattern; + + LogicalResult + matchAndRewrite(memref::ExtractAlignedPointerAsIndexOp extractOp, + OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + BaseMemRefType sourceTy = extractOp.getSource().getType(); + + Value alignedPtr; + if (sourceTy.hasRank()) { + MemRefDescriptor desc(adaptor.getSource()); + alignedPtr = desc.alignedPtr(rewriter, extractOp->getLoc()); + } else { + auto elementPtrTy = LLVM::LLVMPointerType::get( + rewriter.getContext(), sourceTy.getMemorySpaceAsInt()); + + UnrankedMemRefDescriptor desc(adaptor.getSource()); + Value descPtr = desc.memRefDescPtr(rewriter, extractOp->getLoc()); + + alignedPtr = UnrankedMemRefDescriptor::alignedPtr( + rewriter, extractOp->getLoc(), *getTypeConverter(), descPtr, + elementPtrTy); + } + + rewriter.replaceOpWithNewOp( + extractOp, getTypeConverter()->getIndexType(), alignedPtr); + return success(); + } +}; + +// Copy from llvm-project MemrefToLLVM.cpp +// FIXME: Use ptr dialect to fix the error between un-ranked memref to llvm ptr +struct MemRefCastOpLowering : public ConvertOpToLLVMPattern { + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + +public: + LogicalResult matchCast(memref::CastOp memRefCastOp) const { + Type srcType = memRefCastOp.getOperand().getType(); + Type dstType = memRefCastOp.getType(); + + // memref::CastOp reduce to bitcast in the ranked MemRef case and can be + // used for type erasure. For now they must preserve underlying element type + // and require source and result type to have the same rank. Therefore, + // perform a sanity check that the underlying structs are the same. Once op + // semantics are relaxed we can revisit. + if (isa(srcType) && isa(dstType)) + return success(typeConverter->convertType(srcType) == + typeConverter->convertType(dstType)); + + // At least one of the operands is unranked type + assert(isa(srcType) || + isa(dstType)); + + // Unranked to unranked cast is disallowed + return !(isa(srcType) && + isa(dstType)) + ? success() + : failure(); + } + + LogicalResult + matchAndRewrite(memref::CastOp memRefCastOp, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (failed(matchCast(memRefCastOp))) + return failure(); + + auto srcType = memRefCastOp.getOperand().getType(); + auto dstType = memRefCastOp.getType(); + auto targetStructType = typeConverter->convertType(memRefCastOp.getType()); + auto loc = memRefCastOp.getLoc(); + + // For ranked/ranked case, just keep the original descriptor. + if (isa(srcType) && isa(dstType)) { + rewriter.replaceOp(memRefCastOp, {adaptor.getSource()}); + return success(); + } + + if (isa(srcType) && isa(dstType)) { + // Casting ranked to unranked memref type + // Set the rank in the destination from the memref type + // Allocate space on the stack and copy the src memref descriptor + // Set the ptr in the destination to the stack space + auto srcMemRefType = cast(srcType); + int64_t rank = srcMemRefType.getRank(); + // ptr = AllocaOp sizeof(MemRefDescriptor) + auto ptr = getTypeConverter()->promoteOneMemRefDescriptor( + loc, adaptor.getSource(), rewriter); + + // rank = ConstantOp srcRank + auto rankVal = rewriter.create( + loc, getIndexType(), rewriter.getIndexAttr(rank)); + // poison = PoisonOp + UnrankedMemRefDescriptor memRefDesc = + UnrankedMemRefDescriptor::poison(rewriter, loc, targetStructType); + // d1 = InsertValueOp poison, rank, 0 + memRefDesc.setRank(rewriter, loc, rankVal); + // d2 = InsertValueOp d1, ptr, 1 + memRefDesc.setMemRefDescPtr(rewriter, loc, ptr); + rewriter.replaceOp(memRefCastOp, (Value)memRefDesc); + + } else if (isa(srcType) && isa(dstType)) { + // Casting from unranked type to ranked. + // The operation is assumed to be doing a correct cast. If the destination + // type mismatches the unranked the type, it is undefined behavior. + UnrankedMemRefDescriptor memRefDesc(adaptor.getSource()); + auto ptr = memRefDesc.memRefDescPtr(rewriter, loc); + + auto desc = MemRefDescriptor::poison(rewriter, loc, targetStructType); + // FIXME: workaround, take memRefDescPtr as naked ptr. + desc.setAllocatedPtr(rewriter, loc, ptr); + desc.setAlignedPtr(rewriter, loc, ptr); + + rewriter.replaceOp(memRefCastOp, SmallVector{desc}); + } else { + llvm_unreachable("Unsupported unranked memref to unranked memref cast"); + } + return success(); + } +}; + +} // namespace + +void mlir::triton::populateWaferMemrefToLLVMConversionPatterns( + RewritePatternSet &patterns, LLVMTypeConverter &converter) { + // clang-format off + patterns.add, + MemrefLoadOrStoreOpLowering>( + converter); + // clang-format on +} diff --git a/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVMPass.cpp b/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVMPass.cpp new file mode 100755 index 00000000..217f247d --- /dev/null +++ b/third_party/wafer/lib/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVMPass.cpp @@ -0,0 +1,88 @@ +//===------------------- WaferMemrefToLLVMPass.cpp--------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "mlir/Conversion/LLVMCommon/TypeConverter.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Func/Transforms/FuncConversions.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "wafer/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.h" +#include "wafer/Dialect/IR/WaferDialect.h" +#include "llvm/Support/Debug.h" +#include +#include +#include + +#define DEBUG_TYPE "wafer-memref-to-llvm" + +using namespace mlir; + +namespace mlir { +namespace triton { +#define GEN_PASS_CLASSES +#include "wafer/Conversion/WaferMemrefToLLVM/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class WaferMemrefToLLVMPass + : public mlir::triton::WaferMemrefToLLVMBase { + using WaferMemrefToLLVMBase::WaferMemrefToLLVMBase; + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + MLIRContext *context = &getContext(); + RewritePatternSet patterns(context); + ConversionTarget target(*context); + + target.addIllegalOp< + memref::AllocOp, memref::LoadOp, memref::StoreOp, + memref::ReinterpretCastOp, memref::ExtractStridedMetadataOp, + memref::ExtractAlignedPointerAsIndexOp, memref::CastOp>(); + + target.addLegalDialect(); + + target.addLegalOp(); + + LowerToLLVMOptions options(context); + options.useBarePtrCallConv = false; + LLVMTypeConverter llvmTypeConverter(context, options); + triton::populateWaferMemrefToLLVMConversionPatterns(patterns, + llvmTypeConverter); + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + } +}; + +} // namespace + +std::unique_ptr> triton::createWaferMemrefToLLVMPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/WaferToLLVM/CMakeLists.txt b/third_party/wafer/lib/Conversion/WaferToLLVM/CMakeLists.txt new file mode 100755 index 00000000..a9cd03d8 --- /dev/null +++ b/third_party/wafer/lib/Conversion/WaferToLLVM/CMakeLists.txt @@ -0,0 +1,25 @@ +add_triton_library(WaferToLLVM + WaferToLLVM.cpp + KernelArgBufferPass.cpp + + DEPENDS + WaferToLLVMConversionPassIncGen + KernelArgBufferPassIncGen + MLIRMemRefToLLVM + + LINK_LIBS PUBLIC + MLIRMemRefToLLVM + MLIRArithDialect + MLIRArithToLLVM + MLIRFuncDialect + MLIRFuncToLLVM + MLIRLLVMDialect + MLIRMemRefDialect + MLIRMemRefToLLVM + MLIRArithToLLVM + MLIRAffineToStandard + MLIRLinalgToStandard + MLIRSCFDialect + MLIRSCFToControlFlow + MLIRTransforms +) diff --git a/third_party/wafer/lib/Conversion/WaferToLLVM/KernelArgBufferPass.cpp b/third_party/wafer/lib/Conversion/WaferToLLVM/KernelArgBufferPass.cpp new file mode 100755 index 00000000..18f6f993 --- /dev/null +++ b/third_party/wafer/lib/Conversion/WaferToLLVM/KernelArgBufferPass.cpp @@ -0,0 +1,161 @@ +//===- KernelArgBufferPass.cpp - Convert kernel args to single buffer -----===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This pass transforms kernel function signatures by converting multiple +// arguments into a single void* buffer containing all the arguments. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Conversion/WaferToLLVM/KernelArgBufferPass.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "llvm/ADT/TypeSwitch.h" + +using namespace mlir; + +namespace mlir { +namespace triton { +#define GEN_PASS_CLASSES +#include "wafer/Conversion/WaferToLLVM/KernelArgBufferPass.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class KernelArgBufferPass + : public mlir::triton::KernelArgBufferPassBase { + using KernelArgBufferPassBase::KernelArgBufferPassBase; + +private: + // Check if the function is a kernel function + bool isKernelFunction(LLVM::LLVMFuncOp func); + +public: + StringRef getArgument() const final { return "kernel-arg-buffer"; } + StringRef getDescription() const final { + return "Convert kernel arguments to a single buffer argument"; + } + + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override; + +private: + // Insert load op to get real kernel args from new buffered argument + // Side effect: calculate offset and create ops + Value insertKernelArgLoad(OpBuilder &builder, Location loc, Value argsBuffer, + Type argType, int64_t ¤tOffset); +}; + +bool KernelArgBufferPass::isKernelFunction(LLVM::LLVMFuncOp func) { + // NOTE: Need consider math-to-libm and the declared llvm functions (eg. + // __Print, get_spm_memory_mapping_wrapper) + // + // WORKAROUND: For some reason, func.isDeclaration() always returns false and + // func.getLinkage() always returns External, even for kernel functions with + // body. So we use func.getCallableRegion() to check if the function is a + // kernel function. This is a workaround and should be replaced with a more + // robust solution in the future. + return func.getCallableRegion() != nullptr; +} + +Value KernelArgBufferPass::insertKernelArgLoad(OpBuilder &builder, Location loc, + Value argsBuffer, Type argType, + int64_t ¤tOffset) { + // Get pointer to the current position in args buffer + auto offsetValue = builder.create( + loc, builder.getI64Type(), builder.getI64IntegerAttr(currentOffset)); + + // NOTE: GEPOp need distinguish the scalar and ptr type. So here ptr + offset + Value elementPtr = + builder.create(loc, builder.getI64Type(), argsBuffer); + elementPtr = builder.create(loc, builder.getI64Type(), + elementPtr, offsetValue); + elementPtr = builder.create( + loc, LLVM::LLVMPointerType::get(builder.getContext()), elementPtr); + + // Increment offset. Assume all args are 8 bytes + currentOffset += sizeof(int64_t); + + // Load the real kernel arg value + return builder.create(loc, argType, elementPtr); +} + +void KernelArgBufferPass::runOnOperation() { + ModuleOp module = getOperation(); + OpBuilder builder(module.getContext()); + + // Collect functions to process + SmallVector kernelFuncs; + for (auto func : module.getOps()) { + if (!isKernelFunction(func)) + continue; + kernelFuncs.push_back(func); + } + // NOTE: We move this pass before wafer-to-llvm pass. + // So we assume the func op must be only one and must be the triton kernel + assert(kernelFuncs.size() == 1 && "Only one kernel function expected"); + + // Process each kernel function + // TODO: Delete the for loop if the assert is always true for all examples + for (auto func : kernelFuncs) { + // Create new function with bufferized signature + builder.setInsertionPointAfter(func); + // Save the old block arguments + SmallVector blockArguments = + llvm::to_vector<8>(func.getArguments()); + auto numArguments = blockArguments.size(); + + // New bufferized arg type + auto voidPtrType = LLVM::LLVMPointerType::get(builder.getContext()); + + // New bufferized function type + auto newFuncType = LLVM::LLVMFunctionType::get( + func.getFunctionType().getReturnType(), voidPtrType); + func.setFunctionType(newFuncType); + SmallVector newArgAttrs({DictionaryAttr()}); + func.setAllArgAttrs(newArgAttrs); + + // Add the new bufferized argument + Location loc = func.getLoc(); + Block &entryBlock = func.getBlocks().front(); + entryBlock.insertArgument((unsigned)0, voidPtrType, func.getLoc()); + + OpBuilder builder(&entryBlock, entryBlock.begin()); + // Get the bufferized argument + Value argsBuffer = entryBlock.getArgument(0); + + // Offset tracking for buffer access + int64_t currentOffset = 0; + + // Process each original argument + for (auto argIndex : llvm::seq(0, numArguments)) { + auto oldArg = blockArguments[argIndex]; + Type argType = oldArg.getType(); + Value loadedArg = insertKernelArgLoad(builder, func.getLoc(), argsBuffer, + argType, currentOffset); + + if (blockArguments[argIndex].use_empty()) + continue; + oldArg.replaceAllUsesWith(loadedArg); + } + // Remove the old arguments when replace the use-chain + entryBlock.eraseArguments(1, numArguments); + } +} + +} // namespace + +std::unique_ptr triton::createKernelArgBufferPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/WaferToLLVM/WaferToLLVM.cpp b/third_party/wafer/lib/Conversion/WaferToLLVM/WaferToLLVM.cpp new file mode 100755 index 00000000..22f2c834 --- /dev/null +++ b/third_party/wafer/lib/Conversion/WaferToLLVM/WaferToLLVM.cpp @@ -0,0 +1,2716 @@ +//===--------------------- WaferToLLVM.cpp ---------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file implements the patterns to convert operations from wafer dialect to +// LLVM IR dialect. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Conversion/WaferToLLVM/WaferToLLVM.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Conversion/AffineToStandard/AffineToStandard.h" +#include "mlir/Conversion/ArithToLLVM/ArithToLLVM.h" +#include "mlir/Conversion/ControlFlowToLLVM/ControlFlowToLLVM.h" +#include "mlir/Conversion/FuncToLLVM/ConvertFuncToLLVM.h" +#include "mlir/Conversion/LLVMCommon/Pattern.h" +#include "mlir/Conversion/LLVMCommon/TypeConverter.h" +#include "mlir/Conversion/LinalgToStandard/LinalgToStandard.h" +#include "mlir/Conversion/MathToLLVM/MathToLLVM.h" +#include "mlir/Conversion/MemRefToLLVM/MemRefToLLVM.h" +#include "mlir/Conversion/SCFToControlFlow/SCFToControlFlow.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Func/Transforms/FuncConversions.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton-shared/Utils/Utils.h" +#include "triton/Conversion/TritonGPUToLLVM/Utility.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "wafer/Dialect/IR/WaferDialect.h" +#include "llvm/ADT/TypeSwitch.h" + +#ifdef DEBUG_TYPE +#undef DEBUG_TYPE +#endif +#define DEBUG_TYPE "wafer-to-llvm" + +using namespace mlir; + +#define GEN_PASS_CLASSES +#include "wafer/Conversion/WaferToLLVM/Passes.h.inc" + +namespace { +//===----------------------------------------------------------------------===// +// Helper Functions +//===----------------------------------------------------------------------===// +// Crt func name +const char rdma4dFuncName[] = "__Rdma4d"; +const char wdma4dFuncName[] = "__Wdma4d"; +const char rdma1dFuncName[] = "__Rdma1d"; +const char wdma1dFuncName[] = "__Wdma1d"; +const char rdmaFuncName[] = "__Rdma"; +const char wdmaFuncName[] = "__Wdma"; +const char memcpyFuncName[] = "__Memcpy"; +const char recvFuncName[] = "__Recv"; +const char sendFuncName[] = "__Send"; +const char addVVFuncName[] = "__AddVV"; +const char subVVFuncName[] = "__SubVV"; +const char mulVVFuncName[] = "__MulVV"; +const char divVVFuncName[] = "__DivVV"; +const char absVVFuncName[] = "__AbsVV"; +const char rsqrtVVFuncName[] = "__RsqrtVV"; +const char sqrtVVFuncName[] = "__SqrtVV"; +const char recipVVFuncName[] = "__RecipVV"; +const char negVVFuncName[] = "__NegVV"; +const char lnFuncName[] = "__Ln"; +const char log2FuncName[] = "__Log2"; +const char expFuncName[] = "__Exp"; +const char pow2FuncName[] = "__Pow2"; +const char sinFuncName[] = "__Sin"; +const char cosFuncName[] = "__Cos"; +const char addVSFuncName[] = "__AddVS"; +const char subVSFuncName[] = "__SubVS"; +const char mulVSFuncName[] = "__MulVS"; +const char divVSFuncName[] = "__DivVS"; +const char argMinFuncName[] = "__ArgMin"; +const char argMaxFuncName[] = "__ArgMax"; +const char reduceSumFuncName[] = "__ReduceSum"; +const char reduceMaxFuncName[] = "__ReduceMax"; +const char reduceMinFuncName[] = "__ReduceMin"; +const char reduceMulFuncName[] = "__ReduceMul"; +// Int8 +const char int8ToBf16FuncName[] = "__INT8_BF16"; +const char int8ToFp16FuncName[] = "__INT8_FP16"; +const char int8ToFp32FuncName[] = "__INT8_FP32"; +const char int8ToTf32FuncName[] = "__INT8_TF32"; +// Int16 +const char int16ToFp16FuncName[] = "__INT16_FP16"; +const char int16ToBf16FuncName[] = "__INT16_BF16"; +const char int16ToFp32FuncName[] = "__INT16_FP32"; +const char int16ToTf32FuncName[] = "__INT16_TF32"; +// Int32 +const char int32ToFp16FuncName[] = "__INT32_FP16"; +const char int32ToBf16FuncName[] = "__INT32_BF16"; +const char int32ToFp32FuncName[] = "__INT32_FP32"; +const char int32ToTf32FuncName[] = "__INT32_TF32"; +// BF16 +const char bf16ToInt8FuncName[] = "__BF16_INT8"; +const char bf16ToInt16FuncName[] = "__BF16_INT16"; +const char bf16ToInt32FuncName[] = "__BF16_INT32"; +const char bf16ToFp16FuncName[] = "__BF16_FP16"; +const char bf16ToFp32FuncName[] = "__BF16_FP32"; +const char bf16ToTf32FuncName[] = "__BF16_TF32"; +// FP16 +const char fp16ToBf16FuncName[] = "__FP16_BF16"; +const char fp16ToFp32FuncName[] = "__FP16_FP32"; +const char fp16ToTf32FuncName[] = "__FP16_TF32"; +const char fp16ToInt8FuncName[] = "__FP16_INT8"; +const char fp16ToInt16FuncName[] = "__FP16_INT16"; +const char fp16ToInt32FuncName[] = "__FP16_INT32"; +// FP32 +const char fp32ToInt8FuncName[] = "__FP32_INT8"; +const char fp32ToInt16FuncName[] = "__FP32_INT16"; +const char fp32ToInt32FuncName[] = "__FP32_INT32"; +const char fp32ToFp16FuncName[] = "__FP32_FP16"; +const char fp32ToBf16FuncName[] = "__FP32_BF16"; +const char fp32ToTf32FuncName[] = "__FP32_TF32"; +// TF32 +const char tf32ToInt8FuncName[] = "__TF32_INT8"; +const char tf32ToInt16FuncName[] = "__TF32_INT16"; +const char tf32ToInt32FuncName[] = "__TF32_INT32"; +const char tf32ToFp16FuncName[] = "__TF32_FP16"; +const char tf32ToBf16FuncName[] = "__TF32_BF16"; +const char tf32ToFp32FuncName[] = "__TF32_FP32"; +// MXFP +const char fp8E4M3ToBF16FuncName[] = "__FP8E4M3_BF16"; +const char fp8E4M3FNToBF16FuncName[] = "__FP8E4M3FN_BF16"; +const char fp8E5M2ToBF16FuncName[] = "__FP8E5M2_BF16"; +const char fp4E2M1ToBF16FuncName[] = "__FP4E2M1_BF16"; +const char fp8E4M3ToFP16FuncName[] = "__FP8E4M3_FP16"; +const char fp8E4M3FNToFP16FuncName[] = "__FP8E4M3FN_FP16"; +const char fp8E5M2ToFP16FuncName[] = "__FP8E5M2_FP16"; +const char fp4E2M1ToFP16FuncName[] = "__FP4E2M1_FP16"; + +const char boolEqualVVFuncName[] = "__BoolEqualVV"; +const char boolUnEqualVVFuncName[] = "__BoolUnEqualVV"; +const char boolGreaterEqualVVFuncName[] = "__BoolGreaterEqualVV"; +const char boolGreaterVVFuncName[] = "__BoolGreaterVV"; +const char boolLessEqualVVFuncName[] = "__BoolLessEqualVV"; +const char boolLessThenVVFuncName[] = "__BoolLessThenVV"; +const char equalVVFuncName[] = "__EqualVV"; +const char unEqualVVFuncName[] = "__UnEqualVV"; +const char greaterEqualVVFuncName[] = "__GreaterEqualVV"; +const char greaterVVFuncName[] = "__GreaterVV"; +const char lessEqualVVFuncName[] = "__LessEqualVV"; +const char lessThenVVFuncName[] = "__LessThenVV"; +const char boolEqualVSFuncName[] = "__BoolEqualVS"; +const char boolUnEqualVSFuncName[] = "__BoolUnEqualVS"; +const char boolGreaterEqualVSFuncName[] = "__BoolGreaterEqualVS"; +const char boolGreaterVSFuncName[] = "__BoolGreaterVS"; +const char boolLessEqualVSFuncName[] = "__BoolLessEqualVS"; +const char boolLessThenVSFuncName[] = "__BoolLessThenVS"; +const char equalVSFuncName[] = "__EqualVS"; +const char unEqualVSFuncName[] = "__UnEqualVS"; +const char greaterEqualVSFuncName[] = "__GreaterEqualVS"; +const char greaterVSFuncName[] = "__GreaterVS"; +const char lessEqualVSFuncName[] = "__LessEqualVS"; +const char lessThenVSFuncName[] = "__LessThenVS"; +const char andVVFuncName[] = "__AndVV"; +const char orVVFuncName[] = "__OrVV"; +const char xorVVFuncName[] = "__XorVV"; +const char boolNotVFuncName[] = "__BoolNotV"; +const char boolAndVFuncName[] = "__BoolAndV"; +const char boolOrVFuncName[] = "__BoolOrV"; +const char boolXorVFuncName[] = "__BoolXorV"; +const char MaxVVFuncName[] = "__MaxVV"; +const char MinVVFuncName[] = "__MinVV"; +const char transposeFuncName[] = "__Transpose"; +const char nchw2nhwcFuncName[] = "__Nchw2nhwc"; +const char nhwc2nchwFuncName[] = "__Nhwc2nchw"; +const char tanhFuncName[] = "__Tanh"; +const char atomicBarrierInFuncName[] = "__AtomicBarrierIn"; +const char atomicBarrierOutFuncName[] = "__AtomicBarrierOut"; +const char MXFPScaleBF16FuncName[] = "__mxfpScaleBF16"; +const char MXFPScaleFP16FuncName[] = "__mxfpScaleFP16"; + +static Value adjustElemCountType(ConversionPatternRewriter &rewriter, + Location loc, Value elemCount) { + Value newElemCount = elemCount; + if (isa(elemCount.getType())) { + newElemCount = rewriter.create( + loc, rewriter.getI32Type(), elemCount); + } else if (isa(elemCount.getType())) { + auto elemCountType = dyn_cast(elemCount.getType()); + if (elemCountType.isInteger(64)) + newElemCount = rewriter.create( + loc, rewriter.getI32Type(), elemCount); + } + return newElemCount; +} + +static Value castIndexToInt32(ConversionPatternRewriter &rewriter, Location loc, + Value indexOp) { + return rewriter.create(loc, rewriter.getI32Type(), + indexOp); +} + +static Value createInt32ValueArray(ConversionPatternRewriter &rewriter, + Location loc, SmallVector array, + Operation *currentOp) { + auto i32PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i64Ty = rewriter.getI64Type(); + + // Find the parent function of the current operation + Operation *parentFunc = currentOp->getParentOfType(); + LLVM::LLVMFuncOp funcOp = dyn_cast(parentFunc); + assert(funcOp && + "Expected to find a parent function for the current operation\n"); + + // Save current insertion point + auto savedInsertionPoint = rewriter.saveInsertionPoint(); + + // Insert alloca at the beginning of the function entry block + Block &entryBlock = funcOp.getBody().front(); + rewriter.setInsertionPointToStart(&entryBlock); + + // Allocate memory for array + Value rank = rewriter.create( + loc, i64Ty, rewriter.getI64IntegerAttr(array.size())); + + auto allocaOp = rewriter.create(loc, i32PtrTy, i32Ty, rank); + + rewriter.restoreInsertionPoint(savedInsertionPoint); + + // assert(moduleOp && moduleOp->hasAttr("triton_tsm.spm_use") && + // "ModuleOp should not be null when creating an array"); + // auto spmPointer = + // cast(moduleOp->getAttr("triton_tsm.spm_use")) + // .getValue() + // .getZExtValue(); + // Value spmOffsetOp = rewriter.create( + // loc, rewriter.getI64Type(), rewriter.getI32IntegerAttr(spmPointer)); + // auto elementPtrType = LLVM::LLVMPointerType::get(rewriter.getContext()); + // Value spmAddr = rewriter.create(loc, elementPtrType); + // spmAddr = + // rewriter.create(loc, rewriter.getI64Type(), + // spmAddr); + // spmAddr = rewriter.create(loc, rewriter.getI64Type(), + // spmAddr, + // spmOffsetOp); + + // // Types for function declaration + // SmallVector argTypes = { + // rewriter.getI64Type() // offset + // }; + + // Declare the function + // Value funcPtr = triton::utils::declareWaferRuntimeFunction( + // moduleOp, rewriter, loc, "get_spm_memory_mapping_wrapper", + // elementPtrType, argTypes); + + // Create the call to __Rdma + // auto spmMemoryAddrPtr = rewriter.create( + // loc, TypeRange{elementPtrType}, + // "get_spm_memory_mapping_wrapper", // funcPtr, + // ValueRange{spmAddr}); + + // Restore insertion point + rewriter.restoreInsertionPoint(savedInsertionPoint); + + // Store each dimension in the array + for (size_t i = 0; i < array.size(); i++) { + // Create the index + Value idx = rewriter.create( + loc, i64Ty, rewriter.getI32IntegerAttr(i)); + + // Create GEP to get pointer to array element + Value elemPtr = rewriter.create(loc, i32PtrTy, i32Ty, allocaOp, + ArrayRef{idx}); + + // Store the value + rewriter.create(loc, array[i], elemPtr); + } + + // spmPointer += array.size() * sizeof(int32_t); + // // Record spm usage. + // moduleOp->setAttr( + // "triton_tsm.spm_use", + // mlir::IntegerAttr::get(mlir::IntegerType::get(moduleOp.getContext(), + // 32), + // spmPointer)); + + return allocaOp; +} + +static Value +indexValueArrayToInt32ValueArray(ConversionPatternRewriter &rewriter, + Location loc, ValueRange array, + Operation *currentOp) { + + SmallVector arrayValues; + for (size_t i = 0; i < array.size(); i++) { + // Create the dimension value + arrayValues.push_back(castIndexToInt32(rewriter, loc, array[i])); + } + + return createInt32ValueArray(rewriter, loc, arrayValues, currentOp); +} + +static Value int32ArrayToInt32ValueArray(ConversionPatternRewriter &rewriter, + Location loc, ArrayRef array, + Operation *currentOp) { + + SmallVector arrayValues; + auto i32Ty = rewriter.getI32Type(); + for (size_t i = 0; i < array.size(); i++) { + // Create the dimension value + arrayValues.push_back(rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(array[i]))); + } + return createInt32ValueArray(rewriter, loc, arrayValues, currentOp); +} + +//===----------------------------------------------------------------------===// +// Arith Operation Conversion Patterns +//===----------------------------------------------------------------------===// + +// Convert constant operations to LLVM constants +struct ConstantOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(arith::ConstantOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the constant value + auto constAttr = op.getValue(); + + // Get the result type + auto resultType = getTypeConverter()->convertType(op.getResult().getType()); + + // Handle different attribute types + if (auto intAttr = dyn_cast(constAttr)) { + // Convert integer attribute + rewriter.replaceOpWithNewOp(op, resultType, intAttr); + return success(); + } else if (auto floatAttr = dyn_cast(constAttr)) { + // Convert float attribute + rewriter.replaceOpWithNewOp(op, resultType, floatAttr); + return success(); + } else if (auto boolAttr = dyn_cast(constAttr)) { + // Convert bool attribute to i1 + rewriter.replaceOpWithNewOp( + op, resultType, + rewriter.getIntegerAttr(resultType, boolAttr.getValue())); + return success(); + } + + return failure(); + } +}; + +// Convert arith.index_cast to appropriate LLVM conversions +struct IndexCastOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(arith::IndexCastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get source and result types + auto srcType = adaptor.getIn().getType(); + auto dstType = getTypeConverter()->convertType(op.getResult().getType()); + + // Convert from index to specific integer type + if (isa(srcType) && isa(dstType)) { + rewriter.replaceOpWithNewOp(op, dstType, + adaptor.getIn()); + return success(); + } + + // Convert from specific integer type to index + if (isa(srcType) && isa(dstType)) { + rewriter.replaceOpWithNewOp(op, dstType, + adaptor.getIn()); + return success(); + } + + // Handle integer to integer casts + if (isa(srcType) && isa(dstType)) { + unsigned srcWidth = cast(srcType).getWidth(); + unsigned dstWidth = cast(dstType).getWidth(); + + if (srcWidth < dstWidth) { + // Sign extend if source is signed, zero extend otherwise + rewriter.replaceOpWithNewOp(op, dstType, adaptor.getIn()); + } else if (srcWidth > dstWidth) { + // Truncate + rewriter.replaceOpWithNewOp(op, dstType, + adaptor.getIn()); + } else { + // Same width, just pass through + rewriter.replaceOp(op, adaptor.getIn()); + } + return success(); + } + + return failure(); + } +}; + +// Convert arith.addi to LLVM add +struct AddIOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(arith::AddIOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, adaptor.getLhs(), + adaptor.getRhs()); + return success(); + } +}; + +// Convert arith.muli to LLVM mul +struct MulIOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(arith::MulIOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, adaptor.getLhs(), + adaptor.getRhs()); + return success(); + } +}; + +//===----------------------------------------------------------------------===// +// Wafer Operation Conversion Patterns +//===----------------------------------------------------------------------===// + +struct BarrierConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::BarrierOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __Barrier runtime function if not already declared + /* + void __Barrier() + */ + + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + "__Barrier", i8PtrTy, {}); + + // Create the call to __Barrier + auto call = rewriter.create(op.getLoc(), TypeRange{i8PtrTy}, + "__Barrier", // funcPtr, + ValueRange{}); + + // Replace the op with the call + rewriter.eraseOp(op); + + return success(); + } +}; + +struct RandGenOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::RandGenOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + auto module = op->getParentOfType(); + auto *ctx = rewriter.getContext(); + auto voidTy = LLVM::LLVMVoidType::get(ctx); + auto i8PtrTy = LLVM::LLVMPointerType::get(ctx); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // void __RandGen(uint64_t *src0, uint64_t *src1, uint64_t *dst0, + // uint64_t *dst1, uint64_t *dst2, uint32_t byte_count, + // uint16_t fmt); + SmallVector argTypes = {i8PtrTy, i8PtrTy, i8PtrTy, i8PtrTy, + i8PtrTy, i32Ty, i16Ty}; + (void)triton::declareWaferRuntimeFunction(module, rewriter, loc, "__RandGen", + voidTy, argTypes); + + Value src0 = + rewriter.create(loc, i8PtrTy, adaptor.getSrc0()); + Value src1 = + rewriter.create(loc, i8PtrTy, adaptor.getSrc1()); + Value dst0 = + rewriter.create(loc, i8PtrTy, adaptor.getDst0()); + Value dst1 = + rewriter.create(loc, i8PtrTy, adaptor.getDst1()); + Value dst2 = + rewriter.create(loc, i8PtrTy, adaptor.getDst2()); + Value byteCount = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(op.getElemNum())); + Value fmt = rewriter.create( + loc, i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + rewriter.create( + loc, TypeRange{}, "__RandGen", + ValueRange{src0, src1, dst0, dst1, dst2, byteCount, fmt}); + rewriter.eraseOp(op); + return success(); + } +}; + + +template +struct AtomicBarrierOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto module = op->template getParentOfType(); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, {}); + + auto call = rewriter.create(op.getLoc(), TypeRange{i8PtrTy}, + funcPrefix, // funcPtr, + ValueRange{}); + + // erase the op + rewriter.eraseOp(op); + return success(); + } +}; + +template +struct Rdma4dOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + + // Get the module for function declarations + auto module = op->template getParentOfType(); + auto voidType = rewriter.getType(); + auto LLVMPtrType = rewriter.getType(); + auto i32Type = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = { + LLVMPtrType, // dest + LLVMPtrType, // src + i32Type, // elem_count + i32Type, // stride0 + i32Type, // iteration0 + i32Type, // stride1 + i32Type, // iteration1 + i32Type, // stride2 + i32Type, // iteration2 + i32Type // fmt + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, loc, + funcPrefix, voidType, argTypes); + + Value dstPtr = rewriter.create(loc, LLVMPtrType, + adaptor.getTarget()); + Value srcPtr = rewriter.create(loc, LLVMPtrType, + adaptor.getSource()); + Value elemCount = adaptor.getElemCount(); + Value strides[3] = {adaptor.getStride0(), adaptor.getStride1(), + adaptor.getStride2()}; + Value iterations[3] = {adaptor.getIteration0(), adaptor.getIteration1(), + adaptor.getIteration2()}; + Value fmt = + rewriter.create(loc, i32Type, op.getFmtAttr()); + + // Create the call to __Rdma4d/__Wdma4d + auto call = rewriter.replaceOpWithNewOp( + op, TypeRange{}, funcPrefix, + ValueRange{dstPtr, srcPtr, elemCount, strides[0], iterations[0], + strides[1], iterations[1], strides[2], iterations[2], fmt}); + + return success(); + } +}; + +template +struct Rdma1dOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + + // Get the module for function declarations + auto module = op->template getParentOfType(); + auto voidType = rewriter.getType(); + auto LLVMPtrType = rewriter.getType(); + auto i32Type = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = { + LLVMPtrType, // dest + LLVMPtrType, // src + i32Type, // elem_count + i32Type // fmt + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, loc, + funcPrefix, voidType, argTypes); + + Value dstPtr = rewriter.create(loc, LLVMPtrType, + adaptor.getTarget()); + Value srcPtr = rewriter.create(loc, LLVMPtrType, + adaptor.getSource()); + Value elemCount = adaptor.getElemCount(); + Value fmt = + rewriter.create(loc, i32Type, op.getFmtAttr()); + + // Create the call to __Rdma1d/__Wdma1d + auto call = rewriter.replaceOpWithNewOp( + op, TypeRange{}, funcPrefix, + ValueRange{dstPtr, srcPtr, elemCount, fmt}); + + return success(); + } +}; + +// Resolve wafer.remote_buffer to its destination address. +struct RemoteBufferOpConversion + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::RemoteBufferOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOp(op, adaptor.getOperands()[4]); + return success(); + } +}; + +// Convert wafer.remote_load to LLVM call to __Recv function +struct RemoteLoadOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::RemoteLoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto ctx = rewriter.getContext(); + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __Recv runtime function if not already declared + // Signature: + // void __Recv(int64_t chip_x, int64_t chip_y, int64_t die_id, + // int64_t tile_id, void* dst, + // uint32_t elem_bytes, uint64_t data_size) + auto i8PtrTy = LLVM::LLVMPointerType::get(ctx); + auto i64Ty = rewriter.getI64Type(); + auto i32Ty = rewriter.getI32Type(); + auto voidTy = LLVM::LLVMVoidType::get(ctx); + + // Types for function declaration + SmallVector argTypes = { + i64Ty, // remote_chip_id_x + i64Ty, // remote_chip_id_y + i64Ty, // remote_die_id + i64Ty, // remote_tile_id + i8PtrTy, // dst + i32Ty, // elem_bytes + i64Ty // data_size + }; + + // Declare the function with void return type + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, loc, + recvFuncName, voidTy, argTypes); + + // Get the operands and convert dst to i8* + Value chipX = adaptor.getOperands()[0]; + Value chipY = adaptor.getOperands()[1]; + Value dieId = adaptor.getOperands()[2]; + Value tileId = adaptor.getOperands()[3]; + Value dstAddr = adaptor.getOperands()[4]; + Value elemBytes = adaptor.getOperands()[5]; + Value dataSize = adaptor.getOperands()[6]; + + // Convert destination address (i64) directly to pointer. + Value dst = rewriter.create(loc, i8PtrTy, dstAddr); + + // Create the call to __Recv (void function, so empty TypeRange) + rewriter.create( + loc, TypeRange{}, recvFuncName, + ValueRange{chipX, chipY, dieId, tileId, dst, elemBytes, dataSize}); + + // wafer.remote_load has no results, just erase it + rewriter.eraseOp(op); + + return success(); + } +}; + +// Convert wafer.remote_store to LLVM call to __Send function +struct RemoteStoreOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::RemoteStoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto ctx = rewriter.getContext(); + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __Send runtime function if not already declared + // Signature: + // void __Send(int64_t chip_x, int64_t chip_y, int64_t die_id, + // int64_t tile_id, void* dst, void* src, + // uint32_t elem_bytes, uint64_t data_size) + auto i8PtrTy = LLVM::LLVMPointerType::get(ctx); + auto i64Ty = rewriter.getI64Type(); + auto i32Ty = rewriter.getI32Type(); + auto voidTy = LLVM::LLVMVoidType::get(ctx); + + // Types for function declaration + SmallVector argTypes = { + i64Ty, // remote_chip_id_x + i64Ty, // remote_chip_id_y + i64Ty, // remote_die_id + i64Ty, // remote_tile_id + i8PtrTy, // dst + i8PtrTy, // src + i32Ty, // elem_bytes + i64Ty // data_size + }; + + // Declare the function with void return type + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, loc, + sendFuncName, voidTy, argTypes); + + // Get the operands and convert dst/src to i8* + Value chipX = adaptor.getOperands()[0]; + Value chipY = adaptor.getOperands()[1]; + Value dieId = adaptor.getOperands()[2]; + Value tileId = adaptor.getOperands()[3]; + Value dstAddr = adaptor.getOperands()[4]; + Value src = adaptor.getOperands()[5]; + Value elemBytes = adaptor.getOperands()[6]; + Value dataSize = adaptor.getOperands()[7]; + + // Convert destination and source addresses (i64) directly to pointers. + Value dst = rewriter.create(loc, i8PtrTy, dstAddr); + src = rewriter.create(loc, i8PtrTy, src); + + // Create the call to __Send (void function, so empty TypeRange) + rewriter.create( + loc, TypeRange{}, sendFuncName, + ValueRange{chipX, chipY, dieId, tileId, dst, src, elemBytes, dataSize}); + + // wafer.remote_store has no results, just erase it + rewriter.eraseOp(op); + + return success(); + } +}; + +template +struct RdmaWdmaOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto ctx = rewriter.getContext(); + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the __Rdma runtime function if not already declared + auto i8PtrTy = LLVM::LLVMPointerType::get(ctx); + auto i32Ty = rewriter.getI32Type(); + auto i32PtrTy = LLVM::LLVMPointerType::get(ctx); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // src + i8PtrTy, // target + i32PtrTy, // src_shape array + i32PtrTy, // src_strides array + i32PtrTy, // dst_shape array + i32PtrTy, // dst_strides array + i32Ty, // rank + i32Ty, // elemBytes + i32Ty // fmt + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, loc, + funcPrefix, i8PtrTy, argTypes); + + // Get the operands + Value src = adaptor.getSource(); + src = rewriter.create(loc, i8PtrTy, src); + + Value target = adaptor.getTarget(); + target = rewriter.create(loc, i8PtrTy, target); + + // Create arrays for shapes and strides + + // Create arrays for shapes and strides + Value srcShapeArray = indexValueArrayToInt32ValueArray( + rewriter, loc, adaptor.getSrcShape(), op); + Value srcStridesArray = indexValueArrayToInt32ValueArray( + rewriter, loc, adaptor.getSrcStrides(), op); + Value dstShapeArray = indexValueArrayToInt32ValueArray( + rewriter, loc, adaptor.getDstShape(), op); + Value dstStridesArray = indexValueArrayToInt32ValueArray( + rewriter, loc, adaptor.getDstStrides(), op); + + // Handle rank attribute + Value rank = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(op.getRank())); + + // Handle elem byte attribute + Value elemBytes = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(op.getElemBytes())); + + // Handle format attribute + Value fmt = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(op.getFmt())); + + // Create the call to __Rdma + auto call = rewriter.create( + loc, TypeRange{i8PtrTy}, funcPrefix, + ValueRange{src, target, srcShapeArray, srcStridesArray, dstShapeArray, + dstStridesArray, rank, elemBytes, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +// Convert wafer.mask_move to LLVM call to __MaskMove function +struct MaskMoveOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::MaskMoveOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __MaskMove runtime function if not already declared + // Signature: void* __MaskMove(void* source, void* target, uint32_t + // elem_count, int32_t* masks, uint32_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i32PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // source + i8PtrTy, // target + i32Ty, // elem_count + i32PtrTy, // masks + i32Ty // fmt + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction( + module, rewriter, op.getLoc(), "__MaskMove", i8PtrTy, argTypes); + + // Get the operands + Value src = adaptor.getSource(); + + // Need to bitcast src to i8* + src = rewriter.create(op.getLoc(), i8PtrTy, src); + + Value target = adaptor.getTarget(); + + // Need to bitcast src to i8* + target = rewriter.create(op.getLoc(), i8PtrTy, target); + Value elemCount = adaptor.getElemCount(); + elemCount = castIndexToInt32(rewriter, op->getLoc(), elemCount); + + // Handle mask arrays + Value mask = adaptor.getMask(); + + // Need to bitcast src to i8* + mask = rewriter.create(op.getLoc(), i8PtrTy, mask); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getFmt())); + + // Create the call to __MaskMove + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__MaskMove", // funcPtr, + ArrayRef{src, target, elemCount, mask, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +template +struct TransformOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: + // __Transpose(uint64_t *src, uint64_t *dst, int32_t *src_shape, int32_t + // *dst_shape, uint16_t fmt) + + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32PtrTy, i32PtrTy, + i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value src = adaptor.getSource(); + // Need to bitcast src to i8* + src = rewriter.create(op.getLoc(), i8PtrTy, src); + Value dst = adaptor.getTarget(); + // Need to bitcast src to i8* + dst = rewriter.create(op.getLoc(), i8PtrTy, dst); + + // Convert shape attribute to Value + ArrayRef srcShape = adaptor.getSrcShape(); + ArrayRef dstShape = adaptor.getDstShape(); + + // Get shape llvm array + auto srcArray = + int32ArrayToInt32ValueArray(rewriter, op.getLoc(), srcShape, op); + auto dstArray = + int32ArrayToInt32ValueArray(rewriter, op.getLoc(), dstShape, op); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{src, dst, srcArray, dstArray, fmt}); + + // Erase the old op + rewriter.eraseOp(op); + + return success(); + } +}; + +struct GatherScatterOpConversion + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::GatherScatter op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op->getLoc(); + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __GatherScatter runtime function if not already declared + /* + void __GatherScatter(uint64_t *src, uint64_t *dst, uint32_t bytes, + uint32_t src_strideN, uint32_t src_strideH, + uint32_t src_strideW, uint32_t src_iterN, + uint32_t src_iterH, uint32_t src_iterW, + uint32_t dst_strideN, uint32_t dst_strideH, + uint32_t dst_strideW, uint32_t dst_iterN, + uint32_t dst_iterH, uint32_t dst_ite_W) + */ + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i32PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // src + i8PtrTy, // dst + i32Ty, // bytes + i32Ty, // src_StrideN + i32Ty, // src_StrideH + i32Ty, // src_StrideW + i32Ty, // dst_StrideN + i32Ty, // dst_StrideH + i32Ty, // dst_StrideW + i32Ty, // src_IterN + i32Ty, // src_IterH + i32Ty, // src_IterW + i32Ty, // dst_IterN + i32Ty, // dst_IterH + i32Ty // dst_IterW + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction( + module, rewriter, loc, "__GatherScatter", i8PtrTy, argTypes); + + // Get the operands + Value src = adaptor.getSource(); + src = rewriter.create(loc, i8PtrTy, src); + + // Get the operands + Value dst = adaptor.getTarget(); + dst = rewriter.create(loc, i8PtrTy, dst); + + // Get bytes + auto bytes = + rewriter.create(loc, i32Ty, adaptor.getBytes()); + + // Get strides + auto srcStrideN = + rewriter.create(loc, i32Ty, adaptor.getSrcStrideN()); + auto srcStrideH = + rewriter.create(loc, i32Ty, adaptor.getSrcStrideH()); + auto srcStrideW = + rewriter.create(loc, i32Ty, adaptor.getSrcStrideW()); + auto dstStrideN = + rewriter.create(loc, i32Ty, adaptor.getDstStrideN()); + auto dstStrideH = + rewriter.create(loc, i32Ty, adaptor.getDstStrideH()); + auto dstStrideW = + rewriter.create(loc, i32Ty, adaptor.getDstStrideW()); + + // Get iterator + auto srcIterN = + rewriter.create(loc, i32Ty, adaptor.getSrcIterN()); + auto srcIterH = + rewriter.create(loc, i32Ty, adaptor.getSrcIterH()); + auto srcIterW = + rewriter.create(loc, i32Ty, adaptor.getSrcIterW()); + auto dstIterN = + rewriter.create(loc, i32Ty, adaptor.getDstIterN()); + auto dstIterH = + rewriter.create(loc, i32Ty, adaptor.getDstIterH()); + auto dstIterW = + rewriter.create(loc, i32Ty, adaptor.getDstIterW()); + + // Create the call to __GatherScatter + auto call = rewriter.create( + loc, TypeRange{i8PtrTy}, "__GatherScatter", // funcPtr, + ValueRange{src, dst, bytes, srcStrideN, srcStrideH, srcStrideW, + srcIterN, srcIterH, srcIterW, dstStrideN, dstStrideH, + dstStrideW, dstIterN, dstIterH, dstIterW}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +template +struct ArgMinMaxOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: + + // __ArgMinMax(uint64_t *src, uint64_t *dst0, uint64_t *dst1, + // uint32_t elem_count, uint16_t fmt) + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i8PtrTy, i32Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value src = adaptor.getSrc(); + // Need to bitcast src to i8* + src = rewriter.create(op.getLoc(), i8PtrTy, src); + + // Convert results + Value value = adaptor.getValue(); + Value index = adaptor.getIndex(); + // Need to bitcast `value` and `index` to i8* + value = rewriter.create(op.getLoc(), i8PtrTy, value); + index = rewriter.create(op.getLoc(), i8PtrTy, index); + + // Get elem_count operand, convert Index to I32 + Value elemCount = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getElemCount())); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{src, value, index, elemCount, fmt}); + + // Erase the old op + rewriter.eraseOp(op); + + return success(); + } +}; + +// Convert wafer.binary op to LLVM call +template +struct ReduceOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: + // __ReduceSum(uint64_t *src, uint64_t *dst, uint32_t dim, uint16_t src_n, + // uint16_t src_h, uint16_t src_w, uint16_t src_c, uint16_t fmt) + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32Ty, i16Ty, + i16Ty, i16Ty, i16Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value src = adaptor.getSrc(); + // Need to bitcast src to i8* + src = rewriter.create(op.getLoc(), i8PtrTy, src); + Value srcB = adaptor.getSrc(); + Value dst = adaptor.getDst(); + // Need to bitcast src to i8* + dst = rewriter.create(op.getLoc(), i8PtrTy, dst); + + // Convert dim attribute to Value + Value dim = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getDim())); + + // Convert shape attribute to Value + Value shape_n = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[0]); + Value shape_h = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[1]); + Value shape_w = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[2]); + Value shape_c = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[3]); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{src, dst, dim, shape_n, shape_h, shape_w, shape_c, + fmt}); + + // Erase the old op + rewriter.eraseOp(op); + + return success(); + } +}; + +// Convert wafer.elementwise op to LLVM call +template +struct ElementWiseOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + // using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void* __Add(void* a, void* b, void* out, uint32_t elem_count, + // uint32_t rnd_mode, uint32_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i8PtrTy, + + i32Ty, i32Ty, i32Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value srcA = adaptor.getInput0(); + // Need to bitcast src to i8* + srcA = rewriter.create(op.getLoc(), i8PtrTy, srcA); + Value srcB = adaptor.getInput1(); + // Need to bitcast src to i8* + srcB = rewriter.create(op.getLoc(), i8PtrTy, srcB); + Value out = adaptor.getOut(); + // Need to bitcast src to i8* + out = rewriter.create(op.getLoc(), i8PtrTy, out); + + // Get elem_count operand, convert Index to I32 + Value elemCount = op.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle round attribute + Value rnd_mode = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getRndMode())); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{srcA, srcB, out, elemCount, rnd_mode, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +template +struct UnaryOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void* __Abs(void* src, void* dst, uint32_t elem_count, + // uint16_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value input = adaptor.getInput(); + // Need to bitcast src to i8* + input = rewriter.create(op.getLoc(), i8PtrTy, input); + Value out = adaptor.getOut(); + // Need to bitcast out to i8* + out = rewriter.create(op.getLoc(), i8PtrTy, out); + + // Get elem_count operand, convert Index to I32 + Value elemCount = op.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{input, out, elemCount, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +// FIXME: Use trait to refactor the BinaryVSOpConversion and +// ElementWiseOpConversion +template +struct BinaryVSOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void* __Add(void* a, void* b, void* out, uint32_t elem_count, + // uint32_t rnd_mode, uint32_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i32Ty, i8PtrTy, + i32Ty, i32Ty, i32Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value srcA = adaptor.getInput0(); + // Need to bitcast src to i8* + srcA = rewriter.create(op.getLoc(), i8PtrTy, srcA); + + Value srcB = adaptor.getValue(); + + Value out = adaptor.getOut(); + // Need to bitcast src to i8* + out = rewriter.create(op.getLoc(), i8PtrTy, out); + + // Get elem_count operand, convert Index to I32 + Value elemCount = op.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle round attribute + Value rnd_mode = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getRndMode())); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{srcA, srcB, out, elemCount, rnd_mode, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +template +struct BinaryLogicVVOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void* __XorVV(void* a, void* b, void* out, uint32_t + // elem_count, uint32_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // src0_addr + i8PtrTy, // src1_addr + i8PtrTy, // dst_addr + i32Ty, // elem_count + i32Ty // fmt + }; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value srcA = adaptor.getInput0(); + // Need to bitcast src to i8* + srcA = rewriter.create(op.getLoc(), i8PtrTy, srcA); + Value srcB = adaptor.getInput1(); + // Need to bitcast src to i8* + srcB = rewriter.create(op.getLoc(), i8PtrTy, srcB); + Value out = adaptor.getOut(); + // Need to bitcast src to i8* + out = rewriter.create(op.getLoc(), i8PtrTy, out); + + // Get elem_count operand, convert Index to I32 + Value elemCount = op.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{srcA, srcB, out, elemCount, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +template +struct UnaryBoolLogicVOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void* __BoolNotV(void* src, void* dst, uint32_t elem_count); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // src_addr + i8PtrTy, // dst_addr + i32Ty // elem_count + }; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value src = adaptor.getInput(); + // Need to bitcast src to i8* + src = rewriter.create(op.getLoc(), i8PtrTy, src); + + Value out = adaptor.getOut(); + // Need to bitcast dest to i8* + out = rewriter.create(op.getLoc(), i8PtrTy, out); + + // Get elem_count operand, convert Index to I32 + Value elemCount = op.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{src, out, elemCount}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +template +struct BinaryBoolLogicVOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void* __BoolAndV(void* a, void* b, void* out, uint32_t + // elem_count); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // src0_addr + i8PtrTy, // src1_addr + i8PtrTy, // dst_addr + i32Ty // elem_count + }; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value srcA = adaptor.getInput0(); + // Need to bitcast src to i8* + srcA = rewriter.create(op.getLoc(), i8PtrTy, srcA); + Value srcB = adaptor.getInput1(); + // Need to bitcast src to i8* + srcB = rewriter.create(op.getLoc(), i8PtrTy, srcB); + Value out = adaptor.getOut(); + // Need to bitcast src to i8* + out = rewriter.create(op.getLoc(), i8PtrTy, out); + + // Get elem_count operand, convert Index to I32 + Value elemCount = op.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{srcA, srcB, out, elemCount}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +template +struct RelationVVOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename RelationVVOp::Adaptor; + // using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(RelationVVOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void __BoolLessEqualVV(uint64_t *src0, uint64_t *src1, + // uint64_t *dst, uint32_t elem_count, uint16_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i8PtrTy, i32Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value srcA = adaptor.getInput0(); + // Need to bitcast src to i8* + srcA = rewriter.create(op.getLoc(), i8PtrTy, srcA); + Value srcB = adaptor.getInput1(); + // Need to bitcast src to i8* + srcB = rewriter.create(op.getLoc(), i8PtrTy, srcB); + Value out = adaptor.getOut(); + // Need to bitcast src to i8* + out = rewriter.create(op.getLoc(), i8PtrTy, out); + + // Get elem_count operand + Value elemCount = op.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{srcA, srcB, out, elemCount, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +// FIXME: Use trait to refactor the RelationVSOpConversion and +// ElementWiseOpConversion +template +struct RelationVSOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename WaferOpT::Adaptor; + + LogicalResult + matchAndRewrite(WaferOpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void __BoolEqualVS(uint64_t *src0, uint32_t src1, uint64_t + // *dst,uint32_t elem_count, uint16_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i32Ty, i8PtrTy, i32Ty, i32Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value srcA = adaptor.getInput0(); + // Need to bitcast src to i8* + srcA = rewriter.create(op.getLoc(), i8PtrTy, srcA); + + Value srcB = adaptor.getValue(); + + Value out = adaptor.getOut(); + // Need to bitcast src to i8* + out = rewriter.create(op.getLoc(), i8PtrTy, out); + + // Get elem_count operand, convert Index to I32 + Value elemCount = op.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{srcA, srcB, out, elemCount, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +// Convert wafer.ZeroPointConvertOp op to LLVM +template +struct ZeroPointConvertOpConversion + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename ZeroPointConvertOp::Adaptor; + + LogicalResult + matchAndRewrite(ZeroPointConvertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void __INT8_FP32(uint64_t *src, uint64_t *dst, uint32_t + // zero_point, uint32_t elem_count); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32Ty, i32Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value input = adaptor.getSrc(); + Value output = adaptor.getDst(); + Value zeroPoint = rewriter.create( + op.getLoc(), i32Ty, adaptor.getZeroPointAttr()); + Value elemCount = rewriter.create( + op.getLoc(), i32Ty, adaptor.getElemCountAttr()); + + // Bitcast all pointers to i8* + input = rewriter.create(op.getLoc(), i8PtrTy, input); + output = rewriter.create(op.getLoc(), i8PtrTy, output); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{input, output, zeroPoint, elemCount}); + + rewriter.eraseOp(op); + return success(); + } +}; + +// Convert wafer.NormalConvertOp op to LLVM +template +struct NormalConvertOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename NormalConvertOp::Adaptor; + + LogicalResult + matchAndRewrite(NormalConvertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void __FP16_FP32(uint64_t *src, uint64_t *dst, uint32_t + // elem_count); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value input = adaptor.getInput(); + Value output = adaptor.getOutput(); + Value elemCount = adaptor.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Bitcast all pointers to i8* + input = rewriter.create(op.getLoc(), i8PtrTy, input); + output = rewriter.create(op.getLoc(), i8PtrTy, output); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{input, output, elemCount}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +// Convert wafer.RoundConvertOp op to LLVM +template +struct RoundConvertOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename RoundConvertOp::Adaptor; + + LogicalResult + matchAndRewrite(RoundConvertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void __INT16_FP32(uint64_t *src, uint64_t *dst, uint32_t + // elem_count, RND_MODE round); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Convert operands + Value input = adaptor.getInput(); + Value output = adaptor.getOutput(); + Value elemCount = adaptor.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + Value rnd_mode = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getRndMode())); + + // Bitcast all pointers to i8* + input = rewriter.create(op.getLoc(), i8PtrTy, input); + output = rewriter.create(op.getLoc(), i8PtrTy, output); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{input, output, elemCount, rnd_mode}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +template +struct MXFPScaleOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename MXFPScaleOp::Adaptor; + + LogicalResult + matchAndRewrite(MXFPScaleOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // src + i8PtrTy, // scale + i8PtrTy, // dst + i32Ty, // elem_count + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + funcPrefix, i8PtrTy, argTypes); + + // Get the operands + Value src = adaptor.getSrc(); + Value scale = adaptor.getScale(); + Value dst = adaptor.getDst(); + + // Need to bitcast src to i8* + src = rewriter.create(op.getLoc(), i8PtrTy, src); + scale = rewriter.create(op.getLoc(), i8PtrTy, scale); + dst = rewriter.create(op.getLoc(), i8PtrTy, dst); + + Value elemCount = rewriter.create( + op.getLoc(), i32Ty, adaptor.getElemCountAttr()); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, funcPrefix, // funcPtr, + ArrayRef{src, scale, dst, elemCount}); + + op->erase(); + return success(); + } +}; + +struct BitToFPOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = wafer::Bit2FpOp::Adaptor; + + LogicalResult + matchAndRewrite(wafer::Bit2FpOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: void __Bit2Fp(uint64_t *src, uint64_t *target, uint32_t + // elem_count, uint16_t fmt) + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + "__Bit2Fp", i8PtrTy, argTypes); + + // Convert operands + Value input = adaptor.getSrc(); + Value output = adaptor.getTarget(); + Value elemCount = adaptor.getElemCount(); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Bitcast all pointers to i8* + input = rewriter.create(op.getLoc(), i8PtrTy, input); + output = rewriter.create(op.getLoc(), i8PtrTy, output); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__Bit2Fp", // funcPtr, + ArrayRef{input, output, elemCount, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +// Convert wafer.channel_norm op +struct ChannelNormOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename wafer::ChannelNormOp::Adaptor; + + LogicalResult + matchAndRewrite(wafer::ChannelNormOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: + // __ChannelNorm(uint64_t *src, uint64_t *dst, uint16_t n, + // uint16_t h, uint16_t w, uint16_t c, uint16_t c0, uint16_t + // dtype_size) + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i16Ty, i16Ty, + i16Ty, i16Ty, i16Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction( + module, rewriter, op.getLoc(), "__ChannelNorm", i8PtrTy, argTypes); + + // Convert operands + Value src = adaptor.getSrc(); + // Need to bitcast src to i8* + src = rewriter.create(op.getLoc(), i8PtrTy, src); + Value dst = adaptor.getDst(); + // Need to bitcast dst to i8* + dst = rewriter.create(op.getLoc(), i8PtrTy, dst); + + // Convert shape attribute to Value + Value shape_n = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[0]); + Value shape_h = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[1]); + Value shape_w = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[2]); + Value shape_c = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[3]); + + // Convert c0_align attribute to Value + Value c0Align = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getC0Align())); + + // Convert dtype_size attribute to Value + Value dtypeSize = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI32IntegerAttr(op.getDtypeSize())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__ChannelNorm", // funcPtr, + ArrayRef{src, dst, shape_n, shape_h, shape_w, shape_c, c0Align, + dtypeSize}); + + // Erase the old op + rewriter.eraseOp(op); + + return success(); + } +}; + +struct DechannelNormOpConversion + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpAdaptor = typename wafer::DechannelNormOp::Adaptor; + + LogicalResult + matchAndRewrite(wafer::DechannelNormOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->template getParentOfType(); + + // Declare the runtime function if not already declared + // Signature: + // __DechannelNorm(uint64_t *src, uint64_t *dst, uint16_t n, + // uint16_t h, uint16_t w, uint16_t c, uint16_t c0, uint16_t + // dtype_size) + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i16Ty, i16Ty, + i16Ty, i16Ty, i16Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction( + module, rewriter, op.getLoc(), "__DechannelNorm", i8PtrTy, argTypes); + + // Convert operands + Value src = adaptor.getSrc(); + // Need to bitcast src to i8* + src = rewriter.create(op.getLoc(), i8PtrTy, src); + Value dst = adaptor.getDst(); + // Need to bitcast dst to i8* + dst = rewriter.create(op.getLoc(), i8PtrTy, dst); + + // Convert shape attribute to Value + Value shape_n = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[0]); + Value shape_h = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[1]); + Value shape_w = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[2]); + Value shape_c = + rewriter.create(op.getLoc(), i16Ty, op.getShape()[3]); + + // Convert c0_align attribute to Value + Value c0Align = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getC0Align())); + + // Convert dtype_size attribute to Value + Value dtypeSize = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI32IntegerAttr(op.getDtypeSize())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__DechannelNorm", // funcPtr, + ArrayRef{src, dst, shape_n, shape_h, shape_w, shape_c, c0Align, + dtypeSize}); + + // Erase the old op + rewriter.eraseOp(op); + + return success(); + } +}; + +// Convert wafer.gemm to LLVM call to __Gemm function +struct GemmOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::GemmOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __Gemm runtime function if not already declared + // Signature: void __Gemm(int64_t* srcA, int64_t *srcB, int64_t * srcBias, + // int64_t *dst, int32_t *dims, bool enPsum, int64_t *psum, bool enTransA, + // bool enTransB, int64_t batchSizeA, int64_t batchSizeB, bool enLeakyRelu, + // bool enBias,bool enNegScale, int64_t *negScale, bool enPosScale, int64_t + // *posScale, int64_t srcFmt, int64_t dstFmt) + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i64Ty = rewriter.getI64Type(); + auto i32PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i64PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i1Ty = rewriter.getI1Type(); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // srcA + i8PtrTy, // srcB + i8PtrTy, // srcBias + i8PtrTy, // dst + i32PtrTy, // dims + i1Ty, // enPsum + i8PtrTy, // psum + i1Ty, // enTransA + i1Ty, // enTransB + i32Ty, // batchSizeA + i32Ty, // batchSizeB + i32Ty, // reluMode + i1Ty, // enBias + i1Ty, // enNegScale + i8PtrTy, // negScale + i1Ty, // enPosScale + i8PtrTy, // posScale + i32Ty, // srcFmt + i32Ty // dstFmt + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + "__Gemm", i8PtrTy, argTypes); + + // Convert operands + Value srcA = adaptor.getSrcA(); + Value srcB = adaptor.getSrcB(); + Value srcBias = adaptor.getSrcBias(); + Value dst = adaptor.getDst(); + + Value psumAddr = adaptor.getPsumAddr(); + Value srcNegScale = adaptor.getSrcNegScale(); + Value srcPosScale = adaptor.getSrcPosScale(); + + // Bitcast all pointers to i8* + srcA = rewriter.create(op.getLoc(), i8PtrTy, srcA); + srcB = rewriter.create(op.getLoc(), i8PtrTy, srcB); + srcBias = rewriter.create(op.getLoc(), i8PtrTy, srcBias); + dst = rewriter.create(op.getLoc(), i8PtrTy, dst); + psumAddr = + rewriter.create(op.getLoc(), i8PtrTy, psumAddr); + srcNegScale = + rewriter.create(op.getLoc(), i8PtrTy, srcNegScale); + srcPosScale = + rewriter.create(op.getLoc(), i8PtrTy, srcPosScale); + + // Handle dims array - need to convert from attribute to runtime array + auto dimsAttr = op.getDims(); + SmallVector dimsValues; + for (auto dimAttr : dimsAttr) + dimsValues.push_back(cast(dimAttr).getInt()); + + // Allocate memory for the dims array + Value rank = rewriter.create( + op.getLoc(), i64Ty, rewriter.getI64IntegerAttr(dimsValues.size())); + + auto dimsArrayI32Ptr = + int32ArrayToInt32ValueArray(rewriter, op->getLoc(), dimsValues, op); + + // Convert boolean attributes + Value transA = rewriter.create( + op.getLoc(), i1Ty, rewriter.getBoolAttr(op.getTransSrcA())); + Value transB = rewriter.create( + op.getLoc(), i1Ty, rewriter.getBoolAttr(op.getTransSrcB())); + Value enPSum = rewriter.create( + op.getLoc(), i1Ty, rewriter.getBoolAttr(op.getEnPsum())); + Value reluMode = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getReluMode())); + Value enBias = rewriter.create( + op.getLoc(), i1Ty, rewriter.getBoolAttr(op.getEnBias())); + Value enNegScale = rewriter.create( + op.getLoc(), i1Ty, rewriter.getBoolAttr(op.getEnNegScale())); + Value enPosScale = rewriter.create( + op.getLoc(), i1Ty, rewriter.getBoolAttr(op.getEnPosScale())); + + // Convert integer attributes + Value batchA = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getBatchSrcA())); + Value batchB = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getBatchSrcB())); + Value srcFmt = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getSrcFmt())); + Value dstFmt = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getDstFmt())); + + // Create the call to __Gemm + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__Gemm", // funcPtr, + ArrayRef{srcA, srcB, srcBias, dst, dimsArrayI32Ptr, enPSum, + psumAddr, transA, transB, batchA, batchB, reluMode, + enBias, enNegScale, srcNegScale, enPosScale, + srcPosScale, srcFmt, dstFmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +struct SigmoidOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::Sigmoid op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __Sigmoid runtime function if not already declared + // Signature: void __Sigmoid(int64_t* src, int64_t *dst, + // uint32_t elem_count, uint16_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + "__Sigmoid", i8PtrTy, argTypes); + + // Convert operands + Value input = adaptor.getInput(); + Value output = adaptor.getOut(); + Value elemCount = adaptor.getElemCount(); + + // Bitcast all pointers to i8* + input = rewriter.create(op.getLoc(), i8PtrTy, input); + output = rewriter.create(op.getLoc(), i8PtrTy, output); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__Sigmoid", // funcPtr, + ArrayRef{input, output, elemCount, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +struct GeluNoneOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::GeluNone op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __GeluNone runtime function if not already declared + // Signature: void __GeluNone(int64_t* src, int64_t *dst, + // uint32_t elem_count, uint16_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i32Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction( + module, rewriter, op.getLoc(), "__GeluNone", i8PtrTy, argTypes); + + // Convert operands + Value input = adaptor.getInput(); + Value output = adaptor.getOut(); + Value elemCount = adaptor.getElemCount(); + + // Bitcast all pointers to i8* + input = rewriter.create(op.getLoc(), i8PtrTy, input); + output = rewriter.create(op.getLoc(), i8PtrTy, output); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__GeluNone", // funcPtr, + ArrayRef{input, output, elemCount, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +struct GeluTanhOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::GeluTanh op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __GeluTanh runtime function if not already declared + // Signature: void __GeluTanh(int64_t* src, int64_t *imm, int64_t *dst, + // uint32_t elem_count, uint16_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = {i8PtrTy, i8PtrTy, i8PtrTy, i32Ty, i16Ty}; + + Value funcPtr = triton::declareWaferRuntimeFunction( + module, rewriter, op.getLoc(), "__GeluTanh", i8PtrTy, argTypes); + + // Convert operands + Value input = adaptor.getInput(); + Value imm = adaptor.getBuffer(); + Value output = adaptor.getOut(); + Value elemCount = adaptor.getElemCount(); + + // Bitcast all pointers to i8* + input = rewriter.create(op.getLoc(), i8PtrTy, input); + imm = rewriter.create(op.getLoc(), i8PtrTy, imm); + output = rewriter.create(op.getLoc(), i8PtrTy, output); + elemCount = castIndexToInt32(rewriter, op.getLoc(), elemCount); + + // Handle format attribute + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + // Create the call + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__GeluTanh", // funcPtr, + ArrayRef{input, imm, output, elemCount, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +// Convert wafer.memset to LLVM call to __Memset function +struct MemsetOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(wafer::MemsetOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __Memset runtime function if not already declared + // Signature: void* __Memset(void* dst, int64_t value, uint32_t rank, + // int32_t* strides, int32_t* iterations, uint16_t fmt); + auto i8PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i32Ty = rewriter.getI32Type(); + auto i32PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext()); + auto i16Ty = rewriter.getI16Type(); + + // Types for function declaration + SmallVector argTypes = { + i8PtrTy, // Spm addr + i32Ty, // value + i32PtrTy, // src_shape array + i32PtrTy, // src_strides array + i32Ty, // rank + i16Ty // fmt + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + "__Memset", i8PtrTy, argTypes); + + // Get operands + Value dst = adaptor.getTarget(); + dst = rewriter.create(op.getLoc(), i8PtrTy, dst); + + Value value = adaptor.getValue(); + + // Handle strides and iterations arrays + // Create arrays for shapes and strides + Value dstShapeArray = indexValueArrayToInt32ValueArray( + rewriter, loc, adaptor.getDstShape(), op); + Value dstStridesArray = indexValueArrayToInt32ValueArray( + rewriter, loc, adaptor.getDstStrides(), op); + + // Convert fmt attribute to Value + Value fmt = rewriter.create( + op.getLoc(), i16Ty, rewriter.getI16IntegerAttr(op.getFmt())); + + Value rank = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(op.getRank())); + + // Create the call to __Memset + auto call = rewriter.create( + op.getLoc(), i8PtrTy, "__Memset", // funcPtr, + ArrayRef{dst, value, dstShapeArray, dstStridesArray, rank, fmt}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +// Convert tt.get_program_id to LLVM call to __get_pid function +// Think this as Wafer special action. May can separate to a single pass or use +// wafer.get_program_id op +struct GetProgramIDConversion + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + static uint32_t constexpr LAUNCH_GRID_RANK = + mlir::triton::getMaxEnumValForProgramIDDim() + 1; + + LogicalResult + matchAndRewrite(triton::GetProgramIdOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // Get the module for function declarations + auto module = op->getParentOfType(); + + // Declare the __Memset runtime function if not already declared + // Signature: uint32_t __get_pid(uint32_t); + auto i32Ty = rewriter.getI32Type(); + + // Types for function declaration + SmallVector argTypes = { + i32Ty, // x: 0/y: 1/z: 2, + }; + + // Declare the function + Value funcPtr = triton::declareWaferRuntimeFunction(module, rewriter, op.getLoc(), + "__get_pid", i32Ty, argTypes); + + // Get operands + auto axis = (uint32_t)op.getAxis(); + + assert(axis < LAUNCH_GRID_RANK && "program_id expects " + "axis to be either 0, " + "1, or 2"); + + // Convert fmt attribute to Value + Value src = rewriter.create( + op.getLoc(), i32Ty, rewriter.getI32IntegerAttr(axis)); + + // Create the call to __Memset + auto call = rewriter.create(op.getLoc(), i32Ty, + "__get_pid", // funcPtr, + ArrayRef{src}); + + // Replace the op with the result of the call + rewriter.replaceOp(op, call.getResult()); + + return success(); + } +}; + +struct AssertConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(mk::AssertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op->getLoc(); + + auto i32Ty = rewriter.getI32Type(); + auto context = rewriter.getContext(); + ModuleOp parentModule = op->getParentOfType(); + auto assertRef = getOrInsertAssert(rewriter, parentModule); + auto message = op.getMessage(); + + StringRef file = "unknown"; + int line = 0; + int col = 0; + if (auto fileLineColLoc = dyn_cast(loc)) { + file = fileLineColLoc.getFilename(); + line = fileLineColLoc.getLine(); + col = fileLineColLoc.getColumn(); + } + + SmallVector argTypes = { + i32Ty, // x: 0/y: 1/z: 2, + }; + Value funcPtr = triton::declareWaferRuntimeFunction(parentModule, rewriter, loc, + "__get_pid", i32Ty, argTypes); + auto xDim = rewriter.create(loc, i32Ty, 0); + auto yDim = rewriter.create(loc, i32Ty, 1); + auto zDim = rewriter.create(loc, i32Ty, 2); + auto pidX = rewriter + .create(loc, i32Ty, + "__get_pid", // funcPtr, + ArrayRef{xDim}) + ->getResult(0); + auto pidY = rewriter + .create(loc, i32Ty, + "__get_pid", // funcPtr, + ArrayRef{yDim}) + ->getResult(0); + auto pidZ = rewriter + .create(loc, i32Ty, + "__get_pid", // funcPtr, + ArrayRef{zDim}) + ->getResult(0); + + llvm::SmallString<64> messageString(message), fileString(file); + messageString.push_back('\0'); + fileString.push_back('\0'); + Value messageStringVal = + LLVM::addStringToModule(loc, rewriter, "assertMessage_", messageString); + Value fileStringVal = + LLVM::addStringToModule(loc, rewriter, "assertFile_", fileString); + + auto lineValue = rewriter.create(loc, i32Ty, line); + auto colValue = rewriter.create(loc, i32Ty, col); + + rewriter.create(loc, getAssertType(context), assertRef, + ValueRange{messageStringVal, fileStringVal, + lineValue, colValue, pidX, pidY, + pidZ}); + rewriter.eraseOp(op); + return success(); + } + +private: + static LLVM::LLVMFunctionType getAssertType(MLIRContext *context) { + auto llvmPtr = LLVM::LLVMPointerType::get(context); + // Match CRT's void __Assert(const char *, ...). + return LLVM::LLVMFunctionType::get(LLVM::LLVMVoidType::get(context), llvmPtr, + true); + } + + static FlatSymbolRefAttr getOrInsertAssert(PatternRewriter &rewriter, + ModuleOp module, + StringRef funcName = "__Assert") { + auto *context = module.getContext(); + if (module.lookupSymbol(funcName)) + return SymbolRefAttr::get(context, funcName); + + PatternRewriter::InsertionGuard insertGuard(rewriter); + rewriter.setInsertionPointToStart(module.getBody()); + rewriter.create(module.getLoc(), funcName, + getAssertType(context)); + return SymbolRefAttr::get(context, funcName); + } +}; + +// The conversion pass +class WaferToLLVMPass : public WaferToLLVMBase { +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + ModuleOp module = getOperation(); + MLIRContext *context = &getContext(); + ConversionTarget target(*context); + + // Setup LLVM lowering options object which should live across the call to + // applyFull/PartialConversion. + LowerToLLVMOptions options(context); + options.useBarePtrCallConv = false; + + // Setup conversion target + target.addLegalDialect(); + // Handle the wafer op to llvm.call and support kcore load/store op's spm + // offset + target.addIllegalDialect(); + + // Setup rewrite patterns + RewritePatternSet patterns(context); + + // NOTE: LLVMTypeConverter should be enough for MLIR core dialects. + LLVMTypeConverter llvmTypeConverter(context, options); + + // Add the Wafer to LLVM conversion patterns + // clang-format off + patterns.add, + ZeroPointConvertOpConversion, + ZeroPointConvertOpConversion, + ZeroPointConvertOpConversion, + /* INT16 */ + NormalConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + /* INT32 */ + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + /* BF16 */ + NormalConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + /* FP16 */ + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + /* FP32 */ + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, // NOTE: No op used + /* TF32 */ + RoundConvertOpConversion, + RoundConvertOpConversion, + RoundConvertOpConversion, + NormalConvertOpConversion, + RoundConvertOpConversion, + NormalConvertOpConversion, + /* MXFP */ + NormalConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + NormalConvertOpConversion, + MXFPScaleOpConversion, + MXFPScaleOpConversion, + ArgMinMaxOpConversion, + ArgMinMaxOpConversion, + ReduceOpConversion, + ReduceOpConversion, + ReduceOpConversion, + ReduceOpConversion, + ElementWiseOpConversion, + ElementWiseOpConversion, + ElementWiseOpConversion, + ElementWiseOpConversion, + ElementWiseOpConversion, + ElementWiseOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + UnaryOpConversion, + BinaryVSOpConversion, + BinaryVSOpConversion, + BinaryVSOpConversion, + BinaryVSOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVVOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + RelationVSOpConversion, + BinaryLogicVVOpConversion, + BinaryLogicVVOpConversion, + BinaryLogicVVOpConversion, + UnaryBoolLogicVOpConversion, + BinaryBoolLogicVOpConversion, + BinaryBoolLogicVOpConversion, + BinaryBoolLogicVOpConversion, + Rdma4dOpConversion, + Rdma4dOpConversion, + Rdma1dOpConversion, + Rdma1dOpConversion, + RdmaWdmaOpConversion, + RdmaWdmaOpConversion, + UnaryOpConversion, + TransformOpConversion, + TransformOpConversion, + TransformOpConversion, + AtomicBarrierOpConversion, + AtomicBarrierOpConversion, + MaskMoveOpConversion, + GatherScatterOpConversion, + BitToFPOpConversion, + ChannelNormOpConversion, // NOTE: No op used + DechannelNormOpConversion, // NOTE: No op used + GemmOpConversion, + SigmoidOpConversion, + GeluNoneOpConversion, + GeluTanhOpConversion, + MemsetOpConversion, + GetProgramIDConversion, + BarrierConversion, + RemoteStoreOpConversion, + RemoteLoadOpConversion, + RandGenOpConversion, + AssertConversion>( + context); + // clang-format on + + // Add call op conversion + populateCallOpTypeConversionPattern(patterns, llvmTypeConverter); + + // Add return op conversion + populateReturnOpTypeConversionPattern(patterns, llvmTypeConverter); + + // Apply the conversion + if (failed(applyPartialConversion(module, target, std::move(patterns)))) + signalPassFailure(); + } +}; + +} // namespace + +std::unique_ptr> triton::createWaferToLLVMPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Conversion/WaferToLLVM/WaferToLLVMPass.cpp b/third_party/wafer/lib/Conversion/WaferToLLVM/WaferToLLVMPass.cpp new file mode 100755 index 00000000..edaab9c8 --- /dev/null +++ b/third_party/wafer/lib/Conversion/WaferToLLVM/WaferToLLVMPass.cpp @@ -0,0 +1,78 @@ +//===--------------------- WaferToLLVMPass.cpp -----------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Conversion/LLVMCommon/TypeConverter.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Transforms/DialectConversion.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "wafer/Conversion/WaferToLLVM/WaferToLLVM.h" +#include "wafer/Dialect/IR/WaferDialect.h" +#include "llvm/Support/Debug.h" +#include +#include +#include + +#define DEBUG_TYPE "wafer-to-llvm" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "wafer/Conversion/WaferToLLVM/Passes.h.inc" + +namespace { + +class WaferToLLVMPass : public WaferToLLVMBase { +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + ModuleOp module = getOperation(); + MLIRContext *context = &getContext(); + ConversionTarget target(*context); + + // Setup LLVM lowering options object which should live across the call to + // applyFull/PartialConversion. + LowerToLLVMOptions options(context); + options.useBarePtrCallConv = false; + + // Setup conversion target + target.addLegalDialect(); + target.addIllegalDialect(); + + // Setup rewrite patterns + RewritePatternSet patterns(context); + + // NOTE: LLVMTypeConverter should be enough for MLIR core dialects. + TensorToLLVMTypeConverter converter(context, options); + + triton::populateWaferToLLVMConversionPatterns(patterns, target, converter); + + // Apply the conversion + if (failed(applyPartialConversion(module, target, std::move(patterns)))) + signalPassFailure(); + } +}; + +} // namespace + +std::unique_ptr> triton::createWaferToLLVMPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Dialect/Address/CMakeLists.txt b/third_party/wafer/lib/Dialect/Address/CMakeLists.txt new file mode 100755 index 00000000..9f57627c --- /dev/null +++ b/third_party/wafer/lib/Dialect/Address/CMakeLists.txt @@ -0,0 +1,2 @@ +add_subdirectory(IR) +add_subdirectory(Transforms) diff --git a/third_party/wafer/lib/Dialect/Address/IR/AddressDialect.cpp b/third_party/wafer/lib/Dialect/Address/IR/AddressDialect.cpp new file mode 100755 index 00000000..3bd29d51 --- /dev/null +++ b/third_party/wafer/lib/Dialect/Address/IR/AddressDialect.cpp @@ -0,0 +1,303 @@ +//===- AddressDialect.cpp - Address dialect ---------------------*- C++ -*-===// +// +// This file is licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file implements the Address dialect. +// +//===----------------------------------------------------------------------===// + +#include "Address/Dialect/IR/AddressDialect.h" +#include "mlir/IR/DialectImplementation.h" +#include "mlir/IR/PatternMatch.h" +#include "llvm/ADT/SmallString.h" +#include "llvm/ADT/TypeSwitch.h" + +using namespace mlir; +using namespace mlir::addr; + +//===----------------------------------------------------------------------===// +// Address dialect +//===----------------------------------------------------------------------===// + +void AddressDialect::initialize() { + addOperations< +#define GET_OP_LIST +#include "Address/Dialect/IR/AddressOps.cpp.inc" + >(); + registerTypes(); +} + +void AddressDialect::registerTypes() { + addTypes< +#define GET_TYPEDEF_LIST +#include "Address/Dialect/IR/AddressOpsTypes.cpp.inc" + >(); +} + +namespace { +ParseResult parseAddressType(OpAsmParser &parser, Type &ty) { + if (succeeded(parser.parseOptionalColon()) && parser.parseType(ty)) + return parser.emitError(parser.getNameLoc(), "expected a type"); + if (!ty) + ty = parser.getBuilder().getType(); + return success(); +} +void printAddressType(OpAsmPrinter &p, Operation *op, AddressType ty) { + if (ty.getAddressSpace() != nullptr) + p << " : " << ty; +} + +ParseResult parseIntType(OpAsmParser &parser, Type &ty) { + if (succeeded(parser.parseOptionalColon()) && parser.parseType(ty)) + return parser.emitError(parser.getNameLoc(), "expected a type"); + if (!ty) + ty = parser.getBuilder().getIndexType(); + return success(); +} +void printIntType(OpAsmPrinter &p, Operation *op, Type ty) { + if (!ty.isIndex()) + p << " : " << ty; +} +} // namespace + +//===----------------------------------------------------------------------===// +// CastInt Op +//===----------------------------------------------------------------------===// + +bool CastIntOp::areCastCompatible(mlir::TypeRange lhs, mlir::TypeRange rhs) { + return isa(lhs.front()) != isa(rhs.front()); +} + +//===----------------------------------------------------------------------===// +// Constant Op +//===----------------------------------------------------------------------===// + +void ConstantOp::build(OpBuilder &odsBuilder, OperationState &odsState, + int64_t value, Attribute addressSpace) { + build(odsBuilder, odsState, odsBuilder.getType(addressSpace), + odsBuilder.getIndexAttr(value)); +} + +void ConstantOp::getAsmResultNames(OpAsmSetValueNameFn setNameFn) { + SmallString<32> buffer; + llvm::raw_svector_ostream name(buffer); + name << "addr" << getValueAttr().getValue(); + setNameFn(getResult(), name.str()); +} + +OpFoldResult ConstantOp::fold(FoldAdaptor adaptor) { + return adaptor.getValueAttr(); +} + +//===----------------------------------------------------------------------===// +// TypeOffset Op +//===----------------------------------------------------------------------===// + +void TypeOffsetOp::build(OpBuilder &odsBuilder, OperationState &odsState, + TypeAttr baseType, Type resultTy) { + build(odsBuilder, odsState, + resultTy ? resultTy : odsBuilder.getIndexType(), baseType); +} + +OpFoldResult TypeOffsetOp::fold(FoldAdaptor adaptor) { + return adaptor.getBaseTypeAttr(); +} + +//===----------------------------------------------------------------------===// +// Cast Op +//===----------------------------------------------------------------------===// + +void CastOp::build(OpBuilder &odsBuilder, OperationState &odsState, + Attribute addressSpace, Value input) { + build(odsBuilder, odsState, + odsBuilder.getType(addressSpace), input); +} + +LogicalResult CastOp::canonicalize(CastOp op, PatternRewriter &rewriter) { + if (op.getInput().getType() == op.getType()) { + rewriter.replaceOp(op, op.getInput()); + return success(); + } + return failure(); +} + +void CastIntOp::build(OpBuilder &odsBuilder, OperationState &odsState, + Value input, Type resultTy) { + if (!resultTy) + resultTy = isa(input.getType()) + ? cast(odsBuilder.getIndexType()) + : cast(odsBuilder.getType()); + build(odsBuilder, odsState, resultTy, input); +} + +//===----------------------------------------------------------------------===// +// FromMemRef Op +//===----------------------------------------------------------------------===// + +LogicalResult FromMemRefOp::verify() { + if (getType().getAddressSpace() != getInput().getType().getMemorySpace()) + return emitError("address space mismatch"); + return success(); +} + +LogicalResult FromMemRefOp::canonicalize(FromMemRefOp op, + PatternRewriter &rewriter) { + // Collapse the following patterns to an address: + // 1) Result %a = %addr + // %m = addr.to_memref %addr base %base : memref + // %a = addr.from_memref [%m : memref] + // 2) Result %a = %base + // %m = addr.to_memref %addr base %base : memref + // %a = addr.from_memref extract_base [%m : memref] + // 3) Result %a = %addr + // %m = addr.to_memref %addr : memref + // %a = addr.from_memref extract_base [%m : memref] + auto input = dyn_cast_or_null(op.getInput().getDefiningOp()); + if (!input) + return failure(); + // Handle cases 1 & 3 + if (!op.getExtractBase() || !input.getBase()) + rewriter.replaceOp(op, input.getAddress()); + else + rewriter.replaceOp(op, input.getBase()); + return success(); +} + +namespace { +ParseResult parseFromMemRef(OpAsmParser &parser, Type &inputTy, + Type &resultTy) { + if (parser.parseColonType(inputTy)) + return parser.emitError(parser.getNameLoc(), "expected a type"); + auto memrefTy = dyn_cast(inputTy); + assert(memrefTy && "Expected a memref type."); + resultTy = + parser.getBuilder().getType(memrefTy.getMemorySpace()); + return success(); +} +void printFromMemRef(OpAsmPrinter &p, Operation *op, Type inputTy, + Type resultTy) { + p << " : " << inputTy; +} + +ParseResult parseFromUnrankedMemRef(OpAsmParser &parser, Type &inputTy, + Type &resultTy) { + if (parser.parseColonType(inputTy)) + return parser.emitError(parser.getNameLoc(), "expected a type"); + auto memrefTy = dyn_cast(inputTy); + assert(memrefTy && "Expected a memref type."); + resultTy = + parser.getBuilder().getType(memrefTy.getMemorySpace()); + return success(); +} +void printFromUnrankedMemRef(OpAsmPrinter &p, Operation *op, Type inputTy, + Type resultTy) { + p << " : " << inputTy; +} +} // namespace + +namespace {} // namespace + +//===----------------------------------------------------------------------===// +// ToMemRef Op +//===----------------------------------------------------------------------===// + +void ToMemRefOp::build(OpBuilder &odsBuilder, OperationState &odsState, + MemRefType type, Value address) { + build(odsBuilder, odsState, type, address, nullptr); +} + +void ToUnrankedMemRefOp::build(OpBuilder &odsBuilder, OperationState &odsState, + UnrankedMemRefType type, Value address) { + build(odsBuilder, odsState, type, address, nullptr); +} + +LogicalResult ToMemRefOp::verify() { + Attribute inputAS = getAddress().getType().getAddressSpace(); + if (inputAS != getType().getMemorySpace() || + (getBase() && + dyn_cast(getBase().getType()).getAddressSpace() != inputAS)) + return emitError("address space mismatch"); + return success(); +} + +LogicalResult ToMemRefOp::canonicalize(ToMemRefOp op, + PatternRewriter &rewriter) { + // Collapse the following pattern to a memref, where %m = %memref: + // %a = addr.from_memref [%memref : memref] + // %b = addr.from_memref extract_base [%memref : memref] + // %m = addr.to_memref %a base %b : memref + auto address = + dyn_cast_or_null(op.getAddress().getDefiningOp()); + // Fail if the address doesn't come from a `from_memref` or if the Op doesn't + // have a base. + if (!address || !op.getBase()) + return failure(); + auto base = dyn_cast_or_null(op.getBase().getDefiningOp()); + // Fail if the base doesn't come from a `from_memref` or if the base is + // unknown. + if (!base || address.getInput() != base.getInput() || !base.getExtractBase()) + return failure(); + rewriter.replaceOp(op, address.getInput()); + return success(); +} + +namespace { +ParseResult parseToMemRef(OpAsmParser &parser, + std::optional &base, + Type &baseTy, Type &addressTy, Type &resultTy) { + if (succeeded(parser.parseOptionalKeyword("base")) && + parser.parseOperand(base.emplace())) + return parser.emitError(parser.getNameLoc(), "expected an operand"); + if (parser.parseColonType(resultTy)) + return parser.emitError(parser.getNameLoc(), "expected a type"); + auto memrefTy = dyn_cast(resultTy); + assert(memrefTy && "Expected a memref type."); + addressTy = + parser.getBuilder().getType(memrefTy.getMemorySpace()); + if (base.has_value()) + baseTy = addressTy; + return success(); +} +void printToMemRef(OpAsmPrinter &p, Operation *op, Value base, Type baseTy, + Type addressTy, Type memrefTy) { + if (base) + p << "base " << base; + p << " : " << memrefTy; +} + +ParseResult +parseToUnrankedMemRef(OpAsmParser &parser, + std::optional &base, + Type &baseTy, Type &addressTy, Type &resultTy) { + if (succeeded(parser.parseOptionalKeyword("base")) && + parser.parseOperand(base.emplace())) + return parser.emitError(parser.getNameLoc(), "expected an operand"); + if (parser.parseColonType(resultTy)) + return parser.emitError(parser.getNameLoc(), "expected a type"); + auto memrefTy = dyn_cast(resultTy); + assert(memrefTy && "Expected a memref type."); + addressTy = + parser.getBuilder().getType(memrefTy.getMemorySpace()); + if (base.has_value()) + baseTy = addressTy; + return success(); +} +void printToUnrankedMemRef(OpAsmPrinter &p, Operation *op, Value base, + Type baseTy, Type addressTy, Type memrefTy) { + if (base) + p << "base " << base; + p << " : " << memrefTy; +} +} // namespace + +#include "Address/Dialect/IR/AddressOpsDialect.cpp.inc" + +#define GET_OP_CLASSES +#include "Address/Dialect/IR/AddressOps.cpp.inc" + +#define GET_TYPEDEF_CLASSES +#include "Address/Dialect/IR/AddressOpsTypes.cpp.inc" diff --git a/third_party/wafer/lib/Dialect/Address/IR/CMakeLists.txt b/third_party/wafer/lib/Dialect/Address/IR/CMakeLists.txt new file mode 100755 index 00000000..e280fe39 --- /dev/null +++ b/third_party/wafer/lib/Dialect/Address/IR/CMakeLists.txt @@ -0,0 +1,13 @@ +add_mlir_dialect_library( + MLIRAddress + AddressDialect.cpp + ADDITIONAL_HEADER_DIRS + ${PROJECT_SOURCE_DIR}/mlir/Dialect/Address + DEPENDS + MLIRAddressOpsIncGen + LINK_LIBS + PUBLIC + MLIRIR + MLIRInferTypeOpInterface + MLIRFuncDialect +) diff --git a/third_party/wafer/lib/Dialect/Address/Transforms/AddrToLLVM.cpp b/third_party/wafer/lib/Dialect/Address/Transforms/AddrToLLVM.cpp new file mode 100755 index 00000000..81434ec4 --- /dev/null +++ b/third_party/wafer/lib/Dialect/Address/Transforms/AddrToLLVM.cpp @@ -0,0 +1,366 @@ +//===- AddrToLLVM.cpp - Implementation of Address to LLVM conversion ------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file implements the Address to LLVM conversion pass. +// +//===----------------------------------------------------------------------===// + +#include "Address/Dialect/IR/AddressDialect.h" +#include "Address/Transforms/Passes.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Analysis/DataLayoutAnalysis.h" +#include "mlir/Conversion/ConvertToLLVM/ToLLVMInterface.h" +#include "mlir/Conversion/FuncToLLVM/ConvertFuncToLLVM.h" +#include "mlir/Conversion/LLVMCommon/ConversionTarget.h" +#include "mlir/Conversion/LLVMCommon/MemRefBuilder.h" +#include "mlir/Conversion/LLVMCommon/Pattern.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Interfaces/CallInterfaces.h" +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace addr { +#define GEN_PASS_DEF_ADDRTOLLVM +#include "Address/Transforms/Passes.h.inc" +} // namespace addr +} // namespace mlir + +using namespace mlir; +using namespace mlir::addr; + +namespace { +struct AddrToLLVM : public ::mlir::addr::impl::AddrToLLVMBase { + using Base::Base; + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override; +}; + +struct AddrTypeConverter : public ::mlir::LLVMTypeConverter { + AddrTypeConverter(MLIRContext *ctx, const LowerToLLVMOptions &options, + const DataLayoutAnalysis *analysis = nullptr) + : LLVMTypeConverter(ctx, options, analysis) { + addConversion([&](AddressType type) { + unsigned as = 0; + if (auto attr = dyn_cast_or_null(type.getAddressSpace())) + as = attr.getUInt(); + return LLVM::LLVMPointerType::get(&getContext(), as); + }); + } +}; + +struct ConstantOpConversion : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(ConstantOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; + +struct TypeOffsetOpConversion : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(TypeOffsetOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; + +struct CastOpConversion : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(CastOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; + +struct CastIntOpConversion : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(CastIntOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; + +struct FromMemRefOpConversion : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(FromMemRefOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; + +struct FromUnrankedMemRefOpConversion + : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(FromUnrankedMemRefOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; + +struct ToUnrankedMemRefOpConversion + : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(ToUnrankedMemRefOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; + +struct ToMemRefOpConversion : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(ToMemRefOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; + +struct PtrAddOpConversion : public ConvertOpToLLVMPattern { +protected: + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + LogicalResult + matchAndRewrite(PtrAddOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const final; +}; +} // namespace + +LogicalResult ConstantOpConversion::matchAndRewrite( + ConstantOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + auto cst = rewriter.create( + op.getLoc(), rewriter.getIntegerAttr( + typeConverter->convertType(op.getValueAttr().getType()), + operands.getValue())); + // Convert the constant to a ptr + rewriter.replaceOpWithNewOp( + op, typeConverter->convertType(op.getType()), cst.getResult()); + return success(); +} + +LogicalResult TypeOffsetOpConversion::matchAndRewrite( + TypeOffsetOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + // Use GEP to compute the type offset + const LLVMTypeConverter *tc = + static_cast(typeConverter); + auto ptrTy = LLVM::LLVMPointerType::get(getContext()); + Value nullOp = rewriter.create(op.getLoc(), ptrTy); + auto offset = rewriter.create( + op.getLoc(), ptrTy, tc->convertType(op.getBaseType()), nullOp, + ArrayRef({LLVM::GEPArg(1)})); + rewriter.replaceOpWithNewOp( + op, tc->convertType(op.getType()), offset.getRes()); + return success(); +} + +LogicalResult +CastOpConversion::matchAndRewrite(CastOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + rewriter.replaceOpWithNewOp( + op, typeConverter->convertType(op.getType()), operands.getInput()); + return success(); +} + +LogicalResult CastIntOpConversion::matchAndRewrite( + CastIntOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + if (op.getType().isIntOrIndex()) + rewriter.replaceOpWithNewOp( + op, typeConverter->convertType(op.getType()), operands.getInput()); + else + rewriter.replaceOpWithNewOp( + op, typeConverter->convertType(op.getType()), operands.getInput()); + return success(); +} + +LogicalResult FromMemRefOpConversion::matchAndRewrite( + FromMemRefOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + MemRefDescriptor descriptor(operands.getInput()); + if (op.getExtractBase()) + rewriter.replaceOp(op, descriptor.allocatedPtr(rewriter, op.getLoc())); + else + rewriter.replaceOp(op, descriptor.alignedPtr(rewriter, op.getLoc())); + return success(); +} + +LogicalResult FromUnrankedMemRefOpConversion::matchAndRewrite( + FromUnrankedMemRefOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + UnrankedMemRefDescriptor descriptor(operands.getInput()); + rewriter.replaceOp(op, descriptor.memRefDescPtr(rewriter, op->getLoc())); + return success(); +} + +LogicalResult ToMemRefOpConversion::matchAndRewrite( + ToMemRefOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + const LLVMTypeConverter *tc = + static_cast(typeConverter); + Value descriptor; + if (operands.getBase()) + descriptor = MemRefDescriptor::fromStaticShape( + rewriter, op.getLoc(), *tc, op.getType(), operands.getAddress(), + operands.getBase()); + else + descriptor = MemRefDescriptor::fromStaticShape( + rewriter, op.getLoc(), *tc, op.getType(), operands.getAddress()); + rewriter.replaceOp(op, descriptor); + return success(); +} + +LogicalResult ToUnrankedMemRefOpConversion::matchAndRewrite( + ToUnrankedMemRefOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + auto loc = op->getLoc(); + const LLVMTypeConverter *tc = + static_cast(typeConverter); + + auto elementPtrTy = LLVM::LLVMPointerType::get(rewriter.getContext(), 0); + + UnrankedMemRefDescriptor descriptor = UnrankedMemRefDescriptor::poison( + rewriter, loc, tc->convertType(op.getType())); + descriptor.setMemRefDescPtr(rewriter, loc, operands.getAddress()); + rewriter.replaceOp(op, (Value)descriptor); + return success(); +} + +LogicalResult +PtrAddOpConversion::matchAndRewrite(PtrAddOp op, OpAdaptor operands, + ConversionPatternRewriter &rewriter) const { + rewriter.replaceOpWithNewOp( + op, operands.getBase().getType(), rewriter.getI8Type(), + operands.getBase(), operands.getOffset()); + return success(); +} + +struct AddressTypeConversionPattern : public ConversionPattern { + AddressTypeConversionPattern(TypeConverter &converter, MLIRContext *ctx) + : ConversionPattern(converter, Pattern::MatchAnyOpTypeTag(), 2, ctx) {} + + LogicalResult + matchAndRewrite(Operation *op, ArrayRef operands, + ConversionPatternRewriter &rewriter) const override { + + // Convert result types using the type converter. + SmallVector newResultTypes; + if (failed(typeConverter->convertTypes(op->getResultTypes(), + newResultTypes))) { + return failure(); + } + + // Copy attributes (if needed). + SmallVector newAttrs; + for (auto attr : op->getAttrs()) { + newAttrs.push_back(attr); + } + + // Create a new operation state with converted operands, attributes, and + // result types. + OperationState newOpState(op->getLoc(), op->getName()); + newOpState.addOperands(operands); + newOpState.addAttributes(newAttrs); + newOpState.addTypes(newResultTypes); + + // Copy regions (if any) from the original operation. + for (Region ®ion : op->getRegions()) { + newOpState.addRegion()->takeBody(region); + } + + // Create the new operation and replace the original operation with it. + Operation *newOp = rewriter.create(newOpState); + rewriter.replaceOp(op, newOp->getResults()); + return success(); + } +}; + +struct MKBitCastConversionPattern : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(mk::BitcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + + // Result type should have same layout and address space as the source type. + auto sourceType = op.getSrc().getType(); + assert(isa(sourceType)); + auto rankedMemRefType = cast(sourceType); + + auto memrefToPtr = rewriter.create( + loc, addr::AddressType::get(rewriter.getContext()), adaptor.getSrc()); + + auto resultMemrefType = cast(op->getResultTypes()[0]); + + MemRefType resultType = MemRefType::get( + resultMemrefType.getShape(), resultMemrefType.getElementType(), + resultMemrefType.getLayout(), resultMemrefType.getMemorySpace()); + + rewriter.replaceOpWithNewOp(op, resultType, memrefToPtr); + return success(); + } +}; + +void AddrToLLVM::runOnOperation() { + ModuleOp module = getOperation(); + StringRef dataLayout; + auto dataLayoutAttr = dyn_cast_or_null( + module->getAttr(LLVM::LLVMDialect::getDataLayoutAttrName())); + if (dataLayoutAttr) + dataLayout = dataLayoutAttr.getValue(); + if (failed(LLVM::LLVMDialect::verifyDataLayoutString( + dataLayout, [this](const Twine &message) { + getOperation().emitError() << message.str(); + }))) { + signalPassFailure(); + return; + } + const auto &dataLayoutAnalysis = getAnalysis(); + LowerToLLVMOptions options(&getContext(), + dataLayoutAnalysis.getAtOrAbove(module)); + AddrTypeConverter typeConverter(&getContext(), options, &dataLayoutAnalysis); + LLVMConversionTarget target(getContext()); + std::optional optSymbolTable = std::nullopt; + const SymbolTable *symbolTable = nullptr; + if (!options.useBarePtrCallConv) { + optSymbolTable.emplace(module); + symbolTable = &optSymbolTable.value(); + } + RewritePatternSet patterns(&getContext()); + patterns.insert( + typeConverter); + // WORKAROUND: Bufferize not support addr dialect, so convert mk::bitcast here + patterns.add(patterns.getContext()); + + // populateFuncToLLVMConversionPatterns(typeConverter, patterns, symbolTable); + if (failed(applyPartialConversion(module, target, std::move(patterns)))) + signalPassFailure(); + + RewritePatternSet convertTypes(&getContext()); + convertTypes.insert(typeConverter, + &getContext()); + target.markUnknownOpDynamicallyLegal([&](mlir::Operation *op) { + for (Type resultType : op->getResultTypes()) { + if (isa(resultType)) { + return false; + } + } + return true; + }); + + if (failed(applyPartialConversion(module, target, std::move(convertTypes)))) + signalPassFailure(); +} diff --git a/third_party/wafer/lib/Dialect/Address/Transforms/CMakeLists.txt b/third_party/wafer/lib/Dialect/Address/Transforms/CMakeLists.txt new file mode 100755 index 00000000..982efc97 --- /dev/null +++ b/third_party/wafer/lib/Dialect/Address/Transforms/CMakeLists.txt @@ -0,0 +1,16 @@ +add_mlir_dialect_library( + MLIRAddressTransforms + AddrToLLVM.cpp + ADDITIONAL_HEADER_DIRS + ${MLIR_MAIN_INCLUDE_DIR}/mlir/Dialect/Address + DEPENDS + MLIRAddressPassIncGen + LINK_LIBS + PUBLIC + MLIRArithDialect + MLIRDataLayoutInterfaces + MLIRIR + MLIRIndexDialect + MLIRPass + MLIRSupport +) diff --git a/third_party/wafer/lib/Dialect/CMakeLists.txt b/third_party/wafer/lib/Dialect/CMakeLists.txt new file mode 100755 index 00000000..7740d331 --- /dev/null +++ b/third_party/wafer/lib/Dialect/CMakeLists.txt @@ -0,0 +1,5 @@ +add_subdirectory(Address) +# add_subdirectory(TritonTilingExt) +# add_subdirectory(TritonStructured) +add_subdirectory(MagicKernel) +add_subdirectory(Wafer) diff --git a/third_party/wafer/lib/Dialect/MagicKernel/CMakeLists.txt b/third_party/wafer/lib/Dialect/MagicKernel/CMakeLists.txt new file mode 100755 index 00000000..e2953d90 --- /dev/null +++ b/third_party/wafer/lib/Dialect/MagicKernel/CMakeLists.txt @@ -0,0 +1,12 @@ +add_triton_library(MagicKernelIR + IR/MagicKernelDialect.cpp + Transforms/BufferizableOpInterfaceImpl.cpp + + DEPENDS + MagicKernelTableGen + + LINK_LIBS PUBLIC + MLIRIR +) + +add_subdirectory(Transforms) diff --git a/third_party/wafer/lib/Dialect/MagicKernel/IR/MagicKernelDialect.cpp b/third_party/wafer/lib/Dialect/MagicKernel/IR/MagicKernelDialect.cpp new file mode 100755 index 00000000..a75770a6 --- /dev/null +++ b/third_party/wafer/lib/Dialect/MagicKernel/IR/MagicKernelDialect.cpp @@ -0,0 +1,51 @@ +//===------------------- MagicKernelDialect.cpp ---------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Bufferization/IR/BufferizableOpInterface.h" + +using namespace mlir; +using namespace mlir::mk; + +LogicalResult PrintOp::verify() { + if (getOperands().size() > 1) + return emitOpError("expects at most one operand"); + return success(); +} + +/// Dialect creation, the instance will be owned by the context. This is the +/// point of registration of custom types and operations for the dialect. +void MagicKernelDialect::initialize() { + addOperations< +#define GET_OP_LIST +#include "magic-kernel/Dialect/IR/MagicKernelOps.cpp.inc" + >(); + // TODO: Add BufferizableOpInterface to all ops that can be bufferized + declarePromisedInterfaces< + bufferization::BufferizableOpInterface, mk::DotOp, mk::DotScaledOp, + mk::SigmoidOp, mk::GeluOp, mk::GatherOp, mk::PrintOp, mk::AtomicRMWOp, + mk::AtomicCASOp, mk::ArgMaxOp, mk::ArgMinOp, mk::Bit2FpOp, mk::MaskMoveOp, + mk::UnEqualVV, mk::EqualVV, mk::EqualVS, mk::LessThenVS, mk::BoolEqualVS, + mk::ReduceMaxOp, mk::ReduceMinOp, mk::ReduceSumOp, mk::DequantOp, + mk::BitcastOp, mk::AddVS, mk::SubVS, mk::MulVS>(); +} + +//===----------------------------------------------------------------------===// +// TableGen'd op method definitions +//===----------------------------------------------------------------------===// + +OpFoldResult mk::BitcastOp::fold(FoldAdaptor adaptor) { + if (getOperand().getType() == getResult().getType()) { + return getOperand(); + } + return {}; +} + +#define GET_OP_CLASSES +#include "magic-kernel/Dialect/IR/MagicKernelOps.cpp.inc" + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.cpp.inc" diff --git a/third_party/wafer/lib/Dialect/MagicKernel/Transforms/BufferizableOpInterfaceImpl.cpp b/third_party/wafer/lib/Dialect/MagicKernel/Transforms/BufferizableOpInterfaceImpl.cpp new file mode 100755 index 00000000..f01a878a --- /dev/null +++ b/third_party/wafer/lib/Dialect/MagicKernel/Transforms/BufferizableOpInterfaceImpl.cpp @@ -0,0 +1,338 @@ +//===- BufferizableOpInterfaceImpl.cpp ----------------------------------- ===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM +// Exceptions. See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// +// +// This file implements mk dialect DestinationStyleOp BufferizableOpInterface. +// +//===----------------------------------------------------------------------===// + +#include "magic-kernel/Transforms/BufferizableOpInterfaceImpl.h" +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Bufferization/IR/DstBufferizableOpInterfaceImpl.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/Operation.h" +#include "mlir/Interfaces/DestinationStyleOpInterface.h" + +using namespace mlir; +using namespace mlir::bufferization; + +/// Generic conversion for any DestinationStyleOpInterface on tensors. +static LogicalResult +bufferizeDestinationStyleOpInterface(RewriterBase &rewriter, + DestinationStyleOpInterface op, + const BufferizationOptions &options, + BufferizationState &bufferizationState) { + // Take a guard before anything else. + OpBuilder::InsertionGuard g(rewriter); + rewriter.setInsertionPoint(op); + + // Nothing to do. This op is already bufferized. + if (op.hasPureBufferSemantics()) + return success(); + + // Ensure op has only tensors. Allow mixed tensor-buffer mode on a per-need + // basis. + if (!op.hasPureTensorSemantics()) + return op->emitError() << "op does not have pure tensor semantics"; + + // New input operands for the cloned op. + SmallVector newInputBuffers; + newInputBuffers.reserve(op.getNumDpsInputs()); + for (OpOperand *opOperand : op.getDpsInputOperands()) { + if (op.isScalar(opOperand)) { + newInputBuffers.push_back(opOperand->get()); + continue; + } + FailureOr buffer = + getBuffer(rewriter, opOperand->get(), options, bufferizationState); + if (failed(buffer)) + return failure(); + newInputBuffers.push_back(*buffer); + } + + // New output operands for the cloned op. + SmallVector newOutputBuffers; + for (OpResult opResult : op->getOpResults()) { + OpOperand *opOperand = op.getDpsInitOperand(opResult.getResultNumber()); + FailureOr resultBuffer = getBuffer( + rewriter, opOperand->get(), options, bufferizationState); + if (failed(resultBuffer)) + return failure(); + newOutputBuffers.push_back(*resultBuffer); + } + + // Merge input/output operands. + SmallVector newOperands = newInputBuffers; + newOperands.append(newOutputBuffers.begin(), newOutputBuffers.end()); + + // Set insertion point now that potential alloc/dealloc are introduced. + rewriter.setInsertionPoint(op); + // Clone the op, but use the new operands. Move the existing block into the + // new op. Since the new op does not have any tensor results, it does not + // return anything. + OperationState state(op->getLoc(), op->getName(), newOperands, TypeRange{}, + op->getAttrs()); + + Operation *newOp = Operation::create(state); + + // We don't want the rewriter tracks an incomplete operation, so insert new + // operation after op was fully constructed. + rewriter.insert(newOp); + + // Replace the results of the old op with the new output buffers. + replaceOpWithBufferizedValues(rewriter, op, newOutputBuffers); + + return success(); +} + +/// Bufferization of mk ops. Replace with a new mk op that operates entirely on +/// memrefs. +template +struct MKOpInterface + : public DstBufferizableOpInterfaceExternalModel, + OpTy> { + + bool bufferizesToElementwiseAccess(Operation *op, const AnalysisState &state, + ArrayRef opOperands) const { + return op->hasTrait(); + } + + LogicalResult bufferize(Operation *op, RewriterBase &rewriter, + const BufferizationOptions &options, + BufferizationState &bufferizationState) const { + return bufferizeDestinationStyleOpInterface( + rewriter, cast(op), options, + bufferizationState); + } +}; + +struct AtomicRMWOpInterface + : public DstBufferizableOpInterfaceExternalModel { + // TODO: Check for memory effect + LogicalResult bufferize(Operation *op, RewriterBase &rewriter, + const BufferizationOptions &options, + BufferizationState &bufferizationState) const { + + auto atomicRMWOp = cast(op); + if (!isa(atomicRMWOp.getPtr().getType())) { + return failure(); + } + + FailureOr valBuffer = getBuffer( + rewriter, atomicRMWOp.getVal(), options, bufferizationState); + if (failed(valBuffer)) + return failure(); + FailureOr outputBuffer = getBuffer( + rewriter, atomicRMWOp.getDst(), options, bufferizationState); + if (failed(outputBuffer)) + return failure(); + rewriter.create( + atomicRMWOp.getLoc(), + /*result=*/TypeRange(), atomicRMWOp.getPtr(), *valBuffer, *outputBuffer, + atomicRMWOp.getAtomicRmwOpAttr(), atomicRMWOp.getSemAttr(), + atomicRMWOp.getScopeAttr()); + replaceOpWithBufferizedValues(rewriter, op, *outputBuffer); + return success(); + } +}; + +struct AtomicCASOpInterface + : public DstBufferizableOpInterfaceExternalModel { + // TODO: Check for memory effect + LogicalResult bufferize(Operation *op, RewriterBase &rewriter, + const BufferizationOptions &options, + BufferizationState &bufferizationState) const { + + auto atomicCASOp = cast(op); + if (!isa(atomicCASOp.getPtr().getType())) { + return failure(); + } + FailureOr cmpBuffer = getBuffer( + rewriter, atomicCASOp.getCmp(), options, bufferizationState); + if (failed(cmpBuffer)) + return failure(); + + FailureOr valBuffer = getBuffer( + rewriter, atomicCASOp.getVal(), options, bufferizationState); + if (failed(valBuffer)) + return failure(); + + FailureOr outputBuffer = getBuffer( + rewriter, atomicCASOp.getDst(), options, bufferizationState); + if (failed(outputBuffer)) + return failure(); + rewriter.create( + atomicCASOp.getLoc(), + /*result=*/TypeRange(), atomicCASOp.getPtr(), *cmpBuffer, *valBuffer, + *outputBuffer, atomicCASOp.getSemAttr(), atomicCASOp.getScopeAttr()); + replaceOpWithBufferizedValues(rewriter, op, *outputBuffer); + return success(); + } +}; + +struct BitCastOpInterface + : public BufferizableOpInterface::ExternalModel { + + bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand, + const AnalysisState &state) const { + return false; + } + + bool bufferizesToMemoryWrite(Operation *op, OpOperand &opOperand, + const AnalysisState &state) const { + return false; + } + + AliasingValueList getAliasingValues(Operation *op, OpOperand &opOperand, + const AnalysisState &state) const { + return {{op->getResult(0), BufferRelation::Equivalent}}; + } + + // TODO: Check for memory effect + LogicalResult bufferize(Operation *op, RewriterBase &rewriter, + const BufferizationOptions &options, + BufferizationState &bufferizationState) const { + + auto bitcastOp = cast(op); + + auto inputType = bitcastOp.getSrc().getType(); + auto resType = bitcastOp.getType(); + assert(isa(inputType) && isa(resType) && + "expected ranked tensor type"); + + FailureOr srcBuffer = getBuffer( + rewriter, bitcastOp.getSrc(), options, bufferizationState); + if (failed(srcBuffer)) + return failure(); + + // Result type should have same layout and address space as the source type. + auto sourceType = srcBuffer->getType(); + assert(isa(sourceType) && + "expected memref type for bitcast source"); + auto rankedMemRefType = cast(sourceType); + + auto resultTensorType = cast(resType); + MemRefType resultType = MemRefType::get( + resultTensorType.getShape(), resultTensorType.getElementType(), + rankedMemRefType.getLayout(), rankedMemRefType.getMemorySpace()); + + replaceOpWithNewBufferizedOp(rewriter, op, resultType, + *srcBuffer); + return success(); + } +}; + +struct SendOpInterface + : public BufferizableOpInterface::ExternalModel { + bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand, + const AnalysisState &state) const { + auto sendOp = cast(op); + // mk.send reads the local src buffer. The dst_addr is "addr-like" and + // should not be considered a memory read. + return &opOperand == &sendOp.getSrcMutable(); + } + + bool bufferizesToMemoryWrite(Operation *op, OpOperand &opOperand, + const AnalysisState &state) const { + return false; + } + + AliasingValueList getAliasingValues(Operation *op, OpOperand &opOperand, + const AnalysisState &state) const { + return {}; + } + + LogicalResult bufferize(Operation *op, RewriterBase &rewriter, + const BufferizationOptions &options, + BufferizationState &bufferizationState) const { + auto sendOp = cast(op); + + // Nothing to do. This op is already bufferized. + if (!isa(sendOp.getDstAddr().getType()) && + !isa(sendOp.getSrc().getType())) + return success(); + + OpBuilder::InsertionGuard g(rewriter); + rewriter.setInsertionPoint(sendOp); + + SmallVector newOperands(sendOp->getOperands().begin(), + sendOp->getOperands().end()); + + if (isa(sendOp.getDstAddr().getType())) { + FailureOr dstBuffer = getBuffer( + rewriter, sendOp.getDstAddr(), options, bufferizationState); + if (failed(dstBuffer)) + return failure(); + newOperands[sendOp.getDstAddrMutable().getOperandNumber()] = *dstBuffer; + } + + if (isa(sendOp.getSrc().getType())) { + FailureOr srcBuffer = getBuffer( + rewriter, sendOp.getSrc(), options, bufferizationState); + if (failed(srcBuffer)) + return failure(); + newOperands[sendOp.getSrcMutable().getOperandNumber()] = *srcBuffer; + } + + OperationState state(sendOp->getLoc(), sendOp->getName(), newOperands, + TypeRange{}, sendOp->getAttrs()); + Operation *newOp = Operation::create(state); + rewriter.insert(newOp); + rewriter.eraseOp(sendOp); + return success(); + } +}; + +/// Helper structure that iterates over all mkOps in `OpTys` and registers +/// the `BufferizableOpInterface` with each of them. +template struct MKOpInterfaceHelper { + static void registerOpInterface(MLIRContext *ctx) { + (Ops::template attachInterface>(*ctx), ...); + } +}; + +void mlir::mk::registerBufferizableOpInterfaceExternalModels( + mlir::DialectRegistry ®istry) { + registry.addExtension( + +[](MLIRContext *ctx, mlir::mk::MagicKernelDialect *dialect) { + // TODO: Register all mk ops. + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + mk::AtomicRMWOp::attachInterface(*ctx); + mk::AtomicCASOp::attachInterface(*ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + MKOpInterfaceHelper::registerOpInterface(ctx); + mk::BitcastOp::attachInterface(*ctx); + mk::RemoteStoreOp::attachInterface(*ctx); + }); +} diff --git a/third_party/wafer/lib/Dialect/MagicKernel/Transforms/CMakeLists.txt b/third_party/wafer/lib/Dialect/MagicKernel/Transforms/CMakeLists.txt new file mode 100644 index 00000000..c1da4e0c --- /dev/null +++ b/third_party/wafer/lib/Dialect/MagicKernel/Transforms/CMakeLists.txt @@ -0,0 +1,18 @@ +add_triton_library(MKTransforms + MaterializeStridedLinalgInputsPass.cpp + + DEPENDS + MagicKernelTableGen + MKTransformsPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRFuncDialect + MLIRMemRefDialect + MLIRIR + MLIRPass + MLIRSCFDialect + MLIRSupport + TritonIR + TritonTransforms +) diff --git a/third_party/wafer/lib/Dialect/MagicKernel/Transforms/MaterializeStridedLinalgInputsPass.cpp b/third_party/wafer/lib/Dialect/MagicKernel/Transforms/MaterializeStridedLinalgInputsPass.cpp new file mode 100644 index 00000000..726bbe27 --- /dev/null +++ b/third_party/wafer/lib/Dialect/MagicKernel/Transforms/MaterializeStridedLinalgInputsPass.cpp @@ -0,0 +1,232 @@ + +#include "magic-kernel/Dialect/IR/MagicKernelDialect.h" +#include "magic-kernel/Transforms/Passes.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Visitors.h" +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Pass/Pass.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/Support/Debug.h" +#include + +#define DEBUG_TYPE "materialize-strided-linalg-inputs" + +using namespace mlir; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_MATERIALIZESTRIDEDLINALGINPUTS +#include "magic-kernel/Transforms/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +static bool hasNonContiguousStrides(MemRefType type) { + SmallVector strides; + int64_t offset = 0; + if (failed(type.getStridesAndOffset(strides, offset))) + return true; + + int64_t expected = 1; + ArrayRef shape = type.getShape(); + + for (int64_t i = type.getRank() - 1; i >= 0; --i) { + if (ShapedType::isDynamic(shape[i]) || ShapedType::isDynamic(strides[i])) + return true; + + if (shape[i] == 1) + continue; + + if (strides[i] != expected) + return true; + + expected *= shape[i]; + } + + return false; +} + +static bool isComputeGeneric(linalg::GenericOp op) { + Block &body = op.getRegion().front(); + + Operation *computeOp = nullptr; + for (Operation &inner : body.without_terminator()) { + if (computeOp) + return false; // more than one non-yield op + + computeOp = &inner; + } + + if (!computeOp) + return false; + + return computeOp->getDialect()->getNamespace() == + arith::ArithDialect::getDialectNamespace(); +} + +struct MaterializeStridedLinalgInputsPass + : public triton::impl::MaterializeStridedLinalgInputsBase< + MaterializeStridedLinalgInputsPass> { + using MaterializeStridedLinalgInputsBase< + MaterializeStridedLinalgInputsPass>::MaterializeStridedLinalgInputsBase; + + void process(func::FuncOp &func) { + IRRewriter rewriter(func.getContext()); + + SmallVector generics; + func.walk([&](linalg::GenericOp op) { generics.push_back(op); }); + + for (linalg::GenericOp generic : generics) { + if (!isComputeGeneric(generic)) + continue; + + rewriter.setInsertionPoint(generic); + Location loc = generic.getLoc(); + + for (OpOperand *inputOperand : generic.getDpsInputOperands()) { + Value input = inputOperand->get(); + auto subview = input.getDefiningOp(); + if (!subview) + continue; + + auto subviewType = dyn_cast(subview.getType()); + if (!subviewType || !hasNonContiguousStrides(subviewType)) + continue; + + SmallVector dynamicSizes; + for (auto [idx, dim] : llvm::enumerate(subviewType.getShape())) { + if (ShapedType::isDynamic(dim)) { + dynamicSizes.push_back( + rewriter.create(loc, subview, idx)); + } + } + + auto allocType = MemRefType::get( + subviewType.getShape(), subviewType.getElementType(), + MemRefLayoutAttrInterface{}, subviewType.getMemorySpace()); + + Value alloc = + rewriter.create(loc, allocType, dynamicSizes); + + rewriter.create(loc, subview, alloc); + + generic->setOperand(inputOperand->getOperandNumber(), alloc); + } + } + + SmallVector vsOps; + func.walk([&](Operation *op) { + if (isa(op)) + vsOps.push_back(op); + }); + + for (Operation *op : vsOps) { + auto vsOp = dyn_cast(op); + if (!vsOp) + continue; + + bool allInitsEqualInput = true; + Value input = vsOp.getDpsInputOperands().front()->get(); + for (OpOperand &init : vsOp.getDpsInitsMutable()) { + if (init.get() != input) { + allInitsEqualInput = false; + break; + } + } + + auto isStridedSubview = [](Value v, MemRefType *outType = nullptr) { + auto subview = v.getDefiningOp(); + if (!subview) + return false; + auto subviewType = dyn_cast(subview.getType()); + if (!subviewType || !hasNonContiguousStrides(subviewType)) + return false; + if (outType) + *outType = subviewType; + return true; + }; + auto createContiguousAlloc = [&](MemRefType type, Value operand, + Location loc) { + SmallVector dynamicSizes; + for (auto [idx, dim] : llvm::enumerate(type.getShape())) { + if (ShapedType::isDynamic(dim)) { + dynamicSizes.push_back( + rewriter.create(loc, operand, idx)); + } + } + auto allocType = + MemRefType::get(type.getShape(), type.getElementType(), + MemRefLayoutAttrInterface{}, type.getMemorySpace()); + return rewriter.create(loc, allocType, dynamicSizes); + }; + + if (allInitsEqualInput) { + MemRefType subviewType; + if (!isStridedSubview(input, &subviewType)) + continue; + + rewriter.setInsertionPoint(op); + Location loc = op->getLoc(); + + Value alloc = createContiguousAlloc(subviewType, input, loc); + rewriter.create(loc, input, alloc); + + rewriter.modifyOpInPlace(op, [&]() { + op->setOperand(0, alloc); + for (OpOperand &init : vsOp.getDpsInitsMutable()) { + if (init.get() == input) + init.set(alloc); + } + }); + + rewriter.setInsertionPointAfter(op); + rewriter.create(loc, alloc, input); + continue; + } + + rewriter.setInsertionPoint(op); + Location loc = op->getLoc(); + + MemRefType inputSubviewType; + if (isStridedSubview(input, &inputSubviewType)) { + Value alloc = createContiguousAlloc(inputSubviewType, input, loc); + rewriter.create(loc, input, alloc); + rewriter.modifyOpInPlace(op, [&]() { op->setOperand(0, alloc); }); + } + + for (OpOperand &init : vsOp.getDpsInitsMutable()) { + MemRefType initSubviewType; + if (!isStridedSubview(init.get(), &initSubviewType)) + continue; + + rewriter.setInsertionPoint(op); + Value originalInit = init.get(); + Value alloc = createContiguousAlloc(initSubviewType, originalInit, loc); + rewriter.create(loc, originalInit, alloc); + rewriter.modifyOpInPlace(op, [&]() { init.set(alloc); }); + + rewriter.setInsertionPointAfter(op); + // init now refers to alloc; copy back to the saved original view. + rewriter.create(loc, alloc, originalInit); + } + } + } + + void runOnOperation() override { + ModuleOp mod = getOperation(); + mod->walk([&](func::FuncOp func) { process(func); }); + } +}; +} // namespace + +std::unique_ptr> +triton::createMaterializeStridedLinalgInputsPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Dialect/Wafer/CMakeLists.txt b/third_party/wafer/lib/Dialect/Wafer/CMakeLists.txt new file mode 100755 index 00000000..b183bcb3 --- /dev/null +++ b/third_party/wafer/lib/Dialect/Wafer/CMakeLists.txt @@ -0,0 +1,12 @@ +add_triton_library(WaferIR + IR/WaferDialect.cpp + IR/WaferOps.cpp + + DEPENDS + WaferTableGen + + LINK_LIBS PUBLIC + MLIRIR +) + +add_subdirectory(Transforms) diff --git a/third_party/wafer/lib/Dialect/Wafer/IR/WaferDialect.cpp b/third_party/wafer/lib/Dialect/Wafer/IR/WaferDialect.cpp new file mode 100755 index 00000000..fc33d37d --- /dev/null +++ b/third_party/wafer/lib/Dialect/Wafer/IR/WaferDialect.cpp @@ -0,0 +1,30 @@ +//===-------------------------- WaferDialect.cpp ---------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Dialect/IR/WaferDialect.h" + +using namespace mlir; +using namespace mlir::wafer; + +/// Dialect creation, the instance will be owned by the context. This is the +/// point of registration of custom types and operations for the dialect. +void WaferDialect::initialize() { + addOperations< +#define GET_OP_LIST +#include "wafer/Dialect/IR/WaferOps.cpp.inc" + >(); +} + +//===----------------------------------------------------------------------===// +// TableGen'd op method definitions +//===----------------------------------------------------------------------===// + +#define GET_OP_CLASSES +#include "wafer/Dialect/IR/WaferEnums.cpp.inc" +#include "wafer/Dialect/IR/WaferOps.cpp.inc" + +#include "wafer/Dialect/IR/WaferDialect.cpp.inc" diff --git a/third_party/wafer/lib/Dialect/Wafer/IR/WaferOps.cpp b/third_party/wafer/lib/Dialect/Wafer/IR/WaferOps.cpp new file mode 100755 index 00000000..202926ea --- /dev/null +++ b/third_party/wafer/lib/Dialect/Wafer/IR/WaferOps.cpp @@ -0,0 +1,10 @@ +//===-------------------------- WaferOps.cpp -------------------------------===// +// +// Copyright (C) 2020-2025 Terapines Technology (Wuhan) Co., Ltd +// All rights reserved. +// +//===----------------------------------------------------------------------===// + +#include "wafer/Dialect/IR/WaferOps.h" +using namespace mlir; +using namespace mlir::wafer; diff --git a/third_party/wafer/lib/Dialect/Wafer/Transforms/CMakeLists.txt b/third_party/wafer/lib/Dialect/Wafer/Transforms/CMakeLists.txt new file mode 100644 index 00000000..700aac73 --- /dev/null +++ b/third_party/wafer/lib/Dialect/Wafer/Transforms/CMakeLists.txt @@ -0,0 +1,5 @@ +add_triton_library(WaferTransforms + InsertBarrierPass.cpp + DEPENDS WaferTransformsPassIncGen WaferTableGen + LINK_LIBS PUBLIC ZTCAnalysis WaferIR MLIRSCFDialect MLIRFuncDialect MLIRMemRefDialect MLIRPass +) diff --git a/third_party/wafer/lib/Dialect/Wafer/Transforms/InsertBarrierPass.cpp b/third_party/wafer/lib/Dialect/Wafer/Transforms/InsertBarrierPass.cpp new file mode 100644 index 00000000..f7bacc56 --- /dev/null +++ b/third_party/wafer/lib/Dialect/Wafer/Transforms/InsertBarrierPass.cpp @@ -0,0 +1,553 @@ +//===----------- InsertBarrierPass.cpp - Wafer Barrier Insertion --------===// +// +// This pass implements the behavior described in `BarrierInsertion.md`: +// - Use a membar-style analysis for SPM(shared memory) hazards to insert +// `wafer::BarrierOp` only when needed. +// - Add a minimal DDR hazard check for WDMA->RDMA (and CPU touching DDR after +// WDMA) to avoid stale reads. +// +// `wafer::BarrierOp` is later lowered in WaferToLLVM to `__Barrier` (equivalent to +// TsmWaitfinish) to keep the runtime interface unchanged. +// Some `tx.barrier` in `scf.for` bodies may be hoisted to the preheader; see +// `hoistTxBarriersFromScfForLoops`. +// +//===----------------------------------------------------------------------===// + +#include "Analysis/Allocation.h" +#include "Analysis/Membar.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/Visitors.h" +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Pass/Pass.h" +#include "wafer/Dialect/IR/WaferDialect.h" +#include "wafer/Transforms/Passes.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "wafer-insert-barrier" + +using namespace mlir; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_INSERTBARRIER +#include "wafer/Transforms/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +static bool isTxDialect(Operation *op) { + auto *d = op->getDialect(); + return d && d->getNamespace() == "wafer"; +} + +static bool isExplicitBarrier(Operation *op) { + return isa(op); +} + +/// Tx ops that touch data (NPU/DMA/compute). Excludes barrier-only tx ops so +/// we can tell "producer tx is inside this loop body" vs "only CPU inside". +static bool isTxDataOp(Operation *op) { + if (!isTxDialect(op)) + return false; + if (isExplicitBarrier(op)) + return false; + return true; +} + +/// True if a `tx.barrier` appears in preorder after `lhsOp` and before `rhsOp` +/// in the same function. After NPU→CPU sync once, further CPU ops are serial; +/// a second barrier on the same chain (e.g. before another `scf.for`) is +/// redundant when IR already has `tx.barrier` between the tx write and rhs. +static bool txBarrierBetweenLhsAndRhs(Operation *lhsOp, Operation *rhsOp) { + auto funcOp = lhsOp->getParentOfType(); + if (!funcOp || funcOp != rhsOp->getParentOfType()) + return false; + + bool seenLhs = false; + bool seenBarrierAfterLhs = false; + bool result = false; + + funcOp.getOperation()->walk([&](Operation *op) { + if (op == rhsOp) { + result = seenBarrierAfterLhs; + return WalkResult::interrupt(); + } + if (op == lhsOp) + seenLhs = true; + if (seenLhs && + (isa(op))) + seenBarrierAfterLhs = true; + return WalkResult::advance(); + }); + return result; +} + +static Value getBaseBuffer(Value v) { + while (auto *defOp = v.getDefiningOp()) { + if (auto op = dyn_cast(defOp)) + v = op.getSource(); + else if (auto op = dyn_cast(defOp)) + v = op.getSource(); + else if (auto op = dyn_cast(defOp)) + v = op.getSource(); + else if (auto op = dyn_cast(defOp)) + v = op.getSource(); + else if (auto op = dyn_cast(defOp)) { + if (v == op->getResult(0)) + v = op.getSource(); + else + break; + } else + break; + } + return v; +} + +static Value traceToOriginMemRef(Value v, unsigned maxDepth = 16) { + if (maxDepth == 0) + return {}; + if (isa(v.getType())) + return getBaseBuffer(v); + Operation *defOp = v.getDefiningOp(); + if (!defOp) + return {}; + + if (auto op = dyn_cast(defOp)) + return traceToOriginMemRef(op.getIn(), maxDepth - 1); + if (auto op = dyn_cast(defOp)) + return traceToOriginMemRef(op.getIn(), maxDepth - 1); + if (auto op = dyn_cast(defOp)) + return traceToOriginMemRef(op.getIn(), maxDepth - 1); + if (auto op = dyn_cast(defOp)) + return traceToOriginMemRef(op.getIn(), maxDepth - 1); + if (auto op = dyn_cast(defOp)) + return getBaseBuffer(op.getSource()); + if (isa(defOp)) { + if (Value r = traceToOriginMemRef(defOp->getOperand(0), maxDepth - 1)) + return r; + return traceToOriginMemRef(defOp->getOperand(1), maxDepth - 1); + } + return {}; +} + +static Value resolveOrigin(Value v) { + if (!v) + return {}; + if (isa(v.getType())) + return getBaseBuffer(v); + return traceToOriginMemRef(v); +} + +static bool +getAllocationOffsetInterval(Value v, + triton::alloc::Interval &interval) { + Value base = resolveOrigin(v); + if (!base) + return false; + + auto alloc = base.getDefiningOp(); + if (!alloc) + return false; + + auto offsetAttr = alloc->getAttrOfType("allocation.offset"); + if (!offsetAttr) + return false; + + int64_t signedOffset = offsetAttr.getInt(); + if (signedOffset < 0) + return false; + + MemRefType allocType = alloc.getType(); + if (!allocType.hasStaticShape()) + return false; + + int64_t numElements = allocType.getNumElements(); + unsigned bitWidth = allocType.getElementTypeBitWidth(); + uint64_t elemBytes = (bitWidth + 7) / 8; + if (numElements < 0 || elemBytes == 0) + return false; + if (static_cast(numElements) > + std::numeric_limits::max() / elemBytes) + return false; + + uint64_t bytes = static_cast(numElements) * elemBytes; + uint64_t offset = static_cast(signedOffset); + uint64_t maxSize = std::numeric_limits::max(); + if (offset > maxSize || bytes > maxSize - offset) + return false; + + interval = triton::alloc::Interval( + static_cast(offset), static_cast(offset + bytes)); + return true; +} + +static bool mayShareAllocationOffset(Value a, Value b) { + triton::alloc::Interval lhs, rhs; + if (!getAllocationOffsetInterval(a, lhs) || + !getAllocationOffsetInterval(b, rhs)) + return false; + return lhs.intersects(rhs); +} + +static bool mayAliasOrigin(Value a, Value b) { + if (!a || !b) + return true; + if (a == b) + return true; + // If both are distinct memref allocs/args, and differ, treat as no-alias. + auto isDistinct = [](Value v) -> bool { + if (auto ba = dyn_cast(v)) + return isa(ba.getType()); + if (auto *def = v.getDefiningOp()) + return isa(def); + return false; + }; + if (isDistinct(a) && isDistinct(b)) + return false; + return true; +} + +/// Collect memref "base" values for SPM-style alias checks (same idea as Membar +/// buffer resolution). +static void collectMemrefBasesForOp(Operation *op, + SmallVectorImpl &bases) { + for (Value v : op->getOperands()) { + if (Value base = resolveOrigin(v)) + bases.push_back(base); + else if (isa(v.getType())) + bases.push_back(getBaseBuffer(v)); + } +} + +/// True if `producer` and `consumer` may touch the same SPM allocation. +static bool mayShareSpmMemref(Operation *producer, Operation *consumer) { + SmallVector pa, pb; + collectMemrefBasesForOp(producer, pa); + collectMemrefBasesForOp(consumer, pb); + if (pa.empty() || pb.empty()) + return false; + for (Value a : pa) + for (Value b : pb) + if (mayAliasOrigin(a, b) || mayShareAllocationOffset(a, b)) + return true; + return false; +} + +static bool isCpuDataOp(Operation *op) { + return op && !isTxDialect(op) && !triton::membar::isPureAddressOp(op); +} + +static bool touchesSpmAllocation(Operation *op, + triton::alloc::Allocation *allocation) { + SmallVector bases; + collectMemrefBasesForOp(op, bases); + for (Value base : bases) + if (!allocation->getBufferIds(base).empty()) + return true; + return false; +} + +/// True if, in one iteration template, some CPU-side non-tx data op appears +/// before a tx data op in preorder and they may share SPM. Then iteration i+1's +/// CPU can run before iteration i's async NPU finishes — need a barrier at the +/// end of the body (before yield) so the next iteration's CPU waits. +static bool cpuPrecedesTxOnSharedSpmInSameBody(scf::ForOp forOp) { + SmallVector cpuMemOps; + bool found = false; + forOp.getBody()->walk([&](Operation *op) { + if (isCpuDataOp(op)) { + cpuMemOps.push_back(op); + return WalkResult::advance(); + } + if (isTxDataOp(op)) { + for (Operation *cpu : cpuMemOps) { + if (mayShareSpmMemref(cpu, op)) { + found = true; + return WalkResult::interrupt(); + } + } + } + return WalkResult::advance(); + }); + return found; +} + +/// Insert `tx.barrier` before `scf.yield` when loop-carried SPM sync is needed. +static void insertLoopCarriedSpmBarriers(ModuleOp mod) { + mod.walk([&](scf::ForOp forOp) { + if (!cpuPrecedesTxOnSharedSpmInSameBody(forOp)) + return; + auto yield = dyn_cast(forOp.getBody()->getTerminator()); + if (!yield) + return; + if (Operation *prev = yield->getPrevNode(); + prev && isa(prev)) + return; + OpBuilder b(yield); + b.create(yield->getLoc()); + }); +} + +/// Hoist barrier out of `forOp` only when the hazard is not "tx producer and +/// CPU consumer both inside this loop, same SPM". Membar inserts the barrier +/// immediately before `consumer`; we find tx ops before the barrier that may +/// produce the memref `consumer` reads — if that tx is in the loop region, keep +/// the barrier inside (per-iteration sync). If the tx producer is outside the +/// loop and only `consumer` is inside, one barrier before the loop is correct. +static bool shouldHoistBarrierFromLoop(scf::ForOp forOp, + wafer::BarrierOp barrier) { + Operation *consumer = barrier->getNextNode(); + if (!consumer) + return false; + + Operation *pairedTx = nullptr; + bool anyTxDataBeforeBarrier = false; + + forOp.getBody()->walk([&](Operation *op) { + if (op == barrier) + return WalkResult::interrupt(); + if (isTxDataOp(op)) { + anyTxDataBeforeBarrier = true; + if (mayShareSpmMemref(op, consumer)) + pairedTx = op; + } + return WalkResult::advance(); + }); + + if (pairedTx && forOp->isAncestor(pairedTx) && forOp->isAncestor(consumer)) + return false; + + // Could not match producer↔consumer buffers (e.g. i64 address only): if any + // tx data op precedes the barrier in the loop body, keep barriers inside. + if (!pairedTx && anyTxDataBeforeBarrier) + return false; + + return true; +} + +/// Hoist `tx.barrier` from a `scf.for` body to the preheader only when +/// `shouldHoistBarrierFromLoop` says so (see dependency-based rule there). +static void hoistTxBarriersFromScfForLoops(ModuleOp mod) { + mod.walk([&](scf::ForOp forOp) { + SmallVector barriers; + forOp.getBody()->walk([&](wafer::BarrierOp b) { barriers.push_back(b); }); + if (barriers.empty()) + return; + + SmallVector toHoist; + for (wafer::BarrierOp br : barriers) { + if (!shouldHoistBarrierFromLoop(forOp, br)) + continue; + // Barrier immediately before scf.yield ends the iteration after NPU work; + // must stay in the body (per-iter sync). Hoisting would run once outside. + if (Operation *n = br->getNextNode()) + if (isa(n)) + continue; + toHoist.push_back(br); + } + if (toHoist.empty()) + return; + + OpBuilder b(forOp); + b.setInsertionPoint(forOp); + b.create(forOp.getLoc()); + for (wafer::BarrierOp br : toHoist) + br.erase(); + }); +} + +static bool +topLevelOpHasCpuSpmAccessBeforeBarrier(Operation *topLevelOp, + triton::alloc::Allocation *allocation) { + bool foundCpuSpmAccess = false; + topLevelOp->walk([&](Operation *op) { + if (op != topLevelOp && isExplicitBarrier(op)) { + return WalkResult::interrupt(); + } + if (!isCpuDataOp(op)) + return WalkResult::advance(); + if (touchesSpmAllocation(op, allocation)) { + foundCpuSpmAccess = true; + return WalkResult::interrupt(); + } + return WalkResult::advance(); + }); + return foundCpuSpmAccess; +} + +static void insertKernelEntryBarriersBeforeCpuSpmUse( + ModuleOp mod, triton::alloc::ModuleAllocation &moduleAlloc) { + mod.walk([&](func::FuncOp fn) { + if (fn.isExternal() || fn.getBody().empty()) + return; + auto *allocation = moduleAlloc.getFuncData(fn); + if (!allocation) + return; + + for (Operation &op : fn.getBody().front().getOperations()) { + if (isExplicitBarrier(&op)) + return; + if (isa(&op)) + return; + if (topLevelOpHasCpuSpmAccessBeforeBarrier(&op, allocation)) { + if (Operation *prev = op.getPrevNode(); + prev && isExplicitBarrier(prev)) { + return; + } + OpBuilder b(&op); + b.create(op.getLoc()); + return; + } + } + }); +} + +class InsertBarrierPass + : public triton::impl::InsertBarrierBase { + using InsertBarrierBase::InsertBarrierBase; + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + ModuleOp mod = getOperation(); + + // 1) SPM hazards: membar + filter. One `tx.barrier` per NPU→CPU sync chain: + // NPU→barrier→CPU→CPU needs no second barrier between CPUs. Also suppress + // CPU↔CPU (program order). If a `tx.barrier` already appears on the path + // from the tx op (lhs) to this CPU op (rhs), suppress (see + // `txBarrierBetweenLhsAndRhs`). CPU before tx: host finished prior work. + triton::alloc::ModuleAllocation moduleAlloc(mod); + auto filter = [](Operation *lhsOp, Operation *rhsOp, + triton::membar::MembarHazardKind) -> bool { + // Return true means "suppress barrier". + if (isTxDialect(lhsOp) && isTxDialect(rhsOp)) + return true; + if (!isTxDialect(lhsOp) && isTxDialect(rhsOp)) + return true; + if (!isTxDialect(lhsOp) && !isTxDialect(rhsOp)) + return true; + if (isTxDialect(lhsOp) && !isTxDialect(rhsOp) && + txBarrierBetweenLhsAndRhs(lhsOp, rhsOp)) + return true; + return false; + }; + triton::membar::ModuleMembarAnalysis membar(&moduleAlloc, filter); + membar.run(); + + // 2) DDR hazards: minimal WDMA->RDMA and CPU-touching-DDR-after-WDMA. + // We track pending WDMA writes by their origin memref and insert a + // `wafer::BarrierOp` before any RDMA/CPU op that may read that same origin. + mod.walk([&](func::FuncOp fn) { + SmallVector pendingDdrWrites; + auto clearPending = [&]() { pendingDdrWrites.clear(); }; + + auto markWdmaWrite = [&](Operation *op) { + Value tgt; + if (auto w = dyn_cast(op)) + tgt = w.getTarget(); + else if (auto w = dyn_cast(op)) + tgt = w.getTarget(); + else if (auto w = dyn_cast(op)) + tgt = w.getTarget(); + Value origin = resolveOrigin(tgt); + if (origin) + pendingDdrWrites.push_back(origin); + }; + + auto needsBarrierForDdrRead = [&](Value addrLike) -> bool { + Value origin = resolveOrigin(addrLike); + if (!origin) + return false; + for (Value w : pendingDdrWrites) + if (mayAliasOrigin(origin, w)) + return true; + return false; + }; + + auto insertBarrierBefore = [&](Operation *op) { + OpBuilder b(op); + b.create(op->getLoc()); + clearPending(); + }; + + for (Block &block : fn.getBody()) { + // Pending WDMA writes are only meaningful within straight-line code. + // Start each block conservatively with an empty pending set. + clearPending(); + for (Operation &op : + llvm::make_early_inc_range(block.getOperations())) { + if (isa( + &op)) { + clearPending(); + continue; + } + + if (isa(&op)) { + markWdmaWrite(&op); + continue; + } + + if (isa(&op)) { + Value src; + if (auto r = dyn_cast(&op)) + src = r.getSource(); + else if (auto r = dyn_cast(&op)) + src = r.getSource(); + else if (auto r = dyn_cast(&op)) + src = r.getSource(); + if (!pendingDdrWrites.empty() && needsBarrierForDdrRead(src)) { + insertBarrierBefore(&op); + } + continue; + } + + // CPU ops: any non-tx op that uses a pending DDR region triggers a + // barrier. Skip memref/arith/scf address prep (same as SPM membar). + if (!isTxDialect(&op) && !pendingDdrWrites.empty() && + !triton::membar::isPureAddressOp(&op)) { + bool conflict = false; + for (Value operand : op.getOperands()) { + if (needsBarrierForDdrRead(operand)) { + conflict = true; + break; + } + } + if (conflict) + insertBarrierBefore(&op); + } + } + } + }); + + // 3) Loop-carried SPM: same body has CPU memref access before a tx op on + // aliasing SPM — next iter's CPU can overlap previous iter's async NPU. + // Sync at end of body (before scf.yield); do not hoist that barrier. + insertLoopCarriedSpmBarriers(mod); + + // 4) Kernel-boundary sync: different kernels may reuse the same physical + // SPM. If a kernel starts with CPU-side SPM access before any explicit + // barrier, insert one barrier before that first top-level CPU region (for + // example, before an enclosing `scf.for`). + insertKernelEntryBarriersBeforeCpuSpmUse(mod, moduleAlloc); + + // 5) Hoist barriers out of `scf.for` only when sync is outer-tx vs + // inner-CPU (see `hoistTxBarriersFromScfForLoops`); skip yield-adjacent. + hoistTxBarriersFromScfForLoops(mod); + } +}; + +} // namespace + +std::unique_ptr> triton::createInsertBarrierPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/lib/Registrar/CMakeLists.txt b/third_party/wafer/lib/Registrar/CMakeLists.txt new file mode 100755 index 00000000..af8972aa --- /dev/null +++ b/third_party/wafer/lib/Registrar/CMakeLists.txt @@ -0,0 +1 @@ +add_triton_library(Registrar Registrar.cc) diff --git a/third_party/wafer/lib/Registrar/Registrar.cc b/third_party/wafer/lib/Registrar/Registrar.cc new file mode 100755 index 00000000..0a30b824 --- /dev/null +++ b/third_party/wafer/lib/Registrar/Registrar.cc @@ -0,0 +1,15 @@ +#include "flagtree/Common/UnifiedHardware.h" + +class TsingmicroUnifiedHardware : public mlir::flagtree::UnifiedHardware { +public: + int getDMATag() const override; + int getSharedMemoryTag() const override; +}; + +int TsingmicroUnifiedHardware::getDMATag() const { return 11; } +int TsingmicroUnifiedHardware::getSharedMemoryTag() const { return 8; } + +std::unique_ptr +mlir::flagtree::createUnifiedHardwareManager() { + return std::make_unique(); +} diff --git a/third_party/wafer/name.conf b/third_party/wafer/name.conf new file mode 100755 index 00000000..89a6ae8d --- /dev/null +++ b/third_party/wafer/name.conf @@ -0,0 +1 @@ +wafer diff --git a/third_party/wafer/patches/triton/cache_without_vendor_imports.patch b/third_party/wafer/patches/triton/cache_without_vendor_imports.patch new file mode 100644 index 00000000..f2cbfd59 --- /dev/null +++ b/third_party/wafer/patches/triton/cache_without_vendor_imports.patch @@ -0,0 +1,51 @@ +diff --git a/python/triton/runtime/cache.py b/python/triton/runtime/cache.py +index 0442f00..fd14838 100644 +--- a/python/triton/runtime/cache.py ++++ b/python/triton/runtime/cache.py +@@ -268,9 +268,20 @@ def make_so_cache_key(version_hash, signature, constants, ids, **kwargs): + return _base32(key) + + ++def _module_files(path, prefix): ++ """Find package files without executing vendor package initializers.""" ++ import pkgutil ++ ++ for lib in pkgutil.iter_modules([path], prefix=prefix): ++ spec = lib.module_finder.find_spec(lib.name) ++ yield spec.origin ++ if lib.ispkg: ++ for child_path in spec.submodule_search_locations: ++ yield from _module_files(child_path, lib.name + ".") ++ ++ + @functools.lru_cache() + def triton_key(): +- import pkgutil + TRITON_PATH = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + contents = [] + # frontend +@@ -282,8 +293,8 @@ def triton_key(): + (os.path.join(TRITON_PATH, "backends"), "triton.backends."), + ] + for path, prefix in path_prefixes: +- for lib in pkgutil.walk_packages([path], prefix=prefix): +- with open(lib.module_finder.find_spec(lib.name).origin, "rb") as f: ++ for filename in _module_files(path, prefix): ++ with open(filename, "rb") as f: + contents += [hashlib.sha256(f.read()).hexdigest()] + + # backend +@@ -298,8 +309,11 @@ def triton_key(): + contents.append(libtriton_hash.hexdigest()) + # language + language_path = os.path.join(TRITON_PATH, 'language') +- for lib in pkgutil.walk_packages([language_path], prefix="triton.language."): +- with open(lib.module_finder.find_spec(lib.name).origin, "rb") as f: ++ # walk_packages imports packages to discover their children. CANN's import ++ # replaces tensor members and math functions, including in Wafer kernels. ++ # Hash every vendor's files, but leave importing them to target selection. ++ for filename in _module_files(language_path, "triton.language."): ++ with open(filename, "rb") as f: + contents += [hashlib.sha256(f.read()).hexdigest()] + return f'{__version__}' + '-'.join(contents) + diff --git a/third_party/wafer/patches/triton/profiles.json b/third_party/wafer/patches/triton/profiles.json new file mode 100644 index 00000000..c51b6ed7 --- /dev/null +++ b/third_party/wafer/patches/triton/profiles.json @@ -0,0 +1,32 @@ +{ + "triton_commit": "c3c476f357f1e9768ea4e45aa5c17528449ab9ef", + "profiles": { + "ascend": [ + "patch/triton/CMakeLists_txt.patch", + "patch/triton/lib_Dialect_Triton_IR_Ops_cpp.patch", + "patch/triton/lib_Dialect_Triton_IR_Traits_cpp.patch", + "patch/triton/python_src_ir_cc.patch", + "patch/triton/python_src_ir_h.patch", + "patch/triton/python_triton__utils_py.patch", + "patch/triton/python_triton_backends_compiler_py.patch", + "patch/triton/python_triton_compiler_code_generator_py.patch", + "patch/triton/python_triton_compiler_compiler_py.patch", + "patch/triton/python_triton_language_semantic_py.patch", + "patch/triton/setup_py.patch", + "patch/triton/unittest_googletest_cmake.patch" + ], + "wafer-tools": [ + "third_party/wafer/patches/triton/wafer_builder_optional_gluon.patch", + "third_party/wafer/patches/triton/wafer_proton_backend_filter.patch", + "third_party/wafer/patches/triton/python_triton_compiler_optional_gluon_py.patch", + "third_party/wafer/patches/triton/python_triton_jit_optional_gluon.patch" + ], + "wafer-frontend": [ + "third_party/wafer/patches/triton/wafer_builder_optional_gluon.patch", + "third_party/wafer/patches/triton/wafer_proton_backend_filter.patch", + "third_party/wafer/patches/triton/python_triton_compiler_optional_gluon_py.patch", + "third_party/wafer/patches/triton/python_triton_jit_optional_gluon.patch", + "third_party/wafer/patches/triton/cache_without_vendor_imports.patch" + ] + } +} diff --git a/third_party/wafer/patches/triton/python_triton_compiler_optional_gluon_py.patch b/third_party/wafer/patches/triton/python_triton_compiler_optional_gluon_py.patch new file mode 100644 index 00000000..ca6bd11a --- /dev/null +++ b/third_party/wafer/patches/triton/python_triton_compiler_optional_gluon_py.patch @@ -0,0 +1,49 @@ +diff --git a/python/triton/compiler/code_generator.py b/python/triton/compiler/code_generator.py +index 176b6b5150..19dc08d6d7 100644 +--- a/python/triton/compiler/code_generator.py ++++ b/python/triton/compiler/code_generator.py +@@ -12,7 +12,11 @@ from types import ModuleType + from typing import Any, Callable, Dict, Optional, Tuple, Type, Union, Iterable, List + + from .. import knobs, language +-from .._C.libtriton import ir, gluon_ir ++from .._C.libtriton import ir ++try: ++ from .._C.libtriton import gluon_ir ++except ImportError: ++ gluon_ir = None + from ..language import constexpr, str_to_ty, tensor, tuple as tl_tuple + from ..language.core import _unwrap_if_constexpr, base_value, base_type + # ideally we wouldn't need any runtime component +@@ -300,6 +304,9 @@ class CodeGenerator(ast.NodeVisitor): + self.context = context + self.is_gluon = is_gluon + if is_gluon: ++ if gluon_ir is None: ++ raise RuntimeError("Gluon kernels are not supported in this build: the gluon_ir " ++ "bindings were not compiled. Rebuild Triton with Gluon enabled.") + from triton.experimental.gluon.language._semantic import GluonSemantic + self.builder = gluon_ir.GluonOpBuilder(context) + self.semantic = GluonSemantic(self.builder) +@@ -1569,15 +1576,18 @@ class CodeGenerator(ast.NodeVisitor): + + return ret + +- from ..experimental.gluon import language as ttgl + statically_implemented_functions: Dict[object, Callable[[ast.Call], Any]] = { + language.core.static_assert: execute_static_assert, + language.core.static_print: static_executor(print), +- ttgl.static_assert: execute_static_assert, +- ttgl.static_print: static_executor(print), + int: static_executor(int), + len: static_executor(len), + } ++ if gluon_ir is not None: ++ from ..experimental.gluon import language as ttgl ++ statically_implemented_functions.update({ ++ ttgl.static_assert: execute_static_assert, ++ ttgl.static_print: static_executor(print), ++ }) + + + def ast_to_ttir(fn, src, context, options, codegen_fns, module_map, module=None): diff --git a/third_party/wafer/patches/triton/python_triton_jit_optional_gluon.patch b/third_party/wafer/patches/triton/python_triton_jit_optional_gluon.patch new file mode 100644 index 00000000..a1b8aa8b --- /dev/null +++ b/third_party/wafer/patches/triton/python_triton_jit_optional_gluon.patch @@ -0,0 +1,19 @@ +diff --git a/python/triton/runtime/jit.py b/python/triton/runtime/jit.py +--- a/python/triton/runtime/jit.py ++++ b/python/triton/runtime/jit.py +@@ -348,7 +348,14 @@ specialize_impl_cache = [] + def create_specialize_impl(specialize_extra): + + from ..language import constexpr +- from triton.experimental.gluon.nvidia.hopper import TensorDescriptor as GluonTensorDescriptor ++ try: ++ from triton._C.libtriton import gluon_ir # noqa: F401 ++ except ImportError: ++ # Wafer builds can omit Gluon IR. Ordinary JIT argument binding must ++ # still work without importing Gluon's Python semantic module. ++ GluonTensorDescriptor = () ++ else: ++ from triton.experimental.gluon.nvidia.hopper import TensorDescriptor as GluonTensorDescriptor + + def specialize_impl(arg, is_const=False, specialize_value=True, align=True): + if arg is None: diff --git a/third_party/wafer/patches/triton/wafer_builder_optional_gluon.patch b/third_party/wafer/patches/triton/wafer_builder_optional_gluon.patch new file mode 100644 index 00000000..16cb6492 --- /dev/null +++ b/third_party/wafer/patches/triton/wafer_builder_optional_gluon.patch @@ -0,0 +1,116 @@ +diff --git a/CMakeLists.txt b/CMakeLists.txt +index 47a0f3b175..fa5cdd96de 100644 +--- a/CMakeLists.txt ++++ b/CMakeLists.txt +@@ -19,6 +19,7 @@ list(APPEND CMAKE_MODULE_PATH "${CMAKE_CURRENT_SOURCE_DIR}/cmake") + + # Options + option(TRITON_BUILD_PYTHON_MODULE "Build Python Triton bindings" OFF) ++option(TRITON_BUILD_GLUON_IR "Build Gluon IR Python bindings" ON) + option(TRITON_BUILD_PROTON "Build the Triton Proton profiler" ON) + option(TRITON_BUILD_UT "Build C++ Triton Unit Tests" ON) + option(TRITON_BUILD_WITH_CCACHE "Build with ccache (if available)" ON) +@@ -283,12 +284,18 @@ if(TRITON_BUILD_PYTHON_MODULE) + + set(TRITON_BACKENDS_TUPLE "(${TRITON_BACKENDS_TUPLE})") + add_compile_definitions(TRITON_BACKENDS_TUPLE=${TRITON_BACKENDS_TUPLE}) +- add_library(triton SHARED ${PYTHON_SRC_PATH}/main.cc +- ${PYTHON_SRC_PATH}/ir.cc +- ${PYTHON_SRC_PATH}/gluon_ir.cc +- ${PYTHON_SRC_PATH}/passes.cc +- ${PYTHON_SRC_PATH}/interpreter.cc +- ${PYTHON_SRC_PATH}/llvm.cc) ++ set(TRITON_PYTHON_SOURCES ++ ${PYTHON_SRC_PATH}/main.cc ++ ${PYTHON_SRC_PATH}/ir.cc ++ ${PYTHON_SRC_PATH}/passes.cc ++ ${PYTHON_SRC_PATH}/interpreter.cc ++ ${PYTHON_SRC_PATH}/llvm.cc) ++ if(TRITON_BUILD_GLUON_IR) ++ list(APPEND TRITON_PYTHON_SOURCES ${PYTHON_SRC_PATH}/gluon_ir.cc) ++ endif() ++ add_library(triton SHARED ${TRITON_PYTHON_SOURCES}) ++ target_compile_definitions(triton PRIVATE ++ TRITON_BUILD_GLUON_IR=$) + + # Link triton with its dependencies + target_link_libraries(triton PRIVATE ${TRITON_LIBRARIES}) +diff --git a/python/src/ir.cc b/python/src/ir.cc +index d79a9e70f8..b850c83484 100644 +--- a/python/src/ir.cc ++++ b/python/src/ir.cc +@@ -41,6 +41,14 @@ + + #include "llvm/ADT/SmallVector.h" + ++namespace ir { ++static std::unique_ptr> builderClass; ++ ++pybind11::class_ *getBuilderClass() { ++ return builderClass.get(); ++} ++} // namespace ir ++ + void setAsyncTaskIds(mlir::Operation *op, + llvm::ArrayRef asyncTaskIds) { + llvm::SmallVector sortedAsyncTaskIds(asyncTaskIds.begin(), +@@ -777,9 +785,10 @@ void init_triton_ir(py::module &&m) { + + py::class_(m, "InsertPoint", py::module_local()); + +- py::class_(m, "builder", py::module_local(), +- py::dynamic_attr()) +- .def(py::init()) ++ ir::builderClass = std::make_unique>( ++ m, "builder", py::module_local(), py::dynamic_attr()); ++ ir::builderClass ++ ->def(py::init()) + .def("get_op_builder", &TritonOpBuilder::getBuilder, ret::reference) + // getters + .def("create_module", +diff --git a/python/src/ir.h b/python/src/ir.h +index e1f9ce8481..5fc9c48573 100644 +--- a/python/src/ir.h ++++ b/python/src/ir.h +@@ -4,6 +4,10 @@ + #include "llvm/ADT/ArrayRef.h" + #include + ++namespace pybind11 { ++template class class_; ++} ++ + typedef int AsyncTaskId; + void setAsyncTaskIds(mlir::Operation *op, + llvm::ArrayRef asyncTaskIds); +@@ -103,3 +107,7 @@ private: + bool lineInfoEnabled = + !mlir::triton::tools::getBoolEnv("TRITON_DISABLE_LINE_INFO"); + }; ++ ++namespace ir { ++pybind11::class_ *getBuilderClass(); ++} +diff --git a/python/src/main.cc b/python/src/main.cc +index 6e8f15bd30..95107d7b78 100644 +--- a/python/src/main.cc ++++ b/python/src/main.cc +@@ -42,7 +42,9 @@ void init_triton_llvm(pybind11::module &&m); + void init_triton_interpreter(pybind11::module &&m); + void init_triton_passes(pybind11::module &&m); + void init_triton_stacktrace_hook(pybind11::module &m); ++#if TRITON_BUILD_GLUON_IR + void init_gluon_ir(pybind11::module &&m); ++#endif + FOR_EACH_P(DECLARE_BACKEND, TRITON_BACKENDS_TUPLE) + + PYBIND11_MODULE(libtriton, m) { +@@ -53,6 +55,8 @@ PYBIND11_MODULE(libtriton, m) { + init_triton_passes(m.def_submodule("passes")); + init_triton_interpreter(m.def_submodule("interpreter")); + init_triton_llvm(m.def_submodule("llvm")); ++#if TRITON_BUILD_GLUON_IR + init_gluon_ir(m.def_submodule("gluon_ir")); ++#endif + FOR_EACH_P(INIT_BACKEND, TRITON_BACKENDS_TUPLE) + } diff --git a/third_party/wafer/patches/triton/wafer_proton_backend_filter.patch b/third_party/wafer/patches/triton/wafer_proton_backend_filter.patch new file mode 100644 index 00000000..cd6fd30c --- /dev/null +++ b/third_party/wafer/patches/triton/wafer_proton_backend_filter.patch @@ -0,0 +1,82 @@ +diff --git a/third_party/proton/Dialect/CMakeLists.txt b/third_party/proton/Dialect/CMakeLists.txt +index 9ef6bdc31a..38ce14fcc2 100644 +--- a/third_party/proton/Dialect/CMakeLists.txt ++++ b/third_party/proton/Dialect/CMakeLists.txt +@@ -3,6 +3,20 @@ include_directories(${CMAKE_CURRENT_BINARY_DIR}/include) + add_subdirectory(include) + add_subdirectory(lib) + if(TRITON_BUILD_PYTHON_MODULE) +- add_triton_plugin(TritonProton ${CMAKE_CURRENT_SOURCE_DIR}/triton_proton.cc LINK_LIBS ProtonToProtonGPU ProtonGPUToLLVM ProtonAMDGPUToLLVM ProtonNVIDIAGPUToLLVM ProtonAnalysis) ++ set(PROTON_BACKEND_LIBS) ++ set(PROTON_AMD_ENABLED 0) ++ set(PROTON_NVIDIA_ENABLED 0) ++ if("nvidia" IN_LIST TRITON_CODEGEN_BACKENDS) ++ list(APPEND PROTON_BACKEND_LIBS ProtonNVIDIAGPUToLLVM) ++ set(PROTON_NVIDIA_ENABLED 1) ++ endif() ++ if("amd" IN_LIST TRITON_CODEGEN_BACKENDS) ++ list(APPEND PROTON_BACKEND_LIBS ProtonAMDGPUToLLVM) ++ set(PROTON_AMD_ENABLED 1) ++ endif() ++ add_triton_plugin(TritonProton ${CMAKE_CURRENT_SOURCE_DIR}/triton_proton.cc LINK_LIBS ProtonToProtonGPU ProtonGPUToLLVM ${PROTON_BACKEND_LIBS} ProtonAnalysis) ++ target_compile_definitions(TritonProton PRIVATE ++ TRITON_PROTON_ENABLE_AMD=${PROTON_AMD_ENABLED} ++ TRITON_PROTON_ENABLE_NVIDIA=${PROTON_NVIDIA_ENABLED}) + target_link_libraries(TritonProton PRIVATE Python3::Module pybind11::headers) + endif() +diff --git a/third_party/proton/Dialect/lib/ProtonGPUToLLVM/CMakeLists.txt b/third_party/proton/Dialect/lib/ProtonGPUToLLVM/CMakeLists.txt +index e3d89b8c59..1a06f25498 100644 +--- a/third_party/proton/Dialect/lib/ProtonGPUToLLVM/CMakeLists.txt ++++ b/third_party/proton/Dialect/lib/ProtonGPUToLLVM/CMakeLists.txt +@@ -13,5 +13,9 @@ add_triton_library(ProtonGPUToLLVM + ProtonAnalysis + ) + +-add_subdirectory(ProtonNvidiaGPUToLLVM) +-add_subdirectory(ProtonAMDGPUToLLVM) ++if("nvidia" IN_LIST TRITON_CODEGEN_BACKENDS) ++ add_subdirectory(ProtonNvidiaGPUToLLVM) ++endif() ++if("amd" IN_LIST TRITON_CODEGEN_BACKENDS) ++ add_subdirectory(ProtonAMDGPUToLLVM) ++endif() +diff --git a/third_party/proton/Dialect/triton_proton.cc b/third_party/proton/Dialect/triton_proton.cc +index 00ecb3d9c7..5f8694055d 100644 +--- a/third_party/proton/Dialect/triton_proton.cc ++++ b/third_party/proton/Dialect/triton_proton.cc +@@ -1,7 +1,11 @@ + #include "Analysis/ScopeIdAllocation.h" + #include "Conversion/ProtonGPUToLLVM/Passes.h" ++#if TRITON_PROTON_ENABLE_AMD + #include "Conversion/ProtonGPUToLLVM/ProtonAMDGPUToLLVM/Passes.h" ++#endif ++#if TRITON_PROTON_ENABLE_NVIDIA + #include "Conversion/ProtonGPUToLLVM/ProtonNvidiaGPUToLLVM/Passes.h" ++#endif + #include "Conversion/ProtonToProtonGPU/Passes.h" + #include "Dialect/Proton/IR/Dialect.h" + #include "Dialect/ProtonGPU/IR/Dialect.h" +@@ -96,17 +100,23 @@ void init_triton_proton(py::module &&m) { + profileScratchSize, profileScratchAlignment, clkExt)); + }); + ++#if TRITON_PROTON_ENABLE_NVIDIA + ADD_PASS_WRAPPER_0("add_convert_proton_nvidia_gpu_to_llvm", + proton::gpu::createConvertProtonNvidiaGPUToLLVMPass); ++#endif ++#if TRITON_PROTON_ENABLE_AMD + ADD_PASS_WRAPPER_1("add_convert_proton_amd_gpu_to_llvm", + proton::gpu::createConvertProtonAMDGPUToLLVMPass, + const std::string &); ++#endif + ADD_PASS_WRAPPER_0("add_allocate_proton_shared_memory", + proton::gpu::createAllocateProtonSharedMemoryPass); + ADD_PASS_WRAPPER_0("add_allocate_proton_global_scratch_buffer", + proton::gpu::createAllocateProtonGlobalScratchBufferPass); + ADD_PASS_WRAPPER_0("add_schedule_buffer_store", + proton::gpu::createScheduleBufferStorePass); ++#if TRITON_PROTON_ENABLE_AMD + ADD_PASS_WRAPPER_0("add_sched_barriers", + proton::gpu::createAddSchedBarriersPass); ++#endif + } diff --git a/third_party/wafer/profiler/CMakeLists.txt b/third_party/wafer/profiler/CMakeLists.txt new file mode 100755 index 00000000..4d80a7fa --- /dev/null +++ b/third_party/wafer/profiler/CMakeLists.txt @@ -0,0 +1,70 @@ +cmake_minimum_required(VERSION 3.18) + +set(CMAKE_EXPORT_COMPILE_COMMANDS ON) +set(CMAKE_BUILD_WITH_INSTALL_RPATH ON) + +# Set LLVM_SYSPATH from environment variable +if(NOT DEFINED LLVM_SYSPATH) + if(DEFINED ENV{LLVM_SYSPATH}) + set(LLVM_SYSPATH $ENV{LLVM_SYSPATH}) + else() + message(FATAL_ERROR "LLVM_SYSPATH environment variable is not defined") + endif() +endif() + +# Project name and version +project(Profiler LANGUAGES CXX C) + +# Define standard include directories +include_directories(${LLVM_SYSPATH}/include/) + +# Set build type default +if(NOT CMAKE_BUILD_TYPE) + set(CMAKE_BUILD_TYPE Release CACHE STRING "Build type (default Release)" FORCE) +endif() + +# Collect all source files from the vendor directory +file(GLOB_RECURSE SOURCES ./*.cpp) + +set(CMAKE_SYSTEM_NAME Generic) +set(CMAKE_C_COMPILER ${LLVM_SYSPATH}/bin/clang) +set(CMAKE_CXX_COMPILER ${LLVM_SYSPATH}/bin/clang++) + +# Add the library target +set(BIN_NAME wafer-profiler) +add_executable(${BIN_NAME} ${SOURCES}) + +# Apply RISC-V specific settings to our target +set(COMPILE_OPTIONS + -fPIC + -fno-rtti + --std=c++17 + -O2 +) +target_compile_options(${BIN_NAME} PRIVATE ${COMPILE_OPTIONS}) + +# target_link_directories(${BIN_NAME} ${LLVM_SYSPATH}/lib) +target_link_libraries(${BIN_NAME} PRIVATE + LLVMIRReader + LLVMAsmParser + LLVMBitReader + LLVMCore + LLVMBinaryFormat + LLVMDemangle + LLVMRemarks + LLVMBitstreamReader + LLVMSupport + LLVMTargetParser +) + +# Set properties for the library +set_target_properties(${BIN_NAME} PROPERTIES + POSITION_INDEPENDENT_CODE ON + RUNTIME_OUTPUT_DIRECTORY ${CMAKE_BINARY_DIR}/bin +) + + +# Install targets +install(TARGETS ${BIN_NAME} + RUNTIME DESTINATION ${INSTALL_WAFER_DIR}/bin +) diff --git a/third_party/wafer/profiler/profiler.cpp b/third_party/wafer/profiler/profiler.cpp new file mode 100755 index 00000000..899a6267 --- /dev/null +++ b/third_party/wafer/profiler/profiler.cpp @@ -0,0 +1,232 @@ +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/StringRef.h" +#include "llvm/IR/DerivedTypes.h" +#include "llvm/IR/IRBuilder.h" +#include "llvm/IR/Instructions.h" +#include "llvm/IR/LLVMContext.h" +#include "llvm/IR/Module.h" +#include "llvm/IR/Verifier.h" +#include "llvm/IRReader/IRReader.h" +#include "llvm/Support/Alignment.h" +#include "llvm/Support/CommandLine.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/SourceMgr.h" +#include +#include +#include +#include +using namespace llvm; + +#define DEBUG_TYPE "profiler" + +const static std::string kernel_start = "kernel_s"; +const static std::string kernel_end = "kernel_e"; + +static std::map orderDesc; + +cl::OptionCategory ProfileCommon("Common profile options"); +static cl::OptionCategory *ProfileCategories[] = {&ProfileCommon}; + +cl::list + TracePoints("trace-points", cl::CommaSeparated, + cl::desc("Function name which need do profiling tracing"), + cl::value_desc("func1,func2,func3,..."), cl::Prefix, + cl::cat(ProfileCommon)); + +cl::opt InputIR(cl::Positional, cl::desc("The input IR file"), + cl::cat(ProfileCommon)); + +cl::opt OutFile("o", cl::desc("Set the output IR file"), + cl::value_desc("filename"), + cl::init("/tmp/temp.ll"), cl::cat(ProfileCommon)); + +cl::opt Index("index", cl::desc("The ir index to be processed"), + cl::value_desc("int"), cl::init(0), cl::cat(ProfileCommon)); + +/// Handling Config Generator command options with LLVM CommandLine facilities +void initCommandLine(int argc, char **argv) { + // Hide LLVM command options, display Config Generator only command options. + cl::HideUnrelatedOptions(ArrayRef(ProfileCategories)); + + // Parse command line options. + cl::ParseCommandLineOptions(argc, argv, "Profile command line usage\n\n "); +} + +void parse(const char *path, std::unique_ptr &program, + LLVMContext &ctx) { + SMDiagnostic error; + + program = parseIRFile(path, error, ctx); + if (!program) { + printf("Failed to parse IR file\n"); + error.print(path, errs()); + + exit(-1); + } +} + +void writeOrderDesc(std::map desc) { + std::ofstream outFile("profile_desc.ini"); + if (!outFile.is_open()) { + printf("Failed to open file: profile_desc.ini\n"); + return; + } + + outFile << "[" << "profile_order_desc" << "]" << std::endl; + for (const auto &pair : desc) { + outFile << "o" << pair.first << "=" << pair.second << std::endl; + } + outFile.close(); +} + +void dump(const char *path, std::unique_ptr &program) { + std::string ir; + raw_string_ostream stream(ir); + program->print(stream, nullptr); + + std::ofstream output(path); + output << ir; + output.close(); + // llvm::errs() << "Dumping IR: " << ir << "\n"; +} + +void process(const std::unique_ptr &program, IRBuilder<> &builder) { + SmallVector tracePoints(TracePoints.begin(), TracePoints.end()); + if (tracePoints.empty()) { + printf("No trace points specified, only statistic total kernel execution " + "time\n"); + } + + SmallVector functions; + for (auto &func : program->getFunctionList()) { + if (func.isDeclaration()) { + LLVM_DEBUG(llvm::dbgs() << "Function " << func.getName() + << " is a declaration, skipping.\n"); + continue; + } + functions.push_back(&func); + } + + assert(functions.size() == 1 && "Support only one triton kernel\n"); + auto triton_kernel = functions[0]; + + // Insert profiling trace function in the entry block + auto EntryBlock = &triton_kernel->getEntryBlock(); + builder.SetInsertPoint(EntryBlock, EntryBlock->begin()); + auto group_id = builder.getInt32(Index); // Group ID + auto event_init_value = builder.getInt32(0xFFFFFFFF); // Event ID + auto event_id_ptr = + builder.CreateAlloca(builder.getInt32Ty(), nullptr, "event_id"); + // Store event init value + builder.CreateStore(event_init_value, event_id_ptr); + const auto add_profile_trace_point = program->getFunction("addOrderProfile"); + const auto tsm_wait_finish_point = program->getFunction("TsmWaitfinish"); + assert(add_profile_trace_point && + "Function 'addOrderProfile' not found in the module"); + int order_id = 0; + + orderDesc.insert({order_id, kernel_start}); + builder.CreateCall(add_profile_trace_point, + {group_id, builder.getInt8(order_id++), event_id_ptr}); + + // Insert profiling trace function for each trace point + bool isAddWait = false; + for (auto func : tracePoints) { + const auto function = program->getFunction(func); + + if (!function) { + printf("Function '%s' not found.\n", func.data()); + continue; + } + for (const auto &user : function->users()) { + // 确保该引用实际上是一条调用指令 + if (!isa(user)) + continue; + const auto call_instruction = cast(user); + builder.SetInsertPoint(call_instruction); + + orderDesc.insert({order_id, std::string(func.data()) + "_s"}); + builder.CreateCall(add_profile_trace_point, + {group_id, builder.getInt8(order_id++), event_id_ptr}); + builder.SetInsertPoint(call_instruction->getNextNode()); + + builder.CreateCall(tsm_wait_finish_point); + orderDesc.insert({order_id, std::string(func.data()) + "_e"}); + builder.CreateCall(add_profile_trace_point, + {group_id, builder.getInt8(order_id++), event_id_ptr}); + isAddWait = true; + } + } + + // Insert print function in the exit block + const auto print_order_by_event = program->getFunction("printOrderByEvent"); + assert(print_order_by_event && + "Function 'printOrderByEvent' not found in the module"); + auto EndBlock = &triton_kernel->back(); + builder.SetInsertPoint(EndBlock->getTerminator()); + + if (!isAddWait) { + builder.CreateCall(tsm_wait_finish_point); + } + orderDesc.insert({order_id, kernel_end}); + builder.CreateCall(add_profile_trace_point, + {group_id, builder.getInt8(order_id++), event_id_ptr}); + + builder.CreateCall(print_order_by_event, {group_id, event_id_ptr}); +} + +void create_add_order_profile(const std::unique_ptr &program, + IRBuilder<> &builder) { + std::vector args = { + builder.getInt32Ty() /* groupId */, builder.getInt8Ty() /* orderId */, + PointerType::get(builder.getInt32Ty(), 0) /* eventId */}; + auto function_type = FunctionType::get(builder.getInt64Ty(), args, false); + + program->getOrInsertFunction("addOrderProfile", function_type); +} + +void create_print_profile(const std::unique_ptr &program, + IRBuilder<> &builder) { + // void printOrderByEvent(uint32_t groupId, uint32_t *eventId) + std::vector args = { + builder.getInt32Ty() /* groupId */, + PointerType::get(builder.getInt32Ty(), 0) /* eventId */ + }; + auto function_type = FunctionType::get(builder.getVoidTy(), args, false); + + program->getOrInsertFunction("printOrderByEvent", function_type); +} + +void create_tsm_wait_finish(const std::unique_ptr &program, + IRBuilder<> &builder) { + // uint8_t TsmWaitfinish(); + std::vector args = {}; + auto function_type = FunctionType::get(builder.getInt8Ty(), args, false); + + program->getOrInsertFunction("TsmWaitfinish", function_type); +} + +int main(int argc, char *argv[]) { + initCommandLine(argc, argv); + orderDesc.clear(); + + LLVMContext context; + std::unique_ptr program = nullptr; + parse(InputIR.c_str(), program, context); + + LLVM_DEBUG(llvm::dbgs() << "Loaded IR: " + << program->getModuleIdentifier().data() << "\n"); + // dump(OutFile.c_str(), program); + IRBuilder builder(context); + + create_add_order_profile(program, builder); + create_print_profile(program, builder); + create_tsm_wait_finish(program, builder); + process(program, builder); + + LLVM_DEBUG(llvm::dbgs() << "Verification: " << verifyModule(*program, &dbgs()) + << "\n"); + dump(OutFile.c_str(), program); + + return 0; +} diff --git a/third_party/wafer/python/triton_wafer.cc b/third_party/wafer/python/triton_wafer.cc new file mode 100755 index 00000000..ff3763fe --- /dev/null +++ b/third_party/wafer/python/triton_wafer.cc @@ -0,0 +1,98 @@ +#include "mlir/Conversion/AffineToStandard/AffineToStandard.h" +#include "mlir/Conversion/Passes.h" +#include "mlir/Dialect/Affine/Passes.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Arith/Transforms/BufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Bufferization/Transforms/FuncBufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Bufferization/Transforms/Passes.h" +#include "mlir/Dialect/Func/Extensions/AllExtensions.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/Linalg/Transforms/AllInterfaces.h" +#include "mlir/Dialect/SCF/Transforms/BufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/Dialect/Tensor/Transforms/BufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Vector/IR/ValueBoundsOpInterfaceImpl.h" +#include "mlir/Dialect/Vector/IR/VectorOps.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/Passes.h" +#include "passes.h" +#include "triton-shared/Conversion/TritonToLinalg/TritonToLinalg.h" +// #include +// "triton-shared/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimental.h" +#include "triton-shared/Conversion/TritonToCoreDialects/TritonToCoreDialects.h" + +#include +#include +#include +#include + +namespace py = pybind11; +using namespace mlir; + +void init_triton_tle(py::module &&m); + +void init_triton_wafer_passes_convert(py::module &&m) { + ADD_PASS_WRAPPER_0("add_linalg_to_std", createConvertLinalgToStandardPass); + ADD_PASS_WRAPPER_0("add_one_shot_bufferize", + bufferization::createOneShotBufferizePass); + ADD_PASS_WRAPPER_0("add_triton_to_linalg", triton::createTritonToLinalgPass); + ADD_PASS_WRAPPER_0("add_affine_to_std", createLowerAffinePass); + // ADD_PASS_WRAPPER_0("add_triton_to_linalg_pipeline", + // triton::createTritonToLinalgExperimentalPass); + ADD_PASS_WRAPPER_0("add_triton_to_core", + triton::createTritonToCoreDialectsPass); + + ADD_PASS_WRAPPER_0("add_linalg_to_loops", createConvertLinalgToLoopsPass); + ADD_PASS_WRAPPER_0("add_linalg_to_affine_loops", + createConvertLinalgToAffineLoopsPass); + ADD_PASS_WRAPPER_0("add_lower_affine", createLowerAffinePass); + + m.def("add_affine_vectorize", [](mlir::PassManager &pm, int64_t vecsize) { + affine::AffineVectorizeOptions vectorize_options; + vectorize_options.vectorSizes.push_back(vecsize); + pm.addNestedPass( + affine::createAffineVectorize(vectorize_options)); + }); +} + +void init_triton_wafer_common(py::module &&m) { + m.def("generic_print", [](ModuleOp mod) -> std::string { + std::string str; + llvm::raw_string_ostream os(str); + auto printingFlags = OpPrintingFlags(); + printingFlags.enableDebugInfo(); + printingFlags.printGenericOpForm(); + mod.print(os, printingFlags); + return str; + }); +} + +void init_triton_wafer(py::module &&m) { + init_triton_wafer_common(m.def_submodule("common")); + auto passes = m.def_submodule("passes"); + init_triton_wafer_passes_convert(passes.def_submodule("convert")); + init_triton_tle(m.def_submodule("tle")); + + // load dialects + m.def("load_dialects", [](mlir::MLIRContext &context) { + using namespace mlir; + DialectRegistry registry; + registry.insert(); + + arith::registerBufferizableOpInterfaceExternalModels(registry); + linalg::registerAllDialectInterfaceImplementations(registry); + tensor::registerBufferizableOpInterfaceExternalModels(registry); + bufferization::func_ext::registerBufferizableOpInterfaceExternalModels( + registry); + func::registerAllExtensions(registry); + scf::registerBufferizableOpInterfaceExternalModels(registry); + context.appendDialectRegistry(registry); + context.loadAllAvailableDialects(); + }); + // register passes here +} diff --git a/third_party/wafer/python/triton_wafer_frontend.cc b/third_party/wafer/python/triton_wafer_frontend.cc new file mode 100644 index 00000000..f68d9ee3 --- /dev/null +++ b/third_party/wafer/python/triton_wafer_frontend.cc @@ -0,0 +1,48 @@ +// Frontend-only binding: no FLIR, MK, device lowering, or vendor SDK dependency. +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Arith/Transforms/BufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Bufferization/Transforms/FuncBufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Func/Extensions/AllExtensions.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Transforms/AllInterfaces.h" +#include "mlir/Dialect/SCF/Transforms/BufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/Dialect/Tensor/Transforms/BufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Vector/IR/VectorOps.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/MLIRContext.h" +#include + +namespace py = pybind11; +void init_triton_tle(py::module &&m); + +void init_triton_wafer(py::module &&m) { + m.attr("build_role") = "frontend"; + init_triton_tle(m.def_submodule("tle")); + auto common = m.def_submodule("common"); + common.def("generic_print", [](mlir::ModuleOp mod) { + std::string text; + llvm::raw_string_ostream os(text); + mlir::OpPrintingFlags flags; + flags.enableDebugInfo(); + flags.printGenericOpForm(); + mod.print(os, flags); + return text; + }); + m.def("load_dialects", [](mlir::MLIRContext &context) { + using namespace mlir; + DialectRegistry registry; + registry.insert(); + arith::registerBufferizableOpInterfaceExternalModels(registry); + linalg::registerAllDialectInterfaceImplementations(registry); + tensor::registerBufferizableOpInterfaceExternalModels(registry); + bufferization::func_ext::registerBufferizableOpInterfaceExternalModels(registry); + func::registerAllExtensions(registry); + scf::registerBufferizableOpInterfaceExternalModels(registry); + context.appendDialectRegistry(registry); + context.loadAllAvailableDialects(); + }); +} diff --git a/third_party/wafer/requirements-build.txt b/third_party/wafer/requirements-build.txt new file mode 100644 index 00000000..ef825a23 --- /dev/null +++ b/third_party/wafer/requirements-build.txt @@ -0,0 +1,7 @@ +# Wafer-only build/packaging tools; leave the original DLCompiler requirements intact. +setuptools>=40.8.0 +wheel +cmake>=3.20,<4.0 +ninja>=1.11.1 +pybind11>=2.13.1 +nanobind>=2.4 diff --git a/third_party/wafer/scripts/base/base_run.sh b/third_party/wafer/scripts/base/base_run.sh new file mode 100755 index 00000000..9a04d7fb --- /dev/null +++ b/third_party/wafer/scripts/base/base_run.sh @@ -0,0 +1,110 @@ +#!/bin/bash + +WAFER_DEPS_ROOT=${WAFER_DEPS_ROOT:-} + + +if [ $# -le 2 ]; then + echo "Error: At least two parameters need to be passed!" + exit 1 +fi + +WORKSPACE=$1 +if [ ! -d $WORKSPACE ]; then + echo "Error: $WORKSPACE not exist!" 1>&2 + exit 1 +fi + +args1=$2 +if [ "$args1" != "pytest" ] && [ "$args1" != "python" ]; then + echo "Error: first args is 'pytest' or 'python'" + exit 1 +fi +run_model=$args1 + +shift +shift + +TRITON=$WORKSPACE/FlagTree +WAFER_DEPS_ROOT=$WORKSPACE/wafer_deps +LLVM=$WORKSPACE/llvm-a66376b0-ubuntu-x64 + +if [ ! -d $WAFER_DEPS_ROOT ] || [ ! -d $LLVM ]; then + WORKSPACE="${HOME}/.triton/wafer/" + WAFER_DEPS_ROOT=$WORKSPACE/wafer_deps + LLVM=$WORKSPACE/llvm-a66376b0-ubuntu-x64 +fi + +if [ ! -d $WAFER_DEPS_ROOT ]; then + echo "Error: $WAFER_DEPS_ROOT not exist!" 1>&2 + exit 1 +fi + +if [ ! -d $LLVM ]; then + echo "Error: $LLVM not exist!" 1>&2 + exit 1 +fi + +if [ -f $TRITON/.venv/bin/activate ]; then + source $TRITON/.venv/bin/activate +fi + +wafer_skip_ops="repeat_interleave.self_int,pad,to.dtype,uniform_,sort.values_stable,contiguous,resolve_conj" +wafer_fallback_cpu_ops="random_,quantile,_local_scalar_dense,arange,unfold,index,le,all,ge,pad,to,gather_backward,zero_,view_as_real,resolve_neg,embedding_backward,sort,repeat_interleave,rsub,hstack,vstack,min,uniform_,abs,ne,eq,mul,bitwise_and,masked_select,max,ceil,div,gt,lt,sum,scatter,where,resolve_conj,isclose,isfinite,tile,equal,gather,contiguous" + +# 必须的 +export WAFER_DEPS_ROOT=$WAFER_DEPS_ROOT +export LLVM_SYSPATH=$LLVM +export LLVM_BINARY_DIR=$LLVM/bin + +# 后续需要优化删除的 +export PYTHONPATH=$LLVM/python_packages/mlir_core:$PYTHONPATH +export LD_LIBRARY_PATH=$WAFER_DEPS_ROOT/lib:$LD_LIBRARY_PATH +export TXDA_SKIP_OPS=$wafer_skip_ops +export TXDA_FALLBACK_CPU_OPS=$wafer_fallback_cpu_ops + +# 非必须的 调试相关 +export TRITON_DUMP_PATH=$TRITON/dump +export TRITON_ALWAYS_COMPILE=1 +export TRITON_PRINT_AUTOTUNING=1 +# dump launch调用的所有参数,包括kernel func调用的参数 +export DUMP_KERNEL_ARGS=1 + +# 高精度模式 +export PRECISION_PRIORITY=1 +#multinomial算子编译需要 +export TRITON_ALLOW_NON_CONSTEXPR_GLOBALS=1 + +# export DEBUG=ON +# export ENABLE_PROFILING=1 +# export USE_HOST_PROFILE=1 +# export WAFER_LOG_LEVEL=debug +# export CUSTOMIZED_IR=test_0.mlir,test_1.mlir +# export TRACE_POINTS="__Rdma,__Wdma" + +echo "export WAFER_LOG_LEVEL=$TX_LOG_LEVEL" +echo "export WAFER_DEPS_ROOT=$WAFER_DEPS_ROOT" +echo "export LLVM_SYSPATH=$LLVM_SYSPATH" +echo "export LLVM_BINARY_DIR=$LLVM_BINARY_DIR" +echo "export PYTHONPATH=$PYTHONPATH" +echo "export LD_LIBRARY_PATH=$LD_LIBRARY_PATH" +echo "export PRECISION_PRIORITY=$PRECISION_PRIORITY" + +echo "export DUMP_KERNEL_ARGS=$DUMP_KERNEL_ARGS" +echo "export TRITON_DUMP_PATH=$TRITON_DUMP_PATH" +echo "export TRITON_ALWAYS_COMPILE=$TRITON_ALWAYS_COMPILE" +echo "export TXDA_SKIP_OPS=$TXDA_SKIP_OPS" +echo "export TXDA_FALLBACK_CPU_OPS=$TXDA_FALLBACK_CPU_OPS" + +echo "export ENABLE_PROFILING=$ENABLE_PROFILING" +echo "export USE_HOST_PROFILE=$USE_HOST_PROFILE" +echo "export CUSTOMIZED_IR=$CUSTOMIZED_IR" +echo "export TRACE_POINTS=$TRACE_POINTS" + +pytest_cmd="" +if [ "$args1" == "pytest" ]; then + pytest_cmd="-m pytest -v -s" +fi + +echo "run cmd:python3 $pytest_cmd $@" + +USE_SIM_MODE=${USE_SIM_MODE} python3 $pytest_cmd $@ diff --git a/third_party/wafer/scripts/build_llvm.sh b/third_party/wafer/scripts/build_llvm.sh new file mode 100755 index 00000000..2bc46787 --- /dev/null +++ b/third_party/wafer/scripts/build_llvm.sh @@ -0,0 +1,31 @@ +#!/bin/bash +# hash a66376b0dc3b2ea8a84fda26faca287980986f78 + +if [ -z "${LLVM_PROJECT+x}" ]; then + echo "Please set the environment variable “LLVM_PROJECT”." 1>&2 + exit 1 +fi + +if [ ! -d $LLVM_PROJECT ]; then + echo "Error: $LLVM_PROJECT not exist!" 1>&2 + exit 1 +fi + +BUILD_TYPE=Release + +build_llvm() { + mkdir $LLVM_PROJECT/build + cd $LLVM_PROJECT/build + cmake -G Ninja \ + -DCMAKE_BUILD_TYPE=$BUILD_TYPE \ + -DLLVM_ENABLE_ASSERTIONS=ON \ + -DLLVM_ENABLE_PROJECTS="clang;mlir;llvm;lld" \ + -DLLVM_TARGETS_TO_BUILD="host;NVPTX;AMDGPU;RISCV" \ + -DLLVM_USE_LINKER=lld \ + -DMLIR_ENABLE_BINDINGS_PYTHON=1 \ + -DPython3_EXECUTABLE="$(which python3)" \ + ../llvm + ninja +} + +build_llvm diff --git a/third_party/wafer/scripts/build_wafer.sh b/third_party/wafer/scripts/build_wafer.sh new file mode 100755 index 00000000..bcd45716 --- /dev/null +++ b/third_party/wafer/scripts/build_wafer.sh @@ -0,0 +1,142 @@ +#!/bin/bash + +WAFER_DEPS_ROOT=${WAFER_DEPS_ROOT:-} +WAFER_RT_THREAD_SMP_ROOT=${WAFER_RT_THREAD_SMP_ROOT:-} + + +set -e + +script_path=$(realpath "$0") +script_dir=$(dirname "$script_path") +project_dir=$(realpath "$script_dir/../../..") + +if [ -z "${WORKSPACE+x}" ]; then + WORKSPACE=$(realpath "$project_dir/..") +fi + +WAFER_DEPS_ROOT=$WORKSPACE/wafer_deps +LLVM=$WORKSPACE/llvm-a66376b0-ubuntu-x64 +TRITON=$project_dir + +if [ ! -d $WAFER_DEPS_ROOT ] || [ ! -d $LLVM ]; then + WORKSPACE="${HOME}/.triton/wafer/" + WAFER_DEPS_ROOT=$WORKSPACE/wafer_deps + LLVM=$WORKSPACE/llvm-a66376b0-ubuntu-x64 +fi + +if [ ! -d $WAFER_DEPS_ROOT ]; then + echo "Error: $WAFER_DEPS_ROOT not exist!" 1>&2 + exit 1 +fi + +if [ ! -d $LLVM ]; then + echo "Error: $LLVM not exist!" 1>&2 + exit 1 +fi + +# Default values +BUILD_TYPE="release" +ACTION="install" +SCRIPT_NAME=$(basename "$0") + +# Function to display usage +usage() { + echo "Usage: $SCRIPT_NAME [-t build_type] [-a action]" + echo "Options:" + echo " -t build_type Specify build type: debug or release (default: release)" + echo " -a action Specify action: install or wheel (default: install)" + echo " -h Display this help message" + exit 1 +} + +# Process command-line options with getopts +while getopts ":t:a:h" opt; do + case $opt in + t) + BUILD_TYPE="$OPTARG" + # Convert to lowercase for case-insensitive comparison + BUILD_TYPE=$(echo "$BUILD_TYPE" | tr '[:upper:]' '[:lower:]') + # Validate build_type + if [[ "$BUILD_TYPE" != "debug" && "$BUILD_TYPE" != "release" ]]; then + echo "Error: Invalid build type '$BUILD_TYPE'. Must be 'debug' or 'release'." >&2 + usage + fi + ;; + a) + ACTION="$OPTARG" + ACTION=$(echo "$ACTION" | tr '[:upper:]' '[:lower:]') + # Validate action + if [[ "$ACTION" != "install" && "$ACTION" != "wheel" ]]; then + echo "Error: Invalid action '$ACTION'. Must be 'install' or 'wheel'." >&2 + usage + fi + ;; + h) + usage + ;; + \?) + echo "Error: Invalid option -$OPTARG" >&2 + usage + ;; + :) + echo "Error: Option -$OPTARG requires an argument." >&2 + usage + ;; + esac +done + +# Shift off the options and optional arguments +shift $((OPTIND - 1)) + +echo "Build configuration:" +echo " Build Type: $BUILD_TYPE" +echo " Action: $ACTION" + +build_triton() { + if [ "$BUILD_TYPE" == "debug" ]; then + export DEBUG=ON + else + export REL_WITH_DBG_INFO=ON + fi + + export TRITON_BUILD_WITH_CLANG_LLD=true + export TRITON_BUILD_WITH_CCACHE=true + export TRITON_OFFLINE_BUILD=ON + export TRITON_BUILD_PROTON=OFF + + echo "export TRITON_OFFLINE_BUILD=$TRITON_OFFLINE_BUILD" + echo "export TRITON_BUILD_WITH_CLANG_LLD=$TRITON_BUILD_WITH_CLANG_LLD" + echo "export TRITON_BUILD_WITH_CCACHE=$TRITON_BUILD_WITH_CCACHE" + echo "export TRITON_BUILD_PROTON=$TRITON_BUILD_PROTON" + + cd $TRITON/python + build_opt=install + + if [ "$ACTION" == "wheel" ]; then + build_opt=wheel + fi + + python3 -m pip $build_opt . --no-index --no-deps --no-build-isolation -v --verbose +} + +if [ -f $TRITON/.venv/bin/activate ]; then + source $TRITON/.venv/bin/activate +fi + +export LLVM_SYSPATH=$LLVM +export WAFER_DEPS_ROOT=$WAFER_DEPS_ROOT +export WAFER_RT_THREAD_SMP_ROOT=$WAFER_DEPS_ROOT/tx8-yoc-rt-thread-smp +export FLAGTREE_BACKEND=wafer + +# debug +# export USE_HOST_PROFILE=1 +# export NO_INTRNISIC_RUN=1 + +echo "export WAFER_DEPS_ROOT=$WAFER_DEPS_ROOT" +echo "export LLVM_SYSPATH=$LLVM_SYSPATH" + +# synchronous temporary solution: add waitfinish after every cintrinsic exec +export ENABLE_SYNCHRONOUS_INTRINSIC=1 +echo "export ENABLE_SYNCHRONOUS_INTRINSIC=$ENABLE_SYNCHRONOUS_INTRINSIC" + +build_triton diff --git a/third_party/wafer/scripts/publish/run_flaggems_on_multicards.sh b/third_party/wafer/scripts/publish/run_flaggems_on_multicards.sh new file mode 100755 index 00000000..72d5d862 --- /dev/null +++ b/third_party/wafer/scripts/publish/run_flaggems_on_multicards.sh @@ -0,0 +1,97 @@ +#!/bin/bash + +WAFER_DEPS_ROOT=${WAFER_DEPS_ROOT:-} + +set -e +##.在docker容器内版本包路径下执行 +#bash scripts/run_flaggems_on_multicards.sh ci_ops 1 + + +########################################################################################################################## +## ## +## 在多卡上并行运行Triton算子测试脚本 ## +## param1: test_set, set test set name, default 'ci_ops'. ## +## param2: device_count, set device count number, default 1. ## +## ##param: precision_priority, set 1-triton compiler use high precision mode for special ops, default 1. ## +## param3: quick_mode, set 1-quick mode to run flaggems, set 0-normal mode, default 0. ## +## param4: skip_device, set devices that need to be skipped, when they are unavailable, default []. ## +## ## +########################################################################################################################## + +script_path=$(realpath "$0") +echo $script_path +script_dir=$(dirname "$script_path") +echo $script_dir +project_dir=$(realpath "$script_dir/../") +echo $project_dir +export TRITON_WORKSPACE=$project_dir +test_set=ci_ops +device_count=1 +quick_mode=0 +skip_device= +precision_priority=1 +wafer_skip_ops="repeat_interleave.self_int,pad,to.dtype,uniform_,sort.values_stable,contiguous,resolve_conj" +wafer_fallback_cpu_ops="random_,quantile,_local_scalar_dense,arange,unfold,index,le,all,ge,pad,to,gather_backward,zero_,view_as_real,resolve_neg,embedding_backward,sort,repeat_interleave,rsub,hstack,vstack,min,uniform_,abs,ne,eq,mul,bitwise_and,masked_select,max,ceil,div,gt,lt,sum,scatter,where,resolve_conj,isclose,isfinite,tile,equal,gather,contiguous" + +if [ $# -ge 1 ]; then + test_set=$1 +fi +if [ $# -ge 2 ]; then + device_count=$2 +fi +if [ $# -ge 3 ]; then + quick_mode=$3 +fi +if [ $# -ge 4 ]; then + skip_device=$(echo $4 | tr ',' ' ') +fi +echo "param count:"$# +echo "test_set:"$test_set +echo "device_count:"$device_count +echo "quick_mode:"$quick_mode +echo "skip_device:"$skip_device +echo "precision_priority:"$precision_priority +echo "txda_skip_ops:"$wafer_skip_ops +echo "txda_fallback_cpu_ops:"$wafer_fallback_cpu_ops + +#triton系统相关环境变量 +WAFER_DEPS_ROOT=$project_dir/wafer_deps +LLVM=$project_dir/llvm-a66376b0-ubuntu-x64 +export WAFER_DEPS_ROOT=$WAFER_DEPS_ROOT +export LLVM_SYSPATH=$LLVM +export LLVM_BINARY_DIR=$LLVM/bin +export PYTHONPATH=$LLVM/python_packages/mlir_core:$PYTHONPATH +export LD_LIBRARY_PATH=$WAFER_DEPS_ROOT/lib:$LD_LIBRARY_PATH +export TRITON_ALWAYS_COMPILE=1 +#测试任务相关环境变量 +export JSON_FILE_PATH=$project_dir/flaggems_tests +export PRECISION_PRIORITY=$precision_priority +export TRITON_ALLOW_NON_CONSTEXPR_GLOBALS=1 +export TXDA_SKIP_OPS=$wafer_skip_ops +export TXDA_FALLBACK_CPU_OPS=$wafer_fallback_cpu_ops + +echo "WAFER_DEPS_ROOT="$WAFER_DEPS_ROOT +echo "LLVM_SYSPATH="$LLVM_SYSPATH +echo "LLVM_BINARY_DIR="$LLVM_BINARY_DIR +echo "PYTHONPATH="$PYTHONPATH +echo "LD_LIBRARY_PATH="$LD_LIBRARY_PATH +echo "TRITON_ALWAYS_COMPILE="$TRITON_ALWAYS_COMPILE +echo "JSON_FILE_PATH="$JSON_FILE_PATH +echo "PRECISION_PRIORITY="$PRECISION_PRIORITY +echo "TRITON_ALLOW_NON_CONSTEXPR_GLOBALS="$TRITON_ALLOW_NON_CONSTEXPR_GLOBALS +echo "TXDA_SKIP_OPS="$TXDA_SKIP_OPS +echo "TXDA_FALLBACK_CPU_OPS="$TXDA_FALLBACK_CPU_OPS + +source $project_dir/triton/.venv/bin/activate +if [ $quick_mode -eq 1 ]; then + python3 $project_dir/flaggems_tests/test_flag_gems_ci.py --test_set $test_set --device_count $device_count --skip_device $skip_device --quick +else + python3 $project_dir/flaggems_tests/test_flag_gems_ci.py --test_set $test_set --device_count $device_count --skip_device $skip_device +fi + +if [ $? -eq 0 ]; then + echo "Run test complete!" +else + echo "Run test fail!!!" + exit -1 +fi diff --git a/third_party/wafer/scripts/publish/run_wafer.sh b/third_party/wafer/scripts/publish/run_wafer.sh new file mode 100755 index 00000000..5cca67ba --- /dev/null +++ b/third_party/wafer/scripts/publish/run_wafer.sh @@ -0,0 +1,10 @@ +#!/bin/bash + +script_path=$(realpath "$0") +script_dir=$(dirname "$script_path") + +if [ -z "${WORKSPACE+x}" ]; then + WORKSPACE=$(realpath "$script_dir/..") +fi + +bash $script_dir/base_run.sh $WORKSPACE $@ diff --git a/third_party/wafer/scripts/requirements_ts.txt b/third_party/wafer/scripts/requirements_ts.txt new file mode 100755 index 00000000..71ae29f2 --- /dev/null +++ b/third_party/wafer/scripts/requirements_ts.txt @@ -0,0 +1,7 @@ +gitpython +nanobind +torch==2.7.0 +torchvision +pytest +pyyaml +pybind11 diff --git a/third_party/wafer/scripts/run_wafer.sh b/third_party/wafer/scripts/run_wafer.sh new file mode 100755 index 00000000..ddbeaf6d --- /dev/null +++ b/third_party/wafer/scripts/run_wafer.sh @@ -0,0 +1,11 @@ +#!/bin/bash + +script_path=$(realpath "$0") +script_dir=$(dirname "$script_path") +project_dir=$(realpath "$script_dir/../../..") + +if [ -z "${WORKSPACE+x}" ]; then + WORKSPACE=$(realpath "$project_dir/..") +fi + +bash $script_dir/base/base_run.sh $WORKSPACE $@ diff --git a/third_party/wafer/scripts/tools/suuplement.sh b/third_party/wafer/scripts/tools/suuplement.sh new file mode 100755 index 00000000..835f7edc --- /dev/null +++ b/third_party/wafer/scripts/tools/suuplement.sh @@ -0,0 +1,42 @@ +#!/bin/bash + +# 参数检查 +if [ $# -ne 2 ]; then + echo "用法: $0 " + exit 1 +fi + +REQUIREMENTS=$1 +PACKAGES_DIR=$2 + +# 检查pip和requirements文件 +if ! command -v pip &> /dev/null; then + echo "错误: pip未安装" + exit 1 +fi + +if [ ! -f "$REQUIREMENTS" ]; then + echo "错误: 文件 $REQUIREMENTS 不存在" + exit 1 +fi + +if [ ! -d "$PACKAGES_DIR" ]; then + echo "错误: 目录 $PACKAGES_DIR 不存在" + exit 1 +fi + +echo "正在检查并补充下载缺失的依赖包..." + +# 使用exists-action=i选项,跳过已存在的包 +pip download \ + -r "$REQUIREMENTS" \ + -d "$PACKAGES_DIR" \ + --exists-action=i \ + --no-deps # 假设依赖项已经完整 + +if [ $? -ne 0 ]; then + echo "错误: 下载依赖包失败" + exit 1 +fi + +echo "完成! 缺失的依赖包已补充下载到 $PACKAGES_DIR" diff --git a/third_party/wafer/third_party/flir/.clang-format b/third_party/wafer/third_party/flir/.clang-format new file mode 100755 index 00000000..9b3aa8b7 --- /dev/null +++ b/third_party/wafer/third_party/flir/.clang-format @@ -0,0 +1 @@ +BasedOnStyle: LLVM diff --git a/third_party/wafer/third_party/flir/.github/PULL_REQUEST_TEMPLATE.md b/third_party/wafer/third_party/flir/.github/PULL_REQUEST_TEMPLATE.md new file mode 100755 index 00000000..396245e1 --- /dev/null +++ b/third_party/wafer/third_party/flir/.github/PULL_REQUEST_TEMPLATE.md @@ -0,0 +1,16 @@ + diff --git a/third_party/wafer/third_party/flir/.github/workflows/code-format-check.yml b/third_party/wafer/third_party/flir/.github/workflows/code-format-check.yml new file mode 100755 index 00000000..8639cd61 --- /dev/null +++ b/third_party/wafer/third_party/flir/.github/workflows/code-format-check.yml @@ -0,0 +1,23 @@ +name: Code-Format-Check + +on: + schedule: + - cron: '0 21 * * *' + push: + branches: [ "main" ] + pull_request: + branches: [ "main" ] + +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: true + +jobs: + pre-commit: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: '3.11' + - uses: pre-commit/action@v3.0.1 diff --git a/third_party/wafer/third_party/flir/.gitignore b/third_party/wafer/third_party/flir/.gitignore new file mode 100755 index 00000000..fb37c737 --- /dev/null +++ b/third_party/wafer/third_party/flir/.gitignore @@ -0,0 +1,5 @@ +*.pyc +.cache +compile_commands.json +build/* +.vscode/* \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/.gitmodules b/third_party/wafer/third_party/flir/.gitmodules new file mode 100755 index 00000000..e69de29b diff --git a/third_party/wafer/third_party/flir/.pre-commit-config.yaml b/third_party/wafer/third_party/flir/.pre-commit-config.yaml new file mode 100755 index 00000000..88075d3d --- /dev/null +++ b/third_party/wafer/third_party/flir/.pre-commit-config.yaml @@ -0,0 +1,29 @@ +repos: + - repo: https://github.com/pre-commit/pre-commit-hooks + rev: v4.4.0 + hooks: + - id: destroyed-symlinks + - id: check-yaml + - id: check-toml + - id: check-ast + - id: check-added-large-files + - id: check-merge-conflict + - id: check-shebang-scripts-are-executable + - id: detect-private-key + - id: debug-statements + + - repo: https://github.com/pre-commit/mirrors-clang-format + rev: v16.0.6 + hooks: + - id: clang-format + files: | + (?x)^( + include/incubated/.*\.(h|hpp|cc|cpp)| + include/mlir-ext/.*\.(h|hpp|cc|cpp)| + include/npu/.*\.(h|hpp|cc|cpp)| + lib/UtilsIncubated/.*\.(h|hpp|cc|cpp)| + lib/Dialect/(MathExt|TritonAscend|TritonStructuredIncubated)/.*\.(h|hpp|cc|cpp)| + lib/Conversion/(DiscreteMaskAccessConversion|NoBufferize_FlagTree|TritonToAnnotation|TritonToLinalgIncubated|TritonToStructuredIncubated|TritonToUnstructureIncubated)/.*\.(h|hpp|cc|cpp)| + .*_FlagTree\.(h|hpp|cc|cpp) + )$ + stages: [pre-commit, pre-push, manual] diff --git a/third_party/wafer/third_party/flir/CMakeLists.txt b/third_party/wafer/third_party/flir/CMakeLists.txt new file mode 100755 index 00000000..11578ec8 --- /dev/null +++ b/third_party/wafer/third_party/flir/CMakeLists.txt @@ -0,0 +1,27 @@ +option(TRITON_SHARED_BUILD_CPU_BACKEND "Build triton-shared CPU backend" ON) +option(FLIR_BUILD_INCUBATED "Build FLIR incubated dialects/conversions" ${FLIR_BUILD_INCUBATED}) +set(TRITON_SHARED_SOURCE_DIR "${CMAKE_CURRENT_SOURCE_DIR}") +set(TRITON_SHARED_BINARY_DIR "${CMAKE_CURRENT_BINARY_DIR}") + +include_directories(${CMAKE_CURRENT_SOURCE_DIR}/include) +include_directories(${CMAKE_CURRENT_BINARY_DIR}/include) # Tablegen'd files +include_directories(${Python3_INCLUDE_DIRS}) +include_directories(${pybind11_INCLUDE_DIR}) + +if(FLAGTREE_BACKEND STREQUAL "wafer") + add_compile_definitions(FLAGTREE_BACKEND_WAFER) +endif() + +add_subdirectory(include) +add_subdirectory(lib) + +if(NOT FLAGTREE_BACKEND STREQUAL "wafer") + add_subdirectory(test) + add_subdirectory(tools) +endif() + +if (TRITON_SHARED_BUILD_CPU_BACKEND) + add_triton_plugin(TritonShared ${CMAKE_CURRENT_SOURCE_DIR}/triton_shared.cc LINK_LIBS TritonSharedAnalysis TritonToLinalg TritonTilingExtIR) +endif() + + diff --git a/third_party/wafer/third_party/flir/LICENSE b/third_party/wafer/third_party/flir/LICENSE new file mode 100755 index 00000000..e62005eb --- /dev/null +++ b/third_party/wafer/third_party/flir/LICENSE @@ -0,0 +1,23 @@ +/* +* Copyright 2023- Microsoft Corporation +* Copyright 2025- FlagTree Project Management Committee +* +* Permission is hereby granted, free of charge, to any person obtaining +* a copy of this software and associated documentation files +* (the "Software"), to deal in the Software without restriction, +* including without limitation the rights to use, copy, modify, merge, +* publish, distribute, sublicense, and/or sell copies of the Software, +* and to permit persons to whom the Software is furnished to do so, +* subject to the following conditions: +* +* The above copyright notice and this permission notice shall be +* included in all copies or substantial portions of the Software. +* +* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, +* EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF +* MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. +* IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY +* CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, +* TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE +* SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. +*/ diff --git a/third_party/wafer/third_party/flir/README.md b/third_party/wafer/third_party/flir/README.md new file mode 100755 index 00000000..b4c0d897 --- /dev/null +++ b/third_party/wafer/third_party/flir/README.md @@ -0,0 +1,12 @@ +## FLIR + +FlagTree IR is forked from [microsoft/triton-shared](https://github.com/microsoft/triton-shared), which is a Shared Middle-Layer for Triton Compilation. It is used for FlagTree. + +## Contributing + +Please refer to [https://github.com/FlagTree/flagtree](https://github.com/FlagTree/flagtree). + +## License + +FLIR is licensed under the [MIT license](/LICENSE). + diff --git a/third_party/wafer/third_party/flir/backend/compiler.py b/third_party/wafer/third_party/flir/backend/compiler.py new file mode 100755 index 00000000..2f7b4aaa --- /dev/null +++ b/third_party/wafer/third_party/flir/backend/compiler.py @@ -0,0 +1,219 @@ +from triton.backends.compiler import BaseBackend, GPUTarget +from triton._C.libtriton import ir, passes +from dataclasses import dataclass +from typing import Any, Dict, Tuple +from types import ModuleType +import hashlib +import tempfile +import os +import re +import shutil +import subprocess +import functools +from pathlib import Path + +def _get_triton_shared_opt_path() -> str: + path = os.getenv("TRITON_SHARED_OPT_PATH", "") + if path == "": + raise Exception("TRITON_SHARED_OPT_PATH is not set.") + return path + + +def _get_llvm_bin_path(bin_name: str) -> str: + path = os.getenv("LLVM_BINARY_DIR", "") + if path == "": + raise Exception("LLVM_BINARY_DIR is not set.") + return os.path.join(path, bin_name) + + +def _dump_ir_if_needed(files): + path = os.getenv("TRITON_SHARED_DUMP_PATH", "") + if not path: + return + for f in files: + shutil.copy(f, os.path.join(path, os.path.basename(f))) + + +def _ttir_to_ttsharedir(mod): + # Get Triton-MLIR as string + ttir_code = str(mod) + with tempfile.TemporaryDirectory() as tmpdir: + src_path = os.path.join(tmpdir, "tt.mlir") + dst_path = os.path.join(tmpdir, "ttshared.mlir") + Path(src_path).write_text(ttir_code) + _dump_ir_if_needed([src_path]) + triton_shared_opt_path = _get_triton_shared_opt_path() + subprocess.check_call([triton_shared_opt_path, src_path, "--triton-to-linalg-experimental", "--mlir-print-debuginfo", "-o", dst_path]) + return Path(dst_path).read_text() + + +def _optimize_ttsharedir(ttsharedir: str): + # We don't apply any optimizations now, but we can add passes if needed. + return ttsharedir + + +def _ttsharedir_to_llir(ttsharedir: str): + with tempfile.TemporaryDirectory() as tmpdir: + ttshared_path = os.path.join(tmpdir, "ttshared.mlir") + llmlir_path = os.path.join(tmpdir, "ll.mlir") + llir_path = os.path.join(tmpdir, "ll.ir") + Path(ttshared_path).write_text(ttsharedir) + mlir_opt_path = _get_llvm_bin_path("mlir-opt") + # TritonShared-MLIR to LLVM-MLIR + subprocess.check_call([mlir_opt_path, ttshared_path, + "--convert-linalg-to-affine-loops", + # Note: eliminate-empty-tensors fails when there are multiple func.return ops + # in a single kernel which are the results of early returns. + # See python/examples/test_early_return.py for examples. + # We disable this pass for now since performance on CPU isn't the main + # focus at the moment. + # "--eliminate-empty-tensors", + "--empty-tensor-to-alloc-tensor", + "--one-shot-bufferize=allow-return-allocs-from-loops=true", + "--lower-affine", + "--convert-linalg-to-loops", + "--expand-strided-metadata", + "--convert-scf-to-cf", + "--convert-arith-to-llvm", + "--convert-math-to-llvm", + "--convert-complex-to-llvm", + "--convert-vector-to-llvm", + "--convert-index-to-llvm", + "--memref-expand", + "--finalize-memref-to-llvm", + "--convert-func-to-llvm", + "--convert-cf-to-llvm", + # Lowering memrefs creates more affine.apply ops. + # Lowering these affine ops again creates further arith ops, + # so we have to run these two passes again here. + "--lower-affine", + "--convert-arith-to-llvm", + # Remove all unrealized casts created + "--reconcile-unrealized-casts", + "--mlir-print-debuginfo", + "-o", + llmlir_path]) + + # LLVM-MLIR to LLVM-IR + mlir_translate_path = _get_llvm_bin_path("mlir-translate") + subprocess.check_call([mlir_translate_path, llmlir_path, + "--mlir-to-llvmir", + "-o", + llir_path]) + _dump_ir_if_needed([ttshared_path, llmlir_path, llir_path]) + return Path(llir_path).read_text() + + +def _optimize_llir(llir: str): + # We don't apply any optimizations now, but we can add passes if needed. + return llir + + +def _llir_to_bin(llir: str, metadata): + pattern = r"define void @(\w+)\(.+" + matches = re.findall(pattern, llir) + assert len(matches) == 1 + metadata["name"] = matches[0] + with tempfile.TemporaryDirectory() as tmpdir: + src_path = os.path.join(tmpdir, "kernel.ll") + dst_path = os.path.join(tmpdir, "kernel.o") + Path(src_path).write_text(llir) + llc_path = _get_llvm_bin_path("llc") + subprocess.check_call([llc_path, src_path, "-filetype=obj", "-o", dst_path]) + return Path(dst_path).read_bytes() + + + +@dataclass(frozen=True) +class CPUOptions: + debug: bool = False + arch: str = None + num_warps: int = 0 + num_ctas: int = 0 + num_stages: int = 1 + enable_warp_specialization: bool = False + enable_fp_fusion: bool = False + extern_libs = None + cluster_dims: tuple = (1, 1, 1) + shared: bool = False + # Disable FP8 here since this is a sample CPU backend. + # Target specific backends can eanble it with supported types. + supported_fp8_dtypes: Tuple[str] = () + allow_fp8e4nv: bool = False + allowed_dot_input_precisions: Tuple[str] = ("ieee", ) + sanitize_overflow: bool = True + + def __post_init__(self): + pass + + def hash(self): + key = '_'.join([f'{name}-{val}' for name, val in self.__dict__.items()]) + return hashlib.md5(key.encode("utf-8")).hexdigest() + + +class CPUBackend(BaseBackend): + binary_ext = 'obj' + + @staticmethod + def supports_target(target: GPUTarget): + return target.backend == 'cpu' + + def __init__(self, target: GPUTarget) -> None: + super().__init__(target) + + def parse_options(self, opts) -> Any: + args = {'arch': self.target.arch} + args.update({k: opts[k] for k in CPUOptions.__dataclass_fields__.keys() if k in opts}) + return CPUOptions(**args) + + def get_codegen_implementation(self, options): + codegen_fns = {"min_dot_size": lambda lhsType, rhsType: (1, 1, 1)} + return codegen_fns + + def pack_metadata(self, metadata): + # Note: We actually don't need any of these except for the name which is + # used in the launch function in driver.py. Putting these in so we're + # consistent with other backends + return ( + metadata.num_warps, + metadata.num_ctas, + metadata.shared, + metadata.cluster_dims[0], + metadata.cluster_dims[1], + metadata.cluster_dims[2], + metadata.name + ) + + # Our compilation pipeline isn't in python like nvidia or amd, no need to load + # dialects. See `triton_shared.cc` + def load_dialects(self, ctx): + return + + @staticmethod + def make_ttir(mod, metadata, opt): + pm = ir.pass_manager(mod.context) + pm.enable_debug() + passes.common.add_inliner(pm) + passes.ttir.add_combine(pm) + passes.common.add_canonicalizer(pm) + passes.ttir.add_reorder_broadcast(pm) + passes.common.add_cse(pm) + passes.common.add_licm(pm) + passes.common.add_symbol_dce(pm) + pm.run(mod) + return mod + + def add_stages(self, stages, options): + stages["ttir"] = lambda src, metadata: self.make_ttir(src, metadata, options) + stages["ttsharedir"] = lambda src, metadata: _optimize_ttsharedir(_ttir_to_ttsharedir(src)) + stages["llir"] = lambda src, metadata: _optimize_llir(_ttsharedir_to_llir(src)) + stages["obj"] = lambda src, metadata: _llir_to_bin(src, metadata) + + + @functools.lru_cache() + def hash(self): + return self.target + + # The CPU backend does not use any extra python modules, return an empty dictionary + def get_module_map(self) -> Dict[str, ModuleType]: + return {} diff --git a/third_party/wafer/third_party/flir/backend/driver.py b/third_party/wafer/third_party/flir/backend/driver.py new file mode 100755 index 00000000..4c0f026d --- /dev/null +++ b/third_party/wafer/third_party/flir/backend/driver.py @@ -0,0 +1,397 @@ +import hashlib +import tempfile +import sysconfig + +import os, subprocess, tempfile, platform +import importlib.util +import sys + +from pathlib import Path + +from triton.runtime.cache import get_cache_manager +from triton.backends.driver import DriverBase +from triton.backends.compiler import GPUTarget + +# -------------------- Launcher ---------------------------- +def _ty_to_cpp(ty): + if ty[0] == '*': + return "void*" + if ty == "constexpr": + return "PyObject*" + return { + "i1": "int32_t", + "i8": "int8_t", + "i16": "int16_t", + "i32": "int32_t", + "i64": "int64_t", + "u1": "uint32_t", + "u8": "uint8_t", + "u16": "uint16_t", + "u32": "uint32_t", + "u64": "uint64_t", + "fp16": "float", + "bf16": "float", + "fp32": "float", + "f32": "float", + "fp64": "double", + }[ty] + +def _extracted_type(ty): + if ty[0] == '*': + return "PyObject*" + if ty == "constexpr": + return "PyObject*" + return _ty_to_cpp(ty) + +def _format_of(ty): + return { + "PyObject*": "O", + "constexpr": "O", + "float": "f", + "double": "d", + "long": "l", + "int8_t": "b", + "int16_t": "h", + "int32_t": "i", + "int64_t": "l", + "uint8_t": "B", + "uint16_t": "H", + "uint32_t": "I", + "uint64_t": "K", + }[ty] + +def _generate_launcher(constants, signature, kernel_name): + arg_decls = ', '.join(f"{_ty_to_cpp(ty)} arg{i}" for i, ty in signature.items()) + args_format = ''.join([_format_of(_extracted_type(ty)) for ty in signature.values()]) + format = "iiiOOOO" + args_format + args_list = ', ' + ', '.join(f"&_arg{i}" for i, ty in signature.items()) if len(signature) > 0 else '' + + kernel_arg_decls = ', '.join(_ty_to_cpp(ty) if ty[0] != "*" else f"int64_t, void*" for i, ty in signature.items() if ty != "constexpr") + kernel_arg_decls += ', ' if kernel_arg_decls else '' + + kernel_parameters = ', '.join(f"static_cast<{_ty_to_cpp(ty)}>(arg{i})" if ty[0] != "*" else f"0, &ptr_arg{i}" for i, ty in signature.items() if ty != "constexpr") + kernel_parameters += ', ' if kernel_parameters else '' + + return f""" +#include +#include +#include +#include "ExecutionEngine/CRunnerUtils.h" +#include "ExecutionEngine/CRunnerUtils.cpp" + +extern "C" {{ + // Pointer type (=Memref) becomes int64_t + MemRef struct + // FIXME: understand what this int64_t is used for. + void {kernel_name}({kernel_arg_decls} + int, int, int, int, int, int); +}} + +static void _launch(int gridX, int gridY, int gridZ, {arg_decls}) {{ + if (gridX*gridY*gridZ > 0) {{ + // Cast "function" to the real function type. + for(int x = 0; x < gridX; x++) {{ + for(int y = 0; y < gridY; y++) {{ + for(int z = 0; z < gridZ; z++) {{ + // Use some random type "char" here. + {' '.join(f'StridedMemRefType ptr_arg{i} = {{static_cast(arg{i}), static_cast(arg{i}), 0}};' for i, ty in signature.items() if i not in constants and ty[0] == "*")} + {kernel_name}({kernel_parameters} + gridX, gridY, gridZ, x, y, z); + }} + }} + }} + }} +}} + +typedef struct _DevicePtrInfo {{ + void *dev_ptr; + bool valid; +}} DevicePtrInfo; + +static inline DevicePtrInfo getPointer(PyObject *obj, int idx) {{ + DevicePtrInfo ptr_info; + ptr_info.dev_ptr = 0; + ptr_info.valid = true; + if (PyLong_Check(obj)) {{ + ptr_info.dev_ptr = reinterpret_cast(PyLong_AsUnsignedLongLong(obj)); + return ptr_info; + }} + if (obj == Py_None) {{ + // valid nullptr + return ptr_info; + }} + PyObject *ptr = PyObject_GetAttrString(obj, "data_ptr"); + if(ptr){{ + PyObject *empty_tuple = PyTuple_New(0); + PyObject *ret = PyObject_Call(ptr, empty_tuple, NULL); + Py_DECREF(empty_tuple); + Py_DECREF(ptr); + if (!PyLong_Check(ret)) {{ + PyErr_SetString(PyExc_TypeError, "data_ptr method of Pointer object must return 64-bit int"); + ptr_info.valid = false; + return ptr_info; + }} + ptr_info.dev_ptr = reinterpret_cast(PyLong_AsUnsignedLongLong(ret)); + if(!ptr_info.dev_ptr) + return ptr_info; + Py_DECREF(ret); // Thanks ChatGPT! + return ptr_info; + }} + PyErr_SetString(PyExc_TypeError, "Pointer argument must be either uint64 or have data_ptr method"); + return ptr_info; +}} + +static PyObject* launch(PyObject* self, PyObject* args) {{ + int gridX, gridY, gridZ; + PyObject *launch_enter_hook = NULL; + PyObject *launch_exit_hook = NULL; + PyObject *kernel_metadata = NULL; + PyObject *launch_metadata = NULL; + {' '.join([f"{_extracted_type(ty)} _arg{i}; " for i, ty in signature.items()])} + if(!PyArg_ParseTuple(args, \"{format}\", &gridX, &gridY, &gridZ, + &kernel_metadata, &launch_metadata, + &launch_enter_hook, &launch_exit_hook {args_list})) {{ + return NULL; + }} + + // [CPULauncher-specific]: We don't need the metadata below but just put them + // here anyway to be consistent with others. + // This will make updating the driver easier in the future. + + // int num_warps, num_ctas, shared_memory, clusterDimX, clusterDimY, clusterDimZ; + // if (!PyArg_ParseTuple(kernel_metadata, \"iiiiii\", &num_warps, &num_ctas, &shared_memory, &clusterDimX, &clusterDimY, &clusterDimZ)) {{ + // PyErr_SetString(PyExc_TypeError, "kernel_metadata must be a tuple"); + // return NULL; + // }} + + // extract launch metadata + if (launch_enter_hook != Py_None){{ + PyObject* args = Py_BuildValue("(O)", launch_metadata); + PyObject* ret = PyObject_CallObject(launch_enter_hook, args); + Py_DECREF(args); + if (!ret) + return NULL; + }} + + // raise exception asap + {"; ".join([f"DevicePtrInfo ptr_info{i} = getPointer(_arg{i}, {i}); if (!ptr_info{i}.valid) return NULL;" if ty[0] == "*" else "" for i, ty in signature.items()])}; + _launch(gridX, gridY, gridZ, {', '.join(f"ptr_info{i}.dev_ptr" if ty[0]=="*" else f"_arg{i}"for i, ty in signature.items())}); + + if (PyErr_Occurred()) {{ + return NULL; + }} + if(launch_exit_hook != Py_None){{ + PyObject* args = Py_BuildValue("(O)", launch_metadata); + PyObject* ret = PyObject_CallObject(launch_exit_hook, args); + Py_DECREF(args); + if (!ret) + return NULL; + }} + + // return None + Py_INCREF(Py_None); + return Py_None; +}} + +static PyMethodDef ModuleMethods[] = {{ + {{"launch", launch, METH_VARARGS, "Entry point for all kernels with this signature"}}, + {{NULL, NULL, 0, NULL}} // sentinel +}}; + +static struct PyModuleDef ModuleDef = {{ + PyModuleDef_HEAD_INIT, + \"__triton_shared_ref_cpu_kernel_launcher\", + NULL, //documentation + -1, //size + ModuleMethods +}}; + +PyMODINIT_FUNC PyInit___triton_shared_ref_cpu_kernel_launcher(void) {{ + PyObject *m = PyModule_Create(&ModuleDef); + if(m == NULL) {{ + return NULL; + }} + PyModule_AddFunctions(m, ModuleMethods); + return m; +}} +""" + + +def compile_module(launcher_src, kernel_placeholder_name): + py_version = sys.version_info + if platform.system() == "Windows": + py_include_dir = os.path.join(sys.base_prefix, 'include') + py_lib_dir = os.path.join(sys.base_prefix, 'libs') + py_lib = '{name}{major}{minor}.lib'.format(name="python", major=py_version.major, minor=py_version.minor) + else: + py_include_dir = os.path.join(sys.base_prefix, 'include', f'python{sys.version_info.major}.{sys.version_info.minor}') + py_lib_dir = os.path.join(sys.base_prefix, 'lib') + py_lib = '{name}{major}.{minor}'.format(name="python", major=py_version.major, minor=py_version.minor) + cpu_backend_path = Path(__file__).resolve().parent + include_dir = os.path.join(cpu_backend_path, "include") + + def launch( + gridX, gridY, gridZ, stream, cu_function, + kernel_metadata, launch_metadata, + launch_enter_hook, launch_exit_hook, *args): + # Unlike CUDA/HIP, we cannot easily pass function pointer across different pybind libraries. + # Let's compile one kernel every time. + # The cu_function parameter actually contains our kernel obj. + # See CPUUtils.load_binary method. + kernel_obj = cu_function + kernel_name = kernel_metadata[6] # see pack_metadata in compiler.py + src = launcher_src.replace(kernel_placeholder_name, kernel_name) + + key = hashlib.md5(src.encode("utf-8") + kernel_obj).hexdigest() + cache = get_cache_manager(key) + name = "__triton_shared_ref_cpu_kernel_launcher" + + if platform.system() == "Windows": + filename = f"{name}.pyd" + else: + filename = f"{name}.so" + cache_path = cache.get_file(filename) + + if cache_path is None: + with tempfile.TemporaryDirectory() as tmpdir: + if platform.system() == "Windows": + obj_path = os.path.join(tmpdir, "kernel.obj") + launcher_src_path = os.path.join(tmpdir, "main.cxx") + so_path = os.path.join(tmpdir, "kernel.pyd") + Path(obj_path).write_bytes(kernel_obj) + Path(launcher_src_path).write_text(src) + # Compile it together. + subprocess.check_call([ + "cl", "/LD", "/std:c++17", launcher_src_path, obj_path, + f"-I{py_include_dir}", f"-I{include_dir}", "/link", f"/LIBPATH:{py_lib_dir}", + "/link", f"{py_lib}", f"/OUT:{so_path}" + ]) + else: + obj_path = os.path.join(tmpdir, "kernel.o") + launcher_src_path = os.path.join(tmpdir, "main.cxx") + so_path = os.path.join(tmpdir, "kernel.so") + Path(obj_path).write_bytes(kernel_obj) + Path(launcher_src_path).write_text(src) + # Compile it together. + subprocess.check_call([ + "g++", "-std=c++17", launcher_src_path, obj_path, + f"-I{py_include_dir}", f"-I{include_dir}", f"-L{py_lib_dir}", + "-shared", f"-l{py_lib}", "-fPIC", "-o", so_path + ]) + + with open(so_path, "rb") as f: + cache_path = cache.put(f.read(), filename, binary=True) + + # Load and launch the compiled kernel. + spec = importlib.util.spec_from_file_location(name, cache_path) + if spec is None: + raise RuntimeError(f"Cannot find {name} module in {cache_path}") + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + return mod.launch(gridX, gridY, gridZ, + kernel_metadata, launch_metadata, + launch_enter_hook, launch_exit_hook, + *args) + + return launch + + +class CPULauncher(object): + + def __init__(self, src, metadata): + kernel_placeholder_name = "KERNEL_NAME_PLACEHOLDER" + + constants = src.constants if hasattr(src, "constants") else dict() + cst_key = lambda i: src.fn.arg_names.index(i) if isinstance(i, str) else i + constants = {cst_key(key): value for key, value in constants.items()} + signature = {cst_key(key): value for key, value in src.signature.items()} + launcher_src = _generate_launcher(constants, signature, kernel_placeholder_name) + # Later KERNEL_NAME_PLACEHOLDER will be used to assign the kernel name + # in the following launch function. + self.launch = compile_module(launcher_src, kernel_placeholder_name) + + def __call__(self, *args, **kwargs): + self.launch(*args, **kwargs) + + + +class CPUUtils(object): + def __new__(cls): + if not hasattr(cls, "instance"): + cls.instance = super(CPUUtils, cls).__new__(cls) + return cls.instance + + # Note: + # nvidia and amd backends have their corresponding driver.c file that exposes + # get_device_properties and load_binary using python bindings. + # (see third_party/nvidia/backend/driver.c) + # These methods are then used in compiler.py to initialize handles before running + # the triton kernels. + # Since we recompile the kernel every time (see compile_module above), + # and the metadata generated by these functions aren't applicable to the cpu + # backend, just define the same functions with dummy implementation. + @staticmethod + def get_device_properties(device): + return { + "max_shared_mem": 2 ** 20, + "multiprocessor_count": None, + "sm_clock_rate": None, + "mem_clock_rate": None, + "mem_bus_width": None + } + + # Important note: + # Since we cannot easy pass function pointers around, we pass along the + # obj of the kernel so that compile_module above can recompile the + # module every time. + @staticmethod + def load_binary(name, kernel_obj, shared, device): + return ( + None, # module + kernel_obj, # function + None, # n_regs + None # n_spills + ) + + +class CPUDriver(DriverBase): + + def __init__(self): + super().__init__() + self.utils = CPUUtils() + self.launcher_cls = CPULauncher + self.binary_ext = "obj" + + # CPU driver won't be automatically chosen unless explicitly set through + # triton.runtime.driver.set_active(CPUDriver()) + @staticmethod + def is_active(): + return False + + def get_benchmarker(self): + from triton.testing import do_bench + return do_bench + + def get_device_capability(self): + return ("cpu", 0) + + def get_current_stream(self, device): + return None + + def get_current_device(self): + # CPU doesn't have a device to return. Return something. + return "cpu" + + def set_current_device(self, device): + # CPU doesn't have a device to set + assert device == "cpu" + return + + def get_current_target(self): + return GPUTarget("cpu", 0, 0) + + def get_active_torch_device(self): + import torch + return torch.device("cpu") + + def assemble_tensormap_to_arg(self, tensormaps_info, args): + return args diff --git a/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/CRunnerUtils.cpp b/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/CRunnerUtils.cpp new file mode 100755 index 00000000..48e2afbf --- /dev/null +++ b/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/CRunnerUtils.cpp @@ -0,0 +1,192 @@ +//===- CRunnerUtils.cpp - Utils for MLIR execution ------------------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file implements basic functions to manipulate structured MLIR types at +// runtime. Entities in this file are meant to be retargetable, including on +// targets without a C++ runtime, and must be kept C compatible. +// +//===----------------------------------------------------------------------===// + +#include "CRunnerUtils.h" +#include "Msan.h" + +#ifndef _WIN32 +#if defined(__FreeBSD__) || defined(__NetBSD__) || defined(__OpenBSD__) || \ + defined(__DragonFly__) +#include +#else +#include +#endif +#include +#else +#include "malloc.h" +#endif // _WIN32 + +#include +#include +#include +#include +#include +#include + +#ifdef MLIR_CRUNNERUTILS_DEFINE_FUNCTIONS + +namespace { +template +void stdSort(uint64_t n, V *p) { + std::sort(p, p + n); +} + +} // namespace + +// Small runtime support "lib" for vector.print lowering. +// By providing elementary printing methods only, this +// library can remain fully unaware of low-level implementation +// details of our vectors. Also useful for direct LLVM IR output. +extern "C" void printI64(int64_t i) { fprintf(stdout, "%" PRId64, i); } +extern "C" void printU64(uint64_t u) { fprintf(stdout, "%" PRIu64, u); } +extern "C" void printF32(float f) { fprintf(stdout, "%g", f); } +extern "C" void printF64(double d) { fprintf(stdout, "%lg", d); } +extern "C" void printString(char const *s) { fputs(s, stdout); } +extern "C" void printOpen() { fputs("( ", stdout); } +extern "C" void printClose() { fputs(" )", stdout); } +extern "C" void printComma() { fputs(", ", stdout); } +extern "C" void printNewline() { fputc('\n', stdout); } + +extern "C" void memrefCopy(int64_t elemSize, UnrankedMemRefType *srcArg, + UnrankedMemRefType *dstArg) { + DynamicMemRefType src(*srcArg); + DynamicMemRefType dst(*dstArg); + + int64_t rank = src.rank; + MLIR_MSAN_MEMORY_IS_INITIALIZED(src.sizes, rank * sizeof(int64_t)); + + // Handle empty shapes -> nothing to copy. + for (int rankp = 0; rankp < rank; ++rankp) + if (src.sizes[rankp] == 0) + return; + + char *srcPtr = src.data + src.offset * elemSize; + char *dstPtr = dst.data + dst.offset * elemSize; + + if (rank == 0) { + memcpy(dstPtr, srcPtr, elemSize); + return; + } + + int64_t *indices = static_cast(alloca(sizeof(int64_t) * rank)); + int64_t *srcStrides = static_cast(alloca(sizeof(int64_t) * rank)); + int64_t *dstStrides = static_cast(alloca(sizeof(int64_t) * rank)); + + // Initialize index and scale strides. + for (int rankp = 0; rankp < rank; ++rankp) { + indices[rankp] = 0; + srcStrides[rankp] = src.strides[rankp] * elemSize; + dstStrides[rankp] = dst.strides[rankp] * elemSize; + } + + int64_t readIndex = 0, writeIndex = 0; + for (;;) { + // Copy over the element, byte by byte. + memcpy(dstPtr + writeIndex, srcPtr + readIndex, elemSize); + // Advance index and read position. + for (int64_t axis = rank - 1; axis >= 0; --axis) { + // Advance at current axis. + auto newIndex = ++indices[axis]; + readIndex += srcStrides[axis]; + writeIndex += dstStrides[axis]; + // If this is a valid index, we have our next index, so continue copying. + if (src.sizes[axis] != newIndex) + break; + // We reached the end of this axis. If this is axis 0, we are done. + if (axis == 0) + return; + // Else, reset to 0 and undo the advancement of the linear index that + // this axis had. Then continue with the axis one outer. + indices[axis] = 0; + readIndex -= src.sizes[axis] * srcStrides[axis]; + writeIndex -= dst.sizes[axis] * dstStrides[axis]; + } + } +} + +/// Prints GFLOPS rating. +extern "C" void printFlops(double flops) { + fprintf(stderr, "%lf GFLOPS\n", flops / 1.0E9); +} + +/// Returns the number of seconds since Epoch 1970-01-01 00:00:00 +0000 (UTC). +extern "C" double rtclock() { +#ifndef _WIN32 + struct timeval tp; + int stat = gettimeofday(&tp, nullptr); + if (stat != 0) + fprintf(stderr, "Error returning time from gettimeofday: %d\n", stat); + return (tp.tv_sec + tp.tv_usec * 1.0e-6); +#else + fprintf(stderr, "Timing utility not implemented on Windows\n"); + return 0.0; +#endif // _WIN32 +} + +extern "C" void *mlirAlloc(uint64_t size) { return malloc(size); } + +extern "C" void *mlirAlignedAlloc(uint64_t alignment, uint64_t size) { +#ifdef _WIN32 + return _aligned_malloc(size, alignment); +#elif defined(__APPLE__) + // aligned_alloc was added in MacOS 10.15. Fall back to posix_memalign to also + // support older versions. + void *result = nullptr; + (void)::posix_memalign(&result, alignment, size); + return result; +#else + return aligned_alloc(alignment, size); +#endif +} + +extern "C" void mlirFree(void *ptr) { free(ptr); } + +extern "C" void mlirAlignedFree(void *ptr) { +#ifdef _WIN32 + _aligned_free(ptr); +#else + free(ptr); +#endif +} + +extern "C" void *rtsrand(uint64_t s) { + // Standard mersenne_twister_engine seeded with s. + return new std::mt19937(s); +} + +extern "C" uint64_t rtrand(void *g, uint64_t m) { + std::mt19937 *generator = static_cast(g); + std::uniform_int_distribution distrib(0, m); + return distrib(*generator); +} + +extern "C" void rtdrand(void *g) { + std::mt19937 *generator = static_cast(g); + delete generator; +} + +#define IMPL_STDSORT(VNAME, V) \ + extern "C" void _mlir_ciface_stdSort##VNAME(uint64_t n, \ + StridedMemRefType *vref) { \ + assert(vref); \ + assert(vref->strides[0] == 1); \ + V *values = vref->data + vref->offset; \ + stdSort(n, values); \ + } +IMPL_STDSORT(I64, int64_t) +IMPL_STDSORT(F64, double) +IMPL_STDSORT(F32, float) +#undef IMPL_STDSORT + +#endif // MLIR_CRUNNERUTILS_DEFINE_FUNCTIONS diff --git a/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/CRunnerUtils.h b/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/CRunnerUtils.h new file mode 100755 index 00000000..76b04145 --- /dev/null +++ b/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/CRunnerUtils.h @@ -0,0 +1,499 @@ +//===- CRunnerUtils.h - Utils for debugging MLIR execution ----------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file declares basic classes and functions to manipulate structured MLIR +// types at runtime. Entities in this file must be compliant with C++11 and be +// retargetable, including on targets without a C++ runtime. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_EXECUTIONENGINE_CRUNNERUTILS_H +#define MLIR_EXECUTIONENGINE_CRUNNERUTILS_H + +#ifdef _WIN32 +#ifndef MLIR_CRUNNERUTILS_EXPORT +#ifdef mlir_c_runner_utils_EXPORTS +// We are building this library +#define MLIR_CRUNNERUTILS_EXPORT __declspec(dllexport) +#define MLIR_CRUNNERUTILS_DEFINE_FUNCTIONS +#else +// We are using this library +#define MLIR_CRUNNERUTILS_EXPORT __declspec(dllimport) +#endif // mlir_c_runner_utils_EXPORTS +#endif // MLIR_CRUNNERUTILS_EXPORT +#else // _WIN32 +// Non-windows: use visibility attributes. +#define MLIR_CRUNNERUTILS_EXPORT __attribute__((visibility("default"))) +#define MLIR_CRUNNERUTILS_DEFINE_FUNCTIONS +#endif // _WIN32 + +#include +#include +#include +#include +#include + +//===----------------------------------------------------------------------===// +// Codegen-compatible structures for Vector type. +//===----------------------------------------------------------------------===// +namespace mlir { +namespace detail { + +constexpr bool isPowerOf2(int n) { return (!(n & (n - 1))); } + +constexpr unsigned nextPowerOf2(int n) { + return (n <= 1) ? 1 : (isPowerOf2(n) ? n : (2 * nextPowerOf2((n + 1) / 2))); +} + +template +struct Vector1D; + +template +struct Vector1D { + Vector1D() { + static_assert(detail::nextPowerOf2(sizeof(T[Dim])) == sizeof(T[Dim]), + "size error"); + } + inline T &operator[](unsigned i) { return vector[i]; } + inline const T &operator[](unsigned i) const { return vector[i]; } + +private: + T vector[Dim]; +}; + +// 1-D vector, padded to the next power of 2 allocation. +// Specialization occurs to avoid zero size arrays (which fail in -Werror). +template +struct Vector1D { + Vector1D() { + static_assert(nextPowerOf2(sizeof(T[Dim])) > sizeof(T[Dim]), "size error"); + static_assert(nextPowerOf2(sizeof(T[Dim])) < 2 * sizeof(T[Dim]), + "size error"); + } + inline T &operator[](unsigned i) { return vector[i]; } + inline const T &operator[](unsigned i) const { return vector[i]; } + +private: + T vector[Dim]; + char padding[nextPowerOf2(sizeof(T[Dim])) - sizeof(T[Dim])]; +}; +} // namespace detail +} // namespace mlir + +// N-D vectors recurse down to 1-D. +template +struct Vector { + inline Vector &operator[](unsigned i) { return vector[i]; } + inline const Vector &operator[](unsigned i) const { + return vector[i]; + } + +private: + Vector vector[Dim]; +}; + +// 1-D vectors in LLVM are automatically padded to the next power of 2. +// We insert explicit padding in to account for this. +template +struct Vector + : public mlir::detail::Vector1D { +}; + +template +using Vector1D = Vector; +template +using Vector2D = Vector; +template +using Vector3D = Vector; +template +using Vector4D = Vector; + +template +void dropFront(int64_t arr[N], int64_t *res) { + for (unsigned i = 1; i < N; ++i) + *(res + i - 1) = arr[i]; +} + +//===----------------------------------------------------------------------===// +// Codegen-compatible structures for StridedMemRef type. +//===----------------------------------------------------------------------===// +template +class StridedMemrefIterator; + +/// StridedMemRef descriptor type with static rank. +template +struct StridedMemRefType { + T *basePtr; + T *data; + int64_t offset; + int64_t sizes[N]; + int64_t strides[N]; + + template ().begin())> + T &operator[](Range &&indices) { + assert(indices.size() == N && + "indices should match rank in memref subscript"); + int64_t curOffset = offset; + for (int dim = N - 1; dim >= 0; --dim) { + int64_t currentIndex = *(indices.begin() + dim); + assert(currentIndex < sizes[dim] && "Index overflow"); + curOffset += currentIndex * strides[dim]; + } + return data[curOffset]; + } + + StridedMemrefIterator begin() { return {*this, offset}; } + StridedMemrefIterator end() { return {*this, -1}; } + + // This operator[] is extremely slow and only for sugaring purposes. + StridedMemRefType operator[](int64_t idx) { + StridedMemRefType res; + res.basePtr = basePtr; + res.data = data; + res.offset = offset + idx * strides[0]; + dropFront(sizes, res.sizes); + dropFront(strides, res.strides); + return res; + } +}; + +/// StridedMemRef descriptor type specialized for rank 1. +template +struct StridedMemRefType { + T *basePtr; + T *data; + int64_t offset; + int64_t sizes[1]; + int64_t strides[1]; + + template ().begin())> + T &operator[](Range indices) { + assert(indices.size() == 1 && + "indices should match rank in memref subscript"); + return (*this)[*indices.begin()]; + } + + StridedMemrefIterator begin() { return {*this, offset}; } + StridedMemrefIterator end() { return {*this, -1}; } + + T &operator[](int64_t idx) { return *(data + offset + idx * strides[0]); } +}; + +/// StridedMemRef descriptor type specialized for rank 0. +template +struct StridedMemRefType { + T *basePtr; + T *data; + int64_t offset; + + template ().begin())> + T &operator[](Range indices) { + assert((indices.size() == 0) && + "Expect empty indices for 0-rank memref subscript"); + return data[offset]; + } + + StridedMemrefIterator begin() { return {*this, offset}; } + StridedMemrefIterator end() { return {*this, offset + 1}; } +}; + +/// Iterate over all elements in a strided memref. +template +class StridedMemrefIterator { +public: + using iterator_category = std::forward_iterator_tag; + using value_type = T; + using difference_type = std::ptrdiff_t; + using pointer = T *; + using reference = T &; + + StridedMemrefIterator(StridedMemRefType &descriptor, + int64_t offset = 0) + : offset(offset), descriptor(&descriptor) {} + StridedMemrefIterator &operator++() { + int dim = Rank - 1; + while (dim >= 0 && indices[dim] == (descriptor->sizes[dim] - 1)) { + offset -= indices[dim] * descriptor->strides[dim]; + indices[dim] = 0; + --dim; + } + if (dim < 0) { + offset = -1; + return *this; + } + ++indices[dim]; + offset += descriptor->strides[dim]; + return *this; + } + + reference operator*() { return descriptor->data[offset]; } + pointer operator->() { return &descriptor->data[offset]; } + + const std::array &getIndices() { return indices; } + + bool operator==(const StridedMemrefIterator &other) const { + return other.offset == offset && other.descriptor == descriptor; + } + + bool operator!=(const StridedMemrefIterator &other) const { + return !(*this == other); + } + +private: + /// Offset in the buffer. This can be derived from the indices and the + /// descriptor. + int64_t offset = 0; + + /// Array of indices in the multi-dimensional memref. + std::array indices = {}; + + /// Descriptor for the strided memref. + StridedMemRefType *descriptor; +}; + +/// Iterate over all elements in a 0-ranked strided memref. +template +class StridedMemrefIterator { +public: + using iterator_category = std::forward_iterator_tag; + using value_type = T; + using difference_type = std::ptrdiff_t; + using pointer = T *; + using reference = T &; + + StridedMemrefIterator(StridedMemRefType &descriptor, int64_t offset = 0) + : elt(descriptor.data + offset) {} + + StridedMemrefIterator &operator++() { + ++elt; + return *this; + } + + reference operator*() { return *elt; } + pointer operator->() { return elt; } + + // There are no indices for a 0-ranked memref, but this API is provided for + // consistency with the general case. + const std::array &getIndices() { + // Since this is a 0-array of indices we can keep a single global const + // copy. + static const std::array indices = {}; + return indices; + } + + bool operator==(const StridedMemrefIterator &other) const { + return other.elt == elt; + } + + bool operator!=(const StridedMemrefIterator &other) const { + return !(*this == other); + } + +private: + /// Pointer to the single element in the zero-ranked memref. + T *elt; +}; + +//===----------------------------------------------------------------------===// +// Codegen-compatible structure for UnrankedMemRef type. +//===----------------------------------------------------------------------===// +// Unranked MemRef +template +struct UnrankedMemRefType { + int64_t rank; + void *descriptor; +}; + +//===----------------------------------------------------------------------===// +// DynamicMemRefType type. +//===----------------------------------------------------------------------===// +template +class DynamicMemRefIterator; + +// A reference to one of the StridedMemRef types. +template +class DynamicMemRefType { +public: + int64_t rank; + T *basePtr; + T *data; + int64_t offset; + const int64_t *sizes; + const int64_t *strides; + + explicit DynamicMemRefType(const StridedMemRefType &memRef) + : rank(0), basePtr(memRef.basePtr), data(memRef.data), + offset(memRef.offset), sizes(nullptr), strides(nullptr) {} + template + explicit DynamicMemRefType(const StridedMemRefType &memRef) + : rank(N), basePtr(memRef.basePtr), data(memRef.data), + offset(memRef.offset), sizes(memRef.sizes), strides(memRef.strides) {} + explicit DynamicMemRefType(const ::UnrankedMemRefType &memRef) + : rank(memRef.rank) { + auto *desc = static_cast *>(memRef.descriptor); + basePtr = desc->basePtr; + data = desc->data; + offset = desc->offset; + sizes = rank == 0 ? nullptr : desc->sizes; + strides = sizes + rank; + } + + template ().begin())> + T &operator[](Range &&indices) { + assert(indices.size() == rank && + "indices should match rank in memref subscript"); + if (rank == 0) + return data[offset]; + + int64_t curOffset = offset; + for (int dim = rank - 1; dim >= 0; --dim) { + int64_t currentIndex = *(indices.begin() + dim); + assert(currentIndex < sizes[dim] && "Index overflow"); + curOffset += currentIndex * strides[dim]; + } + return data[curOffset]; + } + + DynamicMemRefIterator begin() { return {*this, offset}; } + DynamicMemRefIterator end() { return {*this, -1}; } + + // This operator[] is extremely slow and only for sugaring purposes. + DynamicMemRefType operator[](int64_t idx) { + assert(rank > 0 && "can't make a subscript of a zero ranked array"); + + DynamicMemRefType res(*this); + --res.rank; + res.offset += idx * res.strides[0]; + ++res.sizes; + ++res.strides; + return res; + } + + // This operator* can be used in conjunction with the previous operator[] in + // order to access the underlying value in case of zero-ranked memref. + T &operator*() { + assert(rank == 0 && "not a zero-ranked memRef"); + return data[offset]; + } +}; + +/// Iterate over all elements in a dynamic memref. +template +class DynamicMemRefIterator { +public: + using iterator_category = std::forward_iterator_tag; + using value_type = T; + using difference_type = std::ptrdiff_t; + using pointer = T *; + using reference = T &; + + DynamicMemRefIterator(DynamicMemRefType &descriptor, int64_t offset = 0) + : offset(offset), descriptor(&descriptor) { + indices.resize(descriptor.rank, 0); + } + + DynamicMemRefIterator &operator++() { + if (descriptor->rank == 0) { + offset = -1; + return *this; + } + + int dim = descriptor->rank - 1; + + while (dim >= 0 && indices[dim] == (descriptor->sizes[dim] - 1)) { + offset -= indices[dim] * descriptor->strides[dim]; + indices[dim] = 0; + --dim; + } + + if (dim < 0) { + offset = -1; + return *this; + } + + ++indices[dim]; + offset += descriptor->strides[dim]; + return *this; + } + + reference operator*() { return descriptor->data[offset]; } + pointer operator->() { return &descriptor->data[offset]; } + + const std::vector &getIndices() { return indices; } + + bool operator==(const DynamicMemRefIterator &other) const { + return other.offset == offset && other.descriptor == descriptor; + } + + bool operator!=(const DynamicMemRefIterator &other) const { + return !(*this == other); + } + +private: + /// Offset in the buffer. This can be derived from the indices and the + /// descriptor. + int64_t offset = 0; + + /// Array of indices in the multi-dimensional memref. + std::vector indices = {}; + + /// Descriptor for the dynamic memref. + DynamicMemRefType *descriptor; +}; + +//===----------------------------------------------------------------------===// +// Small runtime support library for memref.copy lowering during codegen. +//===----------------------------------------------------------------------===// +extern "C" MLIR_CRUNNERUTILS_EXPORT void +memrefCopy(int64_t elemSize, ::UnrankedMemRefType *src, + ::UnrankedMemRefType *dst); + +//===----------------------------------------------------------------------===// +// Small runtime support library for vector.print lowering during codegen. +//===----------------------------------------------------------------------===// +extern "C" MLIR_CRUNNERUTILS_EXPORT void printI64(int64_t i); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printU64(uint64_t u); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printF32(float f); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printF64(double d); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printString(char const *s); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printOpen(); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printClose(); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printComma(); +extern "C" MLIR_CRUNNERUTILS_EXPORT void printNewline(); + +//===----------------------------------------------------------------------===// +// Small runtime support library for timing execution and printing GFLOPS +//===----------------------------------------------------------------------===// +extern "C" MLIR_CRUNNERUTILS_EXPORT void printFlops(double flops); +extern "C" MLIR_CRUNNERUTILS_EXPORT double rtclock(); + +//===----------------------------------------------------------------------===// +// Runtime support library for random number generation. +//===----------------------------------------------------------------------===// +// Uses a seed to initialize a random generator and returns the generator. +extern "C" MLIR_CRUNNERUTILS_EXPORT void *rtsrand(uint64_t s); +// Returns a random number in the range of [0, m). +extern "C" MLIR_CRUNNERUTILS_EXPORT uint64_t rtrand(void *, uint64_t m); +// Deletes the random number generator. +extern "C" MLIR_CRUNNERUTILS_EXPORT void rtdrand(void *); + +//===----------------------------------------------------------------------===// +// Runtime support library to allow the use of std::sort in MLIR program. +//===----------------------------------------------------------------------===// +extern "C" MLIR_CRUNNERUTILS_EXPORT void +_mlir_ciface_stdSortI64(uint64_t n, StridedMemRefType *vref); +extern "C" MLIR_CRUNNERUTILS_EXPORT void +_mlir_ciface_stdSortF64(uint64_t n, StridedMemRefType *vref); +extern "C" MLIR_CRUNNERUTILS_EXPORT void +_mlir_ciface_stdSortF32(uint64_t n, StridedMemRefType *vref); +#endif // MLIR_EXECUTIONENGINE_CRUNNERUTILS_H diff --git a/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/Msan.h b/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/Msan.h new file mode 100755 index 00000000..ee94660a --- /dev/null +++ b/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/Msan.h @@ -0,0 +1,35 @@ +//===- Msan.h - Utils related to the memory sanitizer ---------------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file declares and defines macros related to msan. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_EXECUTIONENGINE_MSAN_H +#define MLIR_EXECUTIONENGINE_MSAN_H + +// Memory sanitizer currently can't be enabled for the jit-compiled code, and +// to suppress msan warnings we need to unpoison pointers and pointed-to +// datastructures before they can be accessed. + +#ifndef __has_feature +#define __has_feature(x) 0 +#endif + +#if __has_feature(memory_sanitizer) && !defined(MLIR_MEMORY_SANITIZER) +#define MLIR_MEMORY_SANITIZER +#endif + +#if defined(MLIR_MEMORY_SANITIZER) +#include +#define MLIR_MSAN_MEMORY_IS_INITIALIZED(p, s) __msan_unpoison((p), (s)) +#else // Memory sanitizer: OFF +#define MLIR_MSAN_MEMORY_IS_INITIALIZED(p, s) +#endif // MLIR_MEMORY_SANITIZER + +#endif // MLIR_EXECUTIONENGINE_MSAN_H diff --git a/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/version.txt b/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/version.txt new file mode 100755 index 00000000..c3f15e55 --- /dev/null +++ b/third_party/wafer/third_party/flir/backend/include/ExecutionEngine/version.txt @@ -0,0 +1 @@ +https://github.com/llvm/llvm-project/commit/3be3883e6d67bf908fd12b51219075293ebb3dff diff --git a/third_party/wafer/third_party/flir/backend/name.conf b/third_party/wafer/third_party/flir/backend/name.conf new file mode 100755 index 00000000..38a5dd30 --- /dev/null +++ b/third_party/wafer/third_party/flir/backend/name.conf @@ -0,0 +1 @@ +triton_shared \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/CMakeLists.txt b/third_party/wafer/third_party/flir/include/CMakeLists.txt new file mode 100755 index 00000000..2f72db53 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/CMakeLists.txt @@ -0,0 +1,6 @@ +add_subdirectory(mlir-ext) +add_subdirectory(triton-shared) +add_subdirectory(npu) +if (FLIR_BUILD_INCUBATED) + add_subdirectory(incubated) +endif() \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/incubated/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/CMakeLists.txt new file mode 100755 index 00000000..629c08af --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/CMakeLists.txt @@ -0,0 +1,2 @@ +add_subdirectory(Conversion) +add_subdirectory(Dialect) diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Conversion/CMakeLists.txt new file mode 100755 index 00000000..45978367 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/CMakeLists.txt @@ -0,0 +1,5 @@ +add_subdirectory(TritonToAnnotation) +add_subdirectory(TritonToLinalgIncubated) +add_subdirectory(DiscreteMaskAccessConversion) +add_subdirectory(TritonToUnstructureIncubated) +add_subdirectory(TritonToStructuredIncubated) diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/CMakeLists.txt new file mode 100755 index 00000000..567a119e --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name DiscreteMaskAccessConversion) +add_public_tablegen_target(DiscreteMaskAccessConversionPassIncGen) \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.h new file mode 100755 index 00000000..fed34acf --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.h @@ -0,0 +1,66 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_DISCRETEMASKACCESSCONVERSION_H +#define TRITON_ADAPTER_DISCRETEMASKACCESSCONVERSION_H + +#include "mlir/Pass/Pass.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/IR/PatternMatch.h" + +#define GEN_PASS_DECL_DISCRETEMASKACCESSCONVERSION +#include "incubated/Conversion/DiscreteMaskAccessConversion/Passes.h.inc" + +#define GEN_PASS_DEF_DISCRETEMASKACCESSCONVERSION +#include "incubated/Conversion/DiscreteMaskAccessConversion/Passes.h.inc" + +extern bool compileOn91095Flag; +extern bool forceSimtTemplateFlag; + +namespace mlir { +namespace triton { + +std::unique_ptr> createDiscreteMaskAccessConversionPass( + const DiscreteMaskAccessConversionOptions &options = {}); + +} // namespace triton +} // namespace mlir + +namespace { + +using namespace mlir; +using namespace triton; + +class DiscreteMaskAccessConversionPass + : public ::impl::DiscreteMaskAccessConversionBase< + DiscreteMaskAccessConversionPass> { +public: + explicit DiscreteMaskAccessConversionPass( + const DiscreteMaskAccessConversionOptions &options); + + void runOnOperation() override; +}; + +} // namespace + +#endif // DISCRETE_MASK_ACCESS_CONVERSION_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/Passes.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/Passes.h new file mode 100755 index 00000000..62d6990b --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/Passes.h @@ -0,0 +1,37 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_DISCRETE_MASK_ACCESS_CONVERSION_PASSES_H +#define TRITON_ADAPTER_DISCRETE_MASK_ACCESS_CONVERSION_PASSES_H + +#include "incubated/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "incubated/Conversion/DiscreteMaskAccessConversion/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // TRITON_ADAPTER_DISCRETE_MASK_ACCESS_CONVERSION_PASSES_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/Passes.td b/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/Passes.td new file mode 100755 index 00000000..bee8cb9c --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/Passes.td @@ -0,0 +1,24 @@ +/* + * Copyright (c) Huawei Technologies Co. + * Licensed under the MIT license. + */ + +#ifndef DISCRETE_MASK_ACCESS_CONVERSION_PASSES +#define DISCRETE_MASK_ACCESS_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def DiscreteMaskAccessConversion : Pass<"discrete-mask-access-conversion", "mlir::ModuleOp"> { + let summary = "Recognize and convert discrete mask memory access"; + let constructor = "triton::createDiscreteMaskAccessConversionPass()"; + let options = [ + Option<"compileOn91095", "compile-on-910-95", + "bool", /*default*/"false", + "compile on 910_95">, + Option<"forceSimtTemplate", "force-simt-template", + "bool", /*default*/"false", + "force to use simt template"> + ]; +} + +#endif // DISCRETE_MASK_ACCESS_CONVERSION_PASSES diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/CMakeLists.txt new file mode 100755 index 00000000..69d72d21 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToAnnotation) +add_public_tablegen_target(TritonToAnnotationConversionPassIncGen) \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/Passes.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/Passes.h new file mode 100755 index 00000000..58275105 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/Passes.h @@ -0,0 +1,43 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_TRITON_TO_ANNOTATION_CONVERSION_PASSES_H +#define TRITON_ADAPTER_TRITON_TO_ANNOTATION_CONVERSION_PASSES_H + +#include "mlir/Pass/Pass.h" + +namespace mlir { +// Forward declarations. +class ModuleOp; + +namespace triton { + +/// Creates a pass to convert Triton dialect to Annotation dialect. +std::unique_ptr> createTritonToAnnotationPass(); + +#define GEN_PASS_REGISTRATION +#include "incubated/Conversion/TritonToAnnotation/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // TRITON_ADAPTER_TRITON_TO_ANNOTATION_CONVERSION_PASSES_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/Passes.td b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/Passes.td new file mode 100755 index 00000000..871de7f4 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/Passes.td @@ -0,0 +1,17 @@ +/* + * Copyright (c) Huawei Technologies Co. + * Licensed under the MIT license. + */ + +#ifndef TRITON_TO_ANNOTATION_CONVERSION_PASSES +#define TRITON_TO_ANNOTATION_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToAnnotation : Pass<"triton-to-annotation", "mlir::ModuleOp"> { + let summary = "Convert Triton to Annotation dialect"; + let constructor = "triton::createTritonToAnnotationPass()"; + let dependentDialects = ["annotation::AnnotationDialect"]; +} + +#endif // TRITON_TO_ANNOTATION_CONVERSION_PASSES diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.h new file mode 100755 index 00000000..449e938e --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.h @@ -0,0 +1,360 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_ARGMINMAXCONVERTER_H +#define TRITON_ADAPTER_ARGMINMAXCONVERTER_H + +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "ConversionPatterns.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Transforms/DialectConversion.h" + +#define DEBUG_TYPE "triton-to-linalg" + +#include "llvm/Support/Debug.h" + +#include + +namespace TTOpConverters { +using namespace mlir; +using namespace triton; + +template +class ArgMinMaxBaseConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchTieBreakResult(Value currValue, Value currIndex, + Value reduceValue, Value reduceIndex, + mlir::Block::iterator &it, + Value &tileBreakValue) const { + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto eqCmpOp = dyn_cast(*it); + if (eqCmpOp) { + if (eqCmpOp.getPredicate() != arith::CmpFPredicate::OEQ || + currValue != eqCmpOp.getLhs() || reduceValue != eqCmpOp.getRhs()) { + return failure(); + } + } + + auto eqCmpIOp = dyn_cast(*it++); + if (eqCmpIOp) { + if (eqCmpIOp.getPredicate() != arith::CmpIPredicate::eq || + currValue != eqCmpIOp.getLhs() || reduceValue != eqCmpIOp.getRhs()) { + return failure(); + } + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto sltCmpOp = dyn_cast(*it++); + if (!sltCmpOp || sltCmpOp.getPredicate() != arith::CmpIPredicate::slt || + currIndex != sltCmpOp.getLhs() || reduceIndex != sltCmpOp.getRhs()) { + return failure(); + } + + // matching: %13 = arith.andi %11, %12 : i1 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto andOp = dyn_cast(*it++); + + Value cmpOp; + if (eqCmpOp) + cmpOp = eqCmpOp; + else + cmpOp = eqCmpIOp; + + if (!andOp || andOp.getLhs() != cmpOp || andOp.getRhs() != sltCmpOp) { + return failure(); + } + + tileBreakValue = andOp; + return success(); + } + + LogicalResult matchShouldUpdateValue(Value currValue, Value currIndex, + Value reduceValue, Value reduceIndex, + mlir::Block::iterator &it, + Value &shouldUpdate) const { + Value tieResult; + if (failed(matchTieBreakResult(currValue, currIndex, reduceValue, + reduceIndex, it, tieResult))) { + LLVM_DEBUG(llvm::dbgs() << "Tie break result match failed\n"); + return failure(); + } + + Value comparisonResult; + if (failed(T::matchComparisonResult(currValue, currIndex, reduceValue, + reduceIndex, it, comparisonResult))) { + LLVM_DEBUG(llvm::dbgs() << "Comparison result match failed\n"); + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto orOp = dyn_cast(*it++); + if (!orOp || orOp.getLhs() != comparisonResult || + orOp.getRhs() != tieResult) { + return failure(); + } + + shouldUpdate = orOp; + return success(); + } + + Value getInitTensor(ConversionPatternRewriter &rewriter, + ArrayRef shape, Value fillValue, + Location loc) const { + Value initTensor = + rewriter.create(loc, shape, fillValue.getType()); + return rewriter + .create(loc, ValueRange{fillValue}, + ValueRange{initTensor}) + .result(); + } + +public: + ArgMinMaxBaseConverter(MLIRContext *context) : OpConversionPattern(context) {} + + LogicalResult match(triton::ReduceOp op) const override final { + if (op.getBody()->getNumArguments() != 4) { + return failure(); + } + + auto block = op.getBody(); + auto ops = block->without_terminator(); + + Value currValue = block->getArgument(0); + Value currIndex = block->getArgument(1); + Value reduceValue = block->getArgument(2); + Value reduceIndex = block->getArgument(3); + + auto opsIt = ops.begin(); + Value shouldUpdate; + if (failed(matchShouldUpdateValue(currValue, currIndex, reduceValue, + reduceIndex, opsIt, shouldUpdate))) { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *opsIt << "\n"); + auto valueSelectOp = dyn_cast(*opsIt++); + if (!valueSelectOp || valueSelectOp.getCondition() != shouldUpdate || + currValue != valueSelectOp.getTrueValue() || + reduceValue != valueSelectOp.getFalseValue()) { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *opsIt << "\n"); + auto indexSelectOp = dyn_cast(*opsIt++); + if (indexSelectOp) { + if (indexSelectOp.getCondition() != shouldUpdate || + currIndex != indexSelectOp.getTrueValue() || + reduceIndex != indexSelectOp.getFalseValue()) { + return failure(); + } + } else { + return failure(); + } + if (!indexSelectOp || indexSelectOp.getCondition() != shouldUpdate || + currIndex != indexSelectOp.getTrueValue() || + reduceIndex != indexSelectOp.getFalseValue()) { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *opsIt << "\n"); + auto termOp = dyn_cast(*opsIt++); + if (!(termOp && termOp == block->getTerminator() && + termOp.getOperands() == + ArrayRef{valueSelectOp, indexSelectOp})) { + return failure(); + } + return success(); + } + + void rewrite(triton::ReduceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override final { + auto loc = op.getLoc(); + auto elemTypes = op.getElementTypes(); + + auto valueType = elemTypes[0]; + // tl.argmin reorder + auto block = op.getBody(); + bool isUnsigned = false; + if (isa(valueType)) { + arith::CmpFOp cmpFOp; + block->walk([&](arith::CmpFOp cmpOp) { + auto pred = cmpOp.getPredicate(); + if (pred == arith::CmpFPredicate::OEQ || + pred == arith::CmpFPredicate::ONE || + pred == arith::CmpFPredicate::UEQ || + pred == arith::CmpFPredicate::UNE) { + return WalkResult::advance(); + } else if (pred == arith::CmpFPredicate::OGT || + pred == arith::CmpFPredicate::OLT || + pred == arith::CmpFPredicate::UGT || + pred == arith::CmpFPredicate::ULT) { + cmpFOp = cmpOp; + return WalkResult::interrupt(); + } + return WalkResult::advance(); + }); + cmpFOp->moveBefore(block, block->getOperations().begin()); + } else if (isa(valueType)) { + arith::CmpIOp cmpIOp; + block->walk([&](arith::CmpIOp cmpOp) { + auto pred = cmpOp.getPredicate(); + if (pred == arith::CmpIPredicate::ugt || + pred == arith::CmpIPredicate::ult) { + isUnsigned = true; + } + if (pred == arith::CmpIPredicate::eq || + pred == arith::CmpIPredicate::ne) { + return WalkResult::advance(); + } else if (pred == arith::CmpIPredicate::sgt || + pred == arith::CmpIPredicate::slt || + pred == arith::CmpIPredicate::ugt || + pred == arith::CmpIPredicate::ult) { + if (cmpOp.getLhs() == block->getArgument(0) && + cmpOp.getRhs() == block->getArgument(2)) { + cmpIOp = cmpOp; + return WalkResult::interrupt(); + } + } + return WalkResult::advance(); + }); + cmpIOp->moveBefore(block, block->getOperations().begin()); + } + + TypedAttr valueAttr; + if (isa(valueType)) { + valueAttr = rewriter.getFloatAttr(valueType, T::getBaseReductionValue()); + } else if (isa(valueType)) { + if (isUnsigned) { + valueAttr = + rewriter.getIntegerAttr(valueType, T::getBaseReductionUIntValue()); + } else { + valueAttr = + rewriter.getIntegerAttr(valueType, T::getBaseReductionIntValue()); + } + } + + auto reduceWithIndexParams = getReduceWithIndexParams(op); + auto valuesAccBaseVal = + rewriter.create(loc, valueType, valueAttr); + int indicesInitValue = + (reduceWithIndexParams.has_value() && + (*reduceWithIndexParams).tieBreakType == TieBreakType::RIGHT) + ? -1 + : std::numeric_limits::max(); + + auto indexType = elemTypes[1]; + auto indicesAccBaseVal = rewriter.create( + loc, indexType, rewriter.getIntegerAttr(indexType, indicesInitValue)); + + auto valueResultType = dyn_cast(op.getType(0)); + const auto isScalarReduce = valueResultType == nullptr; + SmallVector reductionResultShape{ + isScalarReduce ? SmallVector{} + : SmallVector(valueResultType.getShape())}; + + SmallVector outputs{ + getInitTensor(rewriter, reductionResultShape, valuesAccBaseVal, loc), + getInitTensor(rewriter, reductionResultShape, indicesAccBaseVal, loc)}; + + auto linalgOp = rewriter.create( + loc, adaptor.getOperands(), outputs, + SmallVector{adaptor.getAxis()}, + [&](OpBuilder &b, Location loc, ValueRange inputs) { + assert(inputs.size() == 4); + + auto tritonReduceBlock = op.getBody(); + IRMapping mapping; + mapping.map(tritonReduceBlock->getArguments(), inputs); + + for (auto &op : tritonReduceBlock->without_terminator()) { + b.clone(op, mapping); + } + + auto tritonYield = tritonReduceBlock->getTerminator(); + auto results = + llvm::map_to_vector(tritonYield->getOperands(), [&](Value val) { + return mapping.lookup(val); + }); + b.create(loc, results); + }); + + // before we rewrite the argmax reduce op, we know it has return value + // so addReduceWithIndexAttrIfNeeded won't fail + // but ignoring it will lead to compiling failure + if (reduceWithIndexParams.has_value()) { + addReduceWithIndexAttr(*reduceWithIndexParams, rewriter, linalgOp); + } + + if (isScalarReduce) { + SmallVector reduceResults{ + rewriter.create( + loc, valueType, linalgOp.getResults()[0], ValueRange{}), + rewriter.create( + loc, indexType, linalgOp.getResults()[1], ValueRange{})}; + rewriter.replaceOp(op, reduceResults); + } else { + rewriter.replaceOp(op, linalgOp); + } + } +}; + +class ArgMinConverter : public ArgMinMaxBaseConverter { +public: + static LogicalResult matchComparisonResult(Value currValue, Value currIndex, + Value reduceValue, + Value reduceIndex, + mlir::Block::iterator &it, + Value &comparisonResult); + + static float getBaseReductionValue(); + + static int8_t getBaseReductionIntValue(); + static uint8_t getBaseReductionUIntValue(); + + ArgMinConverter(MLIRContext *context) : ArgMinMaxBaseConverter(context) {} +}; + +class ArgMaxConverter : public ArgMinMaxBaseConverter { +public: + static LogicalResult matchComparisonResult(Value currValue, Value currIndex, + Value reduceValue, + Value reduceIndex, + mlir::Block::iterator &it, + Value &comparisonResult); + + static float getBaseReductionValue(); + + static int8_t getBaseReductionIntValue(); + static uint8_t getBaseReductionUIntValue(); + + ArgMaxConverter(MLIRContext *context) : ArgMinMaxBaseConverter(context) {} +}; + +} // namespace TTOpConverters + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h new file mode 100755 index 00000000..458632bd --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h @@ -0,0 +1,296 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ANALYSIS_BLOCKPTRANALYSIS_H +#define TRITON_ANALYSIS_BLOCKPTRANALYSIS_H + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" + +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Value.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/SmallVector.h" + +#include +namespace mlir { + +class ConversionPatternRewriter; + +namespace triton { + +enum class MemAccVal { Undefined = 0, StrucMemAcc = 1, UnstrucMemAcc = 2 }; + +struct MemAccType { + + MemAccVal value; + + explicit constexpr MemAccType(MemAccVal v = MemAccVal::Undefined) + : value(v) {} + + constexpr operator MemAccVal() const { return value; } + explicit operator bool() = delete; + + constexpr bool isUndefined() const { return value == MemAccVal::Undefined; } + constexpr bool isStructured() const { + return value == MemAccVal::StrucMemAcc; + } + constexpr bool isUnstructured() const { + return value == MemAccVal::UnstrucMemAcc; + } + + void merge(MemAccType &other) { + this->value = (this->value > other.value) ? this->value : other.value; + } + + std::string_view toString() const { + static constexpr std::string_view names[] = {"Undefined", "StrucMemAcc", + "UnstrucMemAcc"}; + return names[static_cast(value)]; + } +}; + +class BlockData { +public: + SmallVector &getOffsetsRef(); + SmallVector &getSizesRef(); + SmallVector &getStridesRef(); + Value &getSourceRef(); + OpFoldResult &getScalarRef(); + Type &getResElemTyRef(); + MemAccType &getMemAccTypeRef(); + + SmallVector getOffsets() const; + SmallVector getSizes() const; + SmallVector getStrides() const; + Type getResElemTy() const; + OpFoldResult getOffset(int) const; + OpFoldResult getSize(int) const; + OpFoldResult getStride(int) const; + OpFoldResult getScalar() const; + Value getSource() const; + MemAccType getMemAccType() const; + + bool isScalar() const; + bool isEmpty() const; + bool hasSource() const; + bool hasResElemTy() const; + void removeSource(); + + int64_t getRank() const; + MemRefType getResultMemrefType(int64_t offset, + ArrayRef resultShape) const; + + void addBlock(BlockData &lBlock, BlockData &rBlock, Location loc, + ConversionPatternRewriter &rewriter); + void subBlock(BlockData &lBlock, BlockData &rBlock, Location loc, + ConversionPatternRewriter &rewriter); + void mulBlock(BlockData &lBlock, BlockData &rBlock, Location loc, + ConversionPatternRewriter &rewriter); + void divBlock(BlockData &lBlock, BlockData &rBlock, Location loc, + ConversionPatternRewriter &rewriter); + + memref::ReinterpretCastOp createCastOp(ArrayRef resultShape, + const Location &loc, + OpBuilder &builder) const; + + void setResElemTy(const Type &); + void setSource(const Value &); + void setScalar(const OpFoldResult &); + void setOffsets(const SmallVector &); + void setStrides(const SmallVector &); + void setSizes(const SmallVector &); + void setMemAccTy(const MemAccType &); + void setMemAccVal(const MemAccVal); + + void dump() const; + +private: + SmallVector offsets; + SmallVector sizes; + SmallVector strides; + Value source; + // `Scalar` is a shortcut used when the entire blockdata describes a single + // scalar value + OpFoldResult scalar; + Type resElemTy; + MemAccType memAccTy; + + // Accumulate offsets of each dimension in BlockData to get a total offset + // from source ptr, which is used in memref::ReinterpretCastOp + OpFoldResult inferBlockOffset(const Location &loc, OpBuilder &builder) const; +}; + +class BlockDataParser { +public: + static Value getScalarMemRef(Value ptr, Value memref, const Location &loc, + ConversionPatternRewriter &rewriter); + + static void parse(Value operand, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseAdd(arith::AddIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseSub(arith::SubIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseMul(arith::MulIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseDiv(arith::DivSIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseRem(arith::RemSIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void + parseUnrealizedCast(UnrealizedConversionCastOp op, BlockData &data, + const Location &loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void + parseMakeRange(triton::MakeRangeOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void + parseExpandDims(triton::ExpandDimsOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseBitcast(triton::BitcastOp op, BlockData &data, + const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseExtSI(arith::ExtSIOp op, BlockData &data, + const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void + parseBroadcast(triton::BroadcastOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseSplat(triton::SplatOp op, BlockData &data, + const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void + parseConstSplat(arith::ConstantOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + template + static std::enable_if_t || + std::is_same_v> + parseTensorPtr(T op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseAddPtr(triton::AddPtrOp op, BlockData &data, + const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void + parseExtractSlice(tensor::ExtractSliceOp op, BlockData &data, + const Location &loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void + parseReinterpretCast(memref::ReinterpretCastOp op, BlockData &data, + const Location &loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseReduce(triton::ReduceOp op, BlockData &data, + const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseFill(linalg::FillOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void parseSelect(arith::SelectOp op, BlockData &data, + const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + + static void rewriteAddPtr(triton::AddPtrOp op, + triton::AddPtrOp::Adaptor &adaptor, + ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &known); + + static void + rewriteMakeTensorPtrOp(triton::MakeTensorPtrOp op, Value base, + ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &known); + + static void rewriteAdvanceOp(triton::AdvanceOp op, + ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &known); + + template + static std::enable_if_t || + std::is_same_v> + rewriteTerminator(T op, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseSet &blockArgIdxSet, + ArrayRef iterArgIdxMap, + const llvm::SmallDenseMap &known); + + /// @param known is mainly designed for `rewriteLoop`, and is just non-const + /// in `rewriteLoop`, `rewriteAddPtr` and `rewriteAdvance` + static void rewriteLoopOp(LoopLikeOpInterface op, + ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &known); + + static void rewriteAddPtrToUnstrucMemAcc(triton::AddPtrOp op, + triton::AddPtrOp::Adaptor &adaptor, + ConversionPatternRewriter &rewriter, + BlockData &data); +}; + +template +void parseIndirectLoad(OpTy op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known); + +} // namespace triton + +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/CMakeLists.txt new file mode 100755 index 00000000..6c5f54af --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToLinalgIncubated) +add_public_tablegen_target(TritonToLinalgIncubatedConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/ConversionPatterns.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/ConversionPatterns.h new file mode 100755 index 00000000..e749f750 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/ConversionPatterns.h @@ -0,0 +1,133 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef CONVERSIONPATTERNS_H +#define CONVERSIONPATTERNS_H + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include +#include +#include + +using namespace mlir; +using namespace triton; + +//===----------------------------------------------------------------------===// +// Utilities +//===----------------------------------------------------------------------===// + +static Value getScalarValue(Value operand, Location loc, + ConversionPatternRewriter &rewriter) { + SmallVector ops; + + auto reconstructScalarValue = [&](Value src) { + for (auto op = ops.rbegin(); op != ops.rend(); ++op) { + src = TypeSwitch(*op) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return rewriter.create(loc, resType, src); + }) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return rewriter.create(loc, resType, src); + }) + .Default([](Operation *op) { + llvm_unreachable("unsupported op in generating "); + return nullptr; + }); + } + return src; + }; + + while (true) { + if (!dyn_cast(operand.getType())) { + return reconstructScalarValue(operand); + } else if (auto op = operand.getDefiningOp()) { + if (auto attr = dyn_cast(op.getValue())) { + if (!attr.isSplat()) { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load " + "produced by unsupported instruction"; + return nullptr; + } + auto elemValue = attr.getSplatValue(); + auto constOp = arith::ConstantOp::materialize( + rewriter, elemValue, attr.getElementType(), op.getLoc()); + return reconstructScalarValue(constOp.getResult()); + } + } else if (auto op = operand.getDefiningOp()) { + operand = op.getSrc(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load produced " + "by unsupported instruction"; + return nullptr; + } + } + return nullptr; +} + +static SmallVector getNParallelLoopsAttrs(unsigned n) { + return SmallVector(n, utils::IteratorType::parallel); +} + +// for IntLike and FloatLike types +static std::optional getBitWidth(Type a) { + if (auto type = dyn_cast(a)) { + auto elementType = type.getElementType(); + if (elementType.isIntOrFloat()) { + return type.getElementType().getIntOrFloatBitWidth(); + } + return std::nullopt; + } + + if (a.isIntOrFloat()) { + return a.getIntOrFloatBitWidth(); + } + return std::nullopt; +} +#endif // CONVERSIONPATTERNS_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/DescriptorConverter.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/DescriptorConverter.h new file mode 100755 index 00000000..598f2b91 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/DescriptorConverter.h @@ -0,0 +1,74 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_DESCRIPTORCONVERTER_H +#define TRITON_ADAPTER_DESCRIPTORCONVERTER_H + +#include "incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" + +namespace DescriptorConverter { +using namespace mlir; +using namespace triton; + +struct Descriptor { + Value base; + SmallVector shape; + SmallVector strides; +}; + +bool hasATensorDescriptorType(mlir::TypeRange types); + +class DescriptorLoadConverter + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::DescriptorLoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class DescriptorStoreConverter + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::DescriptorStoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +} // end of namespace DescriptorConverter + +#endif // TRITON_ADAPTER_DESCRIPTORCONVERTER_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/FunctionConverter.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/FunctionConverter.h new file mode 100755 index 00000000..48d57e27 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/FunctionConverter.h @@ -0,0 +1,60 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_FUNCTIONCONVERTER_H +#define TRITON_ADAPTER_FUNCTIONCONVERTER_H + +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace FunctionConverter { +using namespace mlir; +using namespace triton; + +class GetProgramIDConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + static uint32_t constexpr LAUNCH_GRID_RANK = + getMaxEnumValForProgramIDDim() + 1; + +public: + LogicalResult + matchAndRewrite(triton::GetProgramIdOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class GetNumProgramsConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + static uint32_t constexpr LAUNCH_GRID_RANK = + getMaxEnumValForProgramIDDim() + 1; + +public: + LogicalResult + matchAndRewrite(triton::GetNumProgramsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; +} // namespace FunctionConverter +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/HoistBroadcast.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/HoistBroadcast.h new file mode 100755 index 00000000..d37b85e7 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/HoistBroadcast.h @@ -0,0 +1,82 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_TRITONTOLINALG_HOISTBROADCAST_H +#define TRITON_ADAPTER_TRITONTOLINALG_HOISTBROADCAST_H + +#include "incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "triton-to-linalg" + +namespace HoistBroadcast { +using namespace mlir; +using namespace triton; + +class BroadcastConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::BroadcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class BroadcastHoister { +public: + BroadcastHoister(triton::BroadcastOp op); + LogicalResult parse(Value operand, const Location &loc, + ConversionPatternRewriter &rewriter); + LogicalResult parseAddptr(triton::AddPtrOp op, const Location &loc, + ConversionPatternRewriter &rewriter); + LogicalResult parseBroadcast(triton::BroadcastOp op, const Location &loc, + ConversionPatternRewriter &rewriter); + LogicalResult parseSplat(triton::SplatOp op, const Location &loc, + ConversionPatternRewriter &rewriter); + + LogicalResult findSrc(Value operand); + LogicalResult replaceBroadcastOp(triton::BroadcastOp op, + ConversionPatternRewriter &rewriter); + bool canBroadcast(); + +private: + Value source; + triton::BroadcastOp opToHoist; + SmallVector tensorSizes; + llvm::SmallDenseMap broadcastMap; +}; +} // namespace HoistBroadcast + +#endif // TRITON_ADAPTER_TRITONTOLINALG_HOISTBROADCAST_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/LoadStoreConverter.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/LoadStoreConverter.h new file mode 100755 index 00000000..08b32d07 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/LoadStoreConverter.h @@ -0,0 +1,278 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_LOADSTORECONVERTER_H +#define TRITON_ADAPTER_LOADSTORECONVERTER_H + +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/AffineMap.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Value.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Arith/Utils/Utils.h" + +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace LoadStoreConverter { + +using namespace mlir; +using namespace triton; + +class AddPtrConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::AddPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class LoadConverter : public OpConversionPattern { +private: + LogicalResult toTensorAndReplace(triton::LoadOp &op, + RankedTensorType &tensorType, Value localMem, + bool mayImplicitTransposeWithLastAxis, + const Location &loc, + ConversionPatternRewriter &rewriter) const; + + LogicalResult checkModifiedByAddPtrConverter(triton::LoadOp &op) const; + + LogicalResult + continueModifyFromAddPtrConverter(triton::LoadOp &op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const; + + void + fillTensorWithOtherForMaskScenario(Value other, Value localMem, + ArrayRef maskDim, + ConversionPatternRewriter &rewriter) const; + +public: + explicit LoadConverter(MLIRContext *context); + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +// tempate class's impl must in header file +template +class LoadStoreCanonicalizer : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(OpTy op, + PatternRewriter &rewriter) const override { + Value ptrVal = op.getPtr(); + Type ptrTy = ptrVal.getType(); + auto ptrDefOp = ptrVal.getDefiningOp(); + + bool shouldAddZeros = false; + if (!isa(ptrVal)) + shouldAddZeros = !isTensorPointerType(ptrTy) && + !isa_and_nonnull(ptrDefOp); + else if (auto ptrType = dyn_cast(ptrTy)) + shouldAddZeros = ptrType.getPointeeType().isIntOrIndexOrFloat(); + + if (shouldAddZeros) { + if (isa_and_nonnull(ptrDefOp)) { + auto castOp = cast(ptrDefOp); + auto castSrc = castOp.getSrc(); + if (!isa(castSrc)) { + auto castSrcDefOp = castSrc.getDefiningOp(); + if (isa(castSrcDefOp)) { + return rewriter.notifyMatchFailure( + op, "BitcastCanonicalizer handles addptr->bitcast->load!"); + } + } + } + + Type zeroTy = getI32SameShape(ptrTy); + Value zeroVal = + createScalarOrSplatConstant(rewriter, op.getLoc(), zeroTy, 0); + Value addptrVal = rewriter.create(op.getLoc(), ptrTy, + ptrVal, zeroVal); + rewriter.modifyOpInPlace( + op, [&]() { op->replaceUsesOfWith(ptrVal, addptrVal); }); + return success(); + } + return failure(); + } +}; + +class ScalarStoreCanonicalizer : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(triton::StoreOp op, + PatternRewriter &rewriter) const override; +}; + +class StoreConverter : public OpConversionPattern { +public: + explicit StoreConverter(MLIRContext *context); + + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class ScalarAtomicRMWCanonicalizer + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(triton::AtomicRMWOp op, + PatternRewriter &rewriter) const override; +}; + +class ScalarAtomicCASCanonicalizer + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(triton::AtomicCASOp op, + PatternRewriter &rewriter) const override; +}; + +class AtomicCASConverter : public OpConversionPattern { +public: + explicit AtomicCASConverter(MLIRContext *context) + : OpConversionPattern(context) {} + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::AtomicCASOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class AtomicRMWConverter : public OpConversionPattern { +private: + Value createAtomicBinaryOps(OpBuilder &builder, Location loc, + triton::AtomicRMWOp op, Type elementType, + Value lhs, Value rhs) const { + auto rmwOp = op.getAtomicRmwOp(); + + // it has been confirmed in AtomicRMWConverter::matchAndRewrite + // that the ptr of op is of MemRefType + Value binaryOp; + if (rmwOp == triton::RMWOp::FADD) { + binaryOp = builder.create(loc, lhs, rhs); + } else if (rmwOp == triton::RMWOp::ADD) { + binaryOp = builder.create(loc, lhs, rhs); + } else if (rmwOp == triton::RMWOp::XOR) { + binaryOp = builder.create(loc, lhs, rhs); + } else if (rmwOp == triton::RMWOp::OR) { + binaryOp = builder.create(loc, lhs, rhs); + } else if (rmwOp == triton::RMWOp::AND) { + binaryOp = builder.create(loc, lhs, rhs); + } else if (rmwOp == triton::RMWOp::MAX) { + // Max/Min only support f32/i32 for now + // Other type is not supported because of semantic.py + if (isa(elementType)) { + binaryOp = builder.create(loc, lhs, rhs); + } else { + binaryOp = builder.create(loc, lhs, rhs); + } + } else if (rmwOp == triton::RMWOp::MIN) { + if (isa(elementType)) { + binaryOp = builder.create(loc, lhs, rhs); + } else { + binaryOp = builder.create(loc, lhs, rhs); + } + } else if (rmwOp == triton::RMWOp::XCHG) { + binaryOp = rhs; + } else if (rmwOp == triton::RMWOp::UMAX) { + binaryOp = builder.create(loc, lhs, rhs); + } else if (rmwOp == triton::RMWOp::UMIN) { + binaryOp = builder.create(loc, lhs, rhs); + } else { + op.emitOpError("unsupported atomic RMW operation: "); + llvm_unreachable( + "Not implemented. Support fadd, add, max, min for now !"); + } + return binaryOp; + } + + // used when handling scalar + // to verify whether we need to handle this scalar + bool isConstantMaskTrue(Value mask) const { + if (auto denseAttr = + mask.getDefiningOp()->getAttrOfType("value")) { + auto eleType = denseAttr.getType().getElementType(); + if (isa(eleType) && + cast(eleType).getWidth() == 1) { + auto values = denseAttr.getValues(); + return values[0]; + } + } + return false; + } + + DenseSet softwareAtomicKinds = { + triton::RMWOp::AND, triton::RMWOp::OR, triton::RMWOp::XOR}; + +public: + explicit AtomicRMWConverter(MLIRContext *context); + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::AtomicRMWOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class AtomicRMWNewConverter : public OpConversionPattern { +private: + // used when handling scalar + // to verify whether we need to handle this scalar + bool isConstantMaskTrue(Value mask) const { + if (auto denseAttr = + mask.getDefiningOp()->getAttrOfType("value")) { + auto eleType = denseAttr.getType().getElementType(); + if (isa(eleType) && + cast(eleType).getWidth() == 1) { + auto values = denseAttr.getValues(); + return values[0]; + } + } + return false; + } + +public: + explicit AtomicRMWNewConverter(MLIRContext *context); + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::AtomicRMWOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class AtomicMaxMinCanonicalizer : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(triton::AtomicRMWOp op, + PatternRewriter &rewriter) const override; +}; + +} // namespace LoadStoreConverter +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h new file mode 100755 index 00000000..9baca3bf --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h @@ -0,0 +1,151 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ANALYSIS_MASKANALYSIS_H +#define TRITON_ANALYSIS_MASKANALYSIS_H + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include +#include + +namespace mlir { + +// this class helps build Operations +class OpBuilder; + +namespace triton { +// use to decode the pattern in a mask used for load and store +namespace Incubated { + +class MaskState { +public: + OpFoldResult start; + OpFoldResult end; + SmallVector dims; + SmallVector offsets; + OpFoldResult scalar; + + int64_t getRank() const { + assert(dims.size() == offsets.size() && "dims and offsets rank mismatch!"); + return dims.size(); + } + + bool isEmpty() const { return getRank() == 0 && !scalar && !start && !end; } + + bool isMask() const { + return !start && !end && !scalar && dims.size() != 0 && offsets.size() != 0; + } + + // parse value recursively + LogicalResult parse(Value operand, const Location &loc, OpBuilder &builder); + + tensor::ExtractSliceOp getExtractSlice(Value source, const Location &loc, + OpBuilder &builder) const; + + tensor::InsertSliceOp getInsertSlice(Value source, Value dest, + const Location &loc, + OpBuilder &builder) const; + + memref::SubViewOp getSubview(Value source, const Location &loc, + OpBuilder &builder) const; + + void eraseInsertedOps(Operation *rawOp, PatternRewriter &rewriter); + +private: + LogicalResult addStateScalar(const MaskState &state, + const OpFoldResult scalar, const Location &loc, + OpBuilder &builder); + + LogicalResult addStates(const MaskState &lhsState, const MaskState &rhsState, + const Location &loc, OpBuilder &builder); + + LogicalResult divStateScalar(const MaskState &state, + const OpFoldResult scalar, const Location &loc, + OpBuilder &builder); + + LogicalResult divStates(const MaskState &lhsState, const MaskState &rhsState, + const Location &loc, OpBuilder &builder); + + // Helper function to handle operator `and` both mask state + LogicalResult minStates(const MaskState &lhsState, const MaskState &rhsState, + const Location &loc, OpBuilder &builder); + + // Helper functions to parse values to populate MaskState + + LogicalResult parseConstant(arith::ConstantOp constOp, const Location &loc, + OpBuilder &builder); + + // Operand is an integer scalar + LogicalResult parseIntScalar(Value scalar, const Location &loc, + OpBuilder &builder); + + // TODO + LogicalResult parseAdd(arith::AddIOp addOp, const Location &loc, + OpBuilder &builder); + + // operand is the result of divsi + LogicalResult parseDiv(arith::DivSIOp divOp, const Location &loc, + OpBuilder &builder); + + // Operand is the result of andi + LogicalResult parseAnd(arith::AndIOp andOp, const Location &loc, + OpBuilder &builder); + + // Operand is the result of cmpi, necessary method to fuse scalar, start and + // end into dims and offset + LogicalResult parseCmp(arith::CmpIOp cmpOp, const Location &loc, + OpBuilder &builder); + + // Operand is the result of select + LogicalResult parseSel(arith::SelectOp selOp, const Location &loc, + OpBuilder &builder); + + // Operand is the result of make_range + LogicalResult parseMakeRange(triton::MakeRangeOp rangeOp, const Location &loc, + OpBuilder &builder); + + // Operand is the result of broadcast + LogicalResult parseBroadcast(triton::BroadcastOp broadcastOp, + const Location &loc, OpBuilder &builder); + + // Operand is the result of splat + LogicalResult parseSplat(triton::SplatOp splatOp, const Location &loc, + OpBuilder &builder); + + // Operand is the result of expand_dims + LogicalResult parseExpandDims(triton::ExpandDimsOp expandDimsOp, + const Location &loc, OpBuilder &builder); +}; + +std::optional runMaskAnalysis(Operation *op, + OpBuilder &builder); +} // namespace Incubated + +} // namespace triton + +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/Passes.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/Passes.h new file mode 100755 index 00000000..f081f8bf --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/Passes.h @@ -0,0 +1,39 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_TRITON_TO_LINALG_CONVERSION_PASSES_H +#define TRITON_ADAPTER_TRITON_TO_LINALG_CONVERSION_PASSES_H + +#include "incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h" + +namespace mlir { +namespace triton { +namespace Incubated { + +#define GEN_PASS_REGISTRATION +#include "incubated/Conversion/TritonToLinalgIncubated/Passes.h.inc" + +} // namespace Incubated +} // namespace triton +} // namespace mlir + +#endif // TRITON_ADAPTER_TRITON_TO_LINALG_CONVERSION_PASSES_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/Passes.td b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/Passes.td new file mode 100755 index 00000000..ff3691bd --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/Passes.td @@ -0,0 +1,28 @@ +#ifndef TRITON_TO_LINALG_CONVERSION_PASSES +#define TRITON_TO_LINALG_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToLinalgIncubated : Pass<"triton-to-linalg-incubated", "mlir::ModuleOp"> { + let summary = "Convert Triton to Linalg dialect (Incubated)"; + let constructor = "createTritonToLinalgIncubatedPass()"; + let options = [ + Option<"globalKernel", "global-kernel", + "bool", /*default*/"true", + "generate a global kernel">, + Option<"namedOps", "named-ops", + "bool", /*default*/"false", + "use linalg named ops instead of linalg.generic">, + Option<"enableNd2nzOnVector", "enable-nd2nz-on-vector", + "bool", /*default*/"false", + "enable nd2nz on vector">, + Option<"enableSelectAnalysis", "enable-select-analysis", + "bool", /*default*/"true", + "enable select analysis">, + Option<"compileOn91095", "compile-on-910-95", + "bool", /*default*/"false", + "compile on 910_95"> + ]; +} + +#endif // TRITON_TO_LINALG_CONVERSION_PASSES diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/TritonOpConverter.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/TritonOpConverter.h new file mode 100755 index 00000000..9f900cb9 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/TritonOpConverter.h @@ -0,0 +1,696 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_TRITONOPCONVERTER_H +#define TRITON_ADAPTER_TRITONOPCONVERTER_H + +#include "incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h" +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/Interfaces/FunctionInterfaces.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "triton-to-linalg" + +namespace TTOpConverters { +using namespace mlir; +using namespace triton; + +static constexpr unsigned kFuncNameCap = 128; + +/* +Convert `tt.precise_div` operation to `arith.divf` operation. +tensor_x / tensor_y + +```ttir + %11 = tt.precise_divf %7, %10 : tensor<100xf32> +``` + +converts to: + +```mlir + %11 = arith.divf %7, %10 : tensor<100xf32> +``` +*/ +struct PreciseDivConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::PreciseDivFOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +/* +Convert `tt.fp_to_fp` operation with RTNE (default) rounding mode to +`arith.truncf` or `arith.extf` operation. +For fp8 conversions with default RTNE rounding: +- downcast: tt.fp_to_fp -> arith.truncf +- upcast: tt.fp_to_fp -> arith.extf +Note: Non-RTNE rounding modes (e.g., RTZ) are handled by TritonToHFusion pass. +*/ +struct FpToFpCanonicalizer : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(triton::FpToFpOp op, + PatternRewriter &rewriter) const override; +}; + +class SelectCanonicalizer : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(arith::SelectOp op, + PatternRewriter &rewriter) const override; +}; + +/* + * Move tt.bitcast to a previous location if tt.bitcast is not directly applied + * on function arguments + */ +class BitcastCanonicalizer : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(triton::BitcastOp bitcastOp, + PatternRewriter &rewriter) const override; +}; + +template +class ScalarMathCanonicalizer : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(MathOp op, + PatternRewriter &rewriter) const override { + if (op->getNumResults() != 1) { + return rewriter.notifyMatchFailure( + op, "ScalarMathCanonicalizer expects single scalar output."); + } + if (!op->getResult(0).getType().isIntOrIndexOrFloat()) { + return rewriter.notifyMatchFailure( + op, "ScalarMathCanonicalizer handles scalar load scene."); + } + if (auto linalgOp = op->template getParentOfType()) { + return rewriter.notifyMatchFailure( + op, "ScalarMathCanonicalizer handles op not within tt.reduce."); + } + if (auto linalgOp = op->template getParentOfType()) { + return rewriter.notifyMatchFailure( + op, "ScalarMathCanonicalizer handles op not within tt.scan."); + } + auto loc = op.getLoc(); + llvm::SmallVector inputs; + for (auto input : op->getOperands()) { + auto blkTy = RankedTensorType::get({(int64_t)1}, input.getType()); + auto inputSplat = rewriter.create(loc, blkTy, input); + inputs.push_back(inputSplat.getResult()); + } + auto blkOp = rewriter.create(loc, inputs); + Value offset = + rewriter.create(loc, rewriter.getIndexAttr(0)); + auto extractOp = + rewriter.create(loc, blkOp.getResult(), offset); + rewriter.replaceOp(op, extractOp); + return success(); + } +}; + +/* + * Rewrite tt.make_tensor_ptr with non-contiguous order to + * tt.make_tensor_ptr + tt.load + tt.trans. + */ +class MakeTensorPtrCanonicalizer + : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(triton::MakeTensorPtrOp op, + PatternRewriter &rewriter) const override; +}; + +class ReduceSingleCanonicalizer : public OpRewritePattern { +public: + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(triton::ReduceOp reduceOp, + PatternRewriter &rewriter) const override; +}; + +class DenseConstantConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(arith::ConstantOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class MakeRangeConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::MakeRangeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class SplatConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::SplatOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class ReshapeConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ReshapeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class ExpandDimsConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ExpandDimsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class ClampFConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::ClampFOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class BroadcastConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::BroadcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +template +class ReductionOpBaseConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(OpTy op, typename OpTy::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const final { + auto sourceType = + cast(adaptor.getOperands().front().getType()); + assert(sourceType.hasRank() && "Expected input is ranked"); + + int64_t axis = op.getAxis(); + assert(axis >= 0 && axis < sourceType.getRank() && + "Expected reduction axis is within operand's rank"); + + auto reductionOps = this->getRedOps(op); + if (reductionOps.size() == 1) { + return this->convertToTargetOp(op, adaptor, rewriter); + } + return this->convertToTargetOpExtended(op, adaptor, rewriter); + } + +protected: + llvm::SmallVector getRedOps(OpTy redOp) const { + auto redBody = redOp.getBody(); + return llvm::map_to_vector(redBody->without_terminator(), + [](Operation &op) { return &op; }); + } + + arith::ConstantOp getRedBaseConstOp(ConversionPatternRewriter &rewriter, + Operation *redOp, + Type constantType) const { + const int64_t bitWidth = constantType.getIntOrFloatBitWidth(); + + auto attr = + llvm::TypeSwitch(redOp) + .Case([&](arith::AddFOp) { + return rewriter.getFloatAttr(constantType, 0.f); + }) + .Case([&](arith::AddIOp) { + return rewriter.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::MulFOp) { + return rewriter.getFloatAttr(constantType, 1.f); + }) + .template Case([&](auto) { + return rewriter.getFloatAttr( + constantType, -std::numeric_limits::infinity()); + }) + .template Case([&](auto) { + return rewriter.getFloatAttr( + constantType, std::numeric_limits::infinity()); + }) + .Case([&](arith::MinSIOp) { + return rewriter.getIntegerAttr(constantType, + llvm::maxIntN(bitWidth)); + }) + .Case([&](arith::MinUIOp) { + return rewriter.getIntegerAttr(constantType, + llvm::maxUIntN(bitWidth)); + }) + .Case([&](arith::MaxSIOp) { + return rewriter.getIntegerAttr(constantType, + llvm::minIntN(bitWidth)); + }) + .Case([&](arith::MaxUIOp) { + return rewriter.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::OrIOp) { + return rewriter.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::AndIOp) { + return rewriter.getIntegerAttr(constantType, 1); + }) + .Case([&](arith::XOrIOp) { + return rewriter.getIntegerAttr(constantType, 0); + }) + .Default([](Operation *op) { + op->dump(); + llvm_unreachable("Reduction op not supported yet"); + return nullptr; + }); + + return rewriter.create(redOp->getLoc(), constantType, + attr); + } + + bool requiresF32Conversion(const Type elemType, Operation *redOp) const { + unsigned width = + cast(Float32Type::get(elemType.getContext())).getWidth(); + return isa(elemType) && + elemType.getIntOrFloatBitWidth() < width && + // Float32Type::get(elemType.getContext()).getWidth() && + (isa(redOp) || isa(redOp)); + } + + Value getRedElement(Value lhs, Value rhs, const Location loc, + Operation *redOp, OpBuilder &b, + const bool convertLhsToF32Precision) const { + return llvm::TypeSwitch(redOp) + .template Case([&](auto redOp) { + if (convertLhsToF32Precision) { + lhs = b.create(loc, Float32Type::get(b.getContext()), + lhs); + } + return b.create(loc, lhs, rhs); + }) + .template Case( + [&](auto redOp) { + return b.create(loc, lhs, rhs); + }) + .Default([](Operation *op) { + op->dump(); + llvm_unreachable("Reduction op not yet supported"); + return nullptr; + }); + } + + virtual bool isReductionOpSupported(Operation *redOp) const = 0; + + virtual LogicalResult + convertToTargetOp(OpTy op, typename OpTy::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const = 0; + + virtual LogicalResult + convertToTargetOpExtended(OpTy op, typename OpTy::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const = 0; +}; + +class ReduceConverter : public ReductionOpBaseConverter { +public: + explicit ReduceConverter(MLIRContext *context) + : ReductionOpBaseConverter(context) {} + + using ReductionOpBaseConverter::ReductionOpBaseConverter; + +protected: + bool isReductionOpSupported(Operation *redOp) const override; + + LogicalResult + convertToTargetOp(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + + LogicalResult + convertToTargetOpExtended(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class ScanConverter : public ReductionOpBaseConverter { +public: + explicit ScanConverter(MLIRContext *context) + : ReductionOpBaseConverter(context) {} + + using ReductionOpBaseConverter::ReductionOpBaseConverter; + +protected: + bool isReductionOpSupported(Operation *redOp) const override; + + LogicalResult + convertToTargetOp(triton::ScanOp op, typename triton::ScanOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + + LogicalResult + convertToTargetOpExtended(triton::ScanOp op, + typename triton::ScanOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class ExternElementwiseClOpConverter + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::ExternElementwiseOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class UnrealizedCastConverter + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(UnrealizedConversionCastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class JoinConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::JoinOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class SplitConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::SplitOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class CatConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::CatOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class GatherConverter : public OpConversionPattern { +private: + static constexpr llvm::StringRef gatherFuncNameBase = "triton_gather"; + +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::GatherOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class YieldConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(scf::YieldOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +template || + std::is_same_v>> +class LoopConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(LoopOpTy op, + typename OpConversionPattern::OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + llvm::SmallDenseMap known; + + op->removeAttr("UnhandledLoopOp"); + BlockDataParser::rewriteLoopOp(op, rewriter, known); + return success(); + } +}; + +class AdvanceConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::AdvanceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class MakeTensorPtrConverter + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + explicit MakeTensorPtrConverter(MLIRContext *context) + : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(triton::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class TransposeConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::TransOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class BitcastConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::BitcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class TritonMulhiuiConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::MulhiUIOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class TritonPreciseSqrtConverter + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::PreciseSqrtOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class DeviceAssertConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +private: + static constexpr llvm::StringRef printFuncNameBase = "triton_assert"; + static constexpr llvm::StringRef msgAttrName = "msg"; + +public: + LogicalResult + matchAndRewrite(triton::AssertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class DevicePrintConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +private: + static constexpr llvm::StringRef printFuncNameBase = "triton_print"; + static constexpr llvm::StringRef prefixAttrName = "prefix"; + static constexpr llvm::StringRef hexAttrName = "hex"; + +public: + LogicalResult + matchAndRewrite(triton::PrintOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +struct MatmulConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::DotOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +struct FlipOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ascend::FlipOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + + static constexpr StringRef baseFuncName = "triton_flip"; +}; + +struct SortOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ascend::SortOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +struct DotScaledConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::DotScaledOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class PtrToIntConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::PtrToIntOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +class EmbeddingGatherConverter + : public OpConversionPattern { +public: + using OpConversionPattern< + triton::ascend::EmbeddingGatherOp>::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::ascend::EmbeddingGatherOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + +private: + static constexpr llvm::StringRef funcNameBase = "triton_embedding_gather"; +}; + +class IndexPutConverter + : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::ascend::IndexPutOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + +private: + static constexpr llvm::StringRef funcNameBase = "triton_index_put"; +}; + +class GatherOutToUbConverter + : public OpConversionPattern { +public: + using OpConversionPattern< + triton::ascend::GatherOutToUbOp>::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::ascend::GatherOutToUbOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + +private: + static constexpr llvm::StringRef funcNameBase = "triton__gather_out_to_ub"; +}; + +class ScatterUbToOutConverter + : public OpConversionPattern { +public: + using OpConversionPattern< + triton::ascend::ScatterUbToOutOp>::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::ascend::ScatterUbToOutOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + +private: + static constexpr llvm::StringRef funcNameBase = "triton_scatter_ub_to_out"; +}; + +class IndirectLoadConverter + : public OpConversionPattern { +public: + using OpConversionPattern< + triton::ascend::IndirectLoadOp>::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::ascend::IndirectLoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + +private: + static constexpr llvm::StringRef funcNameBase = "triton_indirect_load"; +}; + +class IndirectStoreConverter + : public OpConversionPattern { +public: + using OpConversionPattern< + triton::ascend::IndirectStoreOp>::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::ascend::IndirectStoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; + +private: + static constexpr llvm::StringRef funcNameBase = "triton_indirect_store"; +}; + +class IndexSelectSimdConverter + : public OpConversionPattern { +public: + explicit IndexSelectSimdConverter(MLIRContext *context); + using OpConversionPattern< + triton::ascend::IndexSelectSimdOp>::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ascend::IndexSelectSimdOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override; +}; + +} // end of namespace TTOpConverters + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h new file mode 100755 index 00000000..7fb69886 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h @@ -0,0 +1,126 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_CONVERSION_TRITONTOLINALG_H +#define TRITON_ADAPTER_CONVERSION_TRITONTOLINALG_H + +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#define GEN_PASS_CLASSES +#include "incubated/Conversion/TritonToLinalgIncubated/Passes.h.inc" + +extern int nd2nzFlag; +extern bool compileOn91095Flag; +extern bool existDotFlag; + +namespace mlir { +namespace triton { +namespace Incubated { + +std::unique_ptr> createTritonToLinalgIncubatedPass(); + +std::unique_ptr> +createTritonToLinalgIncubatedPass(bool, bool, bool, bool, bool); + +} // namespace Incubated +} // namespace triton +} // namespace mlir +enum TensorKind { NONE = -1, INPUT = 0, OUTPUT = 1, INPUT_OUTPUT = 2 }; + +using namespace mlir; +using namespace triton; +const std::string globalKernelAttr = "global_kernel"; +const std::string kernelMixModeName = "mix_mode"; +const std::string kernelParallelModeName = "parallel_mode"; +const unsigned INT_BIT_WIDTH = 32; +const unsigned SET_INIT_SIZE = 16; + +class TritonTypeConverter : public mlir::TypeConverter { +public: + explicit TritonTypeConverter(); +}; + +class TritonToLinalgIncubatedPass + : public TritonToLinalgIncubatedBase { + + static auto constexpr LAUNCH_GRID_RANK = getMaxEnumValForProgramIDDim() + 1; + static unsigned int constexpr TRITON_PROGRAM_INFO_ARG_COUNT = + LAUNCH_GRID_RANK * 2; + +private: + // grid构造 num_programs 3维, program_id 3维 + // remember 'xxxOp' is usually a Pointer, so that we can change target memory + // without giving a reference argument + void addProgramInfo(triton::FuncOp func, bool globalKernel); + + template + void addTensorKindToArguments(OpTy op, triton::FuncOp func, + TensorKind tensorKind); + + template + void walkAndMarkTensorKind(triton::FuncOp func); + + void annotateTensorKindForModule(ModuleOp moduleOp); + + void convertTTFunc(triton::FuncOp func, const bool existDot, + const bool existSIMTOp); + + LogicalResult convertMultipleBlockControlFlow(Operation *funcOp, + OpBuilder &builder); + // 处理嵌套的if/else + scf::IfOp transformNestedIfElse(Operation &nestedBranch, OpBuilder &builder); + + void addDynamicLegal(ConversionTarget &target, + TritonTypeConverter &tritonTypeConverter); + + void + populateTritonToLinalgCanonicalizationPatterns(RewritePatternSet &patterns); + + void populateTritonToLinalgConversionPatterns(TypeConverter &typeConverter, + RewritePatternSet &patterns, + unsigned int launchGridRank); + + LogicalResult processDescriptorOperations(ModuleOp moduleOp); + LogicalResult processPtrBroadcastOperations(ModuleOp moduleOp); + +public: + TritonToLinalgIncubatedPass() = default; + + TritonToLinalgIncubatedPass(bool globalKernel, bool namedOps, + bool enableNd2nzOnVector, + bool enableSelectAnalysis, bool compileOn91095) { + this->globalKernel = globalKernel; + this->namedOps = namedOps; + this->enableNd2nzOnVector = enableNd2nzOnVector; + this->enableSelectAnalysis = enableSelectAnalysis; + this->compileOn91095 = compileOn91095; + }; + void getDependentDialects(DialectRegistry ®istry) const override; + + void runOnOperation() override; +}; + +#endif // TRITON_ADAPTER_CONVERSION_TRITONTOLINALG_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/UseAnalysis.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/UseAnalysis.h new file mode 100755 index 00000000..4fa42042 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/UseAnalysis.h @@ -0,0 +1,146 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ANALYSIS_USEANALYSIS_H +#define TRITON_ANALYSIS_USEANALYSIS_H + +#include "mlir/Analysis/DataFlow/SparseAnalysis.h" + +#include "mlir/Transforms/DialectConversion.h" +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { +namespace Incubated { +enum class UseType { + Undefined, // Initial state + DataUse, // value used for tensor computation only + MetaUse, // value used for metadata only + MixUse // value used for both tensor computation and metadata +}; + +struct UseInfo : public dataflow::AbstractSparseLattice { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(UseInfo) + using AbstractSparseLattice::AbstractSparseLattice; + + // Lattice state transfer function + ChangeResult meetUseType(const UseType &other) { + if (other == UseType::Undefined) { + return ChangeResult::NoChange; + } + + switch (type) { + case UseType::Undefined: + type = other; + return ChangeResult::Change; + case UseType::DataUse: + case UseType::MetaUse: + if (type == other) { + return ChangeResult::NoChange; + } else { + type = UseType::MixUse; + return ChangeResult::Change; + } + case UseType::MixUse: + return ChangeResult::NoChange; + default: + llvm_unreachable("bad type"); + } + } + + ChangeResult meet(const AbstractSparseLattice &other) override { + auto rhs = reinterpret_cast(&other); + return meetUseType(rhs->type); + } + + void print(raw_ostream &os) const override { + switch (type) { + case UseType::DataUse: + os << "DataUse"; + break; + case UseType::MetaUse: + os << "MetaUse"; + break; + case UseType::MixUse: + os << "MixUse"; + break; + default: + os << "Undefined"; + } + } + + UseType type = UseType::Undefined; +}; + +class UseAnalysis : public dataflow::SparseBackwardDataFlowAnalysis { +public: + using SparseBackwardDataFlowAnalysis::SparseBackwardDataFlowAnalysis; + +#if LLVM_VERSION_MAJOR >= 20 + LogicalResult visitOperation(Operation *op, ArrayRef operands, + ArrayRef results) override; +#else + void visitOperation(Operation *op, ArrayRef operands, + ArrayRef results) override; +#endif + + void visitBranchOperand(OpOperand &operand) override { return; } + + void visitCallOperand(OpOperand &operand) override { return; } + + void setToExitState(UseInfo *lattice) override { + lattice->type = UseType::Undefined; + } + +private: + void propagateUse(UseInfo *lattice, const UseType &type) { + auto changed = lattice->meetUseType(type); + propagateIfChanged(lattice, changed); + } + + void propagateResults(UseInfo *lattice, ArrayRef results) { + auto changed = ChangeResult::NoChange; + for (auto result : results) { + changed |= lattice->meet(*result); + } + propagateIfChanged(lattice, changed); + } +}; + +class MetaUseEraser : public RewritePattern { +public: + MetaUseEraser(MLIRContext *context); + + LogicalResult matchAndRewrite(Operation *op, + PatternRewriter &rewriter) const final; +}; + +LogicalResult runUseAnalysis(triton::FuncOp &funcOp); + +} // namespace Incubated + +} // namespace triton + +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONTOAFFINE_TRITONUSEANALYSIS_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/CMakeLists.txt new file mode 100755 index 00000000..f19e8ed7 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToStructuredIncubated) +add_public_tablegen_target(TritonToStructuredIncubatedConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.h new file mode 100755 index 00000000..c3b4c02b --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.h @@ -0,0 +1,162 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ +#ifndef TRITON_ADAPTER_CANNONICALIZERCONVERTER_H +#define TRITON_ADAPTER_CANNONICALIZERCONVERTER_H + +#include "mlir/Dialect/Arith/Utils/Utils.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/AffineMap.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Value.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace CannonicalizerConverter { + +using namespace mlir; +using namespace triton; + +class CmpConverter : public OpRewritePattern { +public: + explicit CmpConverter(MLIRContext *context) + : OpRewritePattern(context) {} + + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(arith::CmpIOp cmpOp, + PatternRewriter &rewriter) const override; +}; + +class PromotePointerIterArgsPattern : public OpRewritePattern { +public: + explicit PromotePointerIterArgsPattern(MLIRContext *context) + : OpRewritePattern(context) {} + + LogicalResult matchAndRewrite(scf::ForOp forOp, + PatternRewriter &rewriter) const override; + +private: + // Information about a pointer iteration argument to be promoted + struct PointerArgInfo { + unsigned oldIndex; // Original index in the iteration arguments + Value basePointer; // Base pointer value passed as init arg + Value offsetValue; // Offset value used in addptr operation + Value newIterArg; // New integer iteration argument + Value addPtrValue; // The addptr operation result that updates the pointer + }; + + // Check if the loop meets basic transformation conditions + LogicalResult matchLoop(scf::ForOp forOp) const; + + // Collect all pointer iteration arguments that match the promotion pattern + SmallVector collectPointerIterArgs(scf::ForOp forOp) const; + + // Check if a value has pointer tensor type + bool isPointerIterArg(Value iterArg) const; + + // Analyze a pointer iteration argument to determine if it matches the + // promotion pattern + std::optional analyzePointerIterArg(Value iterArg, + Block &loopBody) const; + + // Check if an index corresponds to a pointer argument being promoted + bool isPointerArgIndex(ArrayRef pointerArgs, + unsigned idx) const; + + // Get pointer argument information for a specific index + const PointerArgInfo *getPointerArgInfo(ArrayRef pointerArgs, + unsigned idx) const; + + // Create a new for loop with updated iteration argument types + scf::ForOp createNewForLoop(scf::ForOp forOp, ArrayRef newInitArgs, + ArrayRef newIterArgTypes, + PatternRewriter &rewriter) const; + + // Rewrite the loop body to use integer iteration arguments instead of + // pointers + LogicalResult rewriteLoopBody(scf::ForOp oldForOp, scf::ForOp newForOp, + SmallVector &pointerArgs, + DenseMap &indexMap, + PatternRewriter &rewriter) const; + + // Create new iteration arguments by replacing pointers with integer offsets + std::tuple, SmallVector, + DenseMap> + createNewIterArgs(scf::ForOp forOp, ArrayRef pointerArgs, + PatternRewriter &rewriter) const; + + // Create IR mapping for cloning operations, rebuilding pointers from integer + // offsets + IRMapping createIRMapping(scf::ForOp oldForOp, scf::ForOp newForOp, + SmallVector &pointerArgs, + DenseMap &indexMap, + PatternRewriter &rewriter) const; + + // Reconstruct a pointer value from base pointer and integer offset + Value rebuildPointer(scf::ForOp forOp, ArrayRef pointerArgs, + unsigned idx, PatternRewriter &rewriter) const; + + // Clone instructions from old loop body to new loop body, skipping + // transformed addptr ops + LogicalResult cloneInstructions(Block &oldBody, Block &newBody, + ArrayRef pointerArgs, + DenseMap &indexMap, + IRMapping &mapping, + PatternRewriter &rewriter) const; + + // Clone and transform the yield operation, converting pointer updates to + // integer additions + LogicalResult cloneYieldOp(scf::YieldOp yieldOp, + ArrayRef pointerArgs, + DenseMap &indexMap, + IRMapping &mapping, + PatternRewriter &rewriter) const; + + // Create integer addition for pointer offset updates in the yield operation + Value createIntegerAdd(unsigned idx, ArrayRef pointerArgs, + DenseMap &indexMap, + PatternRewriter &rewriter) const; + + // Extract constant integer value from offset (handles both scalar and tensor + // constants) + std::optional extractConstantOffset(Value offsetValue) const; + + // Replace the original loop results with reconstructed pointers from integer + // results + LogicalResult replaceResults(scf::ForOp oldForOp, scf::ForOp newForOp, + ArrayRef pointerArgs, + DenseMap &indexMap, + PatternRewriter &rewriter) const; + + // Reconstruct final pointer from integer result after the loop + Value reconstructPointer(scf::ForOp forOp, unsigned idx, Value intResult, + ArrayRef pointerArgs, + PatternRewriter &rewriter) const; +}; + +} // namespace CannonicalizerConverter + +#endif \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/MaskAnalysis.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/MaskAnalysis.h new file mode 100755 index 00000000..5e90fe3f --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/MaskAnalysis.h @@ -0,0 +1,131 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ +#ifndef TRITON_TO_STRUCTURED_MASKANALYSIS_H +#define TRITON_TO_STRUCTURED_MASKANALYSIS_H + +#include +#include + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/Value.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Support/LogicalResult.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace TritonToStructuredIncubated { +using namespace mlir; +using namespace triton; + +struct dimInfo { + OpFoldResult offset; + OpFoldResult shape; + OpFoldResult rhs; + size_t dimIndex; + bool hasBroadCast = false; + + enum class CompareType { slt, sge, ult, uge, deafaultType }; + + CompareType currentType = CompareType::deafaultType; + + dimInfo(size_t dimIndex = 0, bool hasBroadCast = false) + : dimIndex(dimIndex), hasBroadCast(hasBroadCast) {} + + dimInfo(OpFoldResult offset, OpFoldResult shape, size_t dimIndex = 0, + bool hasBroadCast = false, + CompareType Type = CompareType::deafaultType, + OpFoldResult rhs = nullptr) + : offset(offset), shape(shape), dimIndex(dimIndex), + hasBroadCast(hasBroadCast), currentType(Type), rhs(rhs) {} + + bool setType(arith::CmpIPredicate Type); + bool compareTypeIsLess() const; + void dump() const; +}; + +struct MaskState { + SmallVector stateInfo; + OpFoldResult scalar; + Value newMask; + + // Recursively parse a Value; call the corresponding function based on the + // defining operation and Value type + LogicalResult parse(Value operand, const Location loc, OpBuilder &builder); + + bool isEmpty() const { return stateInfo.empty() && !scalar; } + void dump() const; + + // Operand is the result of a constant + // Get the value of the constant and assign it to scalar. + LogicalResult parseConstant(arith::ConstantOp constOp, const Location loc, + OpBuilder &builder); + + LogicalResult parseIntScalar(Value scalar, const Location loc, + OpBuilder &builder); + + LogicalResult parseMakeRange(triton::MakeRangeOp rangeOp, const Location loc, + OpBuilder &builder); + + LogicalResult parseExtSI(arith::ExtSIOp op, const Location loc, + OpBuilder &builder); + + LogicalResult parseSplat(triton::SplatOp splatOp, const Location loc, + OpBuilder &builder); + + LogicalResult parseExpandDims(triton::ExpandDimsOp expandDimsOp, + const Location loc, OpBuilder &builder); + + LogicalResult parseAdd(arith::AddIOp addOp, const Location loc, + OpBuilder &builder); + + LogicalResult parseBroadcast(triton::BroadcastOp broadcastOp, + const Location loc, OpBuilder &builder); + + LogicalResult addStates(const MaskState &lhsState, const MaskState &rhsState, + Location loc, OpBuilder &builder); + + LogicalResult addStateScalar(const MaskState &state, + const OpFoldResult scalar, Location loc, + OpBuilder &builder); + + LogicalResult parseCmp(arith::CmpIOp cmpOp, const Location loc, + OpBuilder &builder); + + LogicalResult parseRem(arith::RemSIOp remOp, const Location loc, + OpBuilder &builder); + + LogicalResult parseDiv(arith::DivSIOp divOp, const Location loc, + OpBuilder &builder); + + LogicalResult parseAnd(arith::AndIOp andOp, const Location loc, + OpBuilder &builder); + + LogicalResult analysisMask(Value operand); + + Value createNewMask(const Location loc, OpBuilder &builder); +}; + +} // namespace TritonToStructuredIncubated + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/MemOpConverter.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/MemOpConverter.h new file mode 100755 index 00000000..2cbf2714 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/MemOpConverter.h @@ -0,0 +1,138 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ +#ifndef TRITON_ADAPTER_MEMOPCONVERTER_H +#define TRITON_ADAPTER_MEMOPCONVERTER_H + +#include "bishengir/Dialect/HIVM/IR/HIVM.h" +#include "incubated/Conversion/TritonToStructuredIncubated/MaskAnalysis.h" +#include "incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h" +#include "mlir/Dialect/Arith/Utils/Utils.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/AffineMap.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Value.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace MemOpConverter { + +using namespace mlir; +using namespace triton; + +class LoadConverter : public OpRewritePattern { +public: + explicit LoadConverter(MLIRContext *context, + bool optimizeDynamicOffset = false, + bool enableMaskFallbackConversion = false, + bool compileOn91095 = false) + : OpRewritePattern(context), + optimizeDynamicOffset(optimizeDynamicOffset), + compileOn91095(compileOn91095), + enableMaskFallbackConversion(enableMaskFallbackConversion){}; + + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(triton::LoadOp op, + PatternRewriter &rewriter) const override; + +private: + bool optimizeDynamicOffset; + bool enableMaskFallbackConversion; + bool compileOn91095; +}; + +class StoreConverter : public OpRewritePattern { +public: + explicit StoreConverter(MLIRContext *context, + bool optimizeDynamicOffset = false, + bool enableMaskFallbackConversion = false, + bool compileOn91095 = false) + : OpRewritePattern(context), + optimizeDynamicOffset(optimizeDynamicOffset), + compileOn91095(compileOn91095), + enableMaskFallbackConversion(enableMaskFallbackConversion){}; + + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(triton::StoreOp op, + PatternRewriter &rewriter) const override; + +private: + bool optimizeDynamicOffset; + bool enableMaskFallbackConversion; + bool compileOn91095; +}; + +class MemOpTransformer { +public: + TritonToStructuredIncubated::PtrState ptrState; + TritonToStructuredIncubated::MaskState maskState; + + enum class MemType { load, store, deafaultType }; + + bool optimizeDynamicOffset; + + bool compileOn91095 = false; + + MemType currentType = MemType::deafaultType; + + MemOpTransformer(MemType memType, bool optimizeDynamicOffset = false, + bool compileOn91095 = false) + : currentType(memType), optimizeDynamicOffset(optimizeDynamicOffset), + compileOn91095(compileOn91095) {} + + Value materializeImplicitBroadcast(Value srcTensor, const Location loc, + PatternRewriter &rewriter); + + Value materializeImplicitReshape(Value srcTensor, const Location loc, + PatternRewriter &rewriter); + + Value materializeImplicitSelect(Value srcTensor, Value mask, Value other, + const Location loc, + PatternRewriter &rewriter); + + Value materializeImplicitPermute(Value srcTensor, const Location loc, + PatternRewriter &rewriter); + + Value createNewPtr(Value oldPtr, const Location loc, + PatternRewriter &rewriter); + + Value createNewMask(Value oldPtr, const Location loc, + PatternRewriter &rewriter); + + Value createNewOther(Value oldOther, const Location loc, + PatternRewriter &rewriter); + + bool applyPermuteOnMask(); +}; + +// Create local lock var +hivm::CreateSyncBlockLockOp createSyncBlockLockVar(OpBuilder &builder, + Location loc); + +} // namespace MemOpConverter + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/Passes.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/Passes.h new file mode 100755 index 00000000..d3e88666 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/Passes.h @@ -0,0 +1,37 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_TRITON_TO_STRUCTURE_CONVERSION_PASSES_H +#define TRITON_ADAPTER_TRITON_TO_STRUCTURE_CONVERSION_PASSES_H + +#include "incubated/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "incubated/Conversion/TritonToStructuredIncubated/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // TRITON_ADAPTER_TRITON_TO_UNSTRUCTURE_CONVERSION_PASSES_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/Passes.td b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/Passes.td new file mode 100755 index 00000000..2af51e81 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/Passes.td @@ -0,0 +1,21 @@ +#ifndef TRITON_TO_STRUCTURED_PASSES +#define TRITON_TO_STRUCTURED_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToStructuredIncubated + : Pass<"triton-to-structured-incubated", "mlir::ModuleOp"> { + let summary = "remove reminder/divider and reproduce addptr/mask expression "; + let constructor = "triton::createTritonToStructuredIncubatedPass()"; + let options = + [Option<"enableMaskFallbackConversion", "enable-mask-fallback-conversion", + "bool", /*default*/ "false", + "If enabled, select will perform a fallback conversion when mask " + "matching fails.">, + Option<"optimizeDynamicOffset", "optimize-dynamic-offset", "bool", + /*default*/ "false", "Enable dynamic offset feature">, + Option<"compileOn91095", "compile-on-910-95", "bool", + /*default*/ "false", "compile on 910_95">]; +} + +#endif // TRITON_TO_STRUCTURED_PASSES diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h new file mode 100755 index 00000000..8960f2f7 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h @@ -0,0 +1,177 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ +#ifndef TRITON_TO_STRUCTURED_PTRANALYSIS_H +#define TRITON_TO_STRUCTURED_PTRANALYSIS_H + +#include +#include + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/Value.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Support/LogicalResult.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace TritonToStructuredIncubated { +using namespace mlir; +using namespace triton; + +struct StateInfo { + OpFoldResult stride; + OpFoldResult shape; // rem value + size_t dimIndex; + + StateInfo() : dimIndex(0) {} + StateInfo(OpFoldResult stride, OpFoldResult shape, size_t dimIndex = 0) + : stride(stride), shape(shape), dimIndex(dimIndex) {} + void dump() const; +}; + +struct PtrState { + SmallVector + stateInfo; // shape info when load, maintained with visitOps + SmallVector sizes; // original shape, maintained with visitOps + SmallVector permuteIds; + Value source; // base address (ptr), maintained with visitOps + OpFoldResult offset; // scalar offset (int), maintained with visitOps + + // whether the record needs to be processed in the current pass, when ignore + // is true, it indicates that this scenario should not be processed within the + // current pass + bool shouldLinearize = false; + bool isPermuted = false; + + void dump() const; + bool isEmpty() const; + bool isScalar() const; + bool hasSource() const; + bool isSameSizeAs(const PtrState &x) const; + void analyzePermute(); + + void updatePtrState(SmallVector stateInfo, + SmallVector sizes, Value source, + OpFoldResult offset, const Location loc, + OpBuilder &builder, bool shouldLinearize = false); + + void normalizeState(const Location loc, OpBuilder &builder); + + LogicalResult mulState(const PtrState &lhsState, const PtrState &rhsState, + Operation *op, OpBuilder &builder); + + LogicalResult subState(const PtrState &lhsState, const PtrState &rhsState, + Operation *op, OpBuilder &builder); + + LogicalResult addState(PtrState &lhsState, PtrState &rhsState, Operation *op, + OpBuilder &builder); + + triton::AddPtrOp createAddPtrOp(OpBuilder &builder, Location loc); +}; + +class PtrAnalysis { +public: + // AddptrOp result -> PtrState + llvm::SmallDenseMap knownPtrs; + IRMapping ptrMap; + + bool operandIsScalar(Value operand); + + bool optimizeDynamicOffset; + + PtrAnalysis(bool optimizeDynamicOffset = false) + : optimizeDynamicOffset(optimizeDynamicOffset) {} + + LogicalResult initStateByScalar(Value operand, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult initStateByPointer(Value operand, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandMul(arith::MulIOp mulOp, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandSub(arith::SubIOp subOp, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandMakeRange(triton::MakeRangeOp rangeOp, + PtrState &state, Location loc, + OpBuilder &builder); + + LogicalResult visitOperandBroadcast(triton::BroadcastOp broadcastOp, + PtrState &state, const Location loc, + OpBuilder &builder); + + LogicalResult visitOperandSplat(triton::SplatOp splatOp, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandExpandDims(triton::ExpandDimsOp expandDimsOp, + PtrState &state, const Location loc, + OpBuilder &builder); + + LogicalResult visitOperandConstSplat(arith::ConstantOp op, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandExtSI(arith::ExtSIOp extOp, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandRem(arith::RemSIOp remOp, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandDiv(arith::DivSIOp divOp, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandAdd(arith::AddIOp addOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Recursively parse a Value; call the corresponding + // function based on the defining operation and argument type. + LogicalResult visitOperand(Value operand, PtrState &state, const Location loc, + OpBuilder &builder); + + // Operand is the result of addptr. + // Main assumptions: + // - The ptr field should populate the source field + // - ptr and offset fields should result in same rank + // Expected result: + // - The resulting state for ptr and offset wil be added + LogicalResult visitOperandAddptr(triton::AddPtrOp addptrOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Parse the state of AddPtrOp, insert any instruction needed to + // calculate strides and offsets, build PtrState for this operand, and record + // PtrState for knownPtrs. + LogicalResult rewriteAddptrOp(triton::AddPtrOp op); +}; + +bool isMultiple(const OpFoldResult ÷nd, const OpFoldResult &divisor); +bool isEqual(const OpFoldResult &ofr1, const OpFoldResult &ofr2); +bool isLess(const OpFoldResult &ofs1, const OpFoldResult &ofs2); +bool isGreater(const OpFoldResult &ofs1, const OpFoldResult &ofs2); +bool isOne(const OpFoldResult ofr); +std::optional +extractDivisibilityFromOpFoldResult(mlir::OpFoldResult ofr); + +} // namespace TritonToStructuredIncubated + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.h new file mode 100755 index 00000000..f5cd61da --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.h @@ -0,0 +1,74 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ +#ifndef TRITON_ADAPTER_CONVERSION_TRITONTOSTRUCTURED_H +#define TRITON_ADAPTER_CONVERSION_TRITONTOSTRUCTURED_H + +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#define GEN_PASS_CLASSES +#include "incubated/Conversion/TritonToStructuredIncubated/Passes.h.inc" + +namespace mlir { +namespace triton { + +std::unique_ptr> +createTritonToStructuredIncubatedPass(); + +std::unique_ptr> +createTritonToStructuredIncubatedPass(bool, bool, bool); + +} // namespace triton +} // namespace mlir + +using namespace mlir; +using namespace triton; + +class TritonToStructuredIncubatedPass + : public TritonToStructuredIncubatedBase { +public: + TritonToStructuredIncubatedPass() = default; + + TritonToStructuredIncubatedPass(bool enableMaskFallbackConversion, + bool optimizeDynamicOffset, + bool compileOn91095) { + this->enableMaskFallbackConversion = enableMaskFallbackConversion; + this->optimizeDynamicOffset = optimizeDynamicOffset; + this->compileOn91095 = compileOn91095; + }; + void getDependentDialects(DialectRegistry ®istry) const override; + void runOnOperation() override; + +private: + void populateTritonToStructuredCanonicalizationPatterns( + RewritePatternSet &patterns); + + void populateTritonToStructuredPatterns(RewritePatternSet &patterns, + bool optimizeDynamicOffset, + bool enableMaskFallbackConversion, + bool compileOn91095); +}; + +#endif // TRITON_ADAPTER_CONVERSION_TRITONTOSTRUCTURED_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.h new file mode 100755 index 00000000..e7ec0938 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.h @@ -0,0 +1,112 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#pragma once + +#include "mlir/Pass/Pass.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/IR/PatternMatch.h" + +#define GEN_PASS_DECL_BUBBLEUPOPERATION +#include "flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.h.inc" + +#define GEN_PASS_DEF_BUBBLEUPOPERATION +#include "flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.h.inc" + +namespace mlir { +namespace triton { + +std::unique_ptr> +createBubbleUpOperationPass(const BubbleUpOperationOptions &options = {}); + +} // namespace triton +} // namespace mlir + +using namespace mlir; +using namespace triton; + +template +class BubbleUpExtract : public OpRewritePattern { + static_assert(std::is_same_v || + std::is_same_v); + +public: + using OpRewritePattern::OpRewritePattern; + + explicit BubbleUpExtract(MLIRContext *context, bool enableAggressiveMode); + + LogicalResult matchAndRewrite(ExtractOpTy op, + PatternRewriter &rewriter) const override; + +private: + Value createExtractOp(ExtractOpTy op, Value value, Location loc, + PatternRewriter &rewriter) const; + template + void bubbleUpIntBinaryOp(ExtractOpTy op, BinOpTy binOp, Location loc, + PatternRewriter &rewriter) const; + template + void bubbleUpFloatBinaryOp(ExtractOpTy op, BinOpTy binOp, Location loc, + PatternRewriter &rewriter) const; + + void bubbleUpOperation(ExtractOpTy op, arith::ExtSIOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, arith::CmpIOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, arith::TruncFOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, arith::ExtFOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, arith::FPToSIOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, arith::SIToFPOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, triton::ClampFOp parentOp, + Location loc, PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, arith::CmpFOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, triton::BroadcastOp parentOp, + Location loc, PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, triton::ExpandDimsOp parentOp, + Location loc, PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, triton::SplatOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, triton::MakeRangeOp parentOp, + Location loc, PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, triton::AddPtrOp parentOp, + Location loc, PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, math::FloorOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, math::CeilOp parentOp, Location loc, + PatternRewriter &rewriter) const; + void bubbleUpOperation(ExtractOpTy op, tensor::ExtractSliceOp parentOp, + Location loc, PatternRewriter &rewriter) const; + + bool enableAggressiveMode; +}; + +class BubbleUpOperationPass + : public ::impl::BubbleUpOperationBase { +public: + explicit BubbleUpOperationPass(const BubbleUpOperationOptions &options); + void runOnOperation() override; +}; diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/CMakeLists.txt new file mode 100755 index 00000000..a609ee2d --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToUnstructureIncubated) +add_public_tablegen_target(TritonToUnstructureConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.h new file mode 100755 index 00000000..0eaf3b28 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.h @@ -0,0 +1,260 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ANALYSIS_OFFSETANALYSIS_H +#define TRITON_ANALYSIS_OFFSETANALYSIS_H + +#include "mlir/Dialect/Arith/IR/Arith.h" + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/IR/Value.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" + +namespace mlir { +namespace triton { + +struct PtrOffsetInfo { + /** + Possible status of the ptr offset: + - ScalarLike: + - Tensor's elements are all the same such as [[2.0,2.0,2.0],[2.0,2.0,2.0]] + - Constant integer or floating-point such as 2, 2.0, and `load + tensor<1xptr>` + - Unstructured: + - Not a `ScalarLike` ptr offset + - Or satisfy any below conditions: + - Incontinuous stride such as + - `muli [0,1,2,3] [0,1,2,3]` => [0,1,4,9] + - `divsi [9,8,7] [3,2,1]` => [3,4,7] + - `minsi [3,4,5] [5,4,3]` => [3,4,3] + - From non-`scalarLike` floating point element type such as + - `fptosi [1.0,2.0,3.0]` => [1,2,3] + - Compilation time unknown value + - `load %ptr, %offset` => %value + - Structured: + - orthongonal to `Unstructured` + - if PtrOffsetInfo isn't `Unstructured`, it is `Structured` + + In short: + ScalarLike ⊆ Structured + Unstructured = {x| x ∉ Structured} + + Example: + ``` + %y = sitofp %x + %z = fptosi %y + ``` + If %x is scalarLike (structured), %z will be scalar (structured) as well. + If %x is non-scalarLike structured, %z will be unstructured. + */ + +public: + explicit PtrOffsetInfo(); + PtrOffsetInfo(const PtrOffsetInfo &other); + + explicit PtrOffsetInfo(const Value &ptr); + explicit PtrOffsetInfo(ArrayRef structured); + explicit PtrOffsetInfo(const Value &ptr, bool structured); + explicit PtrOffsetInfo(const Value &ptr, ArrayRef structured); + explicit PtrOffsetInfo(const Value &ptr, const Value &offset, + bool structured); + explicit PtrOffsetInfo(const Value &ptr, const Value &offset, + ArrayRef structured); + + PtrOffsetInfo &operator=(const PtrOffsetInfo &other); + + Value getPtr() const; + Value getOffset() const; + SmallVector getOffsets() const; + SmallVector &getOffsetsRef(); + bool isScalarLike() const; + SmallVector &getStructuredRef(); + const SmallVector &getStructured() const; + int getRank() const; + + void setPtr(const Value &ptr); + void setOffset(const Value &offset); + void setOffsets(ValueRange offsets); + void setStructured(); + void setStructured(int rank); + void setUnstructured(); + void setUnstructured(int rank); + void setStructured(ArrayRef structured); + void setStructured(const PtrOffsetInfo &other); + void setScalarLike(bool scalarLike); + + bool isStructured(int dim) const; + bool isStructured() const; + bool isUnstructured() const; + + void setZeroOffset(); + +private: + Value ptr; + Value offset; + SmallVector tptOffsets; + + bool scalarLike = false; + + SmallVector structured; +}; + +PtrOffsetInfo combineInfo(const PtrOffsetInfo &lhs, const PtrOffsetInfo &rhs); + +void parse(Value operand, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseLoopRegionIterArg(LoopLikeOpInterface loopOp, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap, + BlockArgument regionIterArg); + +void parseArithOp(Operation *arithOp, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseTritonOp(Operation *tritonOp, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseTritonOp(Operation *tritonOp, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseAddPtr(triton::AddPtrOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseSplat(triton::SplatOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +template +void parseBinaryOp(BinOpTy op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseAddI(arith::AddIOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseSubI(arith::SubIOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseIndexCast(arith::IndexCastOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +template +void parseConstantOp(ConstOpTy dst, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseMakeRange(triton::MakeRangeOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseExtSI(arith::ExtSIOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseBitcast(triton::BitcastOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseLoad(triton::LoadOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseMulI(arith::MulIOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseBroadcast(triton::BroadcastOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseExpandDims(triton::ExpandDimsOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseClampF(triton::ClampFOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseSelect(arith::SelectOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseFPToSI(arith::FPToSIOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseSIToFP(arith::SIToFPOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +// FIXME:Z|wait triton version upgrade to 3.4 +// void parseMakeTensorDesc(triton::MakeTensorDescOp op, const Location &loc, +// RewriterBase &rewriter, +// llvm::DenseMap &offsetMap); + +void parseMakeTensorPtr(triton::MakeTensorPtrOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseAdvance(triton::AdvanceOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseReduce(triton::ReduceOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseReduceReturn(triton::ReduceReturnOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseIf(scf::IfOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap, Value dst); + +void parseYield(scf::YieldOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseLoopOp(LoopLikeOpInterface op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap, Value dst); + +void parseExtractSlice(tensor::ExtractSliceOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseExtract(tensor::ExtractOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); + +void parseIntToPtr(triton::IntToPtrOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap); +} // namespace triton + +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.h new file mode 100755 index 00000000..808c5c90 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.h @@ -0,0 +1,38 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_TRITON_TO_UNSTRUCTURE_CONVERSION_PASSES_H +#define TRITON_ADAPTER_TRITON_TO_UNSTRUCTURE_CONVERSION_PASSES_H + +#include "BubbleUpOperation.h" +#include "UnstructureConversionPass.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "incubated/Conversion/TritonToUnstructureIncubated/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif // TRITON_ADAPTER_TRITON_TO_UNSTRUCTURE_CONVERSION_PASSES_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.td b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.td new file mode 100755 index 00000000..5ac7825a --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.td @@ -0,0 +1,29 @@ +#ifndef TRITON_TO_UNSTRUCTURE_CONVERSION_PASSES +#define TRITON_TO_UNSTRUCTURE_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToUnstructureIncubated : Pass<"triton-to-unstructure-incubated", "mlir::ModuleOp"> { + let summary = "Convert Triton for unstructure case(Incubated)"; + let constructor = "triton::createTritonToUnstructureIncubatedPass()"; + let options = [ + Option<"forceScalarizeMode", "force-scalarize-mode", "bool", "false", + "Scalarize unstructured memory access even if structured dimensions are mixed.">, + Option<"compileOn91095", "compile-on-910-95", + "bool", /*default*/"false", + "compile on 910_95">, + Option<"forceSimtTemplate", "force-simt-template", + "bool", /*default*/"false", + "force to use simt template"> + ]; +} + +def BubbleUpOperation : Pass<"bubble-up-operation", "mlir::ModuleOp"> { + let summary = "Apply bubble up operation optimization"; + let constructor = "triton::createBubbleUpOperationPass()"; + let options = [ + Option<"enableAggressiveMode", "enable-aggressive-mode", "bool", "true", + "Enable aggressive bubble up operation.">, + ]; +} +#endif // TRITON_TO_UNSTRUCTURE_CONVERSION_PASSES diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.h new file mode 100755 index 00000000..2d43cdbe --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.h @@ -0,0 +1,150 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITON_ADAPTER_UNSTRUCTURECONVERSION_H +#define TRITON_ADAPTER_UNSTRUCTURECONVERSION_H + +#include "incubated/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.h" +#include "mlir/Pass/Pass.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/IR/PatternMatch.h" +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h" + +#define GEN_PASS_DECL_TRITONTOUNSTRUCTUREINCUBATED +#include "incubated/Conversion/TritonToUnstructureIncubated/Passes.h.inc" + +#define GEN_PASS_DEF_TRITONTOUNSTRUCTUREINCUBATED +#include "incubated/Conversion/TritonToUnstructureIncubated/Passes.h.inc" + +extern bool compileOn91095Flag; +extern bool forceSimtTemplateFlag; + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonToUnstructureIncubatedPass( + const TritonToUnstructureIncubatedOptions &options = {}); + +} // namespace triton +} // namespace mlir + +namespace { + +using namespace mlir; +using namespace triton; + +// For example, in unstructured load case +// %0 = tt.load %structured : tensor<128x128x!tt.ptr> +// %ptr_2 = tt.splat %arg1 : !tt.ptr -> tensor<128x128x!tt.ptr> +// %1 = tt.addptr %ptr_2, %0 : tensor<128x128x!tt.ptr>, +// tensor<128x128xi32> %2 = tt.load %1 : tensor<128x128x!tt.ptr> tt.store +// %output %2 : tensor<128x128x!tt.ptr> +// +// +// In this case, this will be converted to +// +// %0 = tt.load %structured : tensor<128x128x!tt.ptr> +// %1 = tensor.empty() : tensor<128x128xf32> +// %2 = scf.for %arg2 = %c0 to %c128 step %c1 iter_args(%arg3 = %1) -> +// (tensor<128x128xf32>) { +// %4 = scf.for %arg4 = %c0 to %c128 step %c1 iter_args(%arg5 = %arg3) -> +// (tensor<128x128xf32>) { +// %extracted = tensor.extract %10[%arg3, %arg5] {DiscreteMemAccess} : +// tensor<128x128xi32> %5 = arith.extsi %extracted : i32 to i64 %6 = +// tt.addptr %arg1, %5 : !tt.ptr, i64 %7 = tt.load %6 +// {DiscreteMemAccess} : tt.ptr %inserted_slice = tensor.insert_slice +// %7 into %arg5[%arg2, %arg4] [1, 1] [128, 1] {DiscreteMemAccess} : +// tensor<1x1xf32> into tensor<128x128xf32> scf.yield %inserted_slice : +// tensor<128x128xf32> +// } +// scf.yield %4 : tensor<128x128xf32> +// } +// tt.store %output %2 : tensor<128x128x!tt.ptr> +template +class UnstructuredMemAccessConverter : public OpRewritePattern { + static_assert(std::is_same_v || + std::is_same_v || + std::is_same_v || + std::is_same_v); + +public: + using OpRewritePattern::OpRewritePattern; + + explicit UnstructuredMemAccessConverter( + MLIRContext *context, bool forceScalarizeMode, + const llvm::DenseMap &offsetMap, + const llvm::SmallDenseMap &fromTensorArg); + LogicalResult matchAndRewrite(MemAccOpTy op, + PatternRewriter &rewriter) const override; + +private: + bool checkUnstructureAnnotated(MemAccOpTy op, + PatternRewriter &rewriter) const; + Value createExtractOp(Location loc, Value value, PatternRewriter &rewriter, + ArrayRef iterIdx) const; + Value createExtractOp(Location loc, Value value, PatternRewriter &rewriter, + ArrayRef offsets, + ArrayRef sizes, + ArrayRef strides) const; + template + typename std::enable_if, void>::type + splatAndLoadScenario(MemAccOpTy op, int rank, + PatternRewriter &rewriter) const; + + template + MemAccOpTy createMemAccOp(MemAccOpTy op, Value ptrToAccess, Location loc, + PatternRewriter &rewriter, + Args &&...args) const = delete; + + const llvm::DenseMap &offsetMap; + const llvm::SmallDenseMap &fromTensorArg; + bool forceScalarizeMode; +}; + +class TritonToUnstructureIncubatedPass + : public ::impl::TritonToUnstructureIncubatedBase< + TritonToUnstructureIncubatedPass> { +public: + explicit TritonToUnstructureIncubatedPass( + const TritonToUnstructureIncubatedOptions &options); + void getDependentDialects(DialectRegistry ®istry) const override; + + void runOnOperation() override; + +private: + void runPreparse(LoopLikeOpInterface op); + template || + std::is_same_v || + std::is_same_v || + std::is_same_v>> + void runParse(MemAccOpTy op); + llvm::DenseMap offsetMap; + llvm::DenseMap offsetMapForLoopArgs; + llvm::SmallDenseMap fromTensorArg; +}; + +} // namespace + +#endif // TRITON_ADAPTER_UNSTRUCTURECONVERSION_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/CMakeLists.txt new file mode 100755 index 00000000..e69de29b diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/InterleaveOptimization.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/InterleaveOptimization.h new file mode 100755 index 00000000..b67e0ddf --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/InterleaveOptimization.h @@ -0,0 +1,93 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#pragma once + +#include "incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/UseAnalysis.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Arith/Utils/Utils.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/GPU/IR/GPUDialect.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Matchers.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Value.h" +#include "mlir/Support/LogicalResult.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include +#include +#include +#include +#include + +namespace mlir { +namespace triton { + +enum class IndexMode : int { EVEN_MODE = 0, ODD_MODE = 1 }; + +MemRefType expandInterleaveMemRefType(MemRefType originType); + +std::pair +recountReinterpretCastOffset(OpFoldResult originOffset, Builder &builder); + +LogicalResult +DeinterleaveStatusOptimization(triton::LoadOp op, + triton::LoadOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter); + +LogicalResult DeinterleaveStatusWithMaskOptimization( + triton::LoadOp op, triton::LoadOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter, + mlir::triton::Incubated::MaskState &mstate, Value localMem); + +LogicalResult +InterleaveStatusOptimization(SmallVector materializeVec); + +LogicalResult +InterleaveStatusWithMaskOptimization(SmallVector materializeVec); + +} // namespace triton +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/Utils.h b/third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/Utils.h new file mode 100755 index 00000000..7f890c47 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/Utils.h @@ -0,0 +1,252 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#ifndef TRITONNPU_UTILS_UTILS_H +#define TRITONNPU_UTILS_UTILS_H + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/Dialect/Utils/StructuredOpsUtils.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Operation.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "llvm/ADT/ArrayRef.h" + +#include +#include + +namespace mlir { + +namespace ConverterUtils { + +const std::string GeneratedByMakeTensorPtrTAG = "GeneratedByMakeTensorPtr"; +const std::string discreteMaskAttrName = "DiscreteMask"; +const std::string discreteAttrName = "DiscreteMemAccess"; + +bool isaPermutedMemRefType(MemRefType); + +std::optional +getLastStrideOfReinterpretCastOp(memref::ReinterpretCastOp op); + +Value getTransposedValue(Value source, const Location loc, + ConversionPatternRewriter &rewriter, + llvm::ArrayRef order); + +SmallVector getNParallelLoopsAttrs(unsigned n); + +Value getScalarValue(Value operand, Location loc, + ConversionPatternRewriter &rewriter); + +memref::SubViewOp makeSubViewOp(Value src, + const llvm::SmallVector &sizes, + const Location &loc, + ConversionPatternRewriter &rewriter); + +tensor::ExtractSliceOp +makeExtractSliceOp(Value src, const llvm::SmallVector &sizes, + const Location &loc, ConversionPatternRewriter &rewriter); + +std::optional getFullShapeOp(Value val, + ConversionPatternRewriter &rewriter); + +SmallVector +getBoundarySizes(llvm::ArrayRef boundaryCheck, Value ptr, + const Location &loc, ConversionPatternRewriter &rewriter); + +SmallVector getBroadcastDims(RankedTensorType src, + RankedTensorType dst); + +SmallVector getUnbroadcastDims(RankedTensorType src, + RankedTensorType dst); + +} // namespace ConverterUtils + +class ConversionPatternRewriter; + +namespace triton { + +enum class IndirectLoadInterfaceOpType { Undefined = 0, Load = 1, Calc = 2 }; + +// Traceback from rootOp to find the targetOp with the specified condition +mlir::Operation * +findFirstMatchingOperandDef(mlir::Operation *rootOp, + const std::function &condFn); + +void traverseBackwardUpdateOperandChainIf( + Operation *op, std::function conditionFn, + std::function stopFn, + std::function actionFn, OpBuilder &builder, + DenseSet &handledOperation); + +void traverseBackwardUpdateOperandChainIf( + Operation *rootOp, std::function conditionFn, + std::function stopFn, + std::function actionFn); + +void traverseForwardUpdateUserChainIf( + Operation *op, std::function conditionFn, + std::function stopFn, + std::function actionFn, OpBuilder &builder, + llvm::SmallPtrSet &stopOps); + +void traverseForwardUpdateUserChainIf( + Operation *rootOp, std::function conditionFn, + std::function stopFn, + std::function actionFn, + llvm::SmallPtrSet &stopOps); + +// UseAnalysis will tag operations whose results are used only as meta-data +// with "MetaUse" tag. +bool isMetaUse(Operation *op); + +bool isMixUse(Operation *op); + +IndirectLoadInterfaceOpType getIndirectLoadInterfaceOpType(Operation *op); + +bool opIsIndirectLoad(Operation *op); + +bool opIsIndirectCalc(Operation *op); + +/// Maximum expected rank for loop tiling in tensor operations. +static constexpr int kMaxTiledRank = 4; + +/// This function generates a series of `scf.for` loops for the given dimensions +/// in `loopDims`. Although the loops are created sequentially, nesting is +/// simulated by adjusting the insertion point to the body of the last created +/// loop. This allows the `bodyFunc` to be inserted into the innermost scope. +/// +/// \param rewriter The MLIR OpBuilder used to create operations. +/// \param loc The source location information for debuggability. +/// \param target The memref value whose dimensions are being looped over. +/// \param loopDims An array of dimension indices to create loops for. +/// \param bodyFunc A callable that defines the operations to insert in the +/// innermost loop. +/// It takes a SmallVector of induction variables (one per +/// loop). +/// +template +void createSimpleNestedLoops(OpBuilder &rewriter, Location loc, Value target, + ArrayRef loopDims, Func bodyFunc) { + MemRefType type = cast(target.getType()); + int rank = type.getRank(); + + Value zero = rewriter.create(loc, 0); + Value one = rewriter.create(loc, 1); + + llvm::SmallVector loops; + llvm::SmallVector ivs; + + for (int dim : loopDims) { + Value ub; + if (type.isDynamicDim(dim)) { + ub = rewriter.create(loc, target, dim).getResult(); + } else { + ub = rewriter.create(loc, type.getDimSize(dim)); + } + + auto forOp = rewriter.create(loc, zero, ub, one); + rewriter.setInsertionPointToStart(forOp.getBody()); + loops.push_back(forOp); + ivs.push_back(forOp.getInductionVar()); + } + + bodyFunc(ivs); + + if (!loops.empty()) { + rewriter.setInsertionPointAfter(loops.front()); + } +} + +scf::ForOp createNestedLoops( + OpBuilder &builder, Location loc, unsigned currentDim, unsigned totalDims, + ValueRange LBs, ValueRange UBs, ValueRange steps, SmallVector &ivs, + ValueRange initArgs, + function_ref &, ValueRange)> + bodyBuilder); + +ModuleOp getModuleOpFromOperation(Operation *op); + +} // namespace triton + +class OpBuilder; + +OpFoldResult addOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b); + +OpFoldResult subOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b); + +OpFoldResult mulOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b); + +OpFoldResult divOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b); + +OpFoldResult remOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b); + +OpFoldResult minOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b); + +OpFoldResult maxOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b); + +enum class ReduceWithIndexType { MAX, MIN }; +enum class TieBreakType { LEFT, RIGHT }; + +struct ReduceWithIndexParams { + ReduceWithIndexType withIndexType; + TieBreakType tieBreakType; + bool isUnsignedSrc; +}; + +std::optional +getReduceWithIndexParams(triton::ReduceOp reduceOp); + +void addReduceWithIndexAttr(ReduceWithIndexParams params, + ConversionPatternRewriter &rewriter, + linalg::ReduceOp reduceOp); + +OpFoldResult getOpFoldResultOfLayoutInfo(Value value, OpBuilder &builder); + +enum class TypelessValue { Undefined = 0, Zero = 1, Min = 2, Max = 3 }; + +FailureOr specializeTypelessValueToAttr(TypelessValue, Type, + OpBuilder &); + +FailureOr specializeTypelessValueToConstant(TypelessValue, Type, + Location, OpBuilder &); + +std::optional getIntAttr(const OpFoldResult ofr); + +Value materializeValue(OpBuilder &builder, Location loc, OpFoldResult ofr); + +bool isZero(const OpFoldResult ofr); + +Value convertToIndexIfNeeded(Value intValue, const Location &loc, OpBuilder &b); + +} // namespace mlir + +#endif // TRITONNPU_UTILS_UTILS_H diff --git a/third_party/wafer/third_party/flir/include/incubated/Dialect/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Dialect/CMakeLists.txt new file mode 100755 index 00000000..e3548937 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Dialect/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(TritonStructuredIncubated) diff --git a/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/CMakeLists.txt new file mode 100755 index 00000000..c3c9d2da --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/CMakeLists.txt @@ -0,0 +1,8 @@ +set(LLVM_TARGET_DEFINITIONS TritonStructuredDialectIncubated.td) +mlir_tablegen(TritonStructuredDialectIncubated.h.inc -gen-dialect-decls -dialect=tts) +mlir_tablegen(TritonStructuredDialectIncubated.cpp.inc -gen-dialect-defs -dialect=tts) +mlir_tablegen(TritonStructuredOpsIncubated.h.inc -gen-op-decls) +mlir_tablegen(TritonStructuredOpsIncubated.cpp.inc -gen-op-defs) + + +add_public_tablegen_target(TritonStructuredIncubatedTableGen) diff --git a/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.h b/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.h new file mode 100755 index 00000000..2b310777 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.h @@ -0,0 +1,30 @@ +#ifndef MLIR_DIALECT_TRITON_STRUCTURED_INCUBATED_IR_TRITON_STRUCTURED_DIALECT_INCUBATED_H_ +#define MLIR_DIALECT_TRITON_STRUCTURED_INCUBATED_IR_TRITON_STRUCTURED_DIALECT_INCUBATED_H_ + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/SymbolTable.h" +#include "mlir/IR/TypeSupport.h" +#include "mlir/IR/Types.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/IR/Dialect.h" + +using namespace mlir; +using namespace mlir::triton; +//===----------------------------------------------------------------------===// +// TritonStructured Operations +//===----------------------------------------------------------------------===// +#include "incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.h.inc" + +// Include the auto-generated header file containing the declarations of the +// TritonStructured operations. +#define GET_OP_CLASSES + +#include "incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredOpsIncubated.h.inc" + +#endif diff --git a/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.td b/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.td new file mode 100755 index 00000000..780dcbed --- /dev/null +++ b/third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.td @@ -0,0 +1,66 @@ +#ifndef TRITON_STRUCTURED_DIALECT_INCUBATED +#define TRITON_STRUCTURED_DIALECT_INCUBATED + +include "mlir/IR/OpBase.td" +include "triton/Dialect/Triton/IR/TritonDialect.td" +include "triton/Dialect/Triton/IR/TritonTypes.td" +include "triton/Dialect/Triton/IR/TritonAttrDefs.td" +//include "triton/Dialect/Triton/IR/TritonTypes.td" +include "mlir/Interfaces/SideEffectInterfaces.td" + + + +def Triton_Structured_Dialect_Incubated : Dialect { + let name = "tts"; + + let cppNamespace = "::mlir::tts::Incubated"; + + let summary = "Structured Triton operations"; + + let description = [{ + Triton Structured Dialect. + }]; + + let dependentDialects = [ + "triton::TritonDialect" + ]; + + let usePropertiesForAttributes = 1; +} + +// +// Op Base +// +class TTS_Op traits = []> : + Op { +} + + +// SameVariadicResultSize +// AttrSizedResultSegments +def TTS_GetStructuredStateOp : TTS_Op<"get_structured_state", [AttrSizedResultSegments, Pure]> { + let summary = "Placeholder for the structured pointer states computed during PtrAnalysis."; + let description = "Used to pass the offsets and strides to scf.for op to simplify IR rewrites."; + + let arguments = (ins AnyTypeOf<[TT_PtrLike, I32Tensor, I64Tensor,I16Tensor,I8Tensor,I1Tensor]>:$input); + let results = (outs AnyTypeOf<[TT_PtrLike, I32Tensor, I64Tensor,I16Tensor,I8Tensor,I1Tensor]>:$structured, Variadic:$offsets, Variadic:$strides); + + let builders = [ + OpBuilder<(ins "Value":$input)>, + ]; + + let extraClassDeclaration = [{ + static std::optional, SmallVector>> + getOffsetAndStrideTypes(MLIRContext *context, Type ptrLikeType); + + static std::optional> + getOffsetAndStrideSegmentSizes(Type ptrLikeType); + }]; + + let hasFolder = 0; + let hasVerifier = 1; +} + + + +#endif // TRITON_STRUCTURED_DIALECT_INCUBATED diff --git a/third_party/wafer/third_party/flir/include/mlir-ext/CMakeLists.txt b/third_party/wafer/third_party/flir/include/mlir-ext/CMakeLists.txt new file mode 100755 index 00000000..0ca0f41c --- /dev/null +++ b/third_party/wafer/third_party/flir/include/mlir-ext/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(Dialect) diff --git a/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/CMakeLists.txt b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/CMakeLists.txt new file mode 100755 index 00000000..35948d69 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(MathExt) diff --git a/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/CMakeLists.txt b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/CMakeLists.txt new file mode 100755 index 00000000..7d59dce8 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/CMakeLists.txt new file mode 100755 index 00000000..f3a3f005 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/CMakeLists.txt @@ -0,0 +1,10 @@ +set(LLVM_TARGET_DEFINITIONS MathExtBase.td) +mlir_tablegen(MathExtDialect.h.inc -gen-dialect-decls) +mlir_tablegen(MathExtDialect.cpp.inc -gen-dialect-defs) +add_public_tablegen_target(MLIRMathExtDialectIncGen) + +set(LLVM_TARGET_DEFINITIONS MathExtOps.td) +mlir_tablegen(MathExtOps.h.inc -gen-op-decls) +mlir_tablegen(MathExtOps.cpp.inc -gen-op-defs) + +add_public_tablegen_target(MLIRMathExtOpsIncGen) diff --git a/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExt.h b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExt.h new file mode 100755 index 00000000..7057bcaf --- /dev/null +++ b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExt.h @@ -0,0 +1,22 @@ +#ifndef MLIR_DIALECT_MATHEXT_IR_MATHEXT_H_ +#define MLIR_DIALECT_MATHEXT_IR_MATHEXT_H_ + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/OpImplementation.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "mlir/Interfaces/VectorInterfaces.h" + +//===----------------------------------------------------------------------===// +// MathExt Dialect +//===----------------------------------------------------------------------===// +#include "mlir-ext/Dialect/MathExt/IR/MathExtDialect.h.inc" + +//===----------------------------------------------------------------------===// +// MathExt Dialect Operations +//===----------------------------------------------------------------------===// +#define GET_OP_CLASSES +#include "mlir-ext/Dialect/MathExt/IR/MathExtOps.h.inc" + +#endif // MLIR_DIALECT_MATHEXT_IR_MATHEXT_H_ \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExtBase.td b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExtBase.td new file mode 100755 index 00000000..d3c05ce4 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExtBase.td @@ -0,0 +1,25 @@ +#ifndef MATHEXT_BASE +#define MATHEXT_BASE + +include "mlir/IR/OpBase.td" + +//===----------------------------------------------------------------------===// +// MathExt dialect definition. +//===----------------------------------------------------------------------===// + +def MathExt_Dialect : Dialect { + let name = "mathext"; + let cppNamespace = "::mlir::mathext"; + let summary = "Math extensions dialect"; + let description = [{ + This dialect provides additional mathematical operations that extend + the standard MLIR math dialect with operations like fmod that are + not available in the base math dialect. + }]; + let hasConstantMaterializer = 1; + let dependentDialects = [ + "::mlir::arith::ArithDialect" + ]; +} + +#endif // MATHEXT_BASE \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExtOps.td b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExtOps.td new file mode 100755 index 00000000..33a97985 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExtOps.td @@ -0,0 +1,128 @@ +#ifndef MATHEXT_OPS +#define MATHEXT_OPS + +include "mlir-ext/Dialect/MathExt/IR/MathExtBase.td" +include "mlir/Dialect/Arith/IR/ArithBase.td" +include "mlir/Dialect/Arith/IR/ArithOpsInterfaces.td" +include "mlir/Interfaces/InferTypeOpInterface.td" +include "mlir/Interfaces/VectorInterfaces.td" +include "mlir/Interfaces/SideEffectInterfaces.td" + +// Base class for math dialect ops. +class MathExt_Op traits = []> : + Op] # + ElementwiseMappable.traits>; + +// Base class for unary math operations on integer types. Require an operand +// and result of the same type. This type can be an integer type, vector or +// tensor thereof. +class MathExt_IntegerUnaryOp traits = []> : + MathExt_Op { + let arguments = (ins SignlessIntegerOrIndexLike:$operand); + let results = (outs SignlessIntegerOrIndexLike:$result); + + let assemblyFormat = "$operand attr-dict `:` type($result)"; +} + +// Base class for floating point classification ops. Require an operand and +// result of the same shape, which can be a floating point scalar, a vector or a +// tensor thereof. +class MathExt_FloatClassificationOp traits = []> : + MathExt_Op, + TypesMatchWith< + "result type has i1 element type and same shape as operands", + "operand", "result", "::getI1SameShape($_self)">]> { + let arguments = (ins FloatLike:$operand, + DefaultValuedAttr:$fastmath); + let results = (outs BoolLike:$result); + + let assemblyFormat = "$operand attr-dict `:` type($operand)"; +} + +// Base class for unary math operations on floating point types. Require an +// operand and result of the same type. This type can be a floating point type, +// vector or tensor thereof. +class MathExt_FloatUnaryOp traits = []> : + MathExt_Op]> { + let arguments = (ins FloatLike:$operand, + DefaultValuedAttr:$fastmath); + let results = (outs FloatLike:$result); + + let assemblyFormat = [{ $operand (`fastmath` `` $fastmath^)? + attr-dict `:` type($result) }]; +} + +// Base class for binary math operations on integer types. Require two +// operands and one result of the same type. This type can be an integer +// type, vector or tensor thereof. +class MathExt_IntegerBinaryOp traits = []> : + MathExt_Op { + let arguments = (ins SignlessIntegerOrIndexLike:$lhs, SignlessIntegerOrIndexLike:$rhs); + let results = (outs SignlessIntegerOrIndexLike:$result); + + let assemblyFormat = "$lhs `,` $rhs attr-dict `:` type($result)"; +} + +// Base class for binary math operations on floating point types. Require two +// operands and one result of the same type. This type can be a floating point +// type, vector or tensor thereof. +class MathExt_FloatBinaryOp traits = []> : + MathExt_Op]> { + let arguments = (ins FloatLike:$lhs, FloatLike:$rhs, + DefaultValuedAttr:$fastmath); + let results = (outs FloatLike:$result); + + let assemblyFormat = [{ $lhs `,` $rhs (`fastmath` `` $fastmath^)? + attr-dict `:` type($result) }]; +} + +// Base class for floating point ternary operations. Require three operands and +// one result of the same type. This type can be a floating point type, vector +// or tensor thereof. +class MathExt_FloatTernaryOp traits = []> : + MathExt_Op]> { + let arguments = (ins FloatLike:$a, FloatLike:$b, FloatLike:$c, + DefaultValuedAttr:$fastmath); + let results = (outs FloatLike:$result); + + let assemblyFormat = [{ $a `,` $b `,` $c (`fastmath` `` $fastmath^)? + attr-dict `:` type($result) }]; +} + +//===----------------------------------------------------------------------===// +// FModOp +//===----------------------------------------------------------------------===// + +def MathExt_FModOp : MathExt_FloatBinaryOp<"fmod"> { + let summary = "floating point modulo (remainder) operation"; + let description = [{ + `%r = mathext.fmod %a, %b : f32` + }]; + let hasFolder = 1; +} + +//===----------------------------------------------------------------------===// +// DivRzOp +//===----------------------------------------------------------------------===// + +def MathExt_DivRzOp : MathExt_FloatBinaryOp<"div_rz"> { + let summary = "floating point division operation with round-zero"; + let description = [{ + `%r = mathext.div_rz %a, %b : f32` + }]; + let hasFolder = 1; +} + +#endif // MATHEXT_OPS diff --git a/third_party/wafer/third_party/flir/include/npu/CMakeLists.txt b/third_party/wafer/third_party/flir/include/npu/CMakeLists.txt new file mode 100755 index 00000000..0ca0f41c --- /dev/null +++ b/third_party/wafer/third_party/flir/include/npu/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(Dialect) diff --git a/third_party/wafer/third_party/flir/include/npu/Dialect/CMakeLists.txt b/third_party/wafer/third_party/flir/include/npu/Dialect/CMakeLists.txt new file mode 100755 index 00000000..b7e65956 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/npu/Dialect/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(TritonAscend) diff --git a/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/CMakeLists.txt b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/CMakeLists.txt new file mode 100755 index 00000000..488e9132 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/CMakeLists.txt @@ -0,0 +1,15 @@ +set(MLIR_BINARY_DIR ${CMAKE_BINARY_DIR}) + +set(LLVM_TARGET_DEFINITIONS TritonAscendOps.td) +mlir_tablegen(TritonAscendDialect.h.inc -gen-dialect-decls -dialect=ascend) +mlir_tablegen(TritonAscendDialect.cpp.inc -gen-dialect-defs -dialect=ascend) +mlir_tablegen(TritonAscendOps.h.inc -gen-op-decls) +mlir_tablegen(TritonAscendOps.cpp.inc -gen-op-defs) +add_mlir_doc(TritonAscendDialect TritonAscendDialect dialects/ -gen-dialect-doc) +add_mlir_doc(TritonAscendOps TritonAscendOps dialects/ -gen-op-doc) +add_public_tablegen_target(TritonAscendTableGen) + +set(LLVM_TARGET_DEFINITIONS TritonAscendAttrDefs.td) +mlir_tablegen(TritonAscendOpsAttrDefs.h.inc -gen-attrdef-decls) +mlir_tablegen(TritonAscendOpsAttrDefs.cpp.inc -gen-attrdef-defs) +add_public_tablegen_target(TritonAscendAttrDefsIncGen) diff --git a/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendAttrDefs.td b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendAttrDefs.td new file mode 100755 index 00000000..33743fcd --- /dev/null +++ b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendAttrDefs.td @@ -0,0 +1,25 @@ +//===-- TritonAscendAttrDefs.td - dialect attributes def. ----*- tablegen -*-===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_ASCEND_ATTRDEFS +#define TRITON_ASCEND_ATTRDEFS + +include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.td" + +include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/EnumAttr.td" + +class TritonAscend_Attr traits = []> + : AttrDef { + let mnemonic = attrMnemonic; + let cppNamespace = "::mlir::triton::ascend"; +} + + + +#endif // TRITON_ASCEND_ATTRDEFS diff --git a/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendDialect.h b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendDialect.h new file mode 100755 index 00000000..b8ca9fa2 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendDialect.h @@ -0,0 +1,34 @@ +//===- TritonAscendDialect.h - MLIR TritonAscend dialect --------------*- C++ +//-*-===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file defines the TritonAscend dialect in MLIR, containing Ascend +// operations. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_DIALECT_ASCEND_DIALECT_H +#define TRITON_DIALECT_ASCEND_DIALECT_H + +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/OpDefinition.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h.inc" + +#define GET_ATTRDEF_CLASSES +#include "npu/Dialect/TritonAscend/IR/TritonAscendOpsAttrDefs.h.inc" + +#define GET_OP_CLASSES +#include "npu/Dialect/TritonAscend/IR/TritonAscendOps.h.inc" + +namespace mlir::triton::ascend {} // namespace mlir::triton::ascend + +#endif // TRITON_DIALECT_ASCEND_DIALECT_H diff --git a/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendDialect.td b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendDialect.td new file mode 100755 index 00000000..a7314eff --- /dev/null +++ b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendDialect.td @@ -0,0 +1,31 @@ +//===-- TritonAscendDialect.td - dialect op definitions -------*- tablegen -*-===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_ASCEND_DIALECT +#define TRITON_ASCEND_DIALECT + +include "mlir/IR/OpBase.td" + +def TritonAscend_Dialect : Dialect { + let name = "ascend"; + let cppNamespace = "::mlir::triton::ascend"; + let summary = "The TritonAscend dialect in Triton."; + + let description = [{ + TritonAscend is a dialect for representing operations on Ascend NPUs. + }]; + + let dependentDialects = [ + "mlir::LLVM::LLVMDialect", + "triton::TritonDialect", + ]; + + let extraClassDeclaration = [{}]; +} + +#endif // TRITON_ASCEND_DIALECT diff --git a/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendOps.td b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendOps.td new file mode 100755 index 00000000..d4d49623 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendOps.td @@ -0,0 +1,488 @@ +//===-- TritonAscendOps.td - TritonAscend op definitions ---------*- tablegen -*-===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This is the TritonAscend IR operation definition file. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_ASCEND_OPS +#define TRITON_ASCEND_OPS + +include "mlir/IR/OpBase.td" +include "mlir/IR/EnumAttr.td" +include "mlir/Dialect/LLVMIR/LLVMTypes.td" +include "mlir/Interfaces/SideEffectInterfaces.td" +include "mlir/IR/OpAsmInterface.td" +include "mlir/Interfaces/InferTypeOpInterface.td" // SameOperandsAndResultType + +include "triton/Dialect/Triton/IR/TritonAttrDefs.td" +include "triton/Dialect/Triton/IR/TritonTypes.td" +include "triton/Dialect/Triton/IR/TritonInterfaces.td" + +include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.td" +include "npu/Dialect/TritonAscend/IR/TritonAscendAttrDefs.td" + +//===----------------------------------------------------------------------===// +// TritonAscend op definitions +//===----------------------------------------------------------------------===// + +class TT_Ascend_Op traits = []> : + Op; + + +// +// Interfaces +// +def GlobalMemory : Resource<"::mlir::triton::GlobalMemory">; + + +// +// Annotation Op +// +def AnnotationOp : TT_Ascend_Op<"annotation", [Pure, MemoryEffects<[MemWrite]>]> { + let summary = "Annotate a tensor with key-value attribute pairs"; + let description = [{ + `ascend.annotation` operation can be used to annotate a tensor with + key-value attribute pairs. + + Example: + ```mlir + ascend.annotation %target {key : val} + ``` + }]; + let arguments = (ins TT_Tensor:$src); + let assemblyFormat = [{ + $src attr-dict `:` type($src) + }]; +} + + +// +// Mod Op +// +def ModOp : TT_Ascend_Op<"mod", [Pure]> { + let summary = "Mod operation (%) of input tensors."; + let description = [{ + Performs element-wise division with remainder of input tensors. + }]; + + let arguments = (ins TT_Tensor:$lhs, TT_Tensor:$rhs); + let results = (outs TT_Tensor:$result); + + let assemblyFormat = "$lhs `,` $rhs attr-dict `:` type($lhs) type($rhs) `->` type($result)"; +} + +// +// EmbeddingGather Op +// +def EmbeddingGatherOp : TT_Ascend_Op<"embedding_gather", [ + DeclareOpInterfaceMethods, + SameVariadicOperandSize, +]> { + let summary = "Gather load from a tensor pointer with the embedding semantics"; + + let arguments = ( + ins TT_Ptr:$src, + TT_Tensor:$idx, + AnyTypeOf<[I32, I64]>:$bound, + AnyTypeOf<[I32, I64]>:$blocksize, + Variadic>:$offsets, + Variadic>:$numels + ); + + let results = (outs TT_Tensor:$result); + + let assemblyFormat = [{ + $src `:` type($src) `,` $idx `:` type($idx) `,` + $bound `:` type($bound) `,` $blocksize `:` type($blocksize) `,` + `[` $offsets `:` type($offsets) `]` `,` `[` $numels `:` type($numels) `]` + attr-dict `->` type($result) + }]; + + let builders = [ + OpBuilder<(ins + "Value":$src, + "Value":$idx, + "Value":$bound, + "Value":$blocksize, + "ValueRange":$offsets, + "ValueRange":$numels + )> + ]; + // let hasCanonicalizer = 1; +} + +// +// IndexPut Op +// +def IndexPutOp : TT_Ascend_Op<"index_put", [ + MemoryEffects<[MemWrite]>, + SameVariadicOperandSize, +]> { + let summary = "Scatter store to a tensor pointer with embedding semantics"; + + let description = [{ + Index put values from a tensor into a destination tensor. + + The operation takes: + - ptr: pointer type, the destination tensor pointer (in GM) + - index: tensor, a index to scatter (in UB) + - value: tensor, a value to store (in UB) + - dim: int32, the dimension to scatter along + - index_boundary: int64, the upper boundary for index values + - end_offset: tuple of int, the offsets of each dimension for the end of the scatter region + - start_offset: tuple of int, the offsets of each dimension for the start of the scatter region + - dst_stride: tuple of int, the stride of each dimension of destination tensor + + + Constraints: + - `ptr` and `value` must have the same rank. + - `ptr.dtype` only supports `float16`, `bfloat16`, `float32` currently. + - `index` must be an integer tensor. If `index.rank` != 1, it will be reshaped to 1D. + - `index.numel` must equal `value.shape[dim]`. + - `value` support 2~5D tensors. + - `dim` must be valid (0 <= dim < rank(value) - 1). + }]; + + let arguments = ( + ins TT_Ptr:$ptr, + TT_Tensor:$index, + TT_Tensor:$value, + TT_Int:$dim, + TT_Int:$indexBoundary, + Variadic>:$endOffset, + Variadic>:$startOffset, + Variadic>:$dstStride + ); + + let assemblyFormat = [{ + $ptr `:` type($ptr) `,` $index `:` type($index) `,` + $value `:` type($value) `,` $dim `:` type($dim) `,` $indexBoundary `:` type($indexBoundary) `,` + `[` $endOffset `:` type($endOffset) `]` `,` `[` $startOffset `:` type($startOffset) `]` `,` + `[` $dstStride `:` type($dstStride) `]` + attr-dict + }]; +} + +// +// GatherOutToUb Op +// +def GatherOutToUbOp : TT_Ascend_Op<"gather_out_to_ub", [ + DeclareOpInterfaceMethods, + AttrSizedOperandSegments, +]> { + let summary = "Gather load from a tensor pointer with the embedding semantics"; + + let description = [{ + Gather from a source tensor in Global Memory (GM) to Unified Buffer (UB) + along a specified dimension with out-of-bound handling. + + The operation takes: + - src: pointer type, the source tensor pointer (in GM) + - index: tensor, a tensor to gather (in UB) + - index_boundary: int64, the upper boundary for index values + - dim: int32, the dimension to gather along + - src_stride: tuple of int64, the stride of each dimension of src tensor + - end_offset: tuple of int32, the end offsets of each dimension for index tensor + - start_offset: tuple of int32, the start offsets of each dimension for index tensor + - other(Optional): scalar value, the default value when index is out of boundary (in UB) + + Returns: + a tensor, with the same shape as `index.shape` (in UB) + + Constraints: + - `src` and `index` must have the same rank. + - `src.dtype` only supports `float16`, `bfloat16`, `float32` currently. + - `index` must be an integer tensor, with rank between 1 and 5. + - `dim` must be valid (0 <= dim < rank(index)). + - `other` must be a scalar value. + - For every dimension `i` not equal to `dim`, `index.size[i]` <= `src.size[i]`. + - The output shape is the same as `index.shape`. If `index` is None, \ + the output tensor will be an empty tensor with the same shape as `index`. + }]; + + let arguments = ( + ins TT_Ptr:$src, + TT_Tensor:$index, + TT_Int:$indexBoundary, + TT_Int:$dim, + Variadic>:$srcStride, + Variadic>:$endOffset, + Variadic>:$startOffset, + Optional:$other + ); + + let results = (outs TT_Tensor:$result); + + let assemblyFormat = [{ + $src `:` type($src) `,` $index `:` type($index) `,` + $indexBoundary `:` type($indexBoundary) `,` $dim `:` type($dim) `,` + `[` $srcStride `:` type($srcStride) `]` `,` `[` $endOffset `:` type($endOffset) `]` `,` + `[` $startOffset `:` type($startOffset) `]` (`,` $other^ `:` type($other))? + attr-dict `->` type($result) + }]; +} + +// +// ScatterUbToOut Op +// +def ScatterUbToOutOp : TT_Ascend_Op<"scatter_ub_to_out", [ + MemoryEffects<[MemWrite]>, + SameVariadicOperandSize, +]> { + let summary = "scatter store from a tensor pointer with the embedding semantics"; + + let description = [{ + Scatter a tile from Unified Buffer (UB) into a destination tensor in Global Memory (GM) + along a specified dimension, with index-boundary checking. + + The operation takes: + - ptr: pointer type, the destination tensor pointer (in GM) + - value: tensor, a tile value to store (in UB) + - index: tensor, a tile index to scatter (in UB) + - index_boundary: int, the upper boundary for index values + - dim: int, the dimension to scatter along + - dst_stride: tuple of int, the stride of each dimension of destination tensor + - end_offset: tuple of int32, the end offsets of each dimension for index tensor + - start_offset: tuple of int32, the start offsets of each dimension for index tensor + + Constraints: + - `ptr` and `index` must have the same rank. + - `ptr.dtype` only supports `float16`, `bfloat16`, `float32` currently. + - `index` must be an integer tensor, with rank between 1 and 5. + - `dim` must be valid (0 <= dim < rank(index)). + - For every dimension `i` not equal to `dim`, `index.size[i]` <= `ptr.size[i]`. + - The output shape is the same as `index.shape`. If `index` is None, \ + the output tensor will be an empty tensor with the same shape as `index`. + }]; + + let arguments = ( + ins TT_Ptr:$ptr, + TT_Tensor:$value, + TT_Tensor:$index, + TT_Int:$indexBoundary, + TT_Int:$dim, + Variadic>:$dstStride, + Variadic>:$endOffset, + Variadic>:$startOffset + ); + + let assemblyFormat = [{ + $ptr `:` type($ptr) `,` $value `:` type($value) `,` `,` $index `:` type($index) `,` + $indexBoundary `:` type($indexBoundary) `,` $dim `:` type($dim) `,` + `[` $dstStride `:` type($dstStride) `]` `,` `[` $endOffset `:` type($endOffset) `]` `,` + `[` $startOffset `:` type($startOffset) `]` + attr-dict + }]; +} + + +// +// IndexSelectSimd Op +// +def IndexSelectSimdOp : TT_Ascend_Op<"index_select_simd", [ + MemoryEffects<[MemRead]>, + DeclareOpInterfaceMethods, + AttrSizedOperandSegments +]> { + let summary = "Index select SIMD operation from global memory"; + + let description = [{ + Index select operation (SIMD version) that loads data from multiple indices along a + specified dimension. The operation selects data from GM and loads them + as tiles directly to UB with zero-copy semantics. + + The operation takes: + - src: Source pointer (in GM) + - index: 1D tensor of indices to select (already in UB) + - dim: The dimension along which to select + - src_shape: Complete shape of the source tensor + - src_offset: Starting offset for reading + - read_shape: Size to read (tile shape) + + Constraints: + - read_shape[dim] must be -1 + - src_offset[dim] can be -1 (will be ignored) + }]; + + let arguments = ( + ins + TT_PtrLike:$src, + TT_IntTensor:$index, + I32Attr:$dim, + Variadic:$src_shape, + Variadic:$src_offset, + DenseI32ArrayAttr:$read_shape + ); + + let results = (outs TT_Tensor:$result); + + let assemblyFormat = [{ + $src `,` $index `,` $dim `,` + `[` $src_shape `]` `,` + `[` $src_offset `]` `,` + $read_shape + attr-dict `:` type($src) `,` type($index) `->` type($result) + }]; +} + + +// +// Built-in: IndirectLoad Op +// +def IndirectLoadOp : TT_Ascend_Op<"indirect_load", [ + DeclareOpInterfaceMethods, + AttrSizedOperandSegments +]> { + let summary = "Built-in: indirect load from global memory using per-element offsets with optional mask/other"; + + let description = [{ + Built-in operation emitted by the compiler for unstructured (discrete) memory + accesses.These are not written directly in the user IR. + + Load values from global memory based on per-element offsets. If `mask` + is provided, false lanes return `other`. + + The operation takes: + - src: Source pointer + - offsets: Tensor of per-element offsets (relative to `src`) for accessing source memory + - mask (optional): if mask[idx] is false, do not load the data at address pointer[idx] + - other (optional): if mask[idx] is false, return other[idx] + }]; + + let arguments = ( + ins TT_Ptr:$src, + TT_IntTensor:$offsets, + Optional:$mask, + Optional:$other + ); + + let results = (outs TT_Tensor:$result); + + let assemblyFormat = [{ + $src `:` type($src) `,` + $offsets `:` type($offsets) + (`,` $mask^ `:` type($mask))? + (`,` $other^ `:` type($other))? + attr-dict `->` type($result) + }]; + + let builders = [ + OpBuilder<(ins + "Value":$src, + "Value":$offsets, + "Value":$mask, + "Value":$other + )> + ]; +} + + +// +// Built-in: IndirectStore Op +// +def IndirectStoreOp : TT_Ascend_Op<"indirect_store", [ + MemoryEffects<[MemWrite]> +]> { + let summary = "Built-in: indirect store from UB using per-element offsets with optional mask/other"; + + let description = [{ + Built-in operation emitted by the compiler for unstructured (discrete) memory + accesses.These are not written directly in the user IR. + + Store values from UB based to GM on per-element offsets. + + The operation takes: + - src: Source pointer + - offsets: Tensor of per-element offsets (relative to `src`) for accessing source memory + - value: The tensor of elements to be stored + - mask (optional): If mask[idx] is false, do not store value[idx] at pointer[idx] + }]; + + let arguments = ( + ins TT_Ptr:$src, + TT_IntTensor:$offsets, + TT_Type:$value, + Optional:$mask + ); + + let assemblyFormat = [{ + $src `:` type($src) `,` + $offsets `:` type($offsets) `,` + $value `:` type($value) + (`,` $mask^ `:` type($mask))? + attr-dict + }]; + +} + +// +// Custom Op +// +def CustomOp : TT_Ascend_Op<"custom", [Pure, MemoryEffects<[MemWrite]>]> { + let summary = "self-defined custom operation"; + let description = [{ + `ascend.custom` triton custom op is designed to pass self-defined custom operation. + + Example: + ```ascend.custom {str_args = ["sync_block_wait", "cube"]} + ``` + }]; + let arguments = (ins StrAttr:$op_name, ArrayAttr:$str_args, Variadic:$args); + + let assemblyFormat = "$op_name attr-dict ($args^ `:` type($args))?"; +} + +def FlipOp : TT_Ascend_Op<"flip", [ + NoMemoryEffect, + DeclareOpInterfaceMethods +]> { + let summary = "Reverse a tensor along a given dimension"; + let description = [{ + Reverses the elements of the input tensor along the specified dimension. + The output tensor has the same shape and element type as the input. + }]; + + let arguments = (ins + TT_Tensor:$src, // Input tensor + I64Attr:$dim // Dimension to flip along + ); + + let results = (outs + TT_Tensor:$flipped // Flipped values + ); + + let assemblyFormat = + "$src `,` $dim attr-dict `:` type($src) `->` type($flipped)"; +} + +def SortOp : TT_Ascend_Op<"sort", [ + NoMemoryEffect, + DeclareOpInterfaceMethods +]> { + let summary = "Sorts a tensor along a given dimension and returns sorted values."; + let description = [{ + Sorts the elements of the input tensor along the specified dimension. + Returns one tensor: + The sorted tensor (same shape and element type as input). + }]; + + let arguments = (ins + TT_Tensor:$src, // Input tensor + I64Attr:$dim, // Dimension to sort along + BoolAttr:$descending // Sort order + ); + + let results = (outs + TT_Tensor:$sorted // Sorted values + ); + + let assemblyFormat = "$src `,` $dim `,` $descending attr-dict `:` type($src) `->` type($sorted)"; +} + +#endif // TRITON_ASCEND_OPS diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Analysis/MaskAnalysis.h b/third_party/wafer/third_party/flir/include/triton-shared/Analysis/MaskAnalysis.h new file mode 100755 index 00000000..fefea7a0 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Analysis/MaskAnalysis.h @@ -0,0 +1,238 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_ANALYSIS_MASKANALYSIS_H +#define TRITON_ANALYSIS_MASKANALYSIS_H + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" + + +#include "mlir/Support/LogicalResult.h" +#include "mlir/IR/OpDefinition.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "llvm/Support/LogicalResult.h" + +#include "llvm/ADT/SmallVector.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + + +#include + +namespace mlir { + +class OpBuilder; + +namespace triton { + +struct dimInfo { + OpFoldResult div; + OpFoldResult shape; + bool isSlt; + Value rhs; // rhs value of CmpIOp + bool isRealDim; + int64_t dim; + dimInfo() : isSlt(false), isRealDim(false), dim(0) {} + dimInfo(OpFoldResult div, OpFoldResult shape, int64_t dim = 0) : div(div), shape(shape), isSlt(false), isRealDim(true), dim(dim) {} + bool operator==(const dimInfo &other) const { + auto staticDiv = getIntAttr(div); + auto staticShape = getIntAttr(shape); + auto otherShape = getIntAttr(other.shape); + auto otherDiv = getIntAttr(other.div); + assert(staticDiv.has_value() && staticDiv.has_value() && otherDiv.has_value() && otherShape.has_value() && "MaskAnalysis: do not support dynamic shape/div"); + return staticDiv.value() == otherDiv.value() && staticShape.value() == otherShape.value() && dim == other.dim; + } + void dump() const ; + bool hasModulo() const { + auto intAttr = getIntAttr(shape); + if (!intAttr.has_value()) { + return false; + } + return intAttr.value() != 0; + }; + bool hasDivision() const { + auto intAttr = getIntAttr(div); + if (!intAttr.has_value()) { + return false; + } + return intAttr.value() != 0; + }; +}; +// Data structure used to decode the pattern in a mask used for load and store. +// start and end field represent the start and end index of a range (produced +// by make_range, addi, etc.). While multi-dimensional data is possible, we +// assume range comparison can only be done on 1 dimension at a time (and +// results of range comparions across dimensions can be combined), hence start +// and end are not vectors. dims represents the real access size for ld/st +// (instead of the tensor/memref size specified by the IR). scalar is a shortcut +// used when the entire state contains a single scalar value. +// +// The general lifetime of this data structure is roughly: +// 1. A range is created by make_range and optionally operated on by addi w/ +// result of splat, expand_dims, etc. During this phase, either (1) both start +// and end are populated, or (2) scalar is populated. Only one of the dimensions +// (that contains the range) can have dim > 1. +// 2. Result from step 1 is compared with a another MaskState that represents a +// scalar value. The resulting state only has dims populated. +// 3. Optionally, result from step 2 can be broadcasted and anded with other +// results from step 2. The resulting state only has dims populated. +// +// Example of creating 2D mask: +// mask = (rows[:, None] < M) & (cols[None, :] < N) +struct MaskState { + OpFoldResult start; + OpFoldResult end; + SmallVector dims; + OpFoldResult scalar; + const bool useUnsafeMask; + ///ASCEND + SmallVector stateInfo; + + void dump() const; + + MaskState(bool useUnsafeMask = false) : useUnsafeMask(useUnsafeMask) {} + + int64_t getRank() const { return dims.size(); } + + bool isEmpty() const { return getRank() == 0 && !scalar && !start && !end; } + + bool isMask() const { return !start && !end && !scalar && dims.size() != 0; } + // TODO(FLIR): should be isMask() + bool isMaskWithoutScalar() const { return !start && !end && dims.size() != 0; } + + // Recursively parse a Value; call the coresponding function based on the + // defining operation and Value type + LogicalResult parse(Value operand, const Location loc, OpBuilder &builder); + + tensor::ExtractSliceOp getExtractSlice(Value source, const Location loc, + OpBuilder &builder) const; + + memref::SubViewOp getSubview(Value source, const Location loc, + OpBuilder &builder) const; + + std::pair + getSideBySideSubviews(Value block1, Value block2, const Location loc, + OpBuilder &builder) const; + + std::pair + getStackedSubviews(Value block1, Value block2, const Location loc, + OpBuilder &builder) const; + ////ASCEND + tensor::InsertSliceOp getInsertSlice(Value source, Value dest, + const Location &loc, + OpBuilder &builder) const; + + + void eraseInsertedOps(Operation *rawOp, PatternRewriter &rewriter); + +private: + // ------- + // Utility functions to operate on MaskState + // ------- + LogicalResult addStateScalar(const MaskState &state, + const OpFoldResult scalar, Location loc, + OpBuilder &builder); + + LogicalResult addStates(const MaskState &lhsState, const MaskState &rhsState, + Location loc, OpBuilder &builder); + + LogicalResult subStateScalar(const MaskState &state, + const OpFoldResult scalar, Location loc, + OpBuilder &builder); + + LogicalResult subStates(const MaskState &lhsState, const MaskState &rhsState, + Location loc, OpBuilder &builder); + + LogicalResult minStateScalar(const MaskState &lhsState, const MaskState &rhsState, + Location loc, OpBuilder &builder); + + LogicalResult minStates(const MaskState &lhsState, const MaskState &rhsState, + Location loc, OpBuilder &builder); + // ------- + // Helper functions to parse values to populate MaskState + // ------- + + LogicalResult parseExtSI(arith::ExtSIOp op, const Location loc, + OpBuilder &builder); + + // Operand is the result of a constant + // Get the value of the constant and assign it to scalar. + LogicalResult parseConstant(arith::ConstantOp constOp, const Location loc, + OpBuilder &builder); + + // Operand is an integer scalar + LogicalResult parseIntScalar(Value scalar, const Location loc, + OpBuilder &builder); + + // Operand is the result of addi + // One and only one of the operands should be a scalar. Increment both start + // and end, dims remains unchanged, and scalar is empty. + LogicalResult parseAdd(arith::AddIOp addOp, const Location loc, + OpBuilder &builder); + // Operand is the result of subi + // One and only one of the operands should be a scalar. Decrement both start + // and end, dims remains unchanged, and scalar is empty. + LogicalResult parseSub(arith::SubIOp subOp, const Location loc, + OpBuilder &builder); + // Operand is the result of andi + // Each of the result state dims is smaller of the two operands' dims. + // Insert instruction if needed to get new dims. + LogicalResult parseAnd(arith::AndIOp andOp, const Location loc, + OpBuilder &builder); + + // Operand is the result of cmpi + // Assume only one of the dimensions has size > 1. Only support slt/ult, and + // sge against 0 for now. For that dimension, we have three cases: + // 1. Constant comparison with both left and right-hand sides being scalars. + // Calculate this new dim as a compare and select. + // I.e. dim = lhs < rhs ? end : 0 + // 2. Left-hand side is not a scalar, and the right-hand side is. + // 2.a. Predicate is slt/ult. Calculate this new dim as: + // dim = max(min(end, value), start) - start + // 2.b. Predicate is sge against 0. Mask analysis already has an + // assumption that the mask starts at 0, so evaluate this to true + // and calculate this new dim as: dim = end + LogicalResult parseCmp(arith::CmpIOp cmpOp, const Location loc, + OpBuilder &builder); + // Operand is the result of make_range + // Set start and end accordingly; step size must be 1. + LogicalResult parseMakeRange(triton::MakeRangeOp rangeOp, const Location loc, + OpBuilder &builder); + // Operand is the result of broadcast + // Change dims only; assume only applies to tensors. + LogicalResult parseBroadcast(triton::BroadcastOp broadcastOp, + const Location loc, OpBuilder &builder); + // Operand is the result of splat + // Assume only applies to scalar. start and end are left empty; scalar will + // be assigned, and dims will be updated. + LogicalResult parseSplat(triton::SplatOp splatOp, const Location loc, + OpBuilder &builder); + // Operand is the result of expand_dims + // Insert additional dims; start and end do not change and correspond to the + // dimension that contains the range. + LogicalResult parseExpandDims(triton::ExpandDimsOp expandDimsOp, + const Location loc, OpBuilder &builder); +///////////////////ASCEND + LogicalResult parseRemsi(arith::RemSIOp remsiOp, + const Location loc, + OpBuilder &builder); + + LogicalResult parseDivsi(arith::DivSIOp divsiOp, + const Location loc, + OpBuilder &builder) ; + + LogicalResult parseLoopIterArg(Value v, const Location loc, + OpBuilder &builder); +}; + +} // namespace triton + +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Analysis/OpFoldResultUtils.h b/third_party/wafer/third_party/flir/include/triton-shared/Analysis/OpFoldResultUtils.h new file mode 100755 index 00000000..b1517224 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Analysis/OpFoldResultUtils.h @@ -0,0 +1,77 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_ANALYSIS_OPFOLDRESULT_UTILS_H +#define TRITON_ANALYSIS_OPFOLDRESULT_UTILS_H + +#include "mlir/IR/Location.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/Dialect/Arith/IR/Arith.h" + +#include + +namespace mlir { + +class OpBuilder; +Value materializeValue(OpBuilder &builder, Location loc, OpFoldResult ofr); + +// Return integer if ofr is an IntegerAttr. Note that this function differs +// from getConstantIntValue, which returns an integer if ofr is the constant +// result of an operation too. +std::optional getIntAttr(const OpFoldResult ofr); + +// Return if ofr contains a constant zero, either represented by an integer +// attribute or a constant value. +bool hasConstZero(const OpFoldResult ofr); + +// Create a value of index type if necessary from an OpFoldResult. +Value ofrToIndexValue(const OpFoldResult ofr, const Location loc, OpBuilder &b); + +// Create a vector of values of index type if necessary from an array of +// OpFoldResults. +SmallVector ofrsToIndexValues(ArrayRef ofrs, + const Location loc, OpBuilder &b); + +// Process addition of two OFRs. If both OFRs are Integer Attributes, result +// is an Integer Attribute. Otherwise, insert the arith.addi instruction if +// needed and use its result Value. +OpFoldResult addOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b); + +// Produce result = lhs - rhs. If both OFRs are Integer Attributes, result +// is an Integer Attribute. Otherwise, insert the arith.addi instruction if +// needed and use its result Value. +OpFoldResult subOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b); + +// Process multiplication of two OFRs. If both OFRs are Integer Attributes, +// result is an Integer Attribtue. Otherwise, insert the arith.muli +// instruction if needed and use its result Value. +OpFoldResult mulOFRValue(const OpFoldResult lhs, const Value rhs, + const Location loc, OpBuilder &b); +////////////ASCEND +OpFoldResult divOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b); + +OpFoldResult remOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b); +OpFoldResult mulOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b); + +/////////////ASCEND +OpFoldResult minOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b); + +OpFoldResult maxOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b); + +OpFoldResult compareOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const arith::CmpIPredicate pred, const OpFoldResult trueVal, + const OpFoldResult falseVal, const Location loc, OpBuilder &b); +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Analysis/PtrAnalysis.h b/third_party/wafer/third_party/flir/include/triton-shared/Analysis/PtrAnalysis.h new file mode 100755 index 00000000..5a95ebda --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Analysis/PtrAnalysis.h @@ -0,0 +1,271 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_ANALYSIS_PTRANALYSIS_H +#define TRITON_ANALYSIS_PTRANALYSIS_H + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include + +namespace mlir { + +class ConversionPatternRewriter; + +namespace triton { + +struct ModuloState { + Value size; + + // offset is used to determine the wraparound point for patterns like: + // offset + (tl.arange(0, 256) % 12) + // The current code assumes that the modulo operator always runs last, e.g: + // (offset + tl.arange(0, 256)) % 12 + // This is not used at the moment as there haven't been enough use cases and + // the implementation is quite complex. + // OpFoldResult offset; + + static constexpr char const *WraparoundAttr = "ptr.wraparound_type"; + static constexpr char const *WraparoundStacked = "stacked"; + static constexpr char const *WraparoundSideBySide = "side_by_side"; +}; + +// Data structure used to decode pointer arithmetics and potentially to be +// translate it into memref. offsets, sizes, and strides are in unit of elements +// in a linearly laid-out memory, which is the same as pointer arithmetic +// operations in Triton language. scalar is a shortcut used when the entire +// state describes a single scalar value. source is the base pointer. +class PtrState { + + OpFoldResult + accumulateTargetOffset(Location loc, + ConversionPatternRewriter &rewriter) const; + +public: + SmallVector offsets; + SmallVector sizes; + SmallVector strides; + + SmallVector> modulos; + + Value source; + Value scalar; + + int64_t getRank() const; + + bool isEmpty() const; + + bool hasModulo() const; + + MemRefType getResultMemrefType(MLIRContext *context, int64_t offset, + ArrayRef resultShape, + bool useDynamicStrides = false) const; + + // Process addition of two PtrStates. + void addState(const PtrState &lhsState, const PtrState &rhsState, + Location loc, ConversionPatternRewriter &rewriter); + + // Process multiplication of two PtrStates + void mulState(const PtrState &lhsState, const PtrState &rhsState, + const Location loc, ConversionPatternRewriter &rewriter); + + // Produce a reinterpret cast based on the current PtrState. Additional + // instructions may be inserted in calculating the final offset. + memref::ReinterpretCastOp + createCastOp(ArrayRef resultShape, const Location loc, + ConversionPatternRewriter &rewriter) const; + + SmallVector + createSideBySideCastOps(ArrayRef resultShape, const Location loc, + ConversionPatternRewriter &rewriter) const; + + SmallVector + createStackedCastOps(ArrayRef resultShape, const Location loc, + ConversionPatternRewriter &rewriter) const; +}; + +class PtrAnalysis { +public: + using IndexMapSet = std::map>; + + // Recursively parse a Value; call the corresponding + // function based on the defining operation and argument type. + static void + visitOperand(Value operand, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of arith.addi. Process both arguments and insert any + // arith.addi instruction as needed. + // Main assumptions: + // Only one of lhsState and rhsState has source field set + // Current PtrState should be empty + // Expected result: + // source = lhsState.source ? lhsState.source : rhsState.source + // sizes[i] = lhsState.sizes[i] (which should match rhsState.sizes[i]) + // offsets[i] = lhsState.offsets[i] + rhsState.offsets[i] + // strides[i] = lhsState.strides[i] + rhsState.strides[i] + static void + visitOperandAdd(arith::AddIOp addOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of arith.muli. Process both arguments and insert any + // arith.muli instruction as needed. + // Main assumptions: + // Neither lhsState nor rhsState has source field set + // Current PtrState should be empty + // Currently only support one of the operand is a scalar index + // Expected result (scalar and tensorState represent the two operands): + // source = null + // sizes[i] = tensorState.sizes[i] + // offsets[i] = tensorState.offsets[i] * scalar + // strides[i] = tensorState.strides[i] * scalar + static void + visitOperandMul(arith::MulIOp mulOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + static void + visitOperandRem(arith::RemSIOp mulOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + static void visitOperandUnrealizedCast( + UnrealizedConversionCastOp op, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of make_range. + // Main assumptions: + // start, end, and shape are all statically known + // The output of make_range is 1-dimensional + // Does not check validity of inputs (e.g., stride > 0) + // Expected result: + // source = null + // sizes[0] = shape[0] + // offset[0] = start + // strides[0] = ceiling( (end - start) / shape[0] ) + static void + visitOperandMakeRange(triton::MakeRangeOp rangeOp, PtrState &state, + Location loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of expand_dims + // Main assumptions: + // Only 1 dimension changes for each invocation of reshape + // The changed dimension must have size of 1 + // Expected result: + // Insert a dimension of size 1, stride 0, and offset 0 + static void + visitOperandExpandDims(triton::ExpandDimsOp expandDimsOp, PtrState &state, + const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of broadcast + // Main assumptions: + // Rank of soure and result is the same + // Expected result: + // Update sizes[i] only, no changes to other fields + static void + visitOperandBroadcast(triton::BroadcastOp broadcastOp, PtrState &state, + const Location loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of splat + // Main assumptions: + // Source is a scalar value (i.e., an integer or a pointer, not a tensor) + // Expected result: + // sizes[i] reflect the shape of the result, strides[i] = 0, offsets[i] = 0 + // if source is an integer, offset[0] = scalar = source + static void + visitOperandSplat(triton::SplatOp splatOp, PtrState &state, + const Location loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of arith.constant that is a splat + // Main assumptions: + // Source is a constant op that produces a constant dense tensor where all + // elements are the same (i.e.: a constant that is splatted) + // Expected result: + // sizes[i] reflect the shape of the result, strides[i] = 0, offsets[i] = + // splat value if i == 0, otherwise 0 + static void + visitOperandConstSplat(arith::ConstantOp op, PtrState &state, + const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + static void visitOperandMakeTensorPtr( + triton::MakeTensorPtrOp makeTensorPtrOp, PtrState &state, + const Location loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of addptr. + // Main assumptions: + // The ptr field should populate the source field + // ptr and offset fields should result in same rank + // Expected result: + // The resulting state for ptr and offset wil be added + static void + visitOperandAddptr(triton::AddPtrOp addptrOp, PtrState &state, + const Location loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of reinterpret_cast. + // Main assumptions: + // None + // Expected result: + // Directly grab all corresponding fields from reinterpret_cast. + static void + visitOperandReintCast(memref::ReinterpretCastOp reintCastOp, PtrState &state, + const Location loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs); + + // Operand is the result of tt.advance. + // Main assumptions: + // The source of the tt.advance has been mapped to a reinterpret_cast + // Expected result: + // Directly grab all corresponding fields from reinterpret_cast. + // Add the offsets multiplied by the strides to the final offsets. + static void rewriteAdvanceOp(triton::AdvanceOp op, + ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &knownPtrs); + + // Parse the state of AddPtrOp, insert any instruction needed to + // calculate strides and offsets, build PtrState for this operand, and record + // PtrState for knownPtrs. + static void rewriteAddptrOp(triton::AddPtrOp op, + ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &knownPtrs); + + // Parse the state of YieldOp, insert any instruction needed to calculate + // strides and offsets, build PtrState for this operand, and record PtrState + // in knownPtrs. + static void + rewriteYieldOp(scf::YieldOp op, ConversionPatternRewriter &rewriter, + const IndexMapSet &levelToBlockArgIndex, const int level, + const llvm::SmallDenseMap &knownPtrs); + + static void rewriteForOp(scf::ForOp op, ConversionPatternRewriter &rewriter, + IndexMapSet &levelToBlockArgIndex, const int level, + llvm::SmallDenseMap &knownPtrs); + + static Value getScalarMemRef(Value ptr, Value memRef, const Location loc, + ConversionPatternRewriter &rewriter); +}; + +} // namespace triton + +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Analysis/UseAnalysis.h b/third_party/wafer/third_party/flir/include/triton-shared/Analysis/UseAnalysis.h new file mode 100755 index 00000000..39c3055a --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Analysis/UseAnalysis.h @@ -0,0 +1,119 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_ANALYSIS_USEANALYSIS_H +#define TRITON_ANALYSIS_USEANALYSIS_H + +#include "mlir/Analysis/DataFlow/SparseAnalysis.h" +#include "mlir/Pass/Pass.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +std::unique_ptr createTritonUseAnalysisPass(); + +enum class UseType { + Undefined, // Initial state + DataUse, // value used for tensor computation only + MetaUse, // value used for metadata only + MixUse // value used for both tensor computation and metadata +}; + +struct UseInfo : public dataflow::AbstractSparseLattice { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(UseInfo) + using AbstractSparseLattice::AbstractSparseLattice; + + // Lattice state transfer function + ChangeResult meetUseType(const UseType &other) { + if (other == UseType::Undefined) + return ChangeResult::NoChange; + + switch (type) { + case UseType::Undefined: + type = other; + return ChangeResult::Change; + case UseType::DataUse: + case UseType::MetaUse: + if (type == other) { + return ChangeResult::NoChange; + } else { + type = UseType::MixUse; + return ChangeResult::Change; + } + case UseType::MixUse: + return ChangeResult::NoChange; + default: + llvm_unreachable("bad type"); + } + } + + ChangeResult meet(const AbstractSparseLattice &other) override { + auto rhs = reinterpret_cast(&other); + return meetUseType(rhs->type); + } + + void print(raw_ostream &os) const override { + switch (type) { + case UseType::DataUse: + os << "DataUse"; + break; + case UseType::MetaUse: + os << "MetaUse"; + break; + case UseType::MixUse: + os << "MixUse"; + break; + default: + os << "Undefined"; + } + } + + UseType type = UseType::Undefined; +}; + +class UseAnalysis : public dataflow::SparseBackwardDataFlowAnalysis { +public: + using SparseBackwardDataFlowAnalysis::SparseBackwardDataFlowAnalysis; + LogicalResult visitOperation(Operation *op, ArrayRef operands, + ArrayRef results) override; + + void visitBranchOperand(OpOperand &operand) override { return; } + + void visitCallOperand(OpOperand &operand) override { return; } + + void setToExitState(UseInfo *lattice) override { + lattice->type = UseType::Undefined; + } + +private: + void propagateUse(UseInfo *lattice, const UseType &type) { + auto changed = lattice->meetUseType(type); + propagateIfChanged(lattice, changed); + } + + void propagateResults(UseInfo *lattice, ArrayRef results) { + auto changed = ChangeResult::NoChange; + for (auto result : results) + changed |= lattice->meet(*result); + propagateIfChanged(lattice, changed); + } +}; + +// Use SparseBackwardDataAnalysis to identify operations whose results are used +// as data tensor operations, meta operations (address calculation, +// broadcasting/splating constant, etc.), or both. For operations used as both +// purposes, clone them so that the remaining pass built on +// ConversionPatternRewriter can replace all tensor producers cleanly and simply +// delete meta data producers. +LogicalResult runUseAnalysis(triton::FuncOp &funcOp); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONTOAFFINE_TRITONUSEANALYSIS_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/AnalysisStructured/PtrAnalysis.h b/third_party/wafer/third_party/flir/include/triton-shared/AnalysisStructured/PtrAnalysis.h new file mode 100755 index 00000000..f9c9b60b --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/AnalysisStructured/PtrAnalysis.h @@ -0,0 +1,312 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_ANALYSISSTRUCTURED_PTRANALYSIS_H +#define TRITON_ANALYSISSTRUCTURED_PTRANALYSIS_H + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" + +#include "mlir/IR/Value.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Support/LogicalResult.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include +#include + +namespace mlir { + +class OpBuilder; + +namespace tts { + +const extern std::string ptrAnalysisAttr; + +// Data structure used to decode pointer arithmetics. offsets, sizes, and +// strides are in unit of elements in a linearly laid-out memory, which is the +// same as pointer arithmetic operations in Triton language. scalar is a +// shortcut used when the entire state describes a single scalar value. source +// is the base pointer. If order is present, PtrState describes block pointer; +// otherwise it describes non-block pointers. When it describes block pointer, +// shape field means the same field as tt.make_tensor_ptr; when it describes a +// non-block pointer, shape field indicates how address wraps around (i.e., +// modulo); a constant 0 indicates no modulo for the dimension. +struct PtrState { + + SmallVector offsets; + SmallVector sizes; + SmallVector strides; + SmallVector shape; + SmallVector order; + + Value source; + Value scalar; + + int32_t getRank() const; + + bool isEmpty() const; + + bool hasModulo() const; + + bool dimHasModulo(uint32_t dim) const; + + bool isBlockPtr() const; + + void dump() const; + + // Process addition of two PtrStates. + LogicalResult addState(const PtrState &lhsState, const PtrState &rhsState, + Operation *op, OpBuilder &builder); + + // Process subtraction of two PtrStates + LogicalResult subState(const PtrState &lhsState, const PtrState &rhsState, + Operation *op, OpBuilder &builder); + + // Process multiplication of two PtrStates + LogicalResult mulState(const PtrState &lhsState, const PtrState &rhsState, + Operation *op, OpBuilder &builder); + + tts::MakeTensorPtrOp createTTSMakeTensorPtrOp(OpBuilder &builder, + Location loc); +}; + +class PtrAnalysis { + // This function is internally used by getLoopIterArgPtrState and + // getLoopResultPtrState to get the correct PtrState for either an iter-arg or + // a loop's result. + // + // A PtrState of an scf.for's iter-arg is the same as its corresponding + // init-arg, except that the strides and offsets have to point to the loop's + // iter-args that were created to carry the offsets and strides. + // + // For instance, for a pointer with index i and rank 2, 4 additional args + // starting at index i + 1 are created. The PtrState's strides and offsets + // value of the pointer's iter-arg must point to these 4 additionally created + // iter-args. + // + // A similar process is used for getting the PtrState of the loop's i'th + // result: its strides and offsets have to point to the corresponding stride + // and offset values returned by the loop. + PtrState reconcileLoopPtrState( + scf::ForOp forOp, size_t ptrArgIndex, const PtrState &state, + llvm::function_ref getReplacementVal); + + DenseSet maybeStructuredArgs; + +public: + void initializeMaybeStructuredArgs(Operation *op); + + llvm::SmallDenseMap knownPtrs; + + IRMapping ptrMap; + + // Recursively parse a Value; call the corresponding + // function based on the defining operation and argument type. + LogicalResult visitOperand(Value operand, PtrState &state, const Location loc, + OpBuilder &builder); + + // Operand is a result of an scf.for. Such cases occur when there are multiple + // levels of nested loops where the results of the inner scf.for (pointer) are + // yielded by the outer loop. + LogicalResult visitOperandForOp(scf::ForOp forOp, Value operand, + PtrState &state, const Location loc, + OpBuilder &builder); + + // Operand is the result of arith.addi. Process both arguments and insert any + // arith.addi instruction as needed. + // Main assumptions: + // Only one of lhsState and rhsState has source field set + // Current PtrState should be empty + // Expected result: + // source = lhsState.source ? lhsState.source : rhsState.source + // sizes[i] = lhsState.sizes[i] (which should match rhsState.sizes[i]) + // offsets[i] = lhsState.offsets[i] + rhsState.offsets[i] + // strides[i] = lhsState.strides[i] + rhsState.strides[i] + LogicalResult visitOperandAdd(arith::AddIOp addOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Operand is the result of arith.subi. Process both arguments and insert any + // arith.subi instruction as needed. + // Main assumptions: + // Only one of lhsState and rhsState has source field set + // Current PtrState should be empty + // Expected result: + // source = lhsState.source ? lhsState.source : rhsState.source + // sizes[i] = lhsState.sizes[i] (which should match rhsState.sizes[i]) + // offsets[i] = lhsState.offsets[i] - rhsState.offsets[i] + // strides[i] = lhsState.strides[i] - rhsState.strides[i] + LogicalResult visitOperandSub(arith::SubIOp subOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Operand is the result of arith.muli. Process both arguments and insert any + // arith.muli instruction as needed. + // Main assumptions: + // Neither lhsState nor rhsState has source field set + // Current PtrState should be empty + // Currently only support one of the operand is a scalar index + // Expected result (scalar and tensorState represent the two operands): + // source = null + // sizes[i] = tensorState.sizes[i] + // offsets[i] = tensorState.offsets[i] * scalar + // strides[i] = tensorState.strides[i] * scalar + LogicalResult visitOperandMul(arith::MulIOp mulOp, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandRem(arith::RemSIOp mulOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Operand is the result of make_range. + // Main assumptions: + // start, end, and shape are all statically known + // The output of make_range is 1-dimensional + // Does not check validity of inputs (e.g., stride > 0) + // Expected result: + // source = null + // sizes[0] = shape[0] + // offset[0] = start + // strides[0] = ceiling( (end - start) / shape[0] ) + LogicalResult visitOperandMakeRange(triton::MakeRangeOp rangeOp, + PtrState &state, Location loc, + OpBuilder &builder); + + // Operand is the result of expand_dims + // Main assumptions: + // Only 1 dimension changes for each invocation of reshape + // The changed dimension must have size of 1 + // Expected result: + // Insert a dimension of size 1, stride 0, and offset 0 + LogicalResult visitOperandExpandDims(triton::ExpandDimsOp expandDimsOp, + PtrState &state, const Location loc, + OpBuilder &builder); + + // Operand is the result of broadcast + // Main assumptions: + // Rank of soure and result is the same + // Expected result: + // Update sizes[i] only, no changes to other fields + LogicalResult visitOperandBroadcast(triton::BroadcastOp broadcastOp, + PtrState &state, const Location loc, + OpBuilder &builder); + + // Operand is the result of splat + // Main assumptions: + // Source is a scalar value (i.e., an integer or a pointer, not a tensor) + // Expected result: + // sizes[i] reflect the shape of the result, strides[i] = 0, offsets[i] = 0 + // if source is an integer, offset[0] = scalar = source + LogicalResult visitOperandSplat(triton::SplatOp splatOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Operand is the result of arith.constant that is a splat + // Main assumptions: + // Source is a constant op that produces a constant dense tensor where all + // elements are the same (i.e.: a constant that is splatted) + // Expected result: + // sizes[i] reflect the shape of the result, strides[i] = 0, offsets[i] = + // splat value if i == 0, otherwise 0 + LogicalResult visitOperandConstSplat(arith::ConstantOp op, PtrState &state, + const Location loc, OpBuilder &builder); + + LogicalResult visitOperandExtSI(arith::ExtSIOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Operand is the result of addptr. + // Main assumptions: + // The ptr field should populate the source field + // ptr and offset fields should result in same rank + // Expected result: + // The resulting state for ptr and offset wil be added + LogicalResult visitOperandAddptr(triton::AddPtrOp addptrOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Operand is the result of tts.make_tptr. + // Main assumptions: + // This function is only called when rewriting a loop + // Expected result: + // Directly grab all corresponding fields from tts.make_tptr. + LogicalResult visitOperandMakeTPtr(tts::MakeTensorPtrOp makeTPtrOp, + PtrState &state, const Location loc, + OpBuilder &builder); + + // Operand is the result of tt.make_tensor_ptr. + // Expected result: + // Parse source pointer and grab results + LogicalResult visitOperandMakeTensorPtr(triton::MakeTensorPtrOp makeTPtrOp, + PtrState &state, const Location loc, + OpBuilder &builder); + + // Operand is the result of tt.int_to_ptr. + // Expected result: + // Directly grab op result + LogicalResult visitOperandIntToPtr(triton::IntToPtrOp intToPtrOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Operand is the result of tt.bitcast. + // Expected result: + // Directly grab op result + LogicalResult visitOperandBitcast(triton::BitcastOp bitcastOp, PtrState &state, + const Location loc, OpBuilder &builder); + + // Get the computed PtrState for the forOp's init-arg at the provided index. + FailureOr getLoopInitArgPtrState(scf::ForOp forOp, size_t index); + + // Get the computed PtrState for the forOp's iter-arg at the provided index. + FailureOr getLoopIterArgPtrState(scf::ForOp forOp, size_t index); + + // Get the computed PtrState for the forOp's result at the provided index. + FailureOr getLoopResultPtrState(scf::ForOp forOp, size_t index); + + // After PtrAnalysis finishes, rewrite the GetStructuredStateOp by creating + // the correct initialization ops for offsets and strides and passing them to + // any loop's init-args. + LogicalResult rewriteGetStructuredStateOp(tts::GetStructuredStateOp op); + + // Parse the state of AddPtrOp, insert any instruction needed to + // calculate strides and offsets, build PtrState for this operand, and record + // PtrState for knownPtrs. + LogicalResult rewriteAddptrOp(triton::AddPtrOp op); + + // Move the BitcastOp on tensor of pointers to the source scalar pointer + // tracked by PtrAnalysis. + LogicalResult rewriteBitcastOp(triton::BitcastOp op); + + LogicalResult rewriteMakeTensorPtrOp(triton::MakeTensorPtrOp op); + + LogicalResult rewriteAdvanceOp(triton::AdvanceOp op); + + // Parse the state of YieldOp, insert any instruction needed to calculate + // strides and offsets, build PtrState for this operand, and record PtrState + // in knownPtrs. + LogicalResult + rewriteYieldOp(scf::YieldOp op, + llvm::SmallDenseMap &knownPtrsFor); + + // Rewrite eligible tt.addptr in loop init args so loop can update the such + // pointers over iterations. Insert any instruction needed to calculate + // strides, offsets, and modulos. + LogicalResult rewriteForOp(scf::ForOp op); + + LogicalResult rewriteLoadOp(triton::LoadOp op, bool useUnsafeMask = false); + + LogicalResult rewriteStoreOp(triton::StoreOp op, bool useUnsafeMask = false); + + LogicalResult rewriteAtomicRMWOp(triton::AtomicRMWOp op, + bool useUnsafeMask = false); + + LogicalResult rewriteAtomicCASOp(triton::AtomicCASOp op); + + LogicalResult rewriteOp(Operation *op, bool useUnsafeMask = false); +}; + +} // namespace tts + +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/CMakeLists.txt new file mode 100755 index 00000000..629c08af --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/CMakeLists.txt @@ -0,0 +1,2 @@ +add_subdirectory(Conversion) +add_subdirectory(Dialect) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/CMakeLists.txt new file mode 100755 index 00000000..4fab20cc --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/CMakeLists.txt @@ -0,0 +1,13 @@ +if(NOT FLAGTREE_BACKEND STREQUAL "wafer") + add_subdirectory(TritonToLinalgExperimental) + add_subdirectory(TritonArithToLinalg) + add_subdirectory(StructuredToMemref) + add_subdirectory(ReconcilePtrCasts) +endif() +add_subdirectory(TritonToLinalg) +add_subdirectory(TritonToStructured) +add_subdirectory(TritonPtrToMemref) +add_subdirectory(TritonToUnstructured) +add_subdirectory(UnstructuredToMemref) +add_subdirectory(MemrefCopyToDMA_FlagTree) +add_subdirectory(NoBufferize_FlagTree) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/CMakeLists.txt new file mode 100755 index 00000000..048ccb3d --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name MemrefCopyToDMAFlagTree) +add_public_tablegen_target(MemrefCopyToDMAFlagTreeConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.h new file mode 100755 index 00000000..50153a92 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.h @@ -0,0 +1,24 @@ +#ifndef TRITON_CONVERSION_MEMREFCOPYTODMAFLAGTREE_MEMREFCOPYTODMAFLAGTREE_H +#define TRITON_CONVERSION_MEMREFCOPYTODMAFLAGTREE_MEMREFCOPYTODMAFLAGTREE_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +class TypeConverter; +namespace triton { + +#define GEN_PASS_DECL +#include "triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.h.inc" + +void populateMemrefCopyToDMAFlagTreeConversionPatterns( + RewritePatternSet &patterns, TypeConverter &typeConverter); + +std::unique_ptr> createMemrefCopyToDMAFlagTreePass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_STRUCTUREDTOMEMREF_STRUCTUREDTOMEMREF_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.h new file mode 100755 index 00000000..cfa6c233 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_MEMREF_COPY_TO_DMA_FLAGTREE_CONVERSION_PASSES_H +#define TRITON_MEMREF_COPY_TO_DMA_FLAGTREE_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.td new file mode 100755 index 00000000..3da13b02 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.td @@ -0,0 +1,10 @@ +#ifndef MEMREF_COPY_TO_DMA_FLAGTREE_CONVERSION_PASSES +#define MEMREF_COPY_TO_DMA_FLAGTREE_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def MemrefCopyToDMAFlagTree : Pass<"memref-copy-to-dma-flagtree", "mlir::ModuleOp"> { + let summary = "Convert memrefcopy to DMA"; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/CMakeLists.txt new file mode 100755 index 00000000..91da7dda --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name NoBufferizeFlagTree) +add_public_tablegen_target(NoBufferizeFlagTreeConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.h new file mode 100755 index 00000000..2abb7d72 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.h @@ -0,0 +1,24 @@ +#ifndef TRITON_CONVERSION_NOBUFFERIZEFLAGTREE_NOBUFFERIZEFLAGTREE_H +#define TRITON_CONVERSION_NOBUFFERIZEFLAGTREE_NOBUFFERIZEFLAGTREE_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +class TypeConverter; +namespace triton { + +#define GEN_PASS_DECL +#include "triton-shared/Conversion/NoBufferize_FlagTree/Passes.h.inc" + +void populateNoBufferizeFlagTreeConversionPatterns( + RewritePatternSet &patterns, TypeConverter &typeConverter); + +std::unique_ptr> createNoBufferizeFlagTreePass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_NOBUFFERIZEFLAGTREE_NOBUFFERIZEFLAGTREE_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/Passes.h new file mode 100755 index 00000000..f8000dfd --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_NO_BUFFERIZE_FLAGTREE_CONVERSION_PASSES_H +#define TRITON_NO_BUFFERIZE_FLAGTREE_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/NoBufferize_FlagTree/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/Passes.td new file mode 100755 index 00000000..db3b66c8 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/Passes.td @@ -0,0 +1,10 @@ +#ifndef NO_BUFFERIZE_FLAGTREE_CONVERSION_PASSES +#define NO_BUFFERIZE_FLAGTREE_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def NoBufferizeFlagTree : Pass<"no-bufferize-flagtree", "mlir::ModuleOp"> { + let summary = "Add no_bufferize to memref ops in shared memory space"; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/CMakeLists.txt new file mode 100755 index 00000000..278b4906 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name ReconcilePtrCasts) +add_public_tablegen_target(ReconcilePtrCastsPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.h new file mode 100755 index 00000000..941d5b8b --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.h @@ -0,0 +1,15 @@ +#ifndef RECONCILE_PTR_CASTS_CONVERSION_PASSES_H +#define RECONCILE_PTR_CASTS_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/ReconcilePtrCasts/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.td new file mode 100755 index 00000000..d19c5e8a --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.td @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef RECONCILE_PTR_CASTS_CONVERSION_PASSES +#define RECONCILE_PTR_CASTS_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def ReconcilePtrCasts : Pass<"reconcile-ptr-casts", "mlir::ModuleOp"> { + let summary = "Convert unrealized_cast op between tt.ptr or ptr.ptr to memref to to_memref or from_memref"; + let constructor = "triton::createReconcilePtrCastsPass()"; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h new file mode 100755 index 00000000..bea24e8f --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_TRITONTOLINALG_ReconcilePtrCasts_H +#define TRITON_CONVERSION_TRITONTOLINALG_ReconcilePtrCasts_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createReconcilePtrCastsPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONTOLINALG_ReconcilePtrCasts_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/CMakeLists.txt new file mode 100755 index 00000000..83ff64d3 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name StructuredToMemref) +add_public_tablegen_target(StructuredToMemrefConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/Passes.h new file mode 100755 index 00000000..198675b1 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_STRUCTURED_TO_MEMREF_CONVERSION_PASSES_H +#define TRITON_STRUCTURED_TO_MEMREF_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/StructuredToMemref/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/Passes.td new file mode 100755 index 00000000..ac318f54 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/Passes.td @@ -0,0 +1,10 @@ +#ifndef STRUCTURED_TO_MEMREF_CONVERSION_PASSES +#define STRUCTURED_TO_MEMREF_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def StructuredToMemref : Pass<"structured-to-memref", "mlir::ModuleOp"> { + let summary = "Convert triton structured pointer ops to memref"; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h new file mode 100755 index 00000000..8c67c9ec --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h @@ -0,0 +1,24 @@ +#ifndef TRITON_CONVERSION_STRUCTUREDTOMEMREF_STRUCTUREDTOMEMREF_H +#define TRITON_CONVERSION_STRUCTUREDTOMEMREF_STRUCTUREDTOMEMREF_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +class TypeConverter; +namespace triton { + +#define GEN_PASS_DECL +#include "triton-shared/Conversion/StructuredToMemref/Passes.h.inc" + +void populateStructuredToMemrefConversionPatterns(RewritePatternSet &patterns, + TypeConverter &typeConverter); + +std::unique_ptr> createStructuredToMemrefPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_STRUCTUREDTOMEMREF_STRUCTUREDTOMEMREF_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/CMakeLists.txt new file mode 100755 index 00000000..85076bd1 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonArithToLinalg) +add_public_tablegen_target(TritonArithToLinalgConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.hpp b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.hpp new file mode 100755 index 00000000..eef98b84 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.hpp @@ -0,0 +1,2690 @@ +#ifndef TRITON_CONVERSION_PATTERNS +#define TRITON_CONVERSION_PATTERNS + +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "flagtree/Common/UnifiedHardware.h" + +#include "triton-shared/Analysis/MaskAnalysis.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/Analysis/PtrAnalysis.h" +#include "triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns_FlagTree.hpp" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" +#include "mlir-ext/Dialect/MathExt/IR/MathExt.h" +#include "triton-shared/Utils/Utils.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include +#include +#include + +using namespace mlir; +using namespace triton; + +//===----------------------------------------------------------------------===// +// Utilities +//===----------------------------------------------------------------------===// + +// Extract a scalar value from v. +// If v is a scalar, return that directly. Otherwise, parse through operations +// (currently only support splat, sitofp, and truncf) that produce it to +// extract the underlying scalar value. We then reconstruct the chain of +// operations that can produce this constant with the original type. If no +// scalar value can be extracted, a nullptr is returned. +static Value getScalarValue(Value operand, Location loc, + ConversionPatternRewriter &rewriter) { + SmallVector ops; + + auto reconstructScalarValue = [&](Value src) { + for (auto op = ops.rbegin(); op != ops.rend(); ++op) { + src = TypeSwitch(*op) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return rewriter.create(loc, resType, src); + }) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return rewriter.create(loc, resType, src); + }) + .Default([](Operation *op) { + llvm_unreachable("unsupported op in generating "); + return nullptr; + }); + } + return src; + }; + + while (true) { + if (!dyn_cast(operand.getType())) { + return reconstructScalarValue(operand); + } else if (auto op = operand.getDefiningOp()) { + if (auto attr = dyn_cast(op.getValue())) { + if (!attr.isSplat()) { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load " + "produced by unsupported instruction"; + return nullptr; + } + auto elemValue = attr.getSplatValue(); + auto constOp = arith::ConstantOp::materialize( + rewriter, elemValue, attr.getElementType(), op.getLoc()); + return reconstructScalarValue(constOp.getResult()); + } + } else if (auto op = operand.getDefiningOp()) { + operand = op.getSrc(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load produced " + "by unsupported instruction"; + return nullptr; + } + } + return nullptr; +} + +static SmallVector getNParallelLoopsAttrs(unsigned n) { + return SmallVector(n, utils::IteratorType::parallel); +} + +static Value getTransposedValue(Value source, const Location loc, + ConversionPatternRewriter &rewriter) { + + auto sourceType = cast(source.getType()); + auto sourceRank = sourceType.getRank(); + + SmallVector perm(sourceRank); + std::iota(std::begin(perm), std::end(perm), 0); + std::swap(perm[sourceRank - 1], perm[sourceRank - 2]); + + SmallVector transposedShape(sourceType.getShape()); + std::swap(transposedShape[sourceRank - 1], transposedShape[sourceRank - 2]); + + Value transposeInit = rewriter.create( + loc, transposedShape, sourceType.getElementType()); + + Value transpose = + rewriter.create(loc, source, transposeInit, perm) + .getResults()[0]; + + return transpose; +} + +// for IntLike and FloatLike types +static std::optional getBitWidth(Type a) { + if (auto type = dyn_cast(a)) { + auto elementType = type.getElementType(); + if (elementType.isIntOrFloat()) { + return type.getElementType().getIntOrFloatBitWidth(); + } + return std::nullopt; + } + + if (a.isIntOrFloat()) + return a.getIntOrFloatBitWidth(); + + return std::nullopt; +} + +//===----------------------------------------------------------------------===// +// Op Lowering Patterns +//===----------------------------------------------------------------------===// + +namespace { + +//----------------------------- +// Begin of monolithic only +//----------------------------- +struct AdvanceConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(triton::AdvanceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + llvm::SmallDenseMap knownPtrs; + PtrState pointerState; + PtrAnalysis::rewriteAdvanceOp(op, rewriter, knownPtrs); + return success(); + } +}; + +struct MakeTensorPtrConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + void populateVectorAsIndex(SmallVector &vec, + Operation::operand_range ops, + ConversionPatternRewriter &rewriter, + Location loc) const { + for (auto opnd : ops) { + if (isa(opnd.getType())) { + auto castOp = rewriter.create( + loc, rewriter.getIndexType(), opnd); + vec.push_back(castOp.getResult()); + } else { + assert(isa(opnd.getType())); + vec.push_back(opnd); + } + } + } + + LogicalResult + matchAndRewrite(triton::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + PtrState pointerState; + + auto orderSize = op.getOrder().size(); + if (orderSize > 1) { + for (auto [first, second] : + llvm::zip(op.getOrder().slice(0, orderSize - 2), + op.getOrder().slice(1, orderSize - 1))) { + assert(first == second + 1 && + "Currently only support default order on block pointers"); + } + } + + pointerState.source = rewriter.getRemappedValue(op.getBase()); + populateVectorAsIndex(pointerState.offsets, op.getOffsets(), rewriter, loc); + populateVectorAsIndex(pointerState.strides, op.getStrides(), rewriter, loc); + + SmallVector newOffsets; + for (auto [offset, stride] : + llvm::zip(pointerState.offsets, pointerState.strides)) { + auto mulOp = rewriter.create(loc, cast(offset), + cast(stride)); + newOffsets.push_back(mulOp.getResult()); + } + + pointerState.offsets.clear(); + + for (auto offset : newOffsets) { + pointerState.offsets.push_back(offset); + } + + ArrayRef resultShape; + auto pointerType = + cast(op.getResult().getType()); + if (auto shapedType = dyn_cast(pointerType.getPointeeType())) { + resultShape = shapedType.getShape(); + for (auto dim_size : resultShape) { + pointerState.sizes.push_back( + IntegerAttr::get(IntegerType::get(op.getContext(), 64), dim_size)); + } + } else { + // scalar pointer, should produce a one dimensional memref + SmallVector scalarShape(1, 1); + resultShape = scalarShape; + assert(pointerState.getRank() == 1); + } + + auto castOp = pointerState.createCastOp(resultShape, loc, rewriter); + rewriter.replaceOp(op, castOp.getResult()); + return success(); + } +}; + +struct LegacyAddPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::AddPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + llvm::SmallDenseMap knownPtrs; + PtrAnalysis::rewriteAddptrOp(op, rewriter, knownPtrs); + return success(); + } +}; + +struct LoadConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + void createSideBySideCopies(Value block1, Value block2, Value dst, + Location loc, + ConversionPatternRewriter &rewriter) const { + + auto zero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + + auto one = + rewriter.create(loc, rewriter.getIndexAttr(1)); + + Value block1Row = rewriter.create(loc, block1, 0); + Value block1Col = rewriter.create(loc, block1, 1); + + Value block2Row = rewriter.create(loc, block2, 0); + Value block2Col = rewriter.create(loc, block2, 1); + + auto block1Dst = + rewriter.create(loc, dst, /* offsets */ + ValueRange{zero, zero}, + /* sizes */ + ValueRange{block1Row, block1Col}, + /* strides */ + ValueRange{one, one}); + + auto block2Dst = + rewriter.create(loc, dst, + /* offsets */ + ValueRange{zero, block1Col}, + /* sizes */ + ValueRange{block2Row, block2Col}, + /* strides */ + ValueRange{one, one}); + + rewriter.create(loc, block1, block1Dst); + rewriter.create(loc, block2, block2Dst); + } + + void createStackedCopies(Value block1, Value block2, Value dst, Location loc, + ConversionPatternRewriter &rewriter) const { + + auto zero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + auto one = + rewriter.create(loc, rewriter.getIndexAttr(1)); + + Value block1Row = rewriter.create(loc, block1, 0); + Value block1Col = rewriter.create(loc, block1, 1); + + Value block2Row = rewriter.create(loc, block2, 0); + Value block2Col = rewriter.create(loc, block2, 1); + + auto block1Dst = + rewriter.create(loc, dst, /* offsets */ + ValueRange{zero, zero}, + /* sizes */ + ValueRange{block1Row, block1Col}, + /* strides */ + ValueRange{one, one}); + + auto block2Dst = + rewriter.create(loc, dst, + /* offsets */ + ValueRange{block1Row, zero}, + /* sizes */ + ValueRange{block2Row, block2Col}, + /* strides */ + ValueRange{one, one}); + + rewriter.create(loc, block1, block1Dst); + rewriter.create(loc, block2, block2Dst); + } + +public: + LogicalResult + matchAndRewrite(triton::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto ptr = adaptor.getPtr(); + auto mask = op.getMask(); + auto other = op.getOther(); + auto loc = op.getLoc(); + + // 0. Shortcut for scalar loads + if (!isa(op.getResult().getType())) { + auto sMemRef = PtrAnalysis::getScalarMemRef(op.getPtr(), adaptor.getPtr(), + loc, rewriter); + auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext()); + auto loadOp = rewriter.create( + op.getLoc(), sMemRef, zeroMap, ValueRange{}); + rewriter.replaceOp(op, loadOp.getResult()); + return success(); + } + + // 1. Simple case where no mask is used. + auto type = dyn_cast(ptr.getType()); + if (!type) { + // Seen when implicit broadcasting is done late in a chain of operations. + // The workaround is to broadcast the pointers early in the address + // calculation. A proper fix is complicated, but at least we can provide a + // better error message. + return rewriter.notifyMatchFailure( + op, "LoadOp expects a memref, not a memref of pointers"); + } + + auto tensorType = + RankedTensorType::get(type.getShape(), type.getElementType()); + auto alloc = rewriter.create( + loc, MemRefType::get(type.getShape(), type.getElementType())); + + if (!mask) { + assert(!other && "other value used in non-masked load"); + if (auto unrealizedCast = + ptr.getDefiningOp()) { + if (auto wrapType = unrealizedCast->getAttrOfType( + ModuloState::WraparoundAttr)) { + + auto memrefs = unrealizedCast.getOperands(); + auto block1 = memrefs[0]; + auto block2 = memrefs[1]; + + if (wrapType.getValue() == ModuloState::WraparoundSideBySide) { + createSideBySideCopies(block1, block2, alloc, loc, rewriter); + } else if (wrapType.getValue() == ModuloState::WraparoundStacked) { + createStackedCopies(block1, block2, alloc, loc, rewriter); + } else { + llvm_unreachable("unexpected wraparound type"); + } + } else { + llvm_unreachable("unexpected unrealized cast op"); + } + + } else { + rewriter.create(loc, ptr, alloc); + } + + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + + return success(); + } + + // 2. Continuous masked loads. + // Analyze the mask operand to determine at runtime the size of the data we + // are moving. + MaskState mstate; + auto isContMask = mstate.parse(mask, loc, rewriter); + + if (isContMask.failed()) { + return rewriter.notifyMatchFailure( + op, "Cannot lower continuous masked loads"); + } + + // fill load destination with other value + if (other) { + auto scalarOther = getScalarValue(other, loc, rewriter); + assert(scalarOther && "other value used in masked load produced by " + "unsupported instruction"); + + // For each dimension check if mstate.dims[i] < shape[i], or-accumulate + // the result + auto shape = type.getShape(); + auto accBase = + rewriter.create(loc, rewriter.getBoolAttr(false)) + .getResult(); + for (size_t i = 0; i < type.getShape().size(); i++) { + auto shapei = rewriter.create( + loc, rewriter.getIndexAttr(shape[i])); + + Value dimi = dyn_cast(mstate.dims[i]); + if (!dimi) { + dimi = rewriter.create( + loc, cast(cast(mstate.dims[i]))); + } + + auto cmpOp = rewriter.create( + loc, arith::CmpIPredicate::slt, dimi, shapei); + accBase = rewriter.create(loc, accBase, cmpOp.getResult()) + .getResult(); + } + + // condition the memset on the or-accumulation + // initialize with padding prior to CopyOp + rewriter.create( + loc, accBase, [&](OpBuilder &builder, Location loc) { + builder.create(loc, ValueRange{scalarOther}, + ValueRange{alloc}); + builder.create(loc); + }); + } + + if (auto unrealizedCast = ptr.getDefiningOp()) { + if (auto wrapType = unrealizedCast->getAttrOfType( + ModuloState::WraparoundAttr)) { + + auto memrefs = unrealizedCast.getOperands(); + auto block1 = memrefs[0]; + auto block2 = memrefs[1]; + + if (wrapType.getValue() == ModuloState::WraparoundSideBySide) { + auto [subview1, subview2] = + mstate.getSideBySideSubviews(block1, block2, loc, rewriter); + + createSideBySideCopies(subview1, subview2, alloc, loc, rewriter); + } else if (wrapType.getValue() == ModuloState::WraparoundStacked) { + auto [subview1, subview2] = + mstate.getStackedSubviews(block1, block2, loc, rewriter); + + createStackedCopies(subview1, subview2, alloc, loc, rewriter); + } else { + llvm_unreachable("unexpected wraparound type"); + } + + } else { + llvm_unreachable("unexpected unrealized cast op"); + } + + } else { + memref::SubViewOp srcSubview = mstate.getSubview(ptr, loc, rewriter); + memref::SubViewOp dstSubview = mstate.getSubview(alloc, loc, rewriter); + rewriter.create(loc, srcSubview, dstSubview); + } + + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + + return success(); + } +}; + +struct StoreConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto ptr = adaptor.getPtr(); + auto val = adaptor.getValue(); + auto mask = op.getMask(); + auto loc = op.getLoc(); + + // 0. Shortcut for scalar stores + if (!isa(val.getType())) { + auto sMemRef = + PtrAnalysis::getScalarMemRef(op.getPtr(), ptr, loc, rewriter); + auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext()); + rewriter.create(loc, val, sMemRef, zeroMap, + ValueRange{}); + rewriter.eraseOp(op); + return success(); + } + + // 1. Simple case where no mask is used. + if (!mask) { + auto storeOp = rewriter.create( + loc, val, ptr); + storeOp.setWritable(true); + rewriter.eraseOp(op); + return success(); + } + + // 2. Continuous masked stores. + // Analyze the mask operand to determine at runtime the size of the data we + // are moving. + MaskState mstate; + auto isContMask = mstate.parse(mask, loc, rewriter); + + if (isContMask.failed()) + return failure(); + + auto srcSlice = mstate.getExtractSlice(val, loc, rewriter); + auto dstSubview = mstate.getSubview(ptr, loc, rewriter); + + auto storeOp = rewriter.create( + loc, srcSlice, dstSubview); + storeOp.setWritable(true); + rewriter.eraseOp(op); + + return success(); + } +}; + +struct LoopConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(scf::ForOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + llvm::SmallDenseMap knownPtrs; + PtrAnalysis::IndexMapSet + levelToBlockArgIndex; // level -> set of block arg index to be replaced + + PtrAnalysis::rewriteForOp(op, rewriter, levelToBlockArgIndex, 0, knownPtrs); + return success(); + } +}; + +struct YieldConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(scf::YieldOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, adaptor.getOperands()); + return success(); + } +}; + +// Remove all Meta ops except for AddPtr which is handled by AddPtrConverter. +// Use benefit == 10 to ensure that this pattern always takes precedence over +// other patterns. +struct MetaOpConverter : public RewritePattern { +private: + // UseAnalysis will tag operations whose results are used only as meta-data + // with "MetaUse" tag. + bool isMetaUse(Operation *op) const { return op->hasAttr("MetaUse"); } + +public: + MetaOpConverter(MLIRContext *context) + : RewritePattern(MatchAnyOpTypeTag(), /*benefit=*/10, context) {} + + LogicalResult matchAndRewrite(Operation *op, + PatternRewriter &rewriter) const final { + + if (isa(op)) { + return rewriter.notifyMatchFailure(op, + "AddPtrOp will be handled separately"); + } + + if (isMetaUse(op)) { + rewriter.eraseOp(op); + return success(); + } + + return rewriter.notifyMatchFailure(op, "requires meta ops"); + } +}; + +struct UnrealizedCastConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(UnrealizedConversionCastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.eraseOp(op); + return success(); + } +}; + +//----------------------------- +// End of monolithic only +//----------------------------- + +struct SplatConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::SplatOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto opType = cast(op.getType()); + auto loc = op.getLoc(); + + auto init = rewriter.create(loc, opType.getShape(), + opType.getElementType()); + + auto filledTensor = + rewriter + .create(loc, ValueRange{adaptor.getSrc()}, + ValueRange{init}) + .result(); + + rewriter.replaceOp(op, filledTensor); + return success(); + } +}; + +struct BroadcastConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + SmallVector getBroadcastDims(RankedTensorType src, + RankedTensorType dst) const { + SmallVector broadcastDims; + auto srcShape = src.getShape(); + auto dstShape = dst.getShape(); + + for (size_t i = 0; i < srcShape.size(); i++) { + if (dstShape[i] != srcShape[i]) { + assert(srcShape[i] == 1); + broadcastDims.push_back(i); + } + } + assert(!broadcastDims.empty() && "cannot identify broadcast dimension"); + return broadcastDims; + } + + // Broadcasts input tensor based on TosaToLinalg's broadcastToShape + AffineMap getBroadcastAffineMap(MLIRContext *context, + ArrayRef inputShape, + ArrayRef broadcastToShape) const { + + assert(broadcastToShape.size() >= inputShape.size()); + + // Create affine map and shapes for tensor initialization. + SmallVector outExpr; + + size_t diff = broadcastToShape.size() - inputShape.size(); + for (size_t i = 0; i < broadcastToShape.size(); i++) { + if (i < diff) { + continue; + } + size_t j = i - diff; + if (inputShape[j] == 1) { + // Broadcast singleton dimension + outExpr.push_back(mlir::getAffineConstantExpr(0, context)); + continue; + } + // Non-broadcast case + outExpr.push_back(mlir::getAffineDimExpr(i, context)); + } + return AffineMap::get(broadcastToShape.size(), 0, outExpr, context); + } + +public: + LogicalResult + matchAndRewrite(triton::BroadcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + + assert(op->getNumResults() == 1 && "code assumes single result!"); + RankedTensorType sourceType = + cast(adaptor.getSrc().getType()); + RankedTensorType resultType = cast(op.getType()); + auto elementType = resultType.getElementType(); + size_t resultRank = resultType.getRank(); + + SmallVector indexingMaps; + indexingMaps.reserve(op->getNumOperands() + op->getNumResults()); + + indexingMaps.push_back(getBroadcastAffineMap( + op->getContext(), sourceType.getShape(), resultType.getShape())); + indexingMaps.append(op->getNumResults(), + rewriter.getMultiDimIdentityMap(resultRank)); + + assert(op->getNumResults() == 1 && "code assumes single result!"); + auto init = rewriter.create(loc, resultType.getShape(), + elementType); + + auto linalgOp = rewriter.create( + loc, op->getResultTypes(), ValueRange{adaptor.getSrc()}, + ValueRange{init}, indexingMaps, getNParallelLoopsAttrs(resultRank), + [&](OpBuilder &nestedBuilder, Location nestedLoc, + ValueRange blockArgs) { + Value opResult = blockArgs[0]; + nestedBuilder.create(loc, opResult); + }); + + linalgOp->setAttr("broadcastDims", + rewriter.getDenseI64ArrayAttr( + getBroadcastDims(sourceType, resultType))); + + rewriter.replaceOp(op, linalgOp->getResults()); + return success(); + } +}; + +struct ExpandDimsConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ExpandDimsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto src = adaptor.getSrc(); + auto srcRank = cast(src.getType()).getRank(); + auto resType = cast(op->getResultTypes()[0]); + SmallVector reassoc; + int64_t c = 0; + for (int64_t i = 0; i < srcRank; i++) { + ReassociationIndices g; + g.push_back(c++); + if (op.getAxis() == i) { + g.push_back(c++); + } else if (op.getAxis() == i + 1 && i == srcRank - 1) { + g.push_back(c++); + } + reassoc.push_back(g); + } + + auto expandShapeOp = rewriter.create( + op.getLoc(), resType, src, reassoc); + + rewriter.replaceOp(op, expandShapeOp.getResult()); + return success(); + } +}; + +struct TransposeConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::TransOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto src = adaptor.getSrc(); + auto srcRank = cast(src.getType()).getRank(); + assert(srcRank == 2 && "only expect transposing 2D data"); + + auto res = getTransposedValue(src, op.getLoc(), rewriter); + rewriter.replaceOp(op, res); + return success(); + } +}; + +struct MakeRangeConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::MakeRangeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto type = cast(op.getResult().getType()); + auto shape = type.getShape(); + auto elementType = type.getElementType(); + auto context = rewriter.getContext(); + + assert(type.getShape().size() == 1 && + type.getElementType().getIntOrFloatBitWidth() == 32 && + "make range can only return 1D int32 tensor"); + + SmallVector indexingMaps{AffineMap::get( + /* dimCount */ 1, /* symbolCount */ 0, + SmallVector{mlir::getAffineDimExpr(0, context)}, context)}; + + auto init = rewriter.create(loc, shape, elementType); + auto linalgOp = rewriter.create( + loc, op->getResultTypes(), /* operands */ ValueRange{}, + ValueRange{init}, indexingMaps, getNParallelLoopsAttrs(1), + [&](OpBuilder &nestedBuilder, Location nestedLoc, + ValueRange blockArgs) { + Value index = nestedBuilder.create(loc, 0); + Value res = nestedBuilder.create( + loc, type.getElementType(), index); + nestedBuilder.create(loc, res); + }); + + rewriter.replaceOp(op, linalgOp->getResults()); + return success(); + } +}; + +struct AssertConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::AssertOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value condVal = op.getCondition(); + + if (isa(condVal.getType())) { + auto scalarVal = getScalarValue(op.getCondition(), op.getLoc(), rewriter); + condVal = scalarVal ? scalarVal : condVal; + } + assert(condVal && isa(condVal.getType()) && + "Only asserts on scalars are currently supported"); + + if (!condVal.getType().isInteger(1)) { + auto zero = + rewriter.create(op.getLoc(), 0, 32); + auto newCond = rewriter.create( + op.getLoc(), arith::CmpIPredicate::ne, condVal, zero); + condVal = newCond.getResult(); + } + + auto assertMessage = + llvm::formatv("Assertion `{0}` failed", op.getMessage()); + rewriter.create(op.getLoc(), condVal, + assertMessage.str()); + + rewriter.eraseOp(op); + return success(); + } +}; + +struct BitcastConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::BitcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // arith::bitcast does not support casting pointers + if (triton::isPtrTypeLike(op.getType())) { + return failure(); + } + + auto arithBitcast = rewriter.create( + op.getLoc(), op.getType(), op.getOperand()); + + rewriter.replaceOp(op, arithBitcast.getResult()); + return success(); + } +}; + +struct CallConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::CallOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + SmallVector args = adaptor.getOperands(); + + // We need to pass extra arguments added by addProgramInfo which are num_programs and program_ids + if (FuncOp parentFunc = op->getParentOfType()) { + SymbolRefAttr calleeAttr = op.getCalleeAttr(); + StringRef calleeName = calleeAttr.getRootReference(); + + if (ModuleOp module = op->getParentOfType()) { + if (FuncOp calleeFunc = module.lookupSymbol(calleeName)) { + size_t argsNeed = calleeFunc.getFunctionType().getInputs().size(); + Block &entryBlock = parentFunc.front(); + auto parentInputs = entryBlock.getArguments(); + size_t argsParent = parentInputs.size(); + + if (argsNeed > args.size()) { + int missing = argsNeed - args.size(); + int missingArgsStart = argsParent - missing; + for (int i = 0; i < missing; i++) { + args.push_back(parentInputs[missingArgsStart + i]); + } + } + } + } + } + + auto call = rewriter.create( + op.getLoc(), op.getCallee(), op.getResultTypes(), args); + + if (!call) { + op.emitError("Failed to create func::CallOp"); + return failure(); + } + + rewriter.replaceOp(op, call); + return success(); + } +}; + +struct FpToFpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::FpToFpOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto roundingMode = triton::RoundingMode::RTNE; // default + + auto roundingModeAttr = op.getRounding(); + if (roundingModeAttr.has_value()) { + roundingMode = roundingModeAttr.value(); + } + + assert(roundingMode != triton::RoundingMode::RTZ && + "Rounding Towards Zero is not supported"); + + Type resultType = op.getResult().getType(); + + auto operandWidth = getBitWidth(op.getOperand().getType()); + auto resultWidth = getBitWidth(resultType); + + assert(operandWidth.has_value() && resultWidth.has_value() && + "Not a float-like operand or result"); + + if (operandWidth.value() > resultWidth.value()) { + Value truncatedValue = rewriter.create(op.getLoc(), resultType, op.getOperand()); + rewriter.replaceOp(op, truncatedValue); + return success(); + } + + Value extendedValue = rewriter.create(op.getLoc(), resultType, op.getOperand()); + rewriter.replaceOp(op, extendedValue); + + return success(); + } +}; + +struct ClampConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ClampFOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + bool propagateNan = op.getPropagateNan() == triton::PropagateNan::ALL; + + assert(!propagateNan && + "PropagateNan is not supported"); + + Location loc = op.getLoc(); + Value x = adaptor.getOperands()[0]; + Value min = adaptor.getOperands()[1]; + Value max = adaptor.getOperands()[2]; + + Value maxMin = rewriter.create(loc, x, min); + Value clamp = rewriter.create(loc, maxMin, max); + rewriter.replaceOp(op, clamp); + + return success(); + } +}; + +struct PreciseSqrtConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::PreciseSqrtOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto replacement = rewriter.create( + op.getLoc(), adaptor.getOperands()); + + rewriter.replaceOp(op, replacement); + return success(); + } +}; + +struct PreciseDivConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::PreciseDivFOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto replacement = rewriter.create( + op.getLoc(), adaptor.getOperands()); + + rewriter.replaceOp(op, replacement); + return success(); + } +}; + +struct CatConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::CatOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto replacement = rewriter.create( + op.getLoc(), 0 /* concat dimension */, adaptor.getOperands()); + + rewriter.replaceOp(op, replacement); + + return success(); + } +}; + +struct SplitConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::SplitOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + Value input = op.getOperand(); + auto inputType = cast(input.getType()); + + Type resultType = op.getResults().front().getType(); + auto resultTensor = cast(resultType); + auto shape = inputType.getShape(); + + SmallVector offsets(shape.size(), rewriter.getIndexAttr(0)); + SmallVector strides(shape.size(), rewriter.getIndexAttr(1)); + SmallVector sizes = + llvm::to_vector(llvm::map_range(shape, [&](int64_t dim) -> OpFoldResult { + return rewriter.getIndexAttr(dim); + })); + + SmallVector results; + + for (int i = 0; i < 2; ++i) { + offsets.pop_back(); + sizes.pop_back(); + + offsets.push_back(rewriter.getIndexAttr(i)); + sizes.push_back(rewriter.getIndexAttr(1)); + Value slice = rewriter.create( + loc, resultTensor, input, offsets, sizes, strides); + results.push_back(slice); + } + + rewriter.replaceOp(op, results); + return success(); + } +}; + +struct JoinConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::JoinOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + ValueRange inputs = op.getOperands(); + + auto resultType = cast(op.getResult().getType()); + + auto loc = op.getLoc(); + Value result = rewriter.create(loc, resultType.getShape(), resultType.getElementType()); + + auto shape = resultType.getShape(); + + SmallVector offsets(shape.size(), rewriter.getIndexAttr(0)); + SmallVector strides(shape.size(), rewriter.getIndexAttr(1)); + SmallVector sizes = + llvm::to_vector(llvm::map_range(shape, [&](int64_t dim) -> OpFoldResult { + return rewriter.getIndexAttr(dim); + })); + + for (int i = 0; i < 2; ++i) { + offsets.pop_back(); + sizes.pop_back(); + + offsets.push_back(rewriter.getIndexAttr(i)); + sizes.push_back(rewriter.getIndexAttr(1)); + result = rewriter.create(loc, inputs[i], result, offsets, sizes, strides); + } + + rewriter.replaceOp(op, result); + + return success(); + } +}; + +struct MulHiUIOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::MulhiUIOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + + auto mulResult = rewriter.create(loc, adaptor.getOperands()); + rewriter.replaceOp(op, mulResult.getHigh()); + + return success(); + } +}; + +struct MatmulConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + // true means tensor elements are zeros + // false means not zero or it cannot be determined + bool isZeroTensor(Value &v, bool integers) const { + if (auto splatOp = v.getDefiningOp()) { + if (auto constOp = splatOp.getSrc().getDefiningOp()) { + if (auto val = dyn_cast(constOp.getValue())) { + return val.getValueAsDouble() == 0.; + } + if (auto val = dyn_cast(constOp.getValue())) { + return val.getValue() == 0; + } + } + return false; + } + + if (auto constOp = v.getDefiningOp()) { + if (auto denseAttr = dyn_cast(constOp.getValue())) { + if (denseAttr.isSplat()) { + if (integers) + return denseAttr.getSplatValue().isZero(); + return denseAttr.getSplatValue().isZero(); + } + } + } + + return false; + } + + LogicalResult + matchAndRewrite(triton::DotOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto opa = op.getA(); + auto opb = op.getB(); + auto opc = op.getC(); + + auto dstType = cast(op.getType()); + auto elementType = dstType.getElementType(); + bool integers = elementType.isInteger(); + bool skipC = isZeroTensor(opc, integers); + auto init = + rewriter.create(loc, dstType.getShape(), elementType); + TypedAttr constantAttr = integers ? + static_cast(rewriter.getIntegerAttr(elementType, 0)) : + static_cast(rewriter.getFloatAttr(elementType, 0)); + + auto zero = rewriter.create( + op.getLoc(), elementType, constantAttr); + + auto zeroes = + rewriter.create(loc, ValueRange{zero}, ValueRange{init}) + .result(); + + auto res = rewriter + .create(loc, ValueRange{opa, opb}, + ValueRange{zeroes}) + .getResult(0); + + if (!skipC) { + if (integers) { + res = rewriter.create(loc, opc, res); + } else { + res = rewriter.create(loc, opc, res); + } + } + + rewriter.replaceOp(op, res); + return success(); + } +}; + +struct ReduceConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +private: + llvm::SmallVector getRedOps(triton::ReduceOp redOp) const { + auto reduceBlock = redOp.getBody(); + return llvm::map_to_vector(reduceBlock->without_terminator(), + [](Operation &op) { return &op; }); + } + + bool isReductionOpSupported(Operation *redOp) const { + return isa( + redOp); + } + + arith::ConstantOp getRedBaseConstOp(ConversionPatternRewriter &rewriter, + Operation *redOp, + Type constantType) const { + const int64_t bitWidth = constantType.getIntOrFloatBitWidth(); + + auto attr = + llvm::TypeSwitch(redOp) + .Case([&](arith::AddFOp) { + return rewriter.getFloatAttr(constantType, 0.f); + }) + .Case([&](arith::AddIOp) { + return rewriter.getIntegerAttr(constantType, 0); + }) + .Case([&](auto) { + return rewriter.getFloatAttr( + constantType, -std::numeric_limits::infinity()); + }) + .Case([&](auto) { + return rewriter.getFloatAttr( + constantType, std::numeric_limits::infinity()); + }) + .Case([&](arith::MinSIOp) { + return rewriter.getIntegerAttr(constantType, + llvm::maxIntN(bitWidth)); + }) + .Case([&](arith::MinUIOp) { + return rewriter.getIntegerAttr(constantType, + llvm::maxUIntN(bitWidth)); + }) + .Case([&](arith::MaxSIOp) { + return rewriter.getIntegerAttr(constantType, + llvm::minIntN(bitWidth)); + }) + .Case([&](auto) { + return rewriter.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::MulFOp) { + return rewriter.getFloatAttr(constantType, 1.f); + }) + .Case([&](auto) { + return rewriter.getIntegerAttr(constantType, 1); + }) + .Case([&](arith::OrIOp) { + return rewriter.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::DivFOp) { + return rewriter.getFloatAttr(constantType, 1.f); + }) + .Case([&](arith::SubFOp) { + return rewriter.getFloatAttr(constantType, 0.f); + }) + .Default([](Operation *op) { + op->dump(); + llvm_unreachable("Reduction op not yet supported"); + return nullptr; + }); + + return rewriter.create(redOp->getLoc(), constantType, + attr); + } + + bool requiresF32Conversion(const Type elemType, Operation *redOp) const { + unsigned width = + cast(Float32Type::get(elemType.getContext())).getWidth(); + return isa(elemType) && + elemType.getIntOrFloatBitWidth() < width && + isa(redOp); + } + + Value getRedElement(Value lhs, Value rhs, const Location loc, + Operation *redOp, OpBuilder &b, + const bool convertLhsToF32Precision) const { + return llvm::TypeSwitch(redOp) + .Case([&](auto redOp) { + if (convertLhsToF32Precision) { + lhs = b.create(loc, Float32Type::get(b.getContext()), + lhs); + } + return b.create(loc, lhs, rhs); + }) + .Case([&](auto redOp) { + return b.create(loc, lhs, rhs); + }) + .Default([](Operation *op) { + op->dump(); + llvm_unreachable("Reduction op not yet supported"); + return nullptr; + }); + } + + LogicalResult + convertToLinalgReduce(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto source = adaptor.getOperands().front(); + auto sourceType = cast(source.getType()); + auto elemType = sourceType.getElementType(); + auto resType = op.getResult().front().getType(); + auto loc = op.getLoc(); + auto reductionOps = getRedOps(op); + +#ifdef ORIGIN_TRITON_SHARED + // Reduction of arbitrary operations isn't supported because using the first + // element across the reduction dimension requires us to iterate over a + // subview that skips over each first element. + if (reductionOps.size() != 1 || + !isReductionOpSupported(reductionOps.front())) { + return rewriter.notifyMatchFailure( + op, "Only support lowering reduction with body " + "containing 1 max(i/f), addf, ori, or mulf."); +#else + // flagtree: Use unified hardware manager to determine reduction strategy + auto hardwareManager = mlir::flagtree::createUnifiedHardwareManager(); + auto reduceStrategy = hardwareManager->getReduceStrategy(); + + if (reductionOps.size() != 1) { + if (reduceStrategy=="linalg_reduce") { + return applyLinalgReduce(op, adaptor, rewriter); + } else { + return rewriter.notifyMatchFailure(op, "Reduction with multiple ops and unknown strategy is not supported."); + } + } +#endif + + auto rop = reductionOps.front(); + auto axis = op.getAxis(); + auto isVectorReduce = sourceType.getRank() == 1; + + if (axis == sourceType.getRank() - 1 && !isVectorReduce) { + source = getTransposedValue(source, op.getLoc(), rewriter); + axis = sourceType.getRank() - 2; + } + + bool convertToF32Precision = requiresF32Conversion(resType, rop); + + auto constantType = convertToF32Precision + ? Float32Type::get(rewriter.getContext()) + : elemType; + + auto accBaseConstOp = getRedBaseConstOp(rewriter, rop, constantType); + Value initTensor; + + if (isVectorReduce) { + // The affine vectorizer cannot vectorize affine loops generated from + // linalg.reduce for the vector reduce case, so we must rewrite the + // linalg.reduce to affine loops manually. Here we lower to AllocTensor + // directly instead of EmptyOp so that the subsequent pass can recognize + // the patterns (EmptyOp is susceptible to being CSE'd away, making it + // harder to match the patterns correctly). + initTensor = rewriter.create( + loc, RankedTensorType::get({}, constantType), ValueRange{}); + initTensor = rewriter.create(loc, accBaseConstOp, + initTensor, ValueRange{}); + } else { + Value init = rewriter.create( + loc, cast(resType).getShape(), constantType); + initTensor = rewriter + .create(loc, ValueRange{accBaseConstOp}, + ValueRange{init}) + .result(); + } + + Value finalResult = + rewriter + .create( + loc, ValueRange{source}, ValueRange{initTensor}, + SmallVector{axis}, + [&](OpBuilder &opBuilder, Location loc, ValueRange inputs) { + assert(inputs.size() == 2); + Value result = + getRedElement(inputs[0], inputs[1], loc, rop, opBuilder, + convertToF32Precision); + opBuilder.create(loc, result); + }) + .getResult(0); + + if (sourceType.getRank() == 1) { + finalResult = + rewriter.create(loc, constantType, finalResult); + } + + if (convertToF32Precision) { + finalResult = rewriter.create(loc, resType, finalResult); + } + + rewriter.replaceOp(op, finalResult); + return success(); + } + +public: + LogicalResult + matchAndRewrite(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto sourceType = + cast(adaptor.getOperands().front().getType()); + assert(sourceType.hasRank() && "Expected input is " + "ranked"); + + int64_t axis = op.getAxis(); + assert(axis >= 0 && axis < sourceType.getRank() && + "Expected reduction " + "axis is within " + "operand's rank"); + + return convertToLinalgReduce(op, adaptor, rewriter); + } +}; + +// flagtree: Pattern converter for Triton reduce return operations to Linalg +// yield operations. This converter handles the terminator operation within +// reduce combine regions +struct ReduceReturnConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::ReduceReturnOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, adaptor.getOperands()); + return success(); + } +}; + +// flagtree: Pattern converter for var_mean_welford to Linalg yield operations. +class VarMeanConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + // We're looking for an op that looks like this: + // + // %26:3 = "tt.reduce"(%25#1, %25#2, %25#0) <{axis = 1 : i32}> ({ + // ^bb0(%arg6: f32, %arg7: f32, %arg8: f32, %arg9: f32, %arg10: f32, + // %arg11: f32): + // %33 = arith.addf %arg7, %arg10 : f32 + // %34 = arith.maxnumf %33, %cst : f32 + // %35 = arith.mulf %arg6, %arg7 : f32 + // %36 = arith.mulf %arg9, %arg10 : f32 + // %37 = arith.addf %35, %36 : f32 + // %38 = arith.divf %37, %34 : f32 + // %39 = arith.mulf %35, %arg6 : f32 + // %40 = arith.addf %arg8, %39 : f32 + // %41 = arith.addf %40, %arg11 : f32 + // %42 = arith.mulf %36, %arg9 : f32 + // %43 = arith.addf %41, %42 : f32 + // %44 = arith.mulf %33, %38 : f32 + // %45 = arith.mulf %44, %38 : f32 + // %46 = arith.subf %43, %45 : f32 + // tt.reduce.return %38, %33, %46 : f32, f32, f32 + // }) : (tensor<8x2048xf32>, tensor<8x2048xf32>, tensor<8x2048xf32>) -> + // (tensor<8xf32>, tensor<8xf32>, tensor<8xf32>) + // + // The above mlir code is lowered from this combinator in triton's + // standard.py: + // + // def welford_func(mean_x, count_x, M_x, mean_y, count_y, M_y): + // count = count_x + count_y + // _count = tl.maximum(count, 1) + // mc_x = mean_x * count_x + // mc_y = mean_y * count_y + // mean = (mc_x + mc_y) / _count + // M = M_x + mc_x * mean_x + M_y + mc_y * mean_y - count * mean * mean + // return mean, count, M + + Value getInitTensor(ConversionPatternRewriter &rewriter, + ArrayRef shape, Value fillValue, + Location loc) const { + Value initTensor = + rewriter.create(loc, shape, fillValue.getType()); + return rewriter + .create(loc, ValueRange{fillValue}, + ValueRange{initTensor}) + .result(); + } + + LogicalResult checkConstFloat(Value value, float val) const { + if (auto constOp = dyn_cast(value.getDefiningOp())) { + if (auto floatAttr = dyn_cast(constOp.getValue())) { + if (floatAttr.getValueAsDouble() == val) { + return success(); + } + } + } + return failure(); + } + + LogicalResult matchVarMeanBody(Value mean_x, Value count_x, Value M_x, + Value mean_y, Value count_y, Value M_y, + mlir::Block::iterator &it, + Operation *block_teminator) const { + + // %33 = arith.addf %arg7, %arg10 : f32 // count = count_x + count_y + // %34 = arith.maxnumf %33, %cst : f32 // _count = tl.maximum(count, 1) + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto addOp0 = dyn_cast(*it++); + if (addOp0) { + if (count_x != addOp0.getLhs() || count_y != addOp0.getRhs()) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto maxOp0 = dyn_cast(*it++); + if (maxOp0) { + if (maxOp0.getLhs() != addOp0) { + return failure(); + } + if (failed(checkConstFloat(maxOp0.getRhs(), 1.f))) { + return failure(); + } + } else { + return failure(); + } + + // %35 = arith.mulf %arg6, %arg7 : f32 //mc_x = mean_x * count_x + // %36 = arith.mulf %arg9, %arg10 : f32 //mc_y = mean_y * count_y + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto mulOp0 = dyn_cast(*it++); + if (mulOp0) { + if (mean_x != mulOp0.getLhs() || count_x != mulOp0.getRhs()) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto mulOp1 = dyn_cast(*it++); + if (mulOp1) { + if (mean_y != mulOp1.getLhs() || count_y != mulOp1.getRhs()) { + return failure(); + } + } else { + return failure(); + } + + // mean = (mc_x + mc_y) / _count + // + // %37 = arith.addf %35, %36 : f32 // sum_mc = mc_x + mc_y + // %38 = arith.divf %37, %34 : f32 // mean = sum_mc / _count + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto addOp1 = dyn_cast(*it++); + if (addOp1) { + if (addOp1.getLhs() != mulOp0 || addOp1.getRhs() != mulOp1) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto divOp0 = dyn_cast(*it++); + if (divOp0) { + if (divOp0.getLhs() != addOp1 || divOp0.getRhs() != maxOp0) { + return failure(); + } + } else { + return failure(); + } + + // M = M_x + mc_x * mean_x + M_y + mc_y * mean_y - count * mean * mean + // + // %39 = arith.mulf %35, %arg6 : f32 // part_x = mc_x * mean_x + // %40 = arith.addf %arg8, %39 : f32 // item_1 = M_x + part_x + // %41 = arith.addf %40, %arg11 : f32 // item_2 = item_1 + M_y + // %42 = arith.mulf %36, %arg9 : f32 // part_y = mc_y * mean_y + // %43 = arith.addf %41, %42 : f32 // item_3 = iterm_2 + part_y + // %44 = arith.mulf %33, %38 : f32 // mean_0 = count * mean + // %45 = arith.mulf %44, %38 : f32 // mean_1 = mean_0 * mean + // %46 = arith.subf %43, %45 : f32 // M = item_3 - mean_1 + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto mulOp2 = dyn_cast(*it++); + if (mulOp2) { + if (mulOp2.getLhs() != mulOp0 || mulOp2.getRhs() != mean_x) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto addOp2 = dyn_cast(*it++); + if (addOp2) { + if (addOp2.getLhs() != M_x || addOp2.getRhs() != mulOp2) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto addOp3 = dyn_cast(*it++); + if (addOp3) { + if (addOp3.getLhs() != addOp2 || addOp3.getRhs() != M_y) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto mulOp3 = dyn_cast(*it++); + if (mulOp3) { + if (mulOp3.getLhs() != mulOp1 || mulOp3.getRhs() != mean_y) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto addOp4 = dyn_cast(*it++); + if (addOp4) { + if (addOp4.getLhs() != addOp3 || addOp4.getRhs() != mulOp3) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto mulOp4 = dyn_cast(*it++); + if (mulOp4) { + if (mulOp4.getLhs() != addOp0 || mulOp4.getRhs() != divOp0) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto mulOp5 = dyn_cast(*it++); + if (mulOp5) { + if (mulOp5.getLhs() != mulOp4 || mulOp5.getRhs() != divOp0) { + return failure(); + } + } else { + return failure(); + } + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto subOp = dyn_cast(*it++); + if (subOp) { + if (subOp.getLhs() != addOp4 || subOp.getRhs() != mulOp5) { + return failure(); + } + } else { + return failure(); + } + + // tt.reduce.return %38, %33, %46 : f32, f32, f32 //return mean, count, M + + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto termOp = dyn_cast(*it++); + if (termOp && termOp == block_teminator) { + auto opnds = termOp.getOperands(); + if (opnds != ArrayRef{divOp0, addOp0, subOp}) { + return failure(); + } + } else { + return failure(); + } + + return success(); + } + +public: + VarMeanConverter(MLIRContext *context) : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(ReduceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override final { + // check num_args = 6 : mean_x, count_x, M_x, mean_y, count_y, M_y + if (op.getBody()->getNumArguments() != 6) { + return failure(); + } + + auto block = op.getBody(); + auto ops = block->without_terminator(); + + Value mean_x = block->getArgument(0); + Value count_x = block->getArgument(1); + Value M_x = block->getArgument(2); + Value mean_y = block->getArgument(3); + Value count_y = block->getArgument(4); + Value M_y = block->getArgument(5); + + auto opsIt = ops.begin(); + if (failed(matchVarMeanBody(mean_x, count_x, M_x, mean_y, count_y, M_y, + opsIt, block->getTerminator()))) { + return failure(); + } + auto loc = op.getLoc(); + auto elemTypes = op.getElementTypes(); + + auto meanType = elemTypes[0]; + auto countType = elemTypes[1]; + auto MType = elemTypes[2]; + + Value zeroMean = rewriter.create( + loc, meanType, rewriter.getFloatAttr(meanType, 0.f)); + Value zeroCount = rewriter.create( + loc, countType, rewriter.getFloatAttr(countType, 0.f)); + Value zeroM = rewriter.create( + loc, MType, rewriter.getFloatAttr(MType, 0.f)); + + auto valueResultType = dyn_cast(op.getType(0)); + const auto isScalarReduce = valueResultType == nullptr; + SmallVector reductionResultShape{ + isScalarReduce ? SmallVector{} + : SmallVector(valueResultType.getShape())}; + + auto initTensorMean = + getInitTensor(rewriter, reductionResultShape, zeroMean, loc); + auto initTensorCount = + getInitTensor(rewriter, reductionResultShape, zeroCount, loc); + auto initTensorM = + getInitTensor(rewriter, reductionResultShape, zeroM, loc); + + SmallVector outputs = {initTensorMean, initTensorCount, initTensorM}; + + auto linalgOp = rewriter.create( + loc, adaptor.getOperands(), outputs, + SmallVector{adaptor.getAxis()}, + [&](OpBuilder &b, Location loc, ValueRange inputs) { + assert(inputs.size() == 6 && + "Expected 6 inputs to varmean reduce block"); + + auto tritonReduceBlock = op.getBody(); + IRMapping mapping; + mapping.map(tritonReduceBlock->getArguments(), inputs); + + for (auto &op : tritonReduceBlock->without_terminator()) { + b.clone(op, mapping); + } + + auto tritonYield = tritonReduceBlock->getTerminator(); + auto results = + llvm::map_to_vector(tritonYield->getOperands(), [&](Value val) { + return mapping.lookup(val); + }); + b.create(loc, results); + }); + + if (isScalarReduce) { + SmallVector reduceResults{ + rewriter.create( + loc, meanType, linalgOp.getResults()[0], ValueRange{}), + rewriter.create( + loc, countType, linalgOp.getResults()[1], ValueRange{}), + rewriter.create( + loc, MType, linalgOp.getResults()[2], ValueRange{}), + }; + rewriter.replaceOp(op, reduceResults); + } else { + rewriter.replaceOp(op, linalgOp); + } + + return success(); + } +}; + +template +class ArgMinMaxBaseConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + // We're looking for an op that looks like this: + // + // %9:2 = "tt.reduce"(%8, %3) <{axis = 0 : i32}> ({ + // ^bb0(%arg9: f32, %arg10: i32, %arg11: f32, %arg12: i32): + // ------------------------------------------------- + // `matchTieBreakValue` | + // %11 = arith.cmpf oeq, %arg9, %arg11 : f32 | + // %12 = arith.cmpi slt, %arg10, %arg12 : i32 | 1. + // %13 = arith.andi %11, %12 : i1 | + // ------------------------------------------------- |-> `matchShouldUpdate` + // `matchUpdateCondition` | + // %14 = arith.cmpf ogt, %arg9, %arg11 : f32 | 2. + // ------------------------------------------------- | + // %15 = arith.ori %14, %13 : i1 | + // ------------------------------------------------- + // %16 = arith.select %15, %arg9, %arg11 : f32 + // %17 = arith.select %15, %arg10, %arg12 : i32 + // tt.reduce.return %16, %17 : f32, i32 + // }) : (tensor<4096xf32>, tensor<4096xi32>) -> (f32, i32) + // + // The above mlir code is lowered from this combinator in triton's + // standard.py: + // + // def _argmax_combine(value1, index1, value2, index2, tie_break_left): + // if tie_break_left: + // tie = value1 == value2 and index1 < index2 + // else: + // tie = False + // gt = value1 > value2 or tie + // v_ret = core.where(gt, value1, value2) + // i_ret = core.where(gt, index1, index2) + // return v_ret, i_ret + + LogicalResult matchTieBreakResult(Value currValue, Value currIndex, + Value reduceValue, Value reduceIndex, + mlir::Block::iterator &it, + Value &tileBreakValue) const { + // Match the following (section 1. of the above) + // + // %11 = arith.cmpf/i oeq, %arg9, %arg11 : f32 + // %12 = arith.cmpi slt, %arg10, %arg12 : i32 + // %13 = arith.andi %11, %12 : i1 + // + // which is equivalent to the following python code + // + // tie = value1 == value2 and index1 < index2 + + // matching: %11 = arith.cmpf oeq, %arg9, %arg11 : f32 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + Value eqCmpOp; + if (auto cmpOp = dyn_cast(*it)) { + if (cmpOp.getPredicate() != arith::CmpFPredicate::OEQ) { + return failure(); + } + if (currValue != cmpOp.getLhs() || reduceValue != cmpOp.getRhs()) { + return failure(); + } + eqCmpOp = cmpOp; + } else if (auto cmpOp = dyn_cast(*it)) { + if (cmpOp.getPredicate() != arith::CmpIPredicate::eq) { + return failure(); + } + if (currValue != cmpOp.getLhs() || reduceValue != cmpOp.getRhs()) { + return failure(); + } + eqCmpOp = cmpOp; + } else { + return failure(); + } + it++; + + // matching: %12 = arith.cmpi slt, %arg10, %arg12 : i32 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto sltCmpOp = dyn_cast(*it++); + if (sltCmpOp) { + if (sltCmpOp.getPredicate() != arith::CmpIPredicate::slt) { + return failure(); + } + if (currIndex != sltCmpOp.getLhs() || reduceIndex != sltCmpOp.getRhs()) { + return failure(); + } + } else { + return failure(); + } + + // matching: %13 = arith.andi %11, %12 : i1 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto andOp = dyn_cast(*it++); + if (andOp) { + if (andOp.getLhs() != eqCmpOp || andOp.getRhs() != sltCmpOp) { + return failure(); + } + } else { + return failure(); + } + + tileBreakValue = andOp; + return success(); + } + + LogicalResult matchShouldUpdateValue(Value currValue, Value currIndex, + Value reduceValue, Value reduceIndex, + mlir::Block::iterator &it, + Value &shouldUpdate) const { + Value tieResult; + if (failed(matchTieBreakResult(currValue, currIndex, reduceValue, + reduceIndex, it, tieResult))) { + LLVM_DEBUG(llvm::dbgs() << "Tie break result match failed\n"); + return failure(); + } + + Value comparisonResult; + if (failed(T::matchComparisonResult(currValue, currIndex, reduceValue, + reduceIndex, it, comparisonResult))) { + LLVM_DEBUG(llvm::dbgs() << "Comparison result match failed\n"); + return failure(); + } + + // matching: %15 = arith.ori %14, %13 : i1 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + auto orOp = dyn_cast(*it++); + if (orOp) { + if (orOp.getLhs() != comparisonResult || orOp.getRhs() != tieResult) { + return failure(); + } + } else { + return failure(); + } + + shouldUpdate = orOp; + return success(); + } + + Value getInitTensor(ConversionPatternRewriter &rewriter, + ArrayRef shape, Value fillValue, + Location loc) const { + Value initTensor = + rewriter.create(loc, shape, fillValue.getType()); + return rewriter + .create(loc, ValueRange{fillValue}, + ValueRange{initTensor}) + .result(); + } + +public: + ArgMinMaxBaseConverter(MLIRContext *context) : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(ReduceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override final { + if (op.getBody()->getNumArguments() != 4) { + return failure(); + } + + auto block = op.getBody(); + auto ops = block->without_terminator(); + + Value currValue = block->getArgument(0); + Value currIndex = block->getArgument(1); + Value reduceValue = block->getArgument(2); + Value reduceIndex = block->getArgument(3); + + auto opsIt = ops.begin(); + Value shouldUpdate; + if (failed(matchShouldUpdateValue(currValue, currIndex, reduceValue, + reduceIndex, opsIt, shouldUpdate))) { + return failure(); + } + + // matching: %16 = arith.select %15, %arg9, %arg11 : f32 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *opsIt << "\n"); + auto valueSelectOp = dyn_cast(*opsIt++); + if (valueSelectOp) { + if (valueSelectOp.getCondition() != shouldUpdate || + currValue != valueSelectOp.getTrueValue() || + reduceValue != valueSelectOp.getFalseValue()) { + return failure(); + } + } else { + return failure(); + } + + // matching:%17 = arith.select %15, %arg10, %arg12 : i32 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *opsIt << "\n"); + auto indexSelectOp = dyn_cast(*opsIt++); + if (indexSelectOp) { + if (indexSelectOp.getCondition() != shouldUpdate || + currIndex != indexSelectOp.getTrueValue() || + reduceIndex != indexSelectOp.getFalseValue()) { + return failure(); + } + } else { + return failure(); + } + + // matching: tt.reduce.return %16, %17 : f32, i32 + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *opsIt << "\n"); + auto termOp = dyn_cast(*opsIt++); + if (termOp && termOp == block->getTerminator()) { + auto opnds = termOp.getOperands(); + if (opnds != ArrayRef{valueSelectOp, indexSelectOp}) { + return failure(); + } + } else { + return failure(); + } + + auto loc = op.getLoc(); + + auto elemTypes = op.getElementTypes(); + + // Set the initial value of the rank-0 tensor containing + // the result value to either -inf or +inf depending on + // whether we're dealing with argmax or argmin + auto valueType = elemTypes[0]; + Value valuesAccBaseVal; + if (mlir::isa(valueType)) { + valuesAccBaseVal = rewriter.create( + loc, valueType, + rewriter.getFloatAttr(valueType, T::getBaseReductionValue())); + } else { + valuesAccBaseVal = rewriter.create( + loc, valueType, + rewriter.getIntegerAttr(valueType, T::getBaseReductionValue())); + } + + // Set the initial value of the rank-0 tensor containing the index of the + // min or max value to -1 + auto indexType = elemTypes[1]; + auto indicesAccBaseVal = rewriter.create( + loc, indexType, rewriter.getIntegerAttr(indexType, -1)); + + // Get the shape of the resulting tensors (both for values and indices). If + // we are reducing to a single scalar, then the result's type is a tensor of + // rank-0, otherwise we can reuse the original result shape + auto valueResultType = dyn_cast(op.getType(0)); + const auto isScalarReduce = valueResultType == nullptr; + SmallVector reductionResultShape{ + isScalarReduce ? SmallVector{} + : SmallVector(valueResultType.getShape())}; + + SmallVector outputs{ + getInitTensor(rewriter, reductionResultShape, valuesAccBaseVal, loc), + getInitTensor(rewriter, reductionResultShape, indicesAccBaseVal, loc)}; + + auto linalgOp = rewriter.create( + loc, adaptor.getOperands(), outputs, + SmallVector{adaptor.getAxis()}, + [&](OpBuilder &b, Location loc, ValueRange inputs) { + assert(inputs.size() == 4); + + auto tritonReduceBlock = op.getBody(); + IRMapping mapping; + mapping.map(tritonReduceBlock->getArguments(), inputs); + + for (auto &op : tritonReduceBlock->without_terminator()) { + b.clone(op, mapping); + } + + auto tritonYield = tritonReduceBlock->getTerminator(); + auto results = + llvm::map_to_vector(tritonYield->getOperands(), [&](Value val) { + return mapping.lookup(val); + }); + b.create(loc, results); + }); + + if (isScalarReduce) { + SmallVector reduceResults{ + rewriter.create( + loc, valueType, linalgOp.getResults()[0], ValueRange{}), + rewriter.create( + loc, indexType, linalgOp.getResults()[1], ValueRange{})}; + rewriter.replaceOp(op, reduceResults); + } else { + rewriter.replaceOp(op, linalgOp); + } + return success(); + } +}; + +struct ArgMaxConverter : public ArgMinMaxBaseConverter { + static LogicalResult matchComparisonResult(Value currValue, Value currIndex, + Value reduceValue, + Value reduceIndex, + mlir::Block::iterator &it, + Value &comparisonResult) { + // %14 = arith.cmpf/i ogt, %arg9, %arg11 : f32 + // This corresponds to section 2. of the sample snippet in + // ArgMinMaxBaseConverter + if (auto cmpOp = dyn_cast(*it)) { + if (cmpOp.getPredicate() != arith::CmpFPredicate::OGT || + currValue != cmpOp.getLhs() || reduceValue != cmpOp.getRhs()) { + return failure(); + } + comparisonResult = cmpOp; + } else if (auto cmpOp = dyn_cast(*it)) { + auto predicate = cmpOp.getPredicate(); + if ((predicate != arith::CmpIPredicate::sgt && predicate != arith::CmpIPredicate::ugt) || + currValue != cmpOp.getLhs() || reduceValue != cmpOp.getRhs()) { + return failure(); + } + comparisonResult = cmpOp; + } else { + return failure(); + } + it++; + + return success(); + } + + static float getBaseReductionValue() { + return -std::numeric_limits::infinity(); + } + + ArgMaxConverter(MLIRContext *context) : ArgMinMaxBaseConverter(context) {} +}; + +struct ArgMinConverter : public ArgMinMaxBaseConverter { + static LogicalResult matchComparisonResult(Value currValue, Value currIndex, + Value reduceValue, + Value reduceIndex, + mlir::Block::iterator &it, + Value &comparisonResult) { + // %14 = arith.cmpf/i olt, %arg9, %arg11 : f32 + // This corresponds to section 2. of the sample snippet in + // ArgMinMaxBaseConverter + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + + if (auto cmpOp = dyn_cast(*it)) { + if (cmpOp.getPredicate() != arith::CmpFPredicate::OLT || + currValue != cmpOp.getLhs() || reduceValue != cmpOp.getRhs()) { + return failure(); + } + comparisonResult = cmpOp; + } else if (auto cmpOp = dyn_cast(*it)) { + auto predicate = cmpOp.getPredicate(); + if ((predicate != arith::CmpIPredicate::slt && predicate != arith::CmpIPredicate::ult) || + currValue != cmpOp.getLhs() || reduceValue != cmpOp.getRhs()) { + return failure(); + } + comparisonResult = cmpOp; + } else { + return failure(); + } + it++; + + return success(); + } + + static float getBaseReductionValue() { + return std::numeric_limits::infinity(); + } + + ArgMinConverter(MLIRContext *context) : ArgMinMaxBaseConverter(context) {} +}; + +// get_program_id and get_num_programs: +// When launching triton kernels, we pass 6 additional arguments to indicate +// num_programs and program_id. Amongst those six, we have 3 arguments +// correspond to each axis for num_programs followed by 3 additional arguments +// for program_id. +// +// For instance, with triton kernel example_kernel(a, b, c), we have: +// example_kernel( +// a, b, c, +// num_programs_axis_0, +// num_programs_axis_1, +// num_programs_axis_2, +// program_id_axis_0, +// program_id_axis_1, +// program_id_axis_2, +// ) +// +struct GetProgramIDConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + static uint32_t constexpr LAUNCH_GRID_RANK = + getMaxEnumValForProgramIDDim() + 1; + +public: + LogicalResult + matchAndRewrite(triton::GetProgramIdOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto axis = (uint32_t)op.getAxis(); + assert(axis < LAUNCH_GRID_RANK && "program_id expects " + "axis to be either 0, " + "1, or 2"); + + auto func = op->getParentOfType(); + auto numArgs = func.getNumArguments(); + auto id = func.getArgument(numArgs - LAUNCH_GRID_RANK + axis); + + rewriter.replaceOp(op, id); + return success(); + } +}; + +struct GetNumProgramsConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +private: + static uint32_t constexpr LAUNCH_GRID_RANK = + getMaxEnumValForProgramIDDim() + 1; + +public: + GetNumProgramsConverter(MLIRContext *context) + : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(triton::GetNumProgramsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto axis = (uint32_t)op.getAxis(); + assert(axis < LAUNCH_GRID_RANK && "program_id expects " + "axis to be either 0, " + "1, or 2"); + + auto func = op->getParentOfType(); + auto numArgs = func.getNumArguments(); + auto id = func.getArgument(numArgs - LAUNCH_GRID_RANK * 2 + axis); + + rewriter.replaceOp(op, id); + return success(); + } +}; + +// Convert a pair of cmpf and select to either min or max. +// Leave the pattern as simple as possible because triton has plans to emit +// min and max directly. +template +struct MinMaxConverter : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + MinMaxConverter(MLIRContext *context) + : OpRewritePattern(context, /*benefit=*/10) {} + + LogicalResult matchAndRewrite(CmpOp cmpOp, + PatternRewriter &rewriter) const final { + if (!cmpOp.getResult().hasOneUse()) { + return failure(); + } + auto selectOp = + dyn_cast(*cmpOp.getResult().getUsers().begin()); + if (!selectOp) { + return failure(); + } + + if (!(cmpOp.getResult() == selectOp.getCondition() && + cmpOp.getLhs() == selectOp.getTrueValue() && + cmpOp.getRhs() == selectOp.getFalseValue())) { + return failure(); + } + + rewriteOpWithMinMax(rewriter, cmpOp, selectOp, cmpOp.getPredicate()); + rewriter.eraseOp(cmpOp); + + return success(); + } + + void rewriteOpWithMinMax(PatternRewriter &rewriter, arith::CmpFOp cmpOp, + arith::SelectOp selectOp, + arith::CmpFPredicate pred) const { + switch (pred) { + case arith::CmpFPredicate::OGT: + case arith::CmpFPredicate::OGE: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + case arith::CmpFPredicate::OLT: + case arith::CmpFPredicate::OLE: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + default: + llvm_unreachable("Unhandled predicate"); + } + } + + void rewriteOpWithMinMax(PatternRewriter &rewriter, arith::CmpIOp cmpOp, + arith::SelectOp selectOp, + arith::CmpIPredicate pred) const { + switch (pred) { + case arith::CmpIPredicate::sgt: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + case arith::CmpIPredicate::ugt: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + case arith::CmpIPredicate::slt: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + case arith::CmpIPredicate::ult: + rewriter.replaceOpWithNewOp(selectOp, cmpOp.getLhs(), + cmpOp.getRhs()); + break; + default: + llvm_unreachable("Unhandled predicate"); + } + } +}; + +struct DenseConstantConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult + matchAndRewrite(arith::ConstantOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto attr = cast(op.getValue()); + auto loc = op.getLoc(); + + auto splatConst = arith::ConstantOp::materialize( + rewriter, attr.getSplatValue(), attr.getElementType(), loc); + + auto init = rewriter.create( + loc, cast(op.getResult().getType()).getShape(), + attr.getElementType()); + + rewriter.replaceOpWithNewOp(op, ValueRange{splatConst}, + ValueRange{init}); + + return success(); + } +}; + +class CumSumConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + // CumSum is a specific instance of Scan that looks like the following: + // %1 = "tt.scan"(%0) <{axis = 1 : i32}> ({ + // ^bb0(%arg0: f32, %arg1: f32): + // %2 = arith.addf %arg0, %arg1 : f32 + // tt.scan.return %2 : f32 + // }) : (tensor<4x4xf32>) -> tensor<4x4xf32> + bool isCumSum(triton::ScanOp op) const { + auto scanBlock = op.getBody(); + auto ops = llvm::map_to_vector(scanBlock->without_terminator(), + [](Operation &op) { return &op; }); + + if (ops.size() != 1) { + return false; + } + + auto addOp = ops.front(); + if (isa(addOp)) { + if (addOp->getResult(0) != scanBlock->getTerminator()->getOperand(0)) { + return false; + } + + auto blockArgs = + llvm::map_range(scanBlock->getArguments(), [](BlockArgument arg) { + return dyn_cast(arg); + }); + + auto addArgs = addOp->getOperands(); + + return DenseSet(blockArgs.begin(), blockArgs.end()) == + DenseSet(addArgs.begin(), addArgs.end()); + } + + return false; + } + +public: + LogicalResult + matchAndRewrite(triton::ScanOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!isCumSum(op)) { + return rewriter.notifyMatchFailure( + op, "Only support cumsum variant of scan op"); + } + + auto input = op.getOperand(0); + auto axis = op.getAxis(); + auto type = dyn_cast(input.getType()); + + if (type.getRank() != 1 && type.getRank() != 2 && + axis != type.getRank() - 1) { + return rewriter.notifyMatchFailure( + op, "Only support lowering scan op to cumsum with rank " + "= {1, 2} and axis = rank - 1"); + } + + Value init = rewriter.create(op.getLoc(), type.getShape(), + type.getElementType()); + + rewriter.replaceOpWithNewOp( + op, input, rewriter.getUI32IntegerAttr(axis), init); + + return success(); + } +}; + +class AddPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::AddPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto resType = op.getResult().getType(); + assert(isa(resType)); + auto rank = cast(resType).getRank(); + SmallVector indexingMaps( + /*numResult + numOperands*/ 3, rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + SmallVector outputs = {op.getPtr()}; + rewriter.replaceOpWithNewOp( + op, op->getResultTypes(), op->getOperands(), outputs, indexingMaps, + iteratorTypes, + [&](OpBuilder &builder, Location loc, ValueRange regionArgs) { + auto resultTypes = llvm::map_to_vector( + op->getResultTypes(), [](Type type) { + return cast(type).getElementType(); + }); + auto *scalarOp = + builder.create(loc, op->getName().getIdentifier(), + regionArgs.take_front(op->getNumOperands()), + resultTypes, op->getAttrs()); + builder.create(loc, scalarOp->getResults()); + }); + return success(); + } +}; + +// Convert triton op X operating on tensors of pointers to a linalg.generic +// wrapping op X to operate on single pointer. +// This pattern rewriter is almost identical to AddPtrConverter above, except +// that the out param for the linalg op is an empty op instead of reusing one +// of the existing operands. This is because depending on the templatized op, +// the type of the operands might be different, so we cannot pick a default +// operand to reuse for all cases. +template +class TensorOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(OpType op, typename OpType::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto resultTensorType = dyn_cast(op.getResult().getType()); + if (!resultTensorType) { + return failure(); + } + auto rank = resultTensorType.getRank(); + SmallVector indexingMaps( + /*numResult + numOperands*/ op->getNumResults() + op->getNumOperands(), + rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + SmallVector outputs = {rewriter.create( + op->getLoc(), resultTensorType.getShape(), + resultTensorType.getElementType())}; + rewriter.replaceOpWithNewOp( + op, op->getResultTypes(), op->getOperands(), outputs, indexingMaps, + iteratorTypes, + [&](OpBuilder &builder, Location loc, ValueRange regionArgs) { + auto resultTypes = llvm::map_to_vector( + op->getResultTypes(), [](Type type) { + return cast(type).getElementType(); + }); + auto *scalarOp = + builder.create(loc, op->getName().getIdentifier(), + regionArgs.take_front(op->getNumOperands()), + resultTypes, op->getAttrs()); + builder.create(loc, scalarOp->getResults()); + }); + return success(); + } +}; + +// Convert triton store op operating on tensors of pointers to a linalg.generic +// wrapping op a triton store op on single pointer. +// Note that this linalg.generic op has an empty `out` param. +class StorePtrToLinalgConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto storeTensorType = dyn_cast(op.getValue().getType()); + if (!storeTensorType) { + return failure(); + } + auto rank = storeTensorType.getRank(); + SmallVector indexingMaps( + /*numResult + numOperands*/ op->getNumResults() + op.getNumOperands(), + rewriter.getMultiDimIdentityMap(rank)); + SmallVector iteratorTypes( + rank, utils::IteratorType::parallel); + SmallVector outputs; + rewriter.replaceOpWithNewOp( + op, op->getResultTypes(), op->getOperands(), outputs, indexingMaps, + iteratorTypes, + [&](OpBuilder &builder, Location loc, ValueRange regionArgs) { + auto resultTypes = llvm::map_to_vector(op->getResultTypes(), [](Type type) { + return cast(type).getElementType(); + }); + auto *scalarOp = + builder.create(loc, op->getName().getIdentifier(), + regionArgs.take_front(op->getNumOperands()), + resultTypes, op->getAttrs()); + builder.create(loc, scalarOp->getResults()); + }); + return success(); + } +}; + +class ReshapeConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::ReshapeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto input = op.getSrc(); + auto output = op.getResult(); + + auto inputType = input.getType(); + auto outputType = output.getType(); + if (!outputType.hasStaticShape()) { + return failure(); + } + + if (auto maybeReassociationMap = + getReassociationIndicesForReshape(inputType, outputType)) { + auto reassociationMap = *maybeReassociationMap; + if (outputType.getRank() < inputType.getRank()) { + rewriter.replaceOpWithNewOp( + op, outputType, input, reassociationMap); + } else { + rewriter.replaceOpWithNewOp( + op, outputType, input, reassociationMap); + } + return success(); + } + + ArrayRef outputShape = outputType.getShape(); + + auto shape = rewriter.create( + loc, rewriter.getI64TensorAttr(outputShape)); + rewriter.replaceOpWithNewOp(op, outputType, input, + shape); + + return success(); + } +}; + +class ExternElementwiseBinaryOpConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::ExternElementwiseOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + if (!op.getPure() || op.getSrcs().size() != 2) + return failure(); +#define POPULATE_BINARY_OP(FUNC_NAME, DST_OP) \ + if (!op.getSymbol().compare(FUNC_NAME)) { \ + rewriter.replaceOpWithNewOp(op, op.getSrcs()[0], op.getSrcs()[1]); \ + return success(); \ + } + + POPULATE_BINARY_OP("__nv_atan2f", math::Atan2Op); + POPULATE_BINARY_OP("__nv_atan2", math::Atan2Op); + POPULATE_BINARY_OP("__nv_powf", math::PowFOp); + POPULATE_BINARY_OP("__nv_pow", math::PowFOp); + POPULATE_BINARY_OP("fmod", mathext::FModOp); + POPULATE_BINARY_OP("powf", math::PowFOp); + POPULATE_BINARY_OP("div_rn", arith::DivFOp); + POPULATE_BINARY_OP("div_rz", mathext::DivRzOp); + POPULATE_BINARY_OP("atan2", math::Atan2Op); + +#undef POPULATE_BINARY_OP + return failure(); + } +}; + +class ExternElementwiseUnaryOpConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + +public: + LogicalResult + matchAndRewrite(triton::ExternElementwiseOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + if (!op.getPure() || op.getSrcs().size() != 1) + return failure(); +#define POPULATE_UNARY_OP(FUNC_NAME, DST_OP) \ + if (!op.getSymbol().compare(FUNC_NAME)) { \ + rewriter.replaceOpWithNewOp(op, op.getSrcs()[0]); \ + return success(); \ + } + + POPULATE_UNARY_OP("isnan", math::IsNaNOp); + POPULATE_UNARY_OP("isinf", math::IsInfOp); + POPULATE_UNARY_OP("isfinite", math::IsFiniteOp); + POPULATE_UNARY_OP("fabsf", math::AbsFOp); + POPULATE_UNARY_OP("fabs", math::AbsFOp); + POPULATE_UNARY_OP("sinf", math::SinOp); + POPULATE_UNARY_OP("sin", math::SinOp); + POPULATE_UNARY_OP("cosf", math::CosOp); + POPULATE_UNARY_OP("cos", math::CosOp); + POPULATE_UNARY_OP("tanf", math::TanOp); + POPULATE_UNARY_OP("tan", math::TanOp); + POPULATE_UNARY_OP("asinf", math::AsinOp); + POPULATE_UNARY_OP("asin", math::AsinOp); + POPULATE_UNARY_OP("acosf", math::AcosOp); + POPULATE_UNARY_OP("acos", math::AcosOp); + POPULATE_UNARY_OP("atanf", math::AtanOp); + POPULATE_UNARY_OP("atan", math::AtanOp); + POPULATE_UNARY_OP("sinhf", math::SinhOp); + POPULATE_UNARY_OP("sinh", math::SinhOp); + POPULATE_UNARY_OP("coshf", math::CoshOp); + POPULATE_UNARY_OP("cosh", math::CoshOp); + POPULATE_UNARY_OP("tanhf", math::TanhOp); + POPULATE_UNARY_OP("tanh", math::TanhOp); + POPULATE_UNARY_OP("acoshf", math::AcoshOp); + POPULATE_UNARY_OP("acosh", math::AcoshOp); + POPULATE_UNARY_OP("asinhf", math::AsinhOp); + POPULATE_UNARY_OP("asinh", math::AsinhOp); + POPULATE_UNARY_OP("atanhf", math::AtanhOp); + POPULATE_UNARY_OP("atanhf", math::AtanhOp); + POPULATE_UNARY_OP("logf", math::LogOp); + POPULATE_UNARY_OP("log", math::LogOp); + POPULATE_UNARY_OP("log10f", math::Log10Op); + POPULATE_UNARY_OP("log10", math::Log10Op); + POPULATE_UNARY_OP("log1pf", math::Log1pOp); + POPULATE_UNARY_OP("log1p", math::Log1pOp); + POPULATE_UNARY_OP("expf", math::ExpOp); + POPULATE_UNARY_OP("exp", math::ExpOp); + POPULATE_UNARY_OP("exp2f", math::Exp2Op); + POPULATE_UNARY_OP("exp2", math::Exp2Op); + POPULATE_UNARY_OP("erff", math::ErfOp); + POPULATE_UNARY_OP("erf", math::ErfOp); + POPULATE_UNARY_OP("sqrtf", math::SqrtOp); + POPULATE_UNARY_OP("sqrt", math::SqrtOp); + POPULATE_UNARY_OP("rsqrtf", math::RsqrtOp); + POPULATE_UNARY_OP("rsqrt", math::RsqrtOp); + POPULATE_UNARY_OP("ceilf", math::CeilOp); + POPULATE_UNARY_OP("ceil", math::CeilOp); + POPULATE_UNARY_OP("floorf", math::FloorOp); + POPULATE_UNARY_OP("floor", math::FloorOp); + POPULATE_UNARY_OP("truncf", math::TruncOp); + POPULATE_UNARY_OP("trunc", math::TruncOp); + + POPULATE_UNARY_OP("__nv_fabsf", math::AbsFOp); + POPULATE_UNARY_OP("__nv_fabs", math::AbsFOp); + POPULATE_UNARY_OP("__nv_sinf", math::SinOp); + POPULATE_UNARY_OP("__nv_sin", math::SinOp); + POPULATE_UNARY_OP("__nv_cosf", math::CosOp); + POPULATE_UNARY_OP("__nv_cos", math::CosOp); + POPULATE_UNARY_OP("__nv_tanf", math::TanOp); + POPULATE_UNARY_OP("__nv_tan", math::TanOp); + POPULATE_UNARY_OP("__nv_asinf", math::AsinOp); + POPULATE_UNARY_OP("__nv_asin", math::AsinOp); + POPULATE_UNARY_OP("__nv_acosf", math::AcosOp); + POPULATE_UNARY_OP("__nv_acos", math::AcosOp); + POPULATE_UNARY_OP("__nv_atanf", math::AtanOp); + POPULATE_UNARY_OP("__nv_atan", math::AtanOp); + POPULATE_UNARY_OP("__nv_sinhf", math::SinhOp); + POPULATE_UNARY_OP("__nv_sinh", math::SinhOp); + POPULATE_UNARY_OP("__nv_coshf", math::CoshOp); + POPULATE_UNARY_OP("__nv_cosh", math::CoshOp); + POPULATE_UNARY_OP("__nv_tanhf", math::TanhOp); + POPULATE_UNARY_OP("__nv_tanhf", math::TanhOp); + POPULATE_UNARY_OP("__nv_acoshf", math::AcoshOp); + POPULATE_UNARY_OP("__nv_acosh", math::AcoshOp); + POPULATE_UNARY_OP("__nv_asinhf", math::AsinhOp); + POPULATE_UNARY_OP("__nv_asinh", math::AsinhOp); + POPULATE_UNARY_OP("__nv_atanhf", math::AtanhOp); + POPULATE_UNARY_OP("__nv_atanhf", math::AtanhOp); + POPULATE_UNARY_OP("__nv_logf", math::LogOp); + POPULATE_UNARY_OP("__nv_log", math::LogOp); + POPULATE_UNARY_OP("__nv_log10f", math::Log10Op); + POPULATE_UNARY_OP("__nv_log10", math::Log10Op); + POPULATE_UNARY_OP("__nv_log1pf", math::Log1pOp); + POPULATE_UNARY_OP("__nv_log1p", math::Log1pOp); + POPULATE_UNARY_OP("__nv_expf", math::ExpOp); + POPULATE_UNARY_OP("__nv_exp", math::ExpOp); + POPULATE_UNARY_OP("__nv_exp2f", math::Exp2Op); + POPULATE_UNARY_OP("__nv_exp2", math::Exp2Op); + POPULATE_UNARY_OP("__nv_erff", math::ErfOp); + POPULATE_UNARY_OP("__nv_erf", math::ErfOp); + POPULATE_UNARY_OP("__nv_sqrtf", math::SqrtOp); + POPULATE_UNARY_OP("__nv_sqrt", math::SqrtOp); + POPULATE_UNARY_OP("__nv_rsqrtf", math::RsqrtOp); + POPULATE_UNARY_OP("__nv_rsqrt", math::RsqrtOp); + POPULATE_UNARY_OP("__nv_ceilf", math::CeilOp); + POPULATE_UNARY_OP("__nv_ceil", math::CeilOp); + POPULATE_UNARY_OP("__nv_floorf", math::FloorOp); + POPULATE_UNARY_OP("__nv_floor", math::FloorOp); + POPULATE_UNARY_OP("__nv_truncf", math::TruncOp); + POPULATE_UNARY_OP("__nv_trunc", math::TruncOp); + +#undef POPULATE_UNARY_OP + return failure(); + } +}; + +static void populateExternElementwiseOpToMLIROps(RewritePatternSet &patterns) { + patterns.add(patterns.getContext()); +} + +} // namespace + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns_FlagTree.hpp b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns_FlagTree.hpp new file mode 100755 index 00000000..ef227404 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns_FlagTree.hpp @@ -0,0 +1,261 @@ +#ifndef TRITON_CONVERSION_PATTERNS_FLAGTREE +#define TRITON_CONVERSION_PATTERNS_FLAGTREE + +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Some functions in this hpp file are partially borrowed from the +// open-source project "Triton-Linalg", with the address: +// https://github.com/Cambricon/triton-linalg. +// * The original author's Copyright (C) [2022-2025] by Cambricon. +// * Several modifications have been made based on this. +// +// +//===----------------------------------------------------------------------===// + +#include "triton/Dialect/Triton/IR/Dialect.h" + +using namespace mlir; + +namespace { + +// Copyright (C) [2022-2025] by Cambricon. +// flagtree: Extract the first slice along a specified dimension from +// input tensor. This function creates a tensor slice containing only +// the first element along the given dimension +Value sliceFirst(ConversionPatternRewriter &rewriter, Location loc, Value input, + int64_t dim, bool reverse = false) { + ShapedType inputType = cast(input.getType()); + auto sizes = + llvm::to_vector(llvm::map_range(inputType.getShape(), [&](int64_t t) { + return OpFoldResult(rewriter.getI64IntegerAttr(t)); + })); + int64_t rank = inputType.getRank(); + // Retrieve slice offsets of input. + SmallVector offsets(rank, rewriter.getIndexAttr(0)); + if (reverse) + offsets[dim] = rewriter.getIndexAttr(inputType.getDimSize(dim) - 1); + // Retrieve slice sizes of input. + sizes[dim] = rewriter.getIndexAttr(1); + // Retrieve slice strides of input. + SmallVector strides(rank, rewriter.getIndexAttr(1)); + // Create the slice of input. + return rewriter.create(loc, input, offsets, sizes, + strides); +} + +// Copyright (C) [2022-2025] by Cambricon. +// flagtree: Extract the remaining slices (excluding first) along a specified +// dimension from input tensor. This function creates a tensor slice containing +// all elements except the first along the given dimension +Value sliceRemaining(ConversionPatternRewriter &rewriter, Location loc, + Value input, int64_t dim, bool reverse = false) { + ShapedType inputType = cast(input.getType()); + auto sizes = + llvm::to_vector(llvm::map_range(inputType.getShape(), [&](int64_t t) { + return OpFoldResult(rewriter.getI64IntegerAttr(t)); + })); + int64_t rank = inputType.getRank(); + // Retrieve slice sizes of input. + sizes[dim] = rewriter.getIndexAttr(inputType.getDimSize(dim) - 1); + // Retrieve slice offsets of input. + SmallVector offsets(rank, rewriter.getIndexAttr(0)); + if (!reverse) + offsets[dim] = rewriter.getIndexAttr(1); + // Retrieve slice strides of input. + SmallVector strides(rank, rewriter.getIndexAttr(1)); + // Create the slice of input. + return rewriter.create(loc, input, offsets, sizes, + strides); +} + +// Copyright (C) [2022-2025] by Cambricon. +// flagtree: Create reassociation maps for tensor reshape operations +// between expanded and collapsed shapes. This function generates the +// mapping needed for tensor.collapse_shape operations +bool createReassociationMaps( + OpBuilder &builder, llvm::ArrayRef expandedShape, + llvm::ArrayRef collapsedShape, + llvm::SmallVector &reassociationMap) { + if (collapsedShape.empty()) { + reassociationMap = {}; + return true; + } + + // As tensor.expand_shape/tensor.collapse_shape expected rank + // expansion/reduction. + if (expandedShape.size() == collapsedShape.size()) + return false; + if (ShapedType::isDynamicShape(expandedShape) || + ShapedType::isDynamicShape(collapsedShape)) + return false; + // flagtree: Initialize reassociation map with size equal to + // collapsed dimensions + reassociationMap.resize(collapsedShape.size()); + unsigned currExpandDim = 0, currCollapseDim = 0; + // flagtree: Iterate through dimensions to create mapping between + // expanded and collapsed shapes + while (currExpandDim < expandedShape.size() && + currCollapseDim < collapsedShape.size()) { + int64_t dstSize = collapsedShape[currCollapseDim]; + int64_t srcSize = expandedShape[currExpandDim]; + + // flagtree: Accumulate dimensions until we match the target + // collapsed dimension size + while (srcSize < dstSize && currExpandDim < expandedShape.size()) { + reassociationMap[currCollapseDim].push_back( + builder.getAffineDimExpr(currExpandDim++)); + srcSize *= expandedShape[currExpandDim]; + } + if (srcSize == dstSize) { + reassociationMap[currCollapseDim].push_back( + builder.getAffineDimExpr(currExpandDim++)); + // If the next dim in collapsedShape is not 1, treat subsequent dims in + // expandedShape which are 1 to be collapsed. + if (currCollapseDim == collapsedShape.size() - 1 || + collapsedShape[currCollapseDim + 1] != 1) { + while (currExpandDim < expandedShape.size() && + expandedShape[currExpandDim] == 1) { + reassociationMap[currCollapseDim].push_back( + builder.getAffineDimExpr(currExpandDim++)); + } + } + } + // If the reassociationMap for the currCollapseDim is empty, clear all + // mappings and return false. + if (reassociationMap[currCollapseDim].empty()) { + reassociationMap.clear(); + return false; + } + currCollapseDim++; + } + // If both iterators didn't reach the end, we have leftover dimentions which + // implies that we have a mismatch in shape. + return currExpandDim == expandedShape.size() && + currCollapseDim == collapsedShape.size(); +} + +// flagtree: Lower tt.reduce to linalg.reduce, by initialization with the +// first element. +LogicalResult applyLinalgReduce(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) { + Location loc = op->getLoc(); + + // Derive types. `tt.reduce` treats reducing a 1-D tensor with a special + // case that returns a scalar, but we treat it as a 0-D tensor in these + // types. + auto convertedInputTensorTypes = + llvm::map_range(adaptor.getOperands().getTypes(), + [](Type t) { return cast(t); }); + assert(llvm::all_equal(llvm::map_range( + convertedInputTensorTypes, [](TensorType t) { return t.getShape(); }))); + static_cast(convertedInputTensorTypes); + + auto originalResultTensorTypes = + llvm::map_range(op.getResultTypes(), [](Type t) -> TensorType { + if (auto tensorType = dyn_cast(t)) + return tensorType; + return RankedTensorType::get({}, t); + }); + assert(llvm::all_equal(llvm::map_range( + originalResultTensorTypes, [](TensorType t) { return t.getShape(); }))); + ArrayRef resultShape = + (*originalResultTensorTypes.begin()).getShape(); + auto convertedResultTensorTypes = + llvm::map_range(originalResultTensorTypes, [&](TensorType t) { + return RankedTensorType::get(resultShape, t.getElementType()); + }); + + llvm::SmallVector initVals; + llvm::SmallVector inputVals; + // To lowering to linalg.reduce, we use the first slice of the reduction + // axis of input operands as the init value of init operands. And then, + // reduce the remaining elements of input operands. + // We assume that the number of input operands is same as init operands and + // corresponds one to one. + // TODO: This restriction will need to be relaxed in the future. + + assert(adaptor.getOperands().size() == op.getNumResults() && + "tt.reduce requires the same input number and init number"); + for (auto [inputVal, initTy] : + llvm::zip(adaptor.getOperands(), convertedResultTensorTypes)) { + ShapedType inputTy = cast(inputVal.getType()); + ArrayRef inputShape = inputTy.getShape(); + + // If the size of reduce axis is 1, we will replace init operands by input + // operands, so we should resize the input operands' shape by init + // operands. + if (inputShape[op.getAxis()] <= 1) { + assert(inputVals.empty() && + "tt.reduce requires the same shape of all input operands"); + SmallVector reassociationMap; + [[maybe_unused]] bool res = createReassociationMaps( + rewriter, inputShape, initTy.getShape(), reassociationMap); + assert(res && "attempting to collapse into an incompatible shape"); + auto collapse = rewriter.create( + loc, inputVal, reassociationMap); + initVals.push_back(collapse); + continue; + } + + // 1. Slice the first elements of input operands, and use them as init + // operands' init value. + { + Value slice = sliceFirst(rewriter, loc, inputVal, op.getAxis()); + auto sliceShape = cast(slice.getType()).getShape(); + + // Resize slice value's shape by init operand. + SmallVector reassociationMap; + [[maybe_unused]] bool res = createReassociationMaps( + rewriter, sliceShape, initTy.getShape(), reassociationMap); + assert(res && "attempting to collapse into an incompatible shape"); + auto collapse = rewriter.create( + loc, slice, reassociationMap); + initVals.push_back(collapse); + } + // 2. Slice the remaining elements of input operands, reduce them and + // init value. + { + Value slice = sliceRemaining(rewriter, loc, inputVal, op.getAxis()); + inputVals.push_back(slice); + } + } + + // If the results are scalar, we need to extract the scalar from the + // 0-ranked result tensor. + auto getFinalResults = [&](ValueRange results) -> SmallVector { + if (!resultShape.empty()) + return results; + SmallVector extractResults; + for (auto [tensor, type] : llvm::zip(results, convertedResultTensorTypes)) { + Value scalar = rewriter.create( + loc, type.getElementType(), tensor, /*indices=*/ValueRange{}); + extractResults.push_back(scalar); + } + return extractResults; + }; + + // If the the size of reduce axis is 1, we just replace the init operands by + // input operands. + if (inputVals.empty()) { + rewriter.replaceOp(op, getFinalResults(initVals)); + return success(); + } + + // Create a linalg.reduce on the same input and move the combine region + // there. (ReduceReturnOpConversion will take care of the terminator.) + auto reduceOp = rewriter.create( + loc, /*resultTypes=*/SmallVector(convertedResultTensorTypes), + /*inputs=*/inputVals, /*inits=*/initVals, + /*dimensions=*/ArrayRef{op.getAxis()}); + rewriter.inlineRegionBefore(op.getCombineOp(), reduceOp.getCombiner(), + reduceOp.getCombiner().end()); + + rewriter.replaceOp(op, getFinalResults(reduceOp.getResults())); + return success(); +} + +} // namespace + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/Passes.h new file mode 100755 index 00000000..b95cbde7 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_ARITH_TO_LINALG_CONVERSION_PASSES_H +#define TRITON_ARITH_TO_LINALG_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/Passes.td new file mode 100755 index 00000000..8678cca9 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/Passes.td @@ -0,0 +1,22 @@ +#ifndef TRITON_ARITH_TO_LINALG_CONVERSION_PASSES +#define TRITON_ARITH_TO_LINALG_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonArithToLinalg : Pass<"triton-arith-to-linalg", "mlir::ModuleOp"> { + let summary = "Convert Triton arithmetic operations into linalg"; + let options = [ + Option<"pidsToFuncArgs", "pids-to-func-args", "bool", /*default*/"true", + "Convert tt.get_program_id and tt.get_num_programs to reference to function arguments">, + Option<"ttToFuncFunc", "tt-to-func-func", "bool", /*default*/"true", + "Convert tt.func to func.func">, + Option<"addptrToLinalg", "addptr-to-linalg", "bool", /*default*/"true", + "Convert tt.addptr on tensors to linalg">, + Option<"assertToCf", "assert-to-cf", "bool", /*default*/"true", + "Convert tt.assert to cf.assert">, + Option<"tensorPtrToLinalg", "tensor-ptr-to-linalg", "bool", /*default*/"false", + "Convert triton ops on tensor of pointers to linalg.generic">, + ]; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h new file mode 100755 index 00000000..6df8d4ca --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h @@ -0,0 +1,32 @@ +#ifndef TRITON_CONVERSION_TRITONARITHTOLINALG_TRITONARITHTOLINALG_H +#define TRITON_CONVERSION_TRITONARITHTOLINALG_TRITONARITHTOLINALG_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_DECL +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +void populateTritonArithToLinalgCanonicalizationPatterns( + RewritePatternSet &patterns); + +void populateTritonArithToLinalgConversionPatterns(bool pidsToFuncArgs, + bool addptrToLinalg, + bool assertToCf, + RewritePatternSet &patterns); + +// Expand the triton pointer ops operating on pointers to linalg +void populateTritonTensorPtrConversionPatterns(RewritePatternSet &patterns); + +std::unique_ptr> +createTritonArithToLinalgPass(bool tensorPtrToLinalg = false); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONARITHTOLINALG_TRITONARITHTOLINALG_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/CMakeLists.txt new file mode 100755 index 00000000..07f9ad33 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonPtrToMemref) +add_public_tablegen_target(TritonPtrToMemrefConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/Passes.h new file mode 100755 index 00000000..e1f6f33b --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_PTR_TO_MEMREF_CONVERSION_PASSES_H +#define TRITON_PTR_TO_MEMREF_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonPtrToMemref/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/Passes.td new file mode 100755 index 00000000..c027b098 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/Passes.td @@ -0,0 +1,11 @@ +#ifndef TRITON_PTR_TO_MEMREF_CONVERSION_PASSES +#define TRITON_PTR_TO_MEMREF_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonPtrToMemref : Pass<"triton-ptr-to-memref", "mlir::ModuleOp"> { + let summary = "Convert triton pointer to unranked memref"; + let constructor = "triton::createTritonPtrToMemrefPass()"; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h new file mode 100755 index 00000000..4476f7d6 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h @@ -0,0 +1,17 @@ +#ifndef TRITON_CONVERSION_TRITON_PTR_TO_MEMREF_TRITON_PTR_TO_MEMREF_H +#define TRITON_CONVERSION_TRITON_PTR_TO_MEMREF_TRITON_PTR_TO_MEMREF_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonPtrToMemrefPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITON_PTR_TO_MEMREF_TRITON_PTR_TO_MEMREF_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/CMakeLists.txt new file mode 100755 index 00000000..74ccdd39 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/CMakeLists.txt @@ -0,0 +1,9 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Triton Project Contributors. +# +#===------------------------------------------------------------------------===# + +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToLinalg) +add_public_tablegen_target(TritonToLinalgConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/Passes.h new file mode 100755 index 00000000..404af080 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/Passes.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_TO_LINALG_CONVERSION_PASSES_H +#define TRITON_TO_LINALG_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonToLinalg/TritonToLinalg.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonToLinalg/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/Passes.td new file mode 100755 index 00000000..627077e3 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/Passes.td @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_TO_LINALG_CONVERSION_PASSES +#define TRITON_TO_LINALG_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToLinalg : Pass<"triton-to-linalg", "mlir::ModuleOp"> { + let summary = "Convert Triton to Linalg dialect"; + let constructor = "triton::createTritonToLinalgPass()"; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/TritonToLinalg.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/TritonToLinalg.h new file mode 100755 index 00000000..4c58e992 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/TritonToLinalg.h @@ -0,0 +1,33 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_TRITONTOLINALG_TRITONTOLINALG_H +#define TRITON_CONVERSION_TRITONTOLINALG_TRITONTOLINALG_H + +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonToLinalgPass(); + +void populateTritonToLinalgCanonicalizationPatterns( + RewritePatternSet &patterns); + +void populateTritonToLinalgConversionPatterns(TypeConverter &typeConverter, + RewritePatternSet &patterns, + unsigned int launchGridRank); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONTOLINALG_TRITONTOLINALG_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/CMakeLists.txt new file mode 100755 index 00000000..e38329b1 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/CMakeLists.txt @@ -0,0 +1,9 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Triton Project Contributors. +# +#===------------------------------------------------------------------------===# + +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToLinalgExperimental) +add_public_tablegen_target(TritonToLinalgExperimentalConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/Passes.h new file mode 100755 index 00000000..73b6c5f2 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/Passes.h @@ -0,0 +1,24 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_TO_LINALG_EXPERIMENTAL_CONVERSION_PASSES_H +#define TRITON_TO_LINALG_EXPERIMENTAL_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimental.h" +#include "triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h" +#include "triton-shared/Conversion/TritonToLinalgExperimental/TritonToPtr.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonToLinalgExperimental/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/Passes.td new file mode 100755 index 00000000..048f13f8 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/Passes.td @@ -0,0 +1,23 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_TO_LINALG_EXPERIMENTAL_CONVERSION_PASSES +#define TRITON_TO_LINALG_EXPERIMENTAL_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToLinalgExperimental : Pass<"triton-to-linalg-experimental", "mlir::ModuleOp"> { + let summary = "Convert Triton to Linalg dialect"; + let constructor = "triton::createTritonToLinalgExperimentalPass()"; +} + +def TritonToPtr : Pass<"triton-to-ptr", "mlir::ModuleOp"> { + let summary = "Convert Triton ops on pointers to the Ptr dialect"; + let constructor = "triton::createTritonToPtrPass()"; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimental.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimental.h new file mode 100755 index 00000000..59b7b230 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimental.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_TRITONTOLINALG_TRITONTOLINALGEXPERIMENTAL_H +#define TRITON_CONVERSION_TRITONTOLINALG_TRITONTOLINALGEXPERIMENTAL_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonToLinalgExperimentalPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONTOLINALG_TRITONTOLINALGEXPERIMENTAL_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/TritonToPtr.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/TritonToPtr.h new file mode 100755 index 00000000..305d5aa1 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/TritonToPtr.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_TRITONTOLINALG_TRITONTOPTR_H +#define TRITON_CONVERSION_TRITONTOLINALG_TRITONTOPTR_H + +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonToPtrPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONTOLINALG_TRITONTOPTR_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/CMakeLists.txt new file mode 100755 index 00000000..5762c1f6 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToStructured) +add_public_tablegen_target(TritonToStructuredConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/Passes.h new file mode 100755 index 00000000..3c3b81ca --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_TO_STRUCTURED_CONVERSION_PASSES_H +#define TRITON_TO_STRUCTURED_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonToStructured/TritonToStructured.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonToStructured/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/Passes.td new file mode 100755 index 00000000..79eb6de5 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/Passes.td @@ -0,0 +1,19 @@ +#ifndef TRITON_TO_STRUCTURED_CONVERSION_PASSES +#define TRITON_TO_STRUCTURED_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToStructured : Pass<"triton-to-structured", "mlir::ModuleOp"> { + let summary = "Convert Triton non-block pointer to TritonStructured dialect"; + let constructor = "triton::createWaferTritonToStructuredPass()"; + let options = [ + Option<"runPrepassOnly", "run-prepass-only", "bool", /*default*/"false", + "Only run the pre-processing pass which inserts tts.get_structured_state ops used in scf.for">, + Option<"skipPrepass", "skip-prepass", "bool", /*default*/"false", + "Skip the prepass">, + Option<"useUnsafeMask", "use-unsafe-mask", "bool", /*default*/"false", + "Assume that the mask bounds are never less than starting offsets. May produce incorrect results."> + ]; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/TritonToStructured.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/TritonToStructured.h new file mode 100755 index 00000000..2f31dfcd --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/TritonToStructured.h @@ -0,0 +1,17 @@ +#ifndef TRITON_CONVERSION_TRITONTOSTRUCTURED_TRITONTOSTRUCTURED_H +#define TRITON_CONVERSION_TRITONTOSTRUCTURED_TRITONTOSTRUCTURED_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createWaferTritonToStructuredPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITONTOSTRUCTURED_TRITONTOSTRUCTURED_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/CMakeLists.txt new file mode 100755 index 00000000..116a3e3f --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/CMakeLists.txt @@ -0,0 +1,3 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name TritonToUnstructured) +add_public_tablegen_target(TritonToUnstructuredConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/Passes.h new file mode 100755 index 00000000..a2016c7a --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/Passes.h @@ -0,0 +1,15 @@ +#ifndef TRITON_TO_UNSTRUCTURED_CONVERSION_PASSES_H +#define TRITON_TO_UNSTRUCTURED_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/TritonToUnstructured/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/Passes.td new file mode 100755 index 00000000..542d087c --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/Passes.td @@ -0,0 +1,15 @@ +#ifndef TRITON_TO_UNSTRUCTURED_CONVERSION_PASSES +#define TRITON_TO_UNSTRUCTURED_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def TritonToUnstructured : Pass<"triton-to-unstructured", "mlir::ModuleOp"> { + let summary = "Transforms tt.addptr ops into offset accumulation ops"; + let constructor = "triton::createTritonToUnstructuredPass()"; + let options = [ + Option<"offsetBitWidth", "offset-bit-width", "size_t", /*default*/"32", + "Bitwidth used for the starting offset of each pointer"> + ]; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h new file mode 100755 index 00000000..03ccdcd1 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h @@ -0,0 +1,17 @@ +#ifndef TRITON_CONVERSION_TRITON_TO_UNSTRUCTURED_TRITON_TO_UNSTRUCTURED_H +#define TRITON_CONVERSION_TRITON_TO_UNSTRUCTURED_TRITON_TO_UNSTRUCTURED_H + +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createTritonToUnstructuredPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_TRITON_TO_UNSTRUCTURED_TRITON_TO_UNSTRUCTURED_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/CMakeLists.txt new file mode 100755 index 00000000..f988ac9f --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/CMakeLists.txt @@ -0,0 +1,10 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. +# +#===------------------------------------------------------------------------===# + +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls --name UnstructuredToMemref) +add_public_tablegen_target(UnstructuredToMemrefConversionPassIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/Passes.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/Passes.h new file mode 100755 index 00000000..f2d71174 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/Passes.h @@ -0,0 +1,22 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES_H +#define UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES_H + +#include "triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h" + +namespace mlir { +namespace triton { + +#define GEN_PASS_REGISTRATION +#include "triton-shared/Conversion/UnstructuredToMemref/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/Passes.td b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/Passes.td new file mode 100755 index 00000000..a0bf316d --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/Passes.td @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES +#define UNSTRUCTURED_TO_MEMREF_CONVERSION_PASSES + +include "mlir/Pass/PassBase.td" + +def UnstructuredToMemref : Pass<"unstructured-to-memref", "mlir::ModuleOp"> { + let summary = "Convert unstructured triton ptr (gather / scatter) to memref"; + let constructor = "triton::createUnstructuredToMemrefPass()"; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h new file mode 100755 index 00000000..ad0f5c46 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h @@ -0,0 +1,21 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_CONVERSION_UNSTRUCTUREDTOMEMREF_UNSTRUCTUREDTOMEMREF_H +#define TRITON_CONVERSION_UNSTRUCTUREDTOMEMREF_UNSTRUCTUREDTOMEMREF_H + +#include "mlir/Pass/Pass.h" + +namespace mlir { +namespace triton { + +std::unique_ptr> createUnstructuredToMemrefPass(); + +} // namespace triton +} // namespace mlir + +#endif // TRITON_CONVERSION_UNSTRUCTUREDTOMEMREF_UNSTRUCTUREDTOMEMREF_H diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/CMakeLists.txt new file mode 100755 index 00000000..55bbcc51 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/CMakeLists.txt @@ -0,0 +1,3 @@ +add_subdirectory(TritonTilingExt) +add_subdirectory(TritonStructured) +add_subdirectory(TPtr) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/CMakeLists.txt new file mode 100755 index 00000000..9aa73146 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/CMakeLists.txt @@ -0,0 +1,11 @@ +set(LLVM_TARGET_DEFINITIONS TPtrDialect.td) +mlir_tablegen(TPtrDialect.h.inc -gen-dialect-decls -dialect=tptr) +mlir_tablegen(TPtrDialect.cpp.inc -gen-dialect-defs -dialect=tptr) +mlir_tablegen(TPtrOps.h.inc -gen-op-decls) +mlir_tablegen(TPtrOps.cpp.inc -gen-op-defs) + +set(LLVM_TARGET_DEFINITIONS TPtrDialect.td) +mlir_tablegen(TPtrTypes.h.inc -gen-typedef-decls -typedefs-dialect=tptr) +mlir_tablegen(TPtrTypes.cpp.inc -gen-typedef-defs -typedefs-dialect=tptr) + +add_public_tablegen_target(TPtrTableGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/TPtrDialect.h b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/TPtrDialect.h new file mode 100755 index 00000000..ad3b1855 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/TPtrDialect.h @@ -0,0 +1,23 @@ +#ifndef MLIR_DIALECT_TPTR_IR_TPTR_DIALECT_H_ +#define MLIR_DIALECT_TPTR_IR_TPTR_DIALECT_H_ + +#include "mlir/Interfaces/SideEffectInterfaces.h" // Required for IR/TPtrOps.h.inc +#include "mlir/Bytecode/BytecodeOpInterface.h" + +#include "mlir/Dialect/Ptr/IR/PtrDialect.h" // Required for IR/TPtrOps.h.inc +#include "mlir/Dialect/Ptr/IR/PtrTypes.h" // Required for IR/TPtrOps.h.inc + +//===----------------------------------------------------------------------===// +// Temporary Pointer Dialect Operations +//===----------------------------------------------------------------------===// +#include "triton-shared/Dialect/TPtr/IR/TPtrDialect.h.inc" + +// Include the auto-generated header file containing the declarations of the +// Temporary Pointer Dialect operations. +#define GET_OP_CLASSES +#include "triton-shared/Dialect/TPtr/IR/TPtrOps.h.inc" + +#define GET_TYPEDEF_CLASSES +#include "triton-shared/Dialect/TPtr/IR/TPtrTypes.h.inc" + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/TPtrDialect.td b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/TPtrDialect.td new file mode 100755 index 00000000..4cd71678 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/TPtrDialect.td @@ -0,0 +1,200 @@ +#ifndef TPTR_DIALECT +#define TPTR_DIALECT + +include "mlir/IR/OpBase.td" +include "mlir/Interfaces/SideEffectInterfaces.td" +include "mlir/Dialect/Ptr/IR/PtrDialect.td" +include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/BuiltinTypeInterfaces.td" + +def TPtr_Dialect : Dialect { + let name = "tptr"; + + let cppNamespace = "::mlir::tptr"; + + let summary = "Temporary Pointer Dialect"; + + let description = [{ + Typed Pointer Dialect. + }]; + + let extraClassDeclaration = [{ + void registerTypes(); + }]; + + let dependentDialects = [ + "mlir::ptr::PtrDialect" + ]; + + let usePropertiesForAttributes = 1; +} + +class TPtrTypeDef traits = []> + : TypeDef { + // Used by printer/parser + let mnemonic = _mnemonic; +} + +// +// Op Base +// +class TPTR_Op traits = []> : + Op { +} + +def TPTR_IntToPtrOp : TPTR_Op<"inttoptr", [ + Pure + ]> { + let summary = "Integer to a pointer operation"; + let description = [{ + The `inttoptr` operation casts an int or index value to a pointer. + + Example: + ```mlir + %ptr = ptr.inttoptr %int : i32 to !ptr.ptr<1 : i32> + ``` + }]; + let arguments = (ins AnySignlessIntegerOrIndex:$arg); + let results = (outs Ptr_PtrType:$res); + let assemblyFormat = "$arg attr-dict `:` type($arg) `to` type($res)"; +} + +def TPTR_PtrToIntOp : TPTR_Op<"ptrtoint", [ + Pure + ]> { + let summary = "Pointer to an integer operation"; + let description = [{ + The `ptrtoint` operation casts a pointer value to an int or index. + + Example: + ```mlir + %int = ptr.ptrtoint %ptr : !ptr.ptr<1 : i32> to i32 + ``` + }]; + let arguments = (ins Ptr_PtrType:$arg); + let results = (outs AnySignlessIntegerOrIndex:$res); + let assemblyFormat = "$arg attr-dict `:` type($arg) `to` type($res)"; +} + +def TPTR_TypeOffsetOp : TPTR_Op<"type_offset", [ConstantLike, Pure]> { + let summary = "Creates a type offset constant."; + let description = [{ + The `addr.type_offset` operation produces an int or index-typed SSA value + equal to a target-specific constant representing the offset of a single + element of the given type. The default return type is `index`. + Example: + + ```mlir + %0 = addr.type_offset f32 + %1 = addr.type_offset memref<12 x f64> : i32 + ``` + }]; + + let arguments = (ins TypeAttr:$baseType); + let results = (outs AnySignlessIntegerOrIndex:$result); + let builders = [ + OpBuilder<(ins "TypeAttr":$baseType, CArg<"Type", "nullptr">:$resultTy)> + ]; + let assemblyFormat = [{ + attr-dict $baseType custom(type($result)) + }]; + let hasFolder = 1; +} + +def TPTR_FromMemrefOp : TPTR_Op<"from_memref", [Pure]> { + let arguments = (ins AnyMemRef:$input); + let results = (outs Ptr_PtrType:$result); + let assemblyFormat = "$input attr-dict `:` type($input) `to` type($result)"; +} + +def TPTR_ToMemrefOp : TPTR_Op<"to_memref", [ + Pure ]> { + let arguments = (ins Ptr_PtrType:$arg); + let results = (outs AnyStaticShapeMemRef:$res); + let assemblyFormat = "$arg attr-dict `:` type($arg) `to` type($res)"; +} + +def TPTR_PtrAddOp : TPTR_Op<"ptradd", [Pure, AllTypesMatch<["base", "result"]>]> { + let summary = "Pointer-index add operation"; + let description = [{ + The `ptradd` operation adds an `address` and an integer or index to + produce a new address. + + Example: + ```mlir + %addr = ptr.ptradd %addr : !ptr.ptr<3 : i32>, %c10 : i32 + ``` + }]; + + let arguments = (ins Ptr_PtrType:$base, AnySignlessIntegerOrIndex:$offset); + let results = (outs Ptr_PtrType:$result); + let assemblyFormat = "$base $offset attr-dict `:` type($base) `,` type($offset) `to` type($result)"; +} + +def TPTR_LoadOp : TPTR_Op<"load", [ + DeclareOpInterfaceMethods + ]> { + let summary = "Load operation"; + let description = [{ + The `load` operation is used to read from memory. A load may be marked as + atomic, volatile, and/or nontemporal, and takes a number of optional + attributes that specify aliasing information. + + An atomic load only supports a limited set of pointer, integer, and + floating point types, and requires an explicit alignment. + + Examples: + ```mlir + // A volatile load of a float variable. + %0 = ptr.load volatile %ptr : !ptr.ptr -> f32 + + // A nontemporal load of a float variable. + %0 = ptr.load %ptr {nontemporal} : !ptr.ptr -> f32 + + // An atomic load of an integer variable. + %0 = ptr.load %ptr atomic monotonic {alignment = 8 : i64} + : !ptr.ptr -> i64 + ``` + }]; + let arguments = (ins AnyType:$addr); + let results = (outs AnyType:$res); + let assemblyFormat = [{ + $addr + attr-dict `:` qualified(type($addr)) `->` type($res) + }]; +} + +def TTPTR_StoreOp : TPTR_Op<"store", [ + DeclareOpInterfaceMethods + ]> { + let summary = "Store operation"; + let description = [{ + The `store` operation is used to write to memory. A store may be marked as + atomic, volatile, and/or nontemporal, and takes a number of optional + attributes that specify aliasing information. + + An atomic store only supports a limited set of pointer, integer, and + floating point types, and requires an explicit alignment. + + Examples: + ```mlir + // A volatile store of a float variable. + ptr.store volatile %val, %ptr : f32, !ptr.ptr + + // A nontemporal store of a float variable. + ptr.store %val, %ptr {nontemporal} : f32, !ptr.ptr + + // An atomic store of an integer variable. + ptr.store %val, %ptr atomic monotonic {alignment = 8 : i64} + : i64, !ptr.ptr + ``` + }]; + let arguments = (ins AnyType:$value, + AnyType:$addr); + let assemblyFormat = [{ + $value `,` $addr + attr-dict `:` type($value) `,` qualified(type($addr)) + }]; +} + +#endif // TPTR_DIALECT diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/CMakeLists.txt new file mode 100755 index 00000000..9c32c97c --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/CMakeLists.txt @@ -0,0 +1,8 @@ +set(LLVM_TARGET_DEFINITIONS TritonStructuredDialect.td) +mlir_tablegen(TritonStructuredDialect.h.inc -gen-dialect-decls -dialect=tts) +mlir_tablegen(TritonStructuredDialect.cpp.inc -gen-dialect-defs -dialect=tts) +mlir_tablegen(TritonStructuredOps.h.inc -gen-op-decls) +mlir_tablegen(TritonStructuredOps.cpp.inc -gen-op-defs) + + +add_public_tablegen_target(TritonStructuredTableGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h new file mode 100755 index 00000000..fbead0c5 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h @@ -0,0 +1,29 @@ +#ifndef MLIR_DIALECT_TRITON_STRUCTURED_IR_TRITON_STRUCTURED_DIALECT_H_ +#define MLIR_DIALECT_TRITON_STRUCTURED_IR_TRITON_STRUCTURED_DIALECT_H_ + +#include "mlir/IR/Dialect.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +namespace mlir { +namespace tts { +namespace utils { +mlir::Value getScalarValue(mlir::Value operand, mlir::Location loc, + mlir::OpBuilder &builder); +} +} // namespace tts +} // namespace mlir + +//===----------------------------------------------------------------------===// +// TritonStructured Operations +//===----------------------------------------------------------------------===// +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h.inc" + +// Include the auto-generated header file containing the declarations of the +// TritonStructured operations. +#define GET_OP_CLASSES +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredOps.h.inc" + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.td b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.td new file mode 100755 index 00000000..8412c4d0 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.td @@ -0,0 +1,338 @@ +#ifndef TRITON_STRUCTURED_DIALECT +#define TRITON_STRUCTURED_DIALECT + +include "mlir/IR/OpBase.td" +include "triton/Dialect/Triton/IR/TritonTypes.td" +include "triton/Dialect/Triton/IR/TritonAttrDefs.td" +include "mlir/Interfaces/SideEffectInterfaces.td" + +def Triton_Structured_Dialect : Dialect { + let name = "tts"; + + let cppNamespace = "::mlir::tts"; + + let summary = "Structured Triton operations"; + + let description = [{ + Triton Structured Dialect. + }]; + + let dependentDialects = [ + "triton::TritonDialect" + ]; + + let usePropertiesForAttributes = 1; +} + +// +// Op Base +// +class TTS_Op traits = []> : + Op { +} + +def TTS_MakeTensorPtrOp + : TTS_Op<"make_tptr", [AttrSizedOperandSegments, Pure]> { + let summary = "create a pointer that points to a tensor in memory"; + + // base: Base pointer used to contruct the tensor of pointers or pointer to tensor. + // sizes: Size of the data being loaded or stored. + // strides: The strides of the parent tensor, which means how much to increase the pointer + // by when moving by 1 element in a specific axis. + // order: The order of the block, which means how the block is laid out in memory. + // It contains the same info as order in tt.make_tensor_ptr. + // shape: If order is present, this field signifies the shape of the parent tensor in + // memory; if order is not present, it signifies the boundary by which addresses + // wraps around (constant zero indicates no wrap-around in the corresponding dimension). + // offsets: Offset of the block along each dimension from base. + // result: If order is present, this op produces a pointer to a tensor; otherwise, + // it produces a tensor of pointers. + + let arguments = (ins TT_Ptr:$base, + DenseI64ArrayAttr:$sizes, + Variadic:$strides, + Variadic:$offsets, + Variadic:$shape, + DenseI64ArrayAttr:$static_strides, + DenseI64ArrayAttr:$static_offsets, + DenseI64ArrayAttr:$static_shape, + DenseI32ArrayAttr:$order); + + let results = (outs TT_PtrLike:$result); + + let assemblyFormat = [{ + $base `to` `sizes` `` `:` $sizes + `` `,` `strides` `` `:` + custom($strides, $static_strides) + `` `,` `offsets` `` `:` + custom($offsets, $static_offsets) + `` `,` `shape` `` `:` + custom($shape, $static_shape) + `` `,` `order` `` `:` $order + attr-dict `:` type($base) `to` type($result) + }]; + + + let builders = [ + // Build with mixed static and dynamic entries. + OpBuilder<(ins + "Value":$base, + "ArrayRef":$sizes, + "ArrayRef":$strides, + "ArrayRef":$offsets, + "ArrayRef":$shape, + "ArrayRef":$order)>, + ]; + + let extraClassDeclaration = [{ + /// Return a vector of all the static or dynamic fields + SmallVector getMixedSizes() { + Builder b(getContext()); + SmallVector dynSizes; // sizes are always static + return ::mlir::getMixedValues(getSizes(), dynSizes, b); + } + SmallVector getMixedStrides() { + Builder b(getContext()); + return ::mlir::getMixedValues(getStaticStrides(), getStrides(), b); + } + SmallVector getMixedOffsets() { + Builder b(getContext()); + return ::mlir::getMixedValues(getStaticOffsets(), getOffsets(), b); + } + SmallVector getMixedShape() { + Builder b(getContext()); + return ::mlir::getMixedValues(getStaticShape(), getShape(), b); + } + bool isBlockPtr() { + return !getOrder().empty(); + } + bool isStructuredPtr() { + return !isBlockPtr() && + llvm::all_of(getStaticShape(), [](auto shape) { return shape == 0; }); + } + bool isSplitPtr() { + return !isBlockPtr() && + !isStructuredPtr(); + } + }]; + + // TODO + //let hasVerifier = 1; + //let hasCanonicalizer = 1; +} + +def TTS_GetStructuredStateOp : TTS_Op<"get_structured_state", [AttrSizedResultSegments, Pure]> { + let summary = "Placeholder for the structured pointer states computed during PtrAnalysis."; + let description = "Used to pass the offsets and strides to scf.for op to simplify IR rewrites."; + + let arguments = ( + ins AnyTypeOf<[TT_PtrLike, I1Tensor, I16Tensor, I32Tensor, I64Tensor]>:$input + ); + let results = ( + outs AnyTypeOf<[TT_PtrLike, I1Tensor, I16Tensor, I32Tensor, I64Tensor]>:$structured, + Variadic:$offsets, + Variadic:$strides + ); + + let builders = [ + OpBuilder<(ins "Value":$input)>, + ]; + + let extraClassDeclaration = [{ + static std::optional, SmallVector>> + getOffsetAndStrideTypes(MLIRContext *context, Type ptrLikeType); + + static std::optional> + getOffsetAndStrideSegmentSizes(Type ptrLikeType); + }]; + + let hasFolder = 0; + let hasVerifier = 1; +} + +def TTS_GatherOp : TTS_Op<"gather", [ + MemoryEffects<[MemRead]>, + AttrSizedOperandSegments, + OptionalTypesMatchWith<"mask type matches ptr type", "offset", "mask", "triton::getI1SameShape($_self)">, + OptionalTypesMatchWith<"other matches ptr type", "ptr", "other", "triton::getPointeeType($_self)"> +]> { + let summary = "optionally load data from in memory to fill a portion of the tensor"; + + let arguments = ( + ins + TT_Ptr:$ptr, + TT_IntLike:$offset, + Optional:$mask, + Optional:$other + ); + + let results = (outs TT_Type:$result); + + let assemblyFormat = [{ + $ptr `[` $offset `]` (`mask` `=` $mask^)? (`default` `=` $other^)? + attr-dict `:` `(` type($ptr) `,` type($offset) `)` `->` type($result) + }]; +} + +def TTS_ScatterOp : TTS_Op<"scatter", [ + MemoryEffects<[MemWrite]>, + OptionalTypesMatchWith<"mask type matches offset type", "offset", "mask", + "triton::getI1SameShape($_self)"> +]> { + let summary = "optionally store data from in memory to fill a portion of the tensor"; + + let arguments = ( + ins + TT_Ptr:$ptr, + TT_IntLike:$offset, + TT_Type:$value, + Optional:$mask + ); + + let assemblyFormat = [{ + $value `into` $ptr `[` $offset `]` (`mask` `=` $mask^)? + attr-dict `:` type($value) `into` ` ` `(` type($ptr) `,` type($offset) `)` + }]; +} + +def TTS_LoadOp : TTS_Op<"load", [ + MemoryEffects<[MemRead]>, + AttrSizedOperandSegments +]> { + let summary = "optionally load data from in memory to fill a portion of the tensor"; + + let arguments = (ins TT_PtrLike:$ptr, + Variadic:$mask_dims, + DenseI64ArrayAttr:$static_mask_dims, + Optional>:$other); + + let results = (outs TT_Tensor:$result); + + let builders = [ + OpBuilder<(ins "Value":$ptr, "ArrayRef":$mask_dims, "Value":$other)>, + ]; + + let extraClassDeclaration = [{ + /// Return a vector of all the static or dynamic fields + SmallVector getMixedMaskDims() { + Builder b(getContext()); + return ::mlir::getMixedValues(getStaticMaskDims(), getMaskDims(), b); + } + + bool hasMask() { + return !getMixedMaskDims().empty(); + } + }]; + + // TODO + //let hasCustomAssemblyFormat = 1; + //let hasVerifier = 1; +} + +def TTS_StoreOp : TTS_Op<"store", [ + MemoryEffects<[MemWrite]> +]> { + let summary = "optionally store data from in memory to fill a portion of the tensor"; + + let arguments = (ins TT_PtrLike:$ptr, + TT_Tensor:$value, + Variadic:$mask_dims, + DenseI64ArrayAttr:$static_mask_dims); + + let builders = [ + OpBuilder<(ins "Value":$ptr, "Value":$value, "ArrayRef":$dims)>, + ]; + + let extraClassDeclaration = [{ + /// Return a vector of all the static or dynamic fields + SmallVector getMixedMaskDims() { + Builder b(getContext()); + return ::mlir::getMixedValues(getStaticMaskDims(), getMaskDims(), b); + } + + bool hasMask() { + return !getMixedMaskDims().empty(); + } + }]; + + // TODO + //let hasCustomAssemblyFormat = 1; + //let hasVerifier = 1; +} + +def TTS_AtomicRMWOp : TTS_Op<"atomic_rmw", [ + MemoryEffects<[MemRead, MemWrite]> +]> { + let summary = "perform atomic read-modify-write operation on a pointer"; + + let arguments = (ins TT_PtrLike:$ptr, + TT_Tensor:$value, + Variadic:$mask_dims, + DenseI64ArrayAttr:$static_mask_dims, + TT_AtomicRMWAttr:$atomic_rmw_op, + TT_MemSemanticAttr:$sem, + TT_MemSyncScopeAttr:$scope); + + let results = (outs TT_Tensor:$result); + + let builders = [ + OpBuilder<(ins "mlir::Type":$result, "Value":$ptr, "Value":$value, + "ArrayRef":$dims, + "triton::RMWOpAttr":$atomic_rmw_op, + "triton::MemSemanticAttr":$sem, "triton::MemSyncScopeAttr":$scope)>, + ]; + + let extraClassDeclaration = [{ + /// Return a vector of all the static or dynamic fields + SmallVector getMixedMaskDims() { + Builder b(getContext()); + return ::mlir::getMixedValues(getStaticMaskDims(), getMaskDims(), b); + } + + bool hasMask() { + return !getStaticMaskDims().empty(); + } + }]; + + // TODO + //let hasCustomAssemblyFormat = 1; + //let hasVerifier = 1; +} + +def TTS_IndexedAtomicRMWOp : TTS_Op<"indexed_atomic_rmw", [ + MemoryEffects<[MemRead, MemWrite]> +]> { + let summary = "perform atomic read-modify-write operation on a pointer with index/offset"; + + let arguments = (ins TT_PtrLike:$ptr, + TT_Type:$value, + Optional:$mask, + TT_IntLike:$offset, + TT_AtomicRMWAttr:$atomic_rmw_op, + TT_MemSemanticAttr:$sem, + TT_MemSyncScopeAttr:$scope); + + let results = (outs TT_Type:$result); +} + + +def TTS_AtomicCASOp : TTS_Op<"atomic_cas", [ + MemoryEffects<[MemRead, MemWrite]> +]> { + let summary = "perform atomic compare-and-swap operation on a pointer"; + + let arguments = (ins TT_PtrLike:$ptr, + TT_Type:$cmp, + TT_Type:$value, + Optional:$offset, // For unstructured pointers + TT_MemSemanticAttr:$sem, + TT_MemSyncScopeAttr:$scope); + + let results = (outs TT_Type:$result); + + // TODO + //let hasCustomAssemblyFormat = 1; + //let hasVerifier = 1; +} + +#endif // TRITON_STRUCTURED_DIALECT diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/CMakeLists.txt new file mode 100755 index 00000000..ba67b25a --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/CMakeLists.txt @@ -0,0 +1,11 @@ +set(LLVM_TARGET_DEFINITIONS TritonTilingExtOps.td) +mlir_tablegen(TritonTilingExtOpsDialect.h.inc -gen-dialect-decls -dialect=ttx) +mlir_tablegen(TritonTilingExtOpsDialect.cpp.inc -gen-dialect-defs -dialect=ttx) +mlir_tablegen(TritonTilingExtOps.h.inc -gen-op-decls) +mlir_tablegen(TritonTilingExtOps.cpp.inc -gen-op-defs) +add_public_tablegen_target(TritonTilingExtOpsIncGen) + +set(LLVM_TARGET_DEFINITIONS TritonTilingExtInterfaces.td) +mlir_tablegen(TritonTilingExtInterfaces.h.inc -gen-op-interface-decls) +mlir_tablegen(TritonTilingExtInterfaces.cpp.inc -gen-op-interface-defs) +add_public_tablegen_target(TritonTilingExtInterfacesIncGen) diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h new file mode 100755 index 00000000..53e031db --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h @@ -0,0 +1,107 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_DIALECT_TRITON_TILING_EXT_IR_TRITON_TILING_EXT_DIALECT_H_ +#define MLIR_DIALECT_TRITON_TILING_EXT_IR_TRITON_TILING_EXT_DIALECT_H_ + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/SymbolTable.h" +#include "mlir/IR/TypeSupport.h" +#include "mlir/IR/Types.h" +#include "mlir/Interfaces/DestinationStyleOpInterface.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "mlir/Interfaces/TilingInterface.h" + +//===----------------------------------------------------------------------===// +// TritonTilingExt Operations +//===----------------------------------------------------------------------===// + +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOpsDialect.h.inc" + +// Include the generated interface declarations. +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtInterfaces.h.inc" + +// Include the auto-generated header file containing the declarations of the +// TritonTilingExt operations. +#define GET_OP_CLASSES +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOps.h.inc" + +namespace mlir { + +namespace ttx { + +// ----------------------------------------------------------------------------- +// BufferizableOpInterface +// ----------------------------------------------------------------------------- +// All TritonTilingExtOps need to support bufferization: the process of +// allocating buffers for tensors, thereby converting inputs and outputs of +// tensor type to memref. This process is done by implementing the +// "BufferizableOpInterface". We implement the interface for TritonTilingExtOps +// through an external model instead of directly in TritonTilingExtOps.td to be +// consistent with other ops in the mlir project. See some examples here: +// - mlir/lib/Dialect/Linalg/Transforms/BufferizableOpInterfaceImpl.cpp +// - mlir/lib/Dialect/SCF/Transforms/BufferizableOpInterfaceImpl.cpp +void registerBufferizableOpInterfaceExternalModels(DialectRegistry ®istry); + +// ----------------------------------------------------------------------------- +// TilingInterface +// ----------------------------------------------------------------------------- +// The three methods `getTiledImplementation`, `getResultTilePosition`, and +// `generateResultTileValue` are implemented as part of the TilingInterface. +// (see TilingInterface.td). These three methods are re-used across +// all TritonTilingExtOps, while others method are implemented individually by +// each operator depending on their use cases. +template +FailureOr getTiledImplementation(TritonTilingExtOpTy op, + OpBuilder &b, + ArrayRef offsets, + ArrayRef sizes); + +template +LogicalResult getResultTilePosition(TritonTilingExtOpTy op, OpBuilder &b, + unsigned resultNumber, + ArrayRef offsets, + ArrayRef sizes, + SmallVector &resultOffsets, + SmallVector &resultSizes); + +template +FailureOr +generateResultTileValue(TritonTilingExtOpTy op, OpBuilder &b, + unsigned resultNumber, ArrayRef offsets, + ArrayRef sizes); + +// ----------------------------------------------------------------------------- +// MemoryEffectsOpInterface +// ----------------------------------------------------------------------------- +// Implementation of the MemoryEffectsOpInterface for TritonTilingExtOps. +// This allows DCE pass to determine if a TritonTilingExtOp is safe to be +// removed. see TritonTilingExtOps.td for more details. +template +void getEffects( + TritonTilingExtOpTy op, + SmallVectorImpl> + &effects); + +// ----------------------------------------------------------------------------- +// Utilities +// ----------------------------------------------------------------------------- +// Utility method to extract a slice from the input source using either +// tensor::ExtractSlice or memref::SubView +Value getSlice(OpBuilder &b, Location loc, Value source, + ArrayRef offsets, ArrayRef sizes, + ArrayRef strides); + +} // namespace ttx +} // namespace mlir + +#endif // MLIR_DIALECT_TRITON_TILING_EXT_IR_TRITON_TILING_EXT_DIALECT_H_ diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtInterfaces.td b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtInterfaces.td new file mode 100755 index 00000000..e74fbb6c --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtInterfaces.td @@ -0,0 +1,102 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef MLIR_TRITON_TILING_EXT_DIALECT_INTERFACES +#define MLIR_TRITON_TILING_EXT_DIALECT_INTERFACES + +include "mlir/IR/OpBase.td" + +// +// Linalg operators require providing affine maps that define how input / output +// buffers are accessed together with a region that defines how each output +// element is computed; this requirement doesn't work well for operations such as +// `scan`. +// +// Fortunately, the introduction of the TilingInterface allows us to add tiling +// and fusion support to operations that don't fit into the linalg dialect. +// This fits our purpose perfectly: our `scan` operators can be treated as an +// "opaque" / "completely abstract" operation that can be tiled on the batch +// dimensions -- we don't need to provide any associated body together with it. +// +// However, this doesn't mean that we entirely forgo the "indexing map" concept. +// For example, consider the following: +// +// - ttx.scan ins(%1 : tensor<128x768xbf16>) +// outs(%2 : tensor<128x768xbf16>) -> tensor<128x768xbf16> +// +// Tiling the batch dimension gives us: +// +// for (i = 0 to 128) { +// %sliceIn = extract slice from input: tensor<1x768xbf16> +// %sliceOut = extract slice from output: tensor<1x768xbf16> +// %res = ttx.scan ins(slice : tensor<1x768xbf16>) +// outs(%2 : tensor<1x768xbf16>) -> tensor<1x768xbf16> +// insert %res into output +// } +// +// Now our `scan` op has the semantic of running `scan` on a rank-1 tensor and +// can be lowered further to other hardware-specific ops or external library +// calls. +// +// This tiling pattern is essentially the same as tiling a linalg.generic op +// with an identity map. The only difference is we don't need a body associated +// with our `scan` op. +// +// With this idea in mind, the TritonTilingExtInterface exposes methods +// that will be implemented individually by each TritonTilingExtOp, providing +// the indexing map for each input / output that can then be used to generate +// the correct slices during tiling and fusion. +// +// There might be other ops in the future that won't fit in this "indexing map" +// approach; we will consider making TritonTilingExtInterface an optional +// interface for such ops. +// + +def TritonTilingExtInterface : OpInterface<"TritonTilingExtInterface"> { + let cppNamespace = "::mlir::ttx"; + let methods = [ + InterfaceMethod< + /*desc=*/[{ + Return the indexing map for the input operand with the given `index`. + The `tileSizes` input indicates the requested tile size during tiling + in case the indexing map for the operator is dependent on it. + }], + /*retTy=*/"AffineMap", + /*methodName=*/"getInputIndexingMap", + /*args=*/(ins "MLIRContext*":$context, + "unsigned int":$index, + "ArrayRef":$tileSizes) + >, + InterfaceMethod< + /*desc=*/[{ + Return the indexing map for the output operand with the given `index`. + The `tileSizes` input indicates the requested tile size during tiling + in case the indexing map for the operator is dependent on it. + }], + /*retTy=*/"AffineMap", + /*methodName=*/"getOutputIndexingMap", + /*args=*/(ins "MLIRContext*":$context, + "unsigned int":$index, + "ArrayRef":$tileSizes) + >, + InterfaceMethod< + /*desc=*/[{ + Return the indexing map for the operand with the given `index`. + This method returns the operand in order of inputs followed by outputs. + The `tileSizes` input indicates the requested tile size during tiling + in case the indexing map for the operator is dependent on it. + }], + /*retTy=*/"AffineMap", + /*methodName=*/"getIndexingMap", + /*args=*/(ins "MLIRContext*":$context, + "unsigned int":$index, + "ArrayRef":$tileSizes) + > + ]; +} + +#endif diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOps.td b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOps.td new file mode 100755 index 00000000..d3a4268a --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOps.td @@ -0,0 +1,242 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#ifndef TRITON_TILING_EXT_BASE +#define TRITON_TILING_EXT_BASE + +include "mlir/IR/EnumAttr.td" +include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/OpBase.td" +include "mlir/IR/BuiltinAttributes.td" +include "mlir/IR/SymbolInterfaces.td" +include "mlir/Interfaces/CallInterfaces.td" +include "mlir/Interfaces/DestinationStyleOpInterface.td" +include "mlir/Interfaces/SideEffectInterfaces.td" +include "mlir/Interfaces/TilingInterface.td" +include "mlir/Dialect/Linalg/IR/LinalgBase.td" +include "mlir/Dialect/Linalg/IR/LinalgInterfaces.td" + +include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtInterfaces.td" + + +//===----------------------------------------------------------------------===// +// TritonTilingExt dialect definition +//===----------------------------------------------------------------------===// + +def TritonTilingExt_Dialect : Dialect { + let name = "ttx"; + let cppNamespace = "::mlir::ttx"; +} + +//===----------------------------------------------------------------------===// +// TritonTilingExt op definitions +//===----------------------------------------------------------------------===// + +// Base class for TritonTilingExt dialect ops. +class TritonTilingExt_Op traits = []> + : Op { +} + +class TritonTilingExt_TilingOp : Op, + // All TritonTilingExtOps implement TritonTilingExtInterface, which provides a standardized + // way of providing indexing maps for input and output operands. + DeclareOpInterfaceMethods, + + // MemoryEffectsOpInterface provides analysis passes such as DCE to determine + // whether an operation has no memory side effects and therefore is safe to + // be deleted. This interface is important during tile and fuse where we + // create copies of TilingInterface ops with smaller tile sizes but leave the + // original ops intact. + DeclareOpInterfaceMethods, + + // DestinationStyleOpInterface describes ops that have similar semantics to + // linalg ops, with a separate ins (input) and outs (output) operand groups. + // Implementing this op gives us access to a wide variety of useful methods + // to query the inputs and outputs of an op. + DestinationStyleOpInterface, + + // AttrSizedOperandSegments supports having multiple groups of operands. + // For example, linalg ops (as well as TritonTilingExtOps) all look like this: + // ttx.some_op ins(%1) outs(%2) -> resultType + AttrSizedOperandSegments +]> +{ + let results = (outs Variadic:$result_tensors); + + let hasCustomAssemblyFormat = 1; + + code baseClassDecls = [{ + // Implemented as part of DestinationStyleOpInterface + MutableOperandRange getDpsInitsMutable() { return getOutputsMutable(); } + }]; + + // Custom print() and parse() methods to make the TritonTilingExt ops have similar looks + // to the linalg ops. + // Borrowed from llvm-project/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp + let extraClassDefinition = [{ + void $cppClass::print(OpAsmPrinter &p) { + p.printOptionalAttrDict(this->getOperation()->getAttrs(), + /*elidedAttrs=*/{"operand_segment_sizes"}); + + if (!getInputs().empty()) + p << " ins(" << getInputs() << " : " << getInputs().getTypes() << ")"; + if (!getOutputs().empty()) + p << " outs(" << getOutputs() << " : " << getOutputs().getTypes() << ")"; + + if (!getResultTypes().empty()) + p.printOptionalArrowTypeList(getResultTypes()); + } + + ParseResult $cppClass::parse(OpAsmParser &parser, + OperationState &result) { + SmallVector inputTypes; + SmallVector outputTypes; + SMLoc inputsOperandsLoc, outputsOperandsLoc; + SmallVector inputsOperands, + outputsOperands; + if (parser.parseOptionalAttrDict(result.attributes)) + return failure(); + + if (succeeded(parser.parseOptionalKeyword("ins"))) { + if (parser.parseLParen()) + return failure(); + + inputsOperandsLoc = parser.getCurrentLocation(); + if (parser.parseOperandList(inputsOperands) || + parser.parseColonTypeList(inputTypes) || parser.parseRParen()) + return failure(); + } + + if (succeeded(parser.parseOptionalKeyword("outs"))) { + outputsOperandsLoc = parser.getCurrentLocation(); + if (parser.parseLParen() || parser.parseOperandList(outputsOperands) || + parser.parseColonTypeList(outputTypes) || parser.parseRParen()) + return failure(); + } + + if (parser.resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc, + result.operands) || + parser.resolveOperands(outputsOperands, outputTypes, + outputsOperandsLoc, result.operands)) + return failure(); + + result.addAttribute("operand_segment_sizes", + parser.getBuilder().getDenseI32ArrayAttr( + {static_cast(inputsOperands.size()), + static_cast(outputsOperands.size())})); + + SmallVector resultTypes; + if (parser.parseOptionalArrowTypeList(resultTypes)) + return failure(); + result.addTypes(resultTypes); + + return success(); + } + + AffineMap $cppClass::getIndexingMap(MLIRContext *context, + unsigned int index, + ArrayRef sizes) { + assert(index < this->getNumOperands()); + if (index < getNumDpsInputs()) { + return getInputIndexingMap(context, index, sizes); + } + return getOutputIndexingMap(context, index - getNumDpsInputs(), sizes); + } + + // Forward each of the implementation to the shared implementation + FailureOr $cppClass::getTiledImplementation( + OpBuilder &b, + ArrayRef offsets, + ArrayRef sizes + ) { + return mlir::ttx::getTiledImplementation<$cppClass>( + *this, b, offsets, sizes + ); + } + + // Forward each of the implementation to the shared implementation + LogicalResult $cppClass::getResultTilePosition( + OpBuilder &b, + unsigned resultNumber, + ArrayRef offsets, + ArrayRef sizes, + SmallVector &resultOffsets, + SmallVector &resultSizes + ) { + return mlir::ttx::getResultTilePosition<$cppClass>( + *this, b, resultNumber, offsets, sizes, resultOffsets, resultSizes + ); + } + + // Forward each of the implementation to the shared implementation + FailureOr $cppClass::generateResultTileValue( + OpBuilder &b, + unsigned resultNumber, + ArrayRef offsets, + ArrayRef sizes + ) { + return mlir::ttx::generateResultTileValue<$cppClass>( + *this, b, resultNumber, offsets, sizes + ); + } + + // Implemented as part of MemoryEffectsOpInterface + void $cppClass::getEffects( + SmallVectorImpl> + &effects + ) { + return mlir::ttx::getEffects<$cppClass>(*this, effects); + } + }]; +} + +def TritonTilingExt_CumSumOp : TritonTilingExt_TilingOp<"cumsum"> { + let arguments = (ins + Variadic:$inputs, + Variadic:$outputs, + UI32Attr:$axis + ); + + let hasVerifier = 1; + + let skipDefaultBuilders = 1; + + let builders = [ + OpBuilder<(ins + "Value":$input, + "IntegerAttr":$axis, + "Value":$output, + CArg<"ArrayRef", "{}">:$attributes + )> + ]; + + let extraClassDeclaration = baseClassDecls # [{ + int64_t getRank() { + return cast(getInput().getType()).getRank(); + } + + Value getInput() { + return getInputs()[0]; + } + + Value getOutput() { + return getOutputs()[0]; + } + + static StringRef getAxisAttrStrName() { return "axis"; } + }]; +} + +#endif // TRITON_TILING_EXT_BASE diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Utils/FusionHelper.h b/third_party/wafer/third_party/flir/include/triton-shared/Utils/FusionHelper.h new file mode 100755 index 00000000..82ec5bc7 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Utils/FusionHelper.h @@ -0,0 +1,207 @@ +#ifndef TRITON_FUSION_PATTERNS +#define TRITON_FUSION_PATTERNS + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include +#include +#include +#include +#include + +using namespace mlir; + +namespace { + +//===--------------------------- Match ArgMinMax --------------------------===// + + // We're looking for an op that looks like this: + // + // %9:2 = "tt.reduce"(%8, %3) <{axis = 0 : i32}> ({ + // ^bb0(%arg9: f32, %arg10: i32, %arg11: f32, %arg12: i32): + // ------------------------------------------------- + // `matchTieBreakValue` | + // %11 = arith.cmpf oeq, %arg9, %arg11 : f32 | + // %12 = arith.cmpi slt, %arg10, %arg12 : i32 | 1. + // %13 = arith.andi %11, %12 : i1 | + // ------------------------------------------------- |-> `matchShouldUpdate` + // `matchUpdateCondition` | + // %14 = arith.cmpf ogt, %arg9, %arg11 : f32 | 2. + // ------------------------------------------------- | + // %15 = arith.ori %14, %13 : i1 | + // ------------------------------------------------- + // %16 = arith.select %15, %arg9, %arg11 : f32 + // %17 = arith.select %15, %arg10, %arg12 : i32 + +static LogicalResult matchTieBreakResult(Value currValue, Value currIndex, + Value reduceValue, Value reduceIndex, + mlir::Block::iterator &it, + Value &tileBreakValue) { + // Match the following (section 1. of the above) + // + // %11 = arith.cmpf oeq, %arg9, %arg11 : f32 + // %12 = arith.cmpi slt, %arg10, %arg12 : i32 + // %13 = arith.andi %11, %12 : i1 + // + // which is equivalent to the following python code + // + // tie = value1 == value2 and index1 < index2 + + // matching: %11 = arith.cmpf oeq, %arg9, %arg11 : f32 + auto& cmpOp = *it++; + Value eqCmpOp; + if (auto eqCmpFOp = dyn_cast(cmpOp)) { + if (eqCmpFOp.getPredicate() != arith::CmpFPredicate::OEQ || + currValue != eqCmpFOp.getLhs() || reduceValue != eqCmpFOp.getRhs()) { + return failure(); + } + eqCmpOp = eqCmpFOp; + } else if (auto eqCmpIOp = dyn_cast(cmpOp)) { + if (eqCmpIOp.getPredicate() != arith::CmpIPredicate::eq || + currValue != eqCmpIOp.getLhs() || reduceValue != eqCmpIOp.getRhs()) { + return failure(); + } + eqCmpOp = eqCmpIOp; + } else { + return failure(); + } + + // matching: %12 = arith.cmpi slt, %arg10, %arg12 : i32 + auto sltCmpOp = dyn_cast(*it++); + if (!sltCmpOp || sltCmpOp.getPredicate() != arith::CmpIPredicate::slt || + currIndex != sltCmpOp.getLhs() || reduceIndex != sltCmpOp.getRhs()) { + return failure(); + } + + // matching: %13 = arith.andi %11, %12 : i1 + auto andOp = dyn_cast(*it++); + if (!andOp || andOp.getLhs() != eqCmpOp || andOp.getRhs() != sltCmpOp) { + return failure(); + } + + tileBreakValue = andOp; + return success(); +} + +static LogicalResult matchComparisonResult(Value currValue, Value currIndex, + Value reduceValue, + Value reduceIndex, + mlir::Block::iterator &it, + Value &comparisonResult, + bool isArgMin) { + // %14 = arith.cmpf olt(ogt), %arg9, %arg11 : f32 + auto &cmpOp = *it++; + if (auto eqCmpFOp = dyn_cast(cmpOp)) { + auto predicate = + isArgMin ? arith::CmpFPredicate::OLT : arith::CmpFPredicate::OGT; + if (eqCmpFOp.getPredicate() != predicate || + currValue != eqCmpFOp.getLhs() || reduceValue != eqCmpFOp.getRhs()) { + return failure(); + } + comparisonResult = eqCmpFOp; + } else if (auto eqCmpIOp = dyn_cast(cmpOp)) { + auto predicate = + isArgMin ? arith::CmpIPredicate::slt : arith::CmpIPredicate::sgt; + if (eqCmpIOp.getPredicate() != predicate || + currValue != eqCmpIOp.getLhs() || reduceValue != eqCmpIOp.getRhs()) { + return failure(); + } + comparisonResult = eqCmpIOp; + } else { + return failure(); + } + + return success(); +} + +static LogicalResult matchShouldUpdateValue(Value currValue, Value currIndex, + Value reduceValue, Value reduceIndex, + mlir::Block::iterator &it, + Value &shouldUpdate, + bool isArgMin) { + Value tieResult; + if (failed(matchTieBreakResult(currValue, currIndex, reduceValue, + reduceIndex, it, tieResult))) { + return failure(); + } + + Value comparisonResult; + if (failed(matchComparisonResult(currValue, currIndex, reduceValue, + reduceIndex, it, comparisonResult, + isArgMin))) { + return failure(); + } + + // matching: %15 = arith.ori %14, %13 : i1 + auto orOp = dyn_cast(*it++); + if (!orOp || orOp.getLhs() != comparisonResult + || orOp.getRhs() != tieResult) { + return failure(); + } + + shouldUpdate = orOp; + return success(); +} + +LogicalResult matchSelect(mlir::Block::iterator &opsIt, + Value curr, Value reduce, + Value shouldUpdate, Value &result) { + auto selectOp = dyn_cast(*opsIt++); + if (!selectOp) { + return failure(); + } + + if (selectOp.getCondition() != shouldUpdate || + curr != selectOp.getTrueValue() || + reduce != selectOp.getFalseValue()) { + return failure(); + } + + result = selectOp; + + return success(); +} + +LogicalResult matchArgMinMax(Value currValue, Value currIndex, + Value reduceValue, Value reduceIndex, + mlir::Block::iterator &opsIt, + Value &indexResult, Value& valueResult, + bool isArgMin) { + Value shouldUpdate; + if (failed(matchShouldUpdateValue(currValue, currIndex, reduceValue, + reduceIndex, opsIt, shouldUpdate, + isArgMin))) { + return failure(); + } + + // matching: %16 = arith.select %15, %arg9, %arg11 : f32 + Value valueSelectOp; + if (failed(matchSelect(opsIt, currValue, reduceValue, + shouldUpdate, valueSelectOp))) { + return failure(); + } + + // matching:%17 = arith.select %15, %arg10, %arg12 : i32 + Value indexSelectOp; + if (failed(matchSelect(opsIt, currIndex, reduceIndex, + shouldUpdate, indexSelectOp))) { + return failure(); + } + + indexResult = indexSelectOp; + valueResult = valueSelectOp; + + return success(); +} + +} + +#endif \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Utils/ReduceScanCommon.h b/third_party/wafer/third_party/flir/include/triton-shared/Utils/ReduceScanCommon.h new file mode 100755 index 00000000..b0406468 --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Utils/ReduceScanCommon.h @@ -0,0 +1,353 @@ +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Utils/IndexingUtils.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include + +namespace mlir { +namespace triton { + +// Base class for converting scans and reductions. +// +// It provides accumulation function that clones operations from the +// original combine region and applies them on provided tensor. +// Also, it handles multi-dimensional cases reducing them to two +// possible options: lowering for a 1-D tensor inputs and lowering +// the operation over the leading dimension. +// +// Specialized pattern should implement lower1DInput to handle +// trailing dimension case and lowerLeadingDimension to handle the leading +// dimension case through accumulation of sub-tensors. +template +struct ReduceScanOpConversionBase : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + using OpConversionPattern::getTypeConverter; + using typename OpConversionPattern::OpAdaptor; + + virtual SmallVector + lower1DInput(ValueRange inputs, OpT op, + ConversionPatternRewriter &rewriter) const = 0; + virtual SmallVector + lowerLeadingDimension(ValueRange inputs, OpT op, + ConversionPatternRewriter &rewriter) const = 0; + + virtual uint32_t getAxis(OpT op) const = 0; + + virtual SmallVector getInputs(OpT op) const = 0; + + LogicalResult + matchAndRewrite(OpT op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto rank = cast(op.getOperand(0).getType()).getRank(); + auto axis = getAxis(op); + assert(axis < rank && "Expected axis is within the input rank"); + if (axis == (rank - 1)) + return lowerTrailingDimension(op, rewriter); + + return lowerNonTrailingDimension(op, rewriter); + } + + // To handle the trailing dimension case, we extract all input vectors + // and process them through lower1DInput, then build the resulting + // vector using inserts. + LogicalResult + lowerTrailingDimension(OpT op, ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + SmallVector inputs; + if (failed(rewriter.getRemappedValues(getInputs(op), inputs))) + return failure(); + + auto inputType = cast(inputs[0].getType()); + + // 1-D input case. + if (inputType.getRank() == 1) { + auto res = lower1DInput(inputs, op, rewriter); + rewriter.replaceOp(op, res); + return success(); + } + + // TODO: Optimization: The last two dimensions' data can be read with column + // stride and computed in parallel. + uint32_t axis = getAxis(op); + assert(axis == (inputType.getRank() - 1) && + "Expected reduction axis is the last one"); + SmallVector res = + lowering(inputs, op, rewriter, axis, + &ReduceScanOpConversionBase::lower1DInput); + + rewriter.replaceOp(op, res); + return success(); + } + + // In this case we either call lowerLeadingDimension to process the input + // or extract sub-vectors, call lowerLeadingDimension, and then reconstruct + // the result. + LogicalResult + lowerNonTrailingDimension(OpT op, ConversionPatternRewriter &rewriter) const { + + SmallVector inputs; + if (failed(rewriter.getRemappedValues(getInputs(op), inputs))) + return failure(); + + uint32_t axis = getAxis(op); + if (axis == 0) { + rewriter.replaceOp(op, lowerLeadingDimension(inputs, op, rewriter)); + return success(); + } + + SmallVector res = lowering( + inputs, op, rewriter, axis, + &ReduceScanOpConversionBase::lowerLeadingDimension); + + rewriter.replaceOp(op, res); + return success(); + } + + // Accumulate inputs and existing accumulators into a new accumulators + // applying operations from the combine region. + SmallVector accumulate(ValueRange inputs, ValueRange acc, + Region &combineOp, OpBuilder &rewriter) const { + if (acc.empty()) + return inputs; + + auto type = inputs[0].getType(); + SmallVector shape; + if (isa(type)) { + auto temp = cast(type).getShape(); + shape.insert(shape.end(), temp.begin(), temp.end()); + } // else shape is empty for scalar types + auto &block = combineOp.getBlocks().front(); + IRMapping map; + // Map block arguments to the current inputs and accumulators. + for (unsigned i = 0; i < acc.size(); ++i) { + map.map(block.getArgument(i), acc[i]); + map.map(block.getArgument(acc.size() + i), inputs[i]); + } + for (auto &op : block.getOperations()) { + // Returned values are a new accumulator. + if (isa(op)) { + SmallVector res; + for (auto operand : op.getOperands()) { + res.push_back(map.lookup(operand)); + } + return res; + } + + // Clone operation mapping its inputs and building vector + // result types using the input shape. + OperationState newState(op.getLoc(), op.getName()); + for (auto operand : op.getOperands()) { + newState.operands.push_back( + lookupMappedValue(map, operand, shape, rewriter)); + } + for (auto ty : op.getResultTypes()) { + isa(type) + ? newState.types.push_back(RankedTensorType::get(shape, ty)) + : newState.types.push_back(ty); + } + newState.attributes = op.getAttrs(); + auto newOp = rewriter.create(newState); + + // Add new values to the map. + for (auto [oldVal, newVal] : + llvm::zip(op.getResults(), newOp->getResults())) { + map.map(oldVal, newVal); + } + } + llvm_unreachable("No return op found in scan/reduce region"); + } + + Value lookupMappedValue(IRMapping &localMap, Value val, + ArrayRef shape, OpBuilder &rewriter) const { + + // First check in the local mapping + if (Value localMapped = localMap.lookupOrNull(val)) { + return localMapped; + } + + // Delete invariantsMap lookup: val needs to transform differents shape + // tensor. For example, 64->32 needs tensor<32Xf32> , 32->16 needs + // tensor<16xf32>. + // TODO: Profile it to improve performance. Beacause aboved cases(64->32, + // 32->16) maybe create different buffers. + + // Then, if the value is of the expected shape, return it directly + Type valueType = val.getType(); + if ((!isa(valueType) && shape.empty()) || + (isa(valueType) && + cast(valueType).getShape() == shape)) { + // TODO: Check rank tensor when shape is empty. If shape is empty, should + // add extract op. + return val; + } + + // Finally, if value is not found then it's an invariant defined in the + // outer region. We check if it has been already translated and add a + // linalg.fill operation if value shape is different. + auto ip = rewriter.saveInsertionPoint(); + rewriter.setInsertionPointAfterValue(val); + auto ty = isa(valueType) + ? cast(valueType).getElementType() + : valueType; + auto empty = rewriter.create(val.getLoc(), shape, ty); + Value res = rewriter.create( + val.getLoc(), ValueRange{val}, ValueRange{empty}).getResult(0); + rewriter.restoreInsertionPoint(ip); + return res; + } + + std::tuple, SmallVector, SmallVector, + SmallVector> + tensorTransform(OpBuilder &rewriter, Location loc, + SmallVector loopIndices, + RankedTensorType tensorType, uint32_t axis) const { + auto shape = tensorType.getShape(); + + SmallVector dynamicIndices = loopIndices; + dynamicIndices.insert(dynamicIndices.end(), shape.size() - axis, + rewriter.create(loc, 0)); + + SmallVector staticSize(axis, 1); + staticSize.insert(staticSize.end(), shape.begin() + axis, shape.end()); + SmallVector staticStride(shape.size(), 1); + + // {1,1,..(shape[axis])?,shape[axis+1],..shape[rank]} + SmallVector extractShape = staticSize; + + // {1,1,..(shape[axis])?,shape[axis+1],..shape[rank]} -> + // {(shape[axis])?,shape[axis+1],..shape[rank]} + auto reassociationRank = shape.size() - axis; + SmallVector reassociation(reassociationRank); + if (reassociationRank) { + // The first group: [0, 1, ..., axis - 1] + reassociation[0].resize(axis); + std::iota(reassociation[0].begin(), reassociation[0].end(), 0); + // The remaining groups: [axis, axis+1, axis+2, ..., shape.size()-1] + for (size_t i = axis; i < shape.size(); ++i) { + reassociation[i - axis].push_back(i); + } + } + + return std::tuple, SmallVector, + SmallVector, SmallVector>( + dynamicIndices, staticSize, staticStride, reassociation); + } + + using LoweringFuncType = + SmallVector (ReduceScanOpConversionBase::*)( + ValueRange inputs, OpT op, ConversionPatternRewriter &rewriter) const; + + // Though function ptr to call lower1DInput/lowerLeadingDimension to handle + // the tensor. + SmallVector lowering(ValueRange inputs, OpT op, + ConversionPatternRewriter &rewriter, + uint32_t axis, LoweringFuncType handle) const { + auto loc = op.getLoc(); + auto inputType = cast(inputs[0].getType()); + auto inputShape = inputType.getShape(); + + auto outputType = cast(op.getResults()[0].getType()); + auto outputShape = outputType.getShape(); + SmallVector res(inputs.size()); + std::transform(inputs.begin(), inputs.end(), res.begin(), [&](auto val) { + auto valType = cast(val.getType()); + return rewriter.create(loc, outputShape, + valType.getElementType()); + }); + + SmallVector loops; + SmallVector loopIndices; + + std::function buildLoops = [&](unsigned d) { + // Setup loop bounds and step. + Value lowerBound = rewriter.create(loc, 0); + Value upperBound = + rewriter.create(loc, inputShape[d]); + Value step = rewriter.create(loc, 1); + + SmallVector curRes = + d == 0 ? res : loops[d - 1].getInitArgs().take_front(res.size()); + auto loop = rewriter.create(loc, lowerBound, upperBound, step, + ValueRange{curRes}); + auto loopIdx = loop.getInductionVar(); + loopIndices.push_back(loopIdx); + loops.push_back(loop); + + rewriter.setInsertionPointToStart(loop.getBody()); + if (d == axis - 1) { + SmallVector subInputs(inputs.size()); + // [Dynamic indices, static size, static stride, reassociation] + auto [inputDynamicIndices, inputStaticSize, inputStaticStride, + inputReassociation] = + tensorTransform(rewriter, loc, loopIndices, inputType, axis); + for (size_t i = 0; i < inputs.size(); ++i) { + auto valueType = cast(inputs[i].getType()); + auto extractTensor = rewriter.create( + loc, + RankedTensorType::get(inputStaticSize, + valueType.getElementType()), + inputs[i], inputDynamicIndices, /*sizes*/ ValueRange(), + /*strides*/ ValueRange(), + SmallVector(inputShape.size(), ShapedType::kDynamic), + inputStaticSize, inputStaticStride); + subInputs[i] = rewriter.create( + loc, extractTensor, inputReassociation); + } + + auto resElems = (this->*handle)(subInputs, op, rewriter); + + // [Dynamic indices, static size, static stride, reassociation] + auto [outputDynamicIndices, outputStaticSize, outputStaticStride, + outputReassociation] = + tensorTransform(rewriter, loc, loopIndices, outputType, axis); + + for (size_t i = 0; i < res.size(); ++i) { + auto resType = cast(res[i].getType()); + auto targetType = RankedTensorType::get(outputStaticSize, + resType.getElementType()); + // {shape[axis],shape[axis+1],..shape[rank]} -> + // {1,1,..shape[axis],shape[axis+1],..shape[rank]} + Value reshaped = rewriter.create( + loc, targetType, resElems[i], outputReassociation); + + curRes[i] = rewriter.create( + loc, reshaped, curRes[i], outputDynamicIndices, + /*sizes*/ ValueRange(), + /*strides*/ ValueRange(), + SmallVector(outputShape.size(), ShapedType::kDynamic), + outputStaticSize, outputStaticStride); + } + rewriter.create(loc, curRes); + return; + } + buildLoops(d + 1); + // Terminate the loop body. + rewriter.setInsertionPointToEnd(loop.getBody()); + SmallVector yieldValues = + loops[d + 1].getResults().take_front(res.size()); + rewriter.create(loc, yieldValues); + rewriter.setInsertionPointAfter(loop); + }; + + buildLoops(0); + + // Extract result tensors from forOp; + SmallVector results; + for (size_t i = 0; i < inputs.size(); ++i) { + results.push_back(loops.front().getResult(i)); + } + return results; + } + +private: + mutable IRMapping invariantsMap; +}; + +TypedAttr getRedBaseAttr(OpBuilder &builder, Operation *redOp, + Type constantType); + +arith::ConstantOp getRedBaseConstOp(ConversionPatternRewriter &rewriter, + Operation *redOp, Type constantType); + +} // namespace triton +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/include/triton-shared/Utils/Utils.h b/third_party/wafer/third_party/flir/include/triton-shared/Utils/Utils.h new file mode 100755 index 00000000..7afe8c1c --- /dev/null +++ b/third_party/wafer/third_party/flir/include/triton-shared/Utils/Utils.h @@ -0,0 +1,69 @@ +#ifndef TRITON_SHARED_UTILITY_H +#define TRITON_SHARED_UTILITY_H + +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinOps.h" +namespace mlir { +namespace triton { +// Return true if the input type is a triton pointer or a tensor of triton pointers +bool isPtrTypeLike(Type t); + +// Extract a scalar value from v. +// If v is a scalar, return that directly. Otherwise, parse through operations +// (currently only support splat, sitofp, and truncf) that produce it to +// extract the underlying scalar value. We then reconstruct the chain of +// operations that can produce this constant with the original type. If no +// scalar value can be extracted, a nullptr is returned. +Value getScalarValue(Value operand, Location loc, OpBuilder &builder); + +Value declareWaferRuntimeFunction(ModuleOp module, OpBuilder &builder, Location loc, + StringRef name, Type resultType, + ArrayRef argumentTypes); + +bool isOperandMemorySpaceSPM(Value operand); + +// NOTE: Reduction can lower to target reduction ops or target elementwise ops. + +// Target support reduce instruction +// Different target ops have different supported types restriction +bool isTypeRestrictedTargetSupportedReductionOp(mlir::Operation *redOp); +// Integer and float type reduction op +bool isTargetSupportedReductionOp(mlir::Operation *redOp); + +// Reduce to elementwise op +bool isTypeRestrictedTargetSupportedReduceToElementWiseOp( + mlir::Operation *redOp); +// Integer and float type reduction op +bool isTargetSupportedReduceToElementWiseOp(mlir::Operation *redOp); + +bool isTritonAllowedReductionOp(Operation *redOp); + +// Reduction ops and elementwise op only support float types. +bool isTargetSupportedFloatType(Type elementType); +bool isTargetSupportedType(Type elementType); + +bool isReductionOpAndTypeSupportedByTarget(mlir::Operation *redOp, + Type elementType); +bool isReduceToElementWiseOpAndTypeSupportedByTarget(mlir::Operation *redOp, + Type elementType, + int64_t elemCount, + int64_t rank); + +template llvm::SmallVector getRegionOps(T linalgOp) { + assert((linalgOp->getNumRegions() == 1 && + linalgOp->getRegion(0).hasOneBlock()) && + "It only applies to ops with one region and one block!"); + auto regionBlock = linalgOp.getBlock(); + return llvm::map_to_vector(regionBlock->without_terminator(), + [](Operation &op) { return &op; }); +} + +} // namespace triton + +} // namespace mlir + +// Gelu mode. +enum class GeluMode { None = 0, Tanh = 1 }; + +#endif // TRITON_SHARED_UTILITY_H diff --git a/third_party/wafer/third_party/flir/lib/Analysis/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Analysis/CMakeLists.txt new file mode 100755 index 00000000..71a89f94 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Analysis/CMakeLists.txt @@ -0,0 +1,14 @@ +add_triton_library(TritonSharedAnalysis + MaskAnalysis.cpp + OpFoldResultUtils.cpp + PtrAnalysis.cpp + UseAnalysis.cpp + + DEPENDS + TritonTableGen + TritonStructuredTableGen + TritonGPUAttrDefsIncGen + + LINK_LIBS PUBLIC + MLIRAnalysis +) diff --git a/third_party/wafer/third_party/flir/lib/Analysis/MaskAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Analysis/MaskAnalysis.cpp new file mode 100755 index 00000000..905f7c0a --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Analysis/MaskAnalysis.cpp @@ -0,0 +1,821 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// +#include "flagtree/Common/UnifiedHardware.h" + +#include "triton-shared/Analysis/MaskAnalysis.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Support/LogicalResult.h" + +#include "triton-shared/Analysis/OpFoldResultUtils.h" + +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/LogicalResult.h" +#include +#define DEBUG_TYPE "MaskAnalysis" +#include "llvm/Support/Debug.h" +namespace mlir { + +namespace triton { +///////////ascend +void dimInfo::dump() const { + LLVM_DEBUG({ + llvm::dbgs() << "MaskDimInfo: \n" ; + llvm::dbgs() << "dim = " << dim << "\n"; + llvm::dbgs() << "shape = " << shape << "\n"; + llvm::dbgs() << "div = " << div << "\n"; + llvm::dbgs() << "isSlt = " << isSlt << "\n"; + llvm::dbgs() << "rhs = " << rhs << "\n"; + llvm::dbgs() << "isRealDim = " << isRealDim << "\n"; + }); +}; + + +/////////////ascend +LogicalResult MaskState::parse(Value operand, const Location loc, + OpBuilder &builder) { + auto hardwareManager = mlir::flagtree::createUnifiedHardwareManager(); + auto incubatedTag = hardwareManager -> getIncubatedTag(); + if (auto op = operand.getDefiningOp()) { + return this->parseConstant(op, loc, builder); + } else if (isa(operand.getType())) { + return this->parseIntScalar(operand, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseAdd(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseAnd(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseCmp(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseMakeRange(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseBroadcast(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseSplat(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseExpandDims(op, loc, builder); + } else if (!operand.getDefiningOp()) { + return this->parseLoopIterArg(operand, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseExtSI(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseSub(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + if (incubatedTag) + return this->parseRemsi(op, loc, builder); + else + return failure(); + } else if (auto op = operand.getDefiningOp()) { + if (incubatedTag) + return this->parseDivsi(op, loc, builder); + else + return failure(); + } + else { + return failure(); + } +} + +tensor::ExtractSliceOp MaskState::getExtractSlice(Value source, + const Location loc, + OpBuilder &builder) const { + auto sourceType = cast(source.getType()); + SmallVector offsets(getRank(), builder.getIndexAttr(0)); + SmallVector strides(getRank(), builder.getIndexAttr(1)); + + auto dstType = tensor::ExtractSliceOp::inferResultType(sourceType, offsets, + dims, strides); + + return builder.create(loc, dstType, source, offsets, + dims, strides); +} + +memref::SubViewOp MaskState::getSubview(Value source, const Location loc, + OpBuilder &builder) const { + auto sourceType = cast(source.getType()); + SmallVector offsets(getRank(), builder.getIndexAttr(0)); + SmallVector strides(getRank(), builder.getIndexAttr(1)); + auto dstType = + memref::SubViewOp::inferResultType(sourceType, offsets, dims, strides); + + return builder.create(loc, cast(dstType), + source, offsets, dims, strides); +} + +static memref::SubViewOp createSubview(Value src, Location loc, OpBuilder &b, + ArrayRef offsets, + ArrayRef sizes, + ArrayRef strides) { + auto srcType = cast(src.getType()); + auto dstType = + memref::SubViewOp::inferResultType(srcType, offsets, sizes, strides); + return b.create(loc, cast(dstType), src, + offsets, sizes, strides); +} + +// Assume block1 wraps around and the remainder is block2. +// +// |----------------------| +// | | | +// | block2 | block1 | +// | | | +// |----------------------| +// +// Once we copy the chunks in order, the end result is block1 followed by +// block2. +// +// buffer_tmp: +// +// |----------------------| +// | | | +// | block1 | block2 | +// | | | +// |----------------------| +// +// Assume we have the following subview: +// +// +++++++++++++++++------- +// + + | +// + subview + | +// + + | +// +++++++++++++++++------- +// +// If we simply take the subview of `buffer_tmp`, this requires an extra +// buffer to just hold the temporary result. +// +// So we can subview into block1 and block2 directly. There are 2 cases: +// + subview only spans block1 +// + subview spans both block1 and block2, creating sv1 and sv2 (illustrated +// below for case when we wrap around side-by-side) +// +// |----------------------------------------| +// | | +// | col2 col1 | +// |++++++--------| |+++++++++++++++ +// | sv2 + block2 | | block1 & sv1 + +// |++++++--------| |+++++++++++++++ +// | | +// |----------------------------------------| +// +// For simplicity, assume we only wrap around side-by-side. +// +// Let (row, col1) and (row, col2) be the dimensions of block1 and block2, +// respectively. +// +// Let (rowFull, colFull), (rowView1, colView1) and (rowView2, colView2) be +// the dimensions of the full subview, sv1, and sv2, respectively. +// +// + colView1 = min(colFull, col1) +// + colView2 = colFull - colView1 +// + rowView1 = rowView2 = row = rowFull +std::pair +MaskState::getSideBySideSubviews(Value block1, Value block2, const Location loc, + OpBuilder &builder) const { + OpFoldResult subviewRowFull = dims[0]; + OpFoldResult subviewColFull = dims[1]; + OpFoldResult col1 = builder.create(loc, block1, 1).getResult(); + OpFoldResult subviewCol1 = minOFRs(col1, subviewColFull, loc, builder); + OpFoldResult subviewCol2 = subOFRs(subviewColFull, subviewCol1, loc, builder); + + SmallVector offsets(getRank(), builder.getIndexAttr(0)); + SmallVector strides(getRank(), builder.getIndexAttr(1)); + auto sv1 = createSubview(block1, loc, builder, offsets, + {subviewRowFull, subviewCol1}, strides); + auto sv2 = createSubview(block2, loc, builder, offsets, + {subviewRowFull, subviewCol2}, strides); + + return {sv1, sv2}; +} + +std::pair +MaskState::getStackedSubviews(Value block1, Value block2, const Location loc, + OpBuilder &builder) const { + OpFoldResult subviewRowFull = dims[0]; + OpFoldResult subviewColFull = dims[1]; + OpFoldResult row1 = builder.create(loc, block1, 0).getResult(); + OpFoldResult subviewRow1 = minOFRs(row1, subviewRowFull, loc, builder); + OpFoldResult subviewRow2 = subOFRs(subviewRowFull, subviewRow1, loc, builder); + + SmallVector offsets(getRank(), builder.getIndexAttr(0)); + SmallVector strides(getRank(), builder.getIndexAttr(1)); + auto sv1 = createSubview(block1, loc, builder, offsets, + {subviewRow1, subviewColFull}, strides); + auto sv2 = createSubview(block2, loc, builder, offsets, + {subviewRow2, subviewColFull}, strides); + return {sv1, sv2}; +} + +LogicalResult MaskState::addStateScalar(const MaskState &state, + const OpFoldResult scalar, Location loc, + OpBuilder &builder) { + start = addOFRs(state.start, scalar, loc, builder); + end = addOFRs(state.end, scalar, loc, builder); + dims = state.dims; + return success(); +} + +LogicalResult MaskState::subStateScalar(const MaskState &state, + const OpFoldResult scalar, Location loc, + OpBuilder &builder) { + start = subOFRs(state.start, scalar, loc, builder); + end = subOFRs(state.end, scalar, loc, builder); + dims = state.dims; + return success(); +} + +LogicalResult MaskState::subStates(const MaskState &lhsState, + const MaskState &rhsState, Location loc, + OpBuilder &builder) { + if (lhsState.scalar && rhsState.scalar) { + InFlightDiagnostic diag = + emitError(loc) << "Unexpected case where both lhs and rhs are scalars"; + return failure(); + } + + if (!lhsState.scalar && !rhsState.scalar) { + InFlightDiagnostic diag = + emitError(loc) + << "Unsupported scenario where neither lhs nor rhs is a scalar"; + return failure(); + } + + if (lhsState.scalar) + return subStateScalar(rhsState, lhsState.scalar, loc, builder); + else + return subStateScalar(lhsState, rhsState.scalar, loc, builder); +} + +LogicalResult MaskState::addStates(const MaskState &lhsState, + const MaskState &rhsState, Location loc, + OpBuilder &builder) { + if (lhsState.scalar && rhsState.scalar) { + InFlightDiagnostic diag = + emitError(loc) << "Unexpected case where both lhs and rhs are scalars"; + return failure(); + } + + if (!lhsState.scalar && !rhsState.scalar) { + InFlightDiagnostic diag = + emitError(loc) + << "Unsupported scenario where neither lhs nor rhs is a scalar"; + return failure(); + } + + if (lhsState.scalar) + return addStateScalar(rhsState, lhsState.scalar, loc, builder); + else + return addStateScalar(lhsState, rhsState.scalar, loc, builder); +} + +LogicalResult MaskState::minStateScalar(const MaskState &lhsState, + const MaskState &rhsState, Location loc, + OpBuilder &builder) { + if (lhsState.scalar && rhsState.scalar) { + dims.push_back(minOFRs(lhsState.dims[0], rhsState.dims[0], loc, builder)); + } else if (lhsState.scalar) { + for (uint32_t i = 0; i < rhsState.getRank(); i++) { + auto lhsDim = lhsState.dims[0]; + auto rhsDim = rhsState.dims[i]; + dims.push_back(minOFRs(lhsDim, rhsDim, loc, builder)); + } + } else if (rhsState.scalar) { + for (uint32_t i = 0; i < lhsState.getRank(); i++) { + auto lhsDim = lhsState.dims[i]; + auto rhsDim = rhsState.dims[0]; + dims.push_back(minOFRs(lhsDim, rhsDim, loc, builder)); + } + } else { + InFlightDiagnostic diag = + emitError(loc) << "Unexpected case where both lhs and rhs are not scalars"; + return failure(); + } + return success(); +} + +LogicalResult MaskState::minStates(const MaskState &lhsState, + const MaskState &rhsState, Location loc, + OpBuilder &builder) { + if (lhsState.getRank() != rhsState.getRank()) { + InFlightDiagnostic diag = + emitError(loc) + << "Unexpected case where lhs and rhs have different ranks"; + return failure(); + } + + for (uint32_t i = 0; i < lhsState.getRank(); i++) { + auto lhsDim = lhsState.dims[i]; + auto rhsDim = rhsState.dims[i]; + dims.push_back(minOFRs(lhsDim, rhsDim, loc, builder)); + } + return success(); +} + +LogicalResult MaskState::parseConstant(arith::ConstantOp constOp, + const Location loc, OpBuilder &builder) { + assert(this->isEmpty()); + + if (isa(constOp.getValue())) { + auto attr = cast(constOp.getValue()); + auto elementType = attr.getElementType(); + assert(attr.isSplat() && isa(elementType) && + "All elements must share a single integer constant value"); + auto values = attr.getValues(); + auto value = values[0].getValue(); + auto constAttr = builder.getIndexAttr(value.getSExtValue()); + auto op = arith::ConstantOp::materialize(builder, constAttr, + builder.getIndexType(), loc); + this->scalar = op.getValue(); + } else { + auto value = cast(constOp.getValue()).getInt(); + this->scalar = builder.getIndexAttr(value); + } + + return success(); +} + +LogicalResult MaskState::parseIntScalar(Value scalar, const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + auto castOp = + builder.create(loc, builder.getIndexType(), scalar); + this->scalar = castOp.getResult(); + return success(); +} + +void MaskState::dump() const { + llvm::dbgs() << "start: " << start << "\n"; + llvm::dbgs() << "end: " << end << "\n"; + llvm::dbgs() << "scalar: " << scalar << "\n"; + llvm::dbgs() << "useUnsafeMask: " << useUnsafeMask << "\n"; + llvm::dbgs() << "dims: "; + for (auto dim : dims) + llvm::dbgs() << "\t" << dim << "\n"; + llvm::dbgs() << "\n"; +} + +LogicalResult MaskState::parseAdd(arith::AddIOp addOp, const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + MaskState lhsState; + if (failed(lhsState.parse(addOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(addOp.getRhs(), loc, builder))) + return failure(); + + return this->addStates(lhsState, rhsState, loc, builder); +} + +LogicalResult MaskState::parseSub(arith::SubIOp subOp, const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + MaskState lhsState; + if (failed(lhsState.parse(subOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(subOp.getRhs(), loc, builder))) + return failure(); + + return this->subStates(lhsState, rhsState, loc, builder); +} + +LogicalResult MaskState::parseAnd(arith::AndIOp andOp, const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + MaskState lhsState; + if (failed(lhsState.parse(andOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(andOp.getRhs(), loc, builder))) + return failure(); + + // TODO(FLIR): should be isMask() + if(!lhsState.isMaskWithoutScalar() || !rhsState.isMaskWithoutScalar()) { + return this->minStateScalar(lhsState, rhsState, loc, builder); + } + return this->minStates(lhsState, rhsState, loc, builder); +} + +LogicalResult MaskState::parseExtSI(arith::ExtSIOp op, const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + return parse(op.getIn(), loc, builder); +} + +LogicalResult MaskState::parseCmp(arith::CmpIOp cmpOp, const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + if (cmpOp.getPredicate() != arith::CmpIPredicate::slt && + cmpOp.getPredicate() != arith::CmpIPredicate::ult && + cmpOp.getPredicate() != arith::CmpIPredicate::sge) { + InFlightDiagnostic diag = emitError(loc) << "Unsupported cmpi"; + return failure(); + } + + MaskState lhsState; + if (failed(lhsState.parse(cmpOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(cmpOp.getRhs(), loc, builder))) + return failure(); + + // We only support sge against 0 for lower bounds. Dims already has an + // implicit assumption that the lower bound is 0, so if we see this, assume + // the comparison evaluates to true. + if (cmpOp.getPredicate() == arith::CmpIPredicate::sge + && !(rhsState.scalar && hasConstZero(rhsState.scalar))) { + InFlightDiagnostic diag = emitError(loc) + << "Unsupported cmpi with rhs not equal to 0"; + return failure(); + } + + int32_t cmpDim = lhsState.scalar && rhsState.scalar ? 0 : -1; + for (int32_t i = 0; i < lhsState.getRank(); i++) { + auto dimIntAttr = getIntAttr(lhsState.dims[i]); + if (!dimIntAttr || dimIntAttr.value() != 1) { + if (cmpDim != -1) { + InFlightDiagnostic diag = emitError(loc) + << "Unsupported cmpi with more than one " + "dimension with size larger than 1"; + return failure(); + } + cmpDim = i; + } + } + assert(cmpDim != -1 && + "Unexpected case where no dimension has size larger than 1"); + + OpFoldResult newDim; + if (lhsState.scalar) { + assert(rhsState.scalar && "Unexpected case where rhs is not a scalar"); + // If both lhs and rhs are scalars, we can't just derive the dimension of + // the mask as the minimum value: lhs/rhs could be 0 and then we don't + // load/store anything. + // + // Instead treat the comparison as a scalar that determines if anything + // should be loaded/stored by inserting a comparison + select: + // dim = lhs < rhs ? lhs.dim : 0 + newDim = compareOFRs(lhsState.scalar, rhsState.scalar, cmpOp.getPredicate(), + lhsState.dims[cmpDim], builder.getIndexAttr(0), + loc, builder); + } else if (cmpOp.getPredicate() == arith::CmpIPredicate::slt || + cmpOp.getPredicate() == arith::CmpIPredicate::ult) { + // Important: + // In the case where the values we are loading are entirely masked off like + // the following: + // + // ---|-------|-----------| + // ^ ^ ^ + // scalar start end + // + // newEnd = min(end, scalar) = scalar + // Now scalar < start, so simply doing dim = newEnd - start is incorrect. + // + // The correct formula is to optionally move `newDim` back to `start` using + // max(newEnd, start). + auto newEnd = minOFRs(lhsState.end, rhsState.scalar, loc, builder); + newEnd = maxOFRs(newEnd, lhsState.start, loc, builder); + newDim = subOFRs(newEnd, lhsState.start, loc, builder); + } else { + assert(cmpOp.getPredicate() == arith::CmpIPredicate::sge && rhsState.scalar + && hasConstZero(rhsState.scalar)); + newDim = lhsState.dims[cmpDim]; + } + + for (int32_t i = 0; i < lhsState.getRank(); i++) { + if (i == cmpDim) + this->dims.push_back(newDim); + else + this->dims.push_back(lhsState.dims[i]); + } + + return success(); +} + +LogicalResult MaskState::parseLoopIterArg(Value v, const Location loc, + OpBuilder &builder) { + assert(!v.getDefiningOp()); + + auto forOp = llvm::dyn_cast(v.getParentRegion()->getParentOp()); + + if (!forOp) { + return failure(); + } + + // TODO: This implementation does not work with nested loops + if (forOp->getParentOfType()) { + return failure(); + } + + auto it = llvm::find(forOp.getRegionIterArgs(), v); + if (it == forOp.getRegionIterArgs().end()) { + return failure(); + } + + auto argIndex = std::distance(forOp.getRegionIterArgs().begin(), it); + auto initArg = forOp.getInitArgs()[argIndex]; + if (auto getStateOp = initArg.getDefiningOp()) { + auto tritonValue = getStateOp->getOperand(0); + MaskState lhsState; + if (failed(lhsState.parse(tritonValue, loc, builder))) { + return failure(); + } + +#ifdef FLAGTREE_BACKEND_WAFER + if (llvm::isa_and_nonnull(tritonValue.getDefiningOp())) { + // It's accurately a size 1 tl.arange op + auto constOp = tritonValue.getDefiningOp(); + if (llvm::succeeded(this->parseConstant(constOp, loc, builder))) { + this->dims.push_back(builder.getIndexAttr(1)); + return success(); + } + return failure(); + } +#endif + // This is a bit of a hack!! + // + // The offsets and dimensions of a MaskState can now depend on a loop's + // iter-arg. + // + // Because the PtrAnalysis's pre-pass already sets up the offsets, + // we can create a new MaskState for each loop iteration by adding the + // original MaskState with the current iter-arg, which is at `argIndex + + // 1`. + // + // This will not work for nested loop scenarios, which would need a + // more robust implementation. + if (failed(this->addStateScalar( + lhsState, forOp.getRegionIterArgs()[argIndex + 1], loc, builder))) { + return failure(); + } + + return success(); + } + + return failure(); +} + +LogicalResult MaskState::parseMakeRange(triton::MakeRangeOp rangeOp, + const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + auto shape = cast(rangeOp.getType()).getShape(); + auto start = rangeOp.getStart(); + auto end = rangeOp.getEnd(); + auto stride = (end - start + shape[0] - 1) / shape[0]; + + if (stride != 1) { + InFlightDiagnostic diag = + emitError(loc) + << "stride must be 1 for make_range whose result is used " + "as load or store masks"; + return failure(); + } + + this->start = builder.getIndexAttr(start); + this->end = builder.getIndexAttr(end); + this->dims.push_back(builder.getIndexAttr(shape[0])); + + return success(); +} + +LogicalResult MaskState::parseBroadcast(triton::BroadcastOp broadcastOp, + const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + auto src = broadcastOp.getSrc(); + auto dst = broadcastOp.getResult(); + assert(isa(src.getType()) && + "input to tt.broadcast should be a tensor"); + + auto srcShape = cast(src.getType()).getShape(); + auto dstShape = cast(dst.getType()).getShape(); + assert(srcShape.size() == dstShape.size() && + "rank of source and destination should match"); + + if (failed(parse(src, loc, builder))) + return failure(); + + for (size_t i = 0; i < srcShape.size(); i++) { + if (srcShape[i] == dstShape[i]) + continue; + else if (srcShape[i] < dstShape[i]) + this->dims[i] = builder.getIndexAttr(dstShape[i]); + else + llvm_unreachable("unexpected dimensions used in broadcast"); + } + + return success(); +} + +LogicalResult MaskState::parseSplat(triton::SplatOp splatOp, const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + auto src = splatOp.getSrc(); + auto dst = splatOp.getResult(); + auto dstShape = cast(dst.getType()).getShape(); + + if (!isa(src.getType())) { + InFlightDiagnostic diag = + emitError(loc) + << "splat source must be an integer scalar for load/store masks"; + return failure(); + } + + if (failed(this->parse(src, loc, builder))) + return failure(); + + for (auto s : dstShape) + this->dims.push_back(builder.getIndexAttr(s)); + + return success(); +} + +LogicalResult MaskState::parseExpandDims(triton::ExpandDimsOp expandDimsOp, + const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + if (failed(this->parse(expandDimsOp.getSrc(), loc, builder))) + return failure(); + + auto dstShape = + cast(expandDimsOp.getResult().getType()).getShape(); + auto axis = expandDimsOp.getAxis(); + assert(dstShape[axis] == 1 && + "expect changed dimension to be 1 in expand_dims"); + this->dims.insert(this->dims.begin() + axis, builder.getIndexAttr(1)); + + return success(); +} +////////ASCEND +LogicalResult MaskState::parseRemsi(arith::RemSIOp remsiOp, + const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + auto defaultAttr = builder.getIndexAttr(0); + + MaskState lhsState; + if (failed(lhsState.parse(remsiOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(remsiOp.getRhs(), loc, builder))) + return failure(); + + if(lhsState.scalar || !rhsState.scalar){ + remsiOp->emitRemark("Unsupported remsi scenario"); + return failure(); + } + + int64_t staticShape; + + if(auto value = rhsState.scalar.dyn_cast()){ + auto constop = value.getDefiningOp(); + staticShape = cast(constop.getValue()).getInt(); + }else if(auto rhsIntAttr = getIntAttr(rhsState.scalar)){ + assert(rhsIntAttr.has_value()); + staticShape = rhsIntAttr.value(); + }else{ + remsiOp->emitError("MaskAnalysis: Static compilation cannot determine the value of this parameter"); + return failure(); + } + + start = lhsState.start; + dims = lhsState.dims; + end = minOFRs(lhsState.end, rhsState.scalar, loc, builder); + stateInfo = lhsState.stateInfo; + for(auto &info: stateInfo){ + if(info.isRealDim){ + auto staticDim = getIntAttr(rhsState.scalar); + if(!staticDim.has_value() || (staticDim.value() % staticShape != 0 && staticShape % staticDim.value() != 0)){ + remsiOp->emitError("MaskAnalysis: The shape of the mask is not divisible by the shape of the block"); + return failure(); + } + // if(getIntAttr(info.div).has_value() && getIntAttr(info.div).value() != 0){ + // remsiOp->emitError("MaskAnalysis: do not support remsi after div"); + // return failure(); + // } + info.shape = builder.getIndexAttr(staticShape); + } + } + + return success(); +} + +LogicalResult MaskState::parseDivsi(arith::DivSIOp divsiOp, + const Location loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + auto defaultAttr = builder.getIndexAttr(0); + + MaskState lhsState; + if (failed(lhsState.parse(divsiOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(divsiOp.getRhs(), loc, builder))) + return failure(); + + if(lhsState.scalar || !rhsState.scalar){ + divsiOp->emitRemark("Unsupported divsi scenario"); + return failure(); + } + + int64_t staticDiv; + + if(auto value = rhsState.scalar.dyn_cast()){ + auto constop = value.getDefiningOp(); + staticDiv = cast(constop.getValue()).getInt(); + }else if(auto rhsIntAttr = getIntAttr(rhsState.scalar)){ + assert(rhsIntAttr.has_value()); + staticDiv = rhsIntAttr.value(); + }else{ + divsiOp->emitError("MaskAnalysis: Static compilation cannot determine the value of this parameter"); + return failure(); + } + + start = divOFRs(lhsState.start, rhsState.scalar, loc, builder); + dims = lhsState.dims; + auto minEnd = addOFRs(start, builder.getIndexAttr(1), loc, builder); + end = subOFRs(lhsState.end, builder.getIndexAttr(1), loc, builder); + end = divOFRs(end, rhsState.scalar, loc, builder); + end = addOFRs(end, builder.getIndexAttr(1), loc, builder); + end = maxOFRs(end, minEnd, loc, builder); + stateInfo = lhsState.stateInfo; + for(auto &info: stateInfo){ + if(info.isRealDim){ + auto staticDim = getIntAttr(rhsState.scalar); + if(!staticDim.has_value() || (staticDim.value() % staticDiv != 0 && staticDiv % staticDim.value() != 0)){ + divsiOp->emitError("MaskAnalysis: The shape of the mask is not divisible by the shape of the block"); + return failure(); + } + if(getIntAttr(info.shape).has_value() && getIntAttr(info.shape).value() != 0){ + divsiOp->emitError("MaskAnalysis: do not support div after remsi"); + return failure(); + } + info.div = builder.getIndexAttr(staticDiv); + } + } + + return success(); +} + +void MaskState::eraseInsertedOps(Operation *rawOp, PatternRewriter &rewriter) { + auto moduleOp = rawOp->getParentOfType(); + SmallVector worklist; + moduleOp->walk([&](Operation *op) { + if (isOpTriviallyDead(op)) + worklist.push_back(op); + }); + while (!worklist.empty()) { + Operation *op = worklist.pop_back_val(); + if (!isOpTriviallyDead(op)) + continue; + for (Value value : op->getOperands()) { + if (auto defOp = value.getDefiningOp()) + worklist.push_back(defOp); + } + LLVM_DEBUG({ + llvm::dbgs() << "[MaskState]==> inserted op: \n" + << *op << "\n[MaskState]<== is removed\n"; + }); + rewriter.eraseOp(op); + } +} + +tensor::InsertSliceOp MaskState::getInsertSlice(Value source, Value dest, + const Location &loc, + OpBuilder &builder) const { + auto sourceType = cast(source.getType()); + //fixme, kaixin offsets are class member originally + SmallVector offsets(getRank(), builder.getIndexAttr(0)); + SmallVector strides(getRank(), builder.getIndexAttr(1)); + return builder.create(loc, source, dest, offsets, dims, + strides); +} + +} // namespace triton +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/Analysis/OpFoldResultUtils.cpp b/third_party/wafer/third_party/flir/lib/Analysis/OpFoldResultUtils.cpp new file mode 100755 index 00000000..410023a0 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Analysis/OpFoldResultUtils.cpp @@ -0,0 +1,396 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/Analysis/OpFoldResultUtils.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/Transforms/DialectConversion.h" + +namespace mlir { + +#if !defined(__FLIR_BUILD_INCUBATED__) +std::optional getIntAttr(const OpFoldResult ofr) { + if (isa(ofr) && isa(cast(ofr))) + return dyn_cast(cast(ofr)).getInt(); + + return std::nullopt; +} +#endif + +bool hasConstZero(const OpFoldResult ofr) { + auto intAttr = getIntAttr(ofr); + if (intAttr.has_value()) { + if (intAttr.value() == 0) { + return true; + } + return false; + } + + auto val = dyn_cast(ofr); + assert(val); + auto constOp = val.getDefiningOp(); + if (!constOp) + return false; + + intAttr = getIntAttr(constOp.getValue()); + if (intAttr.has_value()) { + if (intAttr.value() == 0) { + return true; + } + return false; + } + + return false; +} + +Value ofrToIndexValue(const OpFoldResult ofr, const Location loc, + OpBuilder &b) { + if (Value val = dyn_cast(ofr)) { + assert(val.getType().isIntOrIndex()); + if (!val.getType().isIndex()) { + val = b.create(loc, b.getIndexType(), val); + } + return val; + } + + auto intVal = getIntAttr(ofr); + if (intVal.has_value()) { + return b.create(loc, b.getIndexAttr(intVal.value())); + } + llvm_unreachable("Unexpected OpFoldResult state"); + return nullptr; +} + +SmallVector ofrsToIndexValues(ArrayRef ofrs, + const Location loc, OpBuilder &b) { + return llvm::to_vector<4>( + llvm::map_range(ofrs, [&](OpFoldResult ofr) -> Value { + return ofrToIndexValue(ofr, loc, b); + })); +} + +OpFoldResult addOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + auto rhsIntAttr = getIntAttr(rhs); + + // shortcut for special cases + if (!lhsIntAttr && rhsIntAttr && rhsIntAttr.value() == 0) + return lhs; + if (!rhsIntAttr && lhsIntAttr && lhsIntAttr.value() == 0) + return rhs; + + // both lhs and rhs are constants, return result directly + if (lhsIntAttr && rhsIntAttr) + return b.getIndexAttr(lhsIntAttr.value() + rhsIntAttr.value()); + + // otherwise, need to create instructions to calculate new attribute value + auto lhsValue = dyn_cast(lhs); + if (lhsIntAttr) { + auto lhsOp = + b.create(loc, b.getIndexAttr(lhsIntAttr.value())); + lhsValue = lhsOp.getResult(); + } else { + assert(isa(lhsValue.getType())); + } + + auto rhsValue = dyn_cast(rhs); + if (rhsIntAttr) { + auto rhsOp = + b.create(loc, b.getIndexAttr(rhsIntAttr.value())); + rhsValue = rhsOp.getResult(); + } else { + assert(isa(lhsValue.getType())); + } + + return b.create(loc, lhsValue, rhsValue).getResult(); +} + +OpFoldResult subOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + auto rhsIntAttr = getIntAttr(rhs); + + // shortcut for special cases + if (!lhsIntAttr && rhsIntAttr && rhsIntAttr.value() == 0) + return lhs; + + // both lhs and rhs are constants, return result directly + if (lhsIntAttr && rhsIntAttr) + return b.getIndexAttr(lhsIntAttr.value() - rhsIntAttr.value()); + + // otherwise, need to create instructions to calculate new attribute value + auto lhsValue = dyn_cast(lhs); + if (lhsIntAttr) { + auto lhsOp = + b.create(loc, b.getIndexAttr(lhsIntAttr.value())); + lhsValue = lhsOp.getResult(); + } + + auto rhsValue = dyn_cast(rhs); + if (rhsIntAttr) { + auto rhsOp = + b.create(loc, b.getIndexAttr(rhsIntAttr.value())); + rhsValue = rhsOp.getResult(); + } + + auto sumOp = b.create(loc, lhsValue, rhsValue); + return sumOp.getResult(); +} + +OpFoldResult mulOFRValue(const OpFoldResult lhs, const Value rhs, + const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + + auto rhsIsConst = false; + // if rhs is not a const, use max value since min is used to represent + // dynamic size or stride + auto rhsConstValue = std::numeric_limits::max(); + auto rhsOp = rhs.getDefiningOp(); + if (rhsOp) { + rhsIsConst = true; + rhsConstValue = cast(rhsOp.getValue()).getInt(); + } + + // shortcuts for special cases + if (lhsIntAttr) { + if (lhsIntAttr.value() == 0) + return lhs; + if (lhsIntAttr.value() == 1) + return rhs; + } + if (rhsIsConst) { + if (rhsConstValue == 0) + return rhsOp.getResult(); + if (rhsConstValue == 1) + return lhs; + } + + // 0. both lhs and rhs are constants + if (lhsIntAttr && rhsIsConst) + return b.getIndexAttr(lhsIntAttr.value() * rhsConstValue); + + // 1. if lhs is constant but rhs is not + if (lhsIntAttr && !rhsIsConst) { + auto lhsConstOp = + b.create(loc, b.getIndexAttr(lhsIntAttr.value())); + auto mulOp = b.create(loc, lhsConstOp.getResult(), rhs); + return mulOp.getResult(); + } + + // 2. if lhs is not constant + assert(!lhsIntAttr); + auto mulOp = b.create(loc, cast(lhs), rhs); + return mulOp.getResult(); +} +//////ASCEND +OpFoldResult divOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + auto rhsIntAttr = getIntAttr(rhs); + + // both lhs and rhs are constants, return result directly + if (lhsIntAttr && rhsIntAttr) + return b.getIndexAttr(lhsIntAttr.value() / rhsIntAttr.value()); + + // shortcut for special cases + if (rhsIntAttr && rhsIntAttr.value() == 1) + return lhs; + + // otherwise, need to create instructions to calculate new attribute value + auto lhsValue = lhs.dyn_cast(); + if (lhsIntAttr) { + auto lhsOp = + b.create(loc, b.getIndexAttr(lhsIntAttr.value())); + lhsValue = lhsOp.getResult(); + } + + auto rhsValue = rhs.dyn_cast(); + if (rhsIntAttr) { + auto rhsOp = + b.create(loc, b.getIndexAttr(rhsIntAttr.value())); + rhsValue = rhsOp.getResult(); + } + + auto divOp = b.create(loc, lhsValue, rhsValue); + return divOp.getResult(); +} + +OpFoldResult remOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + auto rhsIntAttr = getIntAttr(rhs); + + // both lhs and rhs are constants, return result directly + if (lhsIntAttr && rhsIntAttr) + return b.getIndexAttr(lhsIntAttr.value() % rhsIntAttr.value()); + + // shortcut for special cases + if (rhsIntAttr && rhsIntAttr.value() == 1) + return b.getIndexAttr(0); + + // otherwise, need to create instructions to calculate new attribute value + auto lhsValue = lhs.dyn_cast(); + if (lhsIntAttr) { + auto lhsOp = + b.create(loc, b.getIndexAttr(lhsIntAttr.value())); + lhsValue = lhsOp.getResult(); + } + + auto rhsValue = rhs.dyn_cast(); + if (rhsIntAttr) { + auto rhsOp = + b.create(loc, b.getIndexAttr(rhsIntAttr.value())); + rhsValue = rhsOp.getResult(); + } + + auto remOp = b.create(loc, lhsValue, rhsValue); + return remOp.getResult(); +} + +OpFoldResult mulOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + auto rhsIntAttr = getIntAttr(rhs); + + // both lhs and rhs are constants, return result directly + if (lhsIntAttr && rhsIntAttr) + return b.getIndexAttr(lhsIntAttr.value() * rhsIntAttr.value()); + + // shortcuts for special cases + if (lhsIntAttr) { + if (lhsIntAttr.value() == 0) + return lhs; + if (lhsIntAttr.value() == 1) + return rhs; + } + if (rhsIntAttr) { + if (rhsIntAttr.value() == 0) + return rhs; + if (rhsIntAttr.value() == 1) + return lhs; + } + + + // otherwise, need to create instructions to calculate new attribute value + auto lhsValue = dyn_cast(lhs); + if (lhsIntAttr) { + auto lhsOp = + b.create(loc, b.getIndexAttr(lhsIntAttr.value())); + lhsValue = lhsOp.getResult(); + } + + auto rhsValue = dyn_cast(rhs); + if (rhsIntAttr) { + auto rhsOp = + b.create(loc, b.getIndexAttr(rhsIntAttr.value())); + rhsValue = rhsOp.getResult(); + } + + auto mulOp = b.create(loc, lhsValue, rhsValue); + return mulOp.getResult(); +} +///////ASCNED +OpFoldResult minOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + auto rhsIntAttr = getIntAttr(rhs); + + // both lhs and rhs are constants, return result directly + if (lhsIntAttr && rhsIntAttr) + return b.getIndexAttr(std::min(lhsIntAttr.value(), rhsIntAttr.value())); + + // otherwise, need to create instructions to calculate new attribute value + auto lhsValue = dyn_cast(lhs); + if (lhsIntAttr) { + auto lhsOp = + b.create(loc, b.getIndexAttr(lhsIntAttr.value())); + lhsValue = lhsOp.getResult(); + } + + auto rhsValue = dyn_cast(rhs); + if (rhsIntAttr) { + auto rhsOp = + b.create(loc, b.getIndexAttr(rhsIntAttr.value())); + rhsValue = rhsOp.getResult(); + } + + auto minOp = b.create(loc, lhsValue, rhsValue); + return minOp.getResult(); +} + +OpFoldResult maxOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + auto rhsIntAttr = getIntAttr(rhs); + + // both lhs and rhs are constants, return result directly + if (lhsIntAttr && rhsIntAttr) + return b.getIndexAttr(std::max(lhsIntAttr.value(), rhsIntAttr.value())); + + // otherwise, need to create instructions to calculate new attribute value + auto lhsValue = dyn_cast(lhs); + if (lhsIntAttr) { + auto lhsOp = + b.create(loc, b.getIndexAttr(lhsIntAttr.value())); + lhsValue = lhsOp.getResult(); + } + + auto rhsValue = dyn_cast(rhs); + if (rhsIntAttr) { + auto rhsOp = + b.create(loc, b.getIndexAttr(rhsIntAttr.value())); + rhsValue = rhsOp.getResult(); + } + + auto maxOp = b.create(loc, lhsValue, rhsValue); + return maxOp.getResult(); +} + +OpFoldResult compareOFRs(const OpFoldResult lhs, const OpFoldResult rhs, + const arith::CmpIPredicate pred, const OpFoldResult trueOFR, + const OpFoldResult falseOFR, const Location loc, OpBuilder &b) { + auto lhsIntAttr = getIntAttr(lhs); + auto rhsIntAttr = getIntAttr(rhs); + + // both lhs and rhs are constants, return the result directly + if (lhsIntAttr && rhsIntAttr) { + switch (pred) { + case arith::CmpIPredicate::eq: + return *lhsIntAttr == *rhsIntAttr ? trueOFR : falseOFR; + case arith::CmpIPredicate::ne: + return *lhsIntAttr != *rhsIntAttr ? trueOFR : falseOFR; + case arith::CmpIPredicate::slt: + case arith::CmpIPredicate::ult: + return *lhsIntAttr < *rhsIntAttr ? trueOFR : falseOFR; + case arith::CmpIPredicate::sle: + case arith::CmpIPredicate::ule: + return *lhsIntAttr <= *rhsIntAttr ? trueOFR : falseOFR; + case arith::CmpIPredicate::sgt: + case arith::CmpIPredicate::ugt: + return *lhsIntAttr > *rhsIntAttr ? trueOFR : falseOFR; + case arith::CmpIPredicate::sge: + case arith::CmpIPredicate::uge: + return *lhsIntAttr >= *rhsIntAttr ? trueOFR : falseOFR; + default: + llvm_unreachable("Unsupported predicate"); + } + } + + auto lhsValue = ofrToIndexValue(lhs, loc, b); + auto rhsValue = ofrToIndexValue(rhs, loc, b); + auto trueValue = ofrToIndexValue(trueOFR, loc, b); + auto falseValue = ofrToIndexValue(falseOFR, loc, b); + + auto cmpOp = b.create(loc, pred, lhsValue, rhsValue); + auto selectOp = b.create(loc, cmpOp, trueValue, falseValue); + return selectOp.getResult(); +} +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/Analysis/PtrAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Analysis/PtrAnalysis.cpp new file mode 100755 index 00000000..00715a9d --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Analysis/PtrAnalysis.cpp @@ -0,0 +1,1375 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/Analysis/PtrAnalysis.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" + +#include "mlir/IR/IRMapping.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/Support/Debug.h" +#include + +#define DEBUG_TYPE "triton-ptr-analysis" + +namespace mlir { + +namespace triton { + +static void assertValidUnrealizedCast(UnrealizedConversionCastOp op) { + assert(op && op->hasAttr(ModuloState::WraparoundAttr) && + op.getInputs().size() == 3 && + op.getInputs()[0].getDefiningOp() && + op.getInputs()[1].getDefiningOp() && + op.getInputs()[2].getDefiningOp()); +} + +MemRefType PtrState::getResultMemrefType(MLIRContext *context, int64_t offset, + ArrayRef resultShape, + bool useDynamicStrides) const { + + SmallVector staticStrides; + if (useDynamicStrides) { + staticStrides.append(strides.size(), ShapedType::kDynamic); + } else { + SmallVector dynamicStrides; + dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides); + } + + auto elementType = cast(source.getType()).getElementType(); + auto layout = + StridedLayoutAttr::get(source.getContext(), offset, staticStrides); + + return MemRefType::get(resultShape, elementType, layout); +} + +OpFoldResult +PtrState::accumulateTargetOffset(Location loc, + ConversionPatternRewriter &rewriter) const { + OpFoldResult targetOffset = rewriter.getIndexAttr(0); + for (auto o : offsets) { + targetOffset = addOFRs(targetOffset, o, loc, rewriter); + } + return targetOffset; +} + +int64_t PtrState::getRank() const { + assert(offsets.size() == sizes.size() && offsets.size() == strides.size() && + modulos.size() == offsets.size()); + return offsets.size(); +} + +bool PtrState::isEmpty() const { + return (getRank() == 0 && !source && !scalar); +} + +bool PtrState::hasModulo() const { + return llvm::any_of(modulos, [](auto mod) { return mod.has_value(); }); +} + +void PtrState::addState(const PtrState &lhsState, const PtrState &rhsState, + Location loc, ConversionPatternRewriter &rewriter) { + assert(isEmpty() && lhsState.getRank() == rhsState.getRank()); + + // at most one of lhs and rhs should have valid source, since otherwise we + // will be losing information + assert(!(lhsState.source && rhsState.source)); + source = lhsState.source ? lhsState.source : rhsState.source; + + if (lhsState.scalar && rhsState.scalar) { + auto addOp = + rewriter.create(loc, lhsState.scalar, rhsState.scalar); + scalar = addOp.getResult(); + } else if (lhsState.getRank() == 0) { // both lhs and rhs are scalars + scalar = lhsState.scalar ? lhsState.scalar : rhsState.scalar; + } + + for (uint64_t i = 0; i < lhsState.sizes.size(); i++) { + auto newOffset = + addOFRs(lhsState.offsets[i], rhsState.offsets[i], loc, rewriter); + offsets.push_back(newOffset); + + auto newStride = + addOFRs(lhsState.strides[i], rhsState.strides[i], loc, rewriter); + strides.push_back(newStride); + + sizes.push_back(lhsState.sizes[i]); + + assert(!lhsState.hasModulo() || + !rhsState.hasModulo() && "AddPtr where both lhs and rhs containing " + "modulo operators not supported"); + + modulos.push_back(lhsState.modulos[i].has_value() ? lhsState.modulos[i] + : rhsState.modulos[i]); + } +} + +void PtrState::mulState(const PtrState &lhsState, const PtrState &rhsState, + const Location loc, + ConversionPatternRewriter &rewriter) { + assert(isEmpty() && lhsState.getRank() == rhsState.getRank()); + + // neither lhs nor rhs should have source, since multiplying base pointer + // does not make sense + assert(!(lhsState.source && rhsState.source)); + + assert((lhsState.scalar || rhsState.scalar) && + !(lhsState.scalar && rhsState.scalar) && + "currently does not support both tensors are effectively non-scalar"); + + PtrState const *lhs = &lhsState; + PtrState const *rhs = &rhsState; + + if (!rhs->scalar && lhs->scalar) { + std::swap(lhs, rhs); + } + + for (uint64_t i = 0; i < lhs->sizes.size(); i++) { + OpFoldResult newOffset = + mulOFRValue(lhs->offsets[i], rhs->scalar, loc, rewriter); + OpFoldResult newStride = + mulOFRValue(lhs->strides[i], rhs->scalar, loc, rewriter); + offsets.push_back(newOffset); + strides.push_back(newStride); + sizes.push_back(lhs->sizes[i]); + } + + assert(llvm::all_of(rhsState.modulos, + [](auto rhs) { return !rhs.has_value(); })); + + modulos = lhs->modulos; +} + +SmallVector +PtrState::createStackedCastOps(ArrayRef resultShape, + const Location loc, + ConversionPatternRewriter &rewriter) const { + + assert(resultShape.size() == 2); + assert(getRank() == 2); + assert(modulos[0].has_value() && !modulos[1].has_value()); + + Value targetOffset = + ofrToIndexValue(accumulateTargetOffset(loc, rewriter), loc, rewriter); + + ////////////////////////////////////////////////////////////////////////////// + // + // Handling stacked wraparound + // + // We do not support cases where the target offset has already overflown the + // number of rows. See side-by-side wraparound for details. + // + ////////////////////////////////////////////////////////////////////////////// + // We're loading a tensor of dim (rowSize, colSize) + // d1 + d2 = rowSize + // d2 is the number of rows that overflow + // + // cols + // + // wrappedAroundOff + // --------------*------------*-------- + // | d2 | | | + // | |------------| | + // rows| | + // | | + // | targetOffset | + // | *------------| | + // | | | | + // | d1 | | | + // | | clampedOff | | + // --------------*--------------------- + // | overflow | + // *------------- + // nextOff + // + // wrappedAroundOff = targetOffset % cols + // clampedOff = (rows * strideRows) + wrappedAroundOff + // + // clampedOff - targetOffset + // d1 = -------------------- + // strideRows + + auto resultType = getResultMemrefType( + rewriter.getContext(), /* offset */ ShapedType::kDynamic, + /* result shape */ + SmallVector{ + ShapedType::kDynamic, // Row is dynamic, in most cases, this should be + // the same as the original row. The last chunk + // may be smaller due to wrapping around. + resultShape[1], // Col stays the same. + }, + true /*useDynamicStrides*/); + + Value rowSize = ofrToIndexValue(sizes[0], loc, rewriter); + Value colSize = ofrToIndexValue(sizes[1], loc, rewriter); + + Value strideRow = ofrToIndexValue(strides[0], loc, rewriter); + Value strideCol = ofrToIndexValue(strides[1], loc, rewriter); + + Value modRow = rewriter.create( + loc, rewriter.getIndexType(), modulos[0]->size); + + // First chunk + Value wrappedAroundOff = + rewriter.create(loc, targetOffset, strideRow); + Value clampedOff = rewriter.create(loc, modRow, strideRow); + clampedOff = + rewriter.create(loc, clampedOff, wrappedAroundOff); + Value d1 = rewriter.create(loc, clampedOff, targetOffset); + d1 = rewriter.create(loc, d1, strideRow); + + SmallVector sizes1{d1, colSize}; + memref::ReinterpretCastOp cast1 = rewriter.create( + loc, resultType, source, targetOffset, sizes1, + ValueRange{strideRow, strideCol}); + + // Second chunk + Value d2 = rewriter.create(loc, rowSize, d1); + SmallVector sizes2{d2, colSize}; + memref::ReinterpretCastOp cast2 = rewriter.create( + loc, resultType, source, wrappedAroundOff, sizes2, + ValueRange{strideRow, strideCol}); + + return {cast1, cast2}; +} + +SmallVector +PtrState::createSideBySideCastOps(ArrayRef resultShape, + const Location loc, + ConversionPatternRewriter &rewriter) const { + + assert(resultShape.size() == 2); + assert(getRank() == 2 && !modulos[0].has_value() && modulos[1].has_value()); + + // Accumulate final offset + Value targetOffset = + ofrToIndexValue(accumulateTargetOffset(loc, rewriter), loc, rewriter); + + ////////////////////////////////////////////////////////////////////////////// + // + // Handling side-by-side wraparound + // + // Note: We do not support cases where the target has already overflown the + // number of columns! This is because in PtrAnalysis, the offset has already + // been collapsed into a single dimension, so it is ambiguous to determine + // whether the offset actually overflows or just refers to an element on the + // subsequent rows. + // + // Same limitations apply to the stacked wraparound case. + // + ////////////////////////////////////////////////////////////////////////////// + // + // nextOffset - targetOffset = colSize + // d1 + d2 = colSize + // N + // x clampedOffset + // --------------------------*----------------*-----* + // | | nextOffset (might + // | targetOffset | overflow) + // y *----- *----------------| + // | | | | + // M |----- -----------------| + // | d2 d1 | + // -------------------------------------------- + // + // x = targetOffset % N + // nextOffset = x + colSize + // clampedOffset = min(nextOffset, N) + // d1 = clampedOffset - x + // + ////////////////////////////////////////////////////////////////////////////// + + SmallVector casts; + + auto resultType = getResultMemrefType( + rewriter.getContext(), /* offset */ ShapedType::kDynamic, + /* result shape */ + SmallVector{ + resultShape[0], // Row stays the same + ShapedType::kDynamic // Column is dynamic, in most cases, this should + // be the same as the original column. The last + // chunk may be smaller due to wrapping around. + }, + true /*useDynamicStrides*/); + + Value rowSize = ofrToIndexValue(sizes[0], loc, rewriter); + Value colSize = ofrToIndexValue(sizes[1], loc, rewriter); + + Value modN = rewriter.create(loc, rewriter.getIndexType(), + modulos[1]->size); + + Value x = rewriter.create(loc, targetOffset, modN); + Value y = rewriter.create(loc, targetOffset, x); + + SmallVector strideVals = ofrsToIndexValues(strides, loc, rewriter); + + // First chunk + Value nextOffset = rewriter.create(loc, x, colSize); + Value clampedOffset = rewriter.create(loc, nextOffset, modN); + Value d1 = rewriter.create(loc, clampedOffset, x); + SmallVector sizes1{rowSize, d1}; + + auto cast1 = rewriter.create( + loc, resultType, source, targetOffset, sizes1, strideVals); + + // Second chunk + Value d2 = rewriter.create(loc, colSize, d1); + SmallVector sizes2{rowSize, d2}; + + auto cast2 = rewriter.create( + loc, resultType, source, y, sizes2, strideVals); + + return {cast1, cast2}; +} + +memref::ReinterpretCastOp +PtrState::createCastOp(ArrayRef resultShape, const Location loc, + ConversionPatternRewriter &rewriter) const { + // Accumulate final offset + OpFoldResult targetOffset = accumulateTargetOffset(loc, rewriter); + + // Create result MemRefType + SmallVector staticOffset; + SmallVector dynamicOffset; + dispatchIndexOpFoldResult(targetOffset, dynamicOffset, staticOffset); + + auto resultType = + getResultMemrefType(rewriter.getContext(), staticOffset[0], resultShape); + + // Create reinterpret cast + return rewriter.create( + loc, resultType, source, targetOffset, sizes, strides); +} + +void PtrAnalysis::visitOperandAdd( + arith::AddIOp addOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + PtrState lhsState; + visitOperand(addOp.getLhs(), lhsState, loc, rewriter, knownPtrs); + + PtrState rhsState; + visitOperand(addOp.getRhs(), rhsState, loc, rewriter, knownPtrs); + + if ((lhsState.getRank() == 1 && lhsState.hasModulo()) || + (rhsState.getRank() == 1 && rhsState.hasModulo())) { + assert(0 && "Current do not support this pattern: a + arange(0, K) % M"); + } + + state.addState(lhsState, rhsState, loc, rewriter); +} + +void PtrAnalysis::visitOperandMul( + arith::MulIOp mulOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + PtrState lhsState; + visitOperand(mulOp.getLhs(), lhsState, loc, rewriter, knownPtrs); + + PtrState rhsState; + visitOperand(mulOp.getRhs(), rhsState, loc, rewriter, knownPtrs); + + state.mulState(lhsState, rhsState, loc, rewriter); +} + +void PtrAnalysis::visitOperandRem( + arith::RemSIOp remOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + + PtrState rhsState; + visitOperand(remOp.getRhs(), rhsState, loc, rewriter, knownPtrs); + assert(rhsState.scalar); + + visitOperand(remOp.getLhs(), state, loc, rewriter, knownPtrs); + + // If there are multiple modulo ops on an expression (e.g.: (a % b) % c), we + // would have already populated the modulo states after visiting the lhs. + // Assert that all the modulo states are empty. + assert(llvm::all_of(state.modulos, + [](auto modState) { return !modState.has_value(); }) && + "No support for multiple modulo within an expression"); + + if (state.getRank() == 1) { + // Apply the modulo before expanding shape, the common pattern is + // offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + // a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * + // stride_ak) + state.modulos.back() = ModuloState{rhsState.scalar}; + } else if (state.getRank() == 2) { + // torch inductor expands the tensor shape before applying the modulo. + // + // We only support either: + // - (tl.arange(0, end)[:, None] % mod), or + // - (tl.arange(0, end)[None, :] % mod) + // + // In both cases, we apply the modulo to the non-singleton dimension. + auto shape = cast(remOp.getResult().getType()).getShape(); + if (shape[0] == 1) { + state.modulos[1] = ModuloState{rhsState.scalar}; + } else if (shape[1] == 1) { + state.modulos[0] = ModuloState{rhsState.scalar}; + } else { + assert(false && "Taking modulo on a 2D tensor with no singleton " + "dimension not supported"); + } + } else { + assert(false && "Unsupported modulo pattern"); + } +} + +void PtrAnalysis::visitOperandMakeRange( + triton::MakeRangeOp rangeOp, PtrState &state, Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + + auto shape = cast(rangeOp.getType()).getShape(); + + auto start = rangeOp.getStart(); + auto end = rangeOp.getEnd(); + auto stride = (end - start + shape[0] - 1) / shape[0]; + assert(stride == 1 && + "Expect make_range op to always return tensor of stride 1"); + + state.offsets.push_back(rewriter.getIndexAttr(start)); + state.sizes.push_back(rewriter.getIndexAttr(shape[0])); + state.strides.push_back(rewriter.getIndexAttr(stride)); + state.modulos.push_back(std::nullopt); +} + +void PtrAnalysis::visitOperandExpandDims( + triton::ExpandDimsOp expandDimsOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + + // `getSrc` now returns a TypedValue of RankedTensorType. We modify these + // operands in-place and turn them into memrefs in loops, so we have to bypass + // the cast by using getSrcMutable. These are temporary fix only since + // we will be moving over to StructuredPtrAnalysis soon which separate out the + // memref conversion. + visitOperand(expandDimsOp.getSrcMutable().get(), state, loc, rewriter, + knownPtrs); + + auto dstShape = + cast(expandDimsOp.getResult().getType()).getShape(); + auto axis = expandDimsOp.getAxis(); + + assert(dstShape[axis] == 1 && + "expect changed dimension to be 1 in expand_dims"); + + // insert dimension info + state.offsets.insert(state.offsets.begin() + axis, rewriter.getIndexAttr(0)); + state.sizes.insert(state.sizes.begin() + axis, rewriter.getIndexAttr(1)); + state.strides.insert(state.strides.begin() + axis, rewriter.getIndexAttr(0)); + state.modulos.insert(state.modulos.begin() + axis, std::nullopt); +} + +void PtrAnalysis::visitOperandBroadcast( + triton::BroadcastOp broadcastOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + + // `getSrc` now returns a TypedValue of RankedTensorType. We modify these + // operands in-place and turn them into memrefs in loops, so we have to bypass + // the cast by using getSrcMutable. These are temporary fix only since + // we will be moving over to StructuredPtrAnalysis soon which separate out the + // memref conversion. + auto src = broadcastOp.getSrcMutable().get(); + auto dst = broadcastOp.getResult(); + assert(isa(src.getType()) && + "input to tt.broadcast should be a tensor"); + + auto srcShape = cast(src.getType()).getShape(); + auto dstShape = cast(dst.getType()).getShape(); + assert(srcShape.size() == dstShape.size() && + "rank of source and destination should match"); + + visitOperand(src, state, loc, rewriter, knownPtrs); + + for (size_t i = 0; i < srcShape.size(); i++) { + if (srcShape[i] == dstShape[i]) + continue; + else if (srcShape[i] < dstShape[i]) + state.sizes[i] = rewriter.getIndexAttr(dstShape[i]); + else + llvm_unreachable("unexpected dimensions used in broadcast"); + } +} + +void PtrAnalysis::visitOperandSplat( + triton::SplatOp splatOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + + auto src = splatOp.getSrc(); + auto dst = splatOp.getResult(); + auto dstShape = cast(dst.getType()).getShape(); + + visitOperand(src, state, loc, rewriter, knownPtrs); + + if (isa(src.getType())) { + for (auto s : dstShape) { + state.offsets.push_back(rewriter.getIndexAttr(0)); + state.sizes.push_back(rewriter.getIndexAttr(s)); + state.strides.push_back(rewriter.getIndexAttr(0)); + state.modulos.push_back(std::nullopt); + } + } else { + // src is a memref that represent a scalar pointer; it should have + // one dimension of size 1. This happens inside a for loop that + // originally has an init arg that is a tensor of pointers; this arg + // would have been replaced by rewriteForOp. + auto srcType = cast(src.getType()); + assert(srcType.getRank() == 1 && state.getRank() == 1 && + "splat MemRef source should have rank 1"); + assert(srcType.getShape()[0] == 1 && + getIntAttr(state.sizes[0]).value() == 1 && + "splat MemRef source should have size 1"); + + // Stride[0] will have value of 1 set in visitOperandAddPtr. This + // value will be represented by a constOp. Clear this value. + state.strides[0] = rewriter.getIndexAttr(0); + + for (auto [i, s] : llvm::enumerate(dstShape)) { + if (i == 0) { + state.sizes[i] = rewriter.getIndexAttr(s); + continue; + } + state.offsets.push_back(rewriter.getIndexAttr(0)); + state.sizes.push_back(rewriter.getIndexAttr(s)); + state.strides.push_back(rewriter.getIndexAttr(0)); + state.modulos.push_back(std::nullopt); + } + } + + // If we splat a integer value, scalar should become the offset of the outer + // most dimension + if (state.scalar) + state.offsets[0] = state.scalar; +} + +void PtrAnalysis::visitOperandMakeTensorPtr( + triton::MakeTensorPtrOp makeTensorPtrOp, PtrState &state, + const Location loc, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + auto remappedValue = rewriter.getRemappedValue(makeTensorPtrOp); + if (auto castOp = remappedValue.getDefiningOp()) { + visitOperandReintCast(castOp, state, loc, rewriter, knownPtrs); + } else { + llvm_unreachable("Expect value to me mapped to a memref.reinterpret_cast"); + } +} + +void PtrAnalysis::visitOperandAddptr( + triton::AddPtrOp addptrOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + + PtrState ptrState; + visitOperand(addptrOp.getPtr(), ptrState, addptrOp.getLoc(), rewriter, + knownPtrs); + + PtrState offsetState; + visitOperand(addptrOp.getOffset(), offsetState, addptrOp.getLoc(), rewriter, + knownPtrs); + + assert(ptrState.source && "ptr field should provide source / base pointer"); + + // Handle the special case when we are in a for loop, ptr is originally a + // scalar pointer but replaced with a memref. In this case, ptrState will have + // rank 1 and offsetState will have rank 0. + // TODO: + // Passing a block argument pointer directly into a for loop not supported + if (ptrState.getRank() == 1 && offsetState.getRank() == 0) { + offsetState.sizes.push_back(rewriter.getIndexAttr(1)); + offsetState.offsets.push_back(offsetState.scalar); + offsetState.strides.push_back(rewriter.getIndexAttr(0)); + offsetState.modulos.push_back(std::nullopt); + } + + assert(ptrState.getRank() == offsetState.getRank() && + "ptr and offset field should have the same rank"); + + state.addState(ptrState, offsetState, addptrOp.getLoc(), rewriter); +} + +void PtrAnalysis::visitOperandReintCast( + memref::ReinterpretCastOp reintCastOp, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + + state.offsets = reintCastOp.getMixedOffsets(); + state.sizes = reintCastOp.getMixedSizes(); + state.strides = reintCastOp.getMixedStrides(); + state.source = reintCastOp.getSource(); + state.modulos.append(state.sizes.size(), std::nullopt); + + // getMixedOffsets produces staticOffsets (which is the result of collapsing + // multiple dimensions). Populate the rest of the dimensions with zeroes. + assert(state.offsets.size() == 1); + for (size_t i = 1; i < state.sizes.size(); i++) { + state.offsets.push_back(rewriter.getIndexAttr(0)); + } + + // Regular Triton programs cannot express patterns of size 1 and non-zero + // stride; we only set it that way to make memrefs work. Set stride back to + // zero if this scenario detected. + for (size_t i = 0; i < state.strides.size(); i++) { + auto strideIntAttr = getIntAttr(state.strides[i]); + auto sizeIntAttr = getIntAttr(state.sizes[i]); + + assert(sizeIntAttr); + if (sizeIntAttr.value() == 1 && strideIntAttr) { + state.strides[i] = rewriter.getIndexAttr(0); + } + } +} + +void PtrAnalysis::visitOperand( + Value operand, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + + if (knownPtrs.find(operand) != knownPtrs.end()) { + state = knownPtrs.lookup(operand); + return; + } + + if (isa(operand.getType())) { + auto castOp = rewriter.create( + loc, rewriter.getIndexType(), operand); + state.scalar = castOp.getResult(); + return; + } + + if (isa(operand.getType())) { + auto remappedPtr = rewriter.getRemappedValue(operand); + assert(remappedPtr); + + // A scalar pointer can either be produced by AddPtrOp or a block + // argument + if (auto op = operand.getDefiningOp()) { + if (auto addPtrOp = dyn_cast(op)) { + visitOperandAddptr(cast(op), state, loc, rewriter, + knownPtrs); + } else if (auto makeTensorOp = dyn_cast(op)) { + visitOperandMakeTensorPtr(makeTensorOp, state, loc, rewriter, + knownPtrs); + } else { + llvm_unreachable("Unexpected operand defining operation"); + } + } else { + state.source = remappedPtr; + } + return; + } + + if (auto op = operand.getDefiningOp()) { + visitOperandAdd(op, state, loc, rewriter, knownPtrs); + } else if (auto op = operand.getDefiningOp()) { + visitOperandMul(op, state, loc, rewriter, knownPtrs); + } else if (auto op = operand.getDefiningOp()) { + visitOperandMakeRange(op, state, loc, rewriter, knownPtrs); + } else if (auto op = operand.getDefiningOp()) { + visitOperandBroadcast(op, state, loc, rewriter, knownPtrs); + } else if (auto op = operand.getDefiningOp()) { + visitOperandSplat(op, state, loc, rewriter, knownPtrs); + } else if (auto op = operand.getDefiningOp()) { + visitOperandExpandDims(op, state, loc, rewriter, knownPtrs); + } else if (auto op = operand.getDefiningOp()) { + visitOperandAddptr(op, state, loc, rewriter, knownPtrs); + } else if (auto op = operand.getDefiningOp()) { + visitOperandConstSplat(op, state, loc, rewriter, knownPtrs); + } else if (auto op = operand.getDefiningOp()) { + visitOperandRem(op, state, loc, rewriter, knownPtrs); + } else { + operand.dump(); + llvm_unreachable("encountered addptr operand produced by an " + "unsupported operation"); + } +} + +void PtrAnalysis::visitOperandConstSplat( + arith::ConstantOp op, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assert(state.isEmpty()); + // this condition is to handle cases where tt.broadcast and tt.splat are + // folded + auto attr = cast(op.getValue()); + auto elementType = attr.getElementType(); + assert(attr.isSplat() && isa(elementType)); + auto values = attr.getValues(); + auto value = values[0].getValue(); + auto constAttr = rewriter.getIndexAttr(value.getSExtValue()); + auto constOp = arith::ConstantOp::materialize(rewriter, constAttr, + rewriter.getIndexType(), loc); + + state.scalar = constOp; + + auto resultType = cast(op.getResult().getType()); + for (size_t i = 0; i < resultType.getShape().size(); i++) { + if (i == 0) { + state.offsets.push_back(constOp.getResult()); + } else { + state.offsets.push_back(rewriter.getIndexAttr(0)); + } + + state.sizes.push_back(rewriter.getIndexAttr(resultType.getShape()[i])); + state.strides.push_back(rewriter.getIndexAttr(0)); + state.modulos.push_back(std::nullopt); + } +} + +void PtrAnalysis::rewriteAddptrOp( + triton::AddPtrOp op, ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &knownPtrs) { + // any inserted instruction should be before this addptr + auto origIp = rewriter.saveInsertionPoint(); + rewriter.setInsertionPoint(op); + + PtrState state; + visitOperandAddptr(op, state, op.getLoc(), rewriter, knownPtrs); + + // If the result is a scalar pointer, visitOperandAddptr will not populate + // sizes, strides, and offsets. We need to do it here. + if (state.sizes.size() == 0) { + state.sizes.push_back(rewriter.getIndexAttr(1)); + state.strides.push_back(rewriter.getIndexAttr(0)); + state.offsets.push_back(state.scalar); + state.modulos.push_back(std::nullopt); + } + + SmallVector scalarShape(1, 1); + ArrayRef resultShape; + if (auto shapedType = dyn_cast(op.getResult().getType())) { + resultShape = shapedType.getShape(); + } else { + // scalar pointer, should produce a one dimensional memref + resultShape = scalarShape; + assert(state.getRank() == 1); + } + + knownPtrs[op.getResult()] = state; + + // If there are dimensions with size 1 and stride 0, replace 0 stride with the + // product of sizes of all lower dimensions. This avoids creating memref with + // zero stride. Note that we store the unmodified state into knownPtrs, since + // any following pointer arithmetic operations should use the original 0 + // stride. + auto accum_size = 1; + for (int i = state.sizes.size() - 1; i >= 0; i--) { + auto strideIntAttr = getIntAttr(state.strides[i]); + auto sizeIntAttr = getIntAttr(state.sizes[i]); + + assert(sizeIntAttr); + if (sizeIntAttr.value() == 1 && strideIntAttr && strideIntAttr.value() == 0) + state.strides[i] = rewriter.getIndexAttr(accum_size); + + accum_size *= sizeIntAttr.value(); + } + + Value src; + + if (llvm::any_of(state.modulos, [](auto mod) { return mod.has_value(); })) { + assert(state.modulos.size() == 2); + ConversionPatternRewriter::InsertionGuard guard(rewriter); + rewriter.setInsertionPointAfter(op); + + SmallVector casts; + StringRef type; + + if (!state.modulos[0].has_value() && state.modulos[1].has_value()) { + casts = state.createSideBySideCastOps(resultShape, op.getLoc(), rewriter); + type = ModuloState::WraparoundSideBySide; + } else if (state.modulos[0].has_value() && !state.modulos[1].has_value()) { + casts = state.createStackedCastOps(resultShape, op.getLoc(), rewriter); + type = ModuloState::WraparoundStacked; + } else { + assert(false && "not supported"); + } + + auto resultType = state.getResultMemrefType( + rewriter.getContext(), ShapedType::kDynamic, resultShape); + + UnrealizedConversionCastOp combinedCast = + rewriter.create( + op.getLoc(), resultType, + ValueRange{casts[0].getResult(), casts[1].getResult(), + op.getResult()}); + + combinedCast->setAttr(ModuloState::WraparoundAttr, + rewriter.getStringAttr(type)); + + src = combinedCast.getResult(0); + + LLVM_DEBUG({ + llvm::dbgs() << "combine cast for split pointers:\n"; + combinedCast.getOperation()->print( + llvm::dbgs(), OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); + + } else { + memref::ReinterpretCastOp castOp = + state.createCastOp(resultShape, op.getLoc(), rewriter); + + src = castOp.getResult(); + + LLVM_DEBUG({ + llvm::dbgs() << "cast MemRefType:\n"; + castOp.getOperation()->print(llvm::dbgs(), + OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); + } + + state.source = src; + rewriter.replaceOp(op, src); + rewriter.restoreInsertionPoint(origIp); +} + +void PtrAnalysis::rewriteAdvanceOp( + triton::AdvanceOp op, ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &knownPtrs) { + OpBuilder::InsertionGuard insertionGuard{rewriter}; + rewriter.setInsertionPoint(op); + auto loc = op.getLoc(); + + PtrState ptrState; + visitOperand(op.getOperand(0), ptrState, loc, rewriter, knownPtrs); + + auto incrementOffsets = op.getOffsets(); + + SmallVector newOffsets; + for (auto [increment, offset, stride] : + llvm::zip(incrementOffsets, ptrState.offsets, ptrState.strides)) { + Value offsetValue; + if (auto offsetIntAttr = getIntAttr(offset)) { + auto constOp = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(0)); + offsetValue = constOp.getResult(); + } else { + offsetValue = cast(offset); + } + auto castOp = rewriter.create( + loc, rewriter.getIndexType(), increment); + auto mulOp = rewriter.create(loc, castOp.getResult(), + cast(stride)); + auto addOp = + rewriter.create(loc, mulOp.getResult(), offsetValue); + newOffsets.push_back(addOp.getResult()); + } + + ptrState.offsets.clear(); + + for (auto offset : newOffsets) { + ptrState.offsets.push_back(offset); + } + + SmallVector scalarShape(1, 1); + ArrayRef resultShape; + auto pointerType = cast(op.getResult().getType()); + if (auto shapedType = dyn_cast(pointerType.getPointeeType())) { + resultShape = shapedType.getShape(); + } else { + // scalar pointer, should produce a one dimensional memref + resultShape = scalarShape; + assert(ptrState.getRank() == 1); + } + + auto newOp = ptrState.createCastOp(resultShape, loc, rewriter); + + rewriter.replaceOp(op, newOp.getResult()); + + knownPtrs[newOp.getResult()] = ptrState; +} + +void PtrAnalysis::rewriteYieldOp( + scf::YieldOp op, ConversionPatternRewriter &rewriter, + const IndexMapSet &levelToBlockArgIndex, const int level, + const llvm::SmallDenseMap &knownPtrs) { + // any inserted instruction should be before this yield + OpBuilder::InsertionGuard insertionGuard{rewriter}; + rewriter.setInsertionPoint(op); + + auto adaptor = scf::YieldOp::Adaptor(op); + + SmallVector initArgState; + SmallVector operands(adaptor.getOperands()); + // Track the second chunks of modulo pointers so that we can append them to + // the yield results + SmallVector moduloSecondChunks; + + // For each of the init arg that we added additional Values in for loop, we + // need to add corresponding Values as yield operands. The loop below gathers + // PtrState for those values. + for (auto [i, v] : llvm::enumerate(adaptor.getOperands())) { + if (auto mappedV = rewriter.getRemappedValue(v)) { + // If this value is a tensor of pointers produced by AddPtrOp, + // we should have already converted to a ReinterpretCastOp without + // layout information for the normal cases, or to an + // UnrealizedConversionCastOp for the split pointer case. + if (v.getDefiningOp() || + v.getDefiningOp() || + v.getDefiningOp()) { + if (auto castOp = mappedV.getDefiningOp()) { + assertValidUnrealizedCast(castOp); + auto castInputs = castOp.getInputs(); + v = castOp.getResult(0); + operands[i] = castInputs[0]; + moduloSecondChunks.push_back(castInputs[1]); + } else if (auto castOp = + mappedV.getDefiningOp()) { + v = castOp; + } else { + llvm_unreachable("mapped value defined by an unexpected op"); + } + } else { + // If this value is not a tensor of pointers, we will use the + // mapped value, and rely on the conversion will happen later + // automatically when we legalize loop body. + + // TODO: + // The scenario where a value is a tensor of pointers but not + // produced by AddPtrOp is not supported + if (isa(mappedV.getType()) && + isa( + dyn_cast(mappedV.getType()).getElementType())) + llvm_unreachable("unsupported scenario where a value is a tensor of " + "pointers but not produced by AddPtrOp"); + v = mappedV; + } + } + + if (levelToBlockArgIndex.find(level) == levelToBlockArgIndex.end()) + continue; + auto thisSet = levelToBlockArgIndex.find(level)->second; + if (thisSet.find(i) == thisSet.end()) + continue; + + auto reintCastOp = v.getDefiningOp(); + auto unrealizedCastOp = v.getDefiningOp(); + + assert( + reintCastOp || + (unrealizedCastOp && + unrealizedCastOp->hasAttr(ModuloState::WraparoundAttr)) || + (isa(v.getType()) && + isa(dyn_cast(v.getType()).getElementType()))); + + PtrState state; + if (reintCastOp) { + visitOperandReintCast(reintCastOp, state, op.getLoc(), rewriter, + knownPtrs); + } else if (unrealizedCastOp) { + assertValidUnrealizedCast(unrealizedCastOp); + visitOperandUnrealizedCast(unrealizedCastOp, state, op.getLoc(), rewriter, + knownPtrs); + } else { + visitOperand(v, state, op.getLoc(), rewriter, knownPtrs); + } + initArgState.push_back(state); + } + + // For each of the PtrState recorded in the last step, extract value + // that correspond to offset and stride for each dimension and append + // them to yield operands. + for (auto state : initArgState) { + for (auto s : state.offsets) { + // offsets can be IntAttr zeroes, since reinterpret_cast collapses + // them for the input memref, and the for loop may not update + // offsets other than offsets[0]. Create constants Values for those + // zeroes. + if (auto sIntAttr = getIntAttr(s)) { + assert(sIntAttr.value() == 0 && "attribute offsets should be zeroes"); + auto constOp = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(0)); + operands.push_back(constOp.getResult()); + } else { + operands.push_back(cast(s)); + } + } + + for (auto s : state.strides) { + assert(!getIntAttr(s) && "PtrState strides for yield within for " + "loop not expected to be " + "attribute."); + operands.push_back(cast(s)); + } + } + + for (auto chunk : moduloSecondChunks) { + operands.push_back(chunk); + } + + // Yield is a terminator op that must be at the end of the function + rewriter.setInsertionPointAfter(op); + auto newOp = rewriter.replaceOpWithNewOp(op, operands); + assert(op->getNumResults() == 0); + + LLVM_DEBUG({ + llvm::dbgs() << "new yield:"; + newOp.getOperation()->print(llvm::dbgs(), + OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); +} + +// From an unrealized_conversion_cast which takes in two reinterpret_casts +// representing two chunks, we need to get back the full pointer state. We +// cannot rebuild the original state from the two reinterpret_casts similarly to +// the normal case. To solve this, we attach the original addptr as the third +// operand to the unrealized_cast so that we can manually rebuild the state. +void PtrAnalysis::visitOperandUnrealizedCast( + UnrealizedConversionCastOp op, PtrState &state, const Location loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &knownPtrs) { + assertValidUnrealizedCast(op); + + auto origPtr = op.getInputs()[2]; + if (knownPtrs.contains(origPtr)) { + state = knownPtrs.at(origPtr); + } else { + visitOperandAddptr(origPtr.getDefiningOp(), state, loc, + rewriter, knownPtrs); + } +} + +struct ModuloChunkInitArg { + Value reinterpretCast = nullptr; + // where in the init args is the first chunk placed + size_t initArgIndex = -1; +}; + +void PtrAnalysis::rewriteForOp( + scf::ForOp op, ConversionPatternRewriter &rewriter, + IndexMapSet &levelToBlockArgIndex, const int level, + llvm::SmallDenseMap &knownPtrs) { + SmallVector newInitArgs; + + SmallVector, 5> initArgIndexState; + SmallVector, 5> knownPtrsTmp; + + // If we have a load op that uses a modulo pointer, we need to insert both of + // the memref chunks to the init args. We reuse the sizes from the original + // memrefs. This data structure keeps track of where these additional init + // args should be inserted. + // + // As an example, if we have a 2D memrefs being split, we first put the first + // chunk in the order as it appears. Then, once all of the original init args + // are processed, we insert their offsets and strides, and finally the second + // chunk. + SmallVector, PtrState>, + 6> + moduloStates; + + // Amongst the init args, track the indices that map to the first chunk of a + // modulo pair. This is used to distinguish between the normal + // reinterpret_casts whose return types need to be rewritten to match what the + // for loop is yielding. + DenseSet moduloInitArgIndices; + + // Create a new list of init args + for (auto [i, arg] : llvm::enumerate(op.getInitArgs())) { + auto mappedV = rewriter.getRemappedValue(arg); + memref::ReinterpretCastOp reintCastOp; + UnrealizedConversionCastOp unrealizedCastOp; + + // If this init arg is supposed to be remapped, use the remapped + // value instead. In addition, if this init arg is a memref created + // by a reinterpret_cast or a tensor of index, there is a chance that + // it will be used in addptr. Create PtrState for each such init arg. + if (mappedV) { + // TODO: + // Passing a block argument pointer directly into a for loop not + // supported. + assert(!(dyn_cast(mappedV) && + isa(mappedV.getType())) && + "cannot take pointer block argument as init arg for for loop"); + if (auto op = mappedV.getDefiningOp()) { + reintCastOp = op; + newInitArgs.push_back(mappedV); + } else if (auto op = + mappedV.getDefiningOp()) { + assertValidUnrealizedCast(op); + unrealizedCastOp = op; + auto inputs = unrealizedCastOp.getInputs(); + + SmallVector initArgData{ + ModuloChunkInitArg{inputs[0], i}, + ModuloChunkInitArg{inputs[1]}, + }; + + moduloInitArgIndices.insert(i); + moduloStates.push_back( + std::make_tuple(unrealizedCastOp, initArgData, PtrState{})); + + newInitArgs.push_back(inputs[0]); + } else { + newInitArgs.push_back(mappedV); + } + + } else { + newInitArgs.push_back(arg); + } + + auto indexTensor = + isa(arg.getType()) && + isa(dyn_cast(arg.getType()).getElementType()); + + if (!unrealizedCastOp && !reintCastOp && !indexTensor) + continue; + + PtrState state; + if (reintCastOp) { + visitOperandReintCast(reintCastOp, state, op.getLoc(), rewriter, + llvm::SmallDenseMap(0)); + } else if (unrealizedCastOp) { + visitOperandUnrealizedCast(unrealizedCastOp, state, op.getLoc(), rewriter, + llvm::SmallDenseMap(0)); + std::get<2>(moduloStates.back()) = state; + } else { + visitOperand(arg, state, op.getLoc(), rewriter, + llvm::SmallDenseMap(0)); + } + + // Record the PtrState for later processing + initArgIndexState.push_back(std::make_pair(i, state)); + } + + // Set insertion point to be before the for loop for new variables passed + // into the new loop. + auto origIp = rewriter.saveInsertionPoint(); + rewriter.setInsertionPoint(op); + + // For each of the PtrState recorded in the last step, insert new + // instructions to describe offset and stride for each dimension and append + // them to init args + for (auto [i, state] : initArgIndexState) { + // For each dimension, if the corresponding offset and stride is an + // integer attribute, create a constant value and append them at the + // end of init arg list. + for (auto [j, s] : llvm::enumerate(state.offsets)) { + auto sIntAttr = getIntAttr(s); + if (sIntAttr) { + auto constOp = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(sIntAttr.value())); + newInitArgs.push_back(constOp.getResult()); + state.offsets[j] = constOp.getResult(); + } else { + newInitArgs.push_back(cast(s)); + } + } + + for (auto [j, s] : llvm::enumerate(state.strides)) { + auto sIntAttr = getIntAttr(s); + if (sIntAttr) { + auto constOp = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(sIntAttr.value())); + newInitArgs.push_back(constOp.getResult()); + state.strides[j] = constOp.getResult(); + } else { + newInitArgs.push_back(cast(s)); + } + } + + // Note that we want the knownPtrs to be indexed by block arg, but we + // only have index for now. Also, the state we record is the init + // arg, but want to to use newly created block arg. These block args + // are not created yet. We will translate this mapping later. + knownPtrsTmp.push_back(std::make_pair(i, state)); + levelToBlockArgIndex[level].insert(i); + + // If the original init arg is a memref produced by reinterpret_cast, + // create a new memref using new strides and offsets created above. + // This produces a canonicalized memref, which will match what the + // for loop generates if it modifies the memref. E.g., original + // reinterpret_cast can produce a memref with const stride: + // - memref<4x256xbf16, affine_map<(d0, d1)[s0, s1] -> (d0 * 256 + + // s0 + d1 + // * s1)>> + // The new reinterpret_cast will always have dynamic stride and + // offset: + // - memref<4x256xbf16, affine_map<(d0, d1)[s0, s1, s2] -> (d0 * s1 + // + s0 + d1 * s2)>> + // + // For init args that are the first chunk of a modulo pair, there is + // no need for the type to be rewritten because the strides and + // offsets are already dynamic. + if (!moduloInitArgIndices.contains(i) && + newInitArgs[i].getDefiningOp()) { + SmallVector resultShape; + for (auto s : state.sizes) { + auto sIntAttr = getIntAttr(s); + assert(sIntAttr && "expected constant size"); + resultShape.push_back(sIntAttr.value()); + } + auto castOp = state.createCastOp(resultShape, op.getLoc(), rewriter); + + LLVM_DEBUG({ + llvm::dbgs() << "new reinterpret_cast with dynamic sizes " + "and offsets:"; + castOp->print(llvm::dbgs(), OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); + + newInitArgs[i] = castOp.getResult(); + } + } + + // Pass in the second chunk of each modulo pair + for (auto &[unrealizedCastOp, chunkData, state] : moduloStates) { + chunkData[1].initArgIndex = newInitArgs.size(); + newInitArgs.push_back(chunkData[1].reinterpretCast); + } + + rewriter.restoreInsertionPoint(origIp); + + // Create a new scf::ForOp that uses updated init args and same loop body + auto newOp = rewriter.create( + op.getLoc(), op.getLowerBound(), op.getUpperBound(), op.getStep(), + newInitArgs, [&](OpBuilder &b, Location loc, Value iv, ValueRange args) { + IRMapping mapping; + mapping.map(op.getInductionVar(), iv); + mapping.map(op.getInitArgs(), newInitArgs); + mapping.map(op.getRegionIterArgs(), args); + + for (auto &bodyOp : op.getRegion().getOps()) { + b.clone(bodyOp, mapping); + } + + // Load op is lowered independent of the pointer, if we have a split + // pointer due to modulo, we need to "logically combine" these two + // memrefs into a single one using unrealized_cast_op. This way, when + // lowering the load, the pattern can detect if additional copies are + // inserted. When we are in a loop, it is more complicated because we + // have to insert a new unrealized_cast_op that combines the two memrefs + // in the init arg list. In addition, because init args hold no offset + // and size information, we have to manually insert two additional + // reinterpret_cast ops as input to this unrealized_cast_op so that the + // load have enough information to generate the corresponding copy. + OpBuilder::InsertionGuard g(b); + b.setInsertionPointToStart(b.getBlock()); + + Value zero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + + for (auto &[unrealizedCastOp, chunkData, state] : moduloStates) { + SmallVector newReinterpretCasts; + for (auto &chunk : chunkData) { + newReinterpretCasts.push_back(args[chunk.initArgIndex]); + } + + auto combinedCast = b.create( + loc, unrealizedCastOp.getResult(0).getType(), newReinterpretCasts, + unrealizedCastOp->getAttrs()); + + args[chunkData[0].initArgIndex].replaceUsesWithIf( + combinedCast.getResult(0), [](OpOperand &operand) { + assert(!isa(operand.getOwner()) && + "Storing to split pointers not supported"); + return isa(operand.getOwner()); + }); + } + }); + + // Convert the book-keeping data structure to use the correct key and value. + // Key is converted from init arg index to newly created block arg, and + // Value's PtrState fields are converted from init arg to newly created block + // arg + int cnt = op.getRegionIterArgs().size(); + for (auto [i, state] : knownPtrsTmp) { + for (auto it = state.offsets.begin(); it != state.offsets.end(); it++) { + *it = newOp.getRegionIterArgs()[cnt]; + cnt++; + } + + for (auto it = state.strides.begin(); it != state.strides.end(); it++) { + *it = newOp.getRegionIterArgs()[cnt]; + cnt++; + } + + auto key = newOp.getRegionIterArgs()[i]; + knownPtrs.insert(std::make_pair(key, state)); + } + assert(static_cast(cnt + moduloStates.size()) == + newOp.getRegionIterArgs().size() && + "expect to remap all new block args"); + + // Replace only the results that correspond to the original scf.for + auto resultsToReplaceWith = ResultRange( + newOp.result_begin(), newOp.result_begin() + op.getNumResults()); + rewriter.replaceOp(op, resultsToReplaceWith); + + // Update the loop body. Manually invoke the rewrite logic on addptr and yield + // in the loop body, so we can take advantage of the states we built up + for (auto &bodyOp : newOp.getRegion().getOps()) { + if (auto addptrOp = dyn_cast(bodyOp)) { + rewriteAddptrOp(addptrOp, rewriter, knownPtrs); + } else if (auto advanceOp = dyn_cast(bodyOp)) { + rewriteAdvanceOp(advanceOp, rewriter, knownPtrs); + } else if (auto forOp = dyn_cast(bodyOp)) { + // TODO: + // Nested for loops are not supported at the moment + assert(0 && "nested loops currently not supported"); + // rewriteForOp(forOp, rewriter, levelToBlockArgIndex, level+1, + // knownPtrs); levelToBlockArgIndex.erase(level+1); + } + } + + if (op.getNumRegionIterArgs()) { + auto yieldOp = cast(newOp.getBody()->getTerminator()); + rewriteYieldOp(yieldOp, rewriter, levelToBlockArgIndex, level, knownPtrs); + } + + LLVM_DEBUG({ + llvm::dbgs() << "new for\n"; + newOp.getOperation()->print(llvm::dbgs(), + OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); +} + +Value PtrAnalysis::getScalarMemRef(Value ptr, Value memRef, const Location loc, + ConversionPatternRewriter &rewriter) { + assert(cast(ptr.getType()) && "expected scalar pointer"); + + // If the pointer is generated by tt.addptr, we will have already inserted an + // ReinterpretCastOp to cast its type from tt.ptr to unranked memref. Return + // the result. + if (ptr.getDefiningOp()) { + if (auto castOp = memRef.getDefiningOp()) { + return castOp.getResult(); + } else { + llvm_unreachable("pointer value is defined by an unexpected op"); + } + } + + assert(isa(ptr) && + "pointer is neither produced by addptr nor a block argument"); + PtrState state; + state.source = memRef; + state.offsets.push_back(rewriter.getIndexAttr(0)); + state.sizes.push_back(rewriter.getIndexAttr(1)); + state.strides.push_back(rewriter.getIndexAttr(1)); + state.modulos.push_back(std::nullopt); + auto castOp = state.createCastOp(SmallVector(1, 1), loc, rewriter); + return castOp.getResult(); +} + +} // namespace triton +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/Analysis/UseAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Analysis/UseAnalysis.cpp new file mode 100755 index 00000000..62e45080 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Analysis/UseAnalysis.cpp @@ -0,0 +1,220 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/Analysis/UseAnalysis.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Analysis/DataFlow/ConstantPropagationAnalysis.h" +#include "mlir/Analysis/DataFlow/DeadCodeAnalysis.h" + +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" + +using namespace mlir; +using namespace triton; +using namespace dataflow; + +#define DEBUG_TYPE "triton-use-analysis" + +//===----------------------------------------------------------------------===// +// Use Analysis +// Note that logic below should evolve with triton-to-affine pass +//===----------------------------------------------------------------------===// +LogicalResult +triton::UseAnalysis::visitOperation(Operation *op, ArrayRef operands, + ArrayRef results) { + // If an op only produces pointer, all its operands are used as meta data. + // This accounts for scenarios such as addptr in a loop whose result is + // yielded. In this case, if the loop returns data tensors, addptr will be + // marked correctly as meta use. + if (op->getResults().size() == 1) { + auto resultType = dyn_cast(op->getResult(0).getType()); + if (resultType && isa(resultType.getElementType())) { + for (auto opnd : operands) + propagateUse(opnd, UseType::MetaUse); + } + } + + TypeSwitch(op) + .Case([&](auto load) { + propagateUse(operands[0], UseType::MetaUse); + auto mask = load.getMask(); + auto other = load.getOther(); + if (mask) { + assert(mask != other && "mask and other cannot be the same"); + propagateUse(operands[1], UseType::MetaUse); + } + if (other) { + // TODO: + // More complicated patterns that generate other is unsupported. + propagateUse(operands[2], UseType::MetaUse); + } + }) + .Case([&](auto store) { + propagateUse(operands[0], UseType::MetaUse); + propagateUse(operands[1], UseType::DataUse); + auto value = store.getValue(); + auto mask = store.getMask(); + if (mask) { + assert(mask != value && "mask and data cannot be the same"); + propagateUse(operands[2], UseType::MetaUse); + } + }) + .Case([&](auto dot) { + propagateResults(operands[0], results); + propagateResults(operands[1], results); + + auto opc = dot.getC(); + triton::SplatOp splat; + if (opc) + splat = opc.template getDefiningOp(); + + if (opc && splat && splat.getSrc().getDefiningOp()) + propagateUse(operands[2], UseType::MetaUse); + else + propagateUse(operands[2], UseType::DataUse); + }) + .Default([&](Operation *op) { + // this condition account for tt.addptr + for (auto operand : operands) { + propagateResults(operand, results); + } + }); + return success(); +} + +LogicalResult triton::runUseAnalysis(triton::FuncOp &funcOp) { + MLIRContext *context = funcOp.getContext(); + SymbolTableCollection symbolTable; + + DataFlowSolver solver; + solver.load(); + solver.load(); + solver.load(symbolTable); + if (failed(solver.initializeAndRun(funcOp))) + return failure(); + + // Walk the func op, convert tags on operands to tags on operations + funcOp.walk([&](Operation *op) { + UseType useType = UseType::Undefined; + for (auto result : op->getResults()) { + auto use = solver.lookupState(result); + assert(use && "Lattice value not found"); + auto thisUseType = use->type; + if (thisUseType == UseType::Undefined) + continue; + if (useType == UseType::Undefined) + useType = thisUseType; + if (thisUseType == UseType::MixUse || thisUseType != useType) { + useType = UseType::MixUse; + break; + } + } + + if (useType == UseType::Undefined) { + LLVM_DEBUG({ op->setAttr("Undefined", UnitAttr::get(context)); }); + return; + } else if (useType == UseType::MetaUse) { + assert(op->getNumResults() == 1 && + "Ops used for meta computation are expected to have one result"); + // Only set the tag if the operation uses tensors + if (isa(op->getResult(0).getType())) { + // Setting tag for erasing op later + op->setAttr("MetaUse", UnitAttr::get(context)); + } + return; + } else if (useType == UseType::DataUse) { + LLVM_DEBUG({ op->setAttr("DataUse", UnitAttr::get(context)); }); + return; + } + + assert(useType == UseType::MixUse); + + // If the operation only produces scalars, no need to clone it + bool shapedResult = true; + for (auto result : op->getResults()) + shapedResult &= isa(result.getType()); + if (!shapedResult) { + LLVM_DEBUG({ op->setAttr("MixUse", UnitAttr::get(context)); }); + return; + } + + // Value has MixUse. However, the operation may or may not have direct + // MetaUse. E.g., it may only have MixUse, or only have MixUse and + // DataUse. + // - If the operation has direct MetaUse, clone it, tag the clone as + // MetaUse only and point meta users to use the clone. + // - If not, do nothing; this operation will still be materlized. + llvm::SetVector metaUsers; + for (auto result : op->getResults()) { + for (auto user : result.getUsers()) { + TypeSwitch(user) + .Case([&](auto load) { + auto ptr = load.getPtr(); + auto mask = load.getMask(); + auto other = load.getOther(); + if (result == ptr || result == mask || result == other) + metaUsers.insert(user); + }) + .Case([&](auto store) { + auto ptr = store.getPtr(); + auto mask = store.getMask(); + if (result == ptr || result == mask) + metaUsers.insert(user); + }) + .Case([&](auto dot) { + auto opc = dot.getC(); + triton::SplatOp splat; + if (opc) + splat = opc.template getDefiningOp(); + + if (opc && splat && + splat.getSrc().getDefiningOp()) + metaUsers.insert(user); + }) + .Default([&](Operation *op) { + // if all output of user are used as meta data, user is a meta + // user. This condition account for addptr, or an addi whose + // output only feeds into addptr + bool allMeta = true; + for (auto res : op->getResults()) { + auto resUse = solver.lookupState(res); + if (resUse->type != UseType::MetaUse) { + allMeta = false; + break; + } + } + if (allMeta) + metaUsers.insert(user); + }); + } + } + + // If the operation doesn't have direct meta users, no need to clone it + if (metaUsers.empty()) { + LLVM_DEBUG({ op->setAttr("MixUse", UnitAttr::get(context)); }); + return; + } + + // Clone the operation; switch all meta users to use the clone + OpBuilder builder(op); + auto clone = builder.clone(*op); + LLVM_DEBUG({ op->setAttr("MixUse", UnitAttr::get(context)); }); + + // Setting tag for erasing op later + clone->setAttr("MetaUse", UnitAttr::get(context)); + + for (auto [res_i, result] : llvm::enumerate(op->getResults())) + for (auto user : metaUsers) + for (auto &operand : user->getOpOperands()) + if (operand.get() == result) + operand.set(clone->getResult(res_i)); + }); + + return success(); +} diff --git a/third_party/wafer/third_party/flir/lib/AnalysisStructured/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/AnalysisStructured/CMakeLists.txt new file mode 100755 index 00000000..5bec0d19 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/AnalysisStructured/CMakeLists.txt @@ -0,0 +1,26 @@ +if(FLAGTREE_BACKEND STREQUAL "wafer") +add_triton_library(TritonSharedAnalysisStructured + PtrAnalysisTS.cpp + + DEPENDS + TritonTableGen + TritonStructuredTableGen + TritonGPUAttrDefsIncGen + + LINK_LIBS PUBLIC + TritonStructuredIR + MLIRAnalysis +) +else() +add_triton_library(TritonSharedAnalysisStructured + PtrAnalysis.cpp + + DEPENDS + TritonTableGen + TritonStructuredTableGen + + LINK_LIBS PUBLIC + TritonStructuredIR + MLIRAnalysis +) +endif() \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/lib/AnalysisStructured/PtrAnalysis.cpp b/third_party/wafer/third_party/flir/lib/AnalysisStructured/PtrAnalysis.cpp new file mode 100755 index 00000000..1955b466 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/AnalysisStructured/PtrAnalysis.cpp @@ -0,0 +1,1396 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/AnalysisStructured/PtrAnalysis.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/Visitors.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Support/LogicalResult.h" +#include "triton-shared/Analysis/MaskAnalysis.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" + +#include "mlir/IR/IRMapping.h" +#include "mlir/Transforms/DialectConversion.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/LogicalResult.h" +#include +#include +#include +#include +#include + +#define DEBUG_TYPE "triton-ptr-analysis" + +namespace mlir { + +namespace tts { + +int32_t PtrState::getRank() const { + assert(offsets.size() == sizes.size() && offsets.size() == strides.size() && + shape.size() == offsets.size()); + return offsets.size(); +} + +bool PtrState::isEmpty() const { + return (getRank() == 0 && !source && !scalar); +} + +bool PtrState::hasModulo() const { + for (int32_t i = 0; i < getRank(); i++) { + if (dimHasModulo(i)) { + return true; + } + } + return false; +} + +bool PtrState::dimHasModulo(uint32_t dim) const { + assert( + !isBlockPtr() && + "Analysis should not check modulo if PtrState describes block pointer"); + + assert(dim < getRank()); + + auto intAttr = getIntAttr(shape[dim]); + if (!intAttr.has_value()) { + return true; + } + + return intAttr.value() != 0; +} + +bool PtrState::isBlockPtr() const { return !order.empty(); } + +LogicalResult PtrState::addState(const PtrState &lhsState, + const PtrState &rhsState, Operation *op, + OpBuilder &builder) { + assert(isEmpty() && lhsState.getRank() == rhsState.getRank()); + auto loc = op->getLoc(); + + if (lhsState.source && rhsState.source) { + op->emitRemark( + "PtrAnalysis: do not support adding two pointer states that both " + "have base pointers"); + return failure(); + } + + source = lhsState.source ? lhsState.source : rhsState.source; + + if (lhsState.scalar && rhsState.scalar) { + auto addOp = + builder.create(loc, lhsState.scalar, rhsState.scalar); + scalar = addOp.getResult(); + } else if (lhsState.getRank() == 0) { // both lhs and rhs are scalars + scalar = lhsState.scalar ? lhsState.scalar : rhsState.scalar; + } + + for (uint64_t i = 0; i < lhsState.getRank(); i++) { + auto newOffset = + addOFRs(lhsState.offsets[i], rhsState.offsets[i], loc, builder); + offsets.push_back(newOffset); + + auto newStride = + addOFRs(lhsState.strides[i], rhsState.strides[i], loc, builder); + strides.push_back(newStride); + + sizes.push_back(lhsState.sizes[i]); + } + + // AddPtr where both lhs and rhs containing modulo operators not supported + if (lhsState.hasModulo() && rhsState.hasModulo()) { + op->emitRemark("PtrAnalysis: do not support adding two pointer states " + "that both have modulo"); + return failure(); + } + + if (lhsState.hasModulo() || rhsState.hasModulo()) { + // visitOperandSplat and visitOperandExpandDims should enforce below + assert(lhsState.getRank() <= 2); + } + + // dealing with modulo: + // - If lhs has no modulo, skip + // - If rhs has zero offset on dim i, we can just use lhs's modulo + // - If i == 0 and rhs is the result of a splat, we will allow the add. This + // is because the user may be trying to express adding a constant offset to + // increment dim1, but pointer analysis cannot differentiate dim1 vs dim0 in + // this case. + // - Else, the analysis fails + + // An example for the 3rd condition above can look like: + // %0 = tt.splat %scalar + // %1 = tt.splat %ptr + // %2 = tt.arange + // %3 = arith.remsi %2, %size + // %4 = tt.addptr %1, %3 + // %5 = tt.addptr %4, %0 + // %5 may also occur in a loop to increment %4 every iteration. + + // Note that this is not bullet-proof. E.g., broken IR can actually increment + // dim0 while dim0 already has modulo, since Triton offsets are element-wise + // and not in unit of lower dimensions. However, this is highly unlikely but + // the analysis will provide wrong result. Hence we provide a warning in this + // case. + PtrState const *lhs = &lhsState; + PtrState const *rhs = &rhsState; + + if (rhs->hasModulo()) { + std::swap(lhs, rhs); + } + + for (uint64_t i = 0; i < lhs->getRank(); i++) { + if (!lhs->dimHasModulo(i)) { + shape.push_back(lhs->shape[i]); + } else if (hasConstZero(rhs->offsets[i])) { + shape.push_back(lhs->shape[i]); + } else if (i == 0 && lhs->getRank() == 2 && rhs->scalar) { + shape.push_back(lhs->shape[1]); + shape.push_back(lhs->shape[0]); + op->emitWarning( + "PtrAnalysis: allowing adding pointer state with modulo in dim 0 to " + "another pointer state with offset in dim 0.\nPlease verify the " + "operand that contains a scalar is meant to increment pointers in " + "dim1. If that is not the case it WILL LEAD TO WRONG COMPILATION " + "RESULTS.\n\nTo avoid this warning, use expand_dims (instead of " + "splat) to explicitly specify which dimension contains the scalar."); + break; + } else { + op->emitRemark( + "PtrAnalysis: do not support adding to operand with modulo"); + return failure(); + } + } + + return success(); +} + +void PtrState::dump() const { + llvm::dbgs() << "PtrState: "; + if (source) { + llvm::dbgs() << "source: " << source << "\n"; + } + if (scalar) { + llvm::dbgs() << "scalar: " << scalar << "\n"; + } + + llvm::dbgs() << "offsets: "; + llvm::interleave(offsets, llvm::dbgs(), "\n"); + llvm::dbgs() << "\nstrides: "; + llvm::interleave(strides, llvm::dbgs(), "\n"); + llvm::dbgs() << "\nsizes: "; + llvm::interleave(sizes, llvm::dbgs(), "\n"); + llvm::dbgs() << "\nshape: "; + llvm::interleave(shape, llvm::dbgs(), "\n"); + llvm::dbgs() << "\norder: "; + llvm::interleave(order, llvm::dbgs(), "\n"); + llvm::dbgs() << "\n"; +} + +LogicalResult PtrState::mulState(const PtrState &lhsState, + const PtrState &rhsState, Operation *op, + OpBuilder &builder) { + assert(isEmpty() && lhsState.getRank() == rhsState.getRank()); + + auto loc = op->getLoc(); + + // neither lhs nor rhs should have source, since multiplying base pointer + // does not make sense + if (lhsState.source && rhsState.source) { + op->emitRemark("PtrAnalysis: do not support multiplying base pointers"); + return failure(); + } + + // currently do not support both tensors are effectively non-scalar + if (!lhsState.scalar && !rhsState.scalar) { + op->emitRemark( + "PtrAnalysis: only support multiplying pointer states when one of " + "them represent a scalar"); + return failure(); + } + + PtrState const *lhs = &lhsState; + PtrState const *rhs = &rhsState; + + if (!rhs->scalar && lhs->scalar) { + std::swap(lhs, rhs); + } + + if (lhsState.scalar && rhsState.scalar) { + scalar = builder.create( + loc, lhsState.scalar, rhsState.scalar); + } + + for (uint64_t i = 0; i < lhs->sizes.size(); i++) { + OpFoldResult newOffset = + mulOFRValue(lhs->offsets[i], rhs->scalar, loc, builder); + OpFoldResult newStride = + mulOFRValue(lhs->strides[i], rhs->scalar, loc, builder); + OpFoldResult newShape = + mulOFRValue(lhs->shape[i], rhs->scalar, loc, builder); + offsets.push_back(newOffset); + strides.push_back(newStride); + shape.push_back(newShape); + sizes.push_back(lhs->sizes[i]); + } + + if (rhs->hasModulo()) { + op->emitRemark( + "PtrAnalysis: do not support multiplying pointer states that has " + "modulos"); + return failure(); + } + + return success(); +} + +tts::MakeTensorPtrOp PtrState::createTTSMakeTensorPtrOp(OpBuilder &builder, + Location loc) { + SmallVector staticSizes; + for (size_t i = 0; i < getRank(); i++) { + auto s = getIntAttr(sizes[i]); + assert(s.has_value()); + staticSizes.push_back(s.value()); + } + + auto op = builder.create( + loc, source, staticSizes, strides, offsets, shape, order); + LLVM_DEBUG({ + llvm::dbgs() << "creating tts::make_tensor_ptr:\n"; + op->dump(); + }); + + return op; +} + +LogicalResult PtrAnalysis::visitOperandAdd(arith::AddIOp addOp, PtrState &state, + const Location loc, + OpBuilder &builder) { + PtrState lhsState; + if (visitOperand(addOp.getLhs(), lhsState, loc, builder).failed()) { + return failure(); + } + + PtrState rhsState; + if (visitOperand(addOp.getRhs(), rhsState, loc, builder).failed()) + return failure(); + + // Checking for higher dimension is done in addState below + if ((lhsState.getRank() == 1 && lhsState.hasModulo()) || + (rhsState.getRank() == 1 && rhsState.hasModulo())) { + addOp->emitRemark( + "PtrAnalysis: do not support this pattern: a + arange(0, K) % M"); + return failure(); + } + + return state.addState(lhsState, rhsState, addOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandMul(arith::MulIOp mulOp, PtrState &state, + const Location loc, + OpBuilder &builder) { + PtrState lhsState; + if (visitOperand(mulOp.getLhs(), lhsState, loc, builder).failed()) { + return failure(); + } + + PtrState rhsState; + if (visitOperand(mulOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + return state.mulState(lhsState, rhsState, mulOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandRem(arith::RemSIOp remOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + PtrState rhsState; + if (visitOperand(remOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + if (!rhsState.scalar) { + remOp->emitRemark("PtrAnalysis: only support cases when rhs of remainder " + "contains scalar"); + return failure(); + } + + if (visitOperand(remOp.getLhs(), state, loc, builder).failed()) { + return failure(); + } + + // If there are multiple modulo ops on an expression (e.g.: (a % b) % c), we + // would have already populated the modulo states after visiting the lhs. + // Assert that all the modulo states are empty. + if (state.hasModulo()) { + remOp->emitRemark( + "PtrAnalysis: do not support multiple modulo within an expression"); + return failure(); + } + + if (state.getRank() == 1) { + // Apply the modulo before expanding shape, the common pattern is + // offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + // a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * + // stride_ak) + state.shape.back() = rhsState.scalar; + } else if (state.getRank() == 2) { + // torch inductor expands the tensor shape before applying the modulo. + // + // We only support either: + // - (tl.arange(0, end)[:, None] % mod), or + // - (tl.arange(0, end)[None, :] % mod) + // + // In both cases, we apply the modulo to the non-singleton dimension. + auto shape = cast(remOp.getResult().getType()).getShape(); + if (shape[0] == 1) { + state.shape[1] = rhsState.scalar; + } else if (shape[1] == 1) { + state.shape[0] = rhsState.scalar; + } else { + remOp->emitRemark( + "PtrAnalysis: taking modulo on a 2D tensor with no singleton " + "dimension not supported"); + return failure(); + } + } else { + remOp->emitRemark("PtrAnalysis: unsupported modulo pattern"); + return failure(); + } + return success(); +} + +LogicalResult PtrAnalysis::visitOperandExtSI(arith::ExtSIOp extOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + return visitOperand(extOp.getIn(), state, loc, builder); +} + +LogicalResult PtrAnalysis::visitOperandMakeRange(triton::MakeRangeOp rangeOp, + PtrState &state, Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + auto shape = cast(rangeOp.getType()).getShape(); + + auto start = rangeOp.getStart(); + auto end = rangeOp.getEnd(); + auto stride = (end - start + shape[0] - 1) / shape[0]; + assert(stride == 1 && + "Expect make_range op to always return tensor of stride 1"); + + state.offsets.push_back(builder.getIndexAttr(start)); + state.sizes.push_back(builder.getIndexAttr(shape[0])); + state.strides.push_back(builder.getIndexAttr(stride)); + state.shape.push_back(builder.getIndexAttr(0)); + return success(); +} + +LogicalResult +PtrAnalysis::visitOperandExpandDims(triton::ExpandDimsOp expandDimsOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + if (visitOperand(expandDimsOp.getSrc(), state, loc, builder).failed()) { + return failure(); + } + + auto dstShape = + cast(expandDimsOp.getResult().getType()).getShape(); + auto axis = expandDimsOp.getAxis(); + + assert(dstShape[axis] == 1 && + "expect changed dimension to be 1 in expand_dims"); + + // insert dimension info + state.offsets.insert(state.offsets.begin() + axis, builder.getIndexAttr(0)); + state.sizes.insert(state.sizes.begin() + axis, builder.getIndexAttr(1)); + state.strides.insert(state.strides.begin() + axis, builder.getIndexAttr(0)); + state.shape.insert(state.shape.begin() + axis, builder.getIndexAttr(0)); + + if (state.hasModulo() && state.getRank() > 2) { + expandDimsOp->emitRemark( + "PtrAnalysis: unsupported scenario where expand_dims result " + "has modulo and rank > 2"); + return failure(); + } + + return success(); +} + +LogicalResult +PtrAnalysis::visitOperandBroadcast(triton::BroadcastOp broadcastOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + auto src = broadcastOp.getSrc(); + auto dst = broadcastOp.getResult(); + + if (!isa(src.getType())) { + broadcastOp->emitRemark("PtrAnalysis: Unsupported broadcast source type"); + return failure(); + } + + auto srcShape = cast(src.getType()).getShape(); + auto dstShape = cast(dst.getType()).getShape(); + + assert(srcShape.size() == dstShape.size() && + "rank of source and destination should match"); + + if (visitOperand(src, state, loc, builder).failed()) { + return failure(); + } + + for (size_t i = 0; i < dstShape.size(); i++) { + if (srcShape[i] == dstShape[i]) { + continue; + } else if (srcShape[i] < dstShape[i]) { + state.sizes[i] = builder.getIndexAttr(dstShape[i]); + } else { + llvm_unreachable("unexpected dimensions used in broadcast"); + } + } + return success(); +} + +LogicalResult PtrAnalysis::visitOperandSplat(triton::SplatOp splatOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + auto src = splatOp.getSrc(); + auto dst = splatOp.getResult(); + auto dstShape = cast(dst.getType()).getShape(); + + if (visitOperand(src, state, loc, builder).failed()) { + return failure(); + } + + if (isa(src.getType())) { + for (auto s : dstShape) { + state.offsets.push_back(builder.getIndexAttr(0)); + state.sizes.push_back(builder.getIndexAttr(s)); + state.strides.push_back(builder.getIndexAttr(0)); + state.shape.push_back(builder.getIndexAttr(0)); + } + } else { + splatOp->emitRemark("PtrAnalysis: unsupported splat pattern"); + return failure(); + } + + // If we splat a integer value, scalar should become the offset of the outer + // most dimension + if (state.scalar) + state.offsets[0] = state.scalar; + + if (state.hasModulo() && state.getRank() > 2) { + splatOp->emitRemark("PtrAnalysis: unsupported scenario where splat result " + "has modulo and rank > 2"); + return failure(); + } + + return success(); +} + +LogicalResult PtrAnalysis::visitOperandAddptr(triton::AddPtrOp addptrOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + PtrState ptrState; + if (visitOperand(addptrOp.getPtr(), ptrState, addptrOp.getLoc(), builder) + .failed()) { + // assert(0); + return failure(); + } + + PtrState offsetState; + if (visitOperand(addptrOp.getOffset(), offsetState, addptrOp.getLoc(), + builder) + .failed()) { + return failure(); + } + + assert(ptrState.source && "ptr field should provide source / base pointer"); + + assert(ptrState.getRank() == offsetState.getRank() && + "ptr and offset field should have the same rank"); + + return state.addState(ptrState, offsetState, addptrOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandConstSplat(arith::ConstantOp op, + PtrState &state, + const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + // this condition is to handle cases where tt.broadcast and tt.splat are + // folded + auto attr = cast(op.getValue()); + auto elementType = attr.getElementType(); + assert(attr.isSplat() && isa(elementType)); + auto values = attr.getValues(); + auto value = values[0].getValue(); + auto constAttr = builder.getIndexAttr(value.getSExtValue()); + auto constOp = arith::ConstantOp::materialize(builder, constAttr, + builder.getIndexType(), loc); + + state.scalar = constOp; + + auto resultType = cast(op.getResult().getType()); + for (size_t i = 0; i < resultType.getShape().size(); i++) { + if (i == 0) { + state.offsets.push_back(constOp.getResult()); + } else { + state.offsets.push_back(builder.getIndexAttr(0)); + } + + state.sizes.push_back(builder.getIndexAttr(resultType.getShape()[i])); + state.strides.push_back(builder.getIndexAttr(0)); + state.shape.push_back(builder.getIndexAttr(0)); + } + + return success(); +} + +LogicalResult PtrAnalysis::visitOperandMakeTPtr(tts::MakeTensorPtrOp makeTPtrOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + + assert(state.isEmpty()); + state.source = makeTPtrOp.getBase(); + state.offsets = makeTPtrOp.getMixedOffsets(); + state.sizes = makeTPtrOp.getMixedSizes(); + state.strides = makeTPtrOp.getMixedStrides(); + state.shape = makeTPtrOp.getMixedShape(); + state.order = SmallVector(makeTPtrOp.getOrder()); + + return success(); +} + +LogicalResult +PtrAnalysis::visitOperandMakeTensorPtr(triton::MakeTensorPtrOp makeTPtrOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + state.source = makeTPtrOp.getBase(); + + if (makeTPtrOp.getOrder().empty()) { + makeTPtrOp->emitRemark( + "PtrAnalysis: expect tt.make_tensor_ptr to have order field set"); + return failure(); + } + + auto resType = cast(makeTPtrOp.getResult().getType()); + auto pointeeType = cast(resType.getPointeeType()); + auto shape = pointeeType.getShape(); + + for (int64_t i = 0; i < pointeeType.getRank(); i++) { + state.sizes.push_back(builder.getIndexAttr(shape[i])); + + auto strideCst = builder.create( + loc, builder.getIndexType(), makeTPtrOp.getStrides()[i]); + state.strides.push_back(strideCst.getResult()); + + auto offsetCst = builder.create( + loc, builder.getIndexType(), makeTPtrOp.getOffsets()[i]); + + auto scaledOffset = builder.create( + loc, offsetCst.getResult(), strideCst.getResult()); + state.offsets.push_back(scaledOffset.getResult()); + + auto shapeCst = builder.create( + loc, builder.getIndexType(), makeTPtrOp.getShape()[i]); + state.shape.push_back(shapeCst.getResult()); + } + state.order = SmallVector(makeTPtrOp.getOrder()); + assert(state.isBlockPtr() && + "tt.make_tensor_ptr pointer state should describe a block pointer"); + + return success(); +} + +LogicalResult PtrAnalysis::visitOperandForOp(scf::ForOp forOp, Value operand, + PtrState &state, + const Location loc, + OpBuilder &builder) { + + auto it = llvm::find(forOp->getResults(), operand); + auto index = std::distance(forOp->getResults().begin(), it); + + auto newState = getLoopResultPtrState(forOp, index); + if (failed(newState)) { + forOp.emitError( + "Rewrite for-op failed. Could not find PtrState returned by " + "the loop."); + return failure(); + } + + state = newState.value(); + return success(); +} + +LogicalResult PtrAnalysis::visitOperandIntToPtr(triton::IntToPtrOp op, + PtrState &state, + const Location loc, + OpBuilder &builder) { + state.source = op.getResult(); + return success(); +} + +LogicalResult PtrAnalysis::visitOperandBitcast(triton::BitcastOp op, + PtrState &state, + const Location loc, + OpBuilder &builder) { + auto resType = op.getResult().getType(); + if (isa(resType)) { + return visitOperand(op.getSrc(), state, loc, builder); + } + state.source = op.getResult(); + return success(); +} + +LogicalResult PtrAnalysis::visitOperand(Value operand, PtrState &state, + const Location loc, + OpBuilder &builder) { + + if (knownPtrs.find(operand) != knownPtrs.end()) { + state = knownPtrs.lookup(operand); + return success(); + } + + if (isa(operand.getType())) { + OpBuilder::InsertionGuard guard(builder); + if (!isa(operand) && operand.getDefiningOp()) { + builder.setInsertionPointAfter(operand.getDefiningOp()); + } + auto castOp = builder.create( + loc, builder.getIndexType(), operand); + state.scalar = castOp.getResult(); + return success(); + } else if (isa(operand.getType())) { + state.scalar = operand; + return success(); + } + + if (isa(operand.getType())) { + // A scalar pointer can either be produced by AddPtrOp or a block + // argument + if (auto op = operand.getDefiningOp()) { + if (auto addPtrOp = dyn_cast(op)) { + return visitOperandAddptr(cast(op), state, loc, + builder); + } else if (auto castOp = dyn_cast(op)) { + return visitOperandBitcast(castOp, state, loc, builder); + } else if (auto intToPtrOp = dyn_cast(op)) { + return visitOperandIntToPtr(intToPtrOp, state, loc, builder); + } else if (auto makeTensorOp = dyn_cast(op)) { + llvm_unreachable("Unexpected operand defining operation tts.make_tptr"); + } else { + op->emitRemark("Unexpected defining op for triton pointer operand"); + return failure(); + } + } else { + state.source = operand; + return success(); + } + } + + if (auto op = operand.getDefiningOp()) { + return visitOperandAdd(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandMul(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandMakeRange(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandBroadcast(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandSplat(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandExpandDims(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandAddptr(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandConstSplat(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandRem(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandExtSI(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandForOp(op, operand, state, loc, builder); + } else if (!operand.getDefiningOp()) { + if (!knownPtrs.contains(operand)) { + return failure(); + } + + // This operand must be an iter-arg of an inner-loop in a multiple-level + // nested loop, which means its PtrState must have already been populated + // during rewriteForOp of the parent loop. + state = knownPtrs[operand]; + return success(); + } else { + llvm::dbgs() << "PtrAnalysis: encountered addptr operand produced by an " + "unsupported operation\n"; + operand.dump(); + return failure(); + } +} + +LogicalResult PtrAnalysis::rewriteAddptrOp(triton::AddPtrOp op) { + OpBuilder builder(op); + + PtrState state; + if (visitOperandAddptr(op, state, op.getLoc(), builder).failed()) { + return failure(); + } + + knownPtrs[op.getResult()] = state; + + if (isa(op.getPtr().getType())) { + auto maketptrOp = state.createTTSMakeTensorPtrOp(builder, op.getLoc()); + ptrMap.map(op.getResult(), maketptrOp.getResult()); + } else { + // record the ptr as we have visited and built up the state for this scalar + // pointer, which may be used by rewriteForOp later. + ptrMap.map(op.getResult(), op.getResult()); + } + return success(); +} + +LogicalResult PtrAnalysis::rewriteBitcastOp(triton::BitcastOp op) { + // Only rewrite bitcast on tensor of pointers. + auto resultType = dyn_cast(op.getType()); + if (!resultType) { + return failure(); + } + auto ptrType = dyn_cast(resultType.getElementType()); + if (!ptrType) { + return failure(); + } + + OpBuilder builder(op); + + PtrState state; + if (visitOperandBitcast(op, state, op.getLoc(), builder).failed()) { + return failure(); + } + + if (!isa(state.source.getType())) { + return failure(); + } + Value castedPtr = + builder.create(op.getLoc(), ptrType, state.source) + .getResult(); + state.source = castedPtr; + + knownPtrs[op.getResult()] = state; + + auto maketptrOp = state.createTTSMakeTensorPtrOp(builder, op.getLoc()); + ptrMap.map(op.getResult(), maketptrOp.getResult()); + + return success(); +} + +LogicalResult PtrAnalysis::rewriteMakeTensorPtrOp(triton::MakeTensorPtrOp op) { + OpBuilder builder(op); + + PtrState state; + if (visitOperandMakeTensorPtr(op, state, op.getLoc(), builder).failed()) { + return failure(); + } + + auto maketptrOp = state.createTTSMakeTensorPtrOp(builder, op.getLoc()); + knownPtrs[op.getResult()] = state; + ptrMap.map(op.getResult(), maketptrOp.getResult()); + return success(); +} + +LogicalResult PtrAnalysis::rewriteAdvanceOp(triton::AdvanceOp op) { + OpBuilder builder(op); + auto loc = op.getLoc(); + + PtrState state; + if (visitOperand(op->getOperand(0), state, loc, builder).failed()) { + op->emitRemark("PtrAnalysis: Failed to analyze ptr of tt.advance"); + return failure(); + } + assert(state.isBlockPtr() && + "tt.advance pointer state should describe a block pointer"); + + auto incrementOffsets = op.getOffsets(); + + SmallVector newOffsets; + for (auto [increment, offset, stride] : + llvm::zip(incrementOffsets, state.offsets, state.strides)) { + Value offsetValue; + if (auto offsetIntAttr = getIntAttr(offset)) { + auto constOp = builder.create( + loc, builder.getIndexAttr(offsetIntAttr.value())); + offsetValue = constOp.getResult(); + } else { + offsetValue = cast(offset); + } + auto castOp = builder.create( + loc, builder.getIndexType(), increment); + auto mulOp = builder.create(loc, castOp.getResult(), + cast(stride)); + auto addOp = + builder.create(loc, mulOp.getResult(), offsetValue); + newOffsets.push_back(addOp.getResult()); + } + + state.offsets = SmallVector(newOffsets); + + auto newOp = state.createTTSMakeTensorPtrOp(builder, loc); + knownPtrs[op.getResult()] = state; + ptrMap.map(op.getResult(), newOp.getResult()); + return success(); +} + +static bool isPointerType(Type t) { + if (auto tensor = llvm::dyn_cast(t)) { + return isa(tensor.getElementType()); + } + return isa(t); +} + +FailureOr PtrAnalysis::getLoopInitArgPtrState(scf::ForOp forOp, + size_t index) { + auto ptr = forOp.getInitArgs()[index]; + + // If the pointer into the scf.for was defined by tts.get_structured_state, + // we can get the pointer state from the original pointer (the op's input): + // + // %ptr, %offset_1, %offset_2,..., %stride_1, %stride_2,... = + // tts.get_structured_state %original + // scf.for ... (%ptr) {...} + if (auto getStateOp = ptr.getDefiningOp()) { + auto originalPtr = getStateOp->getOperand(0); + if (knownPtrs.count(originalPtr)) { + return knownPtrs[originalPtr]; + } + } + + // For nested loops scenarios, a pointer in init-args can be returned from + // another loop of the same level: + // e.g.: + // clang-format off + // %22:2 = scf.for %arg4 = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%arg5 = %11, %arg6 = %15) -> (tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr>) : i32 { + // %23 = scf.for %arg7 = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%arg8 = %arg5) -> (tensor<2x2x!tt.ptr>) : i32 { + // %26 = tt.addptr %arg8, %17 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> + // scf.yield %26 : tensor<2x2x!tt.ptr> + // } + // %24:2 = scf.for %arg7 = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%arg8 = %23, %arg9 = %arg6) -> (tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr>) : i32 { + // %26 = tt.load %arg8 : tensor<2x2x!tt.ptr> + // %27 = tt.addptr %arg8, %19 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> + // ... + // } + // ... + // } + // clang-format on + // Notice %arg8 = %23 comes from the return value of the first loop. + if (auto forOp = ptr.getDefiningOp()) { + return getLoopResultPtrState(forOp, index); + } + + // If the pointer isn't defined by tts.get_structured_state nor another loop, + // it means the current pointer is an iterarg of the outer loop. + // In such cases, the outer loops would have already set up the PtrState for + // us already. + // + // scf.for iterargs(%ptr = %init_arg) { + // scf.for iterargs(%ptr1 = %ptr) { <--- we're dealing with `%ptr1` here. + // ... + // } + // } + if (knownPtrs.count(ptr)) { + assert(!ptr.getDefiningOp() && "Expect the ptr to be an iterarg"); + return knownPtrs[ptr]; + } + + return failure(); +} + +PtrState PtrAnalysis::reconcileLoopPtrState( + scf::ForOp forOp, size_t iterArgIndex, const PtrState &state, + llvm::function_ref getReplacementVal) { + PtrState newState = state; + int cnt = iterArgIndex + 1; + if (newState.getRank() == 0) { + assert(newState.scalar); + // for scalar pointers, the scalar contains the offset and is the only + // relevant newState that could be updated by the loop. + newState.scalar = getReplacementVal(forOp, cnt); + } else { + for (auto &offset : newState.offsets) { + offset = getReplacementVal(forOp, cnt++); + } + + for (auto &stride : newState.strides) { + stride = getReplacementVal(forOp, cnt++); + } + } + + return newState; +} + +FailureOr PtrAnalysis::getLoopIterArgPtrState(scf::ForOp forOp, + size_t index) { + auto state = getLoopInitArgPtrState(forOp, index); + if (failed(state)) { + return failure(); + } + + return reconcileLoopPtrState( + forOp, index, state.value(), + [](scf::ForOp op, size_t index) { return op.getRegionIterArg(index); }); +} + +FailureOr PtrAnalysis::getLoopResultPtrState(scf::ForOp forOp, + size_t index) { + auto state = getLoopInitArgPtrState(forOp, index); + if (failed(state)) { + return failure(); + } + + return reconcileLoopPtrState( + forOp, index, state.value(), + [](scf::ForOp op, size_t index) { return op->getResult(index); }); +} + +LogicalResult PtrAnalysis::rewriteForOp(scf::ForOp op) { + for (auto [i, arg] : llvm::enumerate(op.getRegionIterArgs())) { + if (!maybeStructuredArgs.contains(arg)) { + continue; + } + + auto state = getLoopIterArgPtrState(op, i); + if (failed(state)) { + // Because the maybeStructuredArgs may contain values that are not + // considered structured by PtrAnalysis, failing to retrieve the PtrState + // should not fail the rewrite process. + // We emit an error for diagnostics and debugging purposes. + op->emitWarning( + "Rewrite for-op failed. Could not find PtrState for iter-arg index " + + std::to_string(i)); + continue; + } + + // Save the current init arg's PtrState + knownPtrs[arg] = state.value(); + + // For tensors of pointers, create a tts.make_tptr at the beginning of the + // loop body that correspond to this region iter arg. In case it is used + // by tt.load/tt.store in the loop body before pointer updates, this will + // make sure rewriteLoadOp/rewriteStoreOp can use the analysis result. + // E.g., given the following input (%tensor_of_ptr is a block arg): + // scf.for (%tensor_of_ptr) { + // %data = tt.load %tensor_of_ptr + // // more operations to update %tensor_of_ptr + // } + // We may produce the following output: + // scf.for (%base_ptr, %stride, %offset) { + // %tensor_of_ptr = tts.make_tptr(%base_ptr, %stride, %offset) + // %data = tts.load %tensor_of_ptr + // // more operations to update %offset + // } + // If %tensor_of_ptr is not used (i.e., %tensor_of_ptr is updated before + // used in the original IR), it will simply be removed by + // canonicalization. + + // For scalar pointers, there is no need to create a tts.addptr at the + // beginning of the loop body. We don't lower tt.load and tt.store on + // scalars in this pass; pointer arithmetics can also just use the + // original pointer. + // Note that there can be tensor of indices in iter-arg, so we only create + // the make_tensor_ptr op when the arg is of pointer type. + if (isPointerType(arg.getType())) { + if (state->getRank() != 0) { + OpBuilder builder(op.getRegion()); + auto maketptrOp = state->createTTSMakeTensorPtrOp(builder, op.getLoc()); + ptrMap.map(arg, maketptrOp.getResult()); + } + } + } + + // Recursively rewrite the inner ops + if (rewriteOp(op).failed()) { + op->emitRemark( + "PtrAnalysis: update loop body failed when rewriting for op"); + return failure(); + } + + return success(); +} + +LogicalResult +PtrAnalysis::rewriteGetStructuredStateOp(tts::GetStructuredStateOp op) { + auto tritonValue = op->getOperand(0); + + // If this triton value isn't known, it means PtrAnalysis has failed to + // analyze this pointer. In such cases, simply remap all uses of the + // structured value back to its original triton value. + if (!knownPtrs.contains(tritonValue)) { + op.emitRemark( + "Rewrite GetStructuredStateOp failed. Could not find PtrState."); + op.getResult(0).replaceAllUsesWith(tritonValue); + return failure(); + } + + tts::PtrState state = knownPtrs[tritonValue]; + Value remappedValue = + ptrMap.contains(tritonValue) ? ptrMap.lookup(tritonValue) : tritonValue; + + SmallVector replacements{remappedValue}; + OpBuilder builder(op); + + if (state.getRank() == 0) { + // For scalar pointers, the scalar contains the offset and is the only + // relevant state that could be updated by the loop. + if (state.scalar) { + replacements.push_back(state.scalar); + } else { + // This operand is a pointer directly from the kernel arguments. + // Use offset 0. + assert(!tritonValue.getDefiningOp()); + replacements.push_back(builder.create( + op.getLoc(), builder.getIndexAttr(0))); + } + } else { + for (auto [j, s] : llvm::enumerate(state.offsets)) { + auto sIntAttr = getIntAttr(s); + if (sIntAttr) { + auto constOp = builder.create( + op.getLoc(), builder.getIndexAttr(sIntAttr.value())); + replacements.push_back(constOp.getResult()); + } else { + replacements.push_back(cast(s)); + } + } + + for (auto [j, s] : llvm::enumerate(state.strides)) { + auto sIntAttr = getIntAttr(s); + if (sIntAttr) { + auto constOp = builder.create( + op.getLoc(), builder.getIndexAttr(sIntAttr.value())); + replacements.push_back(constOp.getResult()); + } else { + replacements.push_back(cast(s)); + } + } + } + + op->replaceAllUsesWith(replacements); + op->erase(); + return success(); +} + +LogicalResult PtrAnalysis::rewriteLoadOp(triton::LoadOp op, + bool useUnsafeMask) { + auto ptr = ptrMap.lookupOrNull(op.getPtr()); + auto mask = op.getMask(); + auto other = op.getOther(); + auto loc = op.getLoc(); + + if (!ptr) { + op->emitRemark("PtrAnalysis: pointer is not replace with tts.make_tptr so " + "loadOp cannot be rewritten"); + return failure(); + } + + auto ptrType = dyn_cast(ptr.getType()); + if (ptrType && !isa(ptrType.getPointeeType())) { + op->emitRemark("PtrAnalysis: scalar loadOp will not be rewritten"); + return failure(); + } + + ArrayRef dims; + mlir::triton::MaskState mstate(useUnsafeMask); + Value scalarOther; + + OpBuilder builder(op); + // Analyze the mask operand to determine at runtime the size of the data we + // are moving. + if (mask) { + if (mstate.parse(mask, loc, builder).failed()) { + op->emitRemark("MaskAnalysis failed"); + return failure(); + } + dims = mstate.dims; + } + + if (other) { + assert(mask && "other value used while no masks are specified"); + + scalarOther = utils::getScalarValue(other, loc, builder); + if (!scalarOther) { + op->emitRemark("other value used in masked load produced by " + "unsupported instruction"); + return failure(); + } + } + + auto loadOp = builder.create(loc, ptr, dims, scalarOther); + auto strAttr = op->getAttrOfType("flagtree_hints"); + if (strAttr && !strAttr.getValue().empty()) { + loadOp->setAttr("flagtree_hints", strAttr); + } + + if (op->getAttr("flagtree_hints")) { + loadOp->setAttr("flagtree_hints", op->getAttr("flagtree_hints")); + } + + LLVM_DEBUG({ + llvm::dbgs() << "creating tts::load:\n"; + loadOp->dump(); + }); + + op.replaceAllUsesWith(loadOp.getResult()); + op->erase(); + return success(); +} + +// Structured values from the TritonStructuredDialect have offsets and strides +// that might change in each loop iteration and hence will appear in an scf.for +// iter-args like so: +// +// %structured, %offsets, %strides = tts.get_structured_state +// scf.for (%arg0 = %structured, %arg1 = %offsets, %arg2 = %strides) { +// %a = %arg0 + 1 +// %b = %b + 2 +// scf.for (%arg1 = %b) { +// ... +// } +// } +// +// In `rewriteForOp`, we have to recognize such structured values in order to +// rewrite their PtrState accordingly. Previously, only values of Pointer-like +// type (e.g.: tensor> or tt.ptr>), so detecting these values +// is as easy as checking the type. +// +// Now, tensor of indices could also appear in a loop's iter-arg. To reliably +// detect all such cases, we perform a BFS-like traversal of the IR where the +// sources are the results of `tts.get_structured_state`. All values that +// originate from the results of `tts.get_structured_state` are consider +// "maybeStructured". If a loop's iter-arg is considered "maybeStructured", we +// must set up their PtrState during `rewriteForOp`. +void PtrAnalysis::initializeMaybeStructuredArgs(Operation *op) { + std::queue q; + DenseSet visited; + + op->walk([&q, &visited](tts::GetStructuredStateOp getStateOp) { + Value value = getStateOp->getResult(0); + visited.insert(value); + q.push(value); + }); + + while (!q.empty()) { + auto v = q.front(); + q.pop(); + for (auto user : v.getUsers()) { + // scf.for is a special case. We have 2 set of values to consider: + // - iter-args + // - loop results + // for every init arg that originates from a `tts.get_structured_state` + // op, its corresponding iter-arg and loop result will also be considered + // "maybeStructured". + if (auto forOp = dyn_cast(user)) { + auto it = llvm::find(forOp.getInitArgs(), v); + + if (it == forOp.getInitArgs().end()) { + continue; + } + + auto argIndex = std::distance(forOp.getInitArgs().begin(), it); + auto iterArg = forOp.getRegionIterArg(argIndex); + auto tiedLoopRes = forOp.getTiedLoopResult(iterArg); + + SmallVector neighbors{iterArg, tiedLoopRes}; + for (auto neighbor : neighbors) { + maybeStructuredArgs.insert(neighbor); + if (!visited.contains(neighbor)) { + visited.insert(neighbor); + q.push(neighbor); + } + } + + } else { + for (auto res : user->getResults()) { + if (res.getType() != v.getType()) { + continue; + } + maybeStructuredArgs.insert(res); + if (!visited.contains(res)) { + visited.insert(res); + q.push(res); + } + } + } + } + } +} + +LogicalResult PtrAnalysis::rewriteStoreOp(triton::StoreOp op, + bool useUnsafeMask) { + auto ptr = ptrMap.lookupOrNull(op.getPtr()); + auto val = op.getValue(); + auto mask = op.getMask(); + auto loc = op.getLoc(); + + if (!ptr) { + op->emitRemark("PtrAnalysis: pointer is not replace with tts.make_tptr so " + "storeOp cannot be rewritten"); + return failure(); + } + + auto ptrType = dyn_cast(ptr.getType()); + if (ptrType && !isa(ptrType.getPointeeType())) { + op->emitRemark("PtrAnalysis: scalar storeOp will not be rewritten"); + return failure(); + } + + ArrayRef dims; + mlir::triton::MaskState mstate(useUnsafeMask); + + OpBuilder builder(op); + + // Analyze the mask operand to determine at runtime the size of the data + // are moving. + if (mask) { + if (mstate.parse(mask, loc, builder).failed()) { + op->emitRemark("MaskAnalysis failed"); + return failure(); + } + dims = mstate.dims; + } + + auto storeOp = builder.create(loc, ptr, val, dims); + + LLVM_DEBUG({ + llvm::dbgs() << "creating tts::store:\n"; + storeOp->dump(); + }); + + op->erase(); + return success(); +} + +LogicalResult PtrAnalysis::rewriteOp(Operation *rootOp, bool useUnsafeMask) { + LLVM_DEBUG({ + llvm::dbgs() << "rewriting rootOp\n"; + rootOp->dump(); + }); + + rootOp->walk([&](Operation *op) { + if (op == rootOp) { + return WalkResult::advance(); + } + return TypeSwitch(op) + .Case([&](auto addptr) { + if (rewriteAddptrOp(addptr).failed()) { + addptr->emitRemark("PtrAnalysis: Failed to rewrite AddPtrOp"); + } + return WalkResult::advance(); + }) + .Case([&](auto bitcast) { + if (rewriteBitcastOp(bitcast).failed()) { + bitcast->emitRemark("PtrAnalysis: Failed to rewrite BitcastOp"); + } + return WalkResult::advance(); + }) + .Case([&](auto maketptr) { + if (rewriteMakeTensorPtrOp(maketptr).failed()) { + maketptr->emitRemark( + "PtrAnalysis: Failed to rewrite MakeTensorPtrOp"); + } + return WalkResult::advance(); + }) + .Case([&](auto advance) { + if (rewriteAdvanceOp(advance).failed()) { + advance->emitRemark("PtrAnalysis: Failed to rewrite AdvanceOp"); + } + return WalkResult::advance(); + }) + .Case([&](auto load) { + if (rewriteLoadOp(load, useUnsafeMask).failed()) { + load->emitRemark("PtrAnalysis: Failed to rewrite LoadOp"); + return WalkResult::advance(); + } + return WalkResult::skip(); + }) + .Case([&](auto store) { + if (rewriteStoreOp(store, useUnsafeMask).failed()) { + store->emitRemark("PtrAnalysis: Failed to rewrite StoreOp"); + return WalkResult::advance(); + } + return WalkResult::skip(); + }) + .Case([&](auto forOp) { + // `rewriteForOp` recursively visits its children, so regardless + // whether the rewrite succeeds or not, we need to return "skip" so + // that the the walk does not visit the for-op's child operations + // the second time. + if (rewriteForOp(forOp).failed()) { + forOp->emitRemark("PtrAnalysis: Failed to rewrite ForOp"); + } + return WalkResult::skip(); + }) + .Case( + [&](tts::GetStructuredStateOp getStateOp) { + // For tensor of indices potentially being used in pointer + // arithmetic sequence, we need to manually populate the state of + // none already exists. + // This process is necessary because unlike triton pointers in a + // loop which always have a `tt.addptr` that triggers the rewrite + // process which includes generating the ops for updating offsets + // and strides, tensor of indices only have a simple `arith.addi` + // (or other arith ops). + // Without visiting these ops manually, the ops to update the + // offsets and strides would not be generated. + auto tritonValue = getStateOp->getOperand(0); + if (!knownPtrs.contains(tritonValue)) { + PtrState state; + OpBuilder b(getStateOp); + if (succeeded(visitOperand(tritonValue, state, + getStateOp->getLoc(), b))) { + knownPtrs[tritonValue] = state; + } else { + getStateOp->emitRemark("PtrAnalysis: Failed to populate ptr " + "state for tensor of indices"); + } + } + + return WalkResult::skip(); + }) + .Default([&](auto) { return WalkResult::advance(); }); + }); + + return success(); +} + +} // namespace tts +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/AnalysisStructured/PtrAnalysisTS.cpp b/third_party/wafer/third_party/flir/lib/AnalysisStructured/PtrAnalysisTS.cpp new file mode 100755 index 00000000..d72285d6 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/AnalysisStructured/PtrAnalysisTS.cpp @@ -0,0 +1,1618 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/AnalysisStructured/PtrAnalysis.h" +#include "triton-shared/Analysis/MaskAnalysis.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" + +#include "mlir/IR/IRMapping.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" +#include "triton-shared/Utils/Utils.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include +#include +#include +#include +#include +#include + +#define DEBUG_TYPE "triton-ptr-analysis" + +namespace mlir { + +// https://triton-lang.org/main/python-api/generated/triton.language.load.html#triton.language.load +// If pointer is a block pointer defined by make_block_ptr, a tensor is +// loaded. In this case: mask and other must be None, and boundary_check and +// padding_option can be specified to control the behavior of out-of-bound +// access. +// WORKAROUND: Assume the load/store ptr operand defining op is +// triton::MakeTensorPtrOp, convert the boundaryCheck to masked dimension +static void boundaryCheckToMaskDim(OpBuilder &builder, Location loc, + tts::PtrState ptrState, + ArrayRef boundaryCheck, + triton::MaskState &maskState) { + + // TODO: This can be optimized based on the boundaryCheck property, which + // specifies along which axis the mask is required. + for (auto i : llvm::seq(0, ptrState.getRank())) { + // The shape of the tensor is used to determine the size of the data + // being stored, so we need to convert it to index type. + auto dim = ofrToIndexValue(ptrState.shape[i], loc, builder); + auto stride = ofrToIndexValue(ptrState.strides[i], loc, builder); + + auto start = ofrToIndexValue(ptrState.offsets[i], loc, builder); + start = builder.create(loc, start, stride); + + auto end = ofrToIndexValue(ptrState.sizes[i], loc, builder); + end = builder.create(loc, start, end); + + Value upperBound = builder.create(loc, dim, end); + upperBound = builder.create(loc, start, upperBound); + auto maskedShape = builder.create(loc, upperBound, start); + + maskState.dims.push_back(maskedShape.getResult()); + } +} + +namespace tts { + +int32_t PtrState::getRank() const { + assert(offsets.size() == sizes.size() && offsets.size() == strides.size() && + shape.size() == offsets.size()); + return offsets.size(); +} + +bool PtrState::isEmpty() const { + return (getRank() == 0 && !source && !scalar); +} + +bool PtrState::hasModulo() const { + for (int32_t i = 0; i < getRank(); i++) { + if (dimHasModulo(i)) { + return true; + } + } + return false; +} + +bool PtrState::dimHasModulo(uint32_t dim) const { + assert( + !isBlockPtr() && + "Analysis should not check modulo if PtrState describes block pointer"); + + assert(dim < getRank()); + + auto intAttr = getIntAttr(shape[dim]); + if (!intAttr.has_value()) { + return true; + } + + return intAttr.value() != 0; +} + +bool PtrState::isBlockPtr() const { return !order.empty(); } + +LogicalResult PtrState::addState(const PtrState &lhsState, + const PtrState &rhsState, Operation *op, + OpBuilder &builder) { + assert(isEmpty() && lhsState.getRank() == rhsState.getRank()); + auto loc = op->getLoc(); + + if (lhsState.source && rhsState.source) { + op->emitRemark( + "PtrAnalysis: do not support adding two pointer states that both " + "have base pointers"); + return failure(); + } + + source = lhsState.source ? lhsState.source : rhsState.source; + + if (lhsState.scalar && rhsState.scalar) { + auto addOp = + builder.create(loc, lhsState.scalar, rhsState.scalar); + scalar = addOp.getResult(); + } else if (lhsState.getRank() == 0) { // both lhs and rhs are scalars + scalar = lhsState.scalar ? lhsState.scalar : rhsState.scalar; + } + + for (uint64_t i = 0; i < lhsState.getRank(); i++) { + auto newOffset = + addOFRs(lhsState.offsets[i], rhsState.offsets[i], loc, builder); + offsets.push_back(newOffset); + + auto newStride = + addOFRs(lhsState.strides[i], rhsState.strides[i], loc, builder); + strides.push_back(newStride); + + sizes.push_back(lhsState.sizes[i]); + } + + // AddPtr where both lhs and rhs containing modulo operators not supported + if (lhsState.hasModulo() && rhsState.hasModulo()) { + op->emitRemark("PtrAnalysis: do not support adding two pointer states " + "that both have modulo"); + return failure(); + } + + if (lhsState.hasModulo() || rhsState.hasModulo()) { + // visitOperandSplat and visitOperandExpandDims should enforce below + assert(lhsState.getRank() <= 2); + } + + // dealing with modulo: + // - If lhs has no modulo, skip + // - If rhs has zero offset on dim i, we can just use lhs's modulo + // - If i == 0 and rhs is the result of a splat, we will allow the add. This + // is because the user may be trying to express adding a constant offset to + // increment dim1, but pointer analysis cannot differentiate dim1 vs dim0 in + // this case. + // - Else, the analysis fails + + // An example for the 3rd condition above can look like: + // %0 = tt.splat %scalar + // %1 = tt.splat %ptr + // %2 = tt.arange + // %3 = arith.remsi %2, %size + // %4 = tt.addptr %1, %3 + // %5 = tt.addptr %4, %0 + // %5 may also occur in a loop to increment %4 every iteration. + + // Note that this is not bullet-proof. E.g., broken IR can actually increment + // dim0 while dim0 already has modulo, since Triton offsets are element-wise + // and not in unit of lower dimensions. However, this is highly unlikely but + // the analysis will provide wrong result. Hence we provide a warning in this + // case. + PtrState const *lhs = &lhsState; + PtrState const *rhs = &rhsState; + + if (rhs->hasModulo()) { + std::swap(lhs, rhs); + } + + for (uint64_t i = 0; i < lhs->getRank(); i++) { + if (!lhs->dimHasModulo(i)) { + shape.push_back(lhs->shape[i]); + } else if (hasConstZero(rhs->offsets[i])) { + shape.push_back(lhs->shape[i]); + } else if (i == 0 && lhs->getRank() == 2 && rhs->scalar) { + shape.push_back(lhs->shape[1]); + shape.push_back(lhs->shape[0]); + op->emitWarning( + "PtrAnalysis: allowing adding pointer state with modulo in dim 0 to " + "another pointer state with offset in dim 0.\nPlease verify the " + "operand that contains a scalar is meant to increment pointers in " + "dim1. If that is not the case it WILL LEAD TO WRONG COMPILATION " + "RESULTS.\n\nTo avoid this warning, use expand_dims (instead of " + "splat) to explicitly specify which dimension contains the scalar."); + break; + } else { + op->emitRemark( + "PtrAnalysis: do not support adding to operand with modulo"); + return failure(); + } + } + + return success(); +} + +void PtrState::dump() const { + llvm::dbgs() << "PtrState: "; + if (source) { + llvm::dbgs() << "source: " << source << "\n"; + } + if (scalar) { + llvm::dbgs() << "scalar: " << scalar << "\n"; + } + + llvm::dbgs() << "offsets: "; + llvm::interleave(offsets, llvm::dbgs(), "\n"); + llvm::dbgs() << "\nstrides: "; + llvm::interleave(strides, llvm::dbgs(), "\n"); + llvm::dbgs() << "\nsizes: "; + llvm::interleave(sizes, llvm::dbgs(), "\n"); + llvm::dbgs() << "\nshape: "; + llvm::interleave(shape, llvm::dbgs(), "\n"); + llvm::dbgs() << "\norder: "; + llvm::interleave(order, llvm::dbgs(), "\n"); + llvm::dbgs() << "\n"; +} + +LogicalResult PtrState::mulState(const PtrState &lhsState, + const PtrState &rhsState, Operation *op, + OpBuilder &builder) { + assert(isEmpty() && lhsState.getRank() == rhsState.getRank()); + + auto loc = op->getLoc(); + + // neither lhs nor rhs should have source, since multiplying base pointer + // does not make sense + if (lhsState.source && rhsState.source) { + op->emitRemark("PtrAnalysis: do not support multiplying base pointers"); + return failure(); + } + + // currently do not support both tensors are effectively non-scalar + if (!lhsState.scalar && !rhsState.scalar) { + op->emitRemark( + "PtrAnalysis: only support multiplying pointer states when one of " + "them represent a scalar"); + return failure(); + } + + PtrState const *lhs = &lhsState; + PtrState const *rhs = &rhsState; + + if (!rhs->scalar && lhs->scalar) { + std::swap(lhs, rhs); + } + + if (lhsState.scalar && rhsState.scalar) { + scalar = + builder.create(loc, lhsState.scalar, rhsState.scalar); + } + + for (uint64_t i = 0; i < lhs->sizes.size(); i++) { + OpFoldResult newOffset = + mulOFRValue(lhs->offsets[i], rhs->scalar, loc, builder); + OpFoldResult newStride = + mulOFRValue(lhs->strides[i], rhs->scalar, loc, builder); + OpFoldResult newShape = + mulOFRValue(lhs->shape[i], rhs->scalar, loc, builder); + offsets.push_back(newOffset); + strides.push_back(newStride); + shape.push_back(newShape); + sizes.push_back(lhs->sizes[i]); + } + + if (rhs->hasModulo()) { + op->emitRemark( + "PtrAnalysis: do not support multiplying pointer states that has " + "modulos"); + return failure(); + } + + return success(); +} + +LogicalResult PtrState::subState(const PtrState &lhsState, + const PtrState &rhsState, Operation *op, + OpBuilder &builder) { + assert(isEmpty() && lhsState.getRank() == rhsState.getRank()); + auto loc = op->getLoc(); + + if (lhsState.source && rhsState.source) { + if (lhsState.source != rhsState.source) { + op->emitRemark( + "PtrAnalysis: subtracting pointers from different bases is not supported"); + return failure(); + } + + if (lhsState.scalar && rhsState.scalar) { + auto subOp = builder.create(loc, lhsState.scalar, rhsState.scalar); + scalar = subOp.getResult(); + } else if (lhsState.scalar) { + scalar = lhsState.scalar; + } else if (rhsState.scalar) { + auto zero = builder.create( + loc, rhsState.scalar.getType(), 0); + auto negOp = builder.create(loc, zero, rhsState.scalar); + scalar = negOp.getResult(); + } + + source = nullptr; + return success(); + } + + if (!lhsState.source && rhsState.source) { + op->emitRemark("PtrAnalysis: scalar minus pointer is not meaningful"); + return failure(); + } + + source = lhsState.source ? lhsState.source : rhsState.source; + + if (lhsState.scalar && rhsState.scalar) { + auto subOp = builder.create(loc, lhsState.scalar, rhsState.scalar); + scalar = subOp.getResult(); + } else if (lhsState.getRank() == 0) { // both lhs and rhs are scalars + scalar = lhsState.scalar ? lhsState.scalar : rhsState.scalar; + } + + for (uint64_t i = 0; i < lhsState.getRank(); i++) { + auto newOffset = subOFRs(lhsState.offsets[i], rhsState.offsets[i], loc, builder); + offsets.push_back(newOffset); + + auto newStride = subOFRs(lhsState.strides[i], rhsState.strides[i], loc, builder); + strides.push_back(newStride); + + sizes.push_back(lhsState.sizes[i]); + } + + for (uint64_t i = 0; i < lhsState.getRank(); i++) { + shape.push_back(lhsState.shape[i]); + } + + return success(); +} + +tts::MakeTensorPtrOp PtrState::createTTSMakeTensorPtrOp(OpBuilder &builder, + Location loc) { + SmallVector staticSizes; + for (size_t i = 0; i < getRank(); i++) { + auto s = getIntAttr(sizes[i]); + assert(s.has_value()); + staticSizes.push_back(s.value()); + } + + auto op = builder.create( + loc, source, staticSizes, strides, offsets, shape, order); + LLVM_DEBUG({ + llvm::dbgs() << "creating tts::make_tensor_ptr:\n"; + op->dump(); + }); + + return op; +} + +LogicalResult PtrAnalysis::visitOperandAdd(arith::AddIOp addOp, PtrState &state, + const Location loc, + OpBuilder &builder) { + PtrState lhsState; + if (visitOperand(addOp.getLhs(), lhsState, loc, builder).failed()) { + return failure(); + } + + PtrState rhsState; + if (visitOperand(addOp.getRhs(), rhsState, loc, builder).failed()) + return failure(); + + // Checking for higher dimension is done in addState below + if ((lhsState.getRank() == 1 && lhsState.hasModulo()) || + (rhsState.getRank() == 1 && rhsState.hasModulo())) { + addOp->emitRemark( + "PtrAnalysis: do not support this pattern: a + arange(0, K) % M"); + return failure(); + } + + return state.addState(lhsState, rhsState, addOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandSub(arith::SubIOp subOp, PtrState &state, + const Location loc, + OpBuilder &builder) { + PtrState lhsState; + if (visitOperand(subOp.getLhs(), lhsState, loc, builder).failed()) { + return failure(); + } + + PtrState rhsState; + if (visitOperand(subOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + if (lhsState.hasModulo() || rhsState.hasModulo()) { + subOp->emitRemark("PtrAnalysis: do not support modulo for subi op\n"); + return failure(); + } + + // Checking for higher dimension is done in subState below + if ((lhsState.getRank() == 1 && lhsState.hasModulo()) || + (rhsState.getRank() == 1 && rhsState.hasModulo())) { + subOp->emitRemark( + "PtrAnalysis: do not support this pattern: a - arange(0, K) % M"); + return failure(); + } + + return state.subState(lhsState, rhsState, subOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandMul(arith::MulIOp mulOp, PtrState &state, + const Location loc, + OpBuilder &builder) { + PtrState lhsState; + if (visitOperand(mulOp.getLhs(), lhsState, loc, builder).failed()) { + return failure(); + } + + PtrState rhsState; + if (visitOperand(mulOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + return state.mulState(lhsState, rhsState, mulOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandRem(arith::RemSIOp remOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + PtrState rhsState; + if (visitOperand(remOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + if (!rhsState.scalar) { + remOp->emitRemark("PtrAnalysis: only support cases when rhs of remainder " + "contains scalar"); + return failure(); + } + + if (visitOperand(remOp.getLhs(), state, loc, builder).failed()) { + return failure(); + } + + // If there are multiple modulo ops on an expression (e.g.: (a % b) % c), we + // would have already populated the modulo states after visiting the lhs. + // Assert that all the modulo states are empty. + if (state.hasModulo()) { + remOp->emitRemark( + "PtrAnalysis: do not support multiple modulo within an expression"); + return failure(); + } + + if (state.getRank() == 1) { + // Apply the modulo before expanding shape, the common pattern is + // offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + // a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * + // stride_ak) + state.shape.back() = rhsState.scalar; + } else if (state.getRank() == 2) { + // torch inductor expands the tensor shape before applying the modulo. + // + // We only support either: + // - (tl.arange(0, end)[:, None] % mod), or + // - (tl.arange(0, end)[None, :] % mod) + // + // In both cases, we apply the modulo to the non-singleton dimension. + auto shape = cast(remOp.getResult().getType()).getShape(); + if (shape[0] == 1) { + state.shape[1] = rhsState.scalar; + } else if (shape[1] == 1) { + state.shape[0] = rhsState.scalar; + } else { + remOp->emitRemark( + "PtrAnalysis: taking modulo on a 2D tensor with no singleton " + "dimension not supported"); + return failure(); + } + } else { + remOp->emitRemark("PtrAnalysis: unsupported modulo pattern"); + return failure(); + } + return success(); +} + +LogicalResult PtrAnalysis::visitOperandExtSI(arith::ExtSIOp extOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + return visitOperand(extOp.getIn(), state, loc, builder); +} + +LogicalResult PtrAnalysis::visitOperandMakeRange(triton::MakeRangeOp rangeOp, + PtrState &state, Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + auto shape = cast(rangeOp.getType()).getShape(); + + auto start = rangeOp.getStart(); + auto end = rangeOp.getEnd(); + auto stride = (end - start + shape[0] - 1) / shape[0]; + assert(stride == 1 && + "Expect make_range op to always return tensor of stride 1"); + + state.offsets.push_back(builder.getIndexAttr(start)); + state.sizes.push_back(builder.getIndexAttr(shape[0])); + state.strides.push_back(builder.getIndexAttr(stride)); + state.shape.push_back(builder.getIndexAttr(0)); + return success(); +} + +LogicalResult +PtrAnalysis::visitOperandExpandDims(triton::ExpandDimsOp expandDimsOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + if (visitOperand(expandDimsOp.getSrc(), state, loc, builder).failed()) { + return failure(); + } + + auto dstShape = + cast(expandDimsOp.getResult().getType()).getShape(); + auto axis = expandDimsOp.getAxis(); + + assert(dstShape[axis] == 1 && + "expect changed dimension to be 1 in expand_dims"); + + // insert dimension info + state.offsets.insert(state.offsets.begin() + axis, builder.getIndexAttr(0)); + state.sizes.insert(state.sizes.begin() + axis, builder.getIndexAttr(1)); + state.strides.insert(state.strides.begin() + axis, builder.getIndexAttr(0)); + state.shape.insert(state.shape.begin() + axis, builder.getIndexAttr(0)); + + if (state.hasModulo() && state.getRank() > 2) { + expandDimsOp->emitRemark( + "PtrAnalysis: unsupported scenario where expand_dims result " + "has modulo and rank > 2"); + return failure(); + } + + return success(); +} + +LogicalResult +PtrAnalysis::visitOperandBroadcast(triton::BroadcastOp broadcastOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + auto src = broadcastOp.getSrc(); + auto dst = broadcastOp.getResult(); + + if (!isa(src.getType())) { + broadcastOp->emitRemark("PtrAnalysis: Unsupported broadcast source type"); + return failure(); + } + + auto srcShape = cast(src.getType()).getShape(); + auto dstShape = cast(dst.getType()).getShape(); + + assert(srcShape.size() == dstShape.size() && + "rank of source and destination should match"); + + if (visitOperand(src, state, loc, builder).failed()) { + return failure(); + } + + for (size_t i = 0; i < dstShape.size(); i++) { + if (srcShape[i] == dstShape[i]) { + continue; + } else if (srcShape[i] < dstShape[i]) { + state.sizes[i] = builder.getIndexAttr(dstShape[i]); + } else { + llvm_unreachable("unexpected dimensions used in broadcast"); + } + } + return success(); +} + +LogicalResult PtrAnalysis::visitOperandSplat(triton::SplatOp splatOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + auto src = splatOp.getSrc(); + auto dst = splatOp.getResult(); + auto dstShape = cast(dst.getType()).getShape(); + + if (visitOperand(src, state, loc, builder).failed()) { + return failure(); + } + + if (isa(src.getType())) { + for (auto s : dstShape) { + state.offsets.push_back(builder.getIndexAttr(0)); + state.sizes.push_back(builder.getIndexAttr(s)); + state.strides.push_back(builder.getIndexAttr(0)); + state.shape.push_back(builder.getIndexAttr(0)); + } + } else { + splatOp->emitRemark("PtrAnalysis: unsupported splat pattern"); + return failure(); + } + + // If we splat a integer value, scalar should become the offset of the outer + // most dimension + if (state.scalar) + state.offsets[0] = state.scalar; + + if (state.hasModulo() && state.getRank() > 2) { + splatOp->emitRemark("PtrAnalysis: unsupported scenario where splat result " + "has modulo and rank > 2"); + return failure(); + } + + return success(); +} + +LogicalResult PtrAnalysis::visitOperandAddptr(triton::AddPtrOp addptrOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + + PtrState ptrState; + if (visitOperand(addptrOp.getPtr(), ptrState, addptrOp.getLoc(), builder) + .failed()) { + // assert(0); + return failure(); + } + + PtrState offsetState; + if (visitOperand(addptrOp.getOffset(), offsetState, addptrOp.getLoc(), + builder) + .failed()) { + return failure(); + } + + assert(ptrState.source && "ptr field should provide source / base pointer"); + + assert(ptrState.getRank() == offsetState.getRank() && + "ptr and offset field should have the same rank"); + + return state.addState(ptrState, offsetState, addptrOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandConstSplat(arith::ConstantOp op, + PtrState &state, + const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + // this condition is to handle cases where tt.broadcast and tt.splat are + // folded + auto attr = cast(op.getValue()); + auto elementType = attr.getElementType(); + assert(attr.isSplat() && isa(elementType)); + auto values = attr.getValues(); + auto value = values[0].getValue(); + auto constAttr = builder.getIndexAttr(value.getSExtValue()); + auto constOp = arith::ConstantOp::materialize(builder, constAttr, + builder.getIndexType(), loc); + + state.scalar = constOp; + + auto resultType = cast(op.getResult().getType()); + for (size_t i = 0; i < resultType.getShape().size(); i++) { + if (i == 0) { + state.offsets.push_back(constOp.getResult()); + } else { + state.offsets.push_back(builder.getIndexAttr(0)); + } + + state.sizes.push_back(builder.getIndexAttr(resultType.getShape()[i])); + state.strides.push_back(builder.getIndexAttr(0)); + state.shape.push_back(builder.getIndexAttr(0)); + } + + return success(); +} + +LogicalResult PtrAnalysis::visitOperandMakeTPtr(tts::MakeTensorPtrOp makeTPtrOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + + assert(state.isEmpty()); + state.source = makeTPtrOp.getBase(); + state.offsets = makeTPtrOp.getMixedOffsets(); + state.sizes = makeTPtrOp.getMixedSizes(); + state.strides = makeTPtrOp.getMixedStrides(); + state.shape = makeTPtrOp.getMixedShape(); + state.order = SmallVector(makeTPtrOp.getOrder()); + + return success(); +} + +LogicalResult +PtrAnalysis::visitOperandMakeTensorPtr(triton::MakeTensorPtrOp makeTPtrOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + assert(state.isEmpty()); + state.source = makeTPtrOp.getBase(); + + if (makeTPtrOp.getOrder().empty()) { + makeTPtrOp->emitRemark( + "PtrAnalysis: expect tt.make_tensor_ptr to have order field set"); + return failure(); + } + + auto resType = cast(makeTPtrOp.getResult().getType()); + auto pointeeType = cast(resType.getPointeeType()); + auto shape = pointeeType.getShape(); + + for (int64_t i = 0; i < pointeeType.getRank(); i++) { + state.sizes.push_back(builder.getIndexAttr(shape[i])); + + auto strideCst = builder.create( + loc, builder.getIndexType(), makeTPtrOp.getStrides()[i]); + state.strides.push_back(strideCst.getResult()); + + auto offsetCst = builder.create( + loc, builder.getIndexType(), makeTPtrOp.getOffsets()[i]); + + auto scaledOffset = builder.create( + loc, offsetCst.getResult(), strideCst.getResult()); + state.offsets.push_back(scaledOffset.getResult()); + + auto shapeCst = builder.create( + loc, builder.getIndexType(), makeTPtrOp.getShape()[i]); + state.shape.push_back(shapeCst.getResult()); + } + state.order = SmallVector(makeTPtrOp.getOrder()); + assert(state.isBlockPtr() && + "tt.make_tensor_ptr pointer state should describe a block pointer"); + + return success(); +} + +LogicalResult PtrAnalysis::visitOperandForOp(scf::ForOp forOp, Value operand, + PtrState &state, + const Location loc, + OpBuilder &builder) { + + auto it = llvm::find(forOp->getResults(), operand); + auto index = std::distance(forOp->getResults().begin(), it); + + auto newState = getLoopResultPtrState(forOp, index); + if (failed(newState)) { + forOp.emitError( + "Rewrite for-op failed. Could not find PtrState returned by " + "the loop."); + return failure(); + } + + state = newState.value(); + return success(); +} + +LogicalResult PtrAnalysis::visitOperandBitcast(triton::BitcastOp op, + PtrState &state, + const Location loc, + OpBuilder &builder) { + auto resType = op.getResult().getType(); + if (isa(resType)) { + return visitOperand(op.getSrc(), state, loc, builder); + } + state.source = op.getResult(); + return success(); +} + +LogicalResult PtrAnalysis::visitOperand(Value operand, PtrState &state, + const Location loc, + OpBuilder &builder) { + + if (knownPtrs.find(operand) != knownPtrs.end()) { + state = knownPtrs.lookup(operand); + return success(); + } + + if (isa(operand.getType())) { + OpBuilder::InsertionGuard guard(builder); + if (!isa(operand) && operand.getDefiningOp()) { + builder.setInsertionPointAfter(operand.getDefiningOp()); + } + auto castOp = builder.create( + loc, builder.getIndexType(), operand); + state.scalar = castOp.getResult(); + return success(); + } else if (isa(operand.getType())) { + state.scalar = operand; + return success(); + } + + if (isa(operand.getType())) { + // A scalar pointer can either be produced by AddPtrOp or a block + // argument + if (auto op = operand.getDefiningOp()) { + if (auto addPtrOp = dyn_cast(op)) { + return visitOperandAddptr(cast(op), state, loc, + builder); + } else if (auto castOp = dyn_cast(op)) { + return visitOperandBitcast(castOp, state, loc, builder); + } else if (auto makeTensorOp = dyn_cast(op)) { + llvm_unreachable("Unexpected operand defining operation tts.make_tptr"); + } else if (auto selectOp = dyn_cast(op)) { + state.source = selectOp.getResult(); + return success(); + } else { + op->emitRemark("Unexpected operand defining operation"); + return failure(); + } + } else { + state.source = operand; + return success(); + } + } + + if (auto op = operand.getDefiningOp()) { + return visitOperandAdd(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandMul(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandSub(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandMakeRange(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandBroadcast(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandSplat(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandExpandDims(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandAddptr(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandConstSplat(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandRem(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandExtSI(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandForOp(op, operand, state, loc, builder); + } else if (!operand.getDefiningOp()) { + if (!knownPtrs.contains(operand)) { + return failure(); + } + + // This operand must be an iter-arg of an inner-loop in a multiple-level + // nested loop, which means its PtrState must have already been populated + // during rewriteForOp of the parent loop. + state = knownPtrs[operand]; + return success(); + } else { + llvm::dbgs() << "PtrAnalysis: encountered addptr operand produced by an " + "unsupported operation\n"; + operand.dump(); + return failure(); + } +} + +LogicalResult PtrAnalysis::rewriteAddptrOp(triton::AddPtrOp op) { + OpBuilder builder(op); + + PtrState state; + if (visitOperandAddptr(op, state, op.getLoc(), builder).failed()) { + return failure(); + } + + knownPtrs[op.getResult()] = state; + + if (isa(op.getPtr().getType())) { + auto maketptrOp = state.createTTSMakeTensorPtrOp(builder, op.getLoc()); + ptrMap.map(op.getResult(), maketptrOp.getResult()); + } else { + // record the ptr as we have visited and built up the state for this scalar + // pointer, which may be used by rewriteForOp later. + ptrMap.map(op.getResult(), op.getResult()); + } + return success(); +} + +LogicalResult PtrAnalysis::rewriteBitcastOp(triton::BitcastOp op) { + LLVM_DEBUG({ + llvm::dbgs() << "Rewriting bitcast op:\n"; + op->dump(); + }); + OpBuilder builder(op); + + PtrState state; + if (visitOperandBitcast(op, state, op.getLoc(), builder).failed()) { + return failure(); + } + + knownPtrs[op.getResult()] = state; + + if (isa(op.getOperand().getType())) { + auto resType = op.getType(); + if (auto tensorType = llvm::dyn_cast(resType)) { + resType = tensorType.getElementType(); + } + auto newBitcast = + builder.create(op->getLoc(), resType, state.source); + state.source = newBitcast; + auto maketptrOp = state.createTTSMakeTensorPtrOp(builder, op.getLoc()); + ptrMap.map(op.getResult(), maketptrOp.getResult()); + } else { + ptrMap.map(op.getResult(), op.getResult()); + } + return success(); +} + +LogicalResult PtrAnalysis::rewriteMakeTensorPtrOp(triton::MakeTensorPtrOp op) { + OpBuilder builder(op); + + PtrState state; + if (visitOperandMakeTensorPtr(op, state, op.getLoc(), builder).failed()) { + return failure(); + } + + auto maketptrOp = state.createTTSMakeTensorPtrOp(builder, op.getLoc()); + knownPtrs[op.getResult()] = state; + ptrMap.map(op.getResult(), maketptrOp.getResult()); + return success(); +} + +LogicalResult PtrAnalysis::rewriteAdvanceOp(triton::AdvanceOp op) { + OpBuilder builder(op); + auto loc = op.getLoc(); + + PtrState state; + if (visitOperand(op->getOperand(0), state, loc, builder).failed()) { + op->emitRemark("PtrAnalysis: Failed to analyze ptr of tt.advance"); + return failure(); + } + assert(state.isBlockPtr() && + "tt.advance pointer state should describe a block pointer"); + + auto incrementOffsets = op.getOffsets(); + + SmallVector newOffsets; + for (auto [increment, offset, stride] : + llvm::zip(incrementOffsets, state.offsets, state.strides)) { + Value offsetValue; + if (auto offsetIntAttr = getIntAttr(offset)) { + auto constOp = builder.create( + loc, builder.getIndexAttr(offsetIntAttr.value())); + offsetValue = constOp.getResult(); + } else { + offsetValue = cast(offset); + } + auto castOp = builder.create( + loc, builder.getIndexType(), increment); + auto mulOp = builder.create(loc, castOp.getResult(), + cast(stride)); + auto addOp = + builder.create(loc, mulOp.getResult(), offsetValue); + newOffsets.push_back(addOp.getResult()); + } + + state.offsets = SmallVector(newOffsets); + + auto newOp = state.createTTSMakeTensorPtrOp(builder, loc); + knownPtrs[op.getResult()] = state; + ptrMap.map(op.getResult(), newOp.getResult()); + return success(); +} + +static bool isPointerType(Type t) { + if (auto tensor = llvm::dyn_cast(t)) { + return isa(tensor.getElementType()); + } + return isa(t); +} + +FailureOr PtrAnalysis::getLoopInitArgPtrState(scf::ForOp forOp, + size_t index) { + auto ptr = forOp.getInitArgs()[index]; + + // If the pointer into the scf.for was defined by tts.get_structured_state, + // we can get the pointer state from the original pointer (the op's input): + // + // %ptr, %offset_1, %offset_2,..., %stride_1, %stride_2,... = + // tts.get_structured_state %original + // scf.for ... (%ptr) {...} + if (auto getStateOp = ptr.getDefiningOp()) { + auto originalPtr = getStateOp->getOperand(0); + if (knownPtrs.count(originalPtr)) { + return knownPtrs[originalPtr]; + } + } + + // For nested loops scenarios, a pointer in init-args can be returned from + // another loop of the same level: + // e.g.: + // clang-format off + // %22:2 = scf.for %arg4 = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%arg5 = %11, %arg6 = %15) -> (tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr>) : i32 { + // %23 = scf.for %arg7 = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%arg8 = %arg5) -> (tensor<2x2x!tt.ptr>) : i32 { + // %26 = tt.addptr %arg8, %17 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> + // scf.yield %26 : tensor<2x2x!tt.ptr> + // } + // %24:2 = scf.for %arg7 = %c0_i32 to %c2_i32 step %c1_i32 iter_args(%arg8 = %23, %arg9 = %arg6) -> (tensor<2x2x!tt.ptr>, tensor<2x2x!tt.ptr>) : i32 { + // %26 = tt.load %arg8 : tensor<2x2x!tt.ptr> + // %27 = tt.addptr %arg8, %19 : tensor<2x2x!tt.ptr>, tensor<2x2xi32> + // ... + // } + // ... + // } + // clang-format on + // Notice %arg8 = %23 comes from the return value of the first loop. + if (auto forOp = ptr.getDefiningOp()) { + return getLoopResultPtrState(forOp, index); + } + + // If the pointer isn't defined by tts.get_structured_state nor another loop, + // it means the current pointer is an iterarg of the outer loop. + // In such cases, the outer loops would have already set up the PtrState for + // us already. + // + // scf.for iterargs(%ptr = %init_arg) { + // scf.for iterargs(%ptr1 = %ptr) { <--- we're dealing with `%ptr1` here. + // ... + // } + // } + if (knownPtrs.count(ptr)) { + assert(!ptr.getDefiningOp() && "Expect the ptr to be an iterarg"); + return knownPtrs[ptr]; + } + + return failure(); +} + +PtrState PtrAnalysis::reconcileLoopPtrState( + scf::ForOp forOp, size_t iterArgIndex, const PtrState &state, + llvm::function_ref getReplacementVal) { + PtrState newState = state; + int cnt = iterArgIndex + 1; + if (newState.getRank() == 0) { + assert(newState.scalar); + // for scalar pointers, the scalar contains the offset and is the only + // relevant newState that could be updated by the loop. + newState.scalar = getReplacementVal(forOp, cnt); + } else { + for (auto &offset : newState.offsets) { + offset = getReplacementVal(forOp, cnt++); + } + + for (auto &stride : newState.strides) { + stride = getReplacementVal(forOp, cnt++); + } + } + + return newState; +} + +FailureOr PtrAnalysis::getLoopIterArgPtrState(scf::ForOp forOp, + size_t index) { + auto state = getLoopInitArgPtrState(forOp, index); + if (failed(state)) { + return failure(); + } + + return reconcileLoopPtrState( + forOp, index, state.value(), + [](scf::ForOp op, size_t index) { return op.getRegionIterArg(index); }); +} + +FailureOr PtrAnalysis::getLoopResultPtrState(scf::ForOp forOp, + size_t index) { + auto state = getLoopInitArgPtrState(forOp, index); + if (failed(state)) { + return failure(); + } + + return reconcileLoopPtrState( + forOp, index, state.value(), + [](scf::ForOp op, size_t index) { return op->getResult(index); }); +} + +LogicalResult PtrAnalysis::rewriteForOp(scf::ForOp op) { + for (auto [i, arg] : llvm::enumerate(op.getRegionIterArgs())) { + if (!maybeStructuredArgs.contains(arg)) { + continue; + } + + auto state = getLoopIterArgPtrState(op, i); + if (failed(state)) { + // Because the maybeStructuredArgs may contain values that are not + // considered structured by PtrAnalysis, failing to retrieve the PtrState + // should not fail the rewrite process. + // We emit an error for diagnostics and debugging purposes. + op->emitWarning( + "Rewrite for-op failed. Could not find PtrState for iter-arg index " + + std::to_string(i)); + continue; + } + + // Save the current init arg's PtrState + knownPtrs[arg] = state.value(); + + // For tensors of pointers, create a tts.make_tptr at the beginning of the + // loop body that correspond to this region iter arg. In case it is used + // by tt.load/tt.store in the loop body before pointer updates, this will + // make sure rewriteLoadOp/rewriteStoreOp can use the analysis result. + // E.g., given the following input (%tensor_of_ptr is a block arg): + // scf.for (%tensor_of_ptr) { + // %data = tt.load %tensor_of_ptr + // // more operations to update %tensor_of_ptr + // } + // We may produce the following output: + // scf.for (%base_ptr, %stride, %offset) { + // %tensor_of_ptr = tts.make_tptr(%base_ptr, %stride, %offset) + // %data = tts.load %tensor_of_ptr + // // more operations to update %offset + // } + // If %tensor_of_ptr is not used (i.e., %tensor_of_ptr is updated before + // used in the original IR), it will simply be removed by + // canonicalization. + + // For scalar pointers, there is no need to create a tts.addptr at the + // beginning of the loop body. We don't lower tt.load and tt.store on + // scalars in this pass; pointer arithmetics can also just use the + // original pointer. + // Note that there can be tensor of indices in iter-arg, so we only create + // the make_tensor_ptr op when the arg is of pointer type. + if (isPointerType(arg.getType())) { + if (state->getRank() != 0) { + OpBuilder builder(op.getRegion()); + auto maketptrOp = state->createTTSMakeTensorPtrOp(builder, op.getLoc()); + ptrMap.map(arg, maketptrOp.getResult()); + } + } + } + + // Recursively rewrite the inner ops + if (rewriteOp(op).failed()) { + op->emitRemark( + "PtrAnalysis: update loop body failed when rewriting for op"); + return failure(); + } + + return success(); +} + +LogicalResult +PtrAnalysis::rewriteGetStructuredStateOp(tts::GetStructuredStateOp op) { + auto tritonValue = op->getOperand(0); + + OpBuilder builder(op); + + // If this triton value isn't known, it means PtrAnalysis has failed to + // analyze this pointer. In such cases, simply remap all uses of the + // structured value back to its original triton value. + if (!knownPtrs.contains(tritonValue)) { + op.emitRemark( + "Rewrite GetStructuredStateOp failed. Could not find PtrState."); + auto numResults = op.getNumResults(); + SmallVector replacements( + numResults, builder.create(op.getLoc(), + builder.getIndexAttr(0))); + replacements.front() = tritonValue; + op.getResults().replaceAllUsesWith(replacements); + return failure(); + } + + tts::PtrState state = knownPtrs[tritonValue]; + Value remappedValue = + ptrMap.contains(tritonValue) ? ptrMap.lookup(tritonValue) : tritonValue; + + SmallVector replacements{remappedValue}; + + if (state.getRank() == 0) { + // For scalar pointers, the scalar contains the offset and is the only + // relevant state that could be updated by the loop. + if (state.scalar) { + replacements.push_back(state.scalar); + } else { + // This operand is a pointer directly from the kernel arguments. + // Use offset 0. + assert(!tritonValue.getDefiningOp()); + replacements.push_back(builder.create( + op.getLoc(), builder.getIndexAttr(0))); + } + } else { + for (auto [j, s] : llvm::enumerate(state.offsets)) { + auto sIntAttr = getIntAttr(s); + if (sIntAttr) { + auto constOp = builder.create( + op.getLoc(), builder.getIndexAttr(sIntAttr.value())); + replacements.push_back(constOp.getResult()); + } else { + replacements.push_back(cast(s)); + } + } + + for (auto [j, s] : llvm::enumerate(state.strides)) { + auto sIntAttr = getIntAttr(s); + if (sIntAttr) { + auto constOp = builder.create( + op.getLoc(), builder.getIndexAttr(sIntAttr.value())); + replacements.push_back(constOp.getResult()); + } else { + replacements.push_back(cast(s)); + } + } + } + + op->replaceAllUsesWith(replacements); + op->erase(); + return success(); +} + +LogicalResult PtrAnalysis::rewriteLoadOp(triton::LoadOp op, + bool useUnsafeMask) { + auto ptr = ptrMap.lookupOrNull(op.getPtr()); + auto mask = op.getMask(); + auto other = op.getOther(); + auto loc = op.getLoc(); + + if (!ptr) { + op->emitRemark("PtrAnalysis: pointer is not replace with tts.make_tptr so " + "loadOp cannot be rewritten"); + return failure(); + } + + auto ptrType = dyn_cast(ptr.getType()); + if (ptrType && !isa(ptrType.getPointeeType())) { + op->emitRemark("PtrAnalysis: scalar loadOp will not be rewritten"); + return failure(); + } + + ArrayRef dims; + mlir::triton::MaskState mstate(useUnsafeMask); + Value scalarOther; + + OpBuilder builder(op); + // Analyze the mask operand to determine at runtime the size of the data we + // are moving. + if (mask) { + if (mstate.parse(mask, loc, builder).failed()) { + op->emitRemark("MaskAnalysis failed"); + return failure(); + } + dims = mstate.dims; + } + + auto boundaryCheck = op.getBoundaryCheck(); + if (!boundaryCheck.empty()) { + boundaryCheckToMaskDim(builder, loc, knownPtrs.at(op.getPtr()), + boundaryCheck, mstate); + dims = mstate.dims; + } + + if (other) { + assert(mask && "other value used while no masks are specified"); + + scalarOther = triton::getScalarValue(other, loc, builder); + if (!scalarOther) { + op->emitRemark("other value used in masked load produced by " + "unsupported instruction"); + return failure(); + } + } + + auto loadOp = builder.create(loc, ptr, dims, scalarOther); + + LLVM_DEBUG({ + llvm::dbgs() << "creating tts::load:\n"; + loadOp->dump(); + }); + + op.replaceAllUsesWith(loadOp.getResult()); + op->erase(); + return success(); +} + +// Structured values from the TritonStructuredDialect have offsets and strides +// that might change in each loop iteration and hence will appear in an scf.for +// iter-args like so: +// +// %structured, %offsets, %strides = tts.get_structured_state +// scf.for (%arg0 = %structured, %arg1 = %offsets, %arg2 = %strides) { +// %a = %arg0 + 1 +// %b = %b + 2 +// scf.for (%arg1 = %b) { +// ... +// } +// } +// +// In `rewriteForOp`, we have to recognize such structured values in order to +// rewrite their PtrState accordingly. Previously, only values of Pointer-like +// type (e.g.: tensor> or tt.ptr>), so detecting these values +// is as easy as checking the type. +// +// Now, tensor of indices could also appear in a loop's iter-arg. To reliably +// detect all such cases, we perform a BFS-like traversal of the IR where the +// sources are the results of `tts.get_structured_state`. All values that +// originate from the results of `tts.get_structured_state` are consider +// "maybeStructured". If a loop's iter-arg is considered "maybeStructured", we +// must set up their PtrState during `rewriteForOp`. +void PtrAnalysis::initializeMaybeStructuredArgs(Operation *op) { + std::queue q; + DenseSet visited; + + op->walk([&q, &visited](tts::GetStructuredStateOp getStateOp) { + Value value = getStateOp->getResult(0); + visited.insert(value); + q.push(value); + }); + + while (!q.empty()) { + auto v = q.front(); + q.pop(); + for (auto user : v.getUsers()) { + // scf.for is a special case. We have 2 set of values to consider: + // - iter-args + // - loop results + // for every init arg that originates from a `tts.get_structured_state` + // op, its corresponding iter-arg and loop result will also be considered + // "maybeStructured". + if (auto forOp = dyn_cast(user)) { + auto it = llvm::find(forOp.getInitArgs(), v); + + if (it == forOp.getInitArgs().end()) { + continue; + } + + auto argIndex = std::distance(forOp.getInitArgs().begin(), it); + auto iterArg = forOp.getRegionIterArg(argIndex); + auto tiedLoopRes = forOp.getTiedLoopResult(iterArg); + + SmallVector neighbors{iterArg, tiedLoopRes}; + for (auto neighbor : neighbors) { + maybeStructuredArgs.insert(neighbor); + if (!visited.contains(neighbor)) { + visited.insert(neighbor); + q.push(neighbor); + } + } + + } else { + for (auto res : user->getResults()) { + if (res.getType() != v.getType()) { + continue; + } + maybeStructuredArgs.insert(res); + if (!visited.contains(res)) { + visited.insert(res); + q.push(res); + } + } + } + } + } +} + +LogicalResult PtrAnalysis::rewriteStoreOp(triton::StoreOp op, + bool useUnsafeMask) { + auto ptr = ptrMap.lookupOrNull(op.getPtr()); + auto val = op.getValue(); + auto mask = op.getMask(); + auto loc = op.getLoc(); + + if (!ptr) { + op->emitRemark("PtrAnalysis: pointer is not replace with tts.make_tptr so " + "storeOp cannot be rewritten"); + return failure(); + } + + auto ptrType = dyn_cast(ptr.getType()); + if (ptrType && !isa(ptrType.getPointeeType())) { + op->emitRemark("PtrAnalysis: scalar storeOp will not be rewritten"); + return failure(); + } + + ArrayRef dims; + mlir::triton::MaskState mstate(useUnsafeMask); + + OpBuilder builder(op); + + // Analyze the mask operand to determine at runtime the size of the data + // are moving. + if (mask) { + if (mstate.parse(mask, loc, builder).failed()) { + op->emitRemark("MaskAnalysis failed"); + return failure(); + } + dims = mstate.dims; + } + + auto boundaryCheck = op.getBoundaryCheck(); + if (!boundaryCheck.empty()) { + boundaryCheckToMaskDim(builder, loc, knownPtrs.at(op.getPtr()), + boundaryCheck, mstate); + dims = mstate.dims; + } + + auto storeOp = builder.create(loc, ptr, val, dims); + + LLVM_DEBUG({ + llvm::dbgs() << "creating tts::store:\n"; + storeOp->dump(); + }); + + op->erase(); + return success(); +} + +LogicalResult PtrAnalysis::rewriteAtomicRMWOp(triton::AtomicRMWOp op, + bool useUnsafeMask) { + auto ptr = ptrMap.lookupOrNull(op.getPtr()); + auto val = op.getVal(); + auto mask = op.getMask(); + auto loc = op.getLoc(); + + if (!ptr) { + LLVM_DEBUG(op->emitRemark( + "PtrAnalysis: pointer is not replace with tts.make_tptr so " + "AtomicCASOp cannot be rewritten")); + return failure(); + } + + auto ptrType = dyn_cast(ptr.getType()); + if (ptrType && !isa(ptrType.getPointeeType())) { + LLVM_DEBUG(op->emitRemark( + "PtrAnalysis: scalar AtomicCASOp will not be rewritten")); + return failure(); + } + + ArrayRef dims; + mlir::triton::MaskState mstate(useUnsafeMask); + + OpBuilder builder(op); + + // Analyze the mask operand to determine at runtime the size of the data + // are moving. + if (mask) { + if (mstate.parse(mask, loc, builder).failed()) { + LLVM_DEBUG(op->emitRemark("MaskAnalysis failed")); + return failure(); + } + dims = mstate.dims; + } + + auto atomicRMWOp = builder.create( + loc, op.getType(), ptr, val, dims, op.getAtomicRmwOpAttr(), + op.getSemAttr(), op.getScopeAttr()); + + LLVM_DEBUG({ + llvm::dbgs() << "creating tts::atomic_rmw:\n"; + atomicRMWOp->dump(); + }); + + op.replaceAllUsesWith(atomicRMWOp.getResult()); + op->erase(); + return success(); +} + +LogicalResult PtrAnalysis::rewriteAtomicCASOp(triton::AtomicCASOp op) { + auto ptr = ptrMap.lookupOrNull(op.getPtr()); + auto cmp = op.getCmp(); + auto val = op.getVal(); + + auto loc = op.getLoc(); + + if (!ptr) { + LLVM_DEBUG(op->emitRemark( + "PtrAnalysis: pointer is not replace with tts.make_tptr so " + "AtomicCASOp cannot be rewritten")); + return failure(); + } + + auto ptrType = dyn_cast(ptr.getType()); + if (ptrType && !isa(ptrType.getPointeeType())) { + LLVM_DEBUG(op->emitRemark( + "PtrAnalysis: scalar AtomicCASOp will not be rewritten")); + return failure(); + } + + OpBuilder builder(op); + + auto atomicCASOp = builder.create( + loc, op.getType(), ptr, cmp, val, nullptr, op.getSemAttr(), + op.getScopeAttr()); + + LLVM_DEBUG({ + llvm::dbgs() << "creating tts::atomic_cas:\n"; + atomicCASOp->dump(); + }); + + op.replaceAllUsesWith(atomicCASOp.getResult()); + op->erase(); + return success(); +} + +LogicalResult PtrAnalysis::rewriteOp(Operation *rootOp, bool useUnsafeMask) { + LLVM_DEBUG({ + llvm::dbgs() << "rewriting rootOp\n"; + rootOp->dump(); + }); + + rootOp->walk([&](Operation *op) { + if (op == rootOp) { + return WalkResult::advance(); + } + return TypeSwitch(op) + .Case([&](auto addptr) { + if (rewriteAddptrOp(addptr).failed()) { + addptr->emitRemark("PtrAnalysis: Failed to rewrite AddPtrOp"); + } + return WalkResult::advance(); + }) + .Case([&](auto bitcast) { + if (rewriteBitcastOp(bitcast).failed()) { + bitcast->emitRemark("PtrAnalysis: Failed to rewrite BitcastOp"); + } + return WalkResult::advance(); + }) + .Case([&](auto maketptr) { + if (rewriteMakeTensorPtrOp(maketptr).failed()) { + maketptr->emitRemark( + "PtrAnalysis: Failed to rewrite MakeTensorPtrOp"); + } + return WalkResult::advance(); + }) + .Case([&](auto advance) { + if (rewriteAdvanceOp(advance).failed()) { + advance->emitRemark("PtrAnalysis: Failed to rewrite AdvanceOp"); + } + return WalkResult::advance(); + }) + .Case([&](auto load) { + if (rewriteLoadOp(load, useUnsafeMask).failed()) { + load->emitRemark("PtrAnalysis: Failed to rewrite LoadOp"); + return WalkResult::advance(); + } + return WalkResult::skip(); + }) + .Case([&](auto store) { + if (rewriteStoreOp(store, useUnsafeMask).failed()) { + store->emitRemark("PtrAnalysis: Failed to rewrite StoreOp"); + return WalkResult::advance(); + } + return WalkResult::skip(); + }) + .Case([&](auto atomicRMW) { + if (rewriteAtomicRMWOp(atomicRMW, useUnsafeMask).failed()) { + LLVM_DEBUG(atomicRMW->emitRemark( + "PtrAnalysis: Failed to rewrite AtomicRMWOp")); + return WalkResult::advance(); + } + return WalkResult::skip(); + }) + .Case([&](auto atomicCAS) { + if (rewriteAtomicCASOp(atomicCAS).failed()) { + LLVM_DEBUG(atomicCAS->emitRemark( + "PtrAnalysis: Failed to rewrite AtomicCASOp")); + return WalkResult::advance(); + } + return WalkResult::skip(); + }) + .Case([&](auto forOp) { + // `rewriteForOp` recursively visits its children, so regardless + // whether the rewrite succeeds or not, we need to return "skip" so + // that the the walk does not visit the for-op's child operations + // the second time. + if (rewriteForOp(forOp).failed()) { + forOp->emitRemark("PtrAnalysis: Failed to rewrite ForOp"); + } + return WalkResult::skip(); + }) + .Case( + [&](tts::GetStructuredStateOp getStateOp) { + // For tensor of indices potentially being used in pointer + // arithmetic sequence, we need to manually populate the state of + // none already exists. + // This process is necessary because unlike triton pointers in a + // loop which always have a `tt.addptr` that triggers the rewrite + // process which includes generating the ops for updating offsets + // and strides, tensor of indices only have a simple `arith.addi` + // (or other arith ops). + // Without visiting these ops manually, the ops to update the + // offsets and strides would not be generated. + auto tritonValue = getStateOp->getOperand(0); + if (!knownPtrs.contains(tritonValue)) { + PtrState state; + OpBuilder b(getStateOp); + if (succeeded(visitOperand(tritonValue, state, + getStateOp->getLoc(), b))) { + knownPtrs[tritonValue] = state; + } else { + getStateOp->emitRemark("PtrAnalysis: Failed to populate ptr " + "state for tensor of indices"); + } + } + + return WalkResult::skip(); + }) + .Default([&](auto) { return WalkResult::advance(); }); + }); + + return success(); +} + +} // namespace tts +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/CMakeLists.txt new file mode 100755 index 00000000..677b3975 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/CMakeLists.txt @@ -0,0 +1,9 @@ +if (FLIR_BUILD_INCUBATED) + add_definitions(-D__FLIR_BUILD_INCUBATED__) + add_subdirectory(UtilsIncubated) +endif() +add_subdirectory(Analysis) +add_subdirectory(AnalysisStructured) +add_subdirectory(Conversion) +add_subdirectory(Dialect) +add_subdirectory(Utils) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/CMakeLists.txt new file mode 100755 index 00000000..92adb7d0 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/CMakeLists.txt @@ -0,0 +1,20 @@ +if(NOT FLAGTREE_BACKEND STREQUAL "wafer") + add_subdirectory(TritonToLinalgExperimental) + add_subdirectory(TritonArithToLinalg) + add_subdirectory(StructuredToMemref) + add_subdirectory(ReconcilePtrCasts) +endif() +add_subdirectory(TritonToLinalg) +add_subdirectory(TritonToStructured) +add_subdirectory(TritonToUnstructured) +add_subdirectory(TritonPtrToMemref) +add_subdirectory(UnstructuredToMemref) +add_subdirectory(MemrefCopyToDMA_FlagTree) +add_subdirectory(NoBufferize_FlagTree) +if (FLIR_BUILD_INCUBATED) + add_subdirectory(TritonToAnnotation) + add_subdirectory(TritonToLinalgIncubated) + add_subdirectory(DiscreteMaskAccessConversion) + add_subdirectory(TritonToUnstructureIncubated) + add_subdirectory(TritonToStructuredIncubated) +endif() \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/lib/Conversion/DiscreteMaskAccessConversion/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/DiscreteMaskAccessConversion/CMakeLists.txt new file mode 100755 index 00000000..0c7f1456 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/DiscreteMaskAccessConversion/CMakeLists.txt @@ -0,0 +1,14 @@ +add_triton_library(DiscreteMaskAccessConversion + DiscreteMaskAccessConversionPass.cpp + + DEPENDS + DiscreteMaskAccessConversionPassIncGen + + LINK_LIBS + BiShengIRHIVMDialect + MLIRIR + MLIRPass + MLIRTransforms + MLIRSupport + TritonIR +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.cpp new file mode 100755 index 00000000..c688ebc0 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.cpp @@ -0,0 +1,196 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/DiscreteMaskAccessConversion/Passes.h" +#include "incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#if __has_include("bishengir/Dialect/HIVM/IR/HIVM.h") +#include "bishengir/Dialect/HIVM/IR/HIVM.h" +#endif + +#include "mlir/IR/Attributes.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "llvm/ADT/StringRef.h" +#include "llvm/Support/LogicalResult.h" + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_DISCRETEMASKACCESSCONVERSION +#include "incubated/Conversion/DiscreteMaskAccessConversion/Passes.h.inc" + +} // namespace triton +} // namespace mlir + +using namespace mlir; +using namespace hivm; + +LogicalResult isDiscreteMask(Operation *op, Value mask, + PatternRewriter &rewriter) { + if (!mask) + return failure(); + + mlir::triton::Incubated::MaskState mstate; + auto isContMask = mstate.parse(mask, op->getLoc(), rewriter); + if (!isContMask.failed()) { + mstate.eraseInsertedOps(op, rewriter); + return failure(); + } + return success(); +} + +struct DiscreteMaskStoreConversion : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(triton::StoreOp op, + PatternRewriter &rewriter) const final { + auto mask = op.getMask(); + auto loc = op.getLoc(); + auto dst = op.getPtr(); + auto src = op.getValue(); + + if (failed(isDiscreteMask(op, mask, rewriter))) + return failure(); + + auto loadFromDstOp = rewriter.create( + loc, dst, op.getCache(), op.getEvict(), false); + + auto selOp = rewriter.create(loc, mask, src, + loadFromDstOp.getResult()); + auto newStore = rewriter.create( + loc, dst, selOp, op.getCache(), op.getEvict()); + newStore->setAttr(ConverterUtils::discreteMaskAttrName, + UnitAttr::get(rewriter.getContext())); + rewriter.replaceOp(op, newStore); + return success(); + } +}; + +struct DiscreteMaskLoadConversion : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(triton::LoadOp op, + PatternRewriter &rewriter) const final { + auto loc = op.getLoc(); + auto other = op.getOther(); + auto mask = op.getMask(); + auto ptr = op.getPtr(); + + if (failed(isDiscreteMask(op, mask, rewriter))) + return failure(); + if (compileOn91095Flag && forceSimtTemplateFlag) + return failure(); + + if (!other) { + FailureOr constant = specializeTypelessValueToConstant( + TypelessValue::Zero, ptr.getType(), loc, rewriter); + if (failed(constant)) + llvm_unreachable("Unsupported type for constant creation"); + other = *constant; + } + + auto newLoadOp = rewriter.create( + loc, ptr, op.getCache(), op.getEvict(), op.getIsVolatile()); + auto discreteMaskOp = + rewriter.create(loc, mask, newLoadOp, other); + rewriter.replaceOp(op, discreteMaskOp); + return success(); + } +}; + +struct DiscreteMaskAtomicConversion + : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(mlir::triton::AtomicRMWOp op, + PatternRewriter &rewriter) const final { + auto loc = op.getLoc(); + auto ptr = op.getPtr(); + auto src = op.getVal(); + auto mask = op.getMask(); + RMWOp rmwOp = op.getAtomicRmwOp(); + + if (failed(isDiscreteMask(op, mask, rewriter))) + return failure(); + + const std::map initMap = { + {RMWOp::FADD, TypelessValue::Zero}, + {RMWOp::ADD, TypelessValue::Zero}, + {RMWOp::UMAX, TypelessValue::Zero}, + {RMWOp::OR, TypelessValue::Zero}, + {RMWOp::MIN, TypelessValue::Max}, + {RMWOp::UMIN, TypelessValue::Max}, + {RMWOp::AND, TypelessValue::Max}, + {RMWOp::MAX, TypelessValue::Min}, + {RMWOp::XOR, TypelessValue::Zero}, + {RMWOp::XCHG, TypelessValue::Undefined}, + }; + assert(initMap.find(rmwOp) != initMap.end()); + auto typelessVal = initMap.at(rmwOp); + if (typelessVal == TypelessValue::Undefined) { + // Undefined default value atomic op will be decomposed in AscendNPU-IR + op->setAttr(ConverterUtils::discreteMaskAttrName, + UnitAttr::get(rewriter.getContext())); + return failure(); + } + + FailureOr fill = specializeTypelessValueToConstant( + typelessVal, src.getType(), loc, rewriter); + if (failed(fill)) + op->emitError("Unsupported atomic operation."); + + auto maskedValue = rewriter.create(loc, mask, src, *fill); + auto newAtomicOp = rewriter.create( + loc, src.getType(), rmwOp, ptr, maskedValue, mlir::Value(), op.getSem(), + op.getScope()); + rewriter.replaceOp(op, newAtomicOp); + return success(); + } +}; + +DiscreteMaskAccessConversionPass::DiscreteMaskAccessConversionPass( + const DiscreteMaskAccessConversionOptions &options) + : DiscreteMaskAccessConversionBase(options) {} + +void DiscreteMaskAccessConversionPass::runOnOperation() { + compileOn91095Flag = this->compileOn91095; + forceSimtTemplateFlag = this->forceSimtTemplate; + + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + patterns.add(patterns.getContext()); + if (failed(applyPatternsAndFoldGreedily(moduleOp, std::move(patterns)))) { + moduleOp->emitError("failed to apply discrete mask access patterns"); + signalPassFailure(); + } +} + +std::unique_ptr> +mlir::triton::createDiscreteMaskAccessConversionPass( + const DiscreteMaskAccessConversionOptions &options) { + return std::make_unique(options); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/CMakeLists.txt new file mode 100755 index 00000000..b7900b2f --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/CMakeLists.txt @@ -0,0 +1,24 @@ +add_triton_library(MemrefCopyToDMAFlagTree + MemrefCopyToDMAFlagTree.cpp + MemrefCopyToDMAFlagTreePass.cpp + + DEPENDS + MemrefCopyToDMAFlagTreeConversionPassIncGen + Registrar + Common + + LINK_LIBS PUBLIC + MLIRSCFTransforms + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonTilingExtIR + TritonStructuredIR +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.cpp b/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.cpp new file mode 100755 index 00000000..34fbb5c4 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.cpp @@ -0,0 +1,190 @@ +#include "flagtree/Common/UnifiedHardware.h" + +#include "triton/Dialect/Triton/IR/Types.h" + +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR//MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" + +#include +#include +#include + +#define DEBUG_TYPE "memref-copy-to-dma-flagtree" + +using namespace mlir; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +namespace { +struct CopyConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + // Get the parameter list of Strides, Sizes and Offsets + SmallVector getValueList(OpBuilder &builder, Location loc, + ArrayRef ofrs) const { + SmallVector values; + for (OpFoldResult ofr : ofrs) { + if (Attribute attr = ofr.dyn_cast()) { + values.push_back(builder.create( + loc, mlir::cast(attr).getInt())); + } else { + values.push_back(ofr.dyn_cast()); + } + } + return values; + } + + // Calculate the total number of DMA handling elements + Value getTotalElementCount(OpBuilder &builder, Location loc, + ArrayRef sizes) const { + assert(!sizes.empty()); + Value total = sizes.front(); + for (size_t i = 1; i < sizes.size(); ++i) { + total = builder.create(loc, total, sizes[i]); + } + return total; + } + + // Check whether the stride is 1 + bool isAllStrideOne(ArrayRef strides) const { + for (OpFoldResult ofr : strides) { + if (auto attr = ofr.dyn_cast()) { + auto intAttr = dyn_cast(attr); + if (!intAttr || intAttr.getInt() != 1) + return false; + continue; + } + + if (auto val = ofr.dyn_cast()) { + if (auto constOp = val.getDefiningOp()) + if (constOp.value() == 1) + continue; + + if (auto intOp = val.getDefiningOp()) + if (intOp.value() == 1) + continue; + } + return false; + } + return true; + } + + LogicalResult rewriteCopyToDma(memref::CopyOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto hardwareManager = mlir::flagtree::createUnifiedHardwareManager(); + auto dmaTag = hardwareManager->getDMATag(); + if (!dmaTag) + return failure(); + + Location loc = op.getLoc(); + Value src = adaptor.getSource(); + Value dst = adaptor.getTarget(); + + Value zero = rewriter.create(loc, 0); + + SmallVector srcIndices, dstIndices; + Value numElements; + Operation *srcDef = src.getDefiningOp(); + // + // Rewriting memref.copy to asynchronous DMA transfer + // + // This transform replaces memref.copy with memref.dma_start + // and memref.dma_wait, There are two cases: Mask and Structured. + // - memref.subview ops with static offsets and sizes, or + // - memref.reinterpret_cast ops that preserve the shape. + // + // Key DMA parameters: + // + // - srcIndices / dstIndices: + // * For memref.subview, use the offset list from getOffsets(). + // E.g.,memref.subview %src[%i, %j][M, N][1, 1] → indices = [%i,%j] + // * For memref.reinterpret_cast, offsets are not meaningful for DMA. + // indices default to all-zero (e.g., [%c0, %c0, ...]). + // + // - numElements: + // * For both cases, compute as the product of the sizes. + // E.g., [M, N] → M * N + // + // - tag: + // * A synchronization buffer of type memref<1xi32, *>. + // The memory space `*` denotes a hardware-reserved region for DMA + // completion signaling, determined by the unified hardware layer. + int64_t rank = mlir::cast(src.getType()).getRank(); + srcIndices.assign(rank, zero); + dstIndices.assign(rank, zero); + + if (auto srcSubview = dyn_cast_or_null(srcDef)) { + auto dstSubview = dyn_cast(dst.getDefiningOp()); + auto sizes = getValueList(rewriter, loc, srcSubview.getMixedSizes()); + numElements = getTotalElementCount(rewriter, loc, sizes); + } else if (auto castOp = dyn_cast(srcDef)) { + auto sizes = getValueList(rewriter, loc, castOp.getMixedSizes()); + numElements = getTotalElementCount(rewriter, loc, sizes); + } else { + return failure(); + } + + Type i32Type = rewriter.getIntegerType(32); + Attribute tagMemSpace = IntegerAttr::get(i32Type, dmaTag); + + MemRefType tagType = MemRefType::get({1}, i32Type, nullptr, tagMemSpace); + Value tag = rewriter.create(loc, tagType); + SmallVector tagIndices = {zero}; + + rewriter.create(loc, src, srcIndices, dst, dstIndices, + numElements, tag, tagIndices); + + rewriter.create(loc, tag, tagIndices, numElements); + + rewriter.eraseOp(op); + return success(); + } + + // StructuredToMemrefPass will generate the memref.copy operation, which + // can be selectively converted to DMA operation later + LogicalResult + matchAndRewrite(memref::CopyOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + auto strAttr = op->getAttrOfType("flagtree_hints"); + if (!strAttr || strAttr.getValue() != "dma") { + return success(); + } + + bool isStrideOne = false; + Value src = adaptor.getSource(); + Operation *srcDef = src.getDefiningOp(); + if (auto srcSubview = dyn_cast_or_null(srcDef)) { + isStrideOne = isAllStrideOne(srcSubview.getMixedStrides()); + } else if (auto castOp = dyn_cast(srcDef)) { + isStrideOne = isAllStrideOne(castOp.getMixedStrides()); + } + + // Skip stride is not 1 + if (isStrideOne) { + return rewriteCopyToDma(op, adaptor, rewriter); + } + + return success(); + } +}; + +} // namespace + +void mlir::triton::populateMemrefCopyToDMAFlagTreeConversionPatterns( + RewritePatternSet &patterns, TypeConverter &typeConverter) { + patterns.add(patterns.getContext()); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTreePass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTreePass.cpp new file mode 100755 index 00000000..cfd926ee --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTreePass.cpp @@ -0,0 +1,145 @@ +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "triton-shared/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/Support/LogicalResult.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/Pass/PassManager.h" +#include "triton/Dialect/Triton/IR/Types.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/Support/Casting.h" + +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include + +#define DEBUG_TYPE "structured-to-memref-flagtree" + +using namespace mlir; +using namespace triton; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_MEMREFCOPYTODMAFLAGTREE +#include "triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class PtrToUnrankedMemrefConverter : public TypeConverter { +public: + PtrToUnrankedMemrefConverter() { + addConversion([](Type type) { return type; }); + addConversion([](triton::PointerType ptrType) { + return UnrankedMemRefType::get(ptrType.getPointeeType(), 0); + }); + addTargetMaterialization([&](OpBuilder &builder, + UnrankedMemRefType resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + } +}; + +class MemrefCopyToDMAFlagTreePass + : public triton::impl::MemrefCopyToDMAFlagTreeBase< + MemrefCopyToDMAFlagTreePass> { + using MemrefCopyToDMAFlagTreeBase< + MemrefCopyToDMAFlagTreePass>::MemrefCopyToDMAFlagTreeBase; + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + // Rewrite on the copied module. Ignore if there is an failure; otherwise, + // replace the original + void runOnOperation() override { + auto module = getOperation(); + ConversionTarget target(getContext()); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, ttx::TritonTilingExtDialect, + memref::MemRefDialect>(); + + target.addIllegalOp(); + + target.addLegalOp(); + + target.addIllegalOp(); + + PtrToUnrankedMemrefConverter typeConverter; + + SmallVector> replacements; + module.walk([&](mlir::func::FuncOp funcOp) { + mlir::func::FuncOp cloned = funcOp.clone(); + cloned->setAttrs(funcOp->getAttrs()); + cloned.setPublic(); + + RewritePatternSet localPatterns(&getContext()); + { + PtrToUnrankedMemrefConverter typeConverter; + triton::populateMemrefCopyToDMAFlagTreeConversionPatterns( + localPatterns, typeConverter); + } + + // Try to apply partial conversion on the cloned operation. + // If it fails, erase the cloned op and return. + if (failed(applyPartialConversion(cloned.getOperation(), target, + std::move(localPatterns)))) { + cloned.erase(); + return; + } + + replacements.emplace_back(funcOp, cloned); + }); + + // Replace original functions with their cloned versions if symbol + // replacement succeeds. + IRRewriter rewriter(&getContext()); + for (auto &p : replacements) { + mlir::func::FuncOp original = p.first; + mlir::func::FuncOp cloned = p.second; + + rewriter.setInsertionPointAfter(original); + rewriter.insert(cloned.getOperation()); + + if (failed(SymbolTable::replaceAllSymbolUses(original.getNameAttr(), + cloned.getNameAttr(), + module.getOperation()))) { + cloned.erase(); + original.emitError("failed to replace symbol uses, keeping original"); + continue; + } + + rewriter.eraseOp(original); + } + } +}; +} // namespace + +std::unique_ptr> +triton::createMemrefCopyToDMAFlagTreePass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/CMakeLists.txt new file mode 100755 index 00000000..262d2b72 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/CMakeLists.txt @@ -0,0 +1,24 @@ +add_triton_library(NoBufferizeFlagTree + NoBufferizeFlagTree.cpp + NoBufferizeFlagTreePass.cpp + + DEPENDS + NoBufferizeFlagTreeConversionPassIncGen + Registrar + Common + + LINK_LIBS PUBLIC + MLIRSCFTransforms + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonTilingExtIR + TritonStructuredIR +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.cpp b/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.cpp new file mode 100755 index 00000000..227ee2ff --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.cpp @@ -0,0 +1,65 @@ +#include "flagtree/Common/UnifiedHardware.h" + +#include "triton/Dialect/Triton/IR/Types.h" + +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" + +#include +#include +#include +#include + +#define DEBUG_TYPE "no-bufferize-flagtree" + +using namespace mlir; +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/NoBufferize_FlagTree/Passes.h.inc" +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +namespace { +struct NoBufferizeConverter : public RewritePattern { + NoBufferizeConverter(MLIRContext *ctx) + : RewritePattern(MatchAnyOpTypeTag(), /*benefit=*/1, ctx) {} + + LogicalResult matchAndRewrite(Operation *op, + PatternRewriter &rewriter) const override { + auto memrefNoBufferize = [](Type t) -> bool { + Attribute msAttr; + if (auto memTy = dyn_cast(t)) + msAttr = memTy.getMemorySpace(); + else if (auto unranked = dyn_cast(t)) + msAttr = unranked.getMemorySpace(); + else + return false; + if (auto intAttr = dyn_cast_or_null(msAttr)) + return intAttr.getValue().getZExtValue() == 8; + return false; + }; + + bool noBuf = llvm::any_of(op->getResultTypes(), memrefNoBufferize); + if (!noBuf) + return failure(); + + op->setAttr("no_bufferize", BoolAttr::get(op->getContext(), true)); + return success(); + } +}; +} // namespace + +void mlir::triton::populateNoBufferizeFlagTreeConversionPatterns( + RewritePatternSet &patterns, TypeConverter &typeConverter) { + patterns.add(patterns.getContext()); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTreePass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTreePass.cpp new file mode 100755 index 00000000..944fffa2 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTreePass.cpp @@ -0,0 +1,64 @@ +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "triton-shared/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/Pass/PassManager.h" +#include "triton/Dialect/Triton/IR/Types.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/Support/Casting.h" +#include +#include + +#define DEBUG_TYPE "no-bufferize-flagtree" +using namespace mlir; +using namespace triton; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_NOBUFFERIZEFLAGTREE +#include "triton-shared/Conversion/NoBufferize_FlagTree/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class NoBufferizeFlagTreePass + : public triton::impl::NoBufferizeFlagTreeBase { + using NoBufferizeFlagTreeBase< + NoBufferizeFlagTreePass>::NoBufferizeFlagTreeBase; + +public: + void runOnOperation() override { + ModuleOp module = getOperation(); + RewritePatternSet localPatterns(&getContext()); + { + TypeConverter typeConverter; + triton::populateNoBufferizeFlagTreeConversionPatterns(localPatterns, + typeConverter); + } + (void)applyPatternsGreedily(module, std::move(localPatterns)); + } +}; +} // namespace + +std::unique_ptr> +triton::createNoBufferizeFlagTreePass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/ReconcilePtrCasts/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/ReconcilePtrCasts/CMakeLists.txt new file mode 100755 index 00000000..3777579f --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/ReconcilePtrCasts/CMakeLists.txt @@ -0,0 +1,18 @@ +add_triton_library(ReconcilePtrCasts + ReconcilePtrCastsPass.cpp + + DEPENDS + ReconcilePtrCastsPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + MLIRReconcileUnrealizedCasts + TritonIR +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/ReconcilePtrCasts/ReconcilePtrCastsPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/ReconcilePtrCasts/ReconcilePtrCastsPass.cpp new file mode 100755 index 00000000..661f65b9 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/ReconcilePtrCasts/ReconcilePtrCastsPass.cpp @@ -0,0 +1,167 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// +// Throughout the conversion process, we convert !tt.ptr -> {!ptr.ptr or memref<*>}. +// This process leaves around unrealized_conversion_cast ops between these types. +// We want to remove these unrealized casts and use the proper conversion ops +// in the PtrDialect: to_memref or from_memref. +// To do this, we use a pattern that simplifies the chain of conversions by +// removing intermediate conversion cast ops. At the end, we are left with just +// pointer to memref or vice versa. We then convert the unrealized cast to +// to_memref or from_memref accordingly. +//===----------------------------------------------------------------------===// + +#include "triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h" +#include "mlir/Conversion/ReconcileUnrealizedCasts/ReconcileUnrealizedCasts.h" +#include "mlir/Dialect/Ptr/IR/PtrTypes.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinDialect.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" + +#include "triton-shared/Dialect/TPtr/IR/TPtrDialect.h" + +#include "mlir/Dialect/MemRef/IR/MemRef.h" + +#include "mlir/Pass/PassManager.h" +#include "triton/Dialect/Triton/IR/Types.h" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/ReconcilePtrCasts/Passes.h.inc" + +namespace { + +static bool isOneToOneCast(UnrealizedConversionCastOp op) { + return (op.getInputs().size() == 1 && op->getNumResults() == 1); +} + +struct SimplifyUnrealizedCast + : public OpRewritePattern { + SimplifyUnrealizedCast(MLIRContext *context, PatternBenefit benefit = 1) + : OpRewritePattern(context, benefit) {} + + LogicalResult matchAndRewrite(UnrealizedConversionCastOp op, + PatternRewriter &rewriter) const override { + if (!isOneToOneCast(op)) { + return failure(); + } + auto in = op.getInputs().front(); + + if (auto unrealizedCast = in.getDefiningOp()) { + if (!isOneToOneCast(unrealizedCast)) { + return failure(); + } + + auto prevInput = unrealizedCast.getInputs().front(); + auto newCast = rewriter.create( + op->getLoc(), op->getResultTypes(), ValueRange{prevInput}); + + rewriter.replaceOp(op, newCast); + return success(); + } + return failure(); + } +}; + +struct FromMemrefConverter + : public OpRewritePattern { + FromMemrefConverter(MLIRContext *context, PatternBenefit benefit = 1) + : OpRewritePattern(context, benefit) {} + + LogicalResult matchAndRewrite(UnrealizedConversionCastOp op, + PatternRewriter &rewriter) const override { + if (!isOneToOneCast(op)) { + return failure(); + } + + auto input = op.getInputs().front(); + auto unrankedInput = dyn_cast(input.getType()); + auto output = op.getResult(0); + auto outType = output.getType(); + + if (unrankedInput && isa(outType)) { + // from_memref only takes ranked memref, cast the unranked memref to + // ranked memref first. + auto rankedMemref = rewriter.create( + op.getLoc(), MemRefType::get({}, unrankedInput.getElementType()), + input); + auto memrefToPtr = rewriter.create( + op->getLoc(), ptr::PtrType::get(rewriter.getContext()), rankedMemref); + + rewriter.replaceAllUsesWith(output, memrefToPtr); + rewriter.eraseOp(op); + + return success(); + } + + return failure(); + } +}; + +struct ToMemrefConverter : public OpRewritePattern { + ToMemrefConverter(MLIRContext *context, PatternBenefit benefit = 1) + : OpRewritePattern(context, benefit) {} + + LogicalResult matchAndRewrite(UnrealizedConversionCastOp op, + PatternRewriter &rewriter) const override { + if (!isOneToOneCast(op)) { + return failure(); + } + auto input = op.getInputs().front(); + auto inType = input.getType(); + auto output = op.getResult(0); + auto outUnrankedMemrefType = dyn_cast(output.getType()); + if (isa(inType) && + outUnrankedMemrefType) { + // to_memref can only cast to ranked static shape memref, we have to cast + // the resulting memref back to unranked + auto elemType = outUnrankedMemrefType.getElementType(); + auto ptrToMemref = rewriter.create( + op->getLoc(), MemRefType::get({1}, elemType), input); + + auto newUnrankedMemref = rewriter.create( + op.getLoc(), MemRefType::get({ShapedType::kDynamic}, elemType), + ptrToMemref); + + rewriter.replaceAllUsesWith(output, newUnrankedMemref); + rewriter.eraseOp(op); + return success(); + } + + return failure(); + } +}; + +class ReconcilePtrCastsPass + : public ReconcilePtrCastsBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + RewritePatternSet patterns(&getContext()); + patterns + .add( + &getContext()); + if (failed(applyPatternsGreedily(moduleOp, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> +triton::createReconcilePtrCastsPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/CMakeLists.txt new file mode 100755 index 00000000..0883bca4 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/CMakeLists.txt @@ -0,0 +1,22 @@ +add_triton_library(StructuredToMemref + StructuredToMemref.cpp + StructuredToMemrefPass.cpp + + DEPENDS + StructuredToMemrefConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRSCFTransforms + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonTilingExtIR + TritonStructuredIR +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/StructuredToMemref.cpp b/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/StructuredToMemref.cpp new file mode 100755 index 00000000..21c8cee6 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/StructuredToMemref.cpp @@ -0,0 +1,903 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton/Dialect/Triton/IR/Types.h" + +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" + +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/TypeUtilities.h" +#include "mlir/IR/Types.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR//MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" + +#include +#include +#include + +#define DEBUG_TYPE "structured-to-memref" + +using namespace mlir; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +static const std::string WRAP_SIDE_BY_SIDE = "wrap_side_by_side"; +static const std::string WRAP_STACKED = "wrap_stacked"; + +static memref::SubViewOp getSubview(int rank, ArrayRef dims, + Value source, Location loc, OpBuilder &b) { + auto sourceType = cast(source.getType()); + SmallVector offsets(rank, b.getIndexAttr(0)); + SmallVector strides(rank, b.getIndexAttr(1)); + auto dstType = + memref::SubViewOp::inferResultType(sourceType, offsets, dims, strides); + + return b.create(loc, cast(dstType), source, + offsets, dims, strides); +} + +namespace { + +struct MakeTensorPtrConverter + : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + static Type getElementTypeStructuredPtr(tts::MakeTensorPtrOp op) { + assert(!op.isBlockPtr()); + // tensor<1024x!tt.ptr> + auto ptrType = cast( + cast(op.getType()).getElementType()); + return ptrType.getPointeeType(); + } + + static Type getElementTypeBlockPtr(tts::MakeTensorPtrOp op) { + assert(op.isBlockPtr()); + // !tt.ptr, 1> + auto shapedType = cast( + cast(op.getType()).getPointeeType()); + return shapedType.getElementType(); + } + + static MemRefType getResultMemrefType(tts::MakeTensorPtrOp op, int64_t offset, + ArrayRef staticStrides, + ArrayRef resultShape) { + auto layout = + StridedLayoutAttr::get(op.getContext(), offset, staticStrides); + Type elemType; + if (op.isBlockPtr()) { + elemType = getElementTypeBlockPtr(op); + } else { + elemType = getElementTypeStructuredPtr(op); + } + return MemRefType::get(resultShape, elemType, layout); + } + + // If there are dimensions with size 1 and stride 0, replace 0 stride with + // the product of sizes of all lower dimensions. This avoids creating memref + // with zero stride. + static llvm::SmallVector + getMixedStridesForMemref(tts::MakeTensorPtrOp op, OpBuilder &b) { + llvm::SmallVector strides; + auto accumulate = 1; + for (auto [size, stride] : + llvm::reverse(llvm::zip(op.getSizes(), op.getMixedStrides()))) { + auto strideIntAttr = getIntAttr(stride); + if (size == 1 && strideIntAttr && strideIntAttr.value() == 0) { + strides.push_back(b.getIndexAttr(accumulate)); + } else if (auto v = llvm::dyn_cast_if_present(stride)) { + OpFoldResult result = getAsOpFoldResult(v); + strides.push_back(result); + } else { + strides.push_back(stride); + } + accumulate *= size; + } + std::reverse(strides.begin(), strides.end()); + return strides; + } + + static OpFoldResult accumulateTargetOffset(tts::MakeTensorPtrOp op, + OpBuilder &b) { + Location loc = op->getLoc(); + OpFoldResult targetOffset = b.getIndexAttr(0); + for (auto o : op.getMixedOffsets()) { + targetOffset = addOFRs(targetOffset, o, loc, b); + } + return targetOffset; + } + + std::pair + createSideBySideCastOps(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op->getLoc(); + auto resultShape = cast(op.getType()).getShape(); + + auto targetOffset = + ofrToIndexValue(accumulateTargetOffset(op, rewriter), loc, rewriter); + + //////////////////////////////////////////////////////////////////////////// + // + // Handling side-by-side wraparound + // + // Note: We do not support cases where the target has already overflown the + // number of columns! This is because in PtrAnalysis, the offset has already + // been collapsed into a single dimension, so it is ambiguous to determine + // whether the offset actually overflows or just refers to an element on the + // subsequent rows. + // + // Same limitations apply to the stacked wraparound case. + // + //////////////////////////////////////////////////////////////////////////// + // + // nextOffset - targetOffset = colSize + // d1 + d2 = colSize + // N + // x clampedOffset + // --------------------------*----------------*-----* + // | | nextOffset (might + // | targetOffset | overflow) + // y *----- *----------------| + // | | | | + // M |----- -----------------| + // | d2 d1 | + // -------------------------------------------- + // + // x = targetOffset % N + // nextOffset = x + colSize + // clampedOffset = min(nextOffset, N) + // d1 = clampedOffset - x + // + //////////////////////////////////////////////////////////////////////////// + + auto resultType = getResultMemrefType( + op, /* offset */ ShapedType::kDynamic, + /* staticStrides */ + SmallVector(resultShape.size(), ShapedType::kDynamic), + /* result shape */ + SmallVector{ + + // Row stays the same, but mlir doesn't allow this anymore. Put + // dynamic. + ShapedType::kDynamic, + + // Column is dynamic, in most cases, this + // should be the same as the original column. + // The last chunk may be smaller due to + // wrapping around. + ShapedType::kDynamic}); + + Value rowSize = rewriter.create( + loc, rewriter.getIndexAttr(op.getSizes()[0])); + Value colSize = rewriter.create( + loc, rewriter.getIndexAttr(op.getSizes()[1])); + + Value modN = ofrToIndexValue(op.getMixedShape()[1], loc, rewriter); + + Value x = rewriter.create(loc, targetOffset, modN); + Value y = rewriter.create(loc, targetOffset, x); + + SmallVector strideVals = + ofrsToIndexValues(op.getMixedStrides(), loc, rewriter); + + // First chunk + Value nextOffset = rewriter.create(loc, x, colSize); + Value clampedOffset = + rewriter.create(loc, nextOffset, modN); + Value d1 = rewriter.create(loc, clampedOffset, x); + SmallVector sizes1{rowSize, d1}; + + auto cast1 = rewriter.create( + loc, resultType, adaptor.getBase(), targetOffset, sizes1, strideVals); + + // Second chunk + Value d2 = rewriter.create(loc, colSize, d1); + SmallVector sizes2{rowSize, d2}; + + auto cast2 = rewriter.create( + loc, resultType, adaptor.getBase(), y, sizes2, strideVals); + + return {cast1, cast2}; + } + + std::pair + createStackedCastOps(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + + auto loc = op->getLoc(); + auto resultShape = cast(op.getType()).getShape(); + + assert(resultShape.size() == 2); + + auto targetOffset = + ofrToIndexValue(accumulateTargetOffset(op, rewriter), loc, rewriter); + + //////////////////////////////////////////////////////////////////////////// + // + // Handling stacked wraparound + // + // We do not support cases where the target offset has already overflown the + // number of rows. See side-by-side wraparound for details. + // + //////////////////////////////////////////////////////////////////////////// + // We're loading a tensor of dim (rowSize, colSize) + // d1 + d2 = rowSize + // d2 is the number of rows that overflow + // + // cols + // + // wrappedAroundOff + // --------------*------------*-------- + // | d2 | | | + // | |------------| | + // rows| | + // | | + // | targetOffset | + // | *------------| | + // | | | | + // | d1 | | | + // | | clampedOff | | + // --------------*--------------------- + // | overflow | + // *------------- + // nextOff + // + // wrappedAroundOff = targetOffset % cols + // clampedOff = (rows * strideRows) + wrappedAroundOff + // ~~~~~~~~~~~~~~~~~ + // ^ + // | + // We have already computed + // rows * strideRows = modRow = shape[1] + // in TritonToStructured + // + // clampedOff - targetOffset + // d1 = -------------------- + // strideRows + + auto resultType = getResultMemrefType( + op, /* offset */ ShapedType::kDynamic, + /* staticStrides */ + SmallVector(resultShape.size(), ShapedType::kDynamic), + /* result shape */ + SmallVector{ + // Row is dynamic, in most cases, this should + // be the same as the original row. The last + // chunk may be smaller due to wrapping + // around. + ShapedType::kDynamic, + + // Col stays the same, which is resultShape[1], but mlir doesn't + // allow this anymore. So we put dynamic instead. + ShapedType::kDynamic}); + + Value rowSize = rewriter.create( + loc, rewriter.getIndexAttr(op.getSizes()[0])); + Value colSize = rewriter.create( + loc, rewriter.getIndexAttr(op.getSizes()[1])); + + Value strideRow = ofrToIndexValue(op.getMixedStrides()[0], loc, rewriter); + Value strideCol = ofrToIndexValue(op.getMixedStrides()[1], loc, rewriter); + + Value modRow = op.getShape()[0]; + + // First chunk + Value wrappedAroundOff = + rewriter.create(loc, targetOffset, strideRow); + Value clampedOff = + rewriter.create(loc, modRow, wrappedAroundOff); + Value d1 = rewriter.create(loc, clampedOff, targetOffset); + d1 = rewriter.create(loc, d1, strideRow); + + SmallVector sizes1{d1, colSize}; + memref::ReinterpretCastOp cast1 = + rewriter.create( + loc, resultType, adaptor.getBase(), targetOffset, sizes1, + ValueRange{strideRow, strideCol}); + + // Second chunk + Value d2 = rewriter.create(loc, rowSize, d1); + SmallVector sizes2{d2, colSize}; + memref::ReinterpretCastOp cast2 = + rewriter.create( + loc, resultType, adaptor.getBase(), wrappedAroundOff, sizes2, + ValueRange{strideRow, strideCol}); + + return {cast1, cast2}; + } + + LogicalResult rewriteSplitPtr(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + + auto parentShape = op.getStaticShape(); + + SmallVector casts; + StringRef wrapType; + + if (parentShape[0] == ShapedType::kDynamic) { + // Stacked case + assert(parentShape[1] == 0); + auto [cast1, cast2] = createStackedCastOps(op, adaptor, rewriter); + casts = {cast1.getResult(), cast2.getResult()}; + wrapType = WRAP_STACKED; + } else { + assert(parentShape[0] == 0); + auto [cast1, cast2] = createSideBySideCastOps(op, adaptor, rewriter); + casts = {cast1.getResult(), cast2.getResult()}; + wrapType = WRAP_SIDE_BY_SIDE; + } + + auto combinedCast = rewriter.create( + op.getLoc(), op.getType(), casts); + + combinedCast->setAttr(wrapType, rewriter.getUnitAttr()); + + rewriter.replaceOp(op, combinedCast); + + return success(); + } + + LogicalResult rewritePtr(ArrayRef resultShape, bool isBlockPtr, + tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + + auto mixedStrides = getMixedStridesForMemref(op, rewriter); + SmallVector staticStrides; + SmallVector dynamicStrides; + dispatchIndexOpFoldResults(mixedStrides, dynamicStrides, staticStrides); + + auto targetOffset = accumulateTargetOffset(op, rewriter); + auto staticTargetOffset = getIntAttr(targetOffset); + auto resultType = getResultMemrefType( + op, staticTargetOffset.value_or(ShapedType::kDynamic), staticStrides, + resultShape); + + auto castOp = rewriter.create( + op.getLoc(), resultType, adaptor.getBase(), targetOffset, + op.getMixedSizes(), mixedStrides); + + rewriter.replaceOp(op, castOp); + + return success(); + } + + LogicalResult + rewriteStructuredPtr(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + ArrayRef resultShape = cast(op.getType()).getShape(); + return rewritePtr(resultShape, false, op, adaptor, rewriter); + } + + LogicalResult rewriteBlockPtr(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + // Block pointers are basically the same as structured pointers except that + // the return types are !tt.ptr> instead of + // tensor> + ArrayRef resultShape = + cast( + cast(op.getType()).getPointeeType()) + .getShape(); + return rewritePtr(resultShape, true, op, adaptor, rewriter); + } + +public: + MakeTensorPtrConverter(const TypeConverter &typeConverter, + MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(tts::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!llvm::is_sorted(op.getOrder(), std::greater<>())) { + emitError(op.getLoc()) << "non-decreasing dimension order on tensor " + "pointers are not yet supported"; + return failure(); + } + + if (op.isBlockPtr()) { + return rewriteBlockPtr(op, adaptor, rewriter); + } + + if (op.isStructuredPtr()) { + return rewriteStructuredPtr(op, adaptor, rewriter); + } + + if (op.isSplitPtr()) { + return rewriteSplitPtr(op, adaptor, rewriter); + } + + return failure(); + } +}; + +struct LoadConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + void createSideBySideCopies(Value block1, Value block2, Value dst, + Location loc, + ConversionPatternRewriter &rewriter) const { + + auto zero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + + auto one = + rewriter.create(loc, rewriter.getIndexAttr(1)); + + Value block1Row = rewriter.create(loc, block1, 0); + Value block1Col = rewriter.create(loc, block1, 1); + + Value block2Row = rewriter.create(loc, block2, 0); + Value block2Col = rewriter.create(loc, block2, 1); + + auto block1Dst = + rewriter.create(loc, dst, /* offsets */ + ValueRange{zero, zero}, + /* sizes */ + ValueRange{block1Row, block1Col}, + /* strides */ + ValueRange{one, one}); + + auto block2Dst = + rewriter.create(loc, dst, + /* offsets */ + ValueRange{zero, block1Col}, + /* sizes */ + ValueRange{block2Row, block2Col}, + /* strides */ + ValueRange{one, one}); + + rewriter.create(loc, block1, block1Dst); + rewriter.create(loc, block2, block2Dst); + } + + void createStackedCopies(Value block1, Value block2, Value dst, Location loc, + ConversionPatternRewriter &rewriter) const { + + auto zero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + auto one = + rewriter.create(loc, rewriter.getIndexAttr(1)); + + Value block1Row = rewriter.create(loc, block1, 0); + Value block1Col = rewriter.create(loc, block1, 1); + + Value block2Row = rewriter.create(loc, block2, 0); + Value block2Col = rewriter.create(loc, block2, 1); + + auto block1Dst = + rewriter.create(loc, dst, /* offsets */ + ValueRange{zero, zero}, + /* sizes */ + ValueRange{block1Row, block1Col}, + /* strides */ + ValueRange{one, one}); + + auto block2Dst = + rewriter.create(loc, dst, + /* offsets */ + ValueRange{block1Row, zero}, + /* sizes */ + ValueRange{block2Row, block2Col}, + /* strides */ + ValueRange{one, one}); + + rewriter.create(loc, block1, block1Dst); + rewriter.create(loc, block2, block2Dst); + } + + memref::SubViewOp createSubview(Value src, ArrayRef offsets, + ArrayRef sizes, + ArrayRef strides, Location loc, + ConversionPatternRewriter &rewriter) const { + auto srcType = cast(src.getType()); + auto dstType = + memref::SubViewOp::inferResultType(srcType, offsets, sizes, strides); + return rewriter.create(loc, cast(dstType), + src, offsets, sizes, strides); + } + + std::pair + getSideBySideSubviews(ArrayRef dims, Value block1, Value block2, + Location loc, + ConversionPatternRewriter &rewriter) const { + OpFoldResult subviewRowFull = dims[0]; + OpFoldResult subviewColFull = dims[1]; + OpFoldResult col1 = + rewriter.create(loc, block1, 1).getResult(); + OpFoldResult subviewCol1 = minOFRs(col1, subviewColFull, loc, rewriter); + OpFoldResult subviewCol2 = + subOFRs(subviewColFull, subviewCol1, loc, rewriter); + + SmallVector offsets(dims.size(), rewriter.getIndexAttr(0)); + SmallVector strides(dims.size(), rewriter.getIndexAttr(1)); + auto sv1 = createSubview(block1, offsets, {subviewRowFull, subviewCol1}, + strides, loc, rewriter); + auto sv2 = createSubview(block2, offsets, {subviewRowFull, subviewCol2}, + strides, loc, rewriter); + + return {sv1, sv2}; + } + + std::pair + getStackedSubviews(ArrayRef dims, Value block1, Value block2, + const Location loc, + ConversionPatternRewriter &rewriter) const { + OpFoldResult subviewRowFull = dims[0]; + OpFoldResult subviewColFull = dims[1]; + OpFoldResult row1 = + rewriter.create(loc, block1, 0).getResult(); + OpFoldResult subviewRow1 = minOFRs(row1, subviewRowFull, loc, rewriter); + OpFoldResult subviewRow2 = + subOFRs(subviewRowFull, subviewRow1, loc, rewriter); + + SmallVector offsets(dims.size(), rewriter.getIndexAttr(0)); + SmallVector strides(dims.size(), rewriter.getIndexAttr(1)); + auto sv1 = createSubview(block1, offsets, {subviewRow1, subviewColFull}, + strides, loc, rewriter); + auto sv2 = createSubview(block2, offsets, {subviewRow2, subviewColFull}, + strides, loc, rewriter); + return {sv1, sv2}; + } + + // Create a corresponding subview for each TEC and + // specify the space it occupies in the shared memory + memref::SubViewOp + createTensorSubview(tts::LoadOp op, Value alloc, + ConversionPatternRewriter &rewriter) const { + auto loc = op->getLoc(); + auto tensorType = cast(op.getType()); + auto shape = tensorType.getShape(); + auto rank = tensorType.getRank(); + SmallVector tensorShape; + for (int64_t dim : shape) { + tensorShape.push_back(rewriter.getIndexAttr(dim)); + } + SmallVector strides(rank, rewriter.getIndexAttr(1)); + auto func = op->getParentOfType(); + unsigned totalArgs = func.getNumArguments(); + // offsets = (pid % TEC_Number) * block_size + // offstes represents the corresponding initial offset value + // of each TEC in the shared memory + SmallVector modOffsets(rank, rewriter.getIndexAttr(0)); + Value c4 = rewriter.create(loc, 4); + for (unsigned dim = 0; dim < rank; dim++) { + Value pid = func.getArgument(totalArgs - 3 + dim); + Value pidIndex = rewriter.create( + loc, rewriter.getIndexType(), pid); + Value mod = rewriter.create(loc, pidIndex, c4); + Value cShape = rewriter.create(loc, shape[dim]); + modOffsets[dim] = + rewriter.create(loc, mod, cShape).getResult(); + } + + return rewriter.create(loc, alloc, modOffsets, + tensorShape, strides); + } + + LogicalResult + rewriteStructuredLoad(tts::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + assert(!op.hasMask()); + + auto loc = op->getLoc(); + auto ptr = adaptor.getPtr(); + auto other = op.getOther(); + + auto tensorType = cast(op.getType()); + auto elemType = tensorType.getElementType(); + + auto alloc = rewriter.create( + loc, MemRefType::get(tensorType.getShape(), elemType)); + if (op->hasAttr("flagtree_hints")) { + auto hintAttr = dyn_cast(op->getAttr("flagtree_hints")); + if (hintAttr && hintAttr.getValue() == "shared_memory") { + // The size of the allocated shared memory is TEC_NUM times + SmallVector sharedShape(tensorType.getShape().begin(), + tensorType.getShape().end()); + for (auto &dim : sharedShape) { + if (!ShapedType::isDynamic(dim)) + dim *= 4; + } + // TODO: tagMemSpace value 8 is only for aipu backend + auto tagMemSpace = + IntegerAttr::get(IntegerType::get(op.getContext(), 64), 8); + alloc = rewriter.create( + loc, MemRefType::get(sharedShape, elemType, nullptr, tagMemSpace)); + alloc->setAttr("flagtree_hints", hintAttr); + } + } + + // No mask + assert(!other && "other value used in non-masked load"); + + auto ptrDefiningOp = ptr.getDefiningOp(); + if (ptrDefiningOp->hasAttr(WRAP_SIDE_BY_SIDE) || + ptrDefiningOp->hasAttr(WRAP_STACKED)) { + + auto unrealizedCast = cast(ptrDefiningOp); + auto memrefs = unrealizedCast.getOperands(); + assert(memrefs.size() == 2); + auto block1 = memrefs[0]; + auto block2 = memrefs[1]; + + if (unrealizedCast->hasAttr(WRAP_SIDE_BY_SIDE)) { + createSideBySideCopies(block1, block2, alloc, loc, rewriter); + } else if (unrealizedCast->hasAttr(WRAP_STACKED)) { + createStackedCopies(block1, block2, alloc, loc, rewriter); + } else { + llvm_unreachable("unexpected wraparound type"); + } + } else { + if (cast(alloc.getType()).getMemorySpaceAsInt() == 8) { + // tensorSubview represents moving data to the space of the + // corresponding TEC in shared memory and converting it into tensor form + SmallVector strides(tensorType.getRank(), + rewriter.getIndexAttr(1)); + auto tensorSubview = createTensorSubview(op, alloc, rewriter); + tensorSubview->setAttr("flagtree_hints", + rewriter.getStringAttr("shared_memory")); + rewriter.create(loc, ptr, tensorSubview); + Value tensor = rewriter.create( + loc, tensorType, tensorSubview, true /* restrict */, + true /* writable */); + rewriter.replaceOp(op, tensor); + } else { + auto copyOp = rewriter.create(loc, ptr, alloc); + auto strAttr = op->getAttrOfType("flagtree_hints"); + if (strAttr && !strAttr.getValue().empty()) { + copyOp->setAttr("flagtree_hints", strAttr); + } + } + } + + if (cast(alloc.getType()).getMemorySpaceAsInt() != 8) { + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + } + + return success(); + } + + LogicalResult rewriteMaskedLoad(tts::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + assert(op.hasMask()); + + auto loc = op->getLoc(); + auto ptr = adaptor.getPtr(); + + auto tensorType = cast(op.getType()); + auto elemType = tensorType.getElementType(); + + auto alloc = rewriter.create( + loc, MemRefType::get(tensorType.getShape(), elemType)); + if (op->hasAttr("flagtree_hints")) { + auto hintAttr = dyn_cast(op->getAttr("flagtree_hints")); + if (hintAttr && hintAttr.getValue() == "shared_memory") { + // The size of the allocated shared memory is TEC_NUM times + SmallVector sharedShape(tensorType.getShape().begin(), + tensorType.getShape().end()); + for (auto &dim : sharedShape) { + if (!ShapedType::isDynamic(dim)) + dim *= 4; + } + // TODO: tagMemSpace value 8 is only for aipu backend + auto tagMemSpace = + IntegerAttr::get(IntegerType::get(op.getContext(), 64), 8); + alloc = rewriter.create( + loc, MemRefType::get(sharedShape, elemType, nullptr, tagMemSpace)); + alloc->setAttr("flagtree_hints", hintAttr); + } + } + + SmallVector mixedDims = op.getMixedMaskDims(); + + // Fill load destination with other value + if (op.getOther()) { + // For each dimension check if dims[i] < shape[i], or-accumulate + // the result + auto shape = tensorType.getShape(); + auto accBase = + rewriter.create(loc, rewriter.getBoolAttr(false)) + .getResult(); + for (size_t i = 0; i < shape.size(); i++) { + auto shapei = rewriter.create( + loc, rewriter.getIndexAttr(shape[i])); + + Value dimi = dyn_cast(mixedDims[i]); + if (!dimi) { + dimi = rewriter.create( + loc, rewriter.getIndexAttr(op.getStaticMaskDims()[i])); + } + + Value cmp = rewriter.create( + loc, arith::CmpIPredicate::slt, dimi, shapei); + accBase = rewriter.create(loc, accBase, cmp); + } + + // condition the memset on the or-accumulation + // initialize with padding prior to CopyOp + Value fillMem; + // When shared memory is used, only the space occupied + // by the corresponding TEC is filled + if (cast(alloc.getType()).getMemorySpaceAsInt() == 8) { + fillMem = createTensorSubview(op, alloc, rewriter).getResult(); + auto fillDefiningOp = fillMem.getDefiningOp(); + fillDefiningOp->setAttr("flagtree_hints", + rewriter.getStringAttr("shared_memory")); + } else { + fillMem = alloc; + } + rewriter.create(loc, accBase, [&](OpBuilder &b, Location loc) { + b.create(loc, ValueRange{op.getOther()}, + ValueRange{fillMem}); + b.create(loc); + }); + } + + auto ptrDefiningOp = ptr.getDefiningOp(); + if (ptrDefiningOp->hasAttr(WRAP_SIDE_BY_SIDE) || + ptrDefiningOp->hasAttr(WRAP_STACKED)) { + + auto unrealizedCast = cast(ptrDefiningOp); + + auto memrefs = unrealizedCast.getOperands(); + assert(memrefs.size() == 2); + auto block1 = memrefs[0]; + auto block2 = memrefs[1]; + + if (unrealizedCast->hasAttr(WRAP_SIDE_BY_SIDE)) { + auto [subview1, subview2] = + getSideBySideSubviews(mixedDims, block1, block2, loc, rewriter); + createSideBySideCopies(subview1, subview2, alloc, loc, rewriter); + } else if (unrealizedCast->hasAttr(WRAP_STACKED)) { + auto [subview1, subview2] = + getStackedSubviews(mixedDims, block1, block2, loc, rewriter); + createStackedCopies(subview1, subview2, alloc, loc, rewriter); + } else { + llvm_unreachable("unexpected wraparound type"); + } + + rewriter.eraseOp(unrealizedCast); + + } else { + memref::SubViewOp srcSubview = + getSubview(tensorType.getRank(), mixedDims, ptr, loc, rewriter); + if (cast(alloc.getType()).getMemorySpaceAsInt() == 8) { + // The tensorSubview passes bufferization.to_tensor to + // convert memref into tensor form for subsequent computations + SmallVector strides(tensorType.getRank(), + rewriter.getIndexAttr(1)); + auto tensorSubview = createTensorSubview(op, alloc, rewriter); + tensorSubview->setAttr("flagtree_hints", + rewriter.getStringAttr("shared_memory")); + // dstSubview represents moving to the + // specified TEC space of shared memory + auto modOffsets = tensorSubview.getMixedOffsets(); + auto allocType = cast(alloc.getType()); + auto dstType = memref::SubViewOp::inferResultType(allocType, modOffsets, + mixedDims, strides); + auto dstSubview = rewriter.create( + loc, cast(dstType), alloc, modOffsets, mixedDims, + strides); + rewriter.create(loc, srcSubview, dstSubview); + Value tensor = rewriter.create( + loc, tensorType, tensorSubview, true /* restrict */, + true /* writable */); + rewriter.replaceOp(op, tensor); + } else { + memref::SubViewOp dstSubview = + getSubview(tensorType.getRank(), mixedDims, alloc, loc, rewriter); + auto copyOp = + rewriter.create(loc, srcSubview, dstSubview); + auto strAttr = op->getAttrOfType("flagtree_hints"); + if (strAttr && !strAttr.getValue().empty()) { + copyOp->setAttr("flagtree_hints", strAttr); + } + } + } + if (cast(alloc.getType()).getMemorySpaceAsInt() != 8) { + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + } + return success(); + } + +public: + LoadConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(tts::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (op.hasMask()) { + return rewriteMaskedLoad(op, adaptor, rewriter); + } else { + return rewriteStructuredLoad(op, adaptor, rewriter); + } + } +}; + +struct StoreConverter : public OpConversionPattern { +private: + using OpConversionPattern::OpConversionPattern; + + static tensor::ExtractSliceOp + getExtractSlice(int rank, ArrayRef dims, Value source, + const Location loc, OpBuilder &b) { + auto sourceType = cast(source.getType()); + SmallVector offsets(rank, b.getIndexAttr(0)); + SmallVector strides(rank, b.getIndexAttr(1)); + + auto dstType = tensor::ExtractSliceOp::inferResultType(sourceType, offsets, + dims, strides); + + return b.create(loc, dstType, source, offsets, dims, + strides); + } + +public: + StoreConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(tts::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = op.getLoc(); + auto ptr = adaptor.getPtr(); + auto storeValue = op.getValue(); + auto rank = cast(storeValue.getType()).getRank(); + + if (op.hasMask()) { + auto mixedDims = op.getMixedMaskDims(); + + auto srcSlice = + getExtractSlice(rank, mixedDims, storeValue, loc, rewriter); + auto dstSubview = getSubview(rank, mixedDims, ptr, loc, rewriter); + + auto storeOp = rewriter.create( + loc, srcSlice, dstSubview); + storeOp.setWritable(true); + } else { + auto storeOp = rewriter.create( + loc, storeValue, ptr); + storeOp.setWritable(true); + } + + rewriter.eraseOp(op); + return success(); + } +}; + +} // namespace + +void mlir::triton::populateStructuredToMemrefConversionPatterns( + RewritePatternSet &patterns, TypeConverter &typeConverter) { + patterns.add(typeConverter, patterns.getContext()); + patterns.add(patterns.getContext()); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/StructuredToMemrefPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/StructuredToMemrefPass.cpp new file mode 100755 index 00000000..ce5ab690 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/StructuredToMemrefPass.cpp @@ -0,0 +1,183 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/Support/LogicalResult.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/Pass/PassManager.h" +#include "triton/Dialect/Triton/IR/Types.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/Support/Casting.h" + +#include + +#define DEBUG_TYPE "structured-to-memref" + +using namespace mlir; +using namespace triton; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_STRUCTUREDTOMEMREF +#include "triton-shared/Conversion/StructuredToMemref/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class LoopTypeConverter : public TypeConverter { +public: + LoopTypeConverter(MLIRContext *context) { + // The order of type conversion is important: later ones are tried earlier. + addConversion([](Type type) { return type; }); + // addConversion([context](triton::PointerType ptrType) { + // SmallVector strides{1}; + // auto layout = + // StridedLayoutAttr::get(context, ShapedType::kDynamic, strides); + + // auto elemType = ptrType.getPointeeType(); + // auto memrefType = MemRefType::get({1}, elemType, layout); + // return memrefType; + // }); + + // A tensor of pointers can be passed in as scf.for's init-args, in such + // cases, we convert the type to a memref with dynamic offsets and + // strides. + addConversion( + [context](RankedTensorType tensorType) -> std::optional { + if (auto ptrType = llvm::dyn_cast( + tensorType.getElementType())) { + auto layout = StridedLayoutAttr::get( + context, ShapedType::kDynamic, + SmallVector(tensorType.getRank(), + ShapedType::kDynamic)); + Type elemType = ptrType.getPointeeType(); + return MemRefType::get(tensorType.getShape(), elemType, layout); + } + + return std::nullopt; + }); + + // Convert the current memref type to a memref type with dynamic offsets and + // strides through another reinterpret_cast with the same offsets. + // Canonicalization will simplify this sequence by removing the inital + // reinterpret_cast. + addTargetMaterialization([&](OpBuilder &builder, MemRefType memrefType, + ValueRange inputs, + Location loc) -> Value { + auto reinterpretCast = + inputs[0].getDefiningOp(); + if (!reinterpretCast) { + return builder + .create(loc, memrefType, inputs) + .getResult(0); + } + return builder.create( + loc, memrefType, inputs[0], reinterpretCast.getMixedOffsets()[0], + reinterpretCast.getMixedSizes(), reinterpretCast.getMixedStrides()); + }); + + addSourceMaterialization([&](OpBuilder &builder, Type resultType, + ValueRange inputs, + Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + + addArgumentMaterialization([&](OpBuilder &builder, Type resultType, + ValueRange inputs, + Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + } +}; + +class PtrToUnrankedMemrefConverter : public TypeConverter { +public: + PtrToUnrankedMemrefConverter() { + addConversion([](Type type) { return type; }); + addConversion([](triton::PointerType ptrType) { + return UnrankedMemRefType::get(ptrType.getPointeeType(), 0); + }); + addTargetMaterialization([&](OpBuilder &builder, + UnrankedMemRefType resultType, + ValueRange inputs, + Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + } +}; + +class StructuredToMemrefPass + : public triton::impl::StructuredToMemrefBase { + using StructuredToMemrefBase::StructuredToMemrefBase; + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, ttx::TritonTilingExtDialect, + memref::MemRefDialect>(); + + target.addIllegalOp(); + + target.addLegalOp(); + + PtrToUnrankedMemrefConverter typeConverter; + + triton::populateStructuredToMemrefConversionPatterns(patterns, + typeConverter); + + LoopTypeConverter loopTypeConverter(patterns.getContext()); + + mlir::scf::populateSCFStructuralTypeConversionsAndLegality( + loopTypeConverter, patterns, target); + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> +triton::createStructuredToMemrefPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/CMakeLists.txt new file mode 100755 index 00000000..aaa74904 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/CMakeLists.txt @@ -0,0 +1,24 @@ +add_triton_library(TritonArithToLinalg + TritonArithToLinalg.cpp + TritonArithToLinalgPass.cpp + + DEPENDS + TritonArithToLinalgConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRLinalgTransforms + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRMathExtDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonTilingExtIR + TritonStructuredIR + TritonSharedUtils +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp new file mode 100755 index 00000000..a9bad3df --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp @@ -0,0 +1,106 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include +#include + +#define DEBUG_TYPE "triton-arith-to-linalg" +#include "triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.hpp" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" + +void mlir::triton::populateTritonArithToLinalgCanonicalizationPatterns( + RewritePatternSet &patterns) { + patterns.add, MinMaxConverter>( + patterns.getContext()); +} + +void mlir::triton::populateTritonTensorPtrConversionPatterns( + RewritePatternSet &patterns) { + patterns.add, + TensorOpConverter, + TensorOpConverter, + TensorOpConverter>(patterns.getContext()); +} + +void mlir::triton::populateTritonArithToLinalgConversionPatterns( + bool pidsToFuncArgs, bool addptrToLinalg, bool assertToCf, + RewritePatternSet &patterns) { + + if (pidsToFuncArgs) { + patterns.add( + patterns.getContext()); + } + if (addptrToLinalg) { + patterns.add(patterns.getContext()); + } + if (assertToCf) { + patterns.add(patterns.getContext()); + } + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + + populateExternElementwiseOpToMLIROps(patterns); + + // Reduce converters + // Triton's reduce op is idential to linalg.reduce op, so we can clone + // `tt.reduce` body to `linalg.reduce`. Unfortunately, we still need to + // perform pattern matching to know what reduce ops we are dealing with + // so that we know how to initialize the initial reduce values correctly. + // + // We can do this in a generic way without pattern matching by always using + // the first elements along the reduction axis and perform the reduction on + // the remaining elements. However, this results in creatings sub-tensors that + // aren't always multiple of 2s, which are sub-optimal for certain hardwares. + patterns.add(patterns.getContext()); // flagtree + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); // flagtree + + // Note: the ordering here matters! + // These patterns are added last to they will be tried last. + linalg::populateElementwiseToLinalgConversionPatterns(patterns); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/TritonArithToLinalgPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/TritonArithToLinalgPass.cpp new file mode 100755 index 00000000..312d2c0f --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/TritonArithToLinalgPass.cpp @@ -0,0 +1,255 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir-ext/Dialect/MathExt/IR/MathExt.h" +#include "triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" +#include "triton-shared/Utils/Utils.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Tensor/Transforms/Transforms.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" + +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "triton-arith-to-linalg" + +using namespace mlir; +using namespace triton; +using namespace mlir::mathext; + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_TRITONARITHTOLINALG +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h.inc" +} // namespace triton +} // namespace mlir + +namespace { + +class TritonArithToLinalgPass + : public triton::impl::TritonArithToLinalgBase { + using TritonArithToLinalgBase< + TritonArithToLinalgPass>::TritonArithToLinalgBase; + + static auto constexpr LAUNCH_GRID_RANK = getMaxEnumValForProgramIDDim() + 1; + static unsigned int constexpr TRITON_PROGRAM_INFO_ARG_COUNT = + LAUNCH_GRID_RANK * 2; + + // Add additional I32 arguments to represent: + // - num_programs, 3 in total, one for each axis of the launch grid + // - program_id, 3 in total, one for each axis of the launch grid + static void addProgramInfo(triton::FuncOp func) { + OpBuilder b(func); + + auto origFuncType = func.getFunctionType(); + auto origInputTypes = origFuncType.getInputs(); + SmallVector newInputTypes(origInputTypes); + newInputTypes.append(TRITON_PROGRAM_INFO_ARG_COUNT, b.getI32Type()); + + auto newFuncType = + b.getFunctionType(newInputTypes, origFuncType.getResults()); + + func.setFunctionType(newFuncType); + + // Add empty attributes for each new argument if needed + if (func.getAllArgAttrs()) { + SmallVector newArgAttrs; + func.getAllArgAttrs(newArgAttrs); + newArgAttrs.append(TRITON_PROGRAM_INFO_ARG_COUNT, DictionaryAttr()); + func.setAllArgAttrs(newArgAttrs); + } + + // Add the corresponding arguments to function body + for (unsigned int i = 0; i < TRITON_PROGRAM_INFO_ARG_COUNT; i++) { + func.getBody().front().addArgument(b.getI32Type(), func.getLoc()); + } + } + + LogicalResult applyTensorConcatDecomposition() { + auto moduleOp = getOperation(); + MLIRContext *context = &getContext(); + RewritePatternSet patterns(context); + + tensor::populateDecomposeTensorConcatPatterns(patterns); + + if (failed(applyPatternsGreedily(moduleOp, std::move(patterns)))) { + return failure(); + } + return success(); + } + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + { + RewritePatternSet patterns(&getContext()); + populateTritonArithToLinalgCanonicalizationPatterns(patterns); + if (failed(applyPatternsGreedily(moduleOp, std::move(patterns)))) { + signalPassFailure(); + } + } + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + mathext::MathExtDialect, linalg::LinalgDialect, affine::AffineDialect, + scf::SCFDialect, cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, ttx::TritonTilingExtDialect, + tts::TritonStructuredDialect>(); + + target.addLegalOp(); + + target.addLegalOp(); + + target.addDynamicallyLegalDialect< + arith::ArithDialect, math::MathDialect, mathext::MathExtDialect>( + [](Operation *op) { + // Lower dense constant to linalg.fill + if (auto constOp = dyn_cast(op)) { + if (!isa(constOp.getResult().getType())) { + return true; + } + + if (auto denseAttr = + dyn_cast(constOp.getValue())) { + if (denseAttr.isSplat() && + isa(denseAttr.getElementType())) { + return false; + } + } + return true; + } + + bool operateOnTensors = + llvm::all_of(op->getOperandTypes(), [](Type type) { + return isa(type); + }); + + return !operateOnTensors; + }); + + if (pidsToFuncArgs) { + target.addIllegalOp(); + } + + if (addptrToLinalg) { + target.addDynamicallyLegalOp([](triton::AddPtrOp op) { + return !isa(op.getResult().getType()); + }); + } + + target.addDynamicallyLegalOp( + [this](triton::BitcastOp op) { + if (!tensorPtrToLinalg) { + return triton::isPtrTypeLike(op.getType()); + } else { + if (triton::isPtrTypeLike(op.getType())) { + return !isa(op.getType()); + } + return false; + } + }); + + // TODO: Might want to consolidate this flag with addptrToLinalg later. + if (tensorPtrToLinalg) { + target.addDynamicallyLegalOp( + [](auto op) { + return !isa(op->getOperands()[0].getType()); + }); + populateTritonTensorPtrConversionPatterns(patterns); + } + + if (!assertToCf) { + target.addLegalOp(); + } + + triton::populateTritonArithToLinalgConversionPatterns( + pidsToFuncArgs, addptrToLinalg, assertToCf, patterns); + + if (pidsToFuncArgs) { + for (auto func : getOperation().getOps()) { + addProgramInfo(func); + } + } + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + + if (failed(applyTensorConcatDecomposition())) { + signalPassFailure(); + } + + // Convert tt.func and tt.return into func's counterparts + if (ttToFuncFunc) { + moduleOp.walk([&](triton::FuncOp func) { + OpBuilder builder(func); + + auto name = func.getName(); + auto type = func.getFunctionType(); + + SmallVector argAttrs, resAttrs; + func.getAllArgAttrs(argAttrs); + func.getAllResultAttrs(resAttrs); + + auto funcFunc = builder.create(func.getLoc(), name, type); + funcFunc.setAllArgAttrs(argAttrs); + funcFunc.setAllResultAttrs(resAttrs); + + auto &funcFuncBody = funcFunc.getBody(); + auto &funcBody = func.getBody(); + + IRMapping map; + funcBody.cloneInto(&funcFuncBody, map); + + for (Block &block : funcFuncBody.getBlocks()) { + auto term = block.getTerminator(); + // Only convert to func.return if the terminator is a tt.return. + // Otherwise, we will accidentally convert cf.br ops which are also + // considered terminators. + if (isa(term)) { + builder.setInsertionPoint(term); + builder.create(func.getLoc(), term->getOperands()); + term->erase(); + } + } + func.erase(); + }); + } + } +}; + +} // namespace + +std::unique_ptr> +triton::createTritonArithToLinalgPass(bool tensorPtrToLinalg) { + TritonArithToLinalgOptions options; + options.tensorPtrToLinalg = tensorPtrToLinalg; + return std::make_unique(options); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonPtrToMemref/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonPtrToMemref/CMakeLists.txt new file mode 100755 index 00000000..eeef2832 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonPtrToMemref/CMakeLists.txt @@ -0,0 +1,18 @@ +add_triton_library(TritonPtrToMemref + TritonPtrToMemrefPass.cpp + + DEPENDS + TritonPtrToMemrefConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + MLIRReconcileUnrealizedCasts + TritonIR +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonPtrToMemref/TritonPtrToMemrefPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonPtrToMemref/TritonPtrToMemrefPass.cpp new file mode 100755 index 00000000..36263bd7 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonPtrToMemref/TritonPtrToMemrefPass.cpp @@ -0,0 +1,145 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Func/Transforms/FuncConversions.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/IR/Types.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/Passes.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/AnalysisStructured/PtrAnalysis.h" +#include "triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#define DEBUG_TYPE "triton-ptr-to-memref" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonPtrToMemref/Passes.h.inc" + +namespace { + +class TritonFunctionSignatureConverter : public TypeConverter { +public: + TritonFunctionSignatureConverter() { + // The order of type conversion is important: later ones are tried earlier. + addConversion([](Type type) { return type; }); + addConversion([](triton::PointerType ptrType) { + return UnrankedMemRefType::get(ptrType.getPointeeType(), + /*memorySpace=*/0); + }); + addConversion([](RankedTensorType tensorType) -> std::optional { + if (auto ptrType = + dyn_cast(tensorType.getElementType())) { + return MemRefType::get(tensorType.getShape(), ptrType.getPointeeType()); + } + return std::nullopt; + }); + + auto createUnrealizedCast = [&](OpBuilder &builder, Type resultType, + ValueRange inputs, + Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }; + addSourceMaterialization(createUnrealizedCast); + addTargetMaterialization(createUnrealizedCast); + } +}; + +struct BitcastConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(triton::BitcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!isa(op.getSrc().getType())) { + return failure(); + } + + rewriter.replaceAllOpUsesWith(op, adaptor.getSrc()); + return success(); + } +}; + +class TritonPtrToMemrefPass + : public TritonPtrToMemrefBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + TritonFunctionSignatureConverter typeConverter; + + // Update function signature and call ops to use memrefs + target.addDynamicallyLegalOp([&](auto op) { + return typeConverter.isSignatureLegal( + cast(cast(op).getFunctionType())); + }); + + target.addDynamicallyLegalOp([&](func::CallOp op) { + return typeConverter.isLegal(op.getResultTypes()) && + typeConverter.isLegal(op.getOperandTypes()); + }); + + target.addDynamicallyLegalOp( + [&](auto op) { return typeConverter.isLegal(op); }); + target.addLegalOp(); + + populateFunctionOpInterfaceTypeConversionPattern( + patterns, typeConverter); + populateFunctionOpInterfaceTypeConversionPattern( + patterns, typeConverter); + populateCallOpTypeConversionPattern(patterns, typeConverter); + + patterns.add(typeConverter, &getContext()); + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createCanonicalizerPass()); + pm.addPass(createCSEPass()); + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> triton::createTritonPtrToMemrefPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToAnnotation/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonToAnnotation/CMakeLists.txt new file mode 100755 index 00000000..b8f47ced --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToAnnotation/CMakeLists.txt @@ -0,0 +1,15 @@ +add_triton_library(TritonToAnnotation + TritonToAnnotation.cpp + + DEPENDS + TritonToAnnotationConversionPassIncGen + + LINK_LIBS + BiShengIRAnnotationDialect + BiShengIRDialectUtils + MLIRIR + MLIRPass + MLIRTransforms + MLIRSupport + TritonIR +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToAnnotation/TritonToAnnotation.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToAnnotation/TritonToAnnotation.cpp new file mode 100755 index 00000000..c1e10a02 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToAnnotation/TritonToAnnotation.cpp @@ -0,0 +1,78 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToAnnotation/Passes.h" +#if __has_include("bishengir/Dialect/Annotation/IR/Annotation.h") +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#endif +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h" + +namespace mlir { +namespace triton { +#define GEN_PASS_DEF_TRITONTOANNOTATION +#include "incubated/Conversion/TritonToAnnotation/Passes.h.inc" +} // namespace triton +} // namespace mlir + +using namespace mlir; + +namespace { +struct TritonToAnnotationPass + : public mlir::triton::impl::TritonToAnnotationBase< + TritonToAnnotationPass> { + void runOnOperation() override; +}; +} // namespace + +struct TritonAnnotationConversionPattern + : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(mlir::triton::ascend::AnnotationOp op, + PatternRewriter &rewriter) const final { + auto markOp = rewriter.create(op.getLoc(), op.getSrc()); + // Forward all annotations. + markOp->setAttrs(op->getAttrs()); + rewriter.eraseOp(op); + return success(); + } +}; + +void TritonToAnnotationPass::runOnOperation() { + auto module = getOperation(); + ConversionTarget target(getContext()); + target.addLegalDialect(); + + RewritePatternSet patterns(&getContext()); + patterns.add(patterns.getContext()); + if (failed(applyPartialConversion(module, target, std::move(patterns)))) { + signalPassFailure(); + } +} + +std::unique_ptr> +mlir::triton::createTritonToAnnotationPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/CMakeLists.txt new file mode 100755 index 00000000..1511cb67 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/CMakeLists.txt @@ -0,0 +1,23 @@ +add_triton_library(TritonToLinalg + TritonToLinalg.cpp + TritonToLinalgPass.cpp + + DEPENDS + TritonToLinalgConversionPassIncGen + MLIRMathExtDialectIncGen + MLIRMathExtOpsIncGen + + LINK_LIBS PUBLIC + TritonTilingExtIR + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonSharedAnalysis +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/TritonToLinalg.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/TritonToLinalg.cpp new file mode 100755 index 00000000..879a1764 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/TritonToLinalg.cpp @@ -0,0 +1,96 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include "triton-shared/Conversion/TritonToLinalg/TritonToLinalg.h" + +#define DEBUG_TYPE "triton-to-linalg" +#include "triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.hpp" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonToLinalg/Passes.h.inc" + +void mlir::triton::populateTritonToLinalgCanonicalizationPatterns( + RewritePatternSet &patterns) { + patterns.add, MinMaxConverter>( + patterns.getContext()); +} + +void mlir::triton::populateTritonToLinalgConversionPatterns( + TypeConverter &typeConverter, RewritePatternSet &patterns, + unsigned int launchGridRank) { + populateFunctionOpInterfaceTypeConversionPattern( + patterns, typeConverter); + populateFunctionOpInterfaceTypeConversionPattern( + patterns, typeConverter); + + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + + populateExternElementwiseOpToMLIROps(patterns); + + // Reduce converters + // Triton's reduce op is idential to linalg.reduce op, so we can clone + // `tt.reduce` body to `linalg.reduce`. Unfortunately, we still need to + // perform pattern matching to know what reduce ops we are dealing with + // so that we know how to initialize the initial reduce values correctly. + // + // We can do this in a generic way without pattern matching by always using + // the first elements along the reduction axis and perform the reduction on + // the remaining elements. However, this results in creatings sub-tensors that + // aren't always multiple of 2s, which are sub-optimal for certain hardwares. + patterns.add(patterns.getContext()); // flagtree + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); // flagtree + + // Note: the ordering here matters! + // MetaOpConverter has PatternBenefit == 10 which should take precedence over + // these linalg patterns, but to be safe, add these patterns last so that they + // will be tried last. Incorrect ordering or having MetaOpConverter has lower + // PatternBenefit will result in element-wise meta ops being converted to + // linalg.generic ops. + linalg::populateElementwiseToLinalgConversionPatterns(patterns); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/TritonToLinalgPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/TritonToLinalgPass.cpp new file mode 100755 index 00000000..25b7db85 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/TritonToLinalgPass.cpp @@ -0,0 +1,229 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton-shared/Analysis/UseAnalysis.h" +#include "triton-shared/Conversion/TritonToLinalg/TritonToLinalg.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "mlir/Transforms/Passes.h" + +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "triton-to-linalg" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonToLinalg/Passes.h.inc" + +namespace { + +class TritonTypeConverter : public TypeConverter { +public: + TritonTypeConverter() { + // The order of type conversion is important: later ones are tried earlier. + addConversion([](Type type) { return type; }); + addConversion([](triton::PointerType ptrType) { + return UnrankedMemRefType::get(ptrType.getPointeeType(), 0); + }); + addConversion([](TensorType tensorType) -> Type { + auto elemType = tensorType.getElementType(); + if (auto ptrType = dyn_cast(elemType)) { + elemType = ptrType.getPointeeType(); + } + return MemRefType::get(tensorType.getShape(), elemType); + }); + } +}; + +class TritonToLinalgPass : public TritonToLinalgBase { + + static auto constexpr LAUNCH_GRID_RANK = getMaxEnumValForProgramIDDim() + 1; + static unsigned int constexpr TRITON_PROGRAM_INFO_ARG_COUNT = + LAUNCH_GRID_RANK * 2; + + // Add additional I32 arguments to represent: + // - num_programs, 3 in total, one for each axis of the launch grid + // - program_id, 3 in total, one for each axis of the launch grid + static void addProgramInfo(triton::FuncOp func) { + OpBuilder b(func); + + auto origFuncType = func.getFunctionType(); + auto origInputTypes = origFuncType.getInputs(); + SmallVector newInputTypes(origInputTypes); + newInputTypes.append(TRITON_PROGRAM_INFO_ARG_COUNT, b.getI32Type()); + + auto newFuncType = + b.getFunctionType(newInputTypes, origFuncType.getResults()); + + func.setFunctionType(newFuncType); + + // Add empty attributes for each new argument if needed + if (func.getAllArgAttrs()) { + SmallVector newArgAttrs; + func.getAllArgAttrs(newArgAttrs); + newArgAttrs.append(TRITON_PROGRAM_INFO_ARG_COUNT, DictionaryAttr()); + func.setAllArgAttrs(newArgAttrs); + } + + // Add the corresponding arguments to function body + for (unsigned int i = 0; i < TRITON_PROGRAM_INFO_ARG_COUNT; i++) { + func.getBody().front().addArgument(b.getI32Type(), func.getLoc()); + } + } + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + { + RewritePatternSet patterns(&getContext()); + populateTritonToLinalgCanonicalizationPatterns(patterns); + if (failed(applyPatternsGreedily(moduleOp, std::move(patterns)))) { + signalPassFailure(); + } + } + + moduleOp.walk([this](triton::FuncOp op) { + if (failed(runUseAnalysis(op))) { + signalPassFailure(); + } + }); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + TritonTypeConverter tritonTypeConverter; + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, memref::MemRefDialect, + ttx::TritonTilingExtDialect>(); + + target.addLegalOp(); + + // Update function signature to use memrefs + target.addDynamicallyLegalOp([&](triton::FuncOp op) { + return tritonTypeConverter.isSignatureLegal(op.getFunctionType()); + }); + + // Lower dense constant to linalg.fill + target.addDynamicallyLegalOp([](arith::ConstantOp op) { + if (!isa(op.getResult().getType())) { + return true; + } + + if (auto denseAttr = dyn_cast(op.getValue())) { + if (denseAttr.isSplat() && + isa(denseAttr.getElementType())) { + return false; + } + } + return true; + }); + + target.addDynamicallyLegalOp([](Operation *op) { + return llvm::all_of(op->getOperandTypes(), [](Type t) { + if (isa(t)) { + return false; + } + if (auto shapedType = dyn_cast(t)) { + return shapedType.getElementType().isIntOrFloat(); + } + assert(t.isIntOrIndexOrFloat()); + return true; + }); + }); + + target.addDynamicallyLegalDialect( + [](Operation *op) { + if (op->hasAttr("MetaUse")) { + return false; + } + + if (isa(op)) { + return true; + } + + bool operateOnTensors = + llvm::all_of(op->getOperandTypes(), [](Type type) { + return isa(type); + }); + + return !operateOnTensors; + }); + + triton::populateTritonToLinalgConversionPatterns( + tritonTypeConverter, patterns, LAUNCH_GRID_RANK); + + for (auto func : getOperation().getOps()) + addProgramInfo(func); + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) + signalPassFailure(); + + // Convert tt.func and tt.return into func's counterparts + moduleOp.walk([&](triton::FuncOp func) { + OpBuilder builder(func); + + auto name = func.getName(); + auto type = func.getFunctionType(); + + SmallVector argAttrs, resAttrs; + func.getAllArgAttrs(argAttrs); + func.getAllResultAttrs(resAttrs); + + auto funcFunc = builder.create(func.getLoc(), name, type); + funcFunc.setAllArgAttrs(argAttrs); + funcFunc.setAllResultAttrs(resAttrs); + + auto &funcFuncBody = funcFunc.getBody(); + auto &funcBody = func.getBody(); + + IRMapping map; + funcBody.cloneInto(&funcFuncBody, map); + + for (Block &block : funcFuncBody.getBlocks()) { + auto term = block.getTerminator(); + builder.setInsertionPoint(term); + builder.create(func.getLoc(), term->getOperands()); + term->erase(); + } + func.erase(); + }); + + // Erase dead code and fold constants created during lowering + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createCanonicalizerPass()); + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> triton::createTritonToLinalgPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/CMakeLists.txt new file mode 100755 index 00000000..b9d43113 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/CMakeLists.txt @@ -0,0 +1,40 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Triton Project Contributors. +# +#===------------------------------------------------------------------------===# + +add_triton_library(TritonToLinalgExperimental + TritonToLinalgExperimentalPass.cpp + TritonToPtrPass.cpp + + DEPENDS + TritonToLinalgExperimentalConversionPassIncGen + + LINK_LIBS PUBLIC + TritonTilingExtIR + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRMathExtDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TPtrIR + TritonIR + TritonTransforms + TritonSharedAnalysis + TritonSharedUtils + + TritonArithToLinalg + StructuredToMemref + NoBufferizeFlagTree + MemrefCopyToDMAFlagTree + WaferTritonToStructured + TritonToUnstructured + TritonPtrToMemref + UnstructuredToMemref + ReconcilePtrCasts +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimentalPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimentalPass.cpp new file mode 100755 index 00000000..f3717f68 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimentalPass.cpp @@ -0,0 +1,107 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// +#include "flagtree/Common/UnifiedHardware.h" + +#include "mlir/Dialect/Ptr/IR/PtrDialect.h" +#include "mlir-ext/Dialect/MathExt/IR/MathExt.h" +#include "triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h" +#include "triton-shared/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.h" +#include "triton-shared/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.h" +#include "triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h" +#include "triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h" +#include "triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h" +#include "triton-shared/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimental.h" +#include "triton-shared/Conversion/TritonToLinalgExperimental/TritonToPtr.h" +#include "triton-shared/Conversion/TritonToStructured/TritonToStructured.h" +#include "triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h" +#include "triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h" +#include "triton-shared/Dialect/TPtr/IR/TPtrDialect.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "mlir/Conversion/ReconcileUnrealizedCasts/ReconcileUnrealizedCasts.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/Passes.h" + +using namespace mlir; +using namespace triton; +//using namespace mlir::mathext; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonToLinalgExperimental/Passes.h.inc" + +namespace { + +class TritonToLinalgExperimentalPass + : public TritonToLinalgExperimentalBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createWaferTritonToStructuredPass()); + + // Erase dead code and fold constants created during lowering + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + + pm.addPass(createTritonToUnstructuredPass()); + pm.addPass(createTritonArithToLinalgPass(/*tensorPtrToLinalg=*/true)); + + // TODO: structured-to-memref converts the loop iter-args to memref, while + // triton-to-ptr converts the loop iter-args to ptr. These two passes might + // end up conflicting with each other in cases where we have a mixed of + // structured and unstructured accesses. Fortunately, the structured ops do + // not need to use the loop iter-args at all (see code in PtrAnalysis.cpp), + // so if we run remove-dead-values after structured-to-memref, the memref + // iter-args that are used in structured loads and stores should be removed. + // Running this now may be too invasive and cause many IR changes, so + // leave as a TODO for now. + pm.addPass(createStructuredToMemrefPass()); + + // Pass selection is controlled by unified hardware configuration. + auto hardwareManager = mlir::flagtree::createUnifiedHardwareManager(); + auto dmaTag = hardwareManager -> getDMATag(); + if (dmaTag) pm.addPass(createMemrefCopyToDMAFlagTreePass()); + + pm.addPass(createUnstructuredToMemrefPass()); + pm.addPass(createTritonPtrToMemrefPass()); + pm.addPass(createTritonToPtrPass()); + pm.addPass(createReconcileUnrealizedCastsPass()); + pm.addPass(createReconcilePtrCastsPass()); + + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + + auto sharedTag = hardwareManager -> getSharedMemoryTag(); + if (sharedTag) pm.addPass(createNoBufferizeFlagTreePass()); + + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> +triton::createTritonToLinalgExperimentalPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/TritonToPtrPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/TritonToPtrPass.cpp new file mode 100755 index 00000000..4b21d6c5 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/TritonToPtrPass.cpp @@ -0,0 +1,490 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// +// This pass lowers all triton ops on pointer to their equivalent form in the +// proposed Pointer Dialect: +// https://discourse.llvm.org/t/rfc-ptr-dialect-modularizing-ptr-ops-in-the-llvm-dialect/75142 +// +// This pass is intended to be used after all running +// triton-arith-to-linalg="tensor-ptr-to-linalg=true". +// All triton ops on tensors of pointers are expected to have been lowered to +// linalg ops, and that only triton ops on single pointers remain. +// +// Implementation notes: +// Because triton pointers are typed whereas the !ptr.ptr type isn't. The +// lowering for addptr will have to manually scale the offsets by pointee type. +// As a result, bitcasts are no-op after this pass. +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Func/Transforms/FuncConversions.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Ptr/IR/PtrDialect.h" +#include "mlir/Dialect/Ptr/IR/PtrTypes.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/IRMapping.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/IR/Types.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/AnalysisStructured/PtrAnalysis.h" +#include "triton-shared/Conversion/TritonToLinalgExperimental/TritonToPtr.h" +#include "triton-shared/Dialect/TPtr/IR/TPtrDialect.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Utils/Utils.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "llvm/ADT/STLExtras.h" + +#define DEBUG_TYPE "triton-to-ptr" + +using namespace mlir; + +namespace { + +#define GEN_PASS_DEF_TRITONTOPTR +#include "triton-shared/Conversion/TritonToLinalgExperimental/Passes.h.inc" + +// Convert tensor.insert_slice to use ptr.ptr type. This insert_slice op must +// have been lowered from tl.cat +struct InsertSliceConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + InsertSliceConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(tensor::InsertSliceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp( + op, adaptor.getSource(), adaptor.getDest(), op.getMixedOffsets(), + op.getMixedSizes(), op.getMixedStrides()); + return success(); + } +}; + +// Convert tensor.empty with !tt.ptr to tensor.empty with !ptr.ptr +struct EmptyTensorConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + EmptyTensorConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(tensor::EmptyOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp( + op, op.getType().getShape(), ptr::PtrType::get(rewriter.getContext())); + return success(); + } +}; + +// This expand shape op must have been lowered from tt.expand_dims which could +// operate on tensor of pointers. +// Convert expand shape op to operate on !ptr.ptr instead of !tt.ptr. +struct ExpandShapeConverter + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + ExpandShapeConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(tensor::ExpandShapeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp( + op, getTypeConverter()->convertType(op.getType()), adaptor.getSrc(), + op.getReassociationExprs()); + return success(); + } +}; + +// arith.select could operate on triton pointers. Convert to use !ptr.ptr +struct SelectOpConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + SelectOpConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(arith::SelectOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp( + op, getTypeConverter()->convertType(op.getType()), + adaptor.getCondition(), adaptor.getTrueValue(), + adaptor.getFalseValue()); + return success(); + } +}; + +// Convert bitcast which is a no-op because !ptr.ptr is opaque with no pointee +// type. +struct BitCastConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + BitCastConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::BitcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getType())) { + return failure(); + } + // Bitcast is a no-op, simply forward the src + rewriter.replaceOp(op, adaptor.getSrc()); + return success(); + } +}; + +// Convert tt.addptr to ptr.ptradd. Since the !ptr.ptr type is opaque, we scale +// the offset explicitly using type_offset op. This approach means that bitcast +// is a no-op. +struct AddPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + AddPtrConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::AddPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getType())) { + return failure(); + } + auto loc = op->getLoc(); + auto pointeeType = cast(op.getType()).getPointeeType(); + auto offsetType = op.getOffset().getType(); + auto pointeeSizeInBytes = + rewriter.create(loc, offsetType, pointeeType); + auto scaledOffset = + rewriter.create(loc, op.getOffset(), pointeeSizeInBytes); + rewriter.replaceOpWithNewOp( + op, ptr::PtrType::get(rewriter.getContext()), adaptor.getPtr(), + scaledOffset); + return success(); + } +}; + +// Convert tt.load which loads from a single pointer into a pair of +// to_memref and memref.load op. +// In the case of mask, the load is guarded by an scf.if +struct LoadConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LoadConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getType())) { + return failure(); + } + auto ptr = op.getPtr(); + auto pointeeType = + cast(ptr.getType()).getPointeeType(); + + auto memref = rewriter.create( + op->getLoc(), MemRefType::get({1}, pointeeType), adaptor.getPtr()); + + auto zero = rewriter.create(op.getLoc(), 0); + + if (op.getMask()) { + auto ifOp = rewriter.create( + op->getLoc(), op.getMask(), + [&](OpBuilder &b, Location loc) { + // Truthy case, load from the index. + Value memrefLoad = rewriter.create( + op->getLoc(), memref, ValueRange{zero}); + b.create(loc, memrefLoad); + }, + [&](OpBuilder &b, Location loc) { + // Falsy case, yield `other` or 0 as the default value. + if (op.getOther()) { + b.create(loc, op.getOther()); + } else { + auto elemType = op.getType(); + auto zeroAttr = b.getZeroAttr(elemType); + assert(zeroAttr && "unexpected element type"); + Value val = b.create(loc, zeroAttr); + b.create(loc, val); + } + }); + rewriter.replaceOp(op, ifOp); + } else { + auto memrefLoad = rewriter.create(op->getLoc(), memref, + ValueRange{zero}); + + rewriter.replaceOp(op, memrefLoad); + } + return success(); + } +}; + +// Convert tt.store which stores to a single pointer into a pair of +// to_memref and memref.store op. +// In the case of mask, the store is guarded by an scf.if +struct StoreConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + StoreConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getValue().getType())) { + return failure(); + } + auto ptr = op.getPtr(); + auto pointeeType = + cast(ptr.getType()).getPointeeType(); + + IRRewriter::InsertionGuard g(rewriter); + if (op.getMask()) { + auto ifOp = rewriter.create(op->getLoc(), op.getMask(), + /*withElseRegion*/ false); + rewriter.setInsertionPointToStart( + &ifOp.getThenRegion().getBlocks().front()); + } + + auto memref = rewriter.create( + op->getLoc(), MemRefType::get({1}, pointeeType), adaptor.getPtr()); + auto zero = rewriter.create(op.getLoc(), 0); + + rewriter.create(op->getLoc(), op.getValue(), memref, + ValueRange{zero}); + + rewriter.eraseOp(op); + + return success(); + } +}; + +// Convert tt.ptr_to_int to ptr.ptrtoint +struct PtrToIntConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + PtrToIntConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::PtrToIntOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getType())) { + return failure(); + } + rewriter.replaceOpWithNewOp(op, op.getType(), + adaptor.getSrc()); + return success(); + } +}; + +// Convert tt.int_to_ptr to ptr.ptrtoint +struct IntToPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + IntToPtrConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(triton::IntToPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (isa(op.getType())) { + return failure(); + } + rewriter.replaceOpWithNewOp( + op, ptr::PtrType::get(rewriter.getContext()), adaptor.getSrc()); + return success(); + } +}; + +// Convert a linalg op on triton pointer to use !ptr.ptr +// The conversion infrastructrure will recursively handle the inner op +// which could be either tt.load, tt.store, tt.bitcast, tt.int_to_ptr, and +// tt.ptr_to_int and use their corresponding converters. +struct LinalgPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LinalgPtrConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(linalg::GenericOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + SmallVector convertedTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), + convertedTypes))) { + return failure(); + } + + auto replacement = rewriter.create( + op.getLoc(), convertedTypes, adaptor.getInputs(), adaptor.getOutputs(), + op.getIndexingMapsArray(), op.getIteratorTypesArray()); + + Region ®ion = op.getRegion(); + Block &block = region.front(); + + TypeConverter::SignatureConversion mapping(block.getArgumentTypes().size()); + if (failed(typeConverter->convertSignatureArgs(block.getArgumentTypes(), + mapping))) + return failure(); + + // Perform signature conversion on the body block. + rewriter.applySignatureConversion(&block, mapping); + + // Splice the old body region into the new for-op. + Region &dstRegion = replacement.getBodyRegion(); + rewriter.inlineRegionBefore(op.getRegion(), dstRegion, dstRegion.end()); + + rewriter.replaceOp(op, replacement); + + return success(); + } +}; + +// The linalg.yield op is still yielding the original !tt.ptr results, convert +// them to use the new !ptr.ptr results +struct LinalgYieldConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LinalgYieldConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(linalg::YieldOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, adaptor.getOperands()); + return success(); + } +}; + +// Convert linalg.fill to use !ptr.ptr. linalg.fill on triton pointer is lowered +// from tt.splat on a triton pointer. +struct LinalgFillPtrConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LinalgFillPtrConverter(const TypeConverter &typeConverter, + MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult + matchAndRewrite(linalg::FillOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, adaptor.getInputs(), + adaptor.getOutputs()); + return success(); + } +}; + +class TritonPtrTypeConverter : public TypeConverter { +public: + TritonPtrTypeConverter(MLIRContext *context) { + addConversion([](Type type) { return type; }); + addConversion([context](triton::PointerType ptrType) { + return ptr::PtrType::get(context); + }); + addConversion([context](RankedTensorType tensorType) { + if (isa(tensorType.getElementType())) { + return RankedTensorType::get(tensorType.getShape(), + ptr::PtrType::get(context)); + } + return tensorType; + }); + auto createCast = [&](OpBuilder &builder, Type resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }; + addTargetMaterialization(createCast); + addSourceMaterialization(createCast); + addArgumentMaterialization(createCast); + } +}; + +class TritonToPtrPass : public impl::TritonToPtrBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert< + arith::ArithDialect, math::MathDialect, affine::AffineDialect, + scf::SCFDialect, tensor::TensorDialect, triton::TritonDialect, + tts::TritonStructuredDialect, ptr::PtrDialect, tptr::TPtrDialect>(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + TritonPtrTypeConverter typeConverter(&getContext()); + + target.addIllegalOp(); + + // We do not want to lower triton load and store on block pointers + target.addDynamicallyLegalOp([](auto op) { + auto ptrType = op->getOperand(0).getType(); + if (triton::isTensorPointerType(ptrType)) { + return true; + } + return !triton::isPtrTypeLike(ptrType); + }); + + target.addDynamicallyLegalOp< + linalg::FillOp, linalg::GenericOp, linalg::YieldOp, tensor::EmptyOp, + tensor::ExpandShapeOp, tensor::InsertSliceOp, arith::SelectOp>( + [](auto op) { + return llvm::all_of( + llvm::concat(op->getOperands(), op->getResults()), + [&](Value v) { return !triton::isPtrTypeLike(v.getType()); }); + }); + + target.addLegalDialect(); + + patterns + .add( + typeConverter, patterns.getContext()); + + mlir::scf::populateSCFStructuralTypeConversionsAndLegality( + typeConverter, patterns, target); + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> triton::createTritonToPtrPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.cpp new file mode 100755 index 00000000..def61b79 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.cpp @@ -0,0 +1,113 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.h" +#include +#include + +namespace TTOpConverters { +using namespace mlir; +using namespace triton; + +// ArgMinConverter functions +LogicalResult ArgMinConverter::matchComparisonResult( + Value currValue, Value currIndex, Value reduceValue, Value reduceIndex, + mlir::Block::iterator &it, Value &comparisonResult) { + LLVM_DEBUG(llvm::dbgs() << "Matching: " << *it << "\n"); + + auto cmpOp = dyn_cast(*it); + auto cmpIOp = dyn_cast(*it++); + if (!cmpOp && !cmpIOp) + return failure(); + + if (cmpOp) { + if (cmpOp.getPredicate() != arith::CmpFPredicate::OLT || + currValue != cmpOp.getLhs() || reduceValue != cmpOp.getRhs()) { + return failure(); + } + comparisonResult = cmpOp; + } + + if (cmpIOp) { + if ((cmpIOp.getPredicate() != arith::CmpIPredicate::slt && + cmpIOp.getPredicate() != arith::CmpIPredicate::ult) || + currValue != cmpIOp.getLhs() || reduceValue != cmpIOp.getRhs()) { + return failure(); + } + comparisonResult = cmpIOp; + } + + return success(); +} + +float ArgMinConverter::getBaseReductionValue() { + return std::numeric_limits::infinity(); +} + +int8_t ArgMinConverter::getBaseReductionIntValue() { + return std::numeric_limits::max(); +} +uint8_t ArgMinConverter::getBaseReductionUIntValue() { + return std::numeric_limits::max(); +} + +// ArgMaxConverter functions +LogicalResult ArgMaxConverter::matchComparisonResult( + Value currValue, Value currIndex, Value reduceValue, Value reduceIndex, + mlir::Block::iterator &it, Value &comparisonResult) { + auto cmpOp = dyn_cast(*it); + auto cmpIOp = dyn_cast(*it++); + if (!cmpOp && !cmpIOp) + return failure(); + + if (cmpOp) { + if (cmpOp.getPredicate() != arith::CmpFPredicate::OGT || + currValue != cmpOp.getLhs() || reduceValue != cmpOp.getRhs()) { + return failure(); + } + comparisonResult = cmpOp; + } + + if (cmpIOp) { + if ((cmpIOp.getPredicate() != arith::CmpIPredicate::sgt && + cmpIOp.getPredicate() != arith::CmpIPredicate::ugt) || + currValue != cmpIOp.getLhs() || reduceValue != cmpIOp.getRhs()) { + return failure(); + } + comparisonResult = cmpIOp; + } + + return success(); +} + +float ArgMaxConverter::getBaseReductionValue() { + return -std::numeric_limits::infinity(); +} + +int8_t ArgMaxConverter::getBaseReductionIntValue() { + return std::numeric_limits::min(); +} +uint8_t ArgMaxConverter::getBaseReductionUIntValue() { + return std::numeric_limits::min(); +} + +} // namespace TTOpConverters diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.cpp new file mode 100755 index 00000000..c0485ee5 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.cpp @@ -0,0 +1,2176 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#if __has_include("bishengir/Dialect/Annotation/IR/Annotation.h") +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#endif +#if __has_include("bishengir/Dialect/HIVM/IR/HIVM.h") +#include "bishengir/Dialect/HIVM/IR/HIVM.h" +#endif + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Arith/Utils/Utils.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/IRMapping.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/OperationSupport.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/FormatVariadic.h" +#include +#include + +#define DEBUG_TYPE "triton-block-ptr-analysis" +namespace mlir { +namespace triton { + +// MemAccType selectMaxMemAccTy(const MemAccType &v1, const MemAccType &v2) { +// return (v1 > v2) ? v1 : v2; +// } + +SmallVector &BlockData::getOffsetsRef() { return this->offsets; } + +SmallVector &BlockData::getSizesRef() { return this->sizes; } + +SmallVector &BlockData::getStridesRef() { return this->strides; } + +Value &BlockData::getSourceRef() { return this->source; } + +OpFoldResult &BlockData::getScalarRef() { return this->scalar; } + +SmallVector BlockData::getOffsets() const { + return this->offsets; +} + +SmallVector BlockData::getSizes() const { return this->sizes; } + +SmallVector BlockData::getStrides() const { + return this->strides; +} + +OpFoldResult BlockData::getOffset(int index) const { + return this->offsets[index]; +} + +OpFoldResult BlockData::getSize(int index) const { return this->sizes[index]; } + +OpFoldResult BlockData::getStride(int index) const { + return this->strides[index]; +} + +OpFoldResult BlockData::getScalar() const { return this->scalar; } + +Value BlockData::getSource() const { return this->source; } + +MemAccType BlockData::getMemAccType() const { return this->memAccTy; }; + +MemAccType &BlockData::getMemAccTypeRef() { return this->memAccTy; }; + +bool BlockData::isScalar() const { return !(this->scalar).isNull(); } + +bool BlockData::isEmpty() const { + return !(this->getRank() || this->source || !(this->scalar).isNull()); +} + +bool BlockData::hasSource() const { return this->source != nullptr; } + +void BlockData::removeSource() { this->source = nullptr; }; + +bool BlockData::hasResElemTy() const { return this->resElemTy != nullptr; } + +Type &BlockData::getResElemTyRef() { return this->resElemTy; } + +Type BlockData::getResElemTy() const { return this->resElemTy; } + +int64_t BlockData::getRank() const { + assert(offsets.size() == sizes.size() && offsets.size() == strides.size()); + return this->offsets.size(); +} + +void BlockData::setResElemTy(const Type &Ty) { this->resElemTy = Ty; } + +void BlockData::setScalar(const OpFoldResult &scalar) { this->scalar = scalar; } + +void BlockData::setSource(const Value &src) { this->source = src; } + +void BlockData::setOffsets(const SmallVector &offsets) { + this->offsets = offsets; +} + +void BlockData::setStrides(const SmallVector &strides) { + this->strides = strides; +} + +void BlockData::setSizes(const SmallVector &szs) { + this->sizes = szs; +} + +void BlockData::setMemAccTy(const MemAccType &v) { this->memAccTy = v; } + +void BlockData::setMemAccVal(const MemAccVal v) { this->memAccTy.value = v; } + +OpFoldResult BlockData::inferBlockOffset(const Location &loc, + OpBuilder &builder) const { + OpFoldResult retOffset = builder.getIndexAttr(0); + for (auto ofr : offsets) { + retOffset = addOpFoldResult(retOffset, ofr, loc, builder); + } + return retOffset; +} + +MemRefType BlockData::getResultMemrefType(int64_t offset, + ArrayRef resultShape) const { + SmallVector staticStrides; + SmallVector dynamicStrides; + dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides); + + auto baseMemrefType = dyn_cast(this->source.getType()); + assert(baseMemrefType && + "Invalid element type. It should be a base memref type."); + auto elementType = baseMemrefType.getElementType(); + auto layout = + StridedLayoutAttr::get(this->source.getContext(), offset, staticStrides); + return MemRefType::get(resultShape, elementType, layout); +} + +void BlockData::addBlock(BlockData &lBlock, BlockData &rBlock, Location loc, + ConversionPatternRewriter &rewriter) { + assert(this->isEmpty() && lBlock.getRank() == rBlock.getRank()); + // When both left block and right block have source, it is indirect load. + assert(!(lBlock.hasSource() && rBlock.hasSource()) && + "Don't support each BlockData has own base source pointer"); + this->source = + lBlock.hasSource() ? lBlock.getSourceRef() : rBlock.getSourceRef(); + + assert(!(lBlock.hasResElemTy() && rBlock.hasResElemTy())); + if (lBlock.hasResElemTy()) { + assert(lBlock.hasSource()); + this->resElemTy = lBlock.getResElemTyRef(); + } else if (rBlock.hasResElemTy()) { + assert(rBlock.hasSource()); + this->resElemTy = rBlock.getResElemTyRef(); + } + + // Acctually `scalar` should be accumulated into `offset` and `stride` finally + // In addBlock, just pass `scalar` when: + // 1. both lhs and rhs have `scalar` + // 2. otherwise, both lhs and rhs are scalar type with rank 0 + // Except above, original `scalar` has been fused into `offset` under add. + if (lBlock.isScalar() && rBlock.isScalar()) { + auto addScalar = addOpFoldResult(lBlock.getScalarRef(), + rBlock.getScalarRef(), loc, rewriter); + this->scalar = addScalar; + } else if (lBlock.getRank() == 0) { + // When both lhs and rhs are scalar type with rank 0, just try passing + // potential `scalar` + this->scalar = + lBlock.isScalar() ? lBlock.getScalarRef() : rBlock.getScalarRef(); + } + + for (const auto &[lOffset, rOffset] : + llvm::zip(lBlock.getOffsetsRef(), rBlock.getOffsetsRef())) { + this->offsets.push_back(addOpFoldResult(lOffset, rOffset, loc, rewriter)); + } + + for (const auto &[lStride, rStride] : + llvm::zip(lBlock.getStridesRef(), rBlock.getStridesRef())) { + this->strides.push_back(addOpFoldResult(lStride, rStride, loc, rewriter)); + } + + // Both sizes are same implicitly under `add` + this->sizes = lBlock.getSizesRef(); + + this->getMemAccTypeRef().merge(lBlock.getMemAccTypeRef()); + this->getMemAccTypeRef().merge(rBlock.getMemAccTypeRef()); + // this->setMemAccTy(selectMaxMemAccTy(lBlock.getMemAccType(), + // rBlock.getMemAccType())); +} + +void BlockData::subBlock(BlockData &lBlock, BlockData &rBlock, Location loc, + ConversionPatternRewriter &rewriter) { + assert(this->isEmpty() && lBlock.getRank() == rBlock.getRank()); + + if (lBlock.isScalar() && rBlock.isScalar()) { + auto subScalar = subOpFoldResult(lBlock.getScalarRef(), + rBlock.getScalarRef(), loc, rewriter); + this->scalar = subScalar; + } else if (lBlock.getRank() == 0) { + // When both lhs and rhs are scalar type with rank 0, just try passing + // potential `scalar` + this->scalar = + lBlock.isScalar() ? lBlock.getScalarRef() : rBlock.getScalarRef(); + } + + for (const auto &[lOffset, rOffset] : + llvm::zip(lBlock.getOffsetsRef(), rBlock.getOffsetsRef())) { + this->offsets.push_back(subOpFoldResult(lOffset, rOffset, loc, rewriter)); + } + + for (const auto &[lStride, rStride] : + llvm::zip(lBlock.getStridesRef(), rBlock.getStridesRef())) { + this->strides.push_back(subOpFoldResult(lStride, rStride, loc, rewriter)); + } + + // Both sizes are same implicitly under `sub` + this->sizes = lBlock.getSizesRef(); + + this->getMemAccTypeRef().merge(lBlock.getMemAccTypeRef()); + this->getMemAccTypeRef().merge(rBlock.getMemAccTypeRef()); + // this->setMemAccTy(selectMaxMemAccTy(lBlock.getMemAccType(), + // rBlock.getMemAccType())); +} + +void BlockData::mulBlock(BlockData &lBlock, BlockData &rBlock, Location loc, + ConversionPatternRewriter &rewriter) { + assert(this->isEmpty() && lBlock.getRank() == rBlock.getRank()); + + assert(!(lBlock.hasSource() && rBlock.hasSource())); + + if (lBlock.isScalar() && rBlock.isScalar()) { + LLVM_DEBUG({ + llvm::dbgs() << "lBlock.scalar:" << lBlock.getScalar() + << " rBlbock.scalar:" << rBlock.getScalar() << "\n"; + }); + + auto scalar = + mulOpFoldResult(lBlock.getScalar(), rBlock.getScalar(), loc, rewriter); + this->scalar = scalar; + } + + // assert( + // (lBlock.isScalar() ^ rBlock.isScalar()) && + // "Currently only support one and only one scalar in function + // mulBlock()"); + + BlockData *lb = &lBlock; + BlockData *rb = &rBlock; + if (lb->isScalar()) { + std::swap(lb, rb); + } + + // In mulBlock, `scalar` will be accumulated into `offset` and `stride` + OpFoldResult rScalar = rb->getScalarRef(); + for (const auto &lOffset : lb->getOffsetsRef()) { + this->offsets.push_back(mulOpFoldResult(lOffset, rScalar, loc, rewriter)); + } + + for (const auto &lStride : lb->getStridesRef()) { + this->strides.push_back(mulOpFoldResult(lStride, rScalar, loc, rewriter)); + } + + this->sizes = lb->getSizesRef(); + + this->getMemAccTypeRef().merge(lBlock.getMemAccTypeRef()); + this->getMemAccTypeRef().merge(rBlock.getMemAccTypeRef()); + // this->setMemAccTy(selectMaxMemAccTy(lBlock.getMemAccType(), + // rBlock.getMemAccType())); +} + +void BlockData::divBlock(BlockData &lBlock, BlockData &rBlock, Location loc, + ConversionPatternRewriter &rewriter) { + assert(this->isEmpty() && lBlock.getRank() == rBlock.getRank()); + + assert(!(lBlock.hasSource() && rBlock.hasSource())); + assert(lBlock.isScalar() && rBlock.isScalar()); + + auto rScalar = rBlock.getScalar(); + this->scalar = divOpFoldResult(lBlock.getScalar(), rScalar, loc, rewriter); + + for (auto lOffset : lBlock.getOffsetsRef()) { + this->offsets.push_back(divOpFoldResult(lOffset, rScalar, loc, rewriter)); + } + + for (auto lStride : lBlock.getStridesRef()) { + this->strides.push_back(divOpFoldResult(lStride, rScalar, loc, rewriter)); + } + + this->sizes = lBlock.getSizesRef(); + + this->getMemAccTypeRef().merge(lBlock.getMemAccTypeRef()); + this->getMemAccTypeRef().merge(rBlock.getMemAccTypeRef()); + // this->setMemAccTy(selectMaxMemAccTy(lBlock.getMemAccType(), + // rBlock.getMemAccType())); +} + +memref::ReinterpretCastOp BlockData::createCastOp(ArrayRef resultShape, + const Location &loc, + OpBuilder &builder) const { + OpFoldResult resOffset = this->inferBlockOffset(loc, builder); + auto resultType = this->getResultMemrefType( + isa(resOffset) ? getConstantIntValue(resOffset).value() + : ShapedType::kDynamic, + resultShape); + + SmallVector strides(this->strides); + for (size_t i = 0; i < strides.size(); i++) { + if (resultShape[i] == 1) { + if (auto strideValue = dyn_cast(strides[i])) { + auto oneIdx = + builder.create(loc, builder.getIndexAttr(1)); + strides[i] = builder.create(loc, strideValue, oneIdx) + .getResult(); + } + } + } + + return builder.create( + loc, resultType, this->source, resOffset, this->sizes, strides); +} + +void BlockData::dump() const { + llvm::outs() << "[INFO][BEG] BlockData info\n"; + llvm::outs() << "offsets has " << offsets.size() << " items\n"; + int cnt = 0; + for (auto it = offsets.begin(); it != offsets.end(); ++it) { + llvm::outs() << "offsets[" << cnt++ << "] = " << *it << "\n"; + } + llvm::outs() << "sizes has " << sizes.size() << " items\n"; + cnt = 0; + for (auto it = sizes.begin(); it != sizes.end(); ++it) { + llvm::outs() << "sizes[" << cnt++ << "] = " << *it << "\n"; + } + llvm::outs() << "strides has " << strides.size() << " items\n"; + cnt = 0; + for (auto it = strides.begin(); it != strides.end(); ++it) { + llvm::outs() << "strides[" << cnt++ << "] = " << *it << "\n"; + } + llvm::outs() << "source = " << source << "\n"; + llvm::outs() << "scalar = " << scalar << "\n"; + llvm::outs() << "resElemTy = " << resElemTy << "\n"; + llvm::outs() << "memAccTy = " << memAccTy.toString() << "\n"; + llvm::outs() << "[INFO][END] BlockData info\n"; +} + +Value BlockDataParser::getScalarMemRef(Value ptr, Value memref, + const Location &loc, + ConversionPatternRewriter &rewriter) { + assert(isa(ptr.getType()) && "expect a scalar pointer"); + if (ptr.getDefiningOp()) { + if (auto castOp = memref.getDefiningOp()) { + return castOp.getResult(); + } else { + llvm_unreachable("pointer value is defined by an unexpected op"); + } + } + + assert(isa(ptr) && + "pointer should be produced by addptr or block argument"); + BlockData data; + data.setSource(memref); + data.getOffsetsRef().push_back(rewriter.getIndexAttr(0)); + data.getSizesRef().push_back(rewriter.getIndexAttr(1)); + data.getStridesRef().push_back(rewriter.getIndexAttr(1)); + auto castOp = data.createCastOp(SmallVector(1, 1), loc, rewriter); + return castOp.getResult(); +} + +void BlockDataParser::parse( + Value operand, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + if (known.find(operand) != known.end()) { + return data = known.lookup(operand), void(); + } + + if (isa(operand.getType())) { + data.setScalar(getOpFoldResultOfLayoutInfo(operand, rewriter)); + return; + } + + // + if (isa(operand.getType())) { + // Just consider two state: ptr and ptr> + auto remappedPtr = rewriter.getRemappedValue(operand); + assert(remappedPtr); + if (auto op = operand.getDefiningOp()) { + if (auto addPtrOp = dyn_cast(op)) { + parseAddPtr(addPtrOp, data, loc, rewriter, known); + } else if (auto bitcastOp = dyn_cast(op)) { + parseBitcast(bitcastOp, data, loc, rewriter, known); + } else if (auto makeTensorPtrOp = dyn_cast(op)) { + parseTensorPtr(makeTensorPtrOp, data, loc, rewriter, known); + } else if (auto advanceOp = dyn_cast(op)) { + // To support + // ptr_0 = tl.advance(ptr) + // ptr_1 = tl.advance(ptr_0) + parseTensorPtr(advanceOp, data, loc, rewriter, known); + } else if (auto intToPtrOp = dyn_cast(op)) { + data.setSource(remappedPtr); + } else { + LLVM_DEBUG({ llvm::dbgs() << operand << "\n"; }); + llvm_unreachable( + "Unexpected operand defining operation, a scalar " + "pointer can only be produced by AddPtrOp or direct block ptr"); + } + } else { + data.setSource(remappedPtr); + } + return; + } + + // not a scalar pointer + if (auto addOp = operand.getDefiningOp()) { + parseAdd(addOp, data, loc, rewriter, known); + } else if (auto subOp = operand.getDefiningOp()) { + parseSub(subOp, data, loc, rewriter, known); + } else if (auto mulOp = operand.getDefiningOp()) { + parseMul(mulOp, data, loc, rewriter, known); + } else if (auto addPtrOp = operand.getDefiningOp()) { + parseAddPtr(addPtrOp, data, loc, rewriter, known); + } else if (auto constOp = operand.getDefiningOp()) { + parseConstSplat(constOp, data, loc, rewriter, known); + } else if (auto broadcastOp = operand.getDefiningOp()) { + parseBroadcast(broadcastOp, data, loc, rewriter, known); + } else if (auto splatOp = operand.getDefiningOp()) { + parseSplat(splatOp, data, loc, rewriter, known); + } else if (auto expandDimsOp = + operand.getDefiningOp()) { + parseExpandDims(expandDimsOp, data, loc, rewriter, known); + } else if (auto remOp = operand.getDefiningOp()) { + parseRem(remOp, data, loc, rewriter, known); + } else if (auto bitcastOp = operand.getDefiningOp()) { + parseBitcast(bitcastOp, data, loc, rewriter, known); + } else if (auto extsiOp = operand.getDefiningOp()) { + parseExtSI(extsiOp, data, loc, rewriter, known); + } else if (auto divOp = operand.getDefiningOp()) { + parseDiv(divOp, data, loc, rewriter, known); + } else if (auto makeRangeOp = operand.getDefiningOp()) { + parseMakeRange(makeRangeOp, data, loc, rewriter, known); + } else if (auto reduceOp = operand.getDefiningOp()) { + parseReduce(reduceOp, data, loc, rewriter, known); + } else if (auto loadOp = operand.getDefiningOp()) { + parseIndirectLoad(loadOp, data, loc, rewriter, known); + } else if (auto castOp = operand.getDefiningOp()) { + parseIndirectLoad(castOp, data, loc, rewriter, known); + } else if (auto extractSliceOp = + operand.getDefiningOp()) { + parseExtractSlice(extractSliceOp, data, loc, rewriter, known); + } else if (auto forOp = operand.getDefiningOp()) { + parseIndirectLoad(forOp, data, loc, rewriter, known); + } else if (auto tensorCastOp = operand.getDefiningOp()) { + // Used for identity operation. + parse(tensorCastOp.getSource(), data, loc, rewriter, known); + } else if (auto fillOp = operand.getDefiningOp()) { + parseFill(fillOp, data, loc, rewriter, known); + } else if (auto selectOp = operand.getDefiningOp()) { + parseSelect(selectOp, data, loc, rewriter, known); + } else { + operand.dump(); + llvm_unreachable("encountered AddPtrOp produced by unsupported operation"); + } +} + +void BlockDataParser::parseAdd( + arith::AddIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + BlockData lBlock, rBlock; + parse(op.getLhs(), lBlock, loc, rewriter, known); + parse(op.getRhs(), rBlock, loc, rewriter, known); + data.addBlock(lBlock, rBlock, loc, rewriter); +} + +void BlockDataParser::parseSub( + arith::SubIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + BlockData lBlock, rBlock; + parse(op.getLhs(), lBlock, loc, rewriter, known); + parse(op.getRhs(), rBlock, loc, rewriter, known); + data.subBlock(lBlock, rBlock, loc, rewriter); +} + +void BlockDataParser::parseMul( + arith::MulIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + BlockData lBlock, rBlock; + parse(op.getLhs(), lBlock, loc, rewriter, known); + parse(op.getRhs(), rBlock, loc, rewriter, known); + + data.mulBlock(lBlock, rBlock, loc, rewriter); +} + +void BlockDataParser::parseDiv( + arith::DivSIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + BlockData lBlock, rBlock; + parse(op.getLhs(), lBlock, loc, rewriter, known); + parse(op.getRhs(), rBlock, loc, rewriter, known); + data.divBlock(lBlock, rBlock, loc, rewriter); +} + +// TODO : support modulos +void BlockDataParser::parseRem( + arith::RemSIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(false && "Address expression with modulo is not supported yet, it " + "shall be analysis at linearize."); +} + +void BlockDataParser::parseMakeRange( + triton::MakeRangeOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + auto shape = dyn_cast(op.getType()).getShape(); + + auto start = op.getStart(); + auto end = op.getEnd(); + auto stride = (end >= start) && (end - start <= shape[0]); + assert(stride == 1 && + "make_range op should always return a tensor of stride 1"); + + data.getOffsetsRef().push_back(rewriter.getIndexAttr(start)); + data.getSizesRef().push_back(rewriter.getIndexAttr(shape[0])); + data.getStridesRef().push_back(rewriter.getIndexAttr(stride)); +} + +void BlockDataParser::parseExpandDims( + triton::ExpandDimsOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + + parse(op.getSrcMutable().get(), data, loc, rewriter, known); + auto resShape = dyn_cast(op.getResult().getType()).getShape(); + auto axis = op.getAxis(); + + assert(resShape[axis] == 1 && + "The destiny shape of changed dimension should be 1"); + + data.getOffsetsRef().insert(data.getOffsetsRef().begin() + axis, + rewriter.getIndexAttr(0)); + data.getSizesRef().insert(data.getSizesRef().begin() + axis, + rewriter.getIndexAttr(1)); + data.getStridesRef().insert(data.getStridesRef().begin() + axis, + rewriter.getIndexAttr(0)); +} + +void BlockDataParser::parseExtractSlice( + tensor::ExtractSliceOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + const std::string scenarioMessages = + "PtsAnalysis supports indirectly block load in the " + "following scenario\n" + "B = tl.load(Aptr + Aoffset) # B is 1D tensor\n" + "s = tl.extract_slice(indices, offsets= (i,), sizes= " + "(1,), strides= (1,)) # s is a tensor<1x$dtype>\n" + "D = tl.load(Cptr + s + Coffset) # s is used as the " + "scalar offset\n"; // tensor<2x$dtype> will be support + // soon + + auto extract_src = op->getOperand(0); + BlockData srcBlock; + parse(extract_src, srcBlock, loc, rewriter, known); + if (!srcBlock.hasSource()) { + llvm_unreachable(scenarioMessages.c_str()); + } + if (!isa(srcBlock.getSource().getDefiningOp())) { + llvm_unreachable(scenarioMessages.c_str()); + } + + auto extract_result = op->getResult(0); + auto shaped_ty = dyn_cast(extract_result.getType()); + auto shape = shaped_ty.getShape(); + if (shape.size() > 1 || shape[0] > 1) { + llvm_unreachable(scenarioMessages.c_str()); + } + auto castOp = rewriter.create( + loc, RankedTensorType::get(shape, rewriter.getIndexType()), + extract_result); + auto offset = castOp.getResult(); + if (data.isEmpty()) { + data.getOffsetsRef().push_back(offset); + data.getSizesRef().push_back(rewriter.getIndexAttr(shape[0])); + data.getStridesRef().push_back(rewriter.getIndexAttr(1)); + } else { + llvm_unreachable( + "parseExtractSlice with offset already setup not yet supported"); + } +} + +void BlockDataParser::parseBitcast( + triton::BitcastOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + parse(op.getSrc(), data, loc, rewriter, known); + + auto resType = op.getResult().getType(); + Type resElemPointeeTy = nullptr; + if (auto resShapedTy = dyn_cast(resType)) { + auto resElemTy = resShapedTy.getElementType(); + resElemPointeeTy = + dyn_cast(resElemTy).getPointeeType(); + } else { + auto srcPointeeType = + cast(op.getSrc().getType()).getPointeeType(); + auto resPointeeType = cast(resType).getPointeeType(); + + // Handling special case + // If Op is MetaUse or src is i1 block argument and dst is i8, + // it should be converted to UnrealizedConversionCast + if (op->hasAttr("MetaUse") || + (isa(op.getSrc()) && + srcPointeeType == rewriter.getIntegerType(1) && + resPointeeType == rewriter.getIntegerType(8))) { + resElemPointeeTy = resPointeeType; + } else { + auto remappedValue = rewriter.getRemappedValue(op); + data.setSource(remappedValue); + LLVM_DEBUG({ + llvm::dbgs() << "Remapping bitcastOp:\n"; + llvm::dbgs() << op << "\nto \n"; + llvm::dbgs() << remappedValue << "\n"; + }); + } + } + data.setResElemTy(resElemPointeeTy); +} + +void BlockDataParser::parseExtSI( + arith::ExtSIOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + parse(op.getIn(), data, loc, rewriter, known); +} + +void BlockDataParser::parseBroadcast( + triton::BroadcastOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + + auto src = op.getSrcMutable().get(); + auto dst = op.getResult(); + assert(isa(src.getType()) && + "tt.broadcast's input should be a tensor"); + + auto srcShape = dyn_cast(src.getType()).getShape(); + auto dstShape = dyn_cast(dst.getType()).getShape(); + assert(srcShape.size() == dstShape.size() && + "rank of source shoule be equal to destnation"); + + parse(src, data, loc, rewriter, known); + + for (const auto &[idx, src_dst] : + llvm::enumerate(llvm::zip(srcShape, dstShape))) { + const auto &[srcAxis, dstAxis] = src_dst; + if (srcAxis == dstAxis) { + continue; + } + assert(srcAxis < dstAxis && + "srcShape of broadcastOp must be less than dstShape."); + data.getSizesRef()[idx] = rewriter.getIndexAttr(dstAxis); + } +} + +void BlockDataParser::parseSplat( + triton::SplatOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + auto src = op.getSrc(); + auto dst = op.getResult(); + auto dstShape = dyn_cast(dst.getType()).getShape(); + + parse(src, data, loc, rewriter, known); + + if (isa(src.getType()) || + isa(src.getType())) { + if (!data.isEmpty()) { + data.getOffsetsRef().clear(); + data.getSizesRef().clear(); + data.getStridesRef().clear(); + } + for (auto dstAxis : dstShape) { + data.getOffsetsRef().push_back(rewriter.getIndexAttr(0)); + data.getSizesRef().push_back(rewriter.getIndexAttr(dstAxis)); + data.getStridesRef().push_back(rewriter.getIndexAttr(0)); + } + } else { + op->emitError("Block data Analysis: unsupported splat pattern"); + return; + } + if (data.isScalar()) { + data.getOffsetsRef()[0] = data.getScalarRef(); + } +} + +void BlockDataParser::parseConstSplat( + arith::ConstantOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + + DenseElementsAttr denseAttr = dyn_cast(op.getValue()); + assert(denseAttr && denseAttr.isSplat() && + isa(denseAttr.getElementType())); + + auto innerVal = denseAttr.getValues()[0].getValue(); + auto innerValIndexAttr = rewriter.getIndexAttr(innerVal.getSExtValue()); + + // for mul state + data.setScalar(innerValIndexAttr); + + auto resType = dyn_cast(op.getResult().getType()); + size_t loopLimit = resType.getShape().size(); + for (auto i = 0; i < loopLimit; i++) { + // Add original dense val to first dim offset for add state + if (i == 0) { + data.getOffsetsRef().push_back(innerValIndexAttr); + } else { + data.getOffsetsRef().push_back(rewriter.getIndexAttr(0)); + } + data.getSizesRef().push_back(rewriter.getIndexAttr(resType.getShape()[i])); + data.getStridesRef().push_back(rewriter.getIndexAttr(0)); + } +} + +template +std::enable_if_t || + std::is_same_v> +BlockDataParser::parseTensorPtr( + T op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + + Value remappedValue = rewriter.getRemappedValue(op); + if (auto castOp = remappedValue.getDefiningOp()) { + parseReinterpretCast(castOp, data, loc, rewriter, known); + } else { + llvm_unreachable("the value should be mapped to memref.reinterpret_cast"); + } +} + +void BlockDataParser::parseAddPtr( + triton::AddPtrOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + + BlockData ptrBlock, offsetBlock; + parse(op.getPtr(), ptrBlock, op.getLoc(), rewriter, known); + parse(op.getOffset(), offsetBlock, op.getLoc(), rewriter, known); + + assert(ptrBlock.hasSource() && + "Ptr field should provide source/base pointer"); + // offset has source means offset is from tl.load and other ops(TODO) + if (offsetBlock.hasSource()) { + ptrBlock.setMemAccTy(offsetBlock.getMemAccType()); + offsetBlock.removeSource(); + } + + // handle for loop & scalar + if (ptrBlock.getRank() == 1 && offsetBlock.getRank() == 0) { + offsetBlock.getSizesRef().push_back(rewriter.getIndexAttr(1)); + offsetBlock.getOffsetsRef().push_back(offsetBlock.getScalarRef()); + offsetBlock.getStridesRef().push_back(rewriter.getIndexAttr(0)); + } + + assert(ptrBlock.getRank() == offsetBlock.getRank() && + "ptr and offset should have same rank"); + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "[parseAddPtr][BEG] =========================\n"; + os << "[parseAddPtr] op is " << op << "\n"; + for (int i = 0; i < ptrBlock.getRank(); i++) { + os << "ptrBlock.getOffsetsRef()[" << i + << "] = " << ptrBlock.getOffsetsRef()[i] << "\n"; + os << "ptrBlock.getSizesRef()[" << i + << "] = " << ptrBlock.getSizesRef()[i] << "\n"; + os << "ptrBlock.getStridesRef()[" << i + << "] = " << ptrBlock.getStridesRef()[i] << "\n"; + os << "offsetBlock.getOffsetsRef()[" << i + << "] = " << offsetBlock.getOffsetsRef()[i] << "\n"; + os << "offsetBlock.getSizesRef()[" << i + << "] = " << offsetBlock.getSizesRef()[i] << "\n"; + os << "offsetBlock.getStridesRef()[" << i + << "] = " << offsetBlock.getStridesRef()[i] << "\n"; + } + os << "[parseAddPtr][END] -------------------------\n"; + }); + data.addBlock(ptrBlock, offsetBlock, op.getLoc(), rewriter); +} + +void BlockDataParser::parseReinterpretCast( + memref::ReinterpretCastOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + + data.setOffsets(op.getMixedOffsets()); + data.setSizes(op.getMixedSizes()); + data.setStrides(op.getMixedStrides()); + data.setSource(op.getSource()); + + // In memref::ReinterpretCastOp, offset means the total of collapsing multiple + // dimensions, which corresponds to first dim offset in block data. + // Here populate the rest of the dimensions with zeroes. + assert(data.getOffsetsRef().size() == 1); + size_t loopLimit = data.getSizesRef().size(); + for (size_t i = 1; i < loopLimit; i++) { + data.getOffsetsRef().push_back(rewriter.getIndexAttr(0)); + } +} + +void BlockDataParser::parseReduce( + triton::ReduceOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + + const std::string scenarioMessages = + "PtsAnalysis supports indirectly block load in the following scenario\n" + "B = tl.load(Aptr + Aoffset) # B is 1D tensor\n" + "s = tl.min(B) # s is a scalar\n" + "D = tl.load(Cptr + s + Coffset) # s is used as the scalar offset\n"; + + auto reduce_src = op->getOperand(0); + BlockData srcBlock; + parse(reduce_src, srcBlock, loc, rewriter, known); + if (!srcBlock.hasSource()) { + llvm_unreachable(scenarioMessages.c_str()); + } + if (!isa(srcBlock.getSource().getDefiningOp())) { + llvm_unreachable(scenarioMessages.c_str()); + } + + auto reduce_result = op->getResult(0); + auto shaped_ty = dyn_cast(reduce_result.getType()); + auto shape = shaped_ty.getShape(); + auto ops = llvm::map_to_vector(op.getBody()->without_terminator(), + [](Operation &op) { return &op; }); + // Support only the case: scalar = tl.load(1D tensor) + if (shape.size() != 1 || op.getAxis() != 0 || ops.size() != 1 || + !isa(ops.front())) { + llvm_unreachable(scenarioMessages.c_str()); + } + + auto castOp = rewriter.create( + loc, RankedTensorType::get(shape, rewriter.getIndexType()), + reduce_result); + auto offset = castOp.getResult(); + if (data.isEmpty()) { + data.getOffsetsRef().push_back(offset); + data.getSizesRef().push_back(rewriter.getIndexAttr(shape[0])); + data.getStridesRef().push_back(rewriter.getIndexAttr(1)); + } else { + llvm_unreachable("parseReduce with offset already setup not yet supported"); + } +} + +template +void parseIndirectLoad(OpTy op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + // FIXME: assume single result of operation + auto opRes = op->getResult(0); + auto opResTy = opRes.getType(); + std::vector resShape; + if (auto shapedResTy = dyn_cast(opResTy)) { + // For now, we consider this is UnstrucMemAcc because we have no other info. + // Visiting other ops may change the type due to more info. + data.setMemAccVal(MemAccVal::UnstrucMemAcc); + resShape = shapedResTy.getShape().vec(); + } else { + // scalar load means this is used as offset. It is StrucMemAcc. + data.setMemAccVal(MemAccVal::StrucMemAcc); + resShape.push_back(1); + } + for (auto &s : resShape) { + data.getOffsetsRef().push_back(rewriter.getIndexAttr(0)); + data.getSizesRef().push_back(rewriter.getIndexAttr(s)); + data.getStridesRef().push_back(rewriter.getIndexAttr(1)); + } + // set the source in BlockData so that we know an indirect-load op exists in + // the chain. + data.setSource(opRes); +} + +void BlockDataParser::parseFill( + linalg::FillOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + auto src = op.getInputs()[0]; + auto dst = op.getResult(0); + auto dstShape = dyn_cast(dst.getType()).getShape(); + + parse(src, data, loc, rewriter, known); + + if (isa(src.getType())) { + if (!data.isEmpty()) { + data.getOffsetsRef().clear(); + data.getSizesRef().clear(); + data.getStridesRef().clear(); + } + for (auto dstAxis : dstShape) { + data.getOffsetsRef().push_back(rewriter.getIndexAttr(0)); + data.getSizesRef().push_back(rewriter.getIndexAttr(dstAxis)); + data.getStridesRef().push_back(rewriter.getIndexAttr(0)); + } + } else { + op->emitError("Block data Analysis: unsupported fillOp pattern"); + return; + } + if (data.isScalar()) { + data.getOffsetsRef()[0] = data.getScalarRef(); + } +} + +void BlockDataParser::parseSelect( + arith::SelectOp op, BlockData &data, const Location &loc, + ConversionPatternRewriter &rewriter, + const llvm::SmallDenseMap &known) { + assert(data.isEmpty()); + auto res = op.getResult(); + auto resType = dyn_cast(op.getResult().getType()); + + assert( + llvm::all_of(resType.getShape(), [](int64_t dim) { return dim == 1; })); + assert(isa(resType.getElementType()) || + isa(resType.getElementType())); + + size_t loopLimit = resType.getShape().size(); + SmallVector indices; + + for (auto i = 0; i < loopLimit; i++) { + indices.push_back(rewriter.create(loc, 0)); + } + auto extractOp = rewriter.create(loc, res, indices); + OpFoldResult IndexOfr = extractOp.getResult(); + if (isa(extractOp.getType())) { + IndexOfr = getOpFoldResultOfLayoutInfo(extractOp.getResult(), rewriter); + } + // Set scalar for mul state + data.setScalar(IndexOfr); + + for (auto i = 0; i < loopLimit; i++) { + // Add original dense val to first dim offset for add state + if (i == 0) { + data.getOffsetsRef().push_back(IndexOfr); + } else { + data.getOffsetsRef().push_back(rewriter.getIndexAttr(0)); + } + data.getSizesRef().push_back(rewriter.getIndexAttr(resType.getShape()[i])); + data.getStridesRef().push_back(rewriter.getIndexAttr(0)); + } +} + +void BlockDataParser::rewriteAddPtr( + triton::AddPtrOp op, triton::AddPtrOp::Adaptor &adaptor, + ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &known) { + auto insertPoint = rewriter.saveInsertionPoint(); + rewriter.setInsertionPoint(op); + + BlockData data; + parseAddPtr(op, data, op.getLoc(), rewriter, known); + + if (auto src = data.getSource(); + data.getMemAccTypeRef().isUnstructured() && + !(src && isa_and_nonnull(src.getDefiningOp()))) { + // TODO: Based on more info, try to create a performant IR + rewriteAddPtrToUnstrucMemAcc(op, adaptor, rewriter, data); + LLVM_DEBUG({ llvm::dbgs() << *getModuleOpFromOperation(op) << "\n"; }); + return; + } + + if (data.getSizesRef().size() == 0) { + data.getSizesRef().push_back(rewriter.getIndexAttr(1)); + data.getStridesRef().push_back(rewriter.getIndexAttr(0)); + data.getOffsetsRef().push_back(data.getScalarRef()); + } + + ArrayRef resultShape; + // shape {1,} is stub for single ptr + SmallVector stubScalarTypeShape(1, 1); + if (auto shapedType = dyn_cast(op.getResult().getType())) { + resultShape = shapedType.getShape(); + } else { + assert(data.getRank() == 1); + resultShape = stubScalarTypeShape; + } + + known[op.getResult()] = data; + + // If there are dimensions with size 1 and stride 0, replace 0 stride with the + // product of sizes of all lower dimensions. This avoids creating memref with + // zero stride. + // And here store the unmodified state into known ptrs, since any following + // pointer arithmetic operations should still use the original 0 stride. + auto inferedSize = 1; + auto hoistDim = op->getAttrOfType("hoist_dim"); + for (int i = data.getSizesRef().size() - 1; i >= 0; i--) { + auto strideConst = getConstantIntValue(data.getStridesRef()[i]); + auto sizeConst = getConstantIntValue(data.getSizesRef()[i]); + assert(sizeConst.has_value()); + bool shouldReplaceStride = + (sizeConst.value() == 1) || (hoistDim && hoistDim.getValue() == i); + if (shouldReplaceStride && strideConst && strideConst.value() == 0) { + data.getStridesRef()[i] = rewriter.getIndexAttr(inferedSize); + } + inferedSize *= sizeConst.value(); + } + + /* if (auto intToPtrOp = + dyn_cast(data.getSourceRef().getDefiningOp())) { + auto rtype = cast(intToPtrOp.getResult().getType()); + auto memrefType = + MemRefType::get({ShapedType::kDynamic}, rtype.getPointeeType()); + auto hivmPointCastOp = rewriter.create( + intToPtrOp.getLoc(), memrefType, ValueRange{intToPtrOp.getSrc()}); + data.setSource(hivmPointCastOp.getResult()); + }*/ + + // this modify is for null data.getSourceRef().getDefiningOp() + if (data.getSourceRef().getDefiningOp()) { + if (auto intToPtrOp = + dyn_cast(data.getSourceRef().getDefiningOp())) { + auto rtype = cast(intToPtrOp.getResult().getType()); + auto memrefType = + MemRefType::get({ShapedType::kDynamic}, rtype.getPointeeType()); + auto hivmPointCastOp = rewriter.create( + intToPtrOp.getLoc(), memrefType, ValueRange{intToPtrOp.getSrc()}); + data.setSource(hivmPointCastOp.getResult()); + } + } + + if (data.hasResElemTy()) { + // Handle bitcast scenario + auto memrefType = dyn_cast(data.getSourceRef().getType()) + .cloneWith(std::nullopt, data.getResElemTyRef()); + UnrealizedConversionCastOp castOp = + rewriter.create( + op.getLoc(), memrefType, data.getSourceRef()); + data.setSource(castOp.getOutputs()[0]); + } + + // ToDo: need to handle module scenario + + memref::ReinterpretCastOp castOp = + data.createCastOp(resultShape, op.getLoc(), rewriter); + Value src = castOp.getResult(); + LLVM_DEBUG({ + llvm::dbgs() << "cast MemRefType:\n"; + castOp.getOperation()->print(llvm::dbgs(), + OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); + + rewriter.replaceOp(op, src); + rewriter.restoreInsertionPoint(insertPoint); +} + +OpFoldResult +accumulatePotentialOffsetOnBase(triton::MakeTensorPtrOp op, Value base, + OpFoldResult offset, + ConversionPatternRewriter &rewriter) { + if (auto baseRecast = base.getDefiningOp()) { + assert(isa(op.getBase().getDefiningOp()) && + "base of MakeTensorPtrOp only comes from native ptr or AddPtrOp"); + + return addOpFoldResult(offset, baseRecast.getConstifiedMixedOffset(), + op.getLoc(), rewriter); + } + + return offset; +} + +// Design for load/store boundary_check. +memref::ReinterpretCastOp createRedundantOp(triton::MakeTensorPtrOp op, + ConversionPatternRewriter &rewriter, + BlockData &data) { + auto loc = op.getLoc(); + // to do boundary_check in tt.load, we need to keep the parent tensor's + // shape info in the IR. + // use the parent tensor's shape to create a cast + auto resultSizes = data.getSizes(); + auto resultOffsets = data.getOffsets(); + data.getSizesRef().clear(); + data.getOffsetsRef().clear(); + data.getSizesRef() = + std::move(llvm::map_to_vector(op.getShape(), [&](Value v) { + return getOpFoldResultOfLayoutInfo(v, rewriter); + })); + + // This redundant ReinterpretCastOp is to describe full tensor_ptr, so each + // dim offset from base is initialized as zero. + SmallVector curOffsets(op.getOffsets().size(), + rewriter.getIndexAttr(0)); + // Just accumulate base potential offset + curOffsets.front() = accumulatePotentialOffsetOnBase( + op, rewriter.getRemappedValue(op.getBase()), curOffsets.front(), + rewriter); + + for (auto offset : curOffsets) { + data.getOffsetsRef().push_back(offset); + } + + SmallVector staticShapes; + SmallVector dynamicShapes; + dispatchIndexOpFoldResults(data.getSizesRef(), dynamicShapes, staticShapes); + auto castOp = data.createCastOp(staticShapes, loc, rewriter); + // restore sizes and offsets + data.getSizesRef().clear(); + for (auto &s : resultSizes) { + data.getSizesRef().push_back(s); + } + data.getOffsetsRef().clear(); + for (auto &offset : resultOffsets) { + data.getOffsetsRef().push_back(offset); + } + return castOp; +} + +void BlockDataParser::rewriteMakeTensorPtrOp( + triton::MakeTensorPtrOp op, Value base, ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &known) { + Location loc = op.getLoc(); + BlockData data; + + auto orderSize = op.getOrder().size(); + if (orderSize > 1) { + // Declaration of llvm::ArrayRef::slice(n, m) + // - Chop off the first N elements of the array, and keep M elements + // in the array. + // Take care that 'm' means chunk length + for (auto [first, second] : + llvm::zip(op.getOrder().slice(0, orderSize - 1), + op.getOrder().slice(1, orderSize - 1))) { + assert(first == second + 1 && + "Currently only support default order on block pointers"); + } + } + + // Handle base is defined by tt.bitcast + BlockDataParser::parse(op.getBase(), data, loc, rewriter, known); + if (data.hasResElemTy()) { + auto memrefType = dyn_cast(data.getSourceRef().getType()) + .cloneWith(std::nullopt, data.getResElemTyRef()); + UnrealizedConversionCastOp castOp = + rewriter.create(loc, memrefType, + data.getSourceRef()); + data.setSource(castOp.getOutputs()[0]); + } else { + data.setSource(rewriter.getRemappedValue(op.getBase())); + } + + data.getOffsetsRef() = + std::move(llvm::map_to_vector(op.getOffsets(), [&](Value v) { + return getOpFoldResultOfLayoutInfo(v, rewriter); + })); + data.getStridesRef() = + std::move(llvm::map_to_vector(op.getStrides(), [&](Value v) { + return getOpFoldResultOfLayoutInfo(v, rewriter); + })); + + SmallVector newOffsets; + for (auto [offset, stride] : + llvm::zip(data.getOffsetsRef(), data.getStridesRef())) + newOffsets.push_back(mulOpFoldResult(offset, stride, loc, rewriter)); + + // 1. Consider that current base ptr may comes from `triton::AddPtrOp`, + // which have been converted to `memref::ReinterpretCastOp` with 1D + // shape([1,]) by `AddPtrConverter`. + // 2. While here would also convert `triton::MakeTensorPtrOp` to + // `memref::ReinterpretCastOp`, it will create use-def on double recast + // which means offset&size&stride info of first one will be dropped in terms + // of memref recast op fold specification. + // + // Conclusion with above two: + // Base of MakeTensorPtrOp has been seen as origin base, so it should + // reserve offset of first recast if it exists. + // Here extract the offset of first recast and add it to highest dimension + newOffsets.front() = + accumulatePotentialOffsetOnBase(op, base, newOffsets.front(), rewriter); + + data.getOffsetsRef().clear(); + + for (auto offset : newOffsets) { + data.getOffsetsRef().push_back(offset); + } + + ArrayRef resultShape; + auto pointerType = cast(op.getResult().getType()); + if (auto shapedType = dyn_cast(pointerType.getPointeeType())) { + resultShape = shapedType.getShape(); + data.getSizesRef().clear(); + for (auto dim_size : resultShape) { + data.getSizesRef().push_back( + IntegerAttr::get(IntegerType::get(op.getContext(), 64), dim_size)); + } + } else { + // scalar pointer, should produce a one dimensional memref + SmallVector scalarShape(1, 1); + resultShape = scalarShape; + assert(data.getRank() == 1); + } + + known[op.getResult()] = data; + + // special handling for davinci + // create redundant reinterpret_cast op for record shape info + auto redundantOp = createRedundantOp(op, rewriter, data); + redundantOp->setAttr("tensor_ptr_full_shape", rewriter.getUnitAttr()); + + // create reinterpret_cast op for the target block + data.setSource(redundantOp.getResult()); + auto castOp = data.createCastOp(resultShape, loc, rewriter); + rewriter.replaceOp(op, castOp.getResult()); + + if (nd2nzFlag) { + auto basePtr = castOp.getResult(); + int original_rank = op.getShape().size() + 1; + std::string shapeStr; + + auto baseMemrefType = mlir::dyn_cast(basePtr.getType()); + assert(baseMemrefType && "basePtr is not a memref type"); + auto shape = baseMemrefType.getShape(); + + if (auto memrefType = mlir::dyn_cast(basePtr.getType())) { + for (auto dim : memrefType.getShape()) { + shapeStr += llvm::formatv("_{0}", dim); + } + } + std::string elemTypeName; + Type elemType = baseMemrefType.getElementType(); + if (auto intType = mlir::dyn_cast(elemType)) { + elemTypeName = llvm::formatv("i{0}", intType.getWidth()); + } else if (auto floatType = mlir::dyn_cast(elemType)) { + std::string floatTypeName; + llvm::raw_string_ostream os(floatTypeName); + floatType.print(os); + os.flush(); + elemTypeName = floatTypeName; + } else { + std::string typeName; + llvm::raw_string_ostream os(typeName); + elemType.print(os); + os.flush(); + elemTypeName = typeName; + } + + std::string memrefTypeStr; + llvm::raw_string_ostream os(memrefTypeStr); + baseMemrefType.print(os); + os.flush(); + + std::string laydbgsuffix; + for (char c : memrefTypeStr) { + if ((c >= '0' && c <= '9') || (c >= 'a' && c <= 'z') || + (c >= 'A' && c <= 'Z') || c == '_' || c == ',' || c == '[' || + c == ']') { + laydbgsuffix += c; + } + } + auto funcName = rewriter.getStringAttr( + llvm::formatv("__hmf_original_shape{0}d{1}_{2}_{3}", original_rank, + shapeStr, elemTypeName, laydbgsuffix)); + MemRefType targetMemrefType = MemRefType::get( + baseMemrefType.getShape(), baseMemrefType.getElementType(), + baseMemrefType.getLayout()); + const int vectorSize = 4; + SmallVector srcElemTys; + for (auto sz : op.getShape()) { + srcElemTys.push_back(sz.getType()); + } + srcElemTys.push_back(targetMemrefType); + Type dstElemTy = rewriter.getNoneType(); + FunctionType hintFuncType = + FunctionType::get(rewriter.getContext(), srcElemTys, {dstElemTy}); + + auto mod = SymbolTable::getNearestSymbolTable(op); + auto extFunc = dyn_cast_or_null( + SymbolTable::lookupSymbolIn(mod, funcName)); + SmallVector args; + for (auto sz : op.getShape()) { + args.push_back(sz); + } + args.push_back(basePtr); + if (!extFunc) { + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPointToStart(&mod->getRegion(0).front()); + extFunc = rewriter.create(rewriter.getUnknownLoc(), + funcName, hintFuncType); + extFunc.setPrivate(); + extFunc->setAttr(LLVM::LLVMDialect::getReadnoneAttrName(), + UnitAttr::get(rewriter.getContext())); + rewriter.setInsertionPoint(op); + } + rewriter.create(loc, funcName, dstElemTy, args); + } +} + +void BlockDataParser::rewriteAdvanceOp( + triton::AdvanceOp op, ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &known) { + OpBuilder::InsertionGuard insertionGuard(rewriter); + rewriter.setInsertionPoint(op); + auto loc = op.getLoc(); + + BlockData blockData; + parse(op.getOperand(0), blockData, loc, rewriter, known); + + // region [BUGFIX] Add the code block below following the same logic as + // 'BlockDataParser::rewriteAddPtr' function. + known[op.getResult()] = blockData; + auto inferedSize = 1; + for (int i = blockData.getSizesRef().size() - 1; i >= 0; i--) { + auto strideConst = getConstantIntValue(blockData.getStridesRef()[i]); + auto sizeConst = getConstantIntValue(blockData.getSizesRef()[i]); + assert(sizeConst.has_value()); + if (sizeConst.value() == 1 && strideConst && strideConst.value() == 0) { + blockData.getStridesRef()[i] = rewriter.getIndexAttr(inferedSize); + } + inferedSize *= sizeConst.value(); + } + // endregion + + SmallVector incrementOffsets = + llvm::map_to_vector(op.getOffsets(), [&](Value offset) { + return getOpFoldResultOfLayoutInfo(offset, rewriter); + }); + + SmallVector newOffsets; + for (const auto [increment, originalOffset, stride] : + llvm::zip(incrementOffsets, blockData.getOffsetsRef(), + blockData.getStridesRef())) { + auto curDimOffset = + addOpFoldResult(mulOpFoldResult(increment, stride, loc, rewriter), + originalOffset, loc, rewriter); + + newOffsets.push_back(curDimOffset); + } + + blockData.getOffsetsRef().clear(); + + for (auto offset : newOffsets) + blockData.getOffsetsRef().push_back(offset); + + SmallVector scalarShape(1, 1); // Stub shape + ArrayRef resultShape; + auto pointerType = cast(op.getResult().getType()); + + if (auto shapedType = dyn_cast(pointerType.getPointeeType())) { + resultShape = shapedType.getShape(); + } else { + // scalar pointer, should produce a one dimensional memref + resultShape = scalarShape; + assert(blockData.getRank() == 1); + } + + auto newOp = blockData.createCastOp(resultShape, loc, rewriter); + rewriter.replaceOp(op, newOp.getResult()); + + known[newOp.getResult()] = blockData; +} + +template +std::enable_if_t || + std::is_same_v> +BlockDataParser::rewriteTerminator( + T op, ConversionPatternRewriter &rewriter, + const llvm::SmallDenseSet &blockArgIdxSet, + ArrayRef iterArgIdxMap, + const llvm::SmallDenseMap &known) { + // Any inserted instruction should be before this yield + OpBuilder::InsertionGuard insertionGuard{rewriter}; + rewriter.setInsertionPoint(op); + + auto adaptor = typename T::Adaptor(op); + ValueRange args; + if constexpr (std::is_same_v) { + args = adaptor.getOperands(); + } else { + args = adaptor.getArgs(); + } + + SmallVector initArgState; + SmallVector operands; + + operands.reserve(op->getNumOperands()); + for (const auto &[oper, newIterArgIdx] : + llvm::zip_equal(args, iterArgIdxMap)) { + if (newIterArgIdx != -1) + operands.push_back(oper); + } + + // For each of the init arg that we added additional Values in for loop, we + // need to add corresponding Values as yield operands. The loop below gathers + // BlockData for those values. + for (auto [i, v] : llvm::enumerate(args)) { + if (auto mappedV = rewriter.getRemappedValue(v)) { + // If this value is a tensor of pointers produced by AddPtrOp, + // we should have already converted to a ReinterpretCastOp without + // layout information for the normal cases + if (v.getDefiningOp() || + v.getDefiningOp() || + v.getDefiningOp()) { + if (auto castOp = mappedV.getDefiningOp()) { + v = castOp; + } else { + llvm_unreachable("mapped value defined by an unexpected op"); + } + } else { + // If this value is not a tensor of pointers, we will use the + // mapped value, and rely on the conversion will happen later + // automatically when we legalize loop body. + + // TODO: + // The scenario where a value is a tensor of pointers but not + // produced by AddPtrOp is not supported + if (isa(mappedV.getType()) && + isa( + dyn_cast(mappedV.getType()).getElementType())) + llvm_unreachable("unsupported scenario where a value is a tensor of " + "pointers but not produced by AddPtrOp"); + v = mappedV; + } + } + + if (blockArgIdxSet.find(i) == blockArgIdxSet.end()) + continue; + + auto reintCastOp = v.getDefiningOp(); + assert( + reintCastOp || + (isa(v.getType()) && + isa(dyn_cast(v.getType()).getElementType()))); + + BlockData state; + if (reintCastOp) { + parseReinterpretCast(reintCastOp, state, op.getLoc(), rewriter, known); + } else { + parse(v, state, op.getLoc(), rewriter, known); + } + initArgState.push_back(state); + } + + // For each of the BlockData recorded in the last step, extract value + // that correspond to offset and stride for each dimension and append + // them to yield operands. + for (auto state : initArgState) { + for (auto offset : state.getOffsetsRef()) { + // offsets can be IntAttr zeroes, since reinterpret_cast collapses + // them for the input memref, and the for loop may not update + // offsets other than offsets[0]. Create constants Values for those + // zeroes. + if (isa(offset)) { + auto constOffset = offset.get(); + assert(isa(constOffset) && + dyn_cast(constOffset).getInt() == 0 && + "attribute offsets should be zeroes"); + auto constOp = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(0)); + operands.push_back(constOp.getResult()); + } else { + operands.push_back(offset.get()); + } + } + + for (OpFoldResult stride : state.getStridesRef()) { + if (isa(stride)) { + auto constStride = stride.get(); + assert(isa(constStride) && + dyn_cast(constStride).getInt() == 1 && + "attribute strides should be ones"); + auto constOp = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(1)); + operands.push_back(constOp.getResult()); + } else { + operands.push_back(stride.get()); + } + } + } + + // Yield is a terminator op that must be at the end of the function + rewriter.setInsertionPointAfter(op); + Operation *newOp; + if constexpr (std::is_same_v) { + newOp = rewriter.replaceOpWithNewOp(op, operands); + } else { + newOp = rewriter.replaceOpWithNewOp(op, op.getCondition(), + operands); + } + + assert(op->getNumResults() == 0); + + LLVM_DEBUG({ + llvm::dbgs() << "new terminator: "; + newOp->print(llvm::dbgs(), OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); +} + +// This function is util function for rewriteLoopOp that +// check if given regionIterArg is used with given condition +bool isUsedWithCondition(Value v, std::function cond, + int depth = 0) { + for (auto &use : v.getUses()) { + auto *user = use.getOwner(); + if (user->hasAttr(ConverterUtils::discreteAttrName)) + continue; + if (cond(&use)) + return true; + if (auto loopOp = dyn_cast(user); + loopOp && !loopOp->hasAttr("ExtractedLoadOrStore")) { + if (isUsedWithCondition(loopOp.getTiedLoopRegionIterArg(&use), cond, + depth + 1)) + return true; + } else if (auto yieldOp = dyn_cast(user); + yieldOp && !isa(user->getParentOp())) { + if (depth && isUsedWithCondition(yieldOp->getParentOp()->getResult( + use.getOperandNumber()), + cond, depth - 1)) + return true; + } else if (auto conditionOp = dyn_cast(user); + conditionOp && use.getOperandNumber() > 0) { + auto whileOp = cast(conditionOp->getParentOp()); + if (depth && + isUsedWithCondition(whileOp->getResult(use.getOperandNumber() - 1), + cond, depth - 1)) + return true; + if (isUsedWithCondition( + whileOp.getAfterArguments()[use.getOperandNumber() - 1], cond, + depth)) + return true; + } + for (auto res : user->getResults()) { + if (isUsedWithCondition(res, cond, depth)) + return true; + } + } + return false; +} + +// This function is util function for rewriteLoopOp that create value from data. +// Assume data is structured, and from regionIterArg from LoopLikeOpInterface. +// +// For example, +// +// %7 = scf.for %arg2 = %c0_i32 to %c3_i32 step %c1_i32 iter_args(%arg3 = %4) -> +// (tensor<128xi32>) : i32 { +// %8 = tt.addptr %5, %arg3 : tensor<128x!tt.ptr>, tensor<128xi32> +// ... +// } +// +// is converted to +// +// %7 = scf.for %arg2 = %c0_i32 to %c3_i32 step %c1_i32 iter_args(%arg3 = %4, +// %arg4 = %5, %arg5 = %6) -> (tensor<128xi32>) : i32 { +// %scalarOffset = arith.index_cast %arg4 : index to i32 +// %scalarStride = arith.index_cast %arg5 : index to i32 +// ... +// %newRes = arith.addi %offset, %stride : tensor<128xi32> +// %8 = tt.addptr %5, %newRes : tensor<128x!tt.ptr>, tensor<128xi32> +// } +Value createFromData(RankedTensorType resType, const BlockData &data, + const Location &loc, OpBuilder &builder, + bool isMaskIterArg) { + auto resShape = resType.getShape(); + Value newRes = nullptr; + for (size_t i = 0; i < resShape.size(); i++) { + auto axisType = + RankedTensorType::get({resShape[i]}, resType.getElementType()); + auto axisI32Type = + RankedTensorType::get({resShape[i]}, builder.getIntegerType(32)); + Value axisValue = + builder.create(loc, axisI32Type, 0, resShape[i]); + if (axisType != axisI32Type) { + axisValue = builder.create(loc, axisType, axisValue); + } + Value offset = cast(data.getOffset(i)); + Value offsetValue = builder.create( + loc, resType.getElementType(), offset); + offsetValue = builder.create(loc, axisType, offsetValue); + Value stride = cast(data.getStride(i)); + if (!isMaskIterArg) { + Value strideValue = builder.create( + loc, resType.getElementType(), stride); + strideValue = builder.create(loc, axisType, strideValue); + axisValue = builder.create(loc, axisValue, strideValue); + } + axisValue = builder.create(loc, axisValue, offsetValue); + + for (size_t j = 0; j < resShape.size(); j++) { + if (i != j) + axisValue = builder.create(loc, axisValue, j); + } + axisValue = builder.create(loc, resType, axisValue); + if (newRes) { + newRes = builder.create(loc, newRes, axisValue); + } else { + newRes = axisValue; + } + } + return newRes; +} + +void BlockDataParser::rewriteLoopOp( + LoopLikeOpInterface op, ConversionPatternRewriter &rewriter, + llvm::SmallDenseMap &known) { + SmallVector newInitArgs; + SmallVector iterArgIdxMap; + SmallVector maskIterArgs; + int64_t argCnt = 0; + + SmallVector, 5> initArgIndexIfBlockData; + SmallVector, 5> knownPtrsTmp; + llvm::SmallDenseSet blockArgIdxSet; + + // Create a new list of init args + for (auto [i, arg] : llvm::enumerate(op.getInits())) { + auto mappedV = rewriter.getRemappedValue(arg); + memref::ReinterpretCastOp reintCastOp; + maskIterArgs.push_back(false); + + // If this init arg is supposed to be remapped, use the remapped + // value instead. + // In addition, if this init arg is a memref created by a reinterpret_cast + // or a tensor of index, there is a chance that it will be used in addptr. + // Create BlockData for each such init arg. + if (mappedV) { + // TODO: + // Passing a block argument pointer directly into a for loop not + // supported. + assert(!(isa(mappedV) && + isa(mappedV.getType())) && + "cannot take pointer block argument as init arg for for loop"); + if (auto reinterpretCastOp = + mappedV.getDefiningOp()) { + // Record memref::ReinterpretCastOp + reintCastOp = reinterpretCastOp; + newInitArgs.push_back(mappedV); + iterArgIdxMap.push_back(argCnt++); + } else { + newInitArgs.push_back(mappedV); + iterArgIdxMap.push_back(argCnt++); + } + } else { + newInitArgs.push_back(arg); + iterArgIdxMap.push_back(argCnt++); + } + + auto indexTensor = + isa(arg.getType()) && + isa(cast(arg.getType()).getElementType()) && + cast(cast(arg.getType()).getElementType()) + .getWidth() != 1 && + isUsedWithCondition(op.getRegionIterArgs()[i], [](OpOperand *use) { + auto *user = use->getOwner(); + return isa(user) || + (isa(user) && use->getOperandNumber() == 1) || + (isa(user) && use->getOperandNumber() == 2); + }); + + // Handle memref::ReinterpretCastOp and tensor specially + if (!reintCastOp && !indexTensor) + continue; + + BlockData data; + if (reintCastOp) { + parseReinterpretCast(reintCastOp, data, op.getLoc(), rewriter, + llvm::SmallDenseMap(0)); + } else { + parse(arg, data, op.getLoc(), rewriter, + llvm::SmallDenseMap(0)); + } + + maskIterArgs[i] = + indexTensor && + isUsedWithCondition(op.getRegionIterArgs()[i], [](OpOperand *use) { + auto *user = use->getOwner(); + return (isa(user) && use->getOperandNumber() == 1) || + (isa(user) && use->getOperandNumber() == 2); + }); + + if (indexTensor) { + newInitArgs.back() = nullptr; + iterArgIdxMap.back() = -1; + argCnt--; + } + + // Record the BlockData for later processing + initArgIndexIfBlockData.push_back(std::make_pair(i, data)); + } + + // Set insertion point to be before the for loop for new variables passed + // into the new loop. + auto origIp = rewriter.saveInsertionPoint(); + rewriter.setInsertionPoint(op); + + // For each of the BlockData recorded in the last step, insert new + // instructions to describe offset and stride for each dimension and append + // them to init args + for (auto [i, data] : initArgIndexIfBlockData) { + // For each dimension, if the corresponding offset and stride is an + // integer attribute, create a constant value and append them at the + // end of init arg list, which is prepared for calculate layout info with + // loop interation index + for (auto &dataOffset : data.getOffsetsRef()) { + if (isa(dataOffset)) { + auto constDataOffset = dataOffset.get(); + assert(isa(constDataOffset)); + auto constOp = rewriter.create( + op.getLoc(), rewriter.getIndexAttr( + dyn_cast(constDataOffset).getInt())); + newInitArgs.push_back(constOp.getResult()); + dataOffset = constOp.getResult(); + } else { + assert(isa(dataOffset.get().getType())); + newInitArgs.push_back(dataOffset.get()); + } + } + + for (auto &dataStride : data.getStridesRef()) { + if (isa(dataStride)) { + auto constDataStride = dataStride.get(); + assert(isa(constDataStride)); + auto constOp = rewriter.create( + op.getLoc(), rewriter.getIndexAttr( + dyn_cast(constDataStride).getInt())); + newInitArgs.push_back(constOp.getResult()); + dataStride = constOp.getResult(); + } else { + assert(isa(dataStride.get().getType())); + newInitArgs.push_back(dataStride.get()); + } + } + + // Note that we want the knownPtrs to be indexed by block arg, but we + // only have index for now. Also, the blockdata we record is the init + // arg, but want to to use newly created block arg. These block args + // are not created yet. We will translate this mapping later. + knownPtrsTmp.push_back(std::make_pair(i, data)); + blockArgIdxSet.insert(i); + + // If the original init arg is a memref produced by reinterpret_cast, + // create a new memref using new strides and offsets created above. + // This produces a canonicalized memref, which will match what the + // for loop generates if it modifies the memref. E.g., original + // reinterpret_cast can produce a memref with const stride: + // - memref<4x256xbf16, affine_map<(d0, d1)[s0, s1] -> (d0 * 256 + + // s0 + d1 + // * s1)>> + // The new reinterpret_cast will always have dynamic stride and + // offset: + // - memref<4x256xbf16, affine_map<(d0, d1)[s0, s1, s2] -> (d0 * s1 + // + s0 + d1 * s2)>> + if (newInitArgs[i] && + newInitArgs[i].getDefiningOp()) { + SmallVector resultShape; + for (auto size : data.getSizesRef()) { + auto constSize = getConstantIntValue(size); + assert(constSize && "expected constant size"); + resultShape.push_back(constSize.value()); + } + + // In current block data layout info, strides and offsets must be dynamic + // value + auto castOp = data.createCastOp(resultShape, op.getLoc(), rewriter); + if (resultShape.size() > 1) { + auto originalOffset = dyn_cast(data.getOffsetsRef()[0]); + for (auto &offsets : newInitArgs) { + if (offsets == originalOffset) { + offsets = castOp.getOffsets()[0]; + break; + } + } + data.getOffsetsRef()[0] = castOp.getOffsets()[0]; + } + + LLVM_DEBUG({ + llvm::dbgs() << "new reinterpret_cast with dynamic sizes " + "and offsets:"; + castOp->print(llvm::dbgs(), OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); + + newInitArgs[i] = castOp.getResult(); + } + } + + rewriter.restoreInsertionPoint(origIp); + IRMapping mapping; + + // Create a new LoopOp that uses updated init args and same loop body + LoopLikeOpInterface newOp; + auto newInits = to_vector( + make_filter_range(newInitArgs, [](Value v) { return v != nullptr; })); + auto commonBodyBuilder = [&](OpBuilder &b, Location loc, bool useInit, + ValueRange newRegionArgs, Region ®ion, + Block::BlockArgListType regionArgs, + ArrayRef isUsedForRegionArgs, + ArrayRef maskIterArgs) { + auto newArgIter = newRegionArgs.begin(); + for (const auto &[regionArg, isUsedForRegionArg] : + llvm::zip(regionArgs, isUsedForRegionArgs)) { + if (isUsedForRegionArg) { + mapping.map(regionArg, *newArgIter); + ++newArgIter; + } + } + + // Convert the book-keeping data structure to use the correct key and value. + // Key is converted from init arg index to newly created block arg, and + // Value's BlockData fields are converted from init arg to newly created + // block arg + + // TODO: remove (useInit = true) logic after supporting make_tensor_ptr + if (useInit) { + for (auto [i, data] : knownPtrsTmp) { + for (auto &offset : data.getOffsetsRef()) { + offset = *newArgIter; + ++newArgIter; + } + + for (auto &stride : data.getStridesRef()) { + stride = *newArgIter; + ++newArgIter; + } + + auto regionArg = regionArgs[i]; + auto key = mapping.lookupOrNull(regionArg); + if (!key) { + // Create IndexTensor regionArg from computed offset and stride data + key = createFromData(cast(regionArg.getType()), + data, op.getLoc(), rewriter, maskIterArgs[i]); + mapping.map(regionArg, key); + } + known.insert(std::make_pair(key, data)); + } + } else { + for (auto [i, isUsedForRegionArg] : + llvm::enumerate(isUsedForRegionArgs)) { + if (!isUsedForRegionArg) { + BlockData data; + auto regionArg = regionArgs[i]; + auto regionArgType = cast(regionArg.getType()); + data.getOffsetsRef().resize(regionArgType.getRank()); + data.getStridesRef().resize(regionArgType.getRank()); + for (auto &offset : data.getOffsetsRef()) { + offset = *newArgIter; + ++newArgIter; + } + for (auto &dim : regionArgType.getShape()) { + data.getSizesRef().push_back(rewriter.getIndexAttr(dim)); + } + for (auto &stride : data.getStridesRef()) { + stride = *newArgIter; + ++newArgIter; + } + + auto key = mapping.lookupOrNull(regionArg); + if (!key) { + // Create IndexTensor regionArg from computed offset and stride data + key = createFromData(regionArgType, data, op.getLoc(), rewriter, + maskIterArgs[i]); + mapping.map(regionArg, key); + } + known.insert(std::make_pair(key, data)); + } + } + } + + for (auto &bodyOp : region.getOps()) + b.clone(bodyOp, mapping); + }; + for (const auto &[initArg, newInitArg] : + llvm::zip(op.getInits(), newInitArgs)) { + if (newInitArg) { + mapping.map(initArg, newInitArg); + } + } + if (auto forOp = dyn_cast(op.getOperation())) { + SmallVector usedForRegionArgs; + for (auto newInitArg : newInitArgs) { + usedForRegionArgs.push_back(newInitArg ? true : false); + } + newOp = rewriter.create( + forOp.getLoc(), forOp.getLowerBound(), forOp.getUpperBound(), + forOp.getStep(), newInits, + [&](OpBuilder &b, Location loc, Value iv, ValueRange args) { + mapping.map(forOp.getInductionVar(), iv); + commonBodyBuilder(b, loc, true, args, forOp.getRegion(), + op.getRegionIterArgs(), usedForRegionArgs, + maskIterArgs); + }); + + // Replace only the results that correspond to the original scf.for + auto newResultIter = newOp->result_begin(); + rewriter.setInsertionPointAfter(newOp); + for (const auto &[res, regionArg, newIterArgIdx, mask] : + llvm::zip_equal(op->getResults(), op.getRegionIterArgs(), + iterArgIdxMap, maskIterArgs)) { + if (newIterArgIdx != -1) { + rewriter.replaceAllUsesWith(res, *newResultIter); + ++newResultIter; + } else { + auto key = mapping.lookup(regionArg); + auto data = known.at(key); + for (auto &offset : data.getOffsetsRef()) + offset = + newOp.getTiedLoopResult(cast(offset.get())); + for (auto &stride : data.getStridesRef()) + stride = + newOp.getTiedLoopResult(cast(stride.get())); + auto newRes = + createFromData(cast(regionArg.getType()), data, + op.getLoc(), rewriter, mask); + rewriter.replaceAllUsesWith(res, newRes); + } + } + } else if (auto whileOp = dyn_cast(op.getOperation())) { + SmallVector resultTypes; + SmallVector usedForBeforeRegionArgs; + SmallVector usedForAfterRegionArgs; + llvm::SmallDenseSet blockArgIdxSetForAfter; + SmallVector iterArgIdxMapForAfter; + SmallVector maskIterArgsForAfter(whileOp->getNumResults()); + + int64_t indexCnt = 0; + + for (auto newInitArg : newInitArgs) { + usedForBeforeRegionArgs.push_back(newInitArg ? true : false); + } + for (size_t i = 0; i < whileOp->getNumResults(); i++) { + auto resType = whileOp->getResultTypes()[i]; + auto indexTensor = + isa(resType) && + isa(cast(resType).getElementType()) && + isUsedWithCondition(whileOp.getAfterArguments()[i], + [](OpOperand *use) { + auto *user = use->getOwner(); + return isa(user) || + (isa(user) && + use->getOperandNumber() == 1) || + (isa(user) && + use->getOperandNumber() == 2); + }); + if (indexTensor) { + indexCnt += 2 * cast(resType).getRank(); + usedForAfterRegionArgs.push_back(false); + iterArgIdxMapForAfter.push_back(-1); + maskIterArgsForAfter[i] = isUsedWithCondition( + whileOp.getAfterArguments()[i], [](OpOperand *use) { + auto *user = use->getOwner(); + return (isa(user) && + use->getOperandNumber() == 1) || + (isa(user) && + use->getOperandNumber() == 2); + }); + blockArgIdxSetForAfter.insert(i); + } else { + resultTypes.push_back(resType); + usedForAfterRegionArgs.push_back(true); + iterArgIdxMapForAfter.push_back(argCnt++); + } + } + resultTypes.append(indexCnt, rewriter.getIndexType()); + newOp = rewriter.create( + whileOp.getLoc(), resultTypes, newInits, + [&](OpBuilder &b, Location loc, ValueRange args) { + commonBodyBuilder(b, loc, true, args, whileOp.getBefore(), + whileOp.getBeforeArguments(), + usedForBeforeRegionArgs, maskIterArgs); + }, + [&](OpBuilder &b, Location loc, ValueRange args) { + commonBodyBuilder(b, loc, false, args, whileOp.getAfter(), + whileOp.getAfterArguments(), usedForAfterRegionArgs, + maskIterArgsForAfter); + }); + + auto newResultIter = newOp->result_begin(); + rewriter.setInsertionPointAfter(newOp); + for (const auto &[res, regionArg, newIterArgIdx, mask] : + llvm::zip_equal(op->getResults(), whileOp.getAfterArguments(), + iterArgIdxMapForAfter, maskIterArgsForAfter)) { + if (newIterArgIdx != -1) { + rewriter.replaceAllUsesWith(res, *newResultIter); + ++newResultIter; + } else { + auto key = mapping.lookup(regionArg); + auto data = known.at(key); + for (auto &offset : data.getOffsetsRef()) + offset = newOp->getResult( + cast(offset.get()).getArgNumber()); + for (auto &stride : data.getStridesRef()) + stride = newOp->getResult( + cast(stride.get()).getArgNumber()); + auto newRes = + createFromData(cast(regionArg.getType()), data, + op.getLoc(), rewriter, mask); + rewriter.replaceAllUsesWith(res, newRes); + } + } + + auto conditionOp = + cast(newOp.getOperation()).getConditionOp(); + rewriteTerminator(conditionOp, rewriter, blockArgIdxSetForAfter, + iterArgIdxMapForAfter, known); + } + + // Copy all attributes from op to newOp + newOp->setAttrs(op->getAttrs()); + rewriter.eraseOp(op); + + // Update the loop body. Manually invoke the rewrite logic on addptr and yield + // in the loop body, so we can take advantage of the states we built up + for (auto *region : newOp.getLoopRegions()) { + for (auto &bodyOp : region->getOps()) { + if (auto addptrOp = dyn_cast(bodyOp)) { + // FIXME: Constructed adaptor here does not hold the transformed op + // info. + auto adaptor = triton::AddPtrOp::Adaptor(addptrOp); + rewriteAddPtr(addptrOp, adaptor, rewriter, known); + } else if (auto advanceOp = dyn_cast(bodyOp)) { + rewriteAdvanceOp(advanceOp, rewriter, known); + } else if (auto makeTensorPtrOp = + dyn_cast(bodyOp)) { + ConversionPatternRewriter::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(makeTensorPtrOp); + rewriteMakeTensorPtrOp( + makeTensorPtrOp, + rewriter.getRemappedValue(makeTensorPtrOp.getBase()), rewriter, + known); + } else if (auto loopOp = dyn_cast(bodyOp); + loopOp && !loopOp->hasAttr("ExtractedLoadOrStore")) { + ConversionPatternRewriter::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(loopOp); + // Remove UnhandledLoopOp attr before process + loopOp->removeAttr("UnhandledLoopOp"); + rewriteLoopOp(loopOp, rewriter, known); + } + } + } + + if (!op.getRegionIterArgs().empty()) { + auto yieldOp = cast( + newOp.getLoopRegions().back()->back().getTerminator()); + rewriteTerminator(yieldOp, rewriter, blockArgIdxSet, iterArgIdxMap, known); + } + + LLVM_DEBUG({ + llvm::dbgs() << "new loop\n"; + newOp.getOperation()->print(llvm::dbgs(), + OpPrintingFlags().printGenericOpForm()); + llvm::dbgs() << "\n"; + }); +} + +/// @brief Rewrite the triton::AddPtrOp to handle unstructured memory access. +/// @param op The triton::AddPtrOp to be rewritten. +/// @param adaptor The adaptor of the triton::AddPtrOp, used to get operands. +/// @param rewriter The pattern rewriter used to modify the IR. +/// @param data The BlockData containing information about the memory access. +void BlockDataParser::rewriteAddPtrToUnstrucMemAcc( + triton::AddPtrOp op, triton::AddPtrOp::Adaptor &adaptor, + ConversionPatternRewriter &rewriter, BlockData &data) { + auto loc = op.getLoc(); + auto &offsets = data.getOffsetsRef(); + auto &blockSizes = data.getSizesRef(); + auto &strides = data.getStridesRef(); + Value ptrOffset = adaptor.getOffset(); + Value zeroIdx = + rewriter.create(loc, rewriter.getIndexAttr(0)); + Value oneIdx = + rewriter.create(loc, rewriter.getIndexAttr(1)); + auto addptrRes = op.getResult(); + assert(addptrRes.hasOneUse() && "Invalid: tt.addptr has multiple users"); + auto loadOp = *(addptrRes.user_begin()); + + // Prepare empty tensor for loop based scalar load + // FIXME: We use cast here because addptr must return tensor>. + // True? + auto resTy = cast(addptrRes.getType()); + auto resEPtrTy = resTy.getElementType(); + auto resETy = cast(resEPtrTy).getPointeeType(); + Value loaded = rewriter.create(loc, blockSizes, resETy); + SmallVector initArgs; + initArgs.push_back(loaded); + + SmallVector forLBs; + SmallVector forUBs; + SmallVector forSteps; + for (auto &s : offsets) { + forLBs.push_back(zeroIdx); + } + for (auto &s : blockSizes) { + forUBs.push_back(getValueOrCreateConstantIndexOp(rewriter, loc, s)); + } + for (auto &s : strides) { + forSteps.push_back(oneIdx); + } + SmallVector ivs; + OpBuilder builder(op); + auto loop = createNestedLoops( + builder, loc, 0, blockSizes.size(), forLBs, forUBs, forSteps, ivs, + initArgs, + [&](OpBuilder &bB, Location bLoc, SmallVector &allIVs, + ValueRange iterArgs) { + OpBuilder::InsertionGuard g(bB); + bB.setInsertionPointToStart(bB.getBlock()); + + Value scalarOffsetRaw = + bB.create(bLoc, ptrOffset, allIVs); + Value scalarOffset = bB.create( + bLoc, bB.getIndexType(), scalarOffsetRaw); + // Replace offset & size. Only single element. + data.getOffsetsRef().clear(); + data.getOffsetsRef().push_back(scalarOffset); + data.getSizesRef().clear(); + data.getSizesRef().push_back(bB.getIndexAttr(1)); + data.getStridesRef().clear(); + data.getStridesRef().push_back(bB.getIndexAttr(1)); + memref::ReinterpretCastOp castOp = data.createCastOp({1}, bLoc, bB); + rewriter.replaceOp(op, castOp); + // Move tt.load using this tt.addptr into this block + loadOp->moveAfter(castOp); + loadOp->setAttr("IndirectLoad", UnitAttr::get(op.getContext())); + bB.create(bLoc, iterArgs); + }); +} + +} // namespace triton +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/CMakeLists.txt new file mode 100755 index 00000000..aa268d86 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/CMakeLists.txt @@ -0,0 +1,32 @@ +add_triton_library(TritonToLinalgIncubated + TritonToLinalgIncubatedPass.cpp + LoadStoreConverter.cpp + FunctionConverter.cpp + ArgMinMaxConverter.cpp + TritonOpConverter.cpp + HoistBroadcast.cpp + BlockPtrAnalysis.cpp + MaskAnalysis.cpp + UseAnalysis.cpp + DescriptorConverter.cpp + + DEPENDS + TritonToLinalgIncubatedConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonAnalysis + MLIRTritonNPUUtils + MLIRSCFTransforms + MLIRLinalgTransforms + TleToLinalg +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/DescriptorConverter.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/DescriptorConverter.cpp new file mode 100755 index 00000000..fa16e469 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/DescriptorConverter.cpp @@ -0,0 +1,196 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/DescriptorConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/TritonOpConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/LogicalResult.h" +#include "llvm/Support/raw_ostream.h" +#include + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/Transforms/DialectConversion.h" + +namespace DescriptorConverter { +using namespace mlir; +using namespace triton; + +bool hasATensorDescriptorType(mlir::TypeRange types) { + return llvm::any_of(types, [](mlir::Type t) { + return llvm::isa(t); + }); +} + +/** + * @brief Filter out operand segment sizes from the list of attributes since + * this attribute is operation specific and shouldn't be set arbitrarily. + */ +mlir::SmallVector +filterSegmentSizes(mlir::ArrayRef attrs) { + mlir::SmallVector ret; + llvm::copy_if(attrs, std::back_inserter(ret), [](const NamedAttribute &attr) { + auto attrName = attr.getName().getValue(); + return attrName != "operandSegmentSizes"; + }); + return ret; +} + +Descriptor unpackDescriptor(TensorDescType type, Value desc, + ConversionPatternRewriter &rewriter) { + auto makeDescOp = desc.getDefiningOp(); + assert(makeDescOp && "Descriptor must be defined by MakeTensorDescOp"); + + Descriptor res; + + // 直接回溯处理的 tt.make_tensor_descriptor + res.base = makeDescOp.getBase(); + for (auto s : makeDescOp.getShape()) { + res.shape.push_back(rewriter.createOrFold( + makeDescOp.getLoc(), rewriter.getI64Type(), s)); + } + for (auto st : makeDescOp.getStrides()) { + res.strides.push_back(rewriter.createOrFold( + makeDescOp.getLoc(), rewriter.getI64Type(), st)); + } + + return res; +} + +SmallVector computeOrder(ArrayRef shape) { + SmallVector order; + int rank = shape.size(); + order.reserve(rank); + // 默认采用逆序 [dims - 1, ..., 0] + for (int i = rank - 1; i >= 0; --i) { + order.push_back(i); + } + return order; +} + +LogicalResult DescriptorLoadConverter::matchAndRewrite( + triton::DescriptorLoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + const auto blockShape = op.getDesc().getType().getBlockType().getShape(); + auto descTy = op.getDesc().getType(); + auto indices = op.getIndices(); + + // 1. 解包 descriptor + auto desc = unpackDescriptor(descTy, adaptor.getDesc(), rewriter); + + // 2. 新增 make_tensor_ptr + SmallVector tensorShapeValues; + for (auto dim : blockShape) { + tensorShapeValues.push_back(static_cast(dim)); + } + Value tensorPtr = rewriter.create( + loc, + desc.base, // 基址 + desc.shape, // 形状 + desc.strides, // 步长 + indices, // 偏移 + tensorShapeValues, // tensorShape + computeOrder(blockShape) // 使用动态计算的 order + ); + // 3. 替换 tt.load 操作 + auto newLoad = rewriter.replaceOpWithNewOp( + op, descTy.getSignlessBlockType(), tensorPtr); + + // 保留原始操作的其他属性 + newLoad->setAttrs(filterSegmentSizes(op->getAttrs())); + + return success(); +} + +LogicalResult DescriptorStoreConverter::matchAndRewrite( + triton::DescriptorStoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + const auto blockShape = op.getDesc().getType().getBlockType().getShape(); + auto descTy = op.getDesc().getType(); + auto indices = op.getIndices(); + + // 1. 解包 descriptor + auto desc = unpackDescriptor(descTy, adaptor.getDesc(), rewriter); + + // 2. 新增 make_tensor_ptr + SmallVector tensorShapeValues; + for (auto dim : blockShape) { + tensorShapeValues.push_back(static_cast(dim)); + } + Value tensorPtr = rewriter.create( + loc, + desc.base, // 基址 + desc.shape, // 形状 + desc.strides, // 步长 + indices, // 偏移 + tensorShapeValues, // tensorShape + computeOrder(blockShape) // 使用动态计算的 order + ); + + // 3. 替换 tt.store 操作 + Value valueToStore = adaptor.getSrc(); + + auto maskType = RankedTensorType::get(blockShape, rewriter.getI1Type()); + rewriter.create(loc, + DenseElementsAttr::get(maskType, true)); + + // 创建属性 + auto boundaryCheck = rewriter.getDenseI32ArrayAttr({}); // 空的边界检查 + auto cacheModifier = triton::CacheModifierAttr::get( + rewriter.getContext(), triton::CacheModifier::NONE); + auto evictionPolicy = triton::EvictionPolicyAttr::get( + rewriter.getContext(), triton::EvictionPolicy::NORMAL); + + // 创建 store 操作并替换原始操作 + auto newStore = + rewriter.replaceOpWithNewOp(op, // 要替换的操作 + tensorPtr, // 指针 + valueToStore, // 要存储的值 + nullptr, // 掩码 + boundaryCheck, // 边界检查 + cacheModifier, // 缓存修饰符 + evictionPolicy // 驱逐策略 + ); + + // 保留原始操作的其他属性 + newStore->setAttrs(filterSegmentSizes(op->getAttrs())); + return success(); +} + +} // namespace DescriptorConverter diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/FunctionConverter.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/FunctionConverter.cpp new file mode 100755 index 00000000..8353f765 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/FunctionConverter.cpp @@ -0,0 +1,56 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/FunctionConverter.h" + +namespace FunctionConverter { +using namespace mlir; +using namespace triton; + +LogicalResult GetProgramIDConverter::matchAndRewrite( + triton::GetProgramIdOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto axis = (uint32_t)op.getAxis(); + assert(axis < GetProgramIDConverter::LAUNCH_GRID_RANK && + "Invalid axis for GetProgramIdOp"); + auto func = op->getParentOfType(); + auto numArgs = func.getNumArguments(); + auto id = func.getArgument(numArgs - GetProgramIDConverter::LAUNCH_GRID_RANK + + axis); + rewriter.replaceOp(op, id); + return success(); +} + +LogicalResult GetNumProgramsConverter::matchAndRewrite( + triton::GetNumProgramsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto axis = (uint32_t)op.getAxis(); + assert(axis < GetNumProgramsConverter::LAUNCH_GRID_RANK && + "Invalid axis for GetNumProgramsOp"); + auto func = op->getParentOfType(); + auto numArgs = func.getNumArguments(); + auto id = func.getArgument( + numArgs - GetNumProgramsConverter::LAUNCH_GRID_RANK * 2 + axis); + rewriter.replaceOp(op, id); + return success(); +} +} // namespace FunctionConverter diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/HoistBroadcast.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/HoistBroadcast.cpp new file mode 100755 index 00000000..ddcf5da7 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/HoistBroadcast.cpp @@ -0,0 +1,229 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/HoistBroadcast.h" +#include "incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/LogicalResult.h" +#include "llvm/Support/raw_ostream.h" +#include + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/MemRef/Transforms/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/ValueRange.h" + +namespace HoistBroadcast { +using namespace mlir; +using namespace triton; + +LogicalResult +BroadcastConverter::matchAndRewrite(triton::BroadcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + assert(op->getNumResults() == 1 && "BroadcastOp assumes single result"); + + if (!isa(op.getType().getElementType())) { + return rewriter.notifyMatchFailure( + op, "only support hoist broadcast for tt.ptr tensor right now."); + } + auto loc = op.getLoc(); + BroadcastHoister hoister(op); + + if (!hoister.canBroadcast()) { + return failure(); + } + + if (hoister.parse(op.getSrc(), loc, rewriter).failed()) { + return failure(); + } + if (hoister.replaceBroadcastOp(op, rewriter).failed()) { + return failure(); + } + return success(); +} + +BroadcastHoister::BroadcastHoister(triton::BroadcastOp op) { + source = nullptr; + if (findSrc(op.getSrc()).failed()) { + LLVM_DEBUG({ llvm::dbgs() << "No legal source found for broadcast op\n"; }); + } + opToHoist = op; + auto resultType = dyn_cast(op.getType()); + for (size_t i = 0; i < resultType.getShape().size(); ++i) { + tensorSizes.push_back(resultType.getShape()[i]); + } +} + +LogicalResult BroadcastHoister::findSrc(Value operand) { + // ptr tensor can only be defined by AddPtrOp or SplatOp in this converter + // another broadcast to be complemented + if (auto op = operand.getDefiningOp()) { + return findSrc(op.getPtr()); + } else if (auto op = operand.getDefiningOp()) { + source = op.getSrc(); + return success(); + } else { + LLVM_DEBUG({ + llvm::dbgs() << "Unsupported operation in BroadcastHoister::findSrc: " + << *operand.getDefiningOp() << "\n"; + }); + return failure(); + } +} + +LogicalResult BroadcastHoister::parse(Value operand, const Location &loc, + ConversionPatternRewriter &rewriter) { + if (auto op = operand.getDefiningOp()) { + return parseAddptr(op, loc, rewriter); + } else if (auto op = operand.getDefiningOp()) { + return parseSplat(op, loc, rewriter); + } else if (auto op = operand.getDefiningOp()) { + return parseBroadcast(op, loc, rewriter); + } else { + // Handle other cases or throw an error + LLVM_DEBUG({ + llvm::dbgs() << "Unsupported operation in BroadcastHoister::parse: " + << *operand.getDefiningOp() << "\n"; + }); + return failure(); + } +} + +LogicalResult +BroadcastHoister::parseAddptr(triton::AddPtrOp addptrOp, const Location &loc, + ConversionPatternRewriter &rewriter) { + // Implementation for parsing AddptrOp + if (parse(addptrOp.getPtr(), loc, rewriter).failed()) { + return failure(); + } + auto broadcastedPtr = broadcastMap[addptrOp.getPtr()]; + + RankedTensorType offsetType = + dyn_cast(addptrOp.getOffset().getType()); + if (!offsetType || !offsetType.hasStaticShape()) { + LLVM_DEBUG({ + llvm::dbgs() << "Offset must be a ranked tensor with static shape.\n"; + }); + return failure(); + } + + auto elementType = offsetType.getElementType(); + auto broadcastType = RankedTensorType::get({tensorSizes}, elementType); + auto broadcastedOffset = rewriter.create( + loc, broadcastType, addptrOp.getOffset()); + + auto ptrType = dyn_cast(source.getType()); + auto ptrTensorType = RankedTensorType::get({tensorSizes}, ptrType); + + auto newAddPtrOp = rewriter.create( + loc, ptrTensorType, broadcastedPtr, broadcastedOffset); + + size_t hoistDim = -1; + for (size_t i = 0; i < offsetType.getShape().size(); ++i) { + if (offsetType.getShape()[i] == 1 && + offsetType.getShape()[i] != tensorSizes[i]) { + hoistDim = i; + break; + } + } + if (hoistDim == -1) { + LLVM_DEBUG({ + llvm::dbgs() << "No dimension to hoist found in AddPtrOp offset.\n"; + }); + } + newAddPtrOp->setAttr( + "hoist_dim", rewriter.getI64IntegerAttr(static_cast(hoistDim))); + broadcastMap[addptrOp.getResult()] = newAddPtrOp.getResult(); + return success(); +} + +LogicalResult +BroadcastHoister::parseSplat(triton::SplatOp splatOp, const Location &loc, + ConversionPatternRewriter &rewriter) { + // End of parse: splat for ptr + auto src = splatOp.getSrc(); + auto dst = splatOp.getResult(); + if (!isa(src.getType())) { + LLVM_DEBUG( + { llvm::dbgs() << "SplatOp source must be of pointer type.\n"; }); + return failure(); + } + source = src; + + auto ptrType = dyn_cast(source.getType()); + auto ptrTensorType = RankedTensorType::get({tensorSizes}, ptrType); + auto newSplatOp = rewriter.create(loc, ptrTensorType, src); + broadcastMap[splatOp.getResult()] = newSplatOp.getResult(); + return success(); +} + +LogicalResult +BroadcastHoister::parseBroadcast(triton::BroadcastOp broadcastOp, + const Location &loc, + ConversionPatternRewriter &rewriter) { + // Another broadcast for ptr tensor + // To be fixed if needed in future + LLVM_DEBUG({ + llvm::dbgs() << "Now cannot handle multi broadcast for ptr tensor.\n"; + }); + return failure(); +} + +LogicalResult +BroadcastHoister::replaceBroadcastOp(triton::BroadcastOp op, + ConversionPatternRewriter &rewriter) { + auto newOp = broadcastMap[op.getSrc()]; + rewriter.replaceOp(op, newOp); + return success(); +} + +bool BroadcastHoister::canBroadcast() { + auto resultType = dyn_cast(opToHoist.getType()); + auto srcType = dyn_cast(opToHoist.getSrc().getType()); + int broadcastedDims = 0; + for (size_t i = 0; i < resultType.getShape().size(); ++i) { + if (srcType.getShape()[i] == 1 && resultType.getShape()[i] != 1) { + broadcastedDims++; + } + } + if (broadcastedDims != 1) { + LLVM_DEBUG({ + llvm::dbgs() << "Now cannot handle broadcast for ptr tensor with multi " + "broadcasted dimension.\n"; + }); + return false; + } + return source != nullptr && isa(source.getType()); +} + +} // namespace HoistBroadcast diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/LoadStoreConverter.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/LoadStoreConverter.cpp new file mode 100755 index 00000000..39dac035 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/LoadStoreConverter.cpp @@ -0,0 +1,1281 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/LoadStoreConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h" +#include "incubated/Conversion/UtilsIncubated/InterleaveOptimization.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#if __has_include("bishengir/Dialect/Annotation/IR/Annotation.h") +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#endif +#if __has_include("bishengir/Dialect/HFusion/IR/HFusion.h") +#include "bishengir/Dialect/HFusion/IR/HFusion.h" +#endif +#if __has_include("bishengir/Dialect/HIVM/IR/HIVM.h") +#include "bishengir/Dialect/HIVM/IR/HIVM.h" +#endif + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Arith/Utils/Utils.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Location.h" +#include "mlir/IR/Matchers.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Value.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include "llvm/Support/Debug.h" + +#include +#include +#include + +#define DEBUG_TYPE "triton-load-store-converter" + +namespace LoadStoreConverter { +using namespace mlir; +using namespace triton; +using namespace mlir::triton::Incubated; +const std::string MayImplicitTransposeWithLastAxisTAG = + "MayImplicitTransposeWithLastAxis"; + +LogicalResult +AddPtrConverter::matchAndRewrite(triton::AddPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + llvm::SmallDenseMap known; + BlockDataParser::rewriteAddPtr(op, adaptor, rewriter, known); + return success(); +} + +LogicalResult LoadConverter::toTensorAndReplace( + triton::LoadOp &op, RankedTensorType &tensorType, Value localMem, + bool mayImplicitTransposeWithLastAxis, const Location &loc, + ConversionPatternRewriter &rewriter) const { + + Value loadedTensor = rewriter.create( + loc, tensorType, localMem, true, true); + if (mayImplicitTransposeWithLastAxis) { + auto markOp = rewriter.create(loc, loadedTensor); + markOp->setAttr(MayImplicitTransposeWithLastAxisTAG, + UnitAttr::get(rewriter.getContext())); + } + rewriter.replaceOp(op, loadedTensor); + return success(); +} + +/// @brief Check whether the triton::LoadOp has been modified to the specified +/// state by the AddPtrConverter. +/// @param op The triton::LoadOp operation to be checked. +/// @return Return success if the operation conforms to the specified state; +/// otherwise, return failure. +LogicalResult +LoadConverter::checkModifiedByAddPtrConverter(triton::LoadOp &op) const { + if (!isa(op->getParentOp())) { + return failure(); + } + if (!op->hasAttr("IndirectLoad")) { + return failure(); + } + auto ptrOp = op.getPtr().getDefiningOp(); + auto ptrBlock = ptrOp->getBlock(); + auto opBlock = op->getBlock(); + if (ptrBlock == opBlock) { + return failure(); + } + + return success(); +} + +/// @brief Continue to modify the triton::LoadOp from the state modified by the +/// AddPtrConverter. +/// @param op The triton::LoadOp operation to be processed. +/// @param adaptor The adaptor for the operation, used to obtain operands. +/// @param rewriter The pattern rewriter used to rewrite the operation. +/// @return Return success if the operation is successful; otherwise, return +/// failure. +LogicalResult LoadConverter::continueModifyFromAddPtrConverter( + triton::LoadOp &op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto forOp = op->getParentOfType(); + Operation *firstOp = &forOp.getBody()->front(); + auto extractOp = cast(firstOp); + auto ivs = extractOp.getIndices(); + // Single iterArg which is inserted by AddPtrConverter. + auto iterArg = forOp.getRegionIterArg(0); + auto ptr = adaptor.getPtr(); + + rewriter.setInsertionPointAfter(op); + Value castVal = ptr.getDefiningOp(); + Value idxZero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + Value loadVal = + rewriter.create(loc, castVal, ValueRange{idxZero}); + Value insertedVal = + rewriter.create(loc, loadVal, iterArg, ValueRange{ivs}); + // a yield op is already created by AddPtrConverter. + // so we need to replace it with a new yield op. + Operation *terminator = forOp.getBody()->getTerminator(); + scf::YieldOp oldYieldOp = cast(terminator); + auto yieldOp = rewriter.create(loc, ValueRange{insertedVal}); + rewriter.replaceOp(oldYieldOp, yieldOp); + // Now the scf.for is complete, we can replace tt.load with it. + auto rank = cast(op.getResult().getType()).getShape().size(); + Operation *rootForOp = op; + while (rank != 0) { + rank--; + rootForOp = rootForOp->getParentOfType(); + } + rewriter.replaceOp(op, rootForOp); + LLVM_DEBUG({ llvm::dbgs() << *getModuleOpFromOperation(rootForOp) << "\n"; }); + return success(); +} + +void LoadConverter::fillTensorWithOtherForMaskScenario( + Value other, Value localMem, ArrayRef maskDim, + ConversionPatternRewriter &rewriter) const { + auto loc = localMem.getLoc(); + MemRefType originalType = cast(localMem.getType()); + assert(originalType.hasStaticShape() && "only support static shape"); + assert(originalType.getRank() == maskDim.size() && + "shape and mask must have same rank"); + + auto fillFlag = + rewriter.create(loc, rewriter.getBoolAttr(false)) + .getResult(); + + for (size_t i = 0; i < originalType.getShape().size(); ++i) { + // Use dynamic value to judge whether overstep boundary + auto shapeVal = rewriter.create( + loc, rewriter.getIndexAttr(originalType.getDimSize(i))); + + Value maskDimVal; + if (isa(maskDim[i])) + maskDimVal = rewriter.create( + loc, cast(maskDim[i].get())); + else + maskDimVal = maskDim[i].get(); + + auto curCmp = rewriter.create(loc, arith::CmpIPredicate::slt, + maskDimVal, shapeVal); + + fillFlag = rewriter.create(loc, fillFlag, curCmp.getResult()) + .getResult(); + } + auto ifOp = rewriter.create(loc, fillFlag); + { + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPointToStart(&ifOp.getThenRegion().front()); + rewriter.create(loc, ValueRange{other}, + ValueRange{localMem}); + } + ifOp->setAttr(rewriter.getStringAttr("hivm.unlikely_condition"), + UnitAttr::get(rewriter.getContext())); +} + +LoadConverter::LoadConverter(MLIRContext *context) + : OpConversionPattern(context) {} + +LogicalResult +LoadConverter::matchAndRewrite(triton::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + + // Check if tt.load is modified by AddPtrConverter to a specified state. + if (checkModifiedByAddPtrConverter(op).succeeded()) { + return continueModifyFromAddPtrConverter(op, adaptor, rewriter); + } + + auto ptr = adaptor.getPtr(); + auto mask = op.getMask(); + auto other = op.getOther(); + auto loc = op.getLoc(); + + // handling scalar + if (!isa(op.getResult().getType())) { + auto scalarMemref = + BlockDataParser::getScalarMemRef(op.getPtr(), ptr, loc, rewriter); + auto resTy = op.getResult().getType(); + auto idxZero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + auto loadedValue = rewriter + .create(loc, resTy, scalarMemref, + idxZero.getResult()) + .getResult(); + if (mask && other) { + mask = rewriter.create( + loc, RankedTensorType::get({1}, mask.getType()), mask); + loadedValue = rewriter.create( + loc, RankedTensorType::get({1}, loadedValue.getType()), loadedValue); + other = rewriter.create( + loc, RankedTensorType::get({1}, other.getType()), other); + loadedValue = + rewriter.create(loc, mask, loadedValue, other); + rewriter.replaceOpWithNewOp(op, loadedValue, + ValueRange({idxZero})); + } else { + rewriter.replaceOp(op, loadedValue); + } + return success(); + } + + int64_t lastStride = -1; + if (isa(ptr)) { + auto u = ptr; + while (auto blkArg = dyn_cast(u)) { + if (auto forOp = dyn_cast(blkArg.getOwner()->getParentOp())) { + auto prt = forOp->getOperand(3 + blkArg.getArgNumber() - 1); + u = prt; + } else { + u = nullptr; + break; + } + } + if (u && isa(u.getDefiningOp())) { + auto ret = mlir::ConverterUtils::getLastStrideOfReinterpretCastOp( + dyn_cast(u.getDefiningOp())); + if (ret.has_value()) + lastStride = *ret; + } + } + + // handling no mask + auto memRefType = dyn_cast(ptr.getType()); + if (!memRefType) { + return rewriter.notifyMatchFailure( + op, "LoadOp expects a memref, not a memref of pointers"); + } + if (!op->hasAttr(ConverterUtils::GeneratedByMakeTensorPtrTAG)) { + auto memrefOp = dyn_cast(ptr.getDefiningOp()); + auto ret = mlir::ConverterUtils::getLastStrideOfReinterpretCastOp(memrefOp); + if (ret.has_value()) + lastStride = *ret; + } + bool mayImplicitTransposeWithLastAxis = + (existDotFlag) && + (!op->hasAttr(ConverterUtils::GeneratedByMakeTensorPtrTAG)) && + (lastStride != 1 && + mlir::ConverterUtils::isaPermutedMemRefType(memRefType)); + auto memRefShape = memRefType.getShape(); + auto memRefElementType = memRefType.getElementType(); + + Value allocOp; + Value allocOpTmp; + if (op->hasAttr(ConverterUtils::discreteAttrName)) { + Operation *loop = op->getParentOp(); + int extractedLoopCount = 1; + for (auto parentOp = loop->getParentOp(); + parentOp->hasAttr("ExtractedLoadOrStore"); + parentOp = parentOp->getParentOp()) { + loop = parentOp; + extractedLoopCount++; + } + rewriter.setInsertionPoint(loop); + auto loopOp = cast(loop); + auto fullMemRefShape = + cast(loopOp.getInitArgs()[0].getType()).getShape(); + auto fullMemRefType = MemRefType::get(fullMemRefShape, memRefElementType); + bool isIndexSelectScenario = + (extractedLoopCount == 1) && (fullMemRefShape.size() > 1u); + if (isIndexSelectScenario) + loopOp->setAttr("hivm.parallel_loop", rewriter.getUnitAttr()); + allocOp = rewriter.create(loc, fullMemRefType); + allocOpTmp = allocOp; + rewriter.setInsertionPointAfter(loop); + auto toTensorOp = rewriter.create( + loc, RankedTensorType::get(fullMemRefShape, memRefElementType), allocOp, + true, true); + rewriter.replaceAllUsesWith(loopOp->getResult(0), toTensorOp->getResult(0)); + tensor::InsertSliceOp insertSliceOp = nullptr; + for (auto *user : op->getUsers()) { + if (auto targetOp = dyn_cast(user)) { + insertSliceOp = targetOp; + break; + } + } + auto offsets = insertSliceOp.getMixedOffsets(); + auto sizes = insertSliceOp.getMixedSizes(); + auto strides = insertSliceOp.getMixedStrides(); + auto allocType = memref::SubViewOp::inferResultType(fullMemRefType, offsets, + sizes, strides); + rewriter.setInsertionPoint(op); + allocOp = rewriter.create( + loc, cast(allocType), allocOp, offsets, sizes, strides); + rewriter.replaceAllUsesExcept(insertSliceOp.getResult(), + insertSliceOp.getDest(), insertSliceOp); + rewriter.eraseOp(insertSliceOp); + } else { + allocOp = rewriter.create( + loc, MemRefType::get(memRefShape, memRefElementType)); + } + + auto tensorType = RankedTensorType::get(memRefShape, memRefElementType); + // boundary check + auto boundaryCheck = op.getBoundaryCheck(); + if (!boundaryCheck.empty()) { + auto boundarySizes = mlir::ConverterUtils::getBoundarySizes( + boundaryCheck, /*remapped*/ ptr, loc, rewriter); + // handle the padding + auto padding = op.getPadding(); + if (padding.has_value()) { + TypedAttr padAttr = rewriter.getZeroAttr(memRefElementType); + // triton already ensure only NAN and ZERO are passed in + if (padding.value() == triton::PaddingOption::PAD_NAN) { + // FIXME: Why NaN requires elemTy to be non-int or non-index? + assert(!memRefElementType.isIntOrIndex()); + auto apNaN = llvm::APFloat::getNaN( + cast(padAttr).getValue().getSemantics()); + padAttr = rewriter.getFloatAttr(memRefElementType, apNaN); + } + auto padVal = rewriter.create(loc, padAttr); + + fillTensorWithOtherForMaskScenario(padVal, allocOp, boundarySizes, + rewriter); + } + + auto srcSubView = + mlir::ConverterUtils::makeSubViewOp(ptr, boundarySizes, loc, rewriter); + auto dstSubview = mlir::ConverterUtils::makeSubViewOp( + allocOp, boundarySizes, loc, rewriter); + rewriter.create(loc, srcSubView, dstSubview); + if (mayImplicitTransposeWithLastAxis && + allocOp.getDefiningOp()) { + auto markOp = rewriter.create(loc, dstSubview); + markOp->setAttr(MayImplicitTransposeWithLastAxisTAG, + UnitAttr::get(rewriter.getContext())); + } else if (mayImplicitTransposeWithLastAxis && + allocOp.getDefiningOp()) { + auto markOp = rewriter.create(loc, allocOpTmp); + markOp->setAttr(MayImplicitTransposeWithLastAxisTAG, + UnitAttr::get(rewriter.getContext())); + } + return this->toTensorAndReplace(op, tensorType, allocOp, + mayImplicitTransposeWithLastAxis, loc, + rewriter); + } + + if (!mask) { + assert(!other && "can not input 'other' when 'mask' is not set"); + if (auto unrealizedCastOp = + ptr.getDefiningOp()) { + // TODO : not support handle associate with "module" + // hint : can be handled in Linearize + op->emitError("meeting unexpected UCC in LoadConverter!"); + return failure(); + } else { + // If last dimension stride equals 2, try deinterleave optimization. +#if LLVM_VERSION_MAJOR < 21 + auto [ptrStrides, ptrOffsets] = getStridesAndOffset(memRefType); +#else // triton_v3.3.x + auto [ptrStrides, ptrOffsets] = memRefType.getStridesAndOffset(); +#endif + if (ptrStrides.back() == 2 && (memRefShape.back() % 2 == 0) && + mlir::triton::DeinterleaveStatusOptimization(op, adaptor, rewriter) + .succeeded()) { + return success(); + } + rewriter.create(loc, ptr, allocOp); + if (mayImplicitTransposeWithLastAxis) { + auto markOp = rewriter.create(loc, allocOp); + markOp->setAttr(MayImplicitTransposeWithLastAxisTAG, + UnitAttr::get(rewriter.getContext())); + } + } + + return this->toTensorAndReplace(op, tensorType, allocOp, + mayImplicitTransposeWithLastAxis, loc, + rewriter); + } + + MaskState mstate; + auto isContMask = mstate.parse(mask, loc, rewriter); + if (isContMask.failed()) { + return rewriter.notifyMatchFailure( + op, "can not lower uncontinuout masked loads"); + } + + if (other) { + auto scalarOther = + mlir::ConverterUtils::getScalarValue(other, loc, rewriter); + assert( + scalarOther && + "other value used in masked load produced by unsupported instruction!"); + + fillTensorWithOtherForMaskScenario(scalarOther, allocOp, mstate.dims, + rewriter); + } + + // To enable deinterleave optimization with mask load, mask state along last + // dimension couldn't be split, which means `dims.back()` must be equal to + // origin type last dimension constant size and `offsets.back()` must be 0. + // + // The basis is that last dimension range comparison would generate + // unaccepted discontinuous mask. + if (mstate.getRank() == memRefType.getRank() && + isConstantIntValue(mstate.offsets.back(), 0) && + isConstantIntValue(mstate.dims.back(), memRefType.getShape().back())) { +#if LLVM_VERSION_MAJOR < 21 + auto [ptrStrides, ptrOffsets] = getStridesAndOffset(memRefType); +#else // triton_v3.3.x + auto [ptrStrides, ptrOffsets] = memRefType.getStridesAndOffset(); +#endif + if (ptrStrides.back() == 2 && (memRefType.getShape().back() % 2 == 0) && + DeinterleaveStatusWithMaskOptimization(op, adaptor, rewriter, mstate, + allocOp) + .succeeded()) { + return success(); + } + } + + if (auto unrealizedCastOp = ptr.getDefiningOp()) { + // TODO : not support handle associate with "module" + // hint : can be handled in Linearize + op->emitError("meeting unexpected UCC in LoadConverter!"); + return failure(); + } else { + memref::SubViewOp srcSubView = mstate.getSubview(ptr, loc, rewriter); + memref::SubViewOp dstSubView = mstate.getSubview(allocOp, loc, rewriter); + MemRefType dstSubViewType = mlir::cast(dstSubView.getType()); +#if LLVM_VERSION_MAJOR < 21 + auto [srcStrides, srcOffset] = getStridesAndOffset(dstSubViewType); +#else // triton_v3.3.x + auto [srcStrides, srcOffset] = dstSubViewType.getStridesAndOffset(); +#endif + MemRefType castType = MemRefType::get( + dstSubViewType.getShape(), dstSubViewType.getElementType(), + makeStridedLinearLayoutMap(srcStrides, srcOffset, + rewriter.getContext())); + auto castOp = rewriter.create(loc, castType, dstSubView); + rewriter.create(loc, srcSubView, castOp); + + if (mayImplicitTransposeWithLastAxis && + allocOp.getDefiningOp()) { + auto markOp = rewriter.create(loc, allocOp); + markOp->setAttr(MayImplicitTransposeWithLastAxisTAG, + UnitAttr::get(rewriter.getContext())); + } else if (mayImplicitTransposeWithLastAxis && + allocOp.getDefiningOp()) { + auto markOp = rewriter.create(loc, allocOpTmp); + markOp->setAttr(MayImplicitTransposeWithLastAxisTAG, + UnitAttr::get(rewriter.getContext())); + } + } + return this->toTensorAndReplace( + op, tensorType, allocOp, mayImplicitTransposeWithLastAxis, loc, rewriter); +} + +AtomicRMWConverter::AtomicRMWConverter(MLIRContext *context) + : OpConversionPattern(context) {} + +// lowering tt.atomicRMW to linalg.generic +// If atomic op's return value is used by other op as it's the old value stored +// at the ptrwe will use tt.load to get it +// +// example: +// input: +// %return_value = tt.atomic_rmw fadd, acq_rel, gpu, +// %output_memref, %input_tensor, %mask : +// (tensor<256x!tt.ptr>, tensor<256xf32>, tensor<256xi1>) +// -> tensor<256xf32> +// +// output: +// memref.copy %output_memref, %ub_buf : memref to memref +// %17 = bufferization.to_tensor %alloc_3 restrict writable : memref<256xf32> +// linalg.generic +// {indexing_maps = [#map, #map, #map], iterator_types = ["parallel"]} +// ins(%output_memref, %masked_input_memref : memref, memref) +// outs(%subview_2 : memref) +// attrs = {GenericAtomicRMW = "fadd", MemSemantic = "acq_rel", +// MemSyncScope = "gpu"} { +// ^bb0(%in: f32, %in_9: f32, %out: f32): +// %25 = arith.addf %in, %in_9 : f32 +// linalg.yield %25 : f32 +// } +LogicalResult +AtomicRMWConverter::matchAndRewrite(triton::AtomicRMWOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + // If the result of AtomicRMWOp is not used, we don't need to load the old + // data stored at the ptr + auto ptr = adaptor.getPtr(); + auto val = op.getVal(); + auto loc = op.getLoc(); + + auto resType = dyn_cast(op.getResult().getType()); + if (!resType) { + return rewriter.notifyMatchFailure( + op, "atomicRMWConverter: scalar will be handled by " + "ScalarAtomicRMWCanonicalizer"); + } + + auto rmwOp = op.getAtomicRmwOp(); + + // 1. Simple case where no mask is used. + auto type = dyn_cast(ptr.getType()); + if (!type) { + // Seen when implicit broadcasting is done late in a chain of + // operations. The workaround is to broadcast the pointers early in the + // address calculation. A proper fix is complicated, but at least we can + // provide a better error message. + return rewriter.notifyMatchFailure( + op, "AtomicRMWOp expects a memref, not a memref of pointers"); + } + + auto dstMemref = ptr; + // Well, linalg structure op wouldn't support mixed tensor/buffer semantics + // any more in latest LLVM(triton LLVM dependency has involed this), so we + // need to convert tensor to buffer early. + auto dstOriType = cast(dstMemref.getType()); + MemRefType dstType = + MemRefType::get(dstOriType.getShape(), dstOriType.getElementType()); + Value inputMemref = + rewriter.create(loc, dstType, val); + + // 2. handle the mask for the atomic op + // When the dsl do not pass the mask to this op like + // `tl.atomic_add(out_ptr0 + xindex, tmp2)`, it will create a constant mask + // for this op by default, which is not supported by maskAnalysis, so we + // need to handle this situation + // + // This logic come from semantic.py: + // + // if not mask: + // mask_ir = builder.get_int1(True) + // mask_ty = tl.int1 + // if ptr.type.is_block(): + // mask_ir = \ + // builder.create_splat(mask_ir, ptr.type.get_block_shapes()) + // mask_ty = tl.block_type(tl.int1, ptr.type.get_block_shapes()) + // mask = tl.tensor(mask_ir, mask_ty) + // + // ... + // + // return ptr, val, mask + // + if (auto mask = op.getMask()) { + MaskState mstate; + auto constantMask = mask.getDefiningOp(); + if (!constantMask) { + auto isContMask = mstate.parse(mask, loc, rewriter); + + if (isContMask.failed()) { + return rewriter.notifyMatchFailure( + op, "Cannot lower continuous masked loads"); + } + dstMemref = mstate.getSubview(ptr, loc, rewriter); + inputMemref = mstate.getSubview(inputMemref, loc, rewriter); + } else { + if (!isConstantMaskTrue(mask)) { + rewriter.eraseOp(op); + return success(); + } + } + } + + // create element-wise map + int64_t rank = type.getRank(); + SmallVector inputDims; + auto context = rewriter.getContext(); + + for (int i = 0; i < rank; i++) { + inputDims.push_back(getAffineDimExpr(i, context)); + } + + SmallVector indexingMaps; + // As mask has been erased for now + // the number of input must be 2 + // the input memref is also the output memref + // Thus, there are a total of three inputs and outputs. + // so here we have 3 map to create + for (int i = 0; i < 3; i++) { + indexingMaps.push_back(AffineMap::get(rank, 0, inputDims, context)); + } + + Value tensorToReplace; + if (!op.getResult().use_empty()) { + auto tensorType = + RankedTensorType::get(type.getShape(), type.getElementType()); + auto alloc = rewriter.create( + loc, MemRefType::get(type.getShape(), type.getElementType())); + // For the return value, don't need to care about mask for now + // this op don't support other, so we best not fill it + rewriter.create(loc, ptr, alloc); + tensorToReplace = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + } + + auto linalgOp = rewriter.create( + loc, /* operands */ ValueRange{dstMemref, inputMemref}, + ValueRange{dstMemref}, indexingMaps, + mlir::ConverterUtils::getNParallelLoopsAttrs(rank), + [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange blockArgs) { + Value opResult = createAtomicBinaryOps(nestedBuilder, nestedLoc, op, + type.getElementType(), + blockArgs[0], blockArgs[1]); + nestedBuilder.create(nestedLoc, opResult); + }); + + // "library_call" + // indicating the actual semantic of this op + // TODO: If the hardware support the MemSemantic/MemSyncScope + // We pass them down + // otherwise they need to be deleted + const StringRef genericAtomicRMW = "GenericAtomicRMW"; + const StringRef memSemantic = "MemSemantic"; + const StringRef memSyncScope = "MemSyncScope"; + linalgOp->setAttr(genericAtomicRMW, + rewriter.getStringAttr(stringifyEnum(op.getAtomicRmwOp()))); + linalgOp->setAttr(memSemantic, + rewriter.getStringAttr(stringifyEnum(op.getSem()))); + linalgOp->setAttr(memSyncScope, + rewriter.getStringAttr(stringifyEnum(op.getScope()))); + + // Mark atomic_and/or/xor specially which need software simulation in terms + // of backend restriction + if (softwareAtomicKinds.contains(op.getAtomicRmwOp())) + linalgOp->setAttr("Software", rewriter.getUnitAttr()); + + // tt.atomicRMW op has two part of feature + // 1. load the old data at the ptr + // 2. atomically store the data on ub to the ptr + // at the same time it perform the action it has been assigned + // So we lower this op to load + atomically store + // + // The first part is not necessary when the returned value of atomic op + // is not used, it will be deleted cause it's meaningless + // Here, we preemptively determine whether it will be used + // and decide whether it is necessary to create the load process based on + // this assessment. + // + // logic of handling is copied + // TODO: decoupling the logic of load, put it in the Utils + if (!op.getResult().use_empty()) { + rewriter.replaceOp(op, tensorToReplace); + } else { + rewriter.eraseOp(op); + } + return success(); +} + +AtomicRMWNewConverter::AtomicRMWNewConverter(MLIRContext *context) + : OpConversionPattern(context) {} + +// lowering tt.atomicRMW to linalg.generic +// If atomic op's return value is used by other op as it's the old value stored +// at the ptrwe will use tt.load to get it +// +// example: +// input: +// %return_value = tt.atomic_rmw fadd, acq_rel, gpu, +// %output_memref, %input_tensor, %mask : +// (tensor<256x!tt.ptr>, tensor<256xf32>, tensor<256xi1>) +// -> tensor<256xf32> +// +// output: +// memref.copy %output_memref, %ub_buf : memref to memref +// %17 = bufferization.to_tensor %alloc_3 restrict writable : memref<256xf32> +// linalg.generic +// {indexing_maps = [#map, #map, #map], iterator_types = ["parallel"]} +// ins(%output_memref, %masked_input_memref : memref, memref) +// outs(%subview_2 : memref) +// attrs = {GenericAtomicRMW = "fadd", MemSemantic = "acq_rel", +// MemSyncScope = "gpu"} { +// ^bb0(%in: f32, %in_9: f32, %out: f32): +// %25 = arith.addf %in, %in_9 : f32 +// linalg.yield %25 : f32 +// } +LogicalResult AtomicRMWNewConverter::matchAndRewrite( + triton::AtomicRMWOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto ptr = adaptor.getPtr(); + auto val = op.getVal(); + auto loc = op.getLoc(); + auto mask = op.getMask(); + auto rmwOp = op.getAtomicRmwOp(); + auto resType = dyn_cast(op.getResult().getType()); + auto ptrType = dyn_cast(ptr.getType()); + + if (!resType) + return rewriter.notifyMatchFailure( + op, "atomicRMWConverter: scalar will be handled by " + "ScalarAtomicRMWCanonicalizer"); + if (!ptrType) + return rewriter.notifyMatchFailure( + op, "AtomicRMWOp expects a memref, not a memref of pointers"); + + const std::map atomicKindMap = { + {RMWOp::ADD, hivm::AtomicKind::ADD}, + {RMWOp::FADD, hivm::AtomicKind::ADD}, + {RMWOp::OR, hivm::AtomicKind::OR}, + {RMWOp::XOR, hivm::AtomicKind::XOR}, + {RMWOp::AND, hivm::AtomicKind::AND}, + {RMWOp::MIN, hivm::AtomicKind::MIN}, + {RMWOp::UMIN, hivm::AtomicKind::UMIN}, + {RMWOp::MAX, hivm::AtomicKind::MAX}, + {RMWOp::UMAX, hivm::AtomicKind::UMAX}, + {RMWOp::XCHG, hivm::AtomicKind::XCHG}, + }; + + assert(atomicKindMap.find(rmwOp) != atomicKindMap.end()); + auto atomicKind = + hivm::AtomicKindAttr::get(rewriter.getContext(), atomicKindMap.at(rmwOp)); + + auto dstMemref = ptr; + Value inputVal = val; + + // Lazily materialize a memref view only when we truly need buffer + // semantics (e.g., mask subview or XCHG lowering). Otherwise keep tensor + // inputs to avoid redundant to_memref conversions before hivm.store. + Value inputMemref; + auto getInputMemref = [&]() -> Value { + if (isa(inputVal.getType())) + return inputVal; + if (inputMemref) + return inputMemref; + inputMemref = + rewriter.create(loc, ptrType, inputVal); + return inputMemref; + }; + + bool isDiscreteMask = false; + if (mask) { + auto constantMask = mask.getDefiningOp(); + if (constantMask && !isConstantMaskTrue(mask)) { + rewriter.eraseOp(op); + return success(); + } + MaskState mstate; + isDiscreteMask = mstate.parse(mask, loc, rewriter).failed(); + if (!constantMask && !isDiscreteMask) { + // For dstMemref (store output), use subview to maintain reference to + // original memref. For inputVal (store input), use tensor.extract_slice + // to keep tensor semantics. + dstMemref = mstate.getSubview(ptr, loc, rewriter); + + auto inputMemrefVal = getInputMemref(); + auto inputMemrefType = cast(inputMemrefVal.getType()); + auto inputTensorType = RankedTensorType::get( + inputMemrefType.getShape(), inputMemrefType.getElementType()); + Value inputTensor = rewriter.create( + loc, inputTensorType, inputMemrefVal, true, true); + inputVal = mstate.getExtractSlice(inputTensor, loc, rewriter); + } + } + + if (!op.getResult().use_empty()) { + auto tensorType = + RankedTensorType::get(ptrType.getShape(), ptrType.getElementType()); + auto alloc = rewriter.create( + loc, MemRefType::get(ptrType.getShape(), ptrType.getElementType())); + rewriter.create(loc, ptr, alloc); + Value tensorToReplace = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensorToReplace); + } + + if (isDiscreteMask) { + if (rmwOp != RMWOp::XCHG) { + return op.emitError( + "Discrete mask is only expected for XCHG; other atomics " + "should be lowered without discrete masks"); + } + Value memrefMask = mask; + if (auto maskTypeT = dyn_cast(mask.getType())) { + MemRefType maskTypeM = + MemRefType::get(maskTypeT.getShape(), maskTypeT.getElementType()); + memrefMask = + rewriter.create(loc, maskTypeM, mask); + } + rewriter.create( + op.getLoc(), TypeRange(), getInputMemref(), dstMemref, memrefMask); + } else { + if (rmwOp == RMWOp::XCHG) + rewriter.create(op.getLoc(), TypeRange(), + getInputMemref(), dstMemref); + else + rewriter.create(op.getLoc(), TypeRange{}, inputVal, + dstMemref, atomicKind); + } + + if (op.getResult().use_empty()) { + rewriter.eraseOp(op); + } + return success(); +} + +LogicalResult +AtomicCASConverter::matchAndRewrite(triton::AtomicCASOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + // If the result of AtomicCASOp is not used, we don't need to load the old + // data stored at the ptr + auto ptr = adaptor.getPtr(); + auto cmp = op.getCmp(); + auto val = op.getVal(); + auto loc = op.getLoc(); + + auto resType = dyn_cast(op.getResult().getType()); + if (!resType) { + return rewriter.notifyMatchFailure( + op, "atomicCASConverter: scalar will be handled by " + "ScalarAtomicCASCanonicalizer"); + } + + // 1. Simple case where no mask is used. + auto type = dyn_cast(ptr.getType()); + if (!type) { + // Seen when implicit broadcasting is done late in a chain of + // operations. The workaround is to broadcast the pointers early in the + // address calculation. A proper fix is complicated, but at least we can + // provide a better error message. + return rewriter.notifyMatchFailure( + op, "AtomicCASOp expects a memref, not a memref of pointers"); + } + + auto dstMemref = ptr; + // Well, linalg structure op wouldn't support mixed tensor/buffer semantics + // any more in latest LLVM(triton LLVM dependency has involed this), so we + // need to convert tensor to buffer early. + auto dstOriType = cast(dstMemref.getType()); + MemRefType dstType = + MemRefType::get(dstOriType.getShape(), dstOriType.getElementType()); + Value inputMemref = + rewriter.create(loc, dstType, val); + + Value cmpMemref = + rewriter.create(loc, dstType, cmp); + + // create element-wise map + int64_t rank = type.getRank(); + SmallVector inputDims; + auto context = rewriter.getContext(); + + for (int i = 0; i < rank; i++) { + inputDims.push_back(getAffineDimExpr(i, context)); + } + + SmallVector indexingMaps; + // As mask has been erased for now + // the number of input must be 2 + // the input memref is also the output memref + // Thus, there are a total of four inputs and outputs. + // so here we have 4 map to create + for (int i = 0; i < 4; i++) { // 4: 3 input and 1 output + indexingMaps.push_back(AffineMap::get(rank, 0, inputDims, context)); + } + + if (!op.getResult().use_empty()) { + auto tensorType = + RankedTensorType::get(type.getShape(), type.getElementType()); + auto alloc = rewriter.create( + loc, MemRefType::get(type.getShape(), type.getElementType())); + + // For the return value, don't need to care about mask for now + // this op don't support other, so we best not fill it + rewriter.create(loc, ptr, alloc); + Value tensor = rewriter.create( + loc, tensorType, alloc, true /* restrict */, true /* writable */); + rewriter.replaceOp(op, tensor); + } + + auto linalgOp = rewriter.create( + loc, ValueRange{dstMemref, cmpMemref, inputMemref}, + mlir::ValueRange{dstMemref}, indexingMaps, + mlir::ConverterUtils::getNParallelLoopsAttrs(rank), + [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange blockArgs) { + Value lhs = blockArgs[0]; + Value rhs = blockArgs[1]; + Value setValue = blockArgs[2]; + Value cond; + if (mlir::isa(lhs.getType())) { + cond = nestedBuilder.create( + nestedLoc, arith::CmpFPredicate::UEQ, lhs, rhs); + } else { + cond = nestedBuilder.create( + nestedLoc, arith::CmpIPredicate::eq, lhs, rhs); + } + auto ifOp = nestedBuilder.create( + nestedLoc, TypeRange{setValue.getType()}, cond, true); + { + OpBuilder::InsertionGuard guard(nestedBuilder); + nestedBuilder.setInsertionPointToEnd(&ifOp.getThenRegion().front()); + nestedBuilder.create(nestedLoc, setValue); + } + { + OpBuilder::InsertionGuard guard(nestedBuilder); + nestedBuilder.setInsertionPointToEnd(&ifOp.getElseRegion().front()); + nestedBuilder.create(nestedLoc, lhs); + } + nestedBuilder.setInsertionPointToEnd(nestedBuilder.getBlock()); + nestedBuilder.create(nestedLoc, + ifOp.getResult(0)); + }); + + const StringRef genericAtomicRMW = "GenericAtomicRMW"; + const StringRef memSemantic = "MemSemantic"; + const StringRef memSyncScope = "MemSyncScope"; + auto attr = mlir::StringAttr::get(context, "cas"); + + linalgOp->setAttr(genericAtomicRMW, attr); + linalgOp->setAttr(memSemantic, + rewriter.getStringAttr(stringifyEnum(op.getSem()))); + linalgOp->setAttr(memSyncScope, + rewriter.getStringAttr(stringifyEnum(op.getScope()))); + + linalgOp->setAttr("Software", rewriter.getUnitAttr()); + + // tt.atomicRMW op has two part of feature + // 1. load the old data at the ptr + // 2. atomically store the data on ub to the ptr + // at the same time it perform the action it has been assigned + // So we lower this op to load + atomically store + // + // The first part is not necessary when the returned value of atomic op + // is not used, it will be deleted cause it's meaningless + // Here, we preemptively determine whether it will be used + // and decide whether it is necessary to create the load process based on + // this assessment. + // + // logic of handling is copied + if (op.getResult().use_empty()) { + rewriter.eraseOp(op); + } + return success(); +} + +LogicalResult +ScalarStoreCanonicalizer::matchAndRewrite(triton::StoreOp op, + PatternRewriter &rewriter) const { + if (!op.getValue().getType().isIntOrIndexOrFloat()) { + return rewriter.notifyMatchFailure( + op, "ScalarStoreCanonicalizer handles scalar store scene!"); + } + auto ptr = op.getPtr(); + auto mask = op.getMask(); + auto value = op.getValue(); + if (mask) { + rewriter.replaceOpWithNewOp( + op, mask, [&](OpBuilder &b, Location loc) { + b.create(loc, ptr, value, op.getCache(), + op.getEvict()); + b.create(loc); + }); + return success(); + } + + auto ptrTy = RankedTensorType::get({(int64_t)1}, ptr.getType()); + auto ptrSplat = rewriter.create(op.getLoc(), ptrTy, ptr); + auto valTy = RankedTensorType::get({(int64_t)1}, value.getType()); + auto valSplat = rewriter.create(op.getLoc(), valTy, value); + auto newStoreOp = rewriter.create( + op.getLoc(), ptrSplat, valSplat, op.getCache(), op.getEvict()); + rewriter.replaceOp(op, newStoreOp); + return success(); +} + +LogicalResult +ScalarAtomicRMWCanonicalizer::matchAndRewrite(triton::AtomicRMWOp op, + PatternRewriter &rewriter) const { + if (!op.getVal().getType().isIntOrIndexOrFloat()) { + return rewriter.notifyMatchFailure( + op, "ScalarAtomicRMWCanonicalizer handles scalar atomic rmw op scene!"); + } + + auto ptr = op.getPtr(); + auto ptrTy = RankedTensorType::get({(int64_t)1}, ptr.getType()); + auto ptrSplat = rewriter.create(op.getLoc(), ptrTy, ptr); + auto valTy = RankedTensorType::get({(int64_t)1}, op.getVal().getType()); + auto valSplat = + rewriter.create(op.getLoc(), valTy, op.getVal()); + auto maskTy = RankedTensorType::get({(int64_t)1}, op.getMask().getType()); + auto maskSplat = + rewriter.create(op.getLoc(), maskTy, op.getMask()); + + auto newAtomicOp = rewriter.create( + op.getLoc(), valTy, op.getAtomicRmwOp(), ptrSplat, valSplat, maskSplat, + op.getSem(), op.getScope()); + auto idxZero = + rewriter.create(op.getLoc(), rewriter.getIndexAttr(0)); + rewriter.replaceOpWithNewOp(op, newAtomicOp, + ValueRange({idxZero})); + return success(); +} + +LogicalResult +ScalarAtomicCASCanonicalizer::matchAndRewrite(triton::AtomicCASOp op, + PatternRewriter &rewriter) const { + if (!op.getVal().getType().isIntOrIndexOrFloat() && + !op.getCmp().getType().isIntOrIndexOrFloat()) { + return rewriter.notifyMatchFailure( + op, "ScalarAtomicCASCanonicalizer handles scalar atomic cas op scene!"); + } + + auto ptr = op.getPtr(); + auto ptrTy = RankedTensorType::get({(int64_t)1}, ptr.getType()); + auto ptrSplat = rewriter.create(op.getLoc(), ptrTy, ptr); + auto cmpTy = RankedTensorType::get({(int64_t)1}, op.getCmp().getType()); + auto cmpSplat = + rewriter.create(op.getLoc(), cmpTy, op.getCmp()); + auto valTy = RankedTensorType::get({(int64_t)1}, op.getVal().getType()); + auto valSplat = + rewriter.create(op.getLoc(), valTy, op.getVal()); + + auto newAtomicOp = rewriter.create( + op.getLoc(), valTy, ptrSplat, cmpSplat, valSplat, op.getSem(), + op.getScope()); + auto idxZero = + rewriter.create(op.getLoc(), rewriter.getIndexAttr(0)); + rewriter.replaceOpWithNewOp(op, newAtomicOp, + ValueRange({idxZero})); + return success(); +} + +// The atomic max op with float input will be devided into +// two atomic max ops with integer input +// One handles the part of the tensor greater than zero +// the other deals with the part less than zero +// It will lead to maskAnalysis failure +// So here we need to revert the procedures in semantics.py +// The triton IR is like +// +// %cst_0 = arith.constant dense<0.000000e+00> : tensor<1x256xf32> +// %1 = tt.bitcast %value : tensor<1x256xf32> -> tensor<1x256xi32> +// %2 = tt.bitcast %ptr : tensor<1x256x!tt.ptr> -> +// tensor<1x256x!tt.ptr> %3 = arith.cmpf oge, %1, %cst_0 %4 = arith.cmpf +// olt, %1, %cst_0 %5 = arith.andi %8, %3 %6 = tt.atomic_rmw max, acq_rel, gpu, +// %2, %1, %5 : +// (tensor<1x256x!tt.ptr>, tensor<1x256xi32>, tensor<1x256xi1>) -> +// tensor<1x256xi32> +// %7 = arith.andi %8, %4 +// %8 = tt.atomic_rmw umin, acq_rel, gpu, %2, %1, %7 : +// (tensor<1x256x!tt.ptr>, tensor<1x256xi32>, tensor<1x256xi1>) -> +// tensor<1x256xi32> +// +// it's hard to handle and meaningless complicated for our device +// so we revert it to +// %0 = tt.atomic_rmw max, acq_rel, gpu, %23, %21, %8 : +// (tensor<1x256x!tt.ptr>, tensor<1x256xf32>, tensor<1x256xi1>) -> +// tensor<1x256xf32> +LogicalResult +AtomicMaxMinCanonicalizer::matchAndRewrite(triton::AtomicRMWOp op, + PatternRewriter &rewriter) const { + // Revert the op to its original form + auto ptrBitcastOp = op.getPtr().getDefiningOp(); + auto valueBitcastOp = op.getVal().getDefiningOp(); + if (!ptrBitcastOp || !valueBitcastOp) { + return failure(); + } + + // We only need to handle the op when the element type is float + auto elementType = + dyn_cast(valueBitcastOp.getSrc().getType()).getElementType(); + if (!isa(elementType)) { + return failure(); + } + + auto rmwOp = op.getAtomicRmwOp(); + // here we know that atomic UMAX/UMIN + // is created by special logic of triton right now + // so we can simply delete it + if (rmwOp == triton::RMWOp::UMAX || rmwOp == triton::RMWOp::UMIN) { + // if the return value of op is used, we can't simply erase it + if (op.getResult().use_empty()) { + rewriter.eraseOp(op); + return success(); + } + return failure(); + } + + if (rmwOp != triton::RMWOp::MAX && rmwOp != triton::RMWOp::MIN) { + return failure(); + } + + // 1. Though semantic interpreter will generate full true tensor as original + // mask if atomicrmwOp don't have it, above float devision process will also + // generate positive and negative comparison mask, which will cause to fold + // true mask. + // 2. While if atomicrmwOp has original mask, there exists andiop between + // original mask and positive/negative comparison mask + // + // Here wanna extract original mask + Value originalMask = op.getMask(); + if (auto andOp = originalMask.getDefiningOp()) + // LHS is convention in semantic interpreter + originalMask = andOp.getLhs(); + else if (auto cmpOp = originalMask.getDefiningOp()) { + if (cmpOp.getPredicate() != mlir::arith::CmpFPredicate::OGE || + !matchPattern(cmpOp.getRhs(), + /*positive float zero matcher*/ m_PosZeroFloat())) + // Here recheck frontend interpreter generation in no manual mask state + return op->emitError("Illegal mask for atomicrmwOp of float type"); + // Restore original true mask + originalMask = rewriter.create( + op->getLoc(), + /*typed attr*/ DenseElementsAttr::get( + cast(originalMask.getType()), true)); + } else + return op->emitError("Illegal mask for atomicrmwOp of float type"); + + auto originAtomicOp = rewriter.create( + op.getLoc(), valueBitcastOp.getSrc().getType(), op.getAtomicRmwOp(), + ptrBitcastOp.getSrc(), valueBitcastOp.getSrc(), originalMask, op.getSem(), + op.getScope()); + + // if the return value of op is used + // we need to handle its usage + // In semantic.py, if the atomic Max/Min with float input is used + // It will use select + bitcast to get float value + // so here we need to revert it too + // + // For example: + // %0 = tt.atomic_rmw max, acq_rel, gpu, %gm, %input, %mask1 : + // (tensor<32x!tt.ptr>... %1 = tt.atomic_rmw umin, acq_rel, gpu, %gm, + // %input, %mask2 : (tensor<32x!tt.ptr>... %2 = arith.select + // %devidedMask, %0, %1 : tensor<32xi1>, tensor<32xi32> %3 = tt.bitcast %2 : + // tensor<32xi32> -> tensor<32xf32> tt.store %outputMemref, %3 : + // tensor<32x!tt.ptr> + // + // will be revert to: + // %0 = tt.atomic_rmw max, acq_rel, gpu, %gm, %input, %mask : + // (tensor<32x!tt.ptr>... tt.store %outputMemref, %0 : + // tensor<32x!tt.ptr> + // + if (!op.getResult().use_empty()) { + for (OpOperand &use : op->getUses()) { + auto selectOp = dyn_cast(use.getOwner()); + if (!selectOp) + continue; + + for (OpOperand &selectUse : selectOp->getUses()) { + if (auto bitcastOp = + dyn_cast(selectUse.getOwner())) { + bitcastOp.getResult().replaceAllUsesWith(originAtomicOp); + } + } + } + rewriter.replaceOp(op, originAtomicOp); + } else { + rewriter.eraseOp(op); + } + + return success(); +} + +StoreConverter::StoreConverter(MLIRContext *context) + : OpConversionPattern(context) {} + +LogicalResult +StoreConverter::matchAndRewrite(triton::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + + // triton store op basic + auto mask = op.getMask(); + auto loc = op.getLoc(); + auto ptr = adaptor.getPtr(); + auto val = adaptor.getValue(); + + if (op->hasAttr(ConverterUtils::discreteAttrName)) { + Operation *loop = op->getParentOp(); + int extractedLoopCount = 1; + for (auto parentOp = loop->getParentOp(); + parentOp->hasAttr("ExtractedLoadOrStore"); + parentOp = parentOp->getParentOp()) { + loop = parentOp; + extractedLoopCount++; + } + + auto valType = dyn_cast(val.getType()); + if (valType) { + auto valShape = valType.getShape(); + bool isIndexPutScenario = + (extractedLoopCount == 1) && (valShape.size() > 1u); + if (isIndexPutScenario) { + auto loopOp = cast(loop); + loopOp->setAttr("hivm.parallel_loop", rewriter.getUnitAttr()); + } + } + } + + // 1. boundary size check + auto boundaryCheck = op.getBoundaryCheck(); + if (!boundaryCheck.empty()) { + auto boundarySizes = mlir::ConverterUtils::getBoundarySizes( + boundaryCheck, /*remapped*/ ptr, loc, rewriter); + auto srcSlice = mlir::ConverterUtils::makeExtractSliceOp(val, boundarySizes, + loc, rewriter); + auto dstSubview = + mlir::ConverterUtils::makeSubViewOp(ptr, boundarySizes, loc, rewriter); + auto storeOp = rewriter.create( + loc, srcSlice, dstSubview); + storeOp.setWritable(true); + rewriter.eraseOp(op); + return success(); + } + + // 2. Simple load with no mask + if (!mask) { + auto storeOp = rewriter.create( + loc, val, ptr); + storeOp.setWritable(true); + rewriter.eraseOp(op); + return success(); + } + + // 3. Continuous masked stores. + // Analyze the mask operand to determine at runtime the size of the data we + // are moving. + MaskState mstate; + auto isContMask = mstate.parse(mask, loc, rewriter); + + if (isContMask.failed()) { + return failure(); + } + LLVM_DEBUG({ llvm::dbgs() << *getModuleOpFromOperation(op) << "\n"; }); + auto srcSlice = mstate.getExtractSlice(val, loc, rewriter); + auto dstSubview = mstate.getSubview(ptr, loc, rewriter); + auto storeOp = rewriter.create( + loc, srcSlice, dstSubview); + storeOp.setWritable(true); + rewriter.eraseOp(op); + return success(); +} + +} // namespace LoadStoreConverter diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/MaskAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/MaskAnalysis.cpp new file mode 100755 index 00000000..762557a4 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/MaskAnalysis.cpp @@ -0,0 +1,667 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Operation.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" +#include +#include + +#define DEBUG_TYPE "mask-analysis" + +namespace mlir { + +namespace triton { + +namespace Incubated { + +template +std::optional runMaskAnalysisImpl(MemAccOpTy op, + OpBuilder &builder) { + auto mask = op.getMask(); + if (!mask) { + return std::nullopt; + } + + PatternRewriter::InsertionGuard insertGuard(builder); + builder.setInsertionPoint(op); + + Incubated::MaskState mstate; + if (mstate.parse(mask, op.getLoc(), builder).failed()) { + return std::nullopt; + } + return mstate; +} + +LogicalResult MaskState::parse(Value operand, const Location &loc, + OpBuilder &builder) { + if (isa(operand.getType())) { + return parseIntScalar(operand, loc, builder); + } + + if (auto blockArgument = dyn_cast(operand)) { + auto parentOp = blockArgument.getOwner()->getParentOp(); + if (auto loopOp = dyn_cast(parentOp)) { + OpOperand *initArgOperand = loopOp.getTiedLoopInit(blockArgument); + if (initArgOperand) { + Value initArg = initArgOperand->get(); + return parse(initArg, loc, builder); + } + } + } + + auto definingOp = operand.getDefiningOp(); + if (!definingOp) + return failure(); + + LLVM_DEBUG({ + llvm::dbgs() << "[MaskState]==> parse op\n" + << *definingOp << "\n[MaskState]<==\n"; + }); + return TypeSwitch(definingOp) + .Case( + [&](auto op) { return this->parseConstant(op, loc, builder); }) + .Case( + [&](auto op) { return this->parseAdd(op, loc, builder); }) + .Case( + [&](auto op) { return this->parseAnd(op, loc, builder); }) + .Case( + [&](auto op) { return this->parseCmp(op, loc, builder); }) + .Case( + [&](auto op) { return this->parseMakeRange(op, loc, builder); }) + .Case( + [&](auto op) { return this->parseBroadcast(op, loc, builder); }) + .Case( + [&](auto op) { return this->parseSplat(op, loc, builder); }) + .Case( + [&](auto op) { return this->parseExpandDims(op, loc, builder); }) + .Case( + [&](auto op) { return this->parse(op.getIn(), loc, builder); }) + .Case( + [&](auto op) { return this->parseDiv(op, loc, builder); }) + .Case( + [&](auto op) { return this->parseSel(op, loc, builder); }) + .Default([&](Operation *op) { return failure(); }); +} + +// extractSlice +tensor::ExtractSliceOp MaskState::getExtractSlice(Value source, + const Location &loc, + OpBuilder &builder) const { + auto sourceRType = cast(source.getType()); + SmallVector strides(getRank(), builder.getIndexAttr(1)); + + auto dstRType = tensor::ExtractSliceOp::inferResultType(sourceRType, offsets, + dims, strides); + return builder.create(loc, dstRType, source, offsets, + dims, strides); +} + +tensor::InsertSliceOp MaskState::getInsertSlice(Value source, Value dest, + const Location &loc, + OpBuilder &builder) const { + SmallVector strides(getRank(), builder.getIndexAttr(1)); + return builder.create(loc, source, dest, offsets, dims, + strides); +} + +memref::SubViewOp MaskState::getSubview(Value source, const Location &loc, + OpBuilder &builder) const { + auto sourceType = cast(source.getType()); + SmallVector strides(getRank(), builder.getIndexAttr(1)); + auto dstType = + memref::SubViewOp::inferResultType(sourceType, offsets, dims, strides); + return builder.create(loc, cast(dstType), + source, offsets, dims, strides); +} + +static memref::SubViewOp createSubview(Value src, const Location &loc, + OpBuilder &builder, + ArrayRef offsets, + ArrayRef sizes, + ArrayRef strides) { + auto srcType = cast(src.getType()); + auto dstType = + memref::SubViewOp::inferResultType(srcType, offsets, sizes, strides); + return builder.create(loc, cast(dstType), src, + offsets, sizes, strides); +} + +LogicalResult MaskState::addStateScalar(const MaskState &state, + const OpFoldResult scalar, + const Location &loc, + OpBuilder &builder) { + start = addOpFoldResult(state.start, scalar, loc, builder); + end = addOpFoldResult(state.end, scalar, loc, builder); + dims = state.dims; + offsets = state.offsets; + return success(); +} + +LogicalResult MaskState::addStates(const MaskState &lhsState, + const MaskState &rhsState, + const Location &loc, OpBuilder &builder) { + if (lhsState.scalar && rhsState.scalar) { + InFlightDiagnostic diag = + emitWarning(loc) + << "Unexpected case where both lhs and rhs are scalars"; + return failure(); + } + if (!lhsState.scalar && !rhsState.scalar) { + InFlightDiagnostic diag = + emitWarning(loc) + << "Unsupported scenario where neither lhs nor rhs is a scalar"; + return failure(); + } + + if (lhsState.scalar) { + return addStateScalar(rhsState, lhsState.scalar, loc, builder); + } else { + return addStateScalar(lhsState, rhsState.scalar, loc, builder); + } +} + +LogicalResult MaskState::divStateScalar(const MaskState &state, + const OpFoldResult scalar, + const Location &loc, + OpBuilder &builder) { + start = divOpFoldResult(state.start, scalar, loc, builder); + end = divOpFoldResult(state.end, scalar, loc, builder); + dims = state.dims; + offsets = state.offsets; + return success(); +} + +LogicalResult MaskState::divStates(const MaskState &lhsState, + const MaskState &rhsState, + const Location &loc, OpBuilder &builder) { + if (!lhsState.scalar && rhsState.scalar) { + if (isZeroIndex(rhsState.scalar)) { + InFlightDiagnostic diag = + emitError(loc) + << "Unsupported scenario where rhs is zero constant in divide!"; + return failure(); + } + + return divStateScalar(lhsState, rhsState.scalar, loc, builder); + } + + InFlightDiagnostic diag = emitWarning(loc) + << "Supported scenario where only rhs is a scalar"; + return failure(); +} + +LogicalResult MaskState::minStates(const MaskState &lhsState, + const MaskState &rhsState, + const Location &loc, OpBuilder &builder) { + if (lhsState.getRank() != rhsState.getRank()) { + InFlightDiagnostic diag = + emitError(loc) + << "Unexpected case where lhs and rhs have different ranks"; + return failure(); + } + + for (uint32_t i = 0; i < lhsState.getRank(); i++) { + auto lhsOffset = lhsState.offsets[i]; + auto rhsOffset = rhsState.offsets[i]; + auto newOffset = maxOpFoldResult(lhsOffset, rhsOffset, loc, builder); + auto lhsDim = lhsState.dims[i]; + auto rhsDim = rhsState.dims[i]; + auto lhsEnd = addOpFoldResult(lhsOffset, lhsDim, loc, builder); + auto rhsEnd = addOpFoldResult(rhsOffset, rhsDim, loc, builder); + auto newEnd = minOpFoldResult(lhsEnd, rhsEnd, loc, builder); + auto newDim = subOpFoldResult(newEnd, newOffset, loc, builder); + + offsets.push_back(newOffset); + dims.push_back(newDim); + } + return success(); +} + +// Helper func for MaskState::parse() +LogicalResult MaskState::parseConstant(arith::ConstantOp constOp, + const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + if (isa(constOp.getValue())) { + auto attr = cast(constOp.getValue()); + auto elementType = attr.getElementType(); + assert(attr.isSplat() && isa(elementType) && + "All elements must share a single integer constant value"); + + if (elementType.isInteger(1) && + isa(constOp.getValue().getType())) { + auto shapedType = cast(constOp.getValue().getType()); + auto shape = shapedType.getShape(); + for (size_t i = 0; i < shape.size(); i++) { + this->dims.push_back(builder.getIndexAttr(shape[i])); + this->offsets.push_back(builder.getIndexAttr(0)); + } + } else { + this->scalar = builder.getIndexAttr( + attr.getSplatValue().getValue().getSExtValue()); + } + } else { + auto value = cast(constOp.getValue()).getInt(); + this->scalar = builder.getIndexAttr(value); + } + return success(); +} + +// parseIntScalar +LogicalResult MaskState::parseIntScalar(Value scalar, const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + this->scalar = getOpFoldResultOfLayoutInfo(scalar, builder); + return success(); +} + +LogicalResult MaskState::parseAdd(arith::AddIOp addOp, const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + MaskState lhsState; + if (failed(lhsState.parse(addOp.getLhs(), loc, builder))) { + return failure(); + } + + MaskState rhsState; + if (failed(rhsState.parse(addOp.getRhs(), loc, builder))) { + return failure(); + } + return this->addStates(lhsState, rhsState, loc, builder); +} + +LogicalResult MaskState::parseDiv(arith::DivSIOp divOp, const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + return failure(); // temporarily disable parseDiv + MaskState lhsState; + if (failed(lhsState.parse(divOp.getLhs(), loc, builder))) { + return failure(); + } + + MaskState rhsState; + if (failed(rhsState.parse(divOp.getRhs(), loc, builder))) { + return failure(); + } + return this->divStates(lhsState, rhsState, loc, builder); +} + +LogicalResult MaskState::parseAnd(arith::AndIOp andOp, const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + MaskState lhsState; + if (failed(lhsState.parse(andOp.getLhs(), loc, builder)) || + !lhsState.isMask()) { + return failure(); + } + + MaskState rhsState; + if (failed(rhsState.parse(andOp.getRhs(), loc, builder)) || + !rhsState.isMask()) { + return failure(); + } + + if (!lhsState.isMask() && !rhsState.isMask()) { + return failure(); + } + + // Only support both lhs and rhs satisfy `isMask` condition + return this->minStates(lhsState, rhsState, loc, builder); +} + +LogicalResult MaskState::parseSel(arith::SelectOp selOp, const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + auto trueValue = selOp.getTrueValue(); + auto falseValue = selOp.getFalseValue(); + + MaskState condState; + auto condition = selOp.getCondition(); + auto cmpOp = condition.getDefiningOp(); + if (!cmpOp || failed(condState.parse(condition, loc, builder))) { + return failure(); + } + + MaskState trueState; + if (failed(trueState.parse(trueValue, loc, builder)) || !trueState.scalar) { + return failure(); + } + + MaskState falseState; + if (failed(falseState.parse(falseValue, loc, builder)) || + !falseState.scalar) { + return failure(); + } + + auto trueScalar = dyn_cast(trueState.scalar.get()); + auto falseScalar = dyn_cast(falseState.scalar.get()); + + if (trueScalar && falseScalar) { + if (trueScalar.getInt() == 1 && falseScalar.getInt() == 0) { + start = condState.start; + end = condState.end; + dims = condState.dims; + offsets = condState.offsets; + return success(); + } + } + + return failure(); +} + +LogicalResult MaskState::parseCmp(arith::CmpIOp cmpOp, const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + auto predicate = cmpOp.getPredicate(); + // Only support <, <=, >=, =, != + if (predicate != arith::CmpIPredicate::slt && + predicate != arith::CmpIPredicate::sle && + predicate != arith::CmpIPredicate::sge && + predicate != arith::CmpIPredicate::eq && + predicate != arith::CmpIPredicate::ne) { + LLVM_DEBUG({ llvm::dbgs() << "Unsupported cmpi predicate\n"; }); + return failure(); + } + + MaskState lhsState; + MaskState rhsState; + auto lhs = cmpOp.getLhs(); + auto rhs = cmpOp.getRhs(); + + if (predicate == arith::CmpIPredicate::ne) { + auto selOp = lhs.getDefiningOp(); + auto constantOp = rhs.getDefiningOp(); + if (!selOp || !constantOp) { + return failure(); + } + } + + if (failed(lhsState.parse(lhs, loc, builder))) { + return failure(); + } + + if (failed(rhsState.parse(rhs, loc, builder))) { + return failure(); + } + + if (!(!lhsState.scalar && rhsState.scalar)) { + InFlightDiagnostic diag = emitWarning(loc) + << "[MaskState] Unsupported cmpi scenario"; + return failure(); + } + + int32_t cmpDim = -1; + for (int32_t i = 0; i < lhsState.getRank(); i++) { + auto constDimLength = getConstantIntValue(lhsState.dims[i]); + if (!constDimLength || constDimLength.value() != 1) { + if (cmpDim != -1) { + InFlightDiagnostic diag = emitWarning(loc) + << "Unsupported cmpi with more than one " + "dimension with size larger than 1"; + return failure(); + } + cmpDim = i; + } + } + + assert(cmpDim != -1 && + "Unexpected case where no dimension has size larger than 1"); + + this->offsets = lhsState.offsets; + this->dims = lhsState.dims; + switch (predicate) { + case arith::CmpIPredicate::slt: { + auto realBound = + maxOpFoldResult(lhsState.start, rhsState.scalar, loc, builder); + auto newEnd = minOpFoldResult(lhsState.end, realBound, loc, builder); + auto newDim = subOpFoldResult(newEnd, lhsState.start, loc, builder); + + this->dims[cmpDim] = newDim; + break; + } + case arith::CmpIPredicate::sle: { + // lhs <= rhs <=> lhs < rhs + 1 + auto rhsPlusOne = + addOpFoldResult(rhsState.scalar, builder.getIndexAttr(1), loc, builder); + auto realBound = maxOpFoldResult(lhsState.start, rhsPlusOne, loc, builder); + auto newEnd = minOpFoldResult(lhsState.end, realBound, loc, builder); + auto newDim = subOpFoldResult(newEnd, lhsState.start, loc, builder); + + this->dims[cmpDim] = newDim; + break; + } + case arith::CmpIPredicate::sge: { + auto realBound = + maxOpFoldResult(lhsState.start, rhsState.scalar, loc, builder); + auto newStart = minOpFoldResult(lhsState.end, realBound, loc, builder); + auto newOffset = subOpFoldResult(newStart, lhsState.start, loc, builder); + auto newDim = subOpFoldResult(lhsState.end, newStart, loc, builder); + + this->offsets[cmpDim] = newOffset; + this->dims[cmpDim] = newDim; + break; + } + case arith::CmpIPredicate::eq: { + auto newOffset = + subOpFoldResult(rhsState.scalar, lhsState.start, loc, builder); + auto newDim = builder.getIndexAttr(1); + + this->offsets[cmpDim] = newOffset; + this->dims[cmpDim] = newDim; + break; + } + case arith::CmpIPredicate::ne: { + // only support lhs != 0 + auto rhsScalar = dyn_cast(rhsState.scalar.get()); + if (!rhsScalar || rhsScalar.getInt() != 0) { + return failure(); + } + + start = lhsState.start; + end = lhsState.end; + break; + } + default: + return failure(); + } + return success(); +} + +LogicalResult MaskState::parseMakeRange(triton::MakeRangeOp rangeOp, + const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + auto shape = cast(rangeOp.getType()).getShape(); + auto start = rangeOp.getStart(); + auto end = rangeOp.getEnd(); + auto stride = (end - start + shape[0] - 1) / shape[0]; + + if (stride != 1) { + InFlightDiagnostic diag = + emitWarning(loc) + << "stride must be 1 for make_range whose result is used " + "as load or store masks"; + return failure(); + } + + this->start = builder.getIndexAttr(start); + this->end = builder.getIndexAttr(end); + this->dims.push_back(builder.getIndexAttr(shape[0])); + this->offsets.push_back(builder.getIndexAttr(0)); + return success(); +} + +LogicalResult MaskState::parseBroadcast(triton::BroadcastOp broadcastOp, + const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + auto src = broadcastOp.getSrc(); + auto dst = broadcastOp.getResult(); + assert(isa(src.getType()) && + "input to tt.broadcast should be a tensor"); + + auto srcShape = cast(src.getType()).getShape(); + auto dstShape = cast(dst.getType()).getShape(); + assert(srcShape.size() == dstShape.size() && + "rank of source and destination should match"); + + if (failed(parse(src, loc, builder))) { + return failure(); + } + for (size_t i = 0; i < srcShape.size(); i++) { + if (srcShape[i] == dstShape[i]) + continue; + else if (srcShape[i] < dstShape[i]) + this->dims[i] = builder.getIndexAttr(dstShape[i]); + else + llvm_unreachable("unexpected dimensions used in broadcast"); + } + return success(); +} + +LogicalResult MaskState::parseSplat(triton::SplatOp splatOp, + const Location &loc, OpBuilder &builder) { + assert(this->isEmpty()); + + auto src = splatOp.getSrc(); + auto dst = splatOp.getResult(); + auto dstShape = cast(dst.getType()).getShape(); + + if (!isa(src.getType())) { + InFlightDiagnostic diag = + emitWarning(loc) + << "splat source must be an integer scalar for load/store masks"; + return failure(); + } + + if (failed(this->parse(src, loc, builder))) + return failure(); + + auto splatAsMask = [&](Operation *userOp) -> bool { + return TypeSwitch(userOp) + .Case([&](arith::AndIOp andOp) { return true; }) + .Case([&](arith::SelectOp selectOp) { + return selectOp.getCondition() == dst; + }) + .Case( + [&](triton::LoadOp loadOp) { return loadOp.getMask() == dst; }) + .Case( + [&](triton::StoreOp storeOp) { return storeOp.getMask() == dst; }) + .Default([&](Operation *op) { return false; }); + }; + + if (src.getType().isInteger(1) && !splatOp->use_empty() && + llvm::all_of(splatOp->getUsers(), splatAsMask)) { + for (auto s : dstShape) { + auto currentDim = + mulOpFoldResult(builder.getIndexAttr(s), this->scalar, loc, builder); + this->dims.push_back(currentDim); + this->offsets.push_back(builder.getIndexAttr(0)); + } + + this->scalar = nullptr; + return success(); + } + + for (auto s : dstShape) { + this->dims.push_back(builder.getIndexAttr(s)); + this->offsets.push_back(builder.getIndexAttr(0)); + } + return success(); +} + +LogicalResult MaskState::parseExpandDims(triton::ExpandDimsOp expandDimsOp, + const Location &loc, + OpBuilder &builder) { + assert(this->isEmpty()); + + if (failed(this->parse(expandDimsOp.getSrc(), loc, builder))) { + return failure(); + } + + auto dstShape = + cast(expandDimsOp.getResult().getType()).getShape(); + auto axis = expandDimsOp.getAxis(); + assert(dstShape[axis] == 1 && + "Expect changed dimention to be 1 in expand_dims"); + this->dims.insert(this->dims.begin() + axis, builder.getIndexAttr(1)); + this->offsets.insert(this->offsets.begin() + axis, builder.getIndexAttr(0)); + + return success(); +} + +void MaskState::eraseInsertedOps(Operation *rawOp, PatternRewriter &rewriter) { + auto moduleOp = rawOp->getParentOfType(); + SmallVector worklist; + moduleOp->walk([&](Operation *op) { + if (isOpTriviallyDead(op)) + worklist.push_back(op); + }); + while (!worklist.empty()) { + Operation *op = worklist.pop_back_val(); + if (!isOpTriviallyDead(op)) + continue; + for (Value value : op->getOperands()) { + if (auto defOp = value.getDefiningOp()) + worklist.push_back(defOp); + } + LLVM_DEBUG({ + llvm::dbgs() << "[MaskState]==> inserted op: \n" + << *op << "\n[MaskState]<== is removed\n"; + }); + rewriter.eraseOp(op); + } +} + +std::optional runMaskAnalysis(Operation *op, + OpBuilder &builder) { + if (auto loadOp = dyn_cast(op)) { + return runMaskAnalysisImpl(loadOp, builder); + } + if (auto storeOp = dyn_cast(op)) { + return runMaskAnalysisImpl(storeOp, builder); + } + if (auto atomicRMWOp = dyn_cast(op)) { + return runMaskAnalysisImpl(atomicRMWOp, builder); + } + return std::nullopt; +} + +} // namespace Incubated + +} // namespace triton + +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/TritonOpConverter.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/TritonOpConverter.cpp new file mode 100755 index 00000000..f7a59b4c --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/TritonOpConverter.cpp @@ -0,0 +1,2670 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/TritonOpConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h" +#include "incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/StringRef.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/LogicalResult.h" +#include "llvm/Support/raw_ostream.h" +#include + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/MemRef/Transforms/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/ValueRange.h" +#if __has_include("bishengir/Dialect/Annotation/IR/Annotation.h") +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#endif +#if __has_include("bishengir/Dialect/HFusion/IR/HFusion.h") +#include "bishengir/Dialect/HFusion/IR/HFusion.h" +#endif + +namespace TTOpConverters { +using namespace mlir; +using namespace triton; + +static llvm::SmallString +generateUniqueFuncName(ModuleOp moduleOp, llvm::StringRef funcNameBase) { + llvm::SmallString funcName = funcNameBase; + int uniqueId = 0; + while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) { + funcName = funcNameBase; + funcName += ("_" + std::to_string(uniqueId++)); + } + return funcName; +} + +LogicalResult +BitcastConverter::matchAndRewrite(triton::BitcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Value result; + auto loc = op.getLoc(); + + if (auto dstPtrTy = dyn_cast(op.getType())) { + auto srcPtrTy = cast(op.getSrc().getType()); + auto resType = + MemRefType::get({ShapedType::kDynamic}, dstPtrTy.getPointeeType()); + + auto i1Ty = rewriter.getIntegerType(1); + auto i8Ty = rewriter.getIntegerType(8); + bool isI1toI8 = (srcPtrTy.getPointeeType() == i1Ty) && + (dstPtrTy.getPointeeType() == i8Ty); + // handling special case: ptr -> ptr, directly forward without + // arith.bitcast + if (isI1toI8) { + // TypeConverter has already converted i1 to i8 memref, + LLVM_DEBUG({ + llvm::dbgs() + << "[BitcastConverter] Special i1->i8 pointer bitcast. Forward " + "without arith.bitcast. srcConvertedTy=" + << adaptor.getSrc().getType() << "\n"; + }); + rewriter.replaceOp(op, adaptor.getSrc()); + return success(); + } + result = rewriter.create(loc, resType, adaptor.getSrc()); + } else { + // handling normal case: bitcast between tensors/memrefs + result = + rewriter.create(loc, op.getType(), adaptor.getSrc()); + } + rewriter.replaceOp(op, result); + return success(); +} + +LogicalResult +TransposeConverter::matchAndRewrite(triton::TransOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto src = adaptor.getSrc(); + auto res = ConverterUtils::getTransposedValue(src, op.getLoc(), rewriter, + op.getOrder()); + rewriter.replaceOp(op, res); + return success(); +} + +LogicalResult +YieldConverter::matchAndRewrite(scf::YieldOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + rewriter.replaceOpWithNewOp(op, adaptor.getOperands()); + return success(); +} + +LogicalResult +AdvanceConverter::matchAndRewrite(triton::AdvanceOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + llvm::SmallDenseMap known; + BlockDataParser::rewriteAdvanceOp(op, rewriter, known); + return success(); +} + +// ToDo: +// 1. Refactor MakeTensorPtrConverter and AdvanceConverter with +// memref::ReinterpretCastOp and memref::SubViewOp. +// Use recast to describe full shape of tensor, and use subview to represent +// current block tensor. +LogicalResult MakeTensorPtrConverter::matchAndRewrite( + triton::MakeTensorPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + llvm::SmallDenseMap known; + BlockDataParser::rewriteMakeTensorPtrOp(op, adaptor.getBase(), rewriter, + known); + return success(); +} + +LogicalResult PreciseDivConverter::matchAndRewrite( + triton::PreciseDivFOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Value opa = op.getX(); + Value opb = op.getY(); + auto loc = op.getLoc(); + + auto resType = dyn_cast(op.getResult().getType()); + auto divOp = rewriter.create(loc, resType, opa, opb); + + rewriter.replaceOp(op, divOp); + return success(); +} + +LogicalResult +SelectCanonicalizer::matchAndRewrite(arith::SelectOp op, + PatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + // 0. Shortcut for scalars + auto type = dyn_cast(op.getResult().getType()); + if (!type) { + // do nothing non-tensor select + return failure(); + } + auto mask = op.getCondition(); + if (!isa(mask.getType())) { + // do nothing for scalar mask + return failure(); + } + + // 1. Check for continuous masked loads. + // Analyze the mask operand to determine at runtime the size of the data we + // are moving. + mlir::triton::Incubated::MaskState mstate; + auto isContMask = mstate.parse(mask, loc, rewriter); + + if (isContMask.failed()) { + mstate.eraseInsertedOps(op, rewriter); + return rewriter.notifyMatchFailure( + op, "Cannot lower continuous masked selects"); + } + + // 2. Slice out the masked part of true tensor + auto trueTensor = op.getTrueValue(); + auto extractSliceOp = mstate.getExtractSlice(trueTensor, loc, rewriter); + + // 3. Insert out the sliced true tensor into false tensor + auto falseTensor = op.getFalseValue(); + auto insertSliceOp = + mstate.getInsertSlice(extractSliceOp, falseTensor, loc, rewriter); + + // 4. Fix if the offset is negative at runtime + rewriter.setInsertionPointAfter(insertSliceOp); + Value zeroIndex = rewriter.create(loc, 0); + auto offsets = mstate.offsets; + SmallVector isInvalidVals; + for (size_t i = 0; i < offsets.size(); i++) { + auto &o = offsets[i]; + if (o.is()) { + auto oVal = o.get(); + int64_t dimSize = type.getShape()[i]; + Value sizeIndex = rewriter.create(loc, dimSize); + Value isNegative = rewriter.create( + loc, arith::CmpIPredicate::slt, oVal, zeroIndex); + Value isOutOfRange = rewriter.create( + loc, arith::CmpIPredicate::sge, oVal, sizeIndex); + isInvalidVals.push_back(isNegative); + isInvalidVals.push_back(isOutOfRange); + } + } + + if (isInvalidVals.empty()) { + rewriter.replaceOp(op, insertSliceOp); + return success(); + } + // At least one value + Value invalidVal = isInvalidVals[0]; + if (isInvalidVals.size() > 1) { + for (int i = 1; i < isInvalidVals.size(); ++i) { + auto tmpOrOp = + rewriter.create(loc, isInvalidVals[i], invalidVal); + invalidVal = tmpOrOp.getResult(); + } + } + // else: what if the number of negative value checks is > 1? + auto ifOp = rewriter.create(loc, TypeRange{falseTensor.getType()}, + invalidVal, true /* addThenBlock */, + true /* addElseBlock */); + // thenBuilder + rewriter.setInsertionPointToStart(&ifOp.getThenRegion().front()); + rewriter.create(loc, ValueRange{falseTensor}); + // elseBuilder + Block *elseBlock = &ifOp.getElseRegion().front(); + extractSliceOp->moveBefore(elseBlock, elseBlock->begin()); + insertSliceOp->moveBefore(elseBlock, elseBlock->end()); + rewriter.setInsertionPointToStart(elseBlock); + { + rewriter.setInsertionPointAfter(insertSliceOp); + rewriter.create(loc, ValueRange{insertSliceOp.getResult()}); + } + + rewriter.replaceOp(op, ifOp); + + return success(); +} + +/* + * Move tt.bitcast to a previous location if tt.bitcast is not directly applied + * on function arguments + */ +LogicalResult +BitcastCanonicalizer::matchAndRewrite(triton::BitcastOp bitcastOp, + PatternRewriter &rewriter) const { + Value castSrc = bitcastOp.getSrc(); + Value castRes = bitcastOp.getResult(); + Type castSrcTy = castSrc.getType(); + Type castSrcPtrTy = isa(castSrcTy) + ? cast(castSrcTy).getElementType() + : castSrcTy; + if (!isa(castSrcPtrTy)) + return failure(); + + auto origBitwidth = getPointeeBitWidth(castSrc.getType()); + auto castBitwidth = getPointeeBitWidth(castRes.getType()); + + if (origBitwidth == 1) + origBitwidth = 8; + if (castBitwidth == 1) + castBitwidth = 8; + if (origBitwidth != castBitwidth) { + bitcastOp.emitError() << "Casting pointers with unmatched bitwidth!\n"; + return failure(); + } + + Operation *beforeCastOp = castSrc.getDefiningOp(); + if (beforeCastOp == nullptr) { + return failure(); + } + + auto newRes = + TypeSwitch>(beforeCastOp) + // before: addptr - bitcast - load/store + // after: bitcast - addptr - load/store + .Case([&](triton::AddPtrOp addptrOp) { + auto newCastOp = rewriter.create( + bitcastOp.getLoc(), castRes.getType(), addptrOp.getPtr()); + return rewriter.create( + bitcastOp.getLoc(), castRes.getType(), newCastOp.getResult(), + addptrOp.getOffset()); + }) + .Case([&](triton::SplatOp splatOp) { + Type newCastSrcTy = + cast(castRes.getType()).getElementType(); + + Value splatSrc = splatOp.getSrc(); + Type splatSrcTy = splatSrc.getType(); + if (auto splatSrcTensorTy = dyn_cast(splatSrcTy)) + newCastSrcTy = + splatSrcTensorTy.cloneWith(std::nullopt, newCastSrcTy); + auto newCastOp = rewriter.create( + bitcastOp.getLoc(), newCastSrcTy, splatSrc); + return rewriter.create( + bitcastOp.getLoc(), castRes.getType(), newCastOp); + }) + // before: bitcast - bitcast + // after(fusion optimization): bitcast + .Case([&](triton::BitcastOp prevCastOp) { + return rewriter.create( + bitcastOp.getLoc(), castRes.getType(), prevCastOp.getSrc()); + }) + .Default([&](Operation *op) { + return rewriter.notifyMatchFailure(bitcastOp, + "Unknown bitcast pattern"); + }); + if (succeeded(newRes)) { + rewriter.replaceOp(bitcastOp, newRes.value()); + if (beforeCastOp->use_empty()) { + rewriter.eraseOp(beforeCastOp); + } + return success(); + } + return failure(); +} + +LogicalResult +FpToFpCanonicalizer::matchAndRewrite(triton::FpToFpOp op, + PatternRewriter &rewriter) const { + auto loc = op.getLoc(); + Value input = op.getSrc(); + auto resultType = op.getResult().getType(); + + // Check if rounding mode is specified + auto roundingMode = op.getRounding(); + if (roundingMode.has_value() && + roundingMode.value() != triton::RoundingMode::RTNE) { + // Non-RTNE rounding modes (e.g., RTZ) should be handled by TritonToHFusion + // pass Return failure here so this pattern doesn't match + return failure(); + } + + // Handle RTNE (default) rounding mode with arith.truncf/extf + auto srcType = cast(input.getType()); + auto dstType = cast(resultType); + auto srcElemType = srcType.getElementType(); + auto dstElemType = dstType.getElementType(); + if (!isa(srcElemType) || !isa(dstElemType)) { + return op.emitError("FpToFp expects floating point types"); + } + + unsigned srcBitwidth = srcElemType.getIntOrFloatBitWidth(); + unsigned dstBitwidth = dstElemType.getIntOrFloatBitWidth(); + + // Create round_mode attribute (RINT for RTNE) + auto roundModeAttr = hfusion::RoundModeAttr::get(rewriter.getContext(), + hfusion::RoundMode::RINT); + + if (srcBitwidth > dstBitwidth) { + // Downcast: use arith.truncf with round_mode=rint + auto truncOp = rewriter.create(loc, dstType, input); + truncOp->setAttr("round_mode", roundModeAttr); + rewriter.replaceOp(op, truncOp.getResult()); + } else if (srcBitwidth < dstBitwidth) { + // Upcast: use arith.extf with round_mode=rint + auto extOp = rewriter.create(loc, dstType, input); + extOp->setAttr("round_mode", roundModeAttr); + rewriter.replaceOp(op, extOp.getResult()); + } else { + // Same bitwidth, should not happen but handle gracefully + rewriter.replaceOp(op, input); + } + + return success(); +} + +void rewriteUserWithNewOrder( + mlir::OpOperand *use, PatternRewriter &rewriter, + llvm::SmallVector &blkShapeI64, // 8: container size + mlir::Location &loc, llvm::ArrayRef &order, size_t &orderSize) { + Operation *user = use->getOwner(); + rewriter.setInsertionPointAfter(user); + if (auto loadOp = dyn_cast(user)) { + auto loadResTy = loadOp.getResult().getType(); + auto loadResShapedTy = cast(loadResTy); + auto newLoadTy = loadResShapedTy.cloneWith( + blkShapeI64, loadResShapedTy.getElementType()); + auto newLoadOp = rewriter.create( + loc, newLoadTy, loadOp->getOperands(), loadOp->getAttrs()); + newLoadOp->setAttr(ConverterUtils::GeneratedByMakeTensorPtrTAG, + UnitAttr::get(rewriter.getContext())); + rewriter.replaceOp(loadOp, newLoadOp); + // load contiguous data then permute. thus the permute order is as + // follows. + SmallVector permuteOrder; // 8: container size + for (auto [i, v] : llvm::enumerate(order)) { + permuteOrder.push_back(orderSize - 1 - order[i]); + } + auto permuteOp = rewriter.create( + loc, newLoadOp.getResult(), + DenseI32ArrayAttr::get(loadOp.getContext(), permuteOrder)); + newLoadOp.getResult().replaceAllUsesExcept(permuteOp.getResult(), + permuteOp); + } else if (auto storeOp = dyn_cast(user)) { + // permute to contiguous then store. thus the permute order is as follows. + SmallVector permuteOrder; // 8: container size + for (auto [i, v] : llvm::enumerate(order)) { + permuteOrder.push_back(order[orderSize - 1 - i]); + } + auto permuteOp = rewriter.create( + loc, storeOp.getValue(), + DenseI32ArrayAttr::get(storeOp.getContext(), permuteOrder)); + storeOp.getValue().replaceAllUsesExcept(permuteOp.getResult(), permuteOp); + auto newStoreOp = rewriter.create( + loc, storeOp.getPtr(), storeOp.getValue(), storeOp.getMask(), + storeOp.getBoundaryCheck(), storeOp.getCache(), storeOp.getEvict()); + rewriter.replaceOp(storeOp, newStoreOp); + } else if (auto advanceOp = dyn_cast(user)) { + auto advanceResPtrTy = + cast(advanceOp.getResult().getType()); + auto advanceResShapedTy = + cast(advanceResPtrTy.getPointeeType()); + auto newAdvanceResShapedTy = advanceResShapedTy.cloneWith( + blkShapeI64, advanceResShapedTy.getElementType()); + auto newAdvanceResPtrTy = triton::PointerType::get( + newAdvanceResShapedTy, advanceResPtrTy.getAddressSpace()); + auto advanceOffsets = advanceOp.getOffsets(); + llvm::SmallVector newAdvanceOffsets; // 8: container size + for (int i = orderSize - 1; i >= 0; i--) { + newAdvanceOffsets.push_back(advanceOffsets[order[i]]); + } + SmallVector resUses; + for (auto &use : advanceOp->getUses()) + resUses.push_back(&use); + auto newAdvanceOp = rewriter.create( + loc, newAdvanceResPtrTy, advanceOp.getPtr(), newAdvanceOffsets); + rewriter.replaceOp(advanceOp, newAdvanceOp); + for (auto resUse : resUses) + rewriteUserWithNewOrder(resUse, rewriter, blkShapeI64, loc, order, + orderSize); + } else if (auto loopOp = dyn_cast(user)) { + auto initArg = use->get(); + auto iterArg = loopOp.getTiedLoopRegionIterArg(use); + auto resultValue = loopOp.getTiedLoopResult(use); + iterArg.setType(initArg.getType()); + resultValue.setType(initArg.getType()); + for (auto &argUse : iterArg.getUses()) + rewriteUserWithNewOrder(&argUse, rewriter, blkShapeI64, loc, order, + orderSize); + for (auto &resUse : resultValue.getUses()) + rewriteUserWithNewOrder(&resUse, rewriter, blkShapeI64, loc, order, + orderSize); + } else if (isa(user)) { + return; + } else { + llvm_unreachable( + "[MakeTensorPtrCanonicalizer] tt.make_tensor_ptr's result is " + "not used by load/store/advance op"); + } +} + +void markLoadUsers(mlir::OpOperand *use, PatternRewriter &rewriter) { + Operation *user = use->getOwner(); + if (auto loadOp = dyn_cast(user)) { + loadOp->setAttr(ConverterUtils::GeneratedByMakeTensorPtrTAG, + UnitAttr::get(rewriter.getContext())); + } else if (auto storeOp = dyn_cast(user)) { + return; + } else if (auto advanceOp = dyn_cast(user)) { + SmallVector resUses; + for (auto &use : advanceOp->getUses()) + resUses.push_back(&use); + for (auto resUse : resUses) + markLoadUsers(resUse, rewriter); + } else if (auto loopOp = dyn_cast(user)) { + auto initArg = use->get(); + auto iterArg = loopOp.getTiedLoopRegionIterArg(use); + auto resultValue = loopOp.getTiedLoopResult(use); + iterArg.setType(initArg.getType()); + resultValue.setType(initArg.getType()); + for (auto &argUse : iterArg.getUses()) + markLoadUsers(&argUse, rewriter); + for (auto &resUse : resultValue.getUses()) + markLoadUsers(&resUse, rewriter); + } else if (isa(user)) { + return; + } else { + llvm_unreachable( + "[MakeTensorPtrCanonicalizer] tt.make_tensor_ptr's result is " + "not used by load/store/advance op"); + } +} + +LogicalResult +MakeTensorPtrCanonicalizer::matchAndRewrite(triton::MakeTensorPtrOp op, + PatternRewriter &rewriter) const { + auto order = op.getOrder(); + auto orderSize = order.size(); + if (orderSize == 1) { + return rewriter.notifyMatchFailure( + op, "make_tensor_ptr's order has single value."); + } + + bool isPermuted = false; + for (auto [first, second] : llvm::zip(order.slice(0, orderSize - 1), + order.slice(1, orderSize - 1))) { + if (first != second + 1) { + isPermuted = true; + break; + } + } + + auto loc = op.getLoc(); + auto base = op.getBase(); + auto shape = op.getShape(); + auto strides = op.getStrides(); + auto offsets = op.getOffsets(); + auto result = op.getResult(); + SmallVector opUses; + + for (auto &use : result.getUses()) + opUses.push_back(&use); + for (auto use : opUses) + markLoadUsers(use, rewriter); + + if (!isPermuted) { + return rewriter.notifyMatchFailure( + op, "make_tensor_ptr's order is contiguous."); + } + + llvm::SmallVector blkShapeI32; + llvm::SmallVector blkShapeI64; + auto resPtrType = cast(result.getType()); + if (auto resShapedTy = dyn_cast(resPtrType.getPointeeType())) { + auto resBlkShape = resShapedTy.getShape(); + for (auto [i, v] : llvm::enumerate(resBlkShape)) { + auto reverseI = orderSize - 1 - i; + blkShapeI32.push_back(resBlkShape[order[reverseI]]); + blkShapeI64.push_back(resBlkShape[order[reverseI]]); + } + } + + llvm::SmallVector newShape; + llvm::SmallVector newStrides; + llvm::SmallVector newOffsets; + for (int i = orderSize - 1; i >= 0; i--) { + newShape.push_back(shape[order[i]]); + newStrides.push_back(strides[order[i]]); + newOffsets.push_back(offsets[order[i]]); + } + + llvm::SmallVector contiguousOrder; + for (int i = orderSize - 1; i >= 0; i--) + contiguousOrder.push_back(i); + + rewriter.setInsertionPoint(op); + auto newMakeTensorPtrOp = rewriter.create( + loc, base, ValueRange(newShape), ValueRange(newStrides), + ValueRange(newOffsets), blkShapeI32, contiguousOrder); + rewriter.replaceOp(op, newMakeTensorPtrOp); + for (auto use : opUses) + rewriteUserWithNewOrder(use, rewriter, blkShapeI64, loc, order, orderSize); + return success(); +} + +LogicalResult +ReduceSingleCanonicalizer::matchAndRewrite(triton::ReduceOp reduceOp, + PatternRewriter &rewriter) const { + assert(reduceOp.getSrcs().size() <= 2 && + "Only reduce or reduce with index are supported"); + auto src = reduceOp.getSrcs()[0]; + auto srcType = cast(src.getType()); + auto srcShape = srcType.getShape(); + if (llvm::any_of(srcShape, [](auto s) { return s != 1; })) + return rewriter.notifyMatchFailure( + reduceOp, "reduce's srcs are not all with single element"); + auto loc = reduceOp->getLoc(); + + // Handle Reduce Value + auto res = reduceOp.getResult()[0]; + Value extracted; + if (srcType.getRank() == 1) { + auto zero = + rewriter.create(loc, rewriter.getIndexAttr(0)); + extracted = rewriter.create(loc, src, zero.getResult()) + .getResult(); + } else { + auto resShape = cast(res.getType()).getShape(); + auto collapseReassociationIndicesOptional = + getReassociationIndicesForCollapse(srcShape, resShape); + if (!collapseReassociationIndicesOptional.has_value()) { + return rewriter.notifyMatchFailure( + reduceOp, "Failure with getReassociationIndicesForCollapse call"); + } + auto collapseReassociationIndices = + collapseReassociationIndicesOptional.value(); + extracted = rewriter + .create( + loc, src, collapseReassociationIndices) + .getResult(); + } + res.replaceAllUsesWith(extracted); + + // Handle Reduce Index + if (reduceOp.getSrcs().size() == 1) + return success(); + + auto resIdx = reduceOp.getResult()[1]; + auto zeroI32 = + rewriter.create(loc, rewriter.getI32IntegerAttr(0)); + if (srcType.getRank() == 1) { + resIdx.replaceAllUsesWith(zeroI32); + } else { + auto resIdxShape = cast(resIdx.getType()).getShape(); + auto initTensor = rewriter.create(loc, resIdxShape, + rewriter.getI32Type()); + auto fillOp = rewriter.create(loc, ValueRange{zeroI32}, + ValueRange{initTensor}); + resIdx.replaceAllUsesWith(fillOp.getResult(0)); + } + + return success(); +} + +LogicalResult DenseConstantConverter::matchAndRewrite( + arith::ConstantOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto denseAttr = cast(op.getValue()); + auto loc = op.getLoc(); + auto constSplatOp = arith::ConstantOp::materialize( + rewriter, denseAttr.getSplatValue(), + denseAttr.getElementType(), loc); + auto emptyOp = rewriter.create( + loc, cast(op.getResult().getType()).getShape(), + denseAttr.getElementType()); + + rewriter.replaceOpWithNewOp(op, ValueRange{constSplatOp}, + ValueRange{emptyOp}); + + return success(); +} + +LogicalResult +MakeRangeConverter::matchAndRewrite(triton::MakeRangeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto type = cast(op.getResult().getType()); + auto shape = type.getShape(); + auto elementType = type.getElementType(); + auto context = op.getContext(); + + assert(type.getShape().size() == 1 && + isa(type.getElementType()) && + type.getElementType().getIntOrFloatBitWidth() == 32 && + "make range can only return 1D int32 tensor"); + + SmallVector indexingMaps{AffineMap::get( + /* dimCount */ 1, /* symbolCount */ 0, + {mlir::getAffineDimExpr(0, context)}, context)}; + + auto init = rewriter.create(loc, shape, elementType); + + auto nestedBody = [&](OpBuilder &nestedBuilder, Location nestedLoc, + ValueRange blockArgs) { + Value index = nestedBuilder.create(loc, 0); + Value res = + nestedBuilder.create(loc, elementType, index); + nestedBuilder.create(loc, res); + }; + + auto linalgOp = rewriter.create( + loc, op->getResultTypes(), /* operands */ ValueRange{}, ValueRange{init}, + indexingMaps, ConverterUtils::getNParallelLoopsAttrs(1), nestedBody); + + int32_t startVal = op.getStartAttr().getInt(); + if (startVal == 0) { + rewriter.replaceOp(op, linalgOp->getResults()); + return success(); + } + + // Apply start offset + Value startScaler = rewriter.create( + loc, rewriter.getI32IntegerAttr(static_cast(startVal))); + auto startInit = rewriter.create(loc, shape, elementType); + Value startTensor = rewriter + .create(loc, ValueRange{startScaler}, + ValueRange{startInit}) + .getResult(0); + auto addOp = + rewriter.create(loc, linalgOp->getResult(0), startTensor); + rewriter.replaceOp(op, addOp); + return success(); +} + +LogicalResult +SplatConverter::matchAndRewrite(triton::SplatOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto shape = op.getType().getShape(); + auto init = rewriter.create(loc, shape, + op.getType().getElementType()); + if (llvm::all_of(shape, [](int64_t dim) { return dim == 1; })) { + SmallVector idx(shape.size(), rewriter.create( + loc, rewriter.getIndexAttr(0))); + rewriter.replaceOpWithNewOp(op, adaptor.getSrc(), init, + idx); + } else { + rewriter.replaceOpWithNewOp( + op, ValueRange{adaptor.getSrc()}, ValueRange{init}); + } + return success(); +} + +LogicalResult +ReshapeConverter::matchAndRewrite(triton::ReshapeOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto src = op.getSrc(); + auto dst = op.getResult(); + Value shape = rewriter.create( + loc, + rewriter.getI64TensorAttr(cast(dst.getType()).getShape())); + auto reshapeOp = + rewriter.create(loc, dst.getType(), src, shape); + rewriter.replaceOp(op, reshapeOp.getResult()); + return success(); +} + +LogicalResult ExpandDimsConverter::matchAndRewrite( + triton::ExpandDimsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto src = op.getSrc(); + auto resShape = cast(op.getResult().getType()).getShape(); + auto axis = op.getAxis(); + + SmallVector reassociation; + + auto src_last_dim = resShape.size() - 2; + auto map_func = [&](unsigned i) -> ReassociationIndices { + if (i < axis) { + return i == src_last_dim ? ReassociationIndices{i, i + 1} + : ReassociationIndices{i}; + } + return i == axis ? ReassociationIndices{i, i + 1} + : ReassociationIndices{i + 1}; + }; + + reassociation = llvm::to_vector( + llvm::map_range(llvm::seq(0, src_last_dim + 1), map_func)); + + auto expandShapeOp = rewriter.create( + op.getLoc(), op.getResult().getType(), src, reassociation); + rewriter.replaceOp(op, expandShapeOp.getResult()); + return success(); +} + +LogicalResult +ClampFConverter::matchAndRewrite(triton::ClampFOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto input = adaptor.getX(); + auto min_para = adaptor.getMin(); + auto max_para = adaptor.getMax(); + auto propagateNan_para = adaptor.getPropagateNan(); + + if (auto input_type = dyn_cast(input.getType())) { + if (isa(min_para.getType())) { + auto minEmptyTensor = rewriter.create( + loc, input_type.getShape(), input_type.getElementType()); + min_para = rewriter + .create(loc, ValueRange{min_para}, + ValueRange{minEmptyTensor}) + .result(); + } + if (isa(max_para.getType())) { + auto maxEmptyTensor = rewriter.create( + loc, input_type.getShape(), input_type.getElementType()); + max_para = rewriter + .create(loc, ValueRange{max_para}, + ValueRange{maxEmptyTensor}) + .result(); + } + } + + if (propagateNan_para == PropagateNan::NONE) { + auto minOp = rewriter.create(loc, input, max_para); + auto maxOp = rewriter.create(loc, min_para, minOp); + rewriter.replaceOp(op, ValueRange{maxOp}); + } else if (propagateNan_para == PropagateNan::ALL) { + auto minOp = rewriter.create(loc, input, max_para); + auto maxOp = rewriter.create(loc, min_para, minOp); + rewriter.replaceOp(op, ValueRange{maxOp}); + } else { + return failure(); + } + + return success(); +} + +// Here convert tt.broadcast to linalg.broadcast +// +// before +// %out = tt.broadcast %in : tensor<1x4x8xf32> -> tensor<128x4x8xf32> +// +// after +// %collpased = tensor.collapse_shape %in [[0, 1], [2]] : +// tensor<1x4x8xf32> into tensor<4x8xf32> +// %out = linalg.broadcast ins(%collpased : tensor<4x8xf32>) +// outs(%empty : tensor<128x4x8xf32>) dimensions = [0] +LogicalResult +BroadcastConverter::matchAndRewrite(triton::BroadcastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + assert(op->getNumResults() == 1 && "BroadcastOp assumes single result"); + + RankedTensorType sourceType = + cast(adaptor.getSrc().getType()); + RankedTensorType resultType = cast(op.getType()); + auto elementType = resultType.getElementType(); + auto loc = op.getLoc(); + + auto initEmpty = + rewriter.create(loc, resultType.getShape(), elementType); + + SmallVector broadcastDims = + ConverterUtils::getBroadcastDims(sourceType, resultType); + SmallVector unbroadcastDims = + ConverterUtils::getUnbroadcastDims(sourceType, resultType); + + SmallVector collapseReassociationIndices; + auto collapseReassociationIndicesOptional = + getReassociationIndicesForCollapse(sourceType.getShape(), + unbroadcastDims); + if (!collapseReassociationIndicesOptional.has_value()) { + return rewriter.notifyMatchFailure( + op, "Failure with getReassociationIndicesForCollapse call"); + } + collapseReassociationIndices = collapseReassociationIndicesOptional.value(); + + RankedTensorType collapseResultType = + RankedTensorType::get(unbroadcastDims, sourceType.getElementType()); + + auto collpasedOp = rewriter.create( + loc, collapseResultType, adaptor.getSrc(), collapseReassociationIndices); + + auto broadcastOp = rewriter.create( + loc, collpasedOp, initEmpty, + rewriter.getDenseI64ArrayAttr(broadcastDims)); + + rewriter.replaceOp(op, broadcastOp.getResults()); + return success(); +} + +// Reduce Converter +bool ReduceConverter::isReductionOpSupported(Operation *redOp) const { + return isa(redOp); +} + +LogicalResult +ReduceConverter::convertToTargetOp(triton::ReduceOp op, + typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto source = adaptor.getOperands().front(); + auto sourceType = cast(source.getType()); + auto elemType = sourceType.getElementType(); + auto resType = op.getResult().front().getType(); + auto loc = op.getLoc(); + auto reductionOps = this->getRedOps(op); + + // Reduction of arbitrary operations isn't supported because using the first + // element across the reduction dimension requires us to iterate over a + // subview that skips over each first element. + if (!this->isReductionOpSupported(reductionOps.front())) { + return rewriter.notifyMatchFailure( + op, "Only support lowering reduction with single op and limited types " + "of reducetion"); + } + + auto rop = reductionOps.front(); + auto axis = op.getAxis(); + auto isVectorReduce = sourceType.getRank() == 1; + + auto constantType = elemType; + + auto accBaseConstOp = this->getRedBaseConstOp(rewriter, rop, constantType); + Value initTensor; + + if (isVectorReduce) { + auto holder = rewriter.create( + loc, RankedTensorType::get({}, constantType), ValueRange{}); + initTensor = rewriter + .create(loc, accBaseConstOp.getResult(), + holder.getResult()) + .getResult(0); + } else { + Value init = rewriter.create( + loc, cast(resType).getShape(), constantType); + initTensor = + rewriter.create(loc, accBaseConstOp.getResult(), init) + .getResult(0); + } + + Value finalResult = + rewriter + .create( + loc, ValueRange{source}, ValueRange{initTensor}, + SmallVector{axis}, + [&](OpBuilder &opBuilder, Location loc, ValueRange inputs) { + assert(inputs.size() == 2); + Value result = this->getRedElement(inputs[0], inputs[1], loc, + rop, opBuilder, false); + opBuilder.create(loc, result); + }) + .getResult(0); + + if (sourceType.getRank() == 1) { + finalResult = + rewriter.create(loc, constantType, finalResult); + } + + rewriter.replaceOp(op, finalResult); + return success(); +} + +LogicalResult ReduceConverter::convertToTargetOpExtended( + triton::ReduceOp op, typename triton::ReduceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto elemTypes = op.getElementTypes(); + + auto valueResultType = dyn_cast(op.getType(0)); + const auto isScalarReduce = valueResultType == nullptr; + + SmallVector outputs; + for (auto i = 0; i < op.getResult().size() && i < elemTypes.size(); i++) { + auto result = dyn_cast(op.getType(i)); + SmallVector resultShape{ + isScalarReduce ? SmallVector{} + : SmallVector(result.getShape())}; + outputs.push_back( + rewriter.create(loc, resultShape, elemTypes[i])); + } + + auto linalgOp = rewriter.create( + loc, adaptor.getOperands(), outputs, + SmallVector{adaptor.getAxis()}, + [&](OpBuilder &b, Location loc, ValueRange inputs) { + auto tritonReduceBlock = op.getBody(); + IRMapping mapping; + mapping.map(tritonReduceBlock->getArguments(), inputs); + + for (auto &op : tritonReduceBlock->without_terminator()) { + b.clone(op, mapping); + } + + auto tritonYield = tritonReduceBlock->getTerminator(); + auto results = + llvm::map_to_vector(tritonYield->getOperands(), + [&](Value val) { return mapping.lookup(val); }); + b.create(loc, results); + }); + + auto reduceWithIndexParams = getReduceWithIndexParams(op); + if (!reduceWithIndexParams.has_value()) { + return rewriter.notifyMatchFailure(op, "meaningless reduce operation"); + } + addReduceWithIndexAttr(*reduceWithIndexParams, rewriter, linalgOp); + + if (isScalarReduce) { + SmallVector reduceResults; + for (auto i = 0; i < linalgOp.getResults().size() && i < elemTypes.size(); + i++) { + reduceResults.push_back(rewriter.create( + loc, elemTypes[i], linalgOp.getResults()[i], ValueRange{})); + } + rewriter.replaceOp(op, reduceResults); + } else { + rewriter.replaceOp(op, linalgOp); + } + return success(); +} + +bool ScanConverter::isReductionOpSupported(Operation *redOp) const { + return isa(redOp); +} + +LogicalResult +ScanConverter::convertToTargetOp(triton::ScanOp op, + typename triton::ScanOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto reductionOps = this->getRedOps(op); + if (reductionOps.empty()) { + return rewriter.notifyMatchFailure(op, + "No reduction op found in scan body"); + } + + llvm::SmallString<64> funcName; + auto rop = reductionOps.front(); + if (this->isReductionOpSupported(reductionOps.front())) { + if (isa(rop)) { + funcName = "triton_cumsum"; + } else if (isa(rop)) { + funcName = "triton_cumprod"; + } + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + + auto loc = op.getLoc(); + auto src = adaptor.getOperands().front(); + auto resTy = op.getResult().front().getType(); + auto libFnType = rewriter.getFunctionType( + {src.getType(), rewriter.getI32Type(), rewriter.getI1Type()}, {resTy}); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + + SymbolTable symTab(moduleOp); + auto maybePrintFuncNameAttr = symTab.renameToUnique(funcOp, {&symTab}); + if (failed(maybePrintFuncNameAttr)) { + return op->emitError( + "failed to create a unique func name for device_print"); + } + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + + rewriter.setInsertionPoint(op); + auto scanAxis = op.getAxis(); + auto scanReverse = op.getReverse(); + Value axis = rewriter.create(loc, scanAxis, 32); + Value reverseVal = + rewriter.create(loc, scanReverse, 1); + auto callOp = rewriter.create( + loc, funcOp.getSymNameAttr(), TypeRange({resTy}), + ValueRange({src, axis, reverseVal})); + + rewriter.replaceOp(op, callOp); + + return success(); + } else { + // This branch is the associative_scan op. + bool reverse = op.getReverse(); + + auto loc = op.getLoc(); + + Value scanInput = op.getOperand(0); + + auto srcType = mlir::dyn_cast(scanInput.getType()); + if (!srcType) { + return rewriter.notifyMatchFailure( + op, "Expected RankedTensorType input for associative_scan"); + } + + auto elementType = srcType.getElementType(); + auto shape = srcType.getShape(); + int rank = shape.size(); + int axis = op.getAxis(); + + if (axis < 0 || axis >= rank) { + return rewriter.notifyMatchFailure(op, "Invalid scan axis: " + + std::to_string(axis)); + } + + if (op->getNumRegions() < 1 || op->getRegion(0).empty()) { + return rewriter.notifyMatchFailure(op, "Missing combine region"); + } + + OpBuilder::InsertionGuard guard(rewriter); + + auto memrefType = MemRefType::get(shape, elementType); + Value inputMemRef = + rewriter.create(loc, memrefType, scanInput); + Value outputMemRef = rewriter.create(loc, memrefType); + + auto processDimension = [&](ArrayRef baseIdxsArray) { + auto startInd = rewriter.create(op.getLoc(), 0); + if (reverse) { + startInd = rewriter.create(op.getLoc(), + shape[axis] - 1); + } + llvm::SmallVector baseIdxs(baseIdxsArray.begin(), + baseIdxsArray.end()); + llvm::SmallVector firstIdx = baseIdxs; + if (axis <= firstIdx.size()) { + firstIdx.insert(firstIdx.begin() + axis, startInd); + } else { + firstIdx.push_back(startInd); + } + + Value firstVal = + rewriter.create(loc, inputMemRef, firstIdx); + rewriter.create(loc, firstVal, outputMemRef, firstIdx); + + Value axisSize = + rewriter.create(loc, inputMemRef, axis).getResult(); + Value one = rewriter.create(loc, 1); + + Value cmp = rewriter.create(loc, arith::CmpIPredicate::sgt, + axisSize, one); + auto ifOp = rewriter.create(loc, cmp, false); + + // Create a loop only when the axis size is greater than 1. + rewriter.setInsertionPointToStart(ifOp.thenBlock()); + + auto forOp = rewriter.create(loc, one, axisSize, one); + rewriter.setInsertionPointToStart(forOp.getBody()); + + Value k = forOp.getInductionVar(); + if (reverse) { + llvm::SmallVector fixInd; + fixInd.push_back( + rewriter + .create(op.getLoc(), shape[axis] - 1) + .getResult()); + fixInd.push_back(k); + auto fixIndVal = rewriter.create(op.getLoc(), fixInd); + k = fixIndVal.getResult(); + } + llvm::SmallVector currIdx = baseIdxs; + if (axis <= currIdx.size()) { + currIdx.insert(currIdx.begin() + axis, k); + } else { + currIdx.push_back(k); + } + + Value km1 = rewriter.create(loc, k, one); + if (reverse) { + km1 = rewriter.create(loc, k, one); + } + llvm::SmallVector prevIdx = baseIdxs; + if (axis <= prevIdx.size()) { + prevIdx.insert(prevIdx.begin() + axis, km1); + } else { + prevIdx.push_back(km1); + } + + Value currentVal = + rewriter.create(loc, inputMemRef, currIdx); + Value prevResult = + rewriter.create(loc, outputMemRef, prevIdx); + + Region &combineRegion = op->getRegion(0); + Block &combineBlock = combineRegion.front(); + IRMapping mapping; + mapping.map(combineBlock.getArgument(0), prevResult); + mapping.map(combineBlock.getArgument(1), currentVal); + + for (Operation &innerOp : combineBlock.without_terminator()) { + rewriter.clone(innerOp, mapping); + } + + Operation *yieldOp = combineBlock.getTerminator(); + Value resultVal = mapping.lookup(yieldOp->getOperand(0)); + + rewriter.create(loc, resultVal, outputMemRef, currIdx); + + rewriter.setInsertionPointAfter(ifOp); + }; + + // Constructing loops for non-scanning dimensions + llvm::SmallVector nonScanDims; + for (int i = 0; i < rank; ++i) { + if (i != axis) + nonScanDims.push_back(i); + } + + createSimpleNestedLoops(rewriter, loc, outputMemRef, nonScanDims, + processDimension); + + rewriter.setInsertionPointAfter(op); + + Value outputTensor = + rewriter.create(loc, outputMemRef, true); + rewriter.replaceOp(op, outputTensor); + return success(); + } +} + +LogicalResult ScanConverter::convertToTargetOpExtended( + triton::ScanOp op, typename triton::ScanOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + bool reverse = op.getReverse(); + + // 1. Extract all input tensors (supports multiple inputs) + auto operands = op->getOperands(); + if (operands.empty()) { + return rewriter.notifyMatchFailure(op, + "No input operands for extended scan"); + } + + // 2. Validate all inputs are of RankedTensorType + llvm::SmallVector inputTensTypes; + for (auto operand : operands) { + auto tensorTy = dyn_cast(operand.getType()); + if (!tensorTy) { + return rewriter.notifyMatchFailure(op, + "All inputs must be RankedTensorType"); + } + inputTensTypes.push_back(tensorTy); + } + + // 3. Validate all input tensors have the same shape (scan operation requires + // matching input dimensions) + auto baseShape = inputTensTypes[0].getShape(); + int rank = baseShape.size(); + int axis = op.getAxis(); + if (axis < 0 || axis >= rank) { + return rewriter.notifyMatchFailure(op, "Invalid scan axis: " + + std::to_string(axis)); + } + for (size_t i = 1; i < inputTensTypes.size(); ++i) { + if (inputTensTypes[i].getShape() != baseShape) { + return rewriter.notifyMatchFailure(op, + "All inputs must have the same shape"); + } + } + + // 4. Prepare MemRefs for multiple inputs/outputs + llvm::SmallVector inputMemRefs; + llvm::SmallVector outputMemRefs; + llvm::SmallVector memRefTypes; + for (size_t i = 0; i < inputTensTypes.size(); ++i) { + auto &tensorTy = inputTensTypes[i]; + auto memRefTy = + MemRefType::get(tensorTy.getShape(), tensorTy.getElementType()); + memRefTypes.push_back(memRefTy); + // Convert input tensors to MemRefs + inputMemRefs.push_back( + rewriter.create(loc, memRefTy, operands[i])); + // Allocate MemRefs for outputs + outputMemRefs.push_back(rewriter.create(loc, memRefTy)); + } + + // 5. Define scanning logic for multiple inputs/outputs + LogicalResult loopResult = success(); + auto processDimension = [&](ArrayRef baseIdxsArray) { + llvm::SmallVector baseIdxs(baseIdxsArray.begin(), + baseIdxsArray.end()); + + auto startInd = rewriter.create(op.getLoc(), 0); + if (reverse) { + startInd = rewriter.create(op.getLoc(), + baseShape[axis] - 1); + } + + llvm::SmallVector firstIdx = baseIdxs; + if (axis <= firstIdx.size()) { + firstIdx.insert(firstIdx.begin() + axis, startInd); + } else { + firstIdx.push_back(startInd); + } + + // 5.1 Process the first element: directly copy multiple inputs to multiple + // outputs (initialize cumulative results) + for (size_t i = 0; i < inputMemRefs.size(); ++i) { + Value firstVal = + rewriter.create(loc, inputMemRefs[i], firstIdx); + rewriter.create(loc, firstVal, outputMemRefs[i], + firstIdx); + } + + Value axisSize = + rewriter.create(loc, baseShape[axis]); + Value one = rewriter.create(loc, 1); + + Value cmp = rewriter.create(loc, arith::CmpIPredicate::sgt, + axisSize, one); + auto ifOp = rewriter.create(loc, cmp, false); + + // Create a loop only when the axis size is greater than 1. + rewriter.setInsertionPointToStart(ifOp.thenBlock()); + + // Use a forward loop, but handle reverse indexing inside the loop. + auto forOp = rewriter.create(loc, one, axisSize, one); + rewriter.setInsertionPointToStart(forOp.getBody()); + + Value k = forOp.getInductionVar(); + + if (reverse) { + // Reverse scanning: Convert the forward loop index to the actual reverse + // index. (axis_size - 1) - k + Value axisSizeVal = + rewriter.create(loc, baseShape[axis]); + Value axisSizeMinusOne = + rewriter.create(loc, axisSizeVal, one); + k = rewriter.create(loc, axisSizeMinusOne, k); + } + + llvm::SmallVector currIdx = baseIdxs; + if (axis <= currIdx.size()) { + currIdx.insert(currIdx.begin() + axis, k); + } else { + currIdx.push_back(k); + } + + Value prevIndex; + if (reverse) { + prevIndex = rewriter.create(loc, k, one); + } else { + prevIndex = rewriter.create(loc, k, one); + } + + llvm::SmallVector prevIdx = baseIdxs; + if (axis <= prevIdx.size()) { + prevIdx.insert(prevIdx.begin() + axis, prevIndex); + } else { + prevIdx.push_back(prevIndex); + } + + // 5.4 Load current elements and previous cumulative results + llvm::SmallVector currentVals; + llvm::SmallVector prevResults; + for (size_t i = 0; i < inputMemRefs.size(); ++i) { + currentVals.push_back( + rewriter.create(loc, inputMemRefs[i], currIdx)); + prevResults.push_back( + rewriter.create(loc, outputMemRefs[i], prevIdx)); + } + + // 5.5 Bind parameters for custom reduction logic + Region &combineRegion = op->getRegion(0); + if (combineRegion.empty()) { + op->emitError("Missing combine region in extended scan"); + loopResult = failure(); + return; + } + Block &combineBlock = combineRegion.front(); + // Validate that the number of reduction region arguments matches (number of + // previous results + number of current elements) + if (combineBlock.getNumArguments() != 2 * inputMemRefs.size()) { + op->emitError("Combine region arguments mismatch with input count"); + loopResult = failure(); + return; + } + IRMapping mapping; + for (size_t i = 0; i < inputMemRefs.size(); ++i) { + // Bind previous results (previous value of the i-th output) to the i-th + // argument of the reduction region + mapping.map(combineBlock.getArgument(i), prevResults[i]); + // Bind current elements (current value of the i-th input) to the i+N-th + // argument of the reduction region (N is the number of inputs) + mapping.map(combineBlock.getArgument(i + inputMemRefs.size()), + currentVals[i]); + } + + // 5.6 Clone all operations within the reduction region + for (Operation &innerOp : combineBlock.without_terminator()) { + rewriter.clone(innerOp, mapping); + } + + // 5.7 Extract reduction results and store them in outputMemRef + Operation *yieldOp = combineBlock.getTerminator(); + if (yieldOp->getNumOperands() != outputMemRefs.size()) { + op->emitError("Combine region returns mismatch with output count"); + loopResult = failure(); + return; + } + for (size_t i = 0; i < outputMemRefs.size(); ++i) { + Value resultVal = mapping.lookup(yieldOp->getOperand(i)); + rewriter.create(loc, resultVal, outputMemRefs[i], + currIdx); + } + + rewriter.setInsertionPointAfter(ifOp); + }; + + // 6. Generate nested loops for non-scan dimensions + llvm::SmallVector nonScanDims; + for (int i = 0; i < rank; ++i) { + if (i != axis) + nonScanDims.push_back(i); + } + createSimpleNestedLoops(rewriter, loc, outputMemRefs[0], nonScanDims, + processDimension); + + if (failed(loopResult)) { + return failure(); + } + + // 7. Convert multiple output MemRefs back to tensors and replace the original + // tt.scan operation + llvm::SmallVector outputTensors; + for (auto outputMemRef : outputMemRefs) { + outputTensors.push_back( + rewriter.create(loc, outputMemRef, true)); + } + rewriter.replaceOp(op, outputTensors); + + return success(); +} + +LogicalResult ExternElementwiseClOpConverter::matchAndRewrite( + triton::ExternElementwiseOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + if (!op.getPure()) { + op->emitWarning() << "impure elementwise op!"; + return failure(); + } + if (op.getSymbol().contains("__hmf_")) { + // 1. get or create the declaration of external elementwise function + Type dstTy = op.getResult().getType(); + bool isDstScalar = !isa(dstTy); + Type dstElemTy = + isDstScalar ? dstTy : cast(dstTy).getElementType(); + SmallVector srcElemTys; + SmallVector srcs; + for (auto src : op.getSrcs()) { + if (!isa(src.getType())) { + src = rewriter.create( + op.getLoc(), RankedTensorType::get({(int64_t)1}, src.getType()), + src); + } + srcs.push_back(src); + srcElemTys.push_back( + cast(src.getType()).getElementType()); + } + FunctionType elemFuncType = + FunctionType::get(rewriter.getContext(), srcElemTys, {dstElemTy}); + auto mod = SymbolTable::getNearestSymbolTable(op); + auto extFunc = dyn_cast_or_null( + SymbolTable::lookupSymbolIn(mod, op.getSymbol())); + if (!extFunc) { + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPointToStart(&mod->getRegion(0).front()); + extFunc = rewriter.create(rewriter.getUnknownLoc(), + op.getSymbol(), elemFuncType); + extFunc.setPrivate(); + extFunc->setAttr(LLVM::LLVMDialect::getReadnoneAttrName(), + UnitAttr::get(rewriter.getContext())); + } + assert(isa( + SymbolTable::lookupSymbolIn(mod, op.getSymbol()))); + // 2. prepare the output tensor + Value output; + if (isDstScalar) { + dstTy = RankedTensorType::get({(int64_t)1}, dstElemTy); + } + bool found = false; + for (Value v : srcs) { + if (v.getType() == dstTy) { + found = true; + output = v; + break; + } + } + if (!found) { + output = rewriter.create( + op.getLoc(), cast(dstTy).getShape(), dstElemTy); + } + // 3. create the linalg.map op + auto mapOp = rewriter.create( + loc, + /*inputs=*/srcs, + /*init=*/output, + /*bodyBuilder=*/ + [&](OpBuilder &builder, Location loc, ValueRange regionArgs) { + auto elemOp = builder.create(loc, + /*name=*/op.getSymbol(), + /*resultType=*/dstElemTy, + /*operands=*/regionArgs); + builder.create(loc, elemOp->getResults()); + }); + if (isDstScalar) { + // need to convert tensor back to scalar + auto indexType = rewriter.getIndexType(); + Value zeroConstant = rewriter.create( + loc, indexType, rewriter.getIntegerAttr(indexType, 0)); + auto extractOp = rewriter.create( + loc, mapOp.getResults()[0], zeroConstant); + rewriter.replaceOp(op, extractOp); + } else { + rewriter.replaceOp(op, mapOp); + } + return success(); + } + return failure(); +} + +LogicalResult UnrealizedCastConverter::matchAndRewrite( + UnrealizedConversionCastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + rewriter.eraseOp(op); + return success(); +} + +LogicalResult +JoinConverter::matchAndRewrite(triton::JoinOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Value opa = op.getLhs(); + Value opb = op.getRhs(); + auto loc = op.getLoc(); + + auto resType = dyn_cast(op.getResult().getType()); + Value emptyOp = rewriter.create(loc, resType.getShape(), + resType.getElementType()); + + auto shape = dyn_cast(opa.getType()).getShape(); + auto sizes = llvm::map_to_vector(shape, [&](int64_t t) { + return OpFoldResult(rewriter.getI64IntegerAttr(t)); + }); + sizes.push_back(rewriter.getI64IntegerAttr(1)); + + int64_t rank = resType.getRank(); + + // Set last dimension stride to 2 in layout + // As last dimension size is always 1, last dimension stride here could be + // either 1 or 2, while stride `2` could carry interleave trait and it's + // convenient for next lower. + SmallVector strides(rank, rewriter.getIndexAttr(1)); + strides.back() = rewriter.getIndexAttr(2); + + SmallVector offsets(rank, rewriter.getIndexAttr(0)); + + auto insert0 = rewriter.create( + loc, opa, emptyOp, offsets, sizes, strides); + + offsets.back() = rewriter.getIndexAttr(1); + auto insert1 = rewriter.create( + loc, opb, insert0, offsets, sizes, strides); + rewriter.replaceOp(op, insert1); + return success(); +} + +LogicalResult +CatConverter::matchAndRewrite(triton::CatOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Value opa = op.getLhs(); + Value opb = op.getRhs(); + auto loc = op.getLoc(); + + auto resType = dyn_cast(op.getResult().getType()); + if (!resType || resType.getRank() != 1) { + return rewriter.notifyMatchFailure(op, "only support 1D cat"); + } + + auto inputTypeA = dyn_cast(opa.getType()); + auto inputTypeB = dyn_cast(opb.getType()); + if (!inputTypeA || !inputTypeB || inputTypeA.getRank() != 1 || + inputTypeB.getRank() != 1) { + return rewriter.notifyMatchFailure(op, "inputs must be 1D tensors"); + } + + int64_t sizeA = inputTypeA.getShape()[0]; + int64_t sizeB = inputTypeB.getShape()[0]; + + // Only handle the case where both inputs have size 1 (i.e., scalar-like) + if (sizeA == 1 && sizeB == 1) { + // Use scalar extract + insert + auto emptyOp = rewriter.create(loc, resType.getShape(), + resType.getElementType()); + + Value zero = rewriter.create(loc, 0); + Value one = rewriter.create(loc, 1); + + Value scalarA = rewriter.create(loc, opa, zero); + Value scalarB = rewriter.create(loc, opb, zero); + + Value inserted0 = + rewriter.create(loc, scalarA, emptyOp, zero); + Value inserted1 = + rewriter.create(loc, scalarB, inserted0, one); + + rewriter.replaceOp(op, inserted1); + return success(); + } + + // General case: use tensor.insert_slice + auto emptyOp = rewriter.create(loc, resType.getShape(), + resType.getElementType()); + + auto rank = resType.getRank(); + SmallVector offsets(rank, rewriter.getIndexAttr(0)); + SmallVector strides(rank, rewriter.getIndexAttr(1)); + + auto inputType = dyn_cast(opa.getType()); + + SmallVector sizes = + llvm::map_to_vector(inputType.getShape(), [&](int64_t t) { + return OpFoldResult(rewriter.getI64IntegerAttr(t)); + }); + + auto insert0 = rewriter.create( + loc, opa, emptyOp, offsets, sizes, strides); + + offsets[0] = + rewriter.getIndexAttr(inputType.getRank() ? inputType.getShape()[0] : 1); + auto insert1 = rewriter.create( + loc, opb, insert0, offsets, sizes, strides); + + rewriter.replaceOp(op, insert1); + return success(); +} + +/// @brief Convert tt.gather to func.call. BiShengIR captures the func +/// with assumed semantics. +/// @param op The `triton::GatherOp` operation to be rewritten. +/// @param adaptor An adaptor for the operation's operands. +/// @param rewriter A pattern rewriter used to modify the IR. +/// @return A `LogicalResult` indicating whether the rewrite was successful. +LogicalResult +GatherConverter::matchAndRewrite(triton::GatherOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + Value src = adaptor.getSrc(); + Value idx = adaptor.getIndices(); + Value res = op.getResult(); + auto gatherAxis = op.getAxis(); + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + + llvm::SmallString funcName = gatherFuncNameBase; + int uniqueId = 0; + while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) { + funcName = gatherFuncNameBase; + funcName += ("_" + std::to_string(uniqueId++)); + } + + auto resTy = res.getType(); + auto libFnType = rewriter.getFunctionType( + {src.getType(), idx.getType(), rewriter.getI32Type()}, {resTy}); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + + rewriter.setInsertionPoint(op); + Value axis = rewriter.create(loc, gatherAxis, 32); + auto callOp = rewriter.create(loc, funcOp.getSymNameAttr(), + TypeRange({resTy}), + ValueRange({src, idx, axis})); + + rewriter.replaceOp(op, callOp); + + return success(); +} + +LogicalResult +SplitConverter::matchAndRewrite(triton::SplitOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Value input = op.getSrc(); + auto loc = op.getLoc(); + auto inputType = cast(input.getType()); + + int64_t rank = inputType.getRank(); + SmallVector offsets(rank, rewriter.getIndexAttr(0)); + // Similar to JoinConverter, here adjust last dimension stride + SmallVector strides(rank, rewriter.getIndexAttr(1)); + strides.back() = rewriter.getIndexAttr(2); + + auto outType = dyn_cast(op.getOutLHS().getType()); + auto sizes = llvm::map_to_vector(outType.getShape(), [&](int64_t t) { + return OpFoldResult(rewriter.getIndexAttr(t)); + }); + sizes.push_back(rewriter.getIndexAttr(1)); + + auto slice0 = rewriter.create( + loc, outType, input, offsets, sizes, strides); + + offsets.back() = rewriter.getIndexAttr(1); + auto slice1 = rewriter.create( + loc, outType, input, offsets, sizes, strides); + + SmallVector slices = {slice0.getResult(), slice1.getResult()}; + rewriter.replaceOp(op, ValueRange(slices)); + return success(); +} + +/* +the element-wise most significant N bits of the 2N-bit product of x and y +%x:2 = arith.mulsi_extended %y, %z : tensor<4x?xi32> +*/ +LogicalResult TritonMulhiuiConverter::matchAndRewrite( + triton::MulhiUIOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + Value opl = op.getX(); + Value opr = op.getY(); + Value res = op.getResult(); + auto newMulOp = rewriter.create( + loc, res.getType(), res.getType(), opl, opr); + // triton only need the high value + rewriter.replaceOp(op, ValueRange{newMulOp.getHigh()}); + return success(); +} + +LogicalResult TritonPreciseSqrtConverter::matchAndRewrite( + triton::PreciseSqrtOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + rewriter.replaceOpWithNewOp(op, adaptor.getOperands()); + return success(); +} + +LogicalResult DevicePrintConverter::matchAndRewrite( + triton::PrintOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + SmallVector inputTypes; + for (auto arg : op.getArgs()) { + inputTypes.push_back(arg.getType()); + } + auto libFnType = rewriter.getFunctionType(inputTypes, {}); + auto funcOp = + rewriter.create(op.getLoc(), printFuncNameBase, libFnType); + SymbolTable symTab(moduleOp); + auto maybePrintFuncNameAttr = symTab.renameToUnique(funcOp, {&symTab}); + if (failed(maybePrintFuncNameAttr)) { + return op->emitError( + "failed to create a unique func name for device_print"); + } + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + auto prefixAttr = op.getPrefixAttr(); + funcOp->setAttr(prefixAttrName, prefixAttr); + auto hexAttr = op.getHexAttr(); + funcOp->setAttr(hexAttrName, hexAttr); + + rewriter.setInsertionPoint(op); + rewriter.create(op.getLoc(), funcOp, op.getArgs()); + + rewriter.eraseOp(op); + return success(); +} + +LogicalResult DeviceAssertConverter::matchAndRewrite( + triton::AssertOp op, OpAdaptor adaptor, + mlir::ConversionPatternRewriter &rewriter) const { + auto msgAttr = op.getMessageAttr(); + // Filter out automatically inserted assert ops + if (auto strAttr = mlir::dyn_cast(msgAttr)) { + llvm::StringRef msg = strAttr.getValue(); + if (msg.contains("overflow detected for operation")) { + rewriter.eraseOp(op); + return success(); + } + } + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + auto conditionType = op.getCondition().getType(); + + auto libFnType = rewriter.getFunctionType({conditionType}, {}); + auto funcOp = + rewriter.create(op.getLoc(), printFuncNameBase, libFnType); + mlir::SymbolTable symTab(moduleOp); + auto maybePrintFuncNameAttr = symTab.renameToUnique(funcOp, {&symTab}); + if (failed(maybePrintFuncNameAttr)) { + return op->emitError( + "failed to create a unique func name for device_assert"); + } + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + funcOp->setAttr(msgAttrName, msgAttr); + + rewriter.setInsertionPoint(op); + rewriter.create(op.getLoc(), funcOp, + ValueRange{op.getCondition()}); + + rewriter.eraseOp(op); + return success(); +} + +LogicalResult +MatmulConverter::matchAndRewrite(triton::DotOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto opa = adaptor.getA(); + auto opb = adaptor.getB(); + auto opc = adaptor.getC(); + auto dstType = cast(op.getType()); + auto inputPrec = op.getInputPrecision(); + + if (dstType.getRank() == 2) { + auto matmulOp = rewriter.replaceOpWithNewOp( + op, ValueRange{opa, opb}, ValueRange{opc}); + matmulOp->setAttr( + "input_precison", + rewriter.getStringAttr(stringifyInputPrecision(inputPrec))); + } else if (dstType.getRank() == 3) { + auto matmulOp = rewriter.replaceOpWithNewOp( + op, ValueRange{opa, opb}, ValueRange{opc}); + matmulOp->setAttr( + "input_precison", + rewriter.getStringAttr(stringifyInputPrecision(inputPrec))); + } else { + llvm_unreachable("Datatype of DotOp operands could only be 2D or 3D"); + } + return success(); +} + +LogicalResult +FlipOpConverter::matchAndRewrite(triton::ascend::FlipOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Value src = adaptor.getSrc(); + auto rankedSrcTy = cast(src.getType()); + + MLIRContext *ctx = rewriter.getContext(); + + Type valuesTy = src.getType(); + Location loc = op.getLoc(); + + auto dimAttr = op->getAttrOfType("dim"); + if (!dimAttr) { + op->emitError("missing 'dim' attribute"); + return failure(); + } + + auto moduleOp = op->getParentOfType(); + if (!moduleOp) { + op->emitError("must be inside a module"); + return failure(); + } + + // Unique callee name: triton_flip, triton_flip_1, … + std::string funcName = baseFuncName.str(); + int uniqueId = 0; + while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) + funcName = (baseFuncName + Twine("_") + Twine(uniqueId++)).str(); + + auto i64Ty = IntegerType::get(ctx, 64); + auto libFnType = + rewriter.getFunctionType({rankedSrcTy, i64Ty}, {rankedSrcTy}); + + // Declare the callee + auto moduleIP = rewriter.saveInsertionPoint(); + rewriter.setInsertionPointToEnd(moduleOp.getBody()); + auto funcOp = rewriter.create(loc, funcName, libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + rewriter.restoreInsertionPoint(moduleIP); + + // dim constant + Value dimVal = + rewriter.create(loc, dimAttr.getInt(), 64); + + // Call the backend function + auto callee = SymbolRefAttr::get(ctx, funcOp.getSymName()); + auto callOp = rewriter.create( + loc, TypeRange({rankedSrcTy}), callee, ValueRange({src, dimVal})); + + Value finalValues = callOp.getResult(0); + + rewriter.replaceOp(op, {finalValues}); + return success(); +} + +LogicalResult +SortOpConverter::matchAndRewrite(triton::ascend::SortOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Value src = adaptor.getSrc(); + auto rankedSrcTy = cast(src.getType()); + auto srcElemTy = rankedSrcTy.getElementType(); + auto srcShape = rankedSrcTy.getShape(); + auto srcEnc = rankedSrcTy.getEncoding(); + + MLIRContext *ctx = rewriter.getContext(); + + Type backendElemTy = srcElemTy; + if (srcElemTy.isInteger(8)) { + backendElemTy = Float16Type::get(ctx); // i8 -> f16 + } else if (srcElemTy.isInteger(16)) { + backendElemTy = Float32Type::get(ctx); // i16 -> f32 + } + Type backendTensorTy = RankedTensorType::get(srcShape, backendElemTy, srcEnc); + + Type valuesTy = src.getType(); + + Location loc = op.getLoc(); + auto dimAttr = op->getAttrOfType("dim"); + auto descAttr = op->getAttrOfType("descending"); + if (!dimAttr || !descAttr) { + op->emitError("missing 'dim' or 'descending' attribute"); + return failure(); + } + + auto moduleOp = op->getParentOfType(); + if (!moduleOp) { + op->emitError("must be inside a module"); + return failure(); + } + + llvm::SmallString<64> baseName("triton_sort"); + llvm::SmallString<64> funcName = baseName; + int uniqueId = 0; + while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) { + funcName = baseName; + funcName += ("_" + std::to_string(uniqueId++)); + } + + auto i64Ty = IntegerType::get(ctx, 64); + auto i1Ty = IntegerType::get(ctx, 1); + auto libFnType = rewriter.getFunctionType({backendTensorTy, i64Ty, i1Ty}, + {backendTensorTy}); + + auto moduleIP = rewriter.saveInsertionPoint(); + rewriter.setInsertionPointToEnd(moduleOp.getBody()); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + rewriter.restoreInsertionPoint(moduleIP); + + Value srcForCall = src; + if (backendElemTy != srcElemTy) { + srcForCall = rewriter.create(loc, backendTensorTy, src); + } + + Value dimVal = + rewriter.create(loc, dimAttr.getInt(), 64); + Value descVal = rewriter.create( + loc, descAttr.getValue() ? 1 : 0, 1); + + auto callee = SymbolRefAttr::get(ctx, funcOp.getSymName()); + auto callOp = + rewriter.create(loc, TypeRange({backendTensorTy}), callee, + ValueRange({srcForCall, dimVal, descVal})); + + Value valuesFloat = callOp.getResult(0); // tensor + + Value finalValues = valuesFloat; + if (backendElemTy != srcElemTy) { + finalValues = rewriter.create(loc, valuesTy, valuesFloat); + } + + rewriter.replaceOp(op, {finalValues}); + + return success(); +} + +LogicalResult +DotScaledConverter::matchAndRewrite(triton::DotScaledOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + Value lhs = adaptor.getLhs(); + Value lhsScale = adaptor.getLhsScale(); + Value rhsScale = adaptor.getRhsScale(); + Value rhs = adaptor.getRhs(); + Value c = adaptor.getC(); + RankedTensorType dstType = cast(op.getType()); + + RankedTensorType lhsTy = cast(lhs.getType()); + RankedTensorType lhsScaleTy = cast(lhsScale.getType()); + RankedTensorType rhsScaleTy = + rhsScale ? cast(rhsScale.getType()) : nullptr; + RankedTensorType rhsTy = cast(rhs.getType()); + + Value lhsScaleOut; + Value rhsScaleOut; + Value c127 = rewriter.create( + op.getLoc(), rewriter.getI16Type(), rewriter.getI16IntegerAttr(127)); + Value c7 = rewriter.create( + op.getLoc(), rewriter.getI16Type(), rewriter.getI16IntegerAttr(7)); + Type i16Ty = rewriter.getI16Type(); + Type bf16Ty = rewriter.getBF16Type(); + Type fp16Ty = rewriter.getF16Type(); + Type fp32Ty = rewriter.getF32Type(); + + if (lhsScaleTy.getElementType().isIntOrIndex()) { + RankedTensorType lhsScaleI16Ty = + RankedTensorType::get(lhsScaleTy.getShape(), i16Ty); + Value lhsScaleI16 = + rewriter.create(op.getLoc(), lhsScaleI16Ty, lhsScale); + + Value lhsShift127Empty = rewriter.create( + op.getLoc(), lhsScaleI16Ty.getShape(), i16Ty); + Value lhsShift127 = + rewriter + .create(op.getLoc(), ValueRange{c127}, + ValueRange{lhsShift127Empty}) + .getResult(0); + + Value lhsScaleI16Add127 = + rewriter.create(op.getLoc(), lhsScaleI16, lhsShift127); + + Value lhsShift7Empty = rewriter.create( + op.getLoc(), lhsScaleI16Ty.getShape(), i16Ty); + Value lhsShift7 = rewriter + .create(op.getLoc(), ValueRange{c7}, + ValueRange{lhsShift7Empty}) + .getResult(0); + Value lhsScaleI16Shifted = rewriter.create( + op.getLoc(), lhsScaleI16Add127, lhsShift7); + + RankedTensorType lhsScaleBF16Ty = + RankedTensorType::get(lhsScaleTy.getShape(), bf16Ty); + Value lhsScaleBF16 = rewriter.create( + op.getLoc(), lhsScaleBF16Ty, lhsScaleI16Shifted); + if (lhsTy.getElementType() == fp16Ty) { + RankedTensorType lhsScaleFp32Ty = + RankedTensorType::get(lhsScaleTy.getShape(), fp32Ty); + Value lhsScaleFp32 = rewriter.create( + op.getLoc(), lhsScaleFp32Ty, lhsScaleBF16); + RankedTensorType lhsScaleFp16Ty = + RankedTensorType::get(lhsScaleTy.getShape(), fp16Ty); + lhsScaleOut = rewriter.create( + op.getLoc(), lhsScaleFp16Ty, lhsScaleFp32); + } else { + lhsScaleOut = lhsScaleBF16; + } + } else { + lhsScaleOut = + rewriter + .create( + op.getLoc(), + RankedTensorType::get(lhsScaleTy.getShape(), fp32Ty), lhsScale) + .getResult(); + } + + if (rhsScale && rhsScaleTy.getElementType().isIntOrIndex()) { + if (rhsScaleTy.getRank() != 2) { + return op.emitError("rhsScale must be 2D for transpose"); + } + + SmallVector transposedShape = {rhsScaleTy.getShape()[1], + rhsScaleTy.getShape()[0]}; + RankedTensorType transposedRhsScaleTy = + RankedTensorType::get(transposedShape, rhsScaleTy.getElementType()); + + Value transposedRhsScale = rewriter.create( + op.getLoc(), transposedRhsScaleTy, rhsScale, + DenseI32ArrayAttr::get(rewriter.getContext(), ArrayRef{1, 0})); + RankedTensorType rhsScaleI16Ty = + RankedTensorType::get(transposedShape, i16Ty); + Value rhsScaleI16 = rewriter.create( + op.getLoc(), rhsScaleI16Ty, transposedRhsScale); + Value rhsShift127Empty = rewriter.create( + op.getLoc(), rhsScaleI16Ty.getShape(), i16Ty); + Value rhsShift127 = + rewriter + .create(op.getLoc(), ValueRange{c127}, + ValueRange{rhsShift127Empty}) + .getResult(0); + + Value rhsScaleI16Add127 = + rewriter.create(op.getLoc(), rhsScaleI16, rhsShift127); + Value rhsShift7Empty = rewriter.create( + op.getLoc(), rhsScaleI16Ty.getShape(), i16Ty); + Value rhsShift7 = rewriter + .create(op.getLoc(), ValueRange{c7}, + ValueRange{rhsShift7Empty}) + .getResult(0); + Value rhsScaleI16Shifted = rewriter.create( + op.getLoc(), rhsScaleI16Add127, rhsShift7); + + RankedTensorType rhsScaleBF16Ty = + RankedTensorType::get(transposedShape, bf16Ty); + Value rhsScaleBF16 = rewriter.create( + op.getLoc(), rhsScaleBF16Ty, rhsScaleI16Shifted); + + if (rhsTy.getElementType() == fp16Ty) { + RankedTensorType rhsScaleFp32Ty = + RankedTensorType::get(transposedShape, fp32Ty); + Value rhsScaleFp32 = rewriter.create( + op.getLoc(), rhsScaleFp32Ty, rhsScaleBF16); + RankedTensorType rhsScaleFp16Ty = + RankedTensorType::get(transposedShape, fp16Ty); + rhsScaleOut = rewriter.create( + op.getLoc(), rhsScaleFp16Ty, rhsScaleFp32); + } else { + rhsScaleOut = rhsScaleBF16; + } + int64_t rhsD0 = rhsScaleTy.getShape()[1]; + int64_t rhsD1 = rhsScaleTy.getShape()[0]; + SmallVector rhsExpandedShape1 = {rhsD0, rhsD1, 1}; + RankedTensorType rhsExpandedTy1 = + RankedTensorType::get(rhsExpandedShape1, rhsTy.getElementType()); + Value rhsExpanded1 = rewriter + .create( + op.getLoc(), rhsExpandedTy1, rhsScaleOut, + rewriter.getI32IntegerAttr(2)) + .getResult(); + + int64_t rhsDim1 = rhsTy.getShape()[0]; + if (rhsDim1 % rhsD0 != 0) { + return op.emitError( + "rhs dim0 must be an integer multiple of rhsScale dim0"); + } + int64_t rhsD2 = rhsDim1 / rhsD0; + SmallVector rhsBroadcastShape = {rhsD0, rhsD1, rhsD2}; + RankedTensorType rhsBroadcastTy = + RankedTensorType::get(rhsBroadcastShape, rhsTy.getElementType()); + Value rhsBroadcasted = rewriter + .create( + op.getLoc(), rhsBroadcastTy, rhsExpanded1) + .getResult(); + + SmallVector transposeOrder = {0, 2, 1}; + Value transposedBroadcasted = rewriter.create( + op.getLoc(), + RankedTensorType::get({rhsD0, rhsD2, rhsD1}, rhsTy.getElementType()), + rhsBroadcasted, + DenseI32ArrayAttr::get(rewriter.getContext(), transposeOrder)); + SmallVector rhsReassociation; + rhsReassociation.push_back({0, 1}); + rhsReassociation.push_back({2}); + + Value scaledRhs = rewriter + .create( + op.getLoc(), + RankedTensorType::get({rhsD0 * rhsD2, rhsD1}, + rhsTy.getElementType()), + transposedBroadcasted, rhsReassociation) + .getResult(); + + rhs = + rewriter.create(op.getLoc(), rhs, scaledRhs).getResult(); + } + + int64_t D0 = lhsScaleTy.getShape()[0]; + int64_t D1 = lhsScaleTy.getShape()[1]; + SmallVector expandedShape1 = {D0, D1, 1}; + RankedTensorType expandedTy1 = + RankedTensorType::get(expandedShape1, lhsTy.getElementType()); + Value expanded1 = + rewriter + .create(op.getLoc(), expandedTy1, lhsScaleOut, + rewriter.getI32IntegerAttr(2)) + .getResult(); + + int64_t lhsDim1 = lhsTy.getShape()[1]; + if (lhsDim1 % D1 != 0) { + return op.emitError( + "lhs dim1 must be an integer multiple of lhsScale dim1"); + } + int64_t D2 = lhsDim1 / D1; + SmallVector broadcastShape = {D0, D1, D2}; + RankedTensorType broadcastTy = + RankedTensorType::get(broadcastShape, lhsTy.getElementType()); + Value broadcasted = + rewriter.create(op.getLoc(), broadcastTy, expanded1) + .getResult(); + + SmallVector reassociation; + reassociation.push_back({0}); + reassociation.push_back({1, 2}); + + Value scaledLhs = + rewriter + .create( + op.getLoc(), + RankedTensorType::get({D0, D1 * D2}, lhsTy.getElementType()), + broadcasted, reassociation) + .getResult(); + + Value scaledLhsFinal = + rewriter.create(op.getLoc(), lhs, scaledLhs).getResult(); + + Operation *matmulOp; + if (dstType.getRank() == 2) { + matmulOp = rewriter.create( + op.getLoc(), ValueRange{scaledLhsFinal, rhs}, ValueRange{c}); + } else if (dstType.getRank() == 3) { + matmulOp = rewriter.create( + op.getLoc(), ValueRange{scaledLhsFinal, rhs}, ValueRange{c}); + } else { + return op.emitError("DotScaledOp only support 2D or 3D tensor"); + } + + rewriter.replaceOp(op, matmulOp->getResults()); + return success(); +} + +LogicalResult +PtrToIntConverter::matchAndRewrite(triton::PtrToIntOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + Value ptr = adaptor.getSrc(); + + if (!mlir::isa(ptr.getType())) { + return rewriter.notifyMatchFailure(op, "input is not a memref type"); + } + + auto resultType = op.getType(); + + // memref.extract_aligned_pointer_as_index is used to obtain the integer + // representation of the base address. + auto ptrToIndexOp = + rewriter.create(loc, ptr); + + Value intResult = + rewriter.create(loc, resultType, ptrToIndexOp); + + rewriter.replaceOp(op, intResult); + return success(); +} + +LogicalResult EmbeddingGatherConverter::matchAndRewrite( + triton::ascend::EmbeddingGatherOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + + auto funcName = generateUniqueFuncName(moduleOp, funcNameBase); + + auto src = adaptor.getSrc(); + auto idx = op.getIdx(); + auto bound = op.getBound(); + auto blksiz = op.getBlocksize(); + auto offsets = op.getOffsets(); + auto numels = op.getNumels(); + auto res = op.getResult(); + auto resTy = res.getType(); + + // convert !tt.ptr to memref + auto srcTy = dyn_cast(src.getType()); + if (!srcTy) { + return rewriter.notifyMatchFailure(op, "expected MemRefType for src"); + } + SmallVector inputTypes( + {srcTy, idx.getType(), bound.getType(), blksiz.getType()}); + inputTypes.append(offsets.getTypes().begin(), offsets.getTypes().end()); + inputTypes.append(numels.getTypes().begin(), numels.getTypes().end()); + auto libFnType = rewriter.getFunctionType(inputTypes, {resTy}); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + + rewriter.setInsertionPoint(op); + SmallVector inputVals({src, idx, bound, blksiz}); + inputVals.append(offsets.begin(), offsets.end()); + inputVals.append(numels.begin(), numels.end()); + auto callOp = rewriter.create(loc, funcOp.getSymNameAttr(), + TypeRange({resTy}), inputVals); + + rewriter.replaceOp(op, callOp); + return success(); +} + +LogicalResult +IndexPutConverter::matchAndRewrite(triton::ascend::IndexPutOp op, + OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + + auto funcName = generateUniqueFuncName(moduleOp, funcNameBase); + + auto ptr = adaptor.getPtr(); + auto index = op.getIndex(); + auto value = op.getValue(); + auto dim = op.getDim(); + auto indexBoundary = op.getIndexBoundary(); + auto endOffset = op.getEndOffset(); + auto startOffset = op.getStartOffset(); + auto dstStride = adaptor.getDstStride(); + + // convert !tt.ptr to memref + auto ptrTy = dyn_cast(ptr.getType()); + if (!ptrTy) { + return rewriter.notifyMatchFailure(op, "expected MemRefType for ptr"); + } + SmallVector inputTypes({ptrTy, index.getType(), value.getType(), + dim.getType(), indexBoundary.getType()}); + inputTypes.append(endOffset.getTypes().begin(), endOffset.getTypes().end()); + inputTypes.append(startOffset.getTypes().begin(), + startOffset.getTypes().end()); + inputTypes.append(dstStride.getTypes().begin(), dstStride.getTypes().end()); + auto libFnType = rewriter.getFunctionType(inputTypes, {}); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + + rewriter.setInsertionPoint(op); + SmallVector inputVals({ptr, index, value, dim, indexBoundary}); + inputVals.append(endOffset.begin(), endOffset.end()); + inputVals.append(startOffset.begin(), startOffset.end()); + inputVals.append(dstStride.begin(), dstStride.end()); + rewriter.create(loc, funcOp.getSymNameAttr(), TypeRange({}), + inputVals); + rewriter.eraseOp(op); + return success(); +} + +LogicalResult GatherOutToUbConverter::matchAndRewrite( + triton::ascend::GatherOutToUbOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + + auto funcName = generateUniqueFuncName(moduleOp, funcNameBase); + + auto src = adaptor.getSrc(); + auto index = op.getIndex(); + auto indexBoundary = op.getIndexBoundary(); + auto dim = op.getDim(); + auto srcStride = op.getSrcStride(); + auto endOffset = op.getEndOffset(); + auto startOffset = op.getStartOffset(); + auto other = op.getOther(); + + auto res = op.getResult(); + auto resTy = res.getType(); + + // convert !tt.ptr to memref + auto srcTy = dyn_cast(src.getType()); + if (!srcTy) { + return rewriter.notifyMatchFailure(op, "expected MemRefType for src"); + } + + SmallVector inputTypes( + {srcTy, index.getType(), indexBoundary.getType(), dim.getType()}); + inputTypes.append(srcStride.getTypes().begin(), srcStride.getTypes().end()); + inputTypes.append(endOffset.getTypes().begin(), endOffset.getTypes().end()); + inputTypes.append(startOffset.getTypes().begin(), + startOffset.getTypes().end()); + if (other) + inputTypes.push_back(other.getType()); + + auto libFnType = rewriter.getFunctionType(inputTypes, {resTy}); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + + rewriter.setInsertionPoint(op); + SmallVector inputVals({src, index, indexBoundary, dim}); + inputVals.append(srcStride.begin(), srcStride.end()); + inputVals.append(endOffset.begin(), endOffset.end()); + inputVals.append(startOffset.begin(), startOffset.end()); + if (other) + inputVals.push_back(other); + auto callOp = rewriter.create(loc, funcOp.getSymNameAttr(), + TypeRange({resTy}), inputVals); + rewriter.replaceOp(op, callOp); + return success(); +} + +LogicalResult ScatterUbToOutConverter::matchAndRewrite( + triton::ascend::ScatterUbToOutOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + + auto funcName = generateUniqueFuncName(moduleOp, funcNameBase); + + auto ptr = adaptor.getPtr(); + auto value = op.getValue(); + auto index = op.getIndex(); + auto indexBoundary = op.getIndexBoundary(); + auto dim = op.getDim(); + auto dstStride = op.getDstStride(); + auto endOffset = op.getEndOffset(); + auto startOffset = op.getStartOffset(); + + // convert !tt.ptr to memref + auto ptrTy = dyn_cast(ptr.getType()); + if (!ptrTy) { + return rewriter.notifyMatchFailure(op, "expected MemRefType for ptr"); + } + + SmallVector inputTypes({ptrTy, value.getType(), index.getType(), + indexBoundary.getType(), dim.getType()}); + inputTypes.append(dstStride.getTypes().begin(), dstStride.getTypes().end()); + inputTypes.append(endOffset.getTypes().begin(), endOffset.getTypes().end()); + inputTypes.append(startOffset.getTypes().begin(), + startOffset.getTypes().end()); + + auto libFnType = rewriter.getFunctionType(inputTypes, {}); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + + rewriter.setInsertionPoint(op); + SmallVector inputVals({ptr, value, index, indexBoundary, dim}); + inputVals.append(dstStride.begin(), dstStride.end()); + inputVals.append(endOffset.begin(), endOffset.end()); + inputVals.append(startOffset.begin(), startOffset.end()); + rewriter.create(loc, funcOp.getSymNameAttr(), TypeRange({}), + inputVals); + rewriter.eraseOp(op); + return success(); +} + +LogicalResult IndirectLoadConverter::matchAndRewrite( + triton::ascend::IndirectLoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + + auto funcName = generateUniqueFuncName(moduleOp, funcNameBase); + + auto src = adaptor.getSrc(); + auto offsets = op.getOffsets(); + auto mask = op.getMask(); + auto other = op.getOther(); + auto res = op.getResult(); + auto resTy = res.getType(); + + // convert !tt.ptr to memref + auto srcTy = dyn_cast(src.getType()); + if (!srcTy) { + return rewriter.notifyMatchFailure(op, "expected MemRefType for src"); + } + SmallVector inputTypes({srcTy, offsets.getType()}); + if (mask) + inputTypes.push_back(mask.getType()); + if (other) + inputTypes.push_back(other.getType()); + auto libFnType = rewriter.getFunctionType(inputTypes, {resTy}); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + + rewriter.setInsertionPoint(op); + SmallVector inputVals({src, offsets}); + if (mask) + inputVals.push_back(mask); + if (other) + inputVals.push_back(other); + auto callOp = rewriter.create(loc, funcOp.getSymNameAttr(), + TypeRange({resTy}), inputVals); + rewriter.replaceOp(op, callOp); + return success(); +} + +LogicalResult IndirectStoreConverter::matchAndRewrite( + triton::ascend::IndirectStoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + auto moduleOp = op->getParentOfType(); + rewriter.setInsertionPoint(moduleOp.getBody(), + std::prev(moduleOp.getBody()->end())); + + auto funcName = generateUniqueFuncName(moduleOp, funcNameBase); + + auto src = adaptor.getSrc(); + auto offsets = op.getOffsets(); + auto value = op.getValue(); + auto mask = op.getMask(); + + // convert !tt.ptr to memref + auto srcTy = dyn_cast(src.getType()); + if (!srcTy) { + return rewriter.notifyMatchFailure(op, "expected MemRefType for src"); + } + SmallVector inputTypes({srcTy, offsets.getType(), value.getType()}); + if (mask) + inputTypes.push_back(mask.getType()); + + auto libFnType = rewriter.getFunctionType(inputTypes, {}); + auto funcOp = rewriter.create(loc, funcName.str(), libFnType); + SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private); + + rewriter.setInsertionPoint(op); + SmallVector inputVals({src, offsets, value}); + if (mask) + inputVals.push_back(mask); + rewriter.create(loc, funcOp.getSymNameAttr(), TypeRange({}), + inputVals); + rewriter.eraseOp(op); + return success(); +} + +IndexSelectSimdConverter::IndexSelectSimdConverter(MLIRContext *context) + : OpConversionPattern(context) {} + +LogicalResult IndexSelectSimdConverter::matchAndRewrite( + triton::ascend::IndexSelectSimdOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + // Get converted operands + Value src = adaptor.getSrc(); + Value indexTensor = adaptor.getIndex(); + auto srcShapeVals = adaptor.getSrcShape(); + auto srcOffsetVals = adaptor.getSrcOffset(); + auto readShapeAttr = op.getReadShape(); + int32_t dim = op.getDim(); + + // Get result type + auto resultTensorType = cast(op.getResult().getType()); + auto elemType = resultTensorType.getElementType(); + auto resultShape = resultTensorType.getShape(); + + // Convert src (tt.ptr -> memref) to the correct memref shape + // src is now memref after type conversion, need to reinterpret to full + // shape + auto srcMemRefType = cast(src.getType()); + + // Build multi-dimensional memref type + SmallVector fullSrcShape; + for (auto shapeVal : srcShapeVals) { + if (auto constOp = shapeVal.getDefiningOp()) { + fullSrcShape.push_back(constOp.value()); + } else { + fullSrcShape.push_back(ShapedType::kDynamic); + } + } + auto fullSrcMemRefType = MemRefType::get(fullSrcShape, elemType); + + // Reinterpret cast from 1D to multi-dimensional + // Build strides: stride[i] = product of all dimensions after i + SmallVector sizes, strides; // offsets are 0 + + // Calculate strides from right to left (row-major layout) + SmallVector stridesList; + Value currentStride = rewriter.create(loc, 1); + + for (int i = fullSrcShape.size() - 1; i >= 0; --i) { + stridesList.insert(stridesList.begin(), currentStride); + + // Update stride for next dimension: stride *= size[i] + if (i > 0) { // Don't need to calculate for the first dimension + if (fullSrcShape[i] != ShapedType::kDynamic) { + // Static dimension: multiply by constant + currentStride = rewriter.create( + loc, currentStride, + rewriter.create(loc, fullSrcShape[i])); + } else { + // Dynamic dimension: multiply by runtime value + Value dimSize = srcShapeVals[i]; + if (!dimSize.getType().isIndex()) { + dimSize = rewriter.create( + loc, rewriter.getIndexType(), dimSize); + } + currentStride = + rewriter.create(loc, currentStride, dimSize); + } + } + } + + // Build offsets, sizes, and strides for ReinterpretCastOp + for (size_t i = 0; i < srcShapeVals.size(); ++i) { + // Convert Value to OpFoldResult for sizes + Value sizeVal = srcShapeVals[i]; + if (!sizeVal.getType().isIndex()) { + sizeVal = rewriter.create( + loc, rewriter.getIndexType(), sizeVal); + } + sizes.push_back(sizeVal); + + // Convert Value to OpFoldResult for strides + strides.push_back(stridesList[i]); + } + + OpFoldResult offset = rewriter.getIndexAttr(0); + + // Use the correct builder method for ReinterpretCastOp + auto srcMemRef = rewriter.create( + loc, fullSrcMemRefType, src, offset, sizes, strides); + + // Allocate output buffer + auto resultMemRefType = MemRefType::get(resultShape, elemType); + auto outputBuffer = rewriter.create(loc, resultMemRefType); + + // Get indices tensor type for extracting + auto indicesTensorType = cast(indexTensor.getType()); + int64_t numIndices = indicesTensorType.getShape()[0]; + + // Create for loop + auto zeroIdx = rewriter.create(loc, 0); + auto numIndicesVal = rewriter.create(loc, numIndices); + auto stepOne = rewriter.create(loc, 1); + auto forOp = + rewriter.create(loc, zeroIdx, numIndicesVal, stepOne); + + // Mark as parallel loop + forOp->setAttr("hivm.parallel_loop", rewriter.getUnitAttr()); + + // Build loop body + Block *loopBody = forOp.getBody(); + auto savedInsertionPoint = rewriter.saveInsertionPoint(); + rewriter.setInsertionPointToStart(loopBody); + + // Remove the terminator temporarily + Operation *terminator = &loopBody->back(); + rewriter.setInsertionPoint(terminator); + + Value iv = forOp.getInductionVar(); + + // Extract index from indices tensor + Value selectedIdx = + rewriter.create(loc, indexTensor, ValueRange{iv}); + Value selectedIdxAsIndex = rewriter.create( + loc, rewriter.getIndexType(), selectedIdx); + + // Build source subview offsets/sizes/strides + SmallVector srcSubviewOffsets, srcSubviewSizes, + srcSubviewStrides; + // DenseI32ArrayAttr can be implicitly converted to ArrayRef + ArrayRef readShape = readShapeAttr; + + for (size_t i = 0; i < srcOffsetVals.size(); ++i) { + if (i == static_cast(dim)) { + // Use the selected index for this dimension + srcSubviewOffsets.push_back(selectedIdxAsIndex); + srcSubviewSizes.push_back(rewriter.getIndexAttr(1)); + } else { + // Use provided offset and read size for other dimensions + Value offsetVal = srcOffsetVals[i]; + if (!offsetVal.getType().isIndex()) { + offsetVal = rewriter.create( + loc, rewriter.getIndexType(), offsetVal); + } + srcSubviewOffsets.push_back(offsetVal); + srcSubviewSizes.push_back(rewriter.getIndexAttr(readShape[i])); + } + srcSubviewStrides.push_back(rewriter.getIndexAttr(1)); + } + + auto srcSubview = rewriter.create( + loc, srcMemRef, srcSubviewOffsets, srcSubviewSizes, srcSubviewStrides); + + // Build destination subview + SmallVector dstSubviewOffsets, dstSubviewSizes, + dstSubviewStrides; + for (size_t i = 0; i < resultShape.size(); ++i) { + if (i == static_cast(dim)) { + dstSubviewOffsets.push_back(iv); + dstSubviewSizes.push_back(rewriter.getIndexAttr(1)); + } else { + dstSubviewOffsets.push_back(rewriter.getIndexAttr(0)); + dstSubviewSizes.push_back(rewriter.getIndexAttr(readShape[i])); + } + dstSubviewStrides.push_back(rewriter.getIndexAttr(1)); + } + + auto dstSubview = rewriter.create( + loc, outputBuffer, dstSubviewOffsets, dstSubviewSizes, dstSubviewStrides); + + // Check if index_select is on the trailing axis (last dimension) + if (static_cast(dim) == fullSrcShape.size() - 1) { + // For index_select on the trailing axis, mark as discrete memory access + // This degrades to scalar read/write handling to avoid alignment issues + auto copyOp = rewriter.create(loc, srcSubview, dstSubview); + copyOp->setAttr(ConverterUtils::discreteAttrName, rewriter.getUnitAttr()); + } else { + // For index_select on non-trailing axes, add stride alignment annotation + // This tells the backend to handle address alignment for DMA operations + auto dstMarkOp = rewriter.create(loc, dstSubview); + dstMarkOp->setAttr( + "hfusion.stride_align_dims", + rewriter.getDenseI32ArrayAttr({static_cast(dim)})); + dstMarkOp->setAttr("hfusion.stride_align_value_in_byte", + rewriter.getDenseI32ArrayAttr({32})); + + // Copy from source to destination + rewriter.create(loc, srcSubview, dstSubview); + } + + // Restore insertion point + rewriter.restoreInsertionPoint(savedInsertionPoint); + + // Convert memref to tensor + auto resultTensor = rewriter.create( + loc, resultTensorType, outputBuffer, true, true); + + // Mark as index_select_simd + resultTensor->setAttr("index_select_simd", rewriter.getUnitAttr()); + + // Replace the original op + rewriter.replaceOp(op, resultTensor); + + return success(); +} + +} // namespace TTOpConverters diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.cpp new file mode 100755 index 00000000..2656ad43 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.cpp @@ -0,0 +1,1213 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h" +#include "incubated/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/DescriptorConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/FunctionConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/HoistBroadcast.h" +#include "incubated/Conversion/TritonToLinalgIncubated/LoadStoreConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/TritonOpConverter.h" +#include "incubated/Conversion/TritonToLinalgIncubated/UseAnalysis.h" +#include "incubated/Conversion/UtilsIncubated/InterleaveOptimization.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h" + +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#if __has_include("bishengir/Dialect/HFusion/IR/HFusion.h") +#include "bishengir/Dialect/HFusion/IR/HFusion.h" +#endif +#include "mlir/IR/Builders.h" +#include "mlir/IR/Operation.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#if __has_include("bishengir/Dialect/HIVM/IR/HIVM.h") +#include "bishengir/Dialect/HIVM/IR/HIVM.h" +#endif +#if __has_include("bishengir/Dialect/Annotation/IR/Annotation.h") +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#endif +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Transforms/Transforms.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Visitors.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "mlir/Transforms/Passes.h" + +#include "llvm/ADT/BitVector.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/Twine.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/LogicalResult.h" + +#include "tle/dsa/dialect/include/Conversion/TleToLinalg/DSACopyConverter.h" +#include "tle/dsa/dialect/include/Conversion/TleToLinalg/MathConverter.h" + +#include +#include +#include + +#define DEBUG_TYPE "triton-to-linalg" + +using namespace mlir; +using namespace triton; +using namespace mlir::triton::Incubated; +int nd2nzFlag = 0; +bool compileOn91095Flag = false; +bool existDotFlag = false; + +// Convert CustomOp after operand type converted, +// for example tt.ptr converted to memref. +class CustomOpConverter : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(hivm::CustomOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto res_types = adaptor.getOutputs().getTypes(); + auto new_op = rewriter.create( + op->getLoc(), res_types, adaptor.getOperands(), op->getAttrs()); + rewriter.replaceOp(op, new_op); + return success(); + } +}; + +static bool isSIMTOp(Operation *op) { + if (auto custom_op = dyn_cast(op)) { + return custom_op.getCoreType() == hivm::TCoreType::VECTOR && + custom_op.getVFMode() == hivm::VFMode::SIMT; + } + return isa( + op); +} + +template struct has_getPtr : std::false_type {}; +template +struct has_getPtr().getPtr())>> + : std::true_type {}; + +template struct has_getSrc : std::false_type {}; +template +struct has_getSrc().getSrc())>> + : std::true_type {}; + +template struct has_getBase : std::false_type {}; +template +struct has_getBase().getBase())>> + : std::true_type {}; + +template static Value extractPointer(OpTy op) { + if constexpr (has_getPtr::value) + return op.getPtr(); + else if constexpr (has_getSrc::value) + return op.getSrc(); + else if constexpr (has_getBase::value) + return op.getBase(); + else { + Operation *raw = op.getOperation(); + if (!raw || raw->getNumOperands() == 0) + return Value(); + return raw->getOperand(0); + } +} + +static void setBlockArgumentAttr(BlockArgument blockArg, triton::FuncOp func, + TensorKind tensorKind) { + unsigned argIdx = blockArg.getArgNumber(); + auto existingAttr = + func.getArgAttrOfType(argIdx, "tt.tensor_kind"); + TensorKind oldVal = existingAttr + ? static_cast(existingAttr.getInt()) + : TensorKind::NONE; + + TensorKind finalVal = tensorKind; + if ((oldVal == TensorKind::INPUT && tensorKind == TensorKind::OUTPUT) || + (oldVal == TensorKind::OUTPUT && tensorKind == TensorKind::INPUT)) { + finalVal = TensorKind::INPUT_OUTPUT; + } else if (oldVal == TensorKind::INPUT_OUTPUT) { + finalVal = oldVal; + } + + LLVM_DEBUG(llvm::dbgs() << "Setting tensor_kind for argument " << argIdx + << ": " << finalVal << "\n";); + + func.setArgAttr( + argIdx, "tt.tensor_kind", + IntegerAttr::get(IntegerType::get(func.getContext(), INT_BIT_WIDTH), + static_cast(finalVal))); +} + +template +void TritonToLinalgIncubatedPass::addTensorKindToArguments( + OpTy op, triton::FuncOp func, TensorKind tensorKind) { + Value ptr = extractPointer(op); + if (!ptr) + return; + + LLVM_DEBUG(llvm::dbgs() << "Processing op: " << *op.getOperation() << "\n";); + + Value cur = ptr; + llvm::SmallPtrSet visited; + // Walk back the def-use chain to find originating BlockArgument + while (visited.insert(cur).second) { + // If reach a BlockArgument, set the attribute + if (auto blockArg = dyn_cast(cur)) { + if (blockArg.getOwner() == &func.getBody().front()) { + auto type = blockArg.getType(); + // Check if it's a triton::PointerType + if (!isa(type)) + break; + setBlockArgumentAttr(blockArg, func, tensorKind); + break; + } + } + + Operation *defOp = cur.getDefiningOp(); + if (!defOp) + break; + cur = defOp->getOperand(0); + } +} + +template +void TritonToLinalgIncubatedPass::walkAndMarkTensorKind(triton::FuncOp func) { + (func.walk([&](Ops op) { this->addTensorKindToArguments(op, func, Kind); }), + ...); +} + +TritonTypeConverter::TritonTypeConverter() { + addConversion([](Type type) { return type; }); + + addConversion([](triton::PointerType ptrType) { + Type elem = ptrType.getPointeeType(); + // Handling special case: ptr -> memref + if (auto it = dyn_cast(elem); it && it.getWidth() == 1) { + elem = IntegerType::get(ptrType.getContext(), 8); + LLVM_DEBUG({ + llvm::dbgs() << "[TritonTypeConverter] Normalize i1 pointer to i8 " + "memref. elemType=" + << elem << "\n"; + }); + } + return MemRefType::get({ShapedType::kDynamic}, elem); + }); + + addConversion([](TensorType tensorType) -> Type { + auto elemType = tensorType.getElementType(); + if (auto ptrType = dyn_cast(elemType)) { + elemType = ptrType.getPointeeType(); + } + // Handling special case: tensor -> memref + if (auto it = dyn_cast(elemType); it && it.getWidth() == 1) { + elemType = IntegerType::get(tensorType.getContext(), 8); + LLVM_DEBUG({ + llvm::dbgs() << "[TritonTypeConverter] Normalize i1 tensor to i8 " + "memref. elemType=" + << elemType << "\n"; + }); + } + return MemRefType::get(tensorType.getShape(), elemType); + }); +} + +void TritonToLinalgIncubatedPass::addProgramInfo(triton::FuncOp func, + bool globalKernel) { + OpBuilder b(func); + + auto origFuncType = func.getFunctionType(); + auto origInputTypes = origFuncType.getInputs(); + SmallVector newInputTypes(origInputTypes); + newInputTypes.append(TRITON_PROGRAM_INFO_ARG_COUNT, b.getI32Type()); + + auto newFuncType = + b.getFunctionType(newInputTypes, origFuncType.getResults()); + + func.setFunctionType(newFuncType); + + // If argument attributes exist, extend attribute list. + if (func.getAllArgAttrs()) { + SmallVector newArgAttrs; + func.getAllArgAttrs(newArgAttrs); + newArgAttrs.append(TRITON_PROGRAM_INFO_ARG_COUNT, DictionaryAttr()); + func.setAllArgAttrs(newArgAttrs); + } + + // Append the arguments to the entry block. + for (unsigned i = 0; i < TRITON_PROGRAM_INFO_ARG_COUNT; i++) { + func.getBody().front().addArgument(b.getI32Type(), func.getLoc()); + } + + if (globalKernel) { + func->setAttr(globalKernelAttr, b.getStringAttr("")); + } else { + func->setAttr(globalKernelAttr, b.getStringAttr("local")); + } +} + +LogicalResult TritonToLinalgIncubatedPass::convertMultipleBlockControlFlow( + Operation *funcOp, OpBuilder &builder) { + if (!isa(funcOp)) { + funcOp->emitError( + "convertMultipleBlockControlFlow can only process func::FuncOp!"); + return failure(); + } + + SmallVector candidate; + SmallVector eraseBlocks; + for (Block &block : dyn_cast(funcOp).getBody()) { + auto curTerminator = block.getTerminator(); + if (isa(curTerminator)) { + candidate.push_back(curTerminator); + } else if (isa(curTerminator)) { + if (candidate.empty()) { + curTerminator->emitError( + "funcOp has more than one Block but got an early 'tt.return' Op."); + return failure(); + } + } else if (!isa(curTerminator)) { + funcOp->emitError( + "funcOp has more than one Block but found unsupported Terminator: ") + << *curTerminator; + return failure(); + } + + if (!block.isEntryBlock()) + eraseBlocks.push_back(&block); + } + + LLVM_DEBUG({ + llvm::dbgs() << "Found " << candidate.size() + << " candidate cond_branch operations to convert.\n"; + }); + + if (candidate.empty()) { + funcOp->emitError("funcOp has more than one Block but no candidate " + "Terminator was found!"); + return failure(); + } + + llvm::BitVector visitFlag(candidate.size(), false); + + // Recursive function to convert all cf::CondBranchOp to scf::IfOp + std::function convertToSCF = + [&](Operation *op, Operation *insertPosOp) -> void { + auto condBranchOp = dyn_cast_if_present(op); + auto iter = llvm::find(candidate, condBranchOp); + if (!(condBranchOp && iter != candidate.end())) { + op->emitError( + "convertToSCF must process with condBranchOp in candidates!"); + return; + } + visitFlag.set(iter - candidate.begin()); + + OpBuilder::InsertionGuard guard(builder); + builder.setInsertionPointAfter(insertPosOp); + + // Well, here force to destory original control flow + builder.create( + condBranchOp->getLoc(), condBranchOp.getCondition(), + /*thenBuilder=*/ + [&](OpBuilder &builder, Location loc) { + SmallVector movedOps = llvm::map_to_vector( + condBranchOp.getTrueDest()->without_terminator(), + [](Operation &op) { return &op; }); + for (auto *innerOp : movedOps) { + innerOp->moveBefore(builder.getInsertionBlock(), + builder.getInsertionPoint()); + } + + auto blockTerm = condBranchOp.getTrueDest()->getTerminator(); + if (auto nextCond = dyn_cast(blockTerm)) { + if (movedOps.empty()) { + blockTerm->emitError("movedOps can not be empty before entering " + "convertToSCF (then)!"); + return; + } + convertToSCF(nextCond, movedOps.back()); + } else if (!isa(blockTerm)) { + blockTerm->emitError( + "Unsupported terminator in then branch after structuring"); + } + + builder.create(loc); + }, + /*elseBuilder=*/ + [&](OpBuilder &builder, Location loc) { + SmallVector movedOps = llvm::map_to_vector( + condBranchOp.getFalseDest()->without_terminator(), + [](Operation &op) { return &op; }); + for (auto *innerOp : movedOps) { + innerOp->moveBefore(builder.getInsertionBlock(), + builder.getInsertionPoint()); + } + + auto blockTerm = condBranchOp.getFalseDest()->getTerminator(); + if (auto nextCond = dyn_cast(blockTerm)) { + if (movedOps.empty()) { + blockTerm->emitError("movedOps can not be empty before entering " + "convertToSCF (else)!"); + return; + } + convertToSCF(nextCond, movedOps.back()); + } else if (!isa(blockTerm)) { + blockTerm->emitError( + "Unsupported terminator in else branch after structuring"); + } + builder.create(loc); + }); + }; + + Block::iterator insertOp(candidate.front()); + if (insertOp == candidate.front()->getBlock()->begin()) { + // if the first operation is a cond_branch, we need to insert before it + convertToSCF(candidate.front(), candidate.front()); + } else { + --insertOp; + convertToSCF(candidate.front(), &(*insertOp)); + } + + if (!visitFlag.all()) { + funcOp->emitError("Not all cf.cond_br converted!"); + return failure(); + } + + OpBuilder::InsertionGuard guard(builder); + builder.setInsertionPoint(candidate.front()); + builder.create(candidate.front()->getLoc()); + + for (Operation *eachTerm : candidate) + eachTerm->erase(); + for (Block *block : llvm::reverse(eraseBlocks)) + block->erase(); + + return success(); +} + +void TritonToLinalgIncubatedPass::convertTTFunc(triton::FuncOp func, + const bool existDot, + const bool existSIMTOp) { + OpBuilder builder(func); + + auto name = func.getName(); + auto type = func.getFunctionType(); + + SmallVector argAttrs, resAttrs; + func.getAllArgAttrs(argAttrs); + func.getAllResultAttrs(resAttrs); + + // Special handling for bit-casted tt.ptr arguments + SmallVector inputTypes{type.getInputs()}; + SmallVector retTypes{type.getResults()}; + if (func.getSymVisibility() == "public" && !func.isDeclaration()) { + for (size_t i = 0; i < func.getNumArguments(); ++i) { + auto arg = func.getArgument(i); + // Special method for i1 arg + if (!isa(arg.getType()) || + dyn_cast(arg.getType()).getElementTypeBitWidth() != + 1) { + continue; + } + + SmallVector argVaildUser{arg.getUsers()}; + llvm::erase_if(argVaildUser, [](Operation *op) -> bool { + return isOpTriviallyDead(op); + }); + + if (!argVaildUser.empty()) { + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << arg << " has users:\n"; + int cnt = 0; + for (auto it : argVaildUser) { + os << "users[" << cnt++ << "] = " << *it; + } + }); + if (llvm::all_of(argVaildUser, [](Operation *userOp) { + return isa(userOp); + })) { + auto castOp = cast(*argVaildUser.begin()); + if (castOp.getInputs().size() == 1 && + castOp.getOutputs().size() == 1) { + arg.setType(castOp.getOutputs()[0].getType()); + inputTypes[i] = arg.getType(); + } + } else { + func->emitError(Twine("Unsupported use of func arg at index ") + + Twine(i)); + } + } else { + // Process unused bool ptr type specially, which guarantees bool pointer + // argument's type is realistic and don't mislead backend compiler. + // realistic memory layout of bool pointer is 8 bit width + auto memType = dyn_cast(arg.getType()) + .cloneWith(std::nullopt, builder.getI8Type()); + arg.setType(memType); + inputTypes[i] = arg.getType(); + } + } + } + auto castType = FunctionType::get(func.getContext(), inputTypes, retTypes); + + auto funcFunc = builder.create(func.getLoc(), name, castType); + funcFunc.setAllArgAttrs(argAttrs); + funcFunc.setAllResultAttrs(resAttrs); + auto kernelAttr = func->getAttr(globalKernelAttr); + if (kernelAttr) { + funcFunc->setAttr(globalKernelAttr, kernelAttr); + } + std::string kernelMixMode = "aiv"; + if (existDot) { + // mix also works for pure cube kernel by using the same MAGIC_ELF keyword + kernelMixMode = "mix"; + } + // Set mix_mode in the func attrs so that the backend could know + // the mix_mode by parse the func attrs. + // The backend needs to know the mix_mode because the host wrapper + // needs to set the devbin.magic. Check npu_utils.cpp. + funcFunc->setAttr(kernelMixModeName, builder.getStringAttr(kernelMixMode)); + + std::string parallelMode = "simd"; + if (existSIMTOp) { + parallelMode = "mix_simd_simt"; + } + funcFunc->setAttr(kernelParallelModeName, + builder.getStringAttr(parallelMode)); + + auto &funcFuncBody = funcFunc.getBody(); + auto &funcBody = func.getBody(); + + IRMapping map; + funcBody.cloneInto(&funcFuncBody, map); + + if (!funcFuncBody.hasOneBlock()) { + if (failed(convertMultipleBlockControlFlow(funcFunc, builder))) { + llvm_unreachable("Encounter unsupported control flow"); + } + } + + for (Block &block : funcFuncBody.getBlocks()) { + auto term = block.getTerminator(); + builder.setInsertionPoint(term); + builder.create(func.getLoc(), term->getOperands()); + term->erase(); + } + func.erase(); +} + +void TritonToLinalgIncubatedPass::addDynamicLegal( + ConversionTarget &target, TritonTypeConverter &tritonTypeConverter) { + target.addLegalDialect(); + + // add legal dialect on condition + target.addLegalOp(); + + // decide which ops need conversion based on uses + target.addDynamicallyLegalOp( + [](mlir::Operation *op) { + if (op->use_empty()) { + return false; + } else { + return true; + } + }); + + target.addDynamicallyLegalOp([&](triton::FuncOp op) { + return tritonTypeConverter.isSignatureLegal(op.getFunctionType()); + }); + + // For CustomOp, tt.ptr should be converted to memref. + target.addDynamicallyLegalOp([&](hivm::CustomOp op) { + return all_of(op->getOperandTypes(), [](Type t) { + if (isa(t)) { + return false; + } + if (auto shapedType = dyn_cast(t)) { + return !isa(shapedType.getElementType()); + } + return true; + }); + }); + + target.addDynamicallyLegalOp([](arith::ConstantOp op) { + auto res = op.getResult(); + if (!isa(res.getType())) { + return true; + } + + if (auto denseAttr = dyn_cast(op.getValue())) { + if (!denseAttr.isSplat() || + !isa(denseAttr.getElementType())) { + return true; + } + if (res.hasOneUse() && isa(*res.user_begin())) { + return true; + } + return false; + } + return true; + }); + + target.addDynamicallyLegalOp([](Operation *op) { + return llvm::all_of(op->getOperandTypes(), [](Type t) { + if (isa(t)) { + return false; + } + if (auto shapedType = dyn_cast(t)) { + return shapedType.getElementType().isIntOrFloat(); + } + assert(t.isIntOrIndexOrFloat()); + return true; + }); + }); + + target.addDynamicallyLegalDialect( + [this](Operation *op) { + if (op->hasAttr("MetaUse")) { + return false; + } + + if (isa(op)) { + return true; + } + + bool operateOnTensors = + llvm::all_of(op->getOperandTypes(), + [](Type type) { return isa(type); }); + + return this->namedOps || !operateOnTensors; + }); +} + +void TritonToLinalgIncubatedPass:: + populateTritonToLinalgCanonicalizationPatterns( + RewritePatternSet &patterns) { + patterns.add, + LoadStoreConverter::LoadStoreCanonicalizer, + LoadStoreConverter::LoadStoreCanonicalizer, + LoadStoreConverter::LoadStoreCanonicalizer>( + patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add, + // TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + // TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer, + TTOpConverters::ScalarMathCanonicalizer + // By test, the following ops do not need canonicalization. + // TTOpConverters::ScalarMathCanonicalizer + // TTOpConverters::ScalarMathCanonicalizer + // TTOpConverters::ScalarMathCanonicalizer + >(patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add( + patterns.getContext()); + if (this->enableSelectAnalysis) { + patterns.add(patterns.getContext()); + } +} + +void TritonToLinalgIncubatedPass::populateTritonToLinalgConversionPatterns( + TypeConverter &typeConverter, RewritePatternSet &patterns, + unsigned int launchGridRank) { + nd2nzFlag = this->enableNd2nzOnVector; + populateFunctionOpInterfaceTypeConversionPattern( + patterns, typeConverter); + + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add(patterns.getContext()); + if (compileOn91095Flag && existDotFlag) { + patterns.add( + patterns.getContext()); + } else { + patterns.add(patterns.getContext()); + } + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + // reduce converters + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + + patterns.add(patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add( + patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add>( + patterns.getContext()); + patterns.add>( + patterns.getContext()); + patterns.add(patterns.getContext()); + + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + patterns.add(patterns.getContext()); + + // Add convert pattern for CustomOp. + patterns.add(patterns.getContext()); + + if (!this->namedOps) { + linalg::populateElementwiseToLinalgConversionPatterns(patterns); + } +} + +void TritonToLinalgIncubatedPass::getDependentDialects( + DialectRegistry ®istry) const { + registry.insert(); +} + +LogicalResult +TritonToLinalgIncubatedPass::processDescriptorOperations(ModuleOp moduleOp) { + // --- ConversionTarget: dynamic legality checks --- + mlir::ConversionTarget target(getContext()); + + // Dialect-level dynamic legality: ops are legal if none of their + // operands/results use TensorDescType. + target.addDynamicallyLegalDialect< + mlir::arith::ArithDialect, mlir::scf::SCFDialect, triton::TritonDialect>( + [](mlir::Operation *op) { + return !DescriptorConverter::hasATensorDescriptorType( + op->getOperandTypes()) && + !DescriptorConverter::hasATensorDescriptorType( + op->getResultTypes()); + }); + // Function signature legality: Triton FuncOp is legal if its inputs/outputs + // contain no TensorDescType. + target.addDynamicallyLegalOp([](triton::FuncOp funcOp) { + return !DescriptorConverter::hasATensorDescriptorType( + funcOp.getFunctionType().getInputs()) && + !DescriptorConverter::hasATensorDescriptorType( + funcOp.getFunctionType().getResults()); + }); + target.addLegalOp(); + target.addIllegalOp(); + + // --- Patterns --- + mlir::RewritePatternSet patterns(&getContext()); + patterns.add( + patterns.getContext()); + patterns.add( + patterns.getContext()); + + mlir::ConversionConfig config; + config.buildMaterializations = true; + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns), + config))) { + moduleOp->emitError("failed to convert tensor descriptor operations"); + return failure(); + } + + return success(); +} + +LogicalResult +TritonToLinalgIncubatedPass::processPtrBroadcastOperations(ModuleOp moduleOp) { + // --- ConversionTarget: dynamic legality checks --- + mlir::ConversionTarget target(getContext()); + target.addLegalOp(); + target.addLegalOp(); + target.addDynamicallyLegalOp([](triton::BroadcastOp op) { + if (op->hasAttr("MetaUse")) { + return true; + } + auto resultType = dyn_cast(op.getType()); + HoistBroadcast::BroadcastHoister hoister(op); + return !(isa(resultType.getElementType()) && + hoister.canBroadcast()); + }); + + // --- Patterns --- + mlir::RewritePatternSet patterns(&getContext()); + patterns.add(patterns.getContext()); + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + moduleOp->emitError("failed to convert ptr broadcast operations"); + return failure(); + } + + return success(); +} + +void TritonToLinalgIncubatedPass::annotateTensorKindForModule( + ModuleOp moduleOp) { + moduleOp.walk([&](triton::FuncOp func) { + // INPUT tensors + this->walkAndMarkTensorKind< + TensorKind::INPUT, triton::LoadOp, triton::ascend::IndexSelectSimdOp, + triton::ascend::EmbeddingGatherOp, triton::ascend::GatherOutToUbOp, + triton::ascend::IndirectLoadOp>(func); + // OUTPUT tensors + this->walkAndMarkTensorKind< + TensorKind::OUTPUT, triton::StoreOp, triton::ascend::IndexPutOp, + triton::ascend::ScatterUbToOutOp, triton::ascend::IndirectStoreOp>( + func); + // INPUT_OUTPUT tensors + this->walkAndMarkTensorKind(func); + }); +} + +void TritonToLinalgIncubatedPass::runOnOperation() { + compileOn91095Flag = this->compileOn91095; + + auto moduleOp = getOperation(); + + // Check if the kernel contains tl.dot. Without tl.dot, + // the kernel would be pure AIV kernel. + bool existDot = false; + moduleOp.walk([&](triton::DotOp dotOp) { + existDot = true; + return WalkResult::interrupt(); + }); + moduleOp.walk([&](triton::DotScaledOp dotScaledOp) { + existDot = true; + return WalkResult::interrupt(); + }); + existDotFlag = existDot; + + bool existSIMTOp = false; + moduleOp.walk([&](Operation *op) { + if (isSIMTOp(op)) { + existSIMTOp = true; + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Found SIMT op in function: "; + os << op->getName(); + os << "\n"; + }); + return WalkResult::interrupt(); + } + return WalkResult::advance(); + }); + + RewritePatternSet canonicalizerPatterns(&getContext()); + + // Execute tensor descriptor operations conversion + if (failed(processDescriptorOperations(moduleOp))) { + signalPassFailure(); + } + + // 0. Annotate Memory-Related Triton FuncOps with tensor_kind (used by + // profiling). + annotateTensorKindForModule(moduleOp); + + // 1. Canonicalize load/store related patterns. + this->populateTritonToLinalgCanonicalizationPatterns(canonicalizerPatterns); + if (failed(applyPatternsAndFoldGreedily(moduleOp, + std::move(canonicalizerPatterns)))) { + moduleOp->emitError("failed to apply Canonicalizer Patterns"); + signalPassFailure(); + } + + // 2. Perform use analysis on FuncOp. + moduleOp.walk([this](triton::FuncOp op) { + if (failed(runUseAnalysis(op))) { + signalPassFailure(); + } + }); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + TritonTypeConverter tritonTypeConverter{}; + + // 3. Mark legal dialects and operations. + this->addDynamicLegal(target, tritonTypeConverter); + + // 4. Mark ops that must be converted explicitly (e.g. tt.scan). + auto loopOpLegalFn = [](LoopLikeOpInterface op) { + return !op.getOperation()->hasAttr("UnhandledLoopOp"); + }; + + target.addIllegalOp(); + target.addDynamicallyLegalOp(loopOpLegalFn); + target.addDynamicallyLegalOp(loopOpLegalFn); + + // 5. Register converters for all illegal Triton ops. + // Execute ptr broadcast operations conversion + if (failed(processPtrBroadcastOperations(moduleOp))) { + signalPassFailure(); + } + this->populateTritonToLinalgConversionPatterns(tritonTypeConverter, patterns, + LAUNCH_GRID_RANK); + triton::tle::populateTleMathOpConversionPatterns(tritonTypeConverter, + patterns); + triton::tle::populateTleCopyOpConversionPatterns(tritonTypeConverter, + patterns); + + // 6. Inject program id / number of programs arguments into each Triton kernel + // function. + for (auto func : getOperation().getOps()) { + addProgramInfo(func, globalKernel); + } + + moduleOp.walk([this](LoopLikeOpInterface loopOp) { + auto *op = loopOp.getOperation(); + if (!op->hasAttr("ExtractedLoadOrStore")) + op->setAttr("UnhandledLoopOp", UnitAttr::get(op->getContext())); + + for (auto res : loopOp->getResults()) { + if (auto tensorType = dyn_cast(res.getType()); + tensorType && + !isa(tensorType.getElementType())) { + IRRewriter rewriter(op->getContext()); + rewriter.setInsertionPointAfter(op); + auto newVal = + rewriter.create(op->getLoc(), res.getType(), res); + rewriter.replaceAllUsesExcept(res, newVal, newVal); + } + } + }); + + // 7. Convert ops. + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { + moduleOp->emitError("failed to apply Conversion Patterns"); + signalPassFailure(); + } + + // 8. Convert function prologue/epilogue. + moduleOp.walk([&](triton::FuncOp func) { + this->convertTTFunc(func, existDot, existSIMTOp); + }); + + // 9. Clean up dead code and simplify IR. + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } + + // Calculate size of PointerCastOp precisely + SmallVector castOps; + + moduleOp.walk([&](hivm::PointerCastOp op) { castOps.push_back(op); }); + + for (auto op : castOps) { + SmallVector userOps(op->getUsers().begin(), + op->getUsers().end()); + IRRewriter rewriter(&getContext()); + rewriter.setInsertionPointAfter(op); + Value addr = op.getAddrs()[0]; + auto elementType = + cast(op.getResult().getType()).getElementType(); + Value elementTypeSize; + if (auto intType = dyn_cast(elementType)) { + elementTypeSize = rewriter.create( + op.getLoc(), + rewriter.getIntegerAttr(addr.getType(), intType.getWidth() / 8)); + } else if (auto floatType = dyn_cast(elementType)) { + elementTypeSize = rewriter.create( + op.getLoc(), + rewriter.getIntegerAttr(addr.getType(), floatType.getWidth() / 8)); + } else { + llvm_unreachable("Cannot get memory size"); + } + + for (auto userOp : userOps) { + auto reinterpretCastOp = cast(userOp); + auto sizes = reinterpretCastOp.getStaticSizes(); + auto staticStrides = reinterpretCastOp.getStaticStrides(); + auto strides = reinterpretCastOp.getStrides(); + if (reinterpretCastOp.getStaticOffsets().size() != 1) + userOp->emitError("IntToPtrOp must converted to PointerCastOp of " + "memref type"); + int64_t castOpSize = 0; + SmallVector dynamicSizes; + for (const auto &[size, stride] : llvm::zip_equal(sizes, staticStrides)) { + assert(!ShapedType::isDynamic(size)); + if (ShapedType::isDynamic(stride)) + dynamicSizes.push_back(size); + else + castOpSize = size * stride; + } + rewriter.setInsertionPoint(reinterpretCastOp); + Value dynamicSize = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(castOpSize)); + for (const auto &[size, stride] : + llvm::zip_equal(dynamicSizes, strides)) { + Value axisSize = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(size)); + axisSize = + rewriter.create(op.getLoc(), stride, axisSize); + dynamicSize = + rewriter.create(op.getLoc(), dynamicSize, axisSize); + } + Value offsetValue; + auto staticOffset = reinterpretCastOp.getStaticOffsets()[0]; + if (ShapedType::isDynamic(staticOffset)) { + offsetValue = reinterpretCastOp.getOffsets()[0]; + if (offsetValue.getType() != addr.getType()) + offsetValue = rewriter.create( + op.getLoc(), addr.getType(), offsetValue); + } else { + offsetValue = rewriter.create( + op.getLoc(), rewriter.getIntegerAttr(addr.getType(), staticOffset)); + } + offsetValue = rewriter.create(op.getLoc(), offsetValue, + elementTypeSize); + Value realAddr = + rewriter.create(op.getLoc(), addr, offsetValue); + auto memrefType = MemRefType::get({ShapedType::kDynamic}, elementType); + auto newCastOp = rewriter.create( + op.getLoc(), memrefType, realAddr, dynamicSize); + auto markOp = rewriter.create(op.getLoc(), + newCastOp.getResult()); + markOp->setAttr(hivm::AddressSpaceAttr::getMnemonic(), + {hivm::AddressSpaceAttr::get(rewriter.getContext(), + hivm::AddressSpace::GM)}); + rewriter.replaceOpWithNewOp( + reinterpretCastOp, + cast(reinterpretCastOp.getResult().getType()), newCastOp, + ValueRange({}), reinterpretCastOp.getSizes(), + reinterpretCastOp.getStrides(), SmallVector({0}), + reinterpretCastOp.getStaticSizes(), + reinterpretCastOp.getStaticStrides()); + } + rewriter.eraseOp(op); + } + + // Try interleave optimization + llvm::DenseMap> interleaveCandidate; + llvm::DenseMap> + interleaveCandidateWithMask; + moduleOp.walk([&](bufferization::MaterializeInDestinationOp materializeOp) { + if (auto reinterpretCastOp = + materializeOp.getDest() + .getDefiningOp()) { + if (llvm::isa(reinterpretCastOp.getSource()) && + reinterpretCastOp.getStaticStrides().back() == 2) { + interleaveCandidate[llvm::cast( + reinterpretCastOp.getSource())] + .push_back(materializeOp); + } + } + + // Difference is that converted op chain of store with mask has + // `memref::SubViewOp` + if (auto subviewOp = + materializeOp.getDest().getDefiningOp()) { + if (!llvm::isa( + materializeOp.getSource().getDefiningOp())) + return WalkResult::advance(); + + if (auto reinterpretCastOp = + subviewOp.getSource() + .getDefiningOp()) { + if (llvm::isa(reinterpretCastOp.getSource()) && + reinterpretCastOp.getStaticStrides().back() == 2) { + interleaveCandidateWithMask[llvm::cast( + reinterpretCastOp.getSource())] + .push_back(materializeOp); + } + } + } + + return WalkResult::advance(); + }); + + for (auto [blockArg, materializeVec] : interleaveCandidate) { + // Just enable optimization where exists double materializeOp with same + // block argument destination. + if (materializeVec.size() != 2) + continue; + auto result = InterleaveStatusOptimization(materializeVec); + } + + for (auto [blockArg, materializeVec] : interleaveCandidateWithMask) { + if (materializeVec.size() != 2) + continue; + auto result = InterleaveStatusWithMaskOptimization(materializeVec); + } + + // Force to add an argument at the beginning of function arguments, which + // represents stub arg for workspace. Default type is memref + for (auto func : getOperation().getOps()) { + if (!func->hasAttr("global_kernel")) + continue; + + auto context = func.getContext(); + constexpr int64_t syncBlockLockArgIdx = 0; + NamedAttribute syncBlockLockArgAttr( + StringAttr::get(context, "syncBlockLock"), UnitAttr::get(context)); + MemRefType syncBlockLockArgType = + MemRefType::get(SmallVector(1, ShapedType::kDynamic), + IntegerType::get(context, 8)); + func.insertArgument(syncBlockLockArgIdx, // argIndex + syncBlockLockArgType, // argType + nullptr, func->getLoc()); // dicAttr + func->setAttr("SyncBlockLockArgIdx", + IntegerAttr::get(IntegerType::get(&getContext(), 64), + 0)); // 64: 64位整型 + + constexpr int64_t workspaceArgIdx = 1; + MemRefType workspaceArgType = + MemRefType::get(SmallVector(1, ShapedType::kDynamic), + IntegerType::get(context, 8)); + NamedAttribute workspaceArgAttr(StringAttr::get(context, "workspace"), + UnitAttr::get(context)); + + func.insertArgument(/*argIndex*/ workspaceArgIdx, + /*argType*/ workspaceArgType, + /*dicAttr*/ nullptr, func->getLoc()); + func->setAttr("WorkspaceArgIdx", + IntegerAttr::get(IntegerType::get(&getContext(), 64), + 1)); // 64: 64位整型 + } + + // Fix the Location info + moduleOp.walk([&](Operation *op) { + auto loc = op->getLoc(); + if (isa(loc)) { + llvm::SmallPtrSet stopOps; + traverseForwardUpdateUserChainIf( + op, + /*conditionFn*/ + [](Operation *curOp) { return false; }, + /*stopFn*/ + [](Operation *curOp) { return !isa(curOp->getLoc()); }, + /*actionFn*/ + nullptr, stopOps); + if (stopOps.empty()) { + op->emitWarning() << *op << " and its users all have no location!"; + } else { + Operation *goodOp = *stopOps.begin(); + op->setLoc(goodOp->getLoc()); + } + } + return WalkResult::advance(); + }); +} + +std::unique_ptr> +triton::Incubated::createTritonToLinalgIncubatedPass(bool globalKernel, + bool namedOps, + bool enableNd2NzOnVector, + bool enableSelectAnalysis, + bool compileOn91095) { + return std::make_unique( + globalKernel, namedOps, enableNd2NzOnVector, enableSelectAnalysis, + compileOn91095); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/UseAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/UseAnalysis.cpp new file mode 100755 index 00000000..a36d867e --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/UseAnalysis.cpp @@ -0,0 +1,531 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToLinalgIncubated/UseAnalysis.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#include "tle/dsa/dialect/include/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Analysis/DataFlow/ConstantPropagationAnalysis.h" +#include "mlir/Analysis/DataFlow/DeadCodeAnalysis.h" + +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Debug.h" + +using namespace mlir; +using namespace triton; +using namespace dataflow; +using namespace mlir::triton::Incubated; +#define DEBUG_TYPE "triton-use-analysis" + +std::string stringifyUseType(UseType useTy) { + std::string ret; + if (useTy == UseType::MetaUse) { + ret = "MetaUse"; + } else if (useTy == UseType::DataUse) { + ret = "DataUse"; + } else if (useTy == UseType::MixUse) { + ret = "MixUse"; + } else if (useTy == UseType::Undefined) { + ret = "Undefined"; + } + return ret; +} + +#if LLVM_VERSION_MAJOR >= 20 +LogicalResult mlir::triton::Incubated::UseAnalysis::visitOperation( + Operation *op, ArrayRef operands, + ArrayRef results) { +#else +void triton::UseAnalysis::visitOperation(Operation *op, + ArrayRef operands, + ArrayRef results) { +#endif + + if (op->getResults().size() == 1) { + auto resultType = dyn_cast(op->getResult(0).getType()); + if (resultType && isa(resultType.getElementType())) { + for (auto opnd : operands) { + propagateUse(opnd, UseType::MetaUse); + } + } + } + + TypeSwitch(op) + .Case([&](auto load) { + propagateUse(operands[0], UseType::MetaUse); + auto mask = load.getMask(); + auto other = load.getOther(); + if (mask) { + assert(mask != other && "mask and other cannot be the same"); + propagateUse(operands[1], UseType::MetaUse); + } + if (other) { + propagateUse(operands[2], UseType::MetaUse); + } + }) + .Case( + [&](auto print) { propagateUse(operands[0], UseType::DataUse); }) + .Case( + [&](auto assert) { propagateUse(operands[0], UseType::DataUse); }) + .Case([&](auto store) { + propagateUse(operands[0], UseType::MetaUse); + propagateUse(operands[1], UseType::DataUse); + auto value = store.getValue(); + auto mask = store.getMask(); + if (mask) { + assert(mask != value && "mask and data cannot be the same"); + propagateUse(operands[2], UseType::MetaUse); + } + }) + .Case([&](auto store) { + propagateUse(operands[0], UseType::MetaUse); + propagateUse(operands[1], UseType::MetaUse); + propagateUse(operands[2], UseType::DataUse); + auto value = store.getValue(); + auto mask = store.getMask(); + if (mask) { + assert(mask != value && "mask and data cannot be the same"); + propagateUse(operands[3], UseType::MetaUse); + } + }) + // Consider triton::AtomicRMWOp as store operation + .Case([&](auto atomicOp) { + propagateUse(operands[0], UseType::MetaUse); + propagateUse(operands[1], UseType::DataUse); + auto value = atomicOp.getVal(); + auto mask = atomicOp.getMask(); + if (mask) { + assert(mask != value && "mask and data cannot be the same"); + propagateUse(operands[2], UseType::MetaUse); + } + }) + .Case([&](auto atomicOp) { + propagateUse(operands[0], UseType::MetaUse); + propagateUse(operands[1], UseType::DataUse); + propagateUse(operands[2], UseType::DataUse); + auto value = atomicOp.getVal(); + }) + .Case([&](auto dot) { + propagateResults(operands[0], results); + propagateResults(operands[1], results); + + auto opc = dot.getC(); + triton::SplatOp splat; + if (opc) { + splat = opc.template getDefiningOp(); + } + + if (opc && splat && splat.getSrc().getDefiningOp()) { + propagateUse(operands[2], UseType::MetaUse); + } else { + propagateUse(operands[2], UseType::DataUse); + } + }) + .Case([&](auto loopOp) { + for (const auto &[yield, init, result] : llvm::zip_equal( + loopOp.getYieldedValues(), loopOp.getInits(), results)) { + propagateResults(getLatticeElement(yield), {result}); + propagateResults(getLatticeElement(init), {result}); + } + }) + .Default([&](Operation *op) { + // this condition account for tt.addptr + for (auto operand : operands) { + propagateResults(operand, results); + } + }); +#if LLVM_VERSION_MAJOR >= 20 + return success(); +#endif +} + +void setMixUseRecursively(Operation *rootOp, bool applyRoot = true) { + traverseBackwardUpdateOperandChainIf( + rootOp, + // ConditionFn + [rootOp, applyRoot](Operation *curOp) { + for (auto res : curOp->getResults()) { + auto tensorType = dyn_cast(res.getType()); + if (tensorType && + isa(tensorType.getElementType())) + return false; + } + return isMetaUse(curOp) && (curOp != rootOp || applyRoot); + }, + // StopFn + [rootOp](Operation *curOp) { + return isa(curOp) && curOp != rootOp; + }, + // ActionFn + [](OpBuilder &b, Operation *op) { + LLVM_DEBUG({ op->setAttr("MixUse", UnitAttr::get(b.getContext())); }); + op->removeAttr("MetaUse"); + }); +} + +std::optional isIterArgMixUse(Value v, Value target, + const DataFlowSolver &solver) { + auto defOp = v.getDefiningOp(); + auto *use = solver.lookupState(v); + if ((use && use->type == UseType::DataUse) || + isa_and_nonnull(defOp)) + return true; + if (v == target) + return false; + if (!defOp) + return std::nullopt; + for (auto oper : defOp->getOperands()) { + auto res = isIterArgMixUse(oper, target, solver); + if (res.has_value()) + return res.value() || !isMetaUse(defOp); + } + return std::nullopt; +} + +void postProcessWhileOp(scf::WhileOp op, const DataFlowSolver &solver) { + for (const auto &[res, arg] : + llvm::zip_equal(op->getResults(), op.getConditionOp().getArgs())) { + auto *defOp = arg.getDefiningOp(); + if (!defOp) + continue; + auto *use = solver.lookupState(res); + if (use && use->type == UseType::DataUse) + setMixUseRecursively(defOp); + } + for (const auto &[yield, regionArg] : llvm::zip_equal( + op.getYieldOp().getOperands(), op.getBeforeArguments())) { + auto *defOp = yield.getDefiningOp(); + if (!defOp) + continue; + if (isIterArgMixUse(yield, regionArg, solver).value_or(false)) + setMixUseRecursively(defOp); + } +} + +void postProcessLoopOp(LoopLikeOpInterface loopOp, + const DataFlowSolver &solver) { + if (auto whileOp = dyn_cast(loopOp.getOperation())) { + postProcessWhileOp(whileOp, solver); + return; + } + for (const auto &[res, yield, regionArg] : + llvm::zip_equal(loopOp->getResults(), loopOp.getYieldedValues(), + loopOp.getRegionIterArgs())) { + auto *defOp = yield.getDefiningOp(); + if (!defOp) + continue; + auto *use = solver.lookupState(res); + if ((use && use->type == UseType::DataUse) || + isIterArgMixUse(yield, regionArg, solver).value_or(false)) + setMixUseRecursively(defOp); + } +} + +LogicalResult mlir::triton::Incubated::runUseAnalysis(triton::FuncOp &funcOp) { + MLIRContext *context = funcOp.getContext(); + SymbolTableCollection symbolTable; + + DataFlowSolver solver; + solver.load(); + solver.load(); + solver.load(symbolTable); + if (failed(solver.initializeAndRun(funcOp))) { + return failure(); + } + auto &os = llvm::dbgs(); + // Walk the func op, convert tags on operands to tags on operations + funcOp.walk([&](Operation *op) { + LLVM_DEBUG({ os << "[UseAnalysis] op is " << *op << "\n"; }); + UseType useType = UseType::Undefined; + for (auto result : op->getResults()) { + LLVM_DEBUG({ os << "[UseAnalysis] ===> result is " << result << "\n"; }); + auto use = solver.lookupState(result); + assert(use && "Lattice value not found"); + auto thisUseType = use->type; + LLVM_DEBUG({ + os << "[UseAnalysis] ==========> useType is " + << stringifyUseType(thisUseType) << "\n"; + }); + if (thisUseType == UseType::Undefined) { + continue; + } + if (useType == UseType::Undefined) { + useType = thisUseType; + } + if (thisUseType == UseType::MixUse || thisUseType != useType) { + useType = UseType::MixUse; + break; + } + } + + if (useType == UseType::Undefined) { + LLVM_DEBUG({ op->setAttr("Undefined", UnitAttr::get(context)); }); + return; + } else if (useType == UseType::MetaUse) { + if (!isa(op)) { + assert(op->getNumResults() == 1 && + "Ops used for meta computation are expected to have one result"); + } + for (auto it = 0; it < op->getNumResults(); ++it) { + // Only set the tag if the operation uses tensors + if (isa(op->getResult(it).getType()) || + (isa(op) && + op->hasAttr(ConverterUtils::discreteAttrName)) || + (isa(op) && + isa(op->getResult(it).getType()))) { + // Setting tag for erasing op later + op->setAttr("MetaUse", UnitAttr::get(context)); + } + } + return; + } else if (useType == UseType::DataUse) { + LLVM_DEBUG({ op->setAttr("DataUse", UnitAttr::get(context)); }); + return; + } + + assert(useType == UseType::MixUse); + + // If the operation only produces scalars, no need to clone it + bool shapedResult = true; + for (auto result : op->getResults()) + shapedResult &= isa(result.getType()); + if (!shapedResult || isa(op)) { + LLVM_DEBUG({ op->setAttr("MixUse", UnitAttr::get(context)); }); + return; + } + + llvm::SetVector metaUsers; + for (auto result : op->getResults()) { + for (auto user : result.getUsers()) { + TypeSwitch(user) + .Case([&](auto load) { + auto ptr = load.getPtr(); + auto mask = load.getMask(); + auto other = load.getOther(); + if (result == ptr || result == mask || result == other) { + metaUsers.insert(user); + } + }) + .Case([&](auto store) { + auto ptr = store.getPtr(); + auto mask = store.getMask(); + if (result == ptr || result == mask) { + metaUsers.insert(user); + } + }) + .Case([&](auto indirectstore) { + auto src = indirectstore.getSrc(); + auto offset = indirectstore.getOffsets(); + auto mask = indirectstore.getMask(); + if (result == src || result == offset || result == mask) { + metaUsers.insert(user); + } + }) + .Case([&](auto atomicOp) { + auto ptr = atomicOp.getPtr(); + auto mask = atomicOp.getMask(); + if (result == ptr || result == mask) + metaUsers.insert(user); + }) + .Case([&](auto atomicOp) { + auto ptr = atomicOp.getPtr(); + if (result == ptr) + metaUsers.insert(user); + }) + .Case([&](auto dot) { + auto opc = dot.getC(); + triton::SplatOp splat; + if (opc) { + splat = opc.template getDefiningOp(); + } + + if (opc && splat && + splat.getSrc().getDefiningOp()) { + metaUsers.insert(user); + } + }) + .Case([&](auto print) {}) + .Default([&](Operation *op) { + bool allMeta = true; + for (auto res : op->getResults()) { + auto resUse = solver.lookupState(res); + if (resUse->type != UseType::MetaUse) { + allMeta = false; + break; + } + } + if (allMeta) { + metaUsers.insert(user); + } + }); + } + } + + // If the operation doesn't have direct meta users, no need to clone it + if (metaUsers.empty()) { + LLVM_DEBUG({ op->setAttr("MixUse", UnitAttr::get(context)); }); + return; + } + + if (isa(op)) + return; + + // Clone the operation; switch all meta users to use the clone + OpBuilder builder(op); + auto clone = builder.clone(*op); + LLVM_DEBUG({ op->setAttr("MixUse", UnitAttr::get(context)); }); + + // Setting tag for erasing op later + clone->setAttr("MetaUse", UnitAttr::get(context)); + + for (auto [res_i, result] : llvm::enumerate(op->getResults())) { + for (auto user : metaUsers) { + for (auto &operand : user->getOpOperands()) { + if (operand.get() == result) { + operand.set(clone->getResult(res_i)); + } + } + } + } + }); + LLVM_DEBUG({ + os << "[UseAnalysis] Before post-process, funcOp is " << *funcOp << "\n"; + }); + // Post-process + funcOp.walk([&](Operation *op) { + // Handle indirect load and store case. + // For example, load(1st) -> computeOp -> load(2nd), + // or load -> computeOp -> store + // The first load is IndirectLoadInterfaceOp. + // Do not inplace replace MetaUse by MixUse. Because the condition checking + // depends on that the op has the attr of MetaUse. + // Handle the indirect load interface op + // We first trace from the 1st load to the 2nd load with the ops between + // them marked as MixUse. Then we traceback from the 2nd load to mark defs + // MixUse. + if (opIsIndirectLoad(op) || opIsIndirectCalc(op) || + isa(op)) { + LLVM_DEBUG({ + os << "[UseAnalysis] Found indirect load interface op: " << *op << "\n"; + }); + llvm::SmallPtrSet stopOps; + // Modify the users of this op's result. + traverseForwardUpdateUserChainIf( + op, + /*conditionFn*/ + [op](Operation *curOp) { return isMetaUse(curOp) && curOp != op; }, + /*stopFn*/ + [&](Operation *curOp) { + // triton::LoadOp or triton::StoreOp without MetaUse means + // it is an indirect load or store + // instead of the load providing the offset. + // The pattern is as follows, + // load -> ops -> load + // load -> ops -> store + // We need to ensure the intermediate ops are marked MixUse + // so that they will be replaced instead of be erased without + // conversion. + return (isa(curOp) || isa(curOp) || + isa(curOp) || + isa(curOp)) && + !isMetaUse(curOp); + }, + /*actionFn*/ + [](OpBuilder &b, Operation *op) { + LLVM_DEBUG( + { op->setAttr("MixUse", UnitAttr::get(b.getContext())); }); + op->removeAttr("MetaUse"); + }, + stopOps); + LLVM_DEBUG({ + os << "[UseAnalysis] stopOps are \n"; + for (auto [idx, stopOp] : llvm::enumerate(stopOps)) + os << idx << ": " << *stopOp << "\n"; + }); + LLVM_DEBUG({ + os << "[UseAnalysis] After trace, funcOp is " << *funcOp << "\n"; + }); + for (auto *stopOp : stopOps) + setMixUseRecursively(stopOp, /*applyRoot=*/false); + LLVM_DEBUG({ + os << "[UseAnalysis] After traceback of stopOp, funcOp is " << *funcOp + << "\n"; + }); + // Modify this op. + LLVM_DEBUG({ op->setAttr("MixUse", UnitAttr::get(context)); }); + op->removeAttr("MetaUse"); + } + if (op->hasAttr(ConverterUtils::discreteAttrName)) + setMixUseRecursively(op); + if (auto loopOp = dyn_cast(op)) { + postProcessLoopOp(loopOp, solver); + } else if (auto ifOp = dyn_cast(op)) { + SmallVector yields(ifOp.thenYield().getOperands()); + if (!ifOp.getElseRegion().empty()) + yields.append(llvm::to_vector(ifOp.elseYield().getOperands())); + for (auto yield : yields) { + if (auto *defOp = yield.getDefiningOp()) + setMixUseRecursively(defOp); + } + } else if (auto atomicRmwOp = dyn_cast(op)) { + auto mask = atomicRmwOp.getMask(); + if (mask && op->hasAttr(ConverterUtils::discreteMaskAttrName)) + setMixUseRecursively(mask.getDefiningOp()); + } + }); + // Remove MetaUse in case of MixUse existing in the op + funcOp.walk([&](Operation *op) { + if (isMetaUse(op) && isMixUse(op)) { + op->removeAttr("MetaUse"); + } + }); + LLVM_DEBUG({ + os << "[UseAnalysis] After post-process, funcOp is " << *funcOp << "\n"; + }); + return success(); +} + +MetaUseEraser::MetaUseEraser(MLIRContext *context) + : RewritePattern(MatchAnyOpTypeTag(), /*benefit=*/10, context) {} + +LogicalResult MetaUseEraser::matchAndRewrite(Operation *op, + PatternRewriter &rewriter) const { + LLVM_DEBUG({ + int64_t count = 0; + for (auto result : op->getResults()) { + count += std::distance(result.use_begin(), result.use_end()); + } + llvm::dbgs() << "Number of user: " << count << "\n"; + }); + if (isa(op)) { + return rewriter.notifyMatchFailure(op, + "AddPtrOp will be handled separately"); + } + if (isMetaUse(op)) { + rewriter.eraseOp(op); + return success(); + } + return rewriter.notifyMatchFailure(op, "requires meta ops"); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructured/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructured/CMakeLists.txt new file mode 100755 index 00000000..f219fcb8 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructured/CMakeLists.txt @@ -0,0 +1,22 @@ +add_triton_library(WaferTritonToStructured + TritonToStructuredPass.cpp + + DEPENDS + TritonStructuredTableGen + TritonToStructuredConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + MLIRReconcileUnrealizedCasts + TritonIR + TritonTransforms + TritonSharedAnalysisStructured + TritonStructuredIR +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructured/TritonToStructuredPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructured/TritonToStructuredPass.cpp new file mode 100755 index 00000000..11be26b1 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructured/TritonToStructuredPass.cpp @@ -0,0 +1,388 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation, Meta Platforms. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "mlir/Conversion/ReconcileUnrealizedCasts/ReconcileUnrealizedCasts.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/TypeRange.h" +#include "mlir/IR/Types.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/Support/LogicalResult.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/AnalysisStructured/PtrAnalysis.h" +#include "triton-shared/Conversion/TritonToStructured/TritonToStructured.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/DialectConversion.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "mlir/Transforms/Passes.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/LogicalResult.h" +#include +#include + +#define DEBUG_TYPE "triton-to-structured" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonToStructured/Passes.h.inc" + +namespace { + +class TritonToStructuredPass + : public TritonToStructuredBase { + + static TupleType getStructuredStateTupleType(MLIRContext *context, Type t) { + SmallVector tupleTypes{t}; + auto [offsetTypes, strideTypes] = + *tts::GetStructuredStateOp::getOffsetAndStrideTypes(context, t); + tupleTypes.append(offsetTypes); + tupleTypes.append(strideTypes); + return TupleType::get(context, tupleTypes); + } + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + LogicalResult convertToPointerTupleWithOffsetsAndStrides() { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + + auto context = &getContext(); + TypeConverter converter; + converter.addConversion([](Type type) { return type; }); + + // We are doing a 1->1 type conversion here, where a triton pointer type + // maps to a tuple of {pointer, offset_0, offset_1,..., stride_0, + // stride_1,...} type. + // + // Case 1: Unstructured pointers (tensor>) + converter.addConversion([context](RankedTensorType tensorType, + SmallVectorImpl &types) + -> std::optional { + // Important note: + // We only care about tensor of index / int (in addition to pointer type) + // because only values of int and index type can potentially be part of a + // pointer arithmetic sequence. + if (!isa(tensorType.getElementType()) && + !tensorType.getElementType().isIntOrIndex()) { + // There's a subtle difference between returning failure() and + // std::nullopt. From the documentation: + // + // If std::nullopt is returned, the converter is allowed to try another + // conversion function to perform the conversion. + // + // Say we have type tensor<4x256xbf16> which is a RankedTensorType. Even + // though this RankedTensorType matches the converter that handles the + // tuple conversion, we want to keep this type as is because the inner + // type isn't a pointer. + // + // By returning failure(), the TypeConverters will stop trying the + // remaining converters. In our case, the last type converter which + // simply returns the same type is skipped. And because the conversion + // for this type has failed, the whole conversion process is also + // skipped. + // + // Relevant links to the implementation: + // + // https://github.com/llvm/llvm-project/blob/cb5dc1faa8b3702e0d03426ee5dfc5e1b903ec47/mlir/lib/Transforms/Utils/DialectConversion.cpp#L2958 + // https://github.com/llvm/llvm-project/blob/cb5dc1faa8b3702e0d03426ee5dfc5e1b903ec47/mlir/lib/Transforms/Utils/DialectConversion.cpp#L3033 + return std::nullopt; + } + types = + SmallVector{getStructuredStateTupleType(context, tensorType)}; + return success(); + }); + + // Case 2: Block pointers (!tt.ptr> or !tt.ptr) + converter.addConversion([context](triton::PointerType ptrType, + SmallVectorImpl &types) + -> std::optional { + types = SmallVector{getStructuredStateTupleType(context, ptrType)}; + return success(); + }); + + // Hooks to compute the correct materialization, "argument" and "source" + // materialization are used when we need to convert the tuple type back to + // the original triton pointer type. These are used when there are ops that + // still need to use the original pointer type. For instance, we convert the + // result of tt.addptr from tt.ptr type to a tuple, but the original ptr + // result is still being used by another tt.load or tt.store. + auto materialize = [](OpBuilder &builder, Type resultType, + ValueRange inputs, Location loc) { + return builder.create(loc, resultType, inputs) + .getResult(0); + }; + + converter.addSourceMaterialization(materialize); + + // Compute the target materialization, given a value with the pointer type, + // convert that value to a tuple type. + converter.addTargetMaterialization([](OpBuilder &builder, + TypeRange resultTypes, + ValueRange inputs, + Location loc) -> SmallVector { + return builder + .create(loc, resultTypes, inputs.front()) + ->getResults(); + }); + + ConversionTarget target(getContext()); + target.markUnknownOpDynamicallyLegal([](Operation *) { return true; }); + scf::populateSCFStructuralTypeConversionsAndLegality(converter, patterns, + target); + + if (failed(applyPartialConversion(getOperation(), target, + std::move(patterns)))) { + return failure(); + } + + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createCanonicalizerPass()); + if (failed(runPipeline(pm, getOperation()))) { + return failure(); + } + + return success(); + } + + LogicalResult decomposePointerTuple() { + auto moduleOp = getOperation(); + + auto context = &getContext(); + TypeConverter converter; + converter.addConversion([](Type type) { return type; }); + + // We are doing a 1->N type conversion here, where a pointer tuple type + // maps to a sequence of {pointer, offset_0, offset_1,..., stride_0, + // stride_1,...} + converter.addConversion( + [context](TupleType tupleType, SmallVectorImpl &types) + -> std::optional { + tupleType.getFlattenedTypes(types); + return success(); + }); + + // Hooks to compute the correct materialization, "argument" and "source" + // materialization are used when we need to convert a series of {pointer, + // offset_0, offset_1,..., stride_0, stride_1,...} type back to the "pointer + // tuple type". + // + // MLIR requires the materialization result to have exactly resultType. + // Keep a typed tuple bridge while SCF rewrites its regions/results. The + // pointer projection is folded below, after the offset/stride values have + // become explicit loop operands; returning inputs[0] here violates the + // conversion contract and loses the bridge before SCF finishes remapping. + auto materialize = [](OpBuilder &builder, Type resultType, + ValueRange inputs, + Location loc) -> Value { + if (inputs.size() == 1 && inputs.front().getType() == resultType) + return inputs.front(); + return builder.create(loc, resultType, inputs) + .getResult(0); + }; + converter.addSourceMaterialization(materialize); + + // For each value of "pointer tuple type" that gets decomposed into a + // sequence of {pointer, offset_0, offset_1,..., stride_0, stride_1,...}, + // create a `tts.get_structured_state` op that serves as a placeholder. + // The return values for this op will be used as the init-args for scf.for. + // At the end of pointer analysis, we will use the PtrState to create the + // correct offsets, strides, and remove these ops. + converter.addTargetMaterialization([](OpBuilder &builder, + TypeRange resultTypes, + ValueRange inputs, Location loc) + -> SmallVector { + if (inputs.size() != 1) + return {}; + auto bridge = inputs.front().getDefiningOp(); + if (!bridge || bridge.getInputs().empty()) + return {}; + // A source materialization can already contain the entire flattened + // state. Reuse it instead of reconstructing offsets from only its head. + if (llvm::equal(bridge.getInputs().getTypes(), resultTypes)) + return SmallVector(bridge.getInputs()); + if (bridge.getInputs().size() != 1 || resultTypes.empty() || + bridge.getInputs().front().getType() != resultTypes.front()) + return {}; + auto placeholder = builder.create( + loc, bridge.getInputs().front()); + assert(llvm::equal(placeholder.getResultTypes(), resultTypes)); + return SmallVector(placeholder.getResults()); + }); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + target.markUnknownOpDynamicallyLegal([](Operation *) { return true; }); + scf::populateSCFStructuralTypeConversionsAndLegality(converter, patterns, + target); + if (failed(applyPartialConversion(getOperation(), target, + std::move(patterns)))) { + return failure(); + } + + // convertToTupleType introduced tuple -> original-value projections. The + // inverse bridge now packs a *sequence*, so generic cast reconciliation + // cannot infer that the projection is its first element. Fold only this + // typed pair; the remaining state values stay in the converted SCF args. + SmallVector projections; + moduleOp.walk([&](UnrealizedConversionCastOp castOp) { + if (castOp.getInputs().size() != 1 || castOp.getResults().size() != 1 || + !isa(castOp.getInputs().front().getType()) || + isa(castOp.getResult(0).getType())) + return; + auto pack = castOp.getInputs().front() + .getDefiningOp(); + if (pack && !pack.getInputs().empty() && + pack.getInputs().front().getType() == castOp.getResult(0).getType()) + projections.push_back(castOp); + }); + for (auto projection : projections) { + auto pack = projection.getInputs().front() + .getDefiningOp(); + projection.getResult(0).replaceAllUsesWith(pack.getInputs().front()); + projection.erase(); + if (pack->use_empty()) + pack.erase(); + } + + // Note: + // Be careful not to run canonicalization here, because the + // tts.get_structured_state ops created above are just placeholders and + // don't have any effects. Canonicalization will remove them altogether. + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(mlir::createReconcileUnrealizedCastsPass()); + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } + + return success(); + } + + // Prepass that inserts `tts.get_structured_state` ops. These ops are used as + // placeholders to make passing structured pointer state into scf.for loop's + // init args easier, especially with multiple levels of loops. + // + // Background: + // + // PtrAnalysis computes a PtrState for every operand (or triton value) + // involved in a sequence of pointer arithmetic; some examples include: triton + // pointer, offsets (which could be a tensor of indices or just a simple index + // value). + // + // If a triton value is updated and returned in a scf.for op, it means + // that we have to carry its offsets and strides in the scf.for's iterargs. + // + // Previously, we have to manually rewrite the loops to include the + // relevant information from a PtrState which was rather involved and + // error-prone; this was also hard to scale up to multiple level of loops + // because there are several book-keeping data structures that we have to + // maintain. + // + // With the introduction of the prepass that inserts + // `tts.get_structured_state`. The return values of these ops, which include a + // triton value with its original result type and its corresponding offsets + // and strides, will be used as "placeholders" into the scf.for's init-args. + // We leverage standard MLIR infrastructure 1->N conversion to perform this + // rewrite, which helps simplify the logic significantly. + // + // After PtrAnalysis finishes, the return values of these + // `tts.get_structured_state` ops will be remapped to the correct + // initialization of the value's offsets and strides through the value's + // computed PtrState. + // + // Implementation details: + // In essence, what we really want to do in the prepass is, for every value + // of triton-pointer-like type (tt.ptr or tensor>) and tensor of + // indices (tensor) which might be used in a sequence of pointer + // arithmetic, we want to create an op `tts.get_structured_state` that takes + // in the original triton value and returns a series of values: + // + // {triton_value, offset_0, offset_1, ..., stride_0, stride_1,...} + // + // Applying the above conversion will also mean that any structural ops such + // as scf.for and scf.yield that originally takes the triton pointer will + // then take {triton_value, offset_0, offset_1, ..., stride_0, stride_1,...}. + // + // The 1->N type conversion is a perfect fit for this transformation. + // Unfortunately, we cannot do this is one pass, because the current 1->N + // type conversion implementation for scf.for ops doesn't provide us with a + // way to detect that a type conversion is recursive. So a triton_value type + // that gets converted to a {triton_value, offset_0, offset_1, ..., stride_0, + // stride_1,...} will recursively trigger other conversions. + // + // To fix this issue, we have to first convert triton_value to + // tuple. + // Finally, we decompose these tuples into the desired sequence. + // + // Note that even though the type conversion happens for every integer tensor + // appearing in loops' iter-args, this conversion is reversible. If the + // integer tensor isn't used in a pointer arithmetic sequence, + // canonicalization will remove all the `tts.get_structured_state` ops and + // revert the IR back to its original form. + LogicalResult runTritonToStructuredPrepass() { + if (failed(convertToPointerTupleWithOffsetsAndStrides())) { + return failure(); + } + + return decomposePointerTuple(); + } + + void runOnOperation() override { + if (!skipPrepass && failed(runTritonToStructuredPrepass())) { + signalPassFailure(); + return; + } + + if (runPrepassOnly) { + return; + } + + auto moduleOp = getOperation(); + mlir::tts::PtrAnalysis ptrAnalysis; + ptrAnalysis.initializeMaybeStructuredArgs(moduleOp); + + if (failed(ptrAnalysis.rewriteOp(moduleOp, useUnsafeMask))) { + moduleOp->emitWarning("PtrAnalysis failed"); + } + + // Now that all the PtrStates have been populated, we can wire up the states + // with the tts.get_structured_state ops inserted in the prepass. + moduleOp.walk([&ptrAnalysis](tts::GetStructuredStateOp op) { + if (failed(ptrAnalysis.rewriteGetStructuredStateOp(op))) { + op.emitWarning("Rewriting GetStructuredStateOp failed."); + } + }); + } +}; +} // namespace + +std::unique_ptr> +triton::createWaferTritonToStructuredPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/CMakeLists.txt new file mode 100755 index 00000000..8d2fb62c --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/CMakeLists.txt @@ -0,0 +1,27 @@ +add_triton_library(TritonToStructuredIncubated + TritonToStructuredIncubatedPass.cpp + PtrAnalysis.cpp + CannonicalizerConverter.cpp + MemOpConverter.cpp + MaskAnalysis.cpp + + DEPENDS + TritonToStructuredIncubatedConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonAnalysis + MLIRTritonNPUUtils + MLIRSCFTransforms + MLIRLinalgTransforms + BiShengIRHIVMDialect +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.cpp new file mode 100755 index 00000000..f4e3c2db --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.cpp @@ -0,0 +1,605 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.h" + +#include +#include +#include + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Arith/Utils/Utils.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Location.h" +#include "mlir/IR/Matchers.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Value.h" +#include "mlir/Support/LLVM.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include "llvm/Support/Debug.h" + +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#include "incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h" +#include "incubated/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.h" +#include "incubated/Conversion/UtilsIncubated/InterleaveOptimization.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#define DEBUG_TYPE "triton-cannonicalizer-converter" + +namespace CannonicalizerConverter { +using namespace mlir; +using namespace triton; + +// Match and rewrite pattern for optimizing cmp.ne (select(cond, 1, 0), 0) -> +// cond This pattern transforms: +// %select = arith.select %cond, %true_val, %false_val +// %cmp = arith.cmpi ne, %select, %zero +// Where: +// - %true_val is a constant splat tensor of 1s +// - %false_val is a constant splat tensor of 0s +// - %zero is a constant splat tensor of 0s +// Into: +// %cond (directly replace the cmp with the select condition) +// +// This optimization is valid because: +// select(cond, 1, 0) != 0 +// is equivalent to: cond != 0 +// Since the result of select is either 1 or 0, the only way it's not equal to +// 0 is when it's 1, which happens exactly when cond is true. +// +// Example: +// Input IR: +// %39 = arith.cmpi slt, %15, %cst_14 : tensor<128xi32> +// %40 = arith.select %39, %cst_13, %cst_12 : tensor<128xi1>, +// tensor<128xi32> %41 = arith.cmpi ne, %40, %cst_12 : tensor<128xi32> +// Where cst_13 is constant dense<1> and cst_12 is constant dense<0> +// Output IR: +// %39 = arith.cmpi slt, %15, %cst_14 : tensor<128xi32> +LogicalResult CmpConverter::matchAndRewrite(arith::CmpIOp cmpOp, + PatternRewriter &rewriter) const { + // Only handle "not equal" comparison + auto cmpType = cmpOp.getPredicate(); + if (cmpType != arith::CmpIPredicate::ne) { + return failure(); + } + + Value rhs = cmpOp.getRhs(); + Value lhs = cmpOp.getLhs(); + + // 1. Check if RHS is a constant zero + APInt rhsValue; + if (!matchPattern(rhs, m_ConstantInt(&rhsValue))) { + return failure(); // RHS is not a constant + } + + if (!rhsValue.isZero()) { + return failure(); // RHS is not zero + } + + // 2. Check if LHS is defined by a select operation + auto selectOp = lhs.getDefiningOp(); + if (!selectOp) { + return failure(); + } + + // 3. Check if select's true and false values are constants + DenseElementsAttr trueAttr; + DenseElementsAttr falseAttr; + if (!matchPattern(selectOp.getTrueValue(), m_Constant(&trueAttr)) || + !matchPattern(selectOp.getFalseValue(), m_Constant(&falseAttr))) { + return failure(); // Either true or false value is not constant + } + + // 4. Check if true value is all 1s and false value is all 0s + if (!trueAttr.isSplat() || !trueAttr.getSplatValue().isOne() || + !falseAttr.isSplat() || !falseAttr.getSplatValue().isZero()) { + return failure(); + } + + // 5. Optimization matched, replace cmp with select's condition + rewriter.replaceOp(cmpOp, selectOp.getCondition()); + return success(); +} + +// Transform a for loop that uses pointer iteration arguments into one that uses +// integer offsets instead. This pattern handles the specific case where: +// 1. The loop has pointer iteration arguments of type like +// tensor<1024x!tt.ptr> +// 2. Each pointer is used in a load/store operation and then incremented by +// a constant offset via tt.addptr +// 3. The updated pointer (from addptr) is yielded back as the next iteration +// value +// +// The transformation converts: +// scf.for iter_args(%ptr = %base_ptr) { +// %val = tt.load %ptr +// tt.store %other_ptr, %val +// %new_ptr = tt.addptr %ptr, %offset +// scf.yield %new_ptr +// } +// +// Into: +// scf.for iter_args(%offset_int = 0) { +// %splat_offset = tt.splat %offset_int +// %current_ptr = tt.addptr %base_ptr, %splat_offset +// %val = tt.load %current_ptr +// tt.store %other_ptr, %val +// %new_offset = arith.addi %offset_int, %const_offset +// scf.yield %new_offset +// } +// +LogicalResult PromotePointerIterArgsPattern::matchAndRewrite( + scf::ForOp forOp, PatternRewriter &rewriter) const { + // 1. Check if the loop meets transformation conditions + if (failed(matchLoop(forOp))) { + return failure(); + } + + // 2. Collect pointer iteration arguments to be processed + auto pointerArgsInfo = collectPointerIterArgs(forOp); + if (pointerArgsInfo.empty()) { + return failure(); + } + + // 3. Create new iteration argument types and initial values + auto [newInitArgs, newIterArgTypes, indexMap] = + createNewIterArgs(forOp, pointerArgsInfo, rewriter); + + // 4. Create the new for loop + auto newForOp = + createNewForLoop(forOp, newInitArgs, newIterArgTypes, rewriter); + + // 5. Rewrite the loop body + if (failed(rewriteLoopBody(forOp, newForOp, pointerArgsInfo, indexMap, + rewriter))) { + return failure(); + } + + // 6. Replace original loop results + return replaceResults(forOp, newForOp, pointerArgsInfo, indexMap, rewriter); +} + +LogicalResult PromotePointerIterArgsPattern::matchLoop(scf::ForOp forOp) const { + auto lowerBound = forOp.getLowerBound(); + auto upperBound = forOp.getUpperBound(); + auto step = forOp.getStep(); + if (!matchPattern(lowerBound, m_Constant()) || + !matchPattern(upperBound, m_Constant()) || + !matchPattern(step, m_Constant())) { + return failure(); + } + return success(); +} + +SmallVector +PromotePointerIterArgsPattern::collectPointerIterArgs(scf::ForOp forOp) const { + SmallVector result; + auto &loopBody = *forOp.getBody(); + + for (auto [idx, iterArg] : llvm::enumerate(forOp.getRegionIterArgs())) { + if (isPointerIterArg(iterArg)) { + auto info = analyzePointerIterArg(iterArg, loopBody); + if (info.has_value()) { + info->oldIndex = static_cast(idx), + info->basePointer = forOp.getInitArgs()[idx], + result.push_back(info.value()); + } + } + } + return result; +} + +bool PromotePointerIterArgsPattern::isPointerIterArg(Value iterArg) const { + auto ptrType = dyn_cast(iterArg.getType()); + return ptrType && isa(ptrType.getElementType()); +} + +std::optional +PromotePointerIterArgsPattern::analyzePointerIterArg(Value iterArg, + Block &loopBody) const { + int memCount = + 0; // Count of memory operations (load/store) using this pointer + int addPtrCount = 0; // Count of addptr operations on this pointer + Value addPtrResult = nullptr; // Result of the addptr operation + Value offset = nullptr; // Offset value used in addptr + Value addPtrValue = nullptr; // The addptr operation result value + + for (auto &op : loopBody) { + TypeSwitch(&op) + .Case([&](auto memoryOp) { + // Check if this memory operation uses the pointer we're analyzing + if (memoryOp.getPtr() == iterArg) + ++memCount; + }) + .Case([&](auto addPtrOp) { + // Check if this addptr operation updates the pointer we're analyzing + if (addPtrOp.getPtr() == iterArg) { + ++addPtrCount; + addPtrResult = addPtrOp.getResult(); + offset = addPtrOp.getOffset(); + addPtrValue = addPtrOp.getResult(); + } + }) + .Default([](auto) {}); // Ignore other operations + } + + // Check the terminator to see if the addptr result is yielded + auto yieldOp = dyn_cast(loopBody.getTerminator()); + if (!yieldOp) + return std::nullopt; + + bool isYielded = false; + for (auto operand : yieldOp.getOperands()) { + if (operand == addPtrResult) { + isYielded = true; + break; + } + } + + // Pattern matched if: + // 1. Exactly one addptr operation on this pointer + // 2. At least one memory operation using this pointer + // 3. The addptr result is yielded + if (addPtrCount == 1 && memCount >= 1 && isYielded) { + return PointerArgInfo{ + .oldIndex = 0, + .basePointer = nullptr, // Will be set in collectPointerIterArgs + .offsetValue = offset, + .newIterArg = nullptr, // Will be set in createNewIterArgs + .addPtrValue = addPtrValue}; + } + return std::nullopt; +} + +std::tuple, SmallVector, DenseMap> +PromotePointerIterArgsPattern::createNewIterArgs( + scf::ForOp forOp, ArrayRef pointerArgs, + PatternRewriter &rewriter) const { + SmallVector newInitArgs; + SmallVector newIterArgTypes; + DenseMap indexMap; + + for (unsigned i = 0; i < forOp.getInitArgs().size(); ++i) { + if (isPointerArgIndex(pointerArgs, i)) { + // Replace pointer with integer offset (initialized to 0) + Value zero = rewriter.create(forOp.getLoc(), 0, 32); + newInitArgs.push_back(zero); + newIterArgTypes.push_back(rewriter.getIntegerType(32)); + } else { + // Preserve original argument unchanged + newInitArgs.push_back(forOp.getInitArgs()[i]); + newIterArgTypes.push_back(forOp.getInitArgs()[i].getType()); + } + + // Identity mapping: argument count and order unchanged, + // may change in future + indexMap[i] = i; + } + + return {newInitArgs, newIterArgTypes, indexMap}; +} + +scf::ForOp PromotePointerIterArgsPattern::createNewForLoop( + scf::ForOp forOp, ArrayRef newInitArgs, + ArrayRef newIterArgTypes, PatternRewriter &rewriter) const { + return rewriter.create(forOp.getLoc(), forOp.getLowerBound(), + forOp.getUpperBound(), forOp.getStep(), + newInitArgs); +} + +LogicalResult PromotePointerIterArgsPattern::rewriteLoopBody( + scf::ForOp oldForOp, scf::ForOp newForOp, + SmallVector &pointerArgs, + DenseMap &indexMap, PatternRewriter &rewriter) const { + Block &oldBody = *oldForOp.getBody(); + Block &newBody = *newForOp.getBody(); + + rewriter.setInsertionPointToStart(&newBody); + + // Create IR mapping that maps original values to their transformed + // equivalents + IRMapping mapping = + createIRMapping(oldForOp, newForOp, pointerArgs, indexMap, rewriter); + + // Clone instructions from original loop body, applying the mapping + return cloneInstructions(oldBody, newBody, pointerArgs, indexMap, mapping, + rewriter); +} + +IRMapping PromotePointerIterArgsPattern::createIRMapping( + scf::ForOp oldForOp, scf::ForOp newForOp, + SmallVector &pointerArgs, + DenseMap &indexMap, PatternRewriter &rewriter) const { + IRMapping mapping; + mapping.map(oldForOp.getInductionVar(), newForOp.getInductionVar()); + + // Process iteration arguments + for (unsigned i = 0; i < oldForOp.getRegionIterArgs().size(); ++i) { + Value oldIterArg = oldForOp.getRegionIterArgs()[i]; + Value newIterArg = newForOp.getRegionIterArgs()[indexMap[i]]; + + if (isPointerArgIndex(pointerArgs, i)) { + // Update the PointerArgInfo with the new integer iteration argument + for (auto &info : pointerArgs) { + if (info.oldIndex == i) { + info.newIterArg = newIterArg; + break; + } + } + + // Map original pointer argument to a reconstructed pointer + mapping.map(oldIterArg, + rebuildPointer(oldForOp, pointerArgs, i, rewriter)); + } else { + // Direct mapping for non-pointer arguments + mapping.map(oldIterArg, newIterArg); + } + } + + return mapping; +} + +bool PromotePointerIterArgsPattern::isPointerArgIndex( + ArrayRef pointerArgs, unsigned idx) const { + for (auto &info : pointerArgs) { + if (info.oldIndex == idx) + return true; + } + return false; +} + +Value PromotePointerIterArgsPattern::rebuildPointer( + scf::ForOp forOp, ArrayRef pointerArgs, unsigned idx, + PatternRewriter &rewriter) const { + const PointerArgInfo *info = nullptr; + for (auto &argInfo : pointerArgs) { + if (argInfo.oldIndex == idx) { + info = &argInfo; + break; + } + } + if (!info) + return nullptr; + + // Create splat operation to broadcast integer offset to tensor shape + auto baseType = info->basePointer.getType(); + Value splatOffset = nullptr; + if (auto rankedType = dyn_cast(baseType)) { + // Get the shape of the original tensor + auto shape = rankedType.getShape(); + + splatOffset = rewriter.create( + forOp.getLoc(), RankedTensorType::get(shape, rewriter.getI32Type()), + info->newIterArg); + } else { + return nullptr; + } + + // Create addptr operation: base pointer + splatted offset + return rewriter.create(forOp.getLoc(), + info->basePointer.getType(), + info->basePointer, splatOffset); +} + +LogicalResult PromotePointerIterArgsPattern::cloneInstructions( + Block &oldBody, Block &newBody, ArrayRef pointerArgs, + DenseMap &indexMap, IRMapping &mapping, + PatternRewriter &rewriter) const { + // Collect all operations from the old loop body except the terminator + SmallVector toClone; + for (auto &op : oldBody.without_terminator()) { + toClone.push_back(&op); + } + + // Build a set of addptr operations to skip (those that update pointer + // iteration arguments) + DenseSet addPtrOpsToSkip; + for (const auto &info : pointerArgs) { + if (info.addPtrValue) { + addPtrOpsToSkip.insert(info.addPtrValue); + } + } + + // Clone all operations except the skipped addptr operations + for (auto *op : toClone) { + // Only skip addptr operations that are updating pointer iteration arguments + if (auto addPtrOp = dyn_cast(op)) { + if (addPtrOpsToSkip.contains(addPtrOp.getResult())) { + continue; + } + } + rewriter.clone(*op, mapping); + } + + // Handle the yield terminator separately + auto yieldOp = dyn_cast(oldBody.getTerminator()); + if (!yieldOp) { + return failure(); + } + + return cloneYieldOp(yieldOp, pointerArgs, indexMap, mapping, rewriter); +} + +LogicalResult PromotePointerIterArgsPattern::cloneYieldOp( + scf::YieldOp yieldOp, ArrayRef pointerArgs, + DenseMap &indexMap, IRMapping &mapping, + PatternRewriter &rewriter) const { + SmallVector newOperands; + // Process each operand of the original yield operation + for (unsigned i = 0; i < yieldOp.getNumOperands(); ++i) { + if (isPointerArgIndex(pointerArgs, i)) { + // For pointer arguments being promoted: create integer addition + Value intResult = createIntegerAdd(i, pointerArgs, indexMap, rewriter); + newOperands.push_back(intResult); + } else { + // For other arguments: use the value from the IR mapping + newOperands.push_back(mapping.lookupOrDefault(yieldOp.getOperand(i))); + } + } + + // Validate that all new operands are non-null + for (auto v : newOperands) { + if (!v) { + return failure(); + } + } + + // Create the new yield operation in the transformed loop + rewriter.create(yieldOp.getLoc(), newOperands); + return success(); +} + +Value PromotePointerIterArgsPattern::createIntegerAdd( + unsigned idx, ArrayRef pointerArgs, + DenseMap &indexMap, PatternRewriter &rewriter) const { + const PointerArgInfo *info = nullptr; + for (auto &argInfo : pointerArgs) { + if (argInfo.oldIndex == idx) { + info = &argInfo; + break; + } + } + if (!info) + return nullptr; + + // Try to extract constant offset value + Attribute offsetAttr; + if (matchPattern(info->offsetValue, m_Constant(&offsetAttr))) { + Location loc = info->offsetValue.getLoc(); + + // Case 1: Integer attribute (scalar constant) + if (auto intAttr = dyn_cast(offsetAttr)) { + Value constOffset = + rewriter.create(loc, intAttr.getInt(), 32); + return rewriter.create(loc, info->newIterArg, constOffset); + } + + // Case 2: DenseElementsAttr (tensor constant) + if (auto denseAttr = dyn_cast(offsetAttr)) { + // Check if it's a splat (all elements are the same) + if (denseAttr.isSplat()) { + // For integer-type DenseElementsAttr + if (denseAttr.getElementType().isInteger(32)) { + auto splatValue = denseAttr.getSplatValue(); + Value constOffset = rewriter.create( + loc, splatValue.getZExtValue(), 32); + return rewriter.create(loc, info->newIterArg, + constOffset); + } + } else { + // If not a splat, but has only one element, we can still handle it + if (denseAttr.getNumElements() == 1) { + auto firstElement = *denseAttr.getValues().begin(); + Value constOffset = rewriter.create( + loc, firstElement.getZExtValue(), 32); + return rewriter.create(loc, info->newIterArg, + constOffset); + } + } + } + } + + // Return nullptr if offset is not a constant (pattern only handles constant + // offsets) + return nullptr; +} + +LogicalResult PromotePointerIterArgsPattern::replaceResults( + scf::ForOp oldForOp, scf::ForOp newForOp, + ArrayRef pointerArgs, + DenseMap &indexMap, PatternRewriter &rewriter) const { + SmallVector newResults; + + for (unsigned i = 0; i < oldForOp.getNumResults(); ++i) { + if (isPointerArgIndex(pointerArgs, i)) { + Value ptrResult = reconstructPointer( + oldForOp, i, newForOp.getResult(indexMap[i]), pointerArgs, rewriter); + newResults.push_back(ptrResult); + } else { + newResults.push_back(newForOp.getResult(indexMap[i])); + } + } + + for (auto v : newResults) { + if (!v) { + return failure(); + } + } + rewriter.replaceOp(oldForOp, newResults); + return success(); +} + +Value PromotePointerIterArgsPattern::reconstructPointer( + scf::ForOp forOp, unsigned idx, Value intResult, + ArrayRef pointerArgs, PatternRewriter &rewriter) const { + const PointerArgInfo *info = nullptr; + for (auto &argInfo : pointerArgs) { + if (argInfo.oldIndex == idx) { + info = &argInfo; + break; + } + } + if (!info) + return nullptr; + + // Create splat operation to broadcast integer result to tensor shape + auto baseType = info->basePointer.getType(); + Value splatOffset = nullptr; + if (auto rankedType = dyn_cast(baseType)) { + // Get the shape of the original tensor + auto shape = rankedType.getShape(); + + splatOffset = rewriter.create( + forOp.getLoc(), RankedTensorType::get(shape, rewriter.getI32Type()), + intResult); + } else { + return nullptr; + } + + // Create a tensor with the same shape, where all elements are the integer + // result + return rewriter.create(forOp.getLoc(), + info->basePointer.getType(), + info->basePointer, splatOffset); +} +} // namespace CannonicalizerConverter diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/MaskAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/MaskAnalysis.cpp new file mode 100755 index 00000000..43d54636 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/MaskAnalysis.cpp @@ -0,0 +1,965 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToStructuredIncubated/MaskAnalysis.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/LogicalResult.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/IRMapping.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/IR/Visitors.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Support/LogicalResult.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#define DEBUG_TYPE "triton-to-structured-mask-analysis" + +namespace TritonToStructuredIncubated { +using namespace mlir; +using namespace triton; + +bool dimInfo::setType(arith::CmpIPredicate Type) { + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Setting compare type for dimIndex " << dimIndex << "\n"; + llvm::dbgs() << "Type: " << Type << "\n"; + llvm::dbgs() << "----------------------------------------------\n"; + }); + + switch (Type) { + case arith::CmpIPredicate::slt: + this->currentType = dimInfo::CompareType::slt; + break; + case arith::CmpIPredicate::ult: + this->currentType = dimInfo::CompareType::ult; + break; + case arith::CmpIPredicate::sge: + this->currentType = dimInfo::CompareType::sge; + break; + case arith::CmpIPredicate::uge: + this->currentType = dimInfo::CompareType::uge; + break; + default: + return false; + } + return true; +} + +bool dimInfo::compareTypeIsLess() const { + return this->currentType == dimInfo::CompareType::slt || + this->currentType == dimInfo::CompareType::ult; +} + +void dimInfo::dump() const { + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "MaskDimInfo: \n"; + llvm::dbgs() << "offset = " << offset << "\n"; + llvm::dbgs() << "shape = " << shape << "\n"; + llvm::dbgs() << "rhs = " << rhs << "\n"; + llvm::dbgs() << "isLessMode = " << compareTypeIsLess() << "\n"; + llvm::dbgs() << "hasBroadCast = " << hasBroadCast << "\n"; + llvm::dbgs() << "----------------------------------------------\n"; +}; + +void MaskState::dump() const { + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "MaskState :\n"; + llvm::dbgs() << "scalar = " << scalar << "\n"; + llvm::dbgs() << "stateInfo.size = " << stateInfo.size() << "\n"; + for (auto info : stateInfo) + info.dump(); + llvm::dbgs() << "----------------------------------------------\n"; +}; + +LogicalResult MaskState::parse(Value operand, const Location loc, + OpBuilder &builder) { + if (isa(operand.getType())) { + return this->parseIntScalar(operand, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseConstant(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseAdd(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseAnd(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseCmp(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseMakeRange(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseBroadcast(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseSplat(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseExpandDims(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseExtSI(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseRem(op, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return this->parseDiv(op, loc, builder); + } + LLVM_DEBUG({ + InFlightDiagnostic diag = emitWarning(loc) + << "MaskAnalysis: compare operand produced by an " + "unsupported operation\n"; + }); + return failure(); +} + +LogicalResult MaskState::parseConstant(arith::ConstantOp constOp, + const Location loc, OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + constOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting constant"); + }); + return failure(); + } + if (isa(constOp.getValue())) { + auto attr = cast(constOp.getValue()); + auto elementType = attr.getElementType(); + if (!attr.isSplat() || !isa(elementType)) { + LLVM_DEBUG({ + constOp.emitError("MaskAnalysis: only support splat integer constant"); + }); + return failure(); + } + auto values = attr.getValues(); + auto value = values[0].getValue(); + auto constAttr = builder.getIndexAttr(value.getSExtValue()); + auto op = arith::ConstantOp::materialize(builder, constAttr, + builder.getIndexType(), loc); + this->scalar = op.getValue(); + } else { + auto value = cast(constOp.getValue()).getInt(); + this->scalar = builder.getIndexAttr(value); + } + return success(); +} + +LogicalResult MaskState::parseIntScalar(Value scalar, const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitError(loc) << "MaskAnalysis: MaskState should be empty when " + "visiting integer scalar"; + }); + return failure(); + } + auto castOp = + builder.create(loc, builder.getIndexType(), scalar); + this->scalar = castOp.getResult(); + return success(); +} + +LogicalResult MaskState::parseMakeRange(triton::MakeRangeOp rangeOp, + const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + rangeOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting make_range"); + }); + return failure(); + } + + auto shape = cast(rangeOp.getType()).getShape(); + auto start = rangeOp.getStart(); + auto end = rangeOp.getEnd(); + auto stride = (end - start + shape[0] - 1) / shape[0]; + + if (stride != 1) { + LLVM_DEBUG( + { + InFlightDiagnostic diag = + emitError(loc) + << "stride must be 1 for make_range whose result is used " + "as load or store masks"; + }); + return failure(); + } + + stateInfo.emplace_back(builder.getIndexAttr(start), + builder.getIndexAttr(shape[0])); + return success(); +} + +LogicalResult MaskState::parseExtSI(arith::ExtSIOp op, const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + op->emitError( + "MaskAnalysis: MaskState should be empty when visiting extsi"); + }); + return failure(); + } + return parse(op.getIn(), loc, builder); +} + +LogicalResult MaskState::parseSplat(triton::SplatOp splatOp, const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + splatOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting splat"); + }); + return failure(); + } + + auto src = splatOp.getSrc(); + auto dst = splatOp.getResult(); + auto dstShape = cast(dst.getType()).getShape(); + + if (!isa(src.getType())) { + LLVM_DEBUG( + { + splatOp.emitError() + << "splat source must be an integer scalar for load/store masks"; + }); + return failure(); + } + + if (failed(this->parse(src, loc, builder))) + return failure(); + + auto zeroAttr = builder.getIndexAttr(0); + for (auto [i, shape] : llvm::enumerate(dstShape)) { + auto shapeAttr = builder.getIndexAttr(shape); + stateInfo.emplace_back(zeroAttr, shapeAttr, i, true); + } + + return success(); +} + +LogicalResult MaskState::parseExpandDims(triton::ExpandDimsOp expandDimsOp, + const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + expandDimsOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting expand_dims"); + }); + return failure(); + } + + auto zeroAttr = builder.getIndexAttr(0); + auto defaultShape = builder.getIndexAttr(1); + + if (failed(this->parse(expandDimsOp.getSrc(), loc, builder))) + return failure(); + + auto dstShape = + cast(expandDimsOp.getResult().getType()).getShape(); + auto axis = expandDimsOp.getAxis(); + if (dstShape[axis] != 1) { + LLVM_DEBUG({ + expandDimsOp.emitError( + "MaskAnalysis: unexpected dimension size in expand_dims"); + }); + return failure(); + } + + size_t insertPos = 0; + for (auto &info : stateInfo) { + if (info.dimIndex >= axis) + ++info.dimIndex; + if (info.dimIndex < axis) + ++insertPos; + } + + dimInfo insertInfo(zeroAttr, defaultShape, axis, true); + stateInfo.insert(stateInfo.begin() + insertPos, insertInfo); + + return success(); +} + +LogicalResult MaskState::parseAdd(arith::AddIOp addOp, const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + addOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting add"); + }); + return failure(); + } + + MaskState lhsState; + if (failed(lhsState.parse(addOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(addOp.getRhs(), loc, builder))) + return failure(); + + return this->addStates(lhsState, rhsState, loc, builder); +} + +LogicalResult MaskState::addStates(const MaskState &lhsState, + const MaskState &rhsState, Location loc, + OpBuilder &builder) { + if (lhsState.scalar && rhsState.scalar) { + LLVM_DEBUG( + { + InFlightDiagnostic diag = + emitWarning(loc) + << "Unexpected case where both lhs and rhs are scalars"; + }); + return failure(); + } + + if (!lhsState.scalar && !rhsState.scalar) { + LLVM_DEBUG( + { + InFlightDiagnostic diag = + emitWarning(loc) + << "Unsupported scenario where neither lhs nor rhs is a scalar"; + }); + return failure(); + } + + if (lhsState.scalar) + return addStateScalar(rhsState, lhsState.scalar, loc, builder); + else + return addStateScalar(lhsState, rhsState.scalar, loc, builder); +} + +LogicalResult MaskState::addStateScalar(const MaskState &state, + const OpFoldResult scalar, Location loc, + OpBuilder &builder) { + for (auto info : state.stateInfo) { + info.offset = addOpFoldResult(info.offset, scalar, loc, builder); + this->stateInfo.emplace_back(info); + } + return success(); +} + +LogicalResult MaskState::parseBroadcast(triton::BroadcastOp broadcastOp, + const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + broadcastOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting broadcast"); + }); + return failure(); + } + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Parsing BROADCAST operation: " << broadcastOp << "\n"; + }); + + auto src = broadcastOp.getSrc(); + auto dst = broadcastOp.getResult(); + if (!isa(dst.getType())) { + LLVM_DEBUG({ + broadcastOp.emitError( + "MaskAnalysis: broadcast dst should be a shaped type"); + }); + return failure(); + } + + auto srcShape = cast(src.getType()).getShape(); + auto dstShape = cast(dst.getType()).getShape(); + if (srcShape.size() != dstShape.size()) { + LLVM_DEBUG({ + broadcastOp.emitError( + "MaskAnalysis: broadcast src and dst should have the same rank"); + }); + return failure(); + } + + if (failed(parse(src, loc, builder))) + return failure(); + + LLVM_DEBUG({ + llvm::dbgs() << "Before BROADCAST MaskState: \n"; + this->dump(); + }); + + for (size_t i = 0; i < srcShape.size(); ++i) { + if (srcShape[i] == dstShape[i]) + continue; + else if (srcShape[i] < dstShape[i] && srcShape[i] == 1) { + for (auto &info : stateInfo) { + if (info.dimIndex != i) + continue; + info.shape = builder.getIndexAttr(dstShape[i]); + info.hasBroadCast = true; + } + } else { + LLVM_DEBUG({ + broadcastOp.emitError( + "MaskAnalysis: unexpected dimensions used in broadcast"); + }); + return failure(); + } + } + + LLVM_DEBUG({ + llvm::dbgs() << "After BROADCAST MaskState: \n"; + this->dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + + return success(); +} + +LogicalResult MaskState::parseCmp(arith::CmpIOp cmpOp, const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + cmpOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting cmpi"); + }); + return failure(); + } + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Parsing CMP operation: " << cmpOp << "\n"; + }); + + MaskState lhsState; + if (failed(lhsState.parse(cmpOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(cmpOp.getRhs(), loc, builder))) + return failure(); + + if (isa(cmpOp.getLhs().getType())) { + LLVM_DEBUG( + { cmpOp.emitRemark("MaskAnalysis: Unsupported cmpi scenario"); }); + return failure(); + } + + // lhs must be a Value and rhs must be scalar + if (lhsState.scalar || !rhsState.scalar) { + LLVM_DEBUG( + { cmpOp.emitWarning("MaskAnalysis: Unsupported cmpi scenario"); }); + return failure(); + } + + LLVM_DEBUG({ + llvm::dbgs() << "LHS MaskState: \n"; + lhsState.dump(); + llvm::dbgs() << "RHS MaskState: \n"; + rhsState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + + // In the case where the values we are loading are entirely masked off like + // the following: + // + // ---|-------|-----------| + // ^ ^ ^ + // scalar start end + // + // newEnd = min(end, scalar) = scalar + // Now scalar < start, so simply doing dim = newEnd - start is incorrect. + // + // The correct formula is to optionally move `newDim` back to `start` using + // max(newEnd, start). + auto cmpType = cmpOp.getPredicate(); + for (auto &info : lhsState.stateInfo) { + if (info.hasBroadCast) + continue; + if (!info.setType(cmpType)) { + LLVM_DEBUG({ cmpOp.emitWarning("MaskAnalysis: Unsupported cmpi type"); }); + return failure(); + } + info.rhs = rhsState.scalar; + } + this->stateInfo = lhsState.stateInfo; + + LLVM_DEBUG({ + llvm::dbgs() << "After CMP MaskState: \n"; + this->dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + return success(); +} + +LogicalResult MaskState::parseRem(arith::RemSIOp remOp, const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + remOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting REMSI"); + }); + return failure(); + } + + MaskState lhsState; + if (failed(lhsState.parse(remOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(remOp.getRhs(), loc, builder))) + return failure(); + + if (lhsState.scalar || !rhsState.scalar) { + LLVM_DEBUG( + { remOp.emitRemark("MaskAnalysis: Unsupported REMSI scenario"); }); + return failure(); + } + + auto divisorAttr = rhsState.scalar; + + if (!getIntAttr(divisorAttr).has_value()) { + LLVM_DEBUG({ + remOp.emitError("MaskAnalysis: do not support dynamic divisor in REMSI."); + }); + return failure(); + } + + SmallVector newStateInfo; + auto zeroAttr = builder.getIndexAttr(0); + for (auto info : lhsState.stateInfo) { + if (info.hasBroadCast) { + newStateInfo.emplace_back(info); + continue; + } + if (!isMultiple(divisorAttr, info.shape) && + !isMultiple(info.shape, divisorAttr)) { + LLVM_DEBUG({ + remOp.emitError( + "MaskAnalysis: do not support dynamic stride before REMSI."); + }); + return failure(); + } + + auto contiguousSize = + minOpFoldResult(divisorAttr, info.shape, loc, builder); + auto nonContiguousSize = + divOpFoldResult(info.shape, contiguousSize, loc, builder); + + auto staticNonContiguousSize = getIntAttr(nonContiguousSize); + if (!staticNonContiguousSize.has_value()) { + LLVM_DEBUG({ + remOp.emitError( + "MaskAnalysis: do not support dynamic size before REMSI."); + }); + return failure(); + } + + if (staticNonContiguousSize.value() != 0) + newStateInfo.emplace_back(zeroAttr, nonContiguousSize, info.dimIndex, + true); + + auto newOffset = remOpFoldResult(info.offset, divisorAttr, loc, builder); + newStateInfo.emplace_back(newOffset, nonContiguousSize, info.dimIndex); + } + + this->stateInfo = newStateInfo; + + return success(); +} + +LogicalResult MaskState::parseDiv(arith::DivSIOp divOp, const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + divOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting DIVSI"); + }); + return failure(); + } + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Parsing DIV operation: " << divOp << "\n"; + }); + + MaskState lhsState; + if (failed(lhsState.parse(divOp.getLhs(), loc, builder))) + return failure(); + + MaskState rhsState; + if (failed(rhsState.parse(divOp.getRhs(), loc, builder))) + return failure(); + + LLVM_DEBUG({ + llvm::dbgs() << "LHS MaskState: \n"; + lhsState.dump(); + llvm::dbgs() << "RHS MaskState: \n"; + rhsState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + + if (lhsState.scalar || !rhsState.scalar) { + LLVM_DEBUG( + { divOp.emitRemark("MaskAnalysis: Unsupported DIVSI scenario"); }); + return failure(); + } + + auto divisorAttr = rhsState.scalar; + + if (!getIntAttr(divisorAttr).has_value()) { + LLVM_DEBUG({ + divOp.emitError("MaskAnalysis: do not support dynamix divisor in DIVSI."); + }); + return failure(); + } + + SmallVector newStateInfo; + auto zeroAttr = builder.getIndexAttr(0); + for (auto info : lhsState.stateInfo) { + if (info.hasBroadCast) { + newStateInfo.emplace_back(info); + continue; + } + if (!isMultiple(divisorAttr, info.shape) && + !isMultiple(info.shape, divisorAttr)) { + LLVM_DEBUG({ + divOp.emitError( + "MaskAnalysis: do not support dynamix stride before DIVSI."); + }); + return failure(); + } + + auto nonContiguousSize = + minOpFoldResult(divisorAttr, info.shape, loc, builder); + auto contiguousSize = + divOpFoldResult(info.shape, nonContiguousSize, loc, builder); + + auto staticContiguousSize = getIntAttr(contiguousSize); + if (!staticContiguousSize.has_value()) { + LLVM_DEBUG({ + divOp.emitError( + "MaskAnalysis: do not support dynamix size before DIVSI."); + }); + return failure(); + } + + if (staticContiguousSize.value() != 0) { + auto newOffset = divOpFoldResult(info.offset, divisorAttr, loc, builder); + newStateInfo.emplace_back(newOffset, contiguousSize, info.dimIndex); + } + + newStateInfo.emplace_back(zeroAttr, nonContiguousSize, info.dimIndex, true); + } + + this->stateInfo = newStateInfo; + + LLVM_DEBUG({ + llvm::dbgs() << "After DIV MaskState: \n"; + this->dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + return success(); +} + +LogicalResult MaskState::parseAnd(arith::AndIOp andOp, const Location loc, + OpBuilder &builder) { + if (!this->isEmpty()) { + LLVM_DEBUG({ + andOp.emitError( + "MaskAnalysis: MaskState should be empty when visiting and"); + }); + return failure(); + } + auto zeroAttr = builder.getIndexAttr(0); + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Parsing AND operation: " << andOp << "\n"; + }); + + MaskState lhsState; + if (failed(lhsState.parse(andOp.getLhs(), loc, builder))) + return failure(); + MaskState rhsState; + if (failed(rhsState.parse(andOp.getRhs(), loc, builder))) + return failure(); + + LLVM_DEBUG({ + llvm::dbgs() << "LHS MaskState: \n"; + lhsState.dump(); + llvm::dbgs() << "RHS MaskState: \n"; + rhsState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + + SmallVector newStateInfo; + auto lIt = lhsState.stateInfo.begin(); + auto rIt = rhsState.stateInfo.begin(); + + while (lIt != lhsState.stateInfo.end() && rIt != rhsState.stateInfo.end()) { + if (lIt->dimIndex != rIt->dimIndex) { + auto newInfo = lIt->dimIndex < rIt->dimIndex ? *lIt++ : *rIt++; + newStateInfo.emplace_back(newInfo); + continue; + } + + if (!isMultiple(lIt->shape, rIt->shape) && + !isMultiple(rIt->shape, lIt->shape)) { + LLVM_DEBUG({ + llvm::dbgs() << "LHS MaskState: \n"; + lhsState.dump(); + llvm::dbgs() << "RHS MaskState: \n"; + rhsState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + LLVM_DEBUG({ + andOp.emitError( + "MaskAnalysis: the add operation have incompatible sizes"); + }); + return failure(); + } + + dimInfo newInfo; + newInfo.dimIndex = lIt->dimIndex; + newInfo.shape = minOpFoldResult(lIt->shape, rIt->shape, loc, builder); + if ((isLess(newInfo.shape, lIt->shape) && !lIt->hasBroadCast || + isLess(newInfo.shape, rIt->shape) && !rIt->hasBroadCast)) { + LLVM_DEBUG({ + llvm::dbgs() << "LHS MaskState: \n"; + lhsState.dump(); + llvm::dbgs() << "RHS MaskState: \n"; + rhsState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + LLVM_DEBUG({ + andOp.emitError( + "MaskAnalysis: the add operation have incompatible sizes." + "Valid dimensions are split."); + }); + return failure(); + } + newInfo.currentType = + lIt->hasBroadCast ? rIt->currentType : lIt->currentType; + if (lIt->currentType != dimInfo::CompareType::deafaultType && + rIt->currentType != dimInfo::CompareType::deafaultType && + lIt->currentType != rIt->currentType) { + LLVM_DEBUG({ + andOp.emitError( + "MaskAnalysis: do not suppport different compare mode within" + "the same dimension."); + }); + return failure(); + } + + if (lIt->hasBroadCast) { + newInfo.offset = rIt->offset; + newInfo.rhs = rIt->rhs; + newInfo.hasBroadCast = rIt->hasBroadCast; + } else if (rIt->hasBroadCast) { + newInfo.offset = lIt->offset; + newInfo.rhs = lIt->rhs; + newInfo.hasBroadCast = lIt->hasBroadCast; + } else { + LLVM_DEBUG({ + andOp.emitError("MaskAnalysis: do not suppport " + "and in the same dimension."); + }); + return failure(); + } + + newStateInfo.emplace_back(newInfo); + + if (isEqual(lIt->shape, newInfo.shape)) + ++lIt; + else + lIt->shape = divOpFoldResult(lIt->shape, newInfo.shape, loc, builder); + if (isEqual(rIt->shape, newInfo.shape)) + ++rIt; + else + rIt->shape = divOpFoldResult(rIt->shape, newInfo.shape, loc, builder); + } + + while (rIt != rhsState.stateInfo.end()) { + newStateInfo.push_back(*rIt++); + } + while (lIt != lhsState.stateInfo.end()) { + newStateInfo.push_back(*lIt++); + } + + this->stateInfo = newStateInfo; + + LLVM_DEBUG({ + llvm::dbgs() << "After AND MaskState: \n"; + this->dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + return success(); +} + +LogicalResult MaskState::analysisMask(Value operand) { + auto op = operand.getDefiningOp(); + if (!op) { + return failure(); + } + auto loc = op->getLoc(); + OpBuilder builder(op); + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Analyzing mask: " << operand << "\n"; + }); + + if (this->parse(operand, loc, builder).failed() || this->isEmpty()) { + return failure(); + } + + LLVM_DEBUG({ + llvm::dbgs() << "Mask analysis result: \n"; + this->dump(); + llvm::dbgs() << "MaskAnalysis: successfully analyzed mask.\n"; + llvm::dbgs() << "----------------------------------------------\n"; + }); + return success(); +} + +Value MaskState::createNewMask(const Location loc, OpBuilder &builder) { + if (isEmpty()) + return nullptr; + + SmallVector shape; + for (auto info : stateInfo) { + auto staticShape = getIntAttr(info.shape); + if (!staticShape.has_value()) { + LLVM_DEBUG( + { + InFlightDiagnostic diag = + emitError(loc) + << "MaskAnalysis: dynamic shape is not supported in mask " + "generation\n"; + }); + return nullptr; + } + shape.emplace_back(staticShape.value()); + } + SmallVector cacheResults; + auto maskShape = RankedTensorType::get(shape, builder.getI1Type()); + + auto createRhsValue = [&](OpFoldResult rhs) -> Value { + if (auto rhsInt = getIntAttr(rhs)) { + auto rhsAttr = + builder.getI32IntegerAttr(static_cast(rhsInt.value())); + return builder.create(loc, rhsAttr).getResult(); + } + Value rhsValue = dyn_cast(rhs); + if (rhsValue.getType().isIndex()) { + rhsValue = builder.create(loc, builder.getI32Type(), + rhsValue); + } + return rhsValue; + }; + for (size_t i = 0; i < stateInfo.size(); ++i) { + auto info = stateInfo[i]; + if (info.hasBroadCast) { + continue; + } + auto indexI32RowType = + RankedTensorType::get(shape[i], builder.getI32Type()); + Value newMask = + builder.create(loc, indexI32RowType, 0, shape[i]); + + Value newOffset = createRhsValue(info.offset); + if (newOffset.getType().isIndex()) { + newOffset = builder.create(loc, builder.getI32Type(), + newOffset); + } + Value splatRhs = + builder.create(loc, indexI32RowType, newOffset); + newMask = builder.create(loc, newMask, splatRhs); + + auto rhsValue = createRhsValue(info.rhs); + auto splatOp = + builder.create(loc, indexI32RowType, rhsValue); + + if (info.currentType == dimInfo::CompareType::deafaultType) { + LLVM_DEBUG({ + InFlightDiagnostic diag = emitError(loc) + << "MaskAnalysis: cannot generate mask when " + "compare type is not set\n"; + }); + return nullptr; + } + auto cmpOp = builder.create(loc, + info.compareTypeIsLess() + ? arith::CmpIPredicate::slt + : arith::CmpIPredicate::sge, + newMask, splatOp.getResult()); + + auto expandValue = cmpOp.getResult(); + for (size_t j = 0; j < stateInfo.size(); ++j) { + if (j == i) + continue; + expandValue = builder.create(loc, expandValue, j); + } + + auto broadcastValue = + builder.create(loc, maskShape, expandValue); + + cacheResults.push_back(broadcastValue); + } + + if (cacheResults.empty()) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) << "MaskAnalysis: cannot generate mask when all " + "dimensions are broadcasted"; + }); + return nullptr; + } + newMask = cacheResults[0]; + for (size_t i = 1; i < cacheResults.size(); ++i) { + newMask = builder.create(loc, newMask, cacheResults[i]); + } + return newMask; +} + +} // namespace TritonToStructuredIncubated diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/MemOpConverter.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/MemOpConverter.cpp new file mode 100755 index 00000000..cb52e532 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/MemOpConverter.cpp @@ -0,0 +1,583 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToStructuredIncubated/MemOpConverter.h" + +#include +#include +#include + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Arith/Utils/Utils.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/Utils/ReshapeOpsUtils.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Location.h" +#include "mlir/IR/Matchers.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Value.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/FormatVariadic.h" +#include "llvm/Support/MathExtras.h" + +#include "llvm/Support/Debug.h" + +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#include "bishengir/Dialect/HIVM/Utils/Utils.h" +#include "incubated/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.h" +#include "incubated/Conversion/TritonToStructuredIncubated/MaskAnalysis.h" +#include "incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h" +#include "incubated/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.h" +#include "incubated/Conversion/UtilsIncubated/InterleaveOptimization.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#define DEBUG_TYPE "triton-mem-op-converter" + +namespace MemOpConverter { +using namespace mlir; +using namespace triton; +using namespace TritonToStructuredIncubated; + +LogicalResult LoadConverter::matchAndRewrite(triton::LoadOp op, + PatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto oldPtr = op.getPtr(); + auto oldMask = op.getMask(); + auto oldOther = op.getOther(); + + MemOpTransformer tf(MemOpTransformer::MemType::load, optimizeDynamicOffset, + compileOn91095); + + auto newPtr = tf.createNewPtr(oldPtr, loc, rewriter); + auto newMask = tf.createNewMask(oldMask, loc, rewriter); + auto newOther = tf.createNewOther(oldOther, loc, rewriter); + + if (!tf.ptrState.shouldLinearize) { + // no need to rewrite + return failure(); + } + + if (!newPtr) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) << "PtrAnalysis: failed to analyze load pointer."; + }); + return failure(); + } + + if (!enableMaskFallbackConversion && oldMask && !newMask) { + LLVM_DEBUG({ + InFlightDiagnostic diag = emitWarning(loc) + << "MaskAnalysis: failed to analyze load mask."; + }); + return failure(); + } + + auto loadOp = rewriter.create(loc, newPtr, newMask, newOther, + op.getCache(), op.getEvict(), + op.getIsVolatile()); + + // insert implicit ops + auto broadCastResult = + tf.materializeImplicitBroadcast(loadOp.getResult(), loc, rewriter); + auto permuteResult = + tf.materializeImplicitPermute(broadCastResult, loc, rewriter); + auto reshapeResult = + tf.materializeImplicitReshape(permuteResult, loc, rewriter); + auto selectResult = tf.materializeImplicitSelect(reshapeResult, oldMask, + oldOther, loc, rewriter); + + rewriter.replaceOp(op, selectResult); + return success(); +} + +LogicalResult StoreConverter::matchAndRewrite(triton::StoreOp op, + PatternRewriter &rewriter) const { + auto loc = op.getLoc(); + auto oldPtr = op.getPtr(); + auto oldMask = op.getMask(); + auto oldValue = op.getValue(); + + MemOpTransformer tf(MemOpTransformer::MemType::store, optimizeDynamicOffset, + compileOn91095); + + auto newPtr = tf.createNewPtr(oldPtr, loc, rewriter); + auto newMask = tf.createNewMask(oldMask, loc, rewriter); + + if (!tf.ptrState.shouldLinearize) { + // no need to rewrite + return failure(); + } + + if (!newPtr) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) << "PtrAnalysis: failed to analyze store pointer."; + }); + return failure(); + } + + if (!enableMaskFallbackConversion && oldMask && !newMask) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) << "MaskAnalysis: failed to analyze store mask."; + }); + return failure(); + } + + // insert sync_block_lock + auto lockVar = createSyncBlockLockVar(rewriter, loc); + if (oldMask && !newMask) { + rewriter.create(loc, lockVar); + } + + auto selectResult = + tf.materializeImplicitSelect(oldValue, oldMask, oldPtr, loc, rewriter); + auto reshapeResult = + tf.materializeImplicitReshape(selectResult, loc, rewriter); + auto permuteResult = + tf.materializeImplicitPermute(reshapeResult, loc, rewriter); + + auto storeOp = rewriter.create( + loc, newPtr, permuteResult, newMask, op.getBoundaryCheck(), op.getCache(), + op.getEvict()); + + // insert sync_block_unlock + if (oldMask && !newMask) { + rewriter.create(loc, lockVar); + } + rewriter.eraseOp(op); + return success(); +} + +Value MemOpTransformer::materializeImplicitBroadcast( + Value srcTensor, const Location loc, PatternRewriter &rewriter) { + SmallVector broadCastIndex; + SmallVector broadCastShape; + for (auto [i, info] : llvm::enumerate(ptrState.stateInfo)) { + if (isZero(info.stride)) { + broadCastIndex.emplace_back(i); + } + auto staticShape = getIntAttr(info.shape); + if (!staticShape.has_value()) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) + << "PtrAnalysis: dynamic shape is not supported in broadcast\n"; + }); + return srcTensor; + } + broadCastShape.emplace_back(staticShape.value()); + } + + if (broadCastIndex.empty()) + return srcTensor; + + // when load is a scalar, we need to use splat to broadcast + auto srcType = srcTensor.getType(); + if (srcType.isIntOrFloat()) { + auto broadCastType = RankedTensorType::get(broadCastShape, srcType); + auto splatOp = + rewriter.create(loc, broadCastType, srcTensor); + return splatOp.getResult(); + } + + auto init = rewriter.create( + loc, broadCastShape, + cast(srcTensor.getType()).getElementType()); + + auto broadCastOp = rewriter.create(loc, srcTensor, init, + broadCastIndex); + + return broadCastOp->getResult(0); +} + +Value MemOpTransformer::materializeImplicitReshape(Value srcTensor, + const Location loc, + PatternRewriter &rewriter) { + if (ptrState.sizes.size() == ptrState.stateInfo.size()) + return srcTensor; + SmallVector targetShape; + if (currentType == MemType::load) { + for (auto size : ptrState.sizes) { + auto staticShape = getIntAttr(size); + if (!staticShape.has_value()) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) + << "PtrAnalysis: dynamic shape is not supported in reshape\n"; + }); + return srcTensor; + } + targetShape.emplace_back(staticShape.value()); + } + } else { + for (auto info : ptrState.stateInfo) { + auto staticShape = getIntAttr(info.shape); + if (!staticShape.has_value()) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) + << "PtrAnalysis: dynamic shape is not supported in reshape\n"; + }); + return srcTensor; + } + targetShape.emplace_back(staticShape.value()); + } + } + + auto targetShapeAttr = DenseIntElementsAttr::get( + RankedTensorType::get({static_cast(targetShape.size())}, + rewriter.getI64Type()), + targetShape); + auto targetShapeType = RankedTensorType::get( + targetShape, cast(srcTensor.getType()).getElementType()); + auto targetShapeValue = + rewriter.create(loc, targetShapeAttr); + auto reshapeOp = rewriter.create( + loc, targetShapeType, srcTensor, targetShapeValue); + return reshapeOp.getResult(); +} + +Value MemOpTransformer::materializeImplicitSelect(Value srcTensor, Value mask, + Value other, + const Location loc, + PatternRewriter &rewriter) { + if (!mask || maskState.newMask) + return srcTensor; + auto TensorType = cast(srcTensor.getType()); + if (cast(mask.getType()).getShape() != TensorType.getShape()) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) << "MaskAnalysis: mask shape is not same as Value"; + }); + return srcTensor; + } + + if (currentType == MemType::store) { + auto loadOp = rewriter.create(loc, other, nullptr, nullptr, + ArrayRef(), nullptr); + other = loadOp.getResult(); + } + + if (!other) { + auto elementType = TensorType.getElementType(); + auto emptyOp = rewriter.create(loc, TensorType.getShape(), + elementType); + other = emptyOp.getResult(); + } + auto selectOp = rewriter.create(loc, mask, srcTensor, other); + return selectOp->getResult(0); +} + +Value MemOpTransformer::materializeImplicitPermute(Value srcTensor, + const Location loc, + PatternRewriter &rewriter) { + auto inTy = dyn_cast(srcTensor.getType()); + if (!inTy || !ptrState.isPermuted) + return srcTensor; + + auto inShape = inTy.getShape(); + SmallVector order(ptrState.permuteIds.size()); + for (size_t i = 0; i < ptrState.permuteIds.size(); ++i) { + if (currentType == MemType::load) { + order[ptrState.permuteIds[i]] = i; + } else { + order[i] = ptrState.permuteIds[i]; + } + } + SmallVector outShape(order.size()); + if (inShape.size() != outShape.size()) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) << "PtrAnalysis: incompatible shape for permute"; + }); + return srcTensor; + } + + for (size_t i = 0; i < outShape.size(); ++i) { + outShape[i] = inShape[order[i]]; + } + + auto outTy = RankedTensorType::get(outShape, inTy.getElementType()); + auto transOp = rewriter.create(loc, outTy, srcTensor, order); + return transOp.getResult(); +} + +Value MemOpTransformer::createNewPtr(Value oldPtr, const Location loc, + PatternRewriter &rewriter) { + TritonToStructuredIncubated::PtrAnalysis ptrAnalysis(optimizeDynamicOffset); + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "PtrAnalysis: analyzing load/store's ptr.\n"; + }); + + if (ptrAnalysis.visitOperand(oldPtr, ptrState, loc, rewriter).failed()) { + ptrState.shouldLinearize = false; + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) << "PtranAlysis: failed to analyze load/store ptr."; + }); + return oldPtr; + } + + // compute missing strides + // if stateinfo.shape is 1 and sizes[dimIndex] is 1, + // then the stride is the accumulated size of all dimensions on the right side + // ie. for shape [1, 128], sizes [1, 128], originally stride is [0, 1], + // after normalization, stride is [128, 1] + OpFoldResult maxStride = rewriter.getIndexAttr(1); + for (auto it = ptrState.stateInfo.rbegin(); it != ptrState.stateInfo.rend(); + ++it) { + if (TritonToStructuredIncubated::isOne(it->shape) && isZero(it->stride)) { + it->stride = maxStride; + } + maxStride = maxOpFoldResult(maxStride, it->stride, loc, rewriter); + } + + for (auto it = ptrState.stateInfo.rbegin(); it != ptrState.stateInfo.rend(); + ++it) { + if (isZero(it->stride)) { + ptrState.shouldLinearize = true; + } + } + + ptrState.analyzePermute(); + + if (ptrState.isPermuted) { + ptrState.shouldLinearize = true; + if (compileOn91095 && currentType == MemType::load) { + ptrState.shouldLinearize = false; + } + } + + return ptrState.createAddPtrOp(rewriter, loc); +} + +Value MemOpTransformer::createNewMask(Value oldMask, const Location loc, + PatternRewriter &rewriter) { + if (!oldMask) + return nullptr; + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "MaskAnalysis: analyzing load/store mask.\n"; + }); + + if (!oldMask || maskState.analysisMask(oldMask).failed()) { + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "MaskAnalysis: no mask or failed to analyze mask.\n"; + llvm::dbgs() << "oldMask:" << oldMask << "\n"; + maskState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + LLVM_DEBUG( + { + InFlightDiagnostic diag = + emitWarning(loc) + << "MaskAnalysis: failed to analyze load/store mask."; + }); + return nullptr; + } + + SmallVector newMaskInfo; + auto itPtr = ptrState.stateInfo.begin(); + auto itMask = maskState.stateInfo.begin(); + + // match and create new mask info + while (itPtr != ptrState.stateInfo.end() && + itMask != maskState.stateInfo.end()) { + // ptr'shape must be multiple of mask'shape or vice versa + if (!isMultiple(itMask->shape, itPtr->shape)) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) + << "MaskAnalysis: incompatible shapes between ptr and mask."; + llvm::dbgs() << "----------------------------------------------\n"; + ptrState.dump(); + llvm::dbgs() << "oldMask:" << oldMask << "\n"; + maskState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + return nullptr; + } + + auto newShape = minOpFoldResult(itMask->shape, itPtr->shape, loc, rewriter); + if (isLess(newShape, itMask->shape) && !itMask->hasBroadCast) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) + << "MaskAnalysis: the mask shape is incompatible with ptr shape."; + }); + return nullptr; + } + + TritonToStructuredIncubated::dimInfo newInfo( + itMask->offset, newShape, itMask->dimIndex, itMask->hasBroadCast, + itMask->currentType, itMask->rhs); + + if (!isZero(itPtr->stride)) { + newMaskInfo.emplace_back(newInfo); + } + + ++itPtr; + if (isEqual(itMask->shape, newShape)) { + ++itMask; + } + } + + if (itPtr != ptrState.stateInfo.end() || + itMask != maskState.stateInfo.end()) { + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "MaskAnalysis: failed to apply permute on mask.\n"; + ptrState.dump(); + llvm::dbgs() << "oldMask:" << oldMask << "\n"; + maskState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + LLVM_DEBUG({ + InFlightDiagnostic diag = emitWarning(loc) + << "MaskAnalysis: incompatible number of " + "dimensions between ptr and mask."; + }); + return nullptr; + } + + maskState.stateInfo = newMaskInfo; + + if (ptrState.isPermuted && !applyPermuteOnMask()) { + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "MaskAnalysis: failed to apply permute on mask.\n"; + ptrState.dump(); + llvm::dbgs() << "oldMask:" << oldMask << "\n"; + maskState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + InFlightDiagnostic diag = + emitWarning(loc) << "MaskAnalysis: failed to apply permute on mask."; + }); + return nullptr; + } + + LLVM_DEBUG({ + llvm::dbgs() << "After matching MaskState: \n"; + for (auto info : newMaskInfo) { + info.dump(); + } + llvm::dbgs() << "----------------------------------------------\n"; + }); + + auto newMask = maskState.createNewMask(loc, rewriter); + return newMask; +} + +Value MemOpTransformer::createNewOther(Value oldOther, const Location loc, + PatternRewriter &rewriter) { + if (!oldOther || !maskState.newMask) + return nullptr; + + auto ptrType = dyn_cast(ptrState.source.getType()); + if (!ptrType) { + LLVM_DEBUG( + { + InFlightDiagnostic diag = + emitWarning(loc) + << "PtrAnalysis: source of ptrState is not a pointer type."; + }); + return nullptr; + } + Type elementType = ptrType.getPointeeType(); + + SmallVector targetShape; + for (auto info : maskState.stateInfo) { + auto staticShape = getIntAttr(info.shape); + if (!staticShape.has_value()) { + LLVM_DEBUG({ + InFlightDiagnostic diag = + emitWarning(loc) + << "MaskAnalysis: dynamic shape is not supported in reshape\n"; + }); + return oldOther; + } + targetShape.emplace_back(staticShape.value()); + } + auto targetShapeAttr = DenseIntElementsAttr::get( + RankedTensorType::get({static_cast(targetShape.size())}, + rewriter.getI64Type()), + targetShape); + auto targetShapeType = RankedTensorType::get(targetShape, elementType); + auto targetShapeValue = + rewriter.create(loc, targetShapeAttr); + + auto reshapeOp = rewriter.create( + loc, targetShapeType, oldOther, targetShapeValue); + + return reshapeOp.getResult(); +} + +bool MemOpTransformer::applyPermuteOnMask() { + if (!ptrState.isPermuted || maskState.isEmpty()) { + return true; + } + if (ptrState.permuteIds.size() != maskState.stateInfo.size()) { + return false; + } + SmallVector newMaskInfo; + for (auto id : ptrState.permuteIds) { + newMaskInfo.push_back(maskState.stateInfo[id]); + } + maskState.stateInfo = newMaskInfo; + return true; +} + +hivm::CreateSyncBlockLockOp createSyncBlockLockVar(OpBuilder &builder, + Location loc) { + SmallVector shape = {1}; + auto elementType = builder.getI64Type(); + Type memrefType = MemRefType::get(shape, elementType); + + auto createSyncBlockLockOp = + builder.create(loc, memrefType, Value()); + return createSyncBlockLockOp; +} +} // namespace MemOpConverter diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/PtrAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/PtrAnalysis.cpp new file mode 100755 index 00000000..7b62ed61 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/PtrAnalysis.cpp @@ -0,0 +1,1361 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * Copyright (c) Microsoft Corporation. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/IR/Visitors.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Support/LogicalResult.h" + +#include "mlir/IR/IRMapping.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/LogicalResult.h" + +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#define DEBUG_TYPE "triton-to-structured-ptr-analysis" + +namespace TritonToStructuredIncubated { +using namespace mlir; +using namespace triton; + +bool isMultiple(const OpFoldResult ÷nd, const OpFoldResult &divisor) { + auto staticDividend = getIntAttr(dividend); + auto staticDivisor = getIntAttr(divisor); + if (!staticDividend || !staticDivisor) { + return false; + } + return staticDividend.value() % staticDivisor.value() == 0; +} + +bool isOne(const OpFoldResult ofr) { + auto staticOfr = getIntAttr(ofr); + return staticOfr.has_value() && staticOfr.value() == 1; +} + +bool isEqual(const OpFoldResult &ofr1, const OpFoldResult &ofr2) { + auto staticOfr1 = getIntAttr(ofr1); + auto staticOfr2 = getIntAttr(ofr2); + return staticOfr1 == staticOfr2; +} + +bool isLess(const OpFoldResult &ofr1, const OpFoldResult &ofr2) { + auto staticOfr1 = getIntAttr(ofr1); + auto staticOfr2 = getIntAttr(ofr2); + // When sorting for permute, the value determined at runtime + // is greater than the value determined at compile time. + if (!staticOfr1) { + return false; + } + if (!staticOfr2) { + return true; + } + return staticOfr1 < staticOfr2; +} + +bool isGreater(const OpFoldResult &ofr1, const OpFoldResult &ofr2) { + auto staticOfr1 = getIntAttr(ofr1); + auto staticOfr2 = getIntAttr(ofr2); + // When sorting for permute, the value determined at runtime + // is greater than the value determined at compile time. + if (!staticOfr2) { + return false; + } + if (!staticOfr1) { + return true; + } + return staticOfr1 > staticOfr2; +} + +void StateInfo::dump() const { + llvm::dbgs() << "StateInfo: \n"; + llvm::dbgs() << "dimIndex = " << dimIndex << "\n"; + llvm::dbgs() << "shape = " << shape << "\n"; + llvm::dbgs() << "stride = " << stride << "\n"; +} + +void PtrState::dump() const { + llvm::dbgs() << "PtrState: \n"; + llvm::dbgs() << "source:" << source << "\n"; + llvm::dbgs() << "scalar:" << offset << "\n"; + llvm::dbgs() << "size: ["; + for (auto size : sizes) + llvm::dbgs() << size << ", "; + llvm::dbgs() << "]\n"; + llvm::dbgs() << "shoueleLinearize: " << shouldLinearize << "\n"; + llvm::dbgs() << "stateInfo:\n"; + llvm::dbgs() << "\n"; + for (auto info : stateInfo) { + llvm::dbgs() << "-----------------------------------------\n"; + info.dump(); + llvm::dbgs() << "-----------------------------------------\n"; + } +} + +bool PtrState::isEmpty() const { + return (stateInfo.empty() && !source && !offset); +} + +bool PtrState::isScalar() const { + bool scalar = true; + for (auto info : stateInfo) { + auto staticStride = getIntAttr(info.stride); + if (!staticStride.has_value() || staticStride.value() != 0) + scalar = false; + } + return scalar && (offset || source); +} + +bool PtrState::hasSource() const { return source != nullptr; } + +bool PtrState::isSameSizeAs(const PtrState &x) const { + if (this->sizes.size() != x.sizes.size()) + return false; + + for (size_t i = 0; i < this->sizes.size(); ++i) { + if (this->sizes[i] != x.sizes[i]) + return false; + } + return true; +} + +void PtrState::updatePtrState(SmallVector stateInfo, + SmallVector sizes, Value source, + OpFoldResult offset, const Location loc, + OpBuilder &builder, bool shouldLinearize) { + this->stateInfo = stateInfo; + this->sizes = sizes; + this->source = source; + this->offset = offset; + this->shouldLinearize = shouldLinearize; + this->normalizeState(loc, builder); +} + +void PtrState::normalizeState(const Location loc, OpBuilder &builder) { + SmallVector newStateInfo; + auto zeroAttr = builder.getIndexAttr(0); + + // merge continuous zero strides + // e.g., stride [0, 0, 1] shape [4, 32, 16] --> stride [0, 1] shape [128, 16] + for (auto it = this->stateInfo.begin(); it != this->stateInfo.end(); ++it) { + while (it != this->stateInfo.end() && isZero(it->stride)) { + auto newShape = it->shape; + auto dimIndex = it->dimIndex; + for (++it; it != this->stateInfo.end() && isZero(it->stride) && + it->dimIndex == dimIndex; + ++it) { + newShape = mulOpFoldResult(newShape, it->shape, loc, builder); + } + newStateInfo.emplace_back(zeroAttr, newShape, dimIndex); + } + if (it == this->stateInfo.end()) + break; + // if the info is the only one with oriSize 1 in this dimension, skip it + // e.g., stride [0, 1] shape [1, 128] sizes [1, 128] do not delete the first + // info + if (isOne(it->shape) && !isOne(sizes[it->dimIndex])) + continue; + newStateInfo.emplace_back(*it); + } + + this->stateInfo = newStateInfo; +} + +LogicalResult PtrAnalysis::visitOperandAddptr(triton::AddPtrOp addptrOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + addptrOp.emitError( + "PtrAnalysis: PtrState should be empty when visiting addptr"); + }); + return failure(); + } + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Visit addptr operation: " << addptrOp << "\n"; + }); + + PtrState ptrState; + if (visitOperand(addptrOp.getPtr(), ptrState, addptrOp.getLoc(), builder) + .failed()) { + return failure(); + } + + PtrState offsetState; + if (visitOperand(addptrOp.getOffset(), offsetState, addptrOp.getLoc(), + builder) + .failed()) { + return failure(); + } + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Before visiting addptr operands: \n"; + llvm::dbgs() << "PtrState: \n"; + ptrState.dump(); + llvm::dbgs() << "OffsetState: \n"; + offsetState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + + if (!ptrState.source) { + LLVM_DEBUG({ + addptrOp.emitError("ptr field should provide source / base pointer"); + }); + return failure(); + } + return state.addState(ptrState, offsetState, addptrOp, builder); +} + +bool PtrAnalysis::operandIsScalar(Value operand) { + auto tensorType = dyn_cast(operand.getType()); + auto elementType = + tensorType ? tensorType.getElementType() : operand.getType(); + bool isScalar = true; + if (tensorType) { + for (size_t i = 0; i < tensorType.getRank() && isScalar; ++i) { + isScalar = tensorType.getDimSize(i) == 1; + } + } + return isScalar && + (isa(elementType) || isa(elementType)); +} + +LogicalResult PtrAnalysis::initStateByScalar(Value operand, PtrState &state, + const Location loc, + OpBuilder &builder) { + OpFoldResult newOffset; + SmallVector newSizes; + SmallVector newStateInfo; + if (isa(operand.getType())) { + OpBuilder::InsertionGuard guard(builder); + if (!isa(operand) && operand.getDefiningOp()) { + builder.setInsertionPointAfter(operand.getDefiningOp()); + } + auto castOp = builder.create( + loc, builder.getIndexType(), operand); + newOffset = castOp.getResult(); + } else if (isa(operand.getType())) { + newOffset = operand; + } else { + auto tensorType = dyn_cast(operand.getType()); + auto index = builder.create(loc, 0); + auto zeroAttr = builder.getIndexAttr(0); + auto oneAttr = builder.getIndexAttr(1); + SmallVector indices; + for (size_t i = 0; i < tensorType.getRank(); ++i) { + indices.push_back(index); + newSizes.emplace_back(oneAttr); + newStateInfo.emplace_back(zeroAttr, oneAttr, i); + } + auto extractedElement = + builder.create(loc, operand, indices); + newOffset = extractedElement.getResult(); + } + state.updatePtrState(newStateInfo, newSizes, nullptr, newOffset, loc, + builder); + return success(); +} + +LogicalResult PtrAnalysis::initStateByPointer(Value operand, PtrState &state, + const Location loc, + OpBuilder &builder) { + Value newSource; + SmallVector newSizes; + SmallVector newStateInfo; + + if (auto op = operand.getDefiningOp()) { + if (auto addPtrOp = dyn_cast(op)) { + return visitOperandAddptr(cast(op), state, loc, + builder); + } else if (auto bitCastOp = dyn_cast(op)) { + newSource = operand; + } else if (auto makeTensorOp = dyn_cast(op)) { + LLVM_DEBUG({ + op->emitWarning("Unexpected operand defining operation tts.make_tptr."); + }); + return failure(); + } else if (auto intToPtrOp = dyn_cast(op)) { + newSource = operand; + } else { + LLVM_DEBUG({ op->emitWarning("PtrAnalysis: Unexpected operand."); }); + return failure(); + } + } else { + newSource = operand; + } + auto newOffset = builder.getIndexAttr(0); + state.updatePtrState(newStateInfo, newSizes, newSource, newOffset, loc, + builder); + return success(); +} + +LogicalResult PtrState::mulState(const PtrState &lhsState, + const PtrState &rhsState, Operation *op, + OpBuilder &builder) { + auto loc = op->getLoc(); + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "mulState: " << op << "\n"; + }); + + if (!isEmpty()) { + LLVM_DEBUG({ + op->emitError("PtrAnalysis: PtrState should be empty when multiplying"); + }); + return failure(); + } + + // neither lhs nor rhs should have source, since multiplying base pointer + // does not make sense + if (lhsState.hasSource() || rhsState.hasSource()) { + LLVM_DEBUG({ + op->emitError("PtrAnalysis: do not support base inters in multiplying"); + }); + return failure(); + } else if (!lhsState.isScalar() && !rhsState.isScalar()) { + // do not support both tensors are effectively non-scalar + LLVM_DEBUG({ + op->emitError( + "PtrAnalysis: only support multiplying pointer states when one of " + "them represent a scalar"); + }); + return failure(); + } + + PtrState const *lhs = &lhsState; + PtrState const *rhs = &rhsState; + + if (!rhs->isScalar() && lhs->isScalar()) { + std::swap(lhs, rhs); + } + + SmallVector newStateInfo; + for (auto info : lhs->stateInfo) { + OpFoldResult newStride = + mulOpFoldResult(info.stride, rhs->offset, loc, builder); + newStateInfo.emplace_back(newStride, info.shape, info.dimIndex); + } + + auto newOffset = + mulOpFoldResult(lhsState.offset, rhsState.offset, loc, builder); + updatePtrState(newStateInfo, lhs->sizes, lhs->source, newOffset, loc, builder, + lhs->shouldLinearize); + + LLVM_DEBUG({ + llvm::dbgs() << "After mulState: \n"; + this->dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + + return success(); +} + +LogicalResult PtrState::subState(const PtrState &lhsState, + const PtrState &rhsState, Operation *op, + OpBuilder &builder) { + auto loc = op->getLoc(); + if (!isEmpty()) { + LLVM_DEBUG({ + op->emitError("PtrAnalysis: PtrState should be empty when subtracting"); + }); + return failure(); + } + + if (lhsState.hasSource() && rhsState.hasSource()) { + LLVM_DEBUG({ + op->emitError( + "PtrAnalysis: do not support both sides have base pointers in sub"); + }); + return failure(); + } + + if (!rhsState.isScalar()) { + LLVM_DEBUG({ + op->emitError("PtrAnalysis: only support sub when one of " + "them represents a scalar"); + }); + return failure(); + } + + auto newOffset = + subOpFoldResult(lhsState.offset, rhsState.offset, loc, builder); + updatePtrState(lhsState.stateInfo, lhsState.sizes, lhsState.source, newOffset, + loc, builder, lhsState.shouldLinearize); + + return success(); +} + +LogicalResult PtrState::addState(PtrState &lhsState, PtrState &rhsState, + Operation *op, OpBuilder &builder) { + auto loc = op->getLoc(); + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "addState: " << op << "\n"; + }); + + if (!isEmpty()) { + LLVM_DEBUG({ + op->emitError("PtrAnalysis: PtrState should be empty when adding"); + }); + return failure(); + } + if (!lhsState.isSameSizeAs(rhsState)) { + LLVM_DEBUG({ + op->emitError( + "PtrAnalysis: The original size of the addition should be the same"); + }); + return failure(); + } + + SmallVector newStateInfo; + auto lIt = lhsState.stateInfo.begin(); + auto rIt = rhsState.stateInfo.begin(); + while (lIt != lhsState.stateInfo.end() && rIt != rhsState.stateInfo.end()) { + if (lIt->dimIndex != rIt->dimIndex) { + auto newInfo = lIt->dimIndex < rIt->dimIndex ? *lIt++ : *rIt++; + newStateInfo.emplace_back(newInfo); + continue; + } + if (!isMultiple(lIt->shape, rIt->shape) && + !isMultiple(rIt->shape, lIt->shape)) { + LLVM_DEBUG({ + llvm::dbgs() << "LHS PtrState: \n"; + lhsState.dump(); + llvm::dbgs() << "RHS PtrState: \n"; + rhsState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + LLVM_DEBUG({ + op->emitError("PtrAnalysis: the add operation have incompatible sizes"); + }); + return failure(); + } + + auto newShape = minOpFoldResult(lIt->shape, rIt->shape, loc, builder); + if ((isLess(newShape, lIt->shape) && !isZero(lIt->stride) || + isLess(newShape, rIt->shape) && !isZero(rIt->stride))) { + LLVM_DEBUG({ + llvm::dbgs() << "LHS PtrState: \n"; + lhsState.dump(); + llvm::dbgs() << "RHS PtrState: \n"; + rhsState.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + LLVM_DEBUG({ + op->emitError("PtrAnalysis: the add operation have incompatible sizes." + "Valid dimensions are split."); + }); + return failure(); + } + + auto newStride = addOpFoldResult(lIt->stride, rIt->stride, loc, builder); + newStateInfo.emplace_back(newStride, newShape, lIt->dimIndex); + + if (isEqual(lIt->shape, newShape)) + ++lIt; + else + lIt->shape = divOpFoldResult(lIt->shape, newShape, loc, builder); + if (isEqual(rIt->shape, newShape)) + ++rIt; + else + rIt->shape = divOpFoldResult(rIt->shape, newShape, loc, builder); + } + + while (rIt != rhsState.stateInfo.end()) { + newStateInfo.push_back(*rIt++); + } + while (lIt != lhsState.stateInfo.end()) { + newStateInfo.push_back(*lIt++); + } + + auto newSource = source = lhsState.source ? lhsState.source : rhsState.source; + auto newOffset = + addOpFoldResult(lhsState.offset, rhsState.offset, loc, builder); + auto newShouldLinearize = + lhsState.shouldLinearize || rhsState.shouldLinearize; + auto newSizes = lhsState.sizes; + + updatePtrState(newStateInfo, newSizes, newSource, newOffset, loc, builder, + newShouldLinearize); + + LLVM_DEBUG({ + llvm::dbgs() << "After addState: \n"; + this->dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + + return success(); +} + +triton::AddPtrOp PtrState::createAddPtrOp(OpBuilder &builder, Location loc) { + SmallVector tensorSizes; + SmallVector tensorStrides; + + auto zeroAttr = builder.getIndexAttr(0); + auto oneAttr = builder.getIndexAttr(1); + + for (auto id : permuteIds) { + auto info = stateInfo[id]; + if (isZero(info.stride)) + continue; + tensorStrides.emplace_back(info.stride); + tensorSizes.emplace_back(getIntAttr(info.shape).value()); + } + + // load a scalar pointer + if (tensorSizes.empty()) { + Value offsetValue = materializeValue(builder, loc, offset); + if (offsetValue.getType().isIndex()) { + offsetValue = builder.create( + loc, builder.getI32Type(), offsetValue); + } + auto addptrOp = builder.create(loc, source.getType(), + source, offsetValue); + return addptrOp; + } + + SmallVector cachedRange; + auto ptrType = cast(source.getType()); + auto ptrTensorType = RankedTensorType::get({tensorSizes}, ptrType); + auto broadCastType = + RankedTensorType::get({tensorSizes}, builder.getI32Type()); + + if (tensorSizes.size() != tensorStrides.size()) { + LLVM_DEBUG( + { + InFlightDiagnostic diag = + emitError(loc) + << "PtrAnalysis: inconsistent tensor sizes and strides"; + }); + return nullptr; + } + for (size_t i = 0; i < tensorSizes.size(); ++i) { + // make range + auto indexI32RowType = + RankedTensorType::get({tensorSizes[i]}, builder.getI32Type()); + Value makeRangeOp = builder.create( + loc, indexI32RowType, 0, tensorSizes[i]); + + // multiply stride + Value strideValue = materializeValue(builder, loc, tensorStrides[i]); + if (strideValue.getType().isIndex()) { + strideValue = builder.create( + loc, builder.getI32Type(), strideValue); + } + Value splatStride = + builder.create(loc, indexI32RowType, strideValue); + auto rangeAfterMul = + builder.create(loc, makeRangeOp, splatStride); + + // reshape + Value expandedValue = rangeAfterMul; + for (size_t j = 0; j < tensorSizes.size(); ++j) { + if (j == i) + continue; + expandedValue = + builder.create(loc, expandedValue, j); + } + + // broadcast + auto broadcastValue = + builder.create(loc, broadCastType, expandedValue); + cachedRange.push_back(broadcastValue); + } + + // combine the cachedRange + Value rangeAfterCombine = cachedRange[0]; + for (size_t i = 1; i < cachedRange.size(); ++i) { + rangeAfterCombine = + builder.create(loc, rangeAfterCombine, cachedRange[i]); + } + + // addOffset + Value addValue = materializeValue(builder, loc, offset); + if (addValue.getType().isIndex()) { + addValue = + builder.create(loc, builder.getI32Type(), addValue); + } + Value splatOffset = + builder.create(loc, broadCastType, addValue); + auto rangeAfterAdd = + builder.create(loc, rangeAfterCombine, splatOffset); + + // addPtr + Value splatPtr = builder.create(loc, ptrTensorType, source); + auto addptrOp = builder.create(loc, ptrTensorType, splatPtr, + rangeAfterAdd); + return addptrOp; +} + +LogicalResult PtrAnalysis::visitOperandMul(arith::MulIOp mulOp, PtrState &state, + const Location loc, + OpBuilder &builder) { + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Visit Mul operation: " << mulOp << "\n"; + }); + + PtrState lhsState; + if (visitOperand(mulOp.getLhs(), lhsState, loc, builder).failed()) { + return failure(); + } + + PtrState rhsState; + if (visitOperand(mulOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + return state.mulState(lhsState, rhsState, mulOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandSub(arith::SubIOp subOp, PtrState &state, + const Location loc, + OpBuilder &builder) { + PtrState lhsState; + if (visitOperand(subOp.getLhs(), lhsState, loc, builder).failed()) { + return failure(); + } + + PtrState rhsState; + if (visitOperand(subOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + return state.subState(lhsState, rhsState, subOp, builder); +} + +LogicalResult PtrAnalysis::visitOperandMakeRange(triton::MakeRangeOp rangeOp, + PtrState &state, Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + rangeOp.emitError( + "PtrAnalysis: PtrState should be empty when visiting make_range"); + }); + return failure(); + } + + auto shape = cast(rangeOp.getType()).getShape(); + + auto start = rangeOp.getStart(); + auto end = rangeOp.getEnd(); + auto stride = (end - start + shape[0] - 1) / shape[0]; + if (stride != 1) { + LLVM_DEBUG({ + rangeOp.emitError( + "PtrAnalysis: make_range op with stride != 1 is not supported"); + }); + return failure(); + } + + auto infoStride = builder.getIndexAttr(stride); + auto size = builder.getIndexAttr(shape[0]); + auto offset = builder.getIndexAttr(start); + + SmallVector stateInfo; + SmallVector sizes; + stateInfo.emplace_back(infoStride, size); + sizes.emplace_back(size); + + state.updatePtrState(stateInfo, sizes, nullptr, offset, loc, builder); + return success(); +} + +LogicalResult +PtrAnalysis::visitOperandBroadcast(triton::BroadcastOp broadcastOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + broadcastOp.emitError( + "PtrAnalysis: PtrState should be empty when visiting broadcast"); + }); + return failure(); + } + + auto src = broadcastOp.getSrc(); + auto dst = broadcastOp.getResult(); + if (!isa(dst.getType())) { + LLVM_DEBUG({ + broadcastOp.emitRemark( + "PtrAnalysis: broadcast dst should be a shaped type"); + }); + return failure(); + } + + auto srcShape = cast(src.getType()).getShape(); + auto dstShape = cast(dst.getType()).getShape(); + if (srcShape.size() != dstShape.size()) { + LLVM_DEBUG({ + broadcastOp.emitRemark( + "PtrAnalysis: broadcast src and dst should have the same rank"); + }); + return failure(); + } + if (visitOperand(src, state, loc, builder).failed()) { + return failure(); + } + + if (state.sizes.size() != dstShape.size()) { + llvm::dbgs() << broadcastOp << "\n"; + state.dump(); + llvm::dbgs() << "dst.size = " << dstShape.size() << "\n"; + for (auto x : dstShape) + llvm::dbgs() << x << ", "; + llvm::dbgs() << "\n"; + } + + SmallVector newStateInfo(state.stateInfo); + SmallVector newSizes; + if (srcShape.size() != dstShape.size()) { + LLVM_DEBUG({ + broadcastOp.emitRemark( + "PtrAnalysis: unexpected state info size in broadcast"); + }); + return failure(); + } + for (size_t i = 0; i < dstShape.size(); ++i) { + newSizes.emplace_back(builder.getIndexAttr(dstShape[i])); + if (srcShape[i] == dstShape[i]) { + continue; + } else if (srcShape[i] < dstShape[i] && srcShape[i] == 1) { + for (auto &info : newStateInfo) { + if (info.dimIndex != i) + continue; + info.shape = builder.getIndexAttr(dstShape[i]); + } + } else { + LLVM_DEBUG({ + broadcastOp.emitRemark("unexpected dimensions used in broadcast"); + }); + return failure(); + } + } + state.updatePtrState(newStateInfo, newSizes, state.source, state.offset, loc, + builder, state.shouldLinearize); + return success(); +} + +LogicalResult PtrAnalysis::visitOperandSplat(triton::SplatOp splatOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + splatOp.emitError( + "PtrAnalysis: PtrState should be empty when visiting splat"); + }); + return failure(); + } + + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Visit SPLAT operation: " << splatOp << "\n"; + }); + + auto src = splatOp.getSrc(); + auto dst = splatOp.getResult(); + auto dstShape = cast(dst.getType()).getShape(); + + if (visitOperand(src, state, loc, builder).failed()) { + return failure(); + } + + LLVM_DEBUG({ + llvm::dbgs() << "splat ptrState: \n"; + state.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + + if (!state.isScalar()) { + LLVM_DEBUG( + { splatOp.emitRemark("PtrAnalysis: splat source should be scalar"); }); + return failure(); + } + + SmallVector newStateInfo; + SmallVector newSizes; + auto zeroAttr = builder.getIndexAttr(0); + if (isa(src.getType())) { + for (size_t i = 0; i < dstShape.size(); ++i) { + auto currentSize = builder.getIndexAttr(dstShape[i]); + newSizes.emplace_back(currentSize); + newStateInfo.emplace_back(zeroAttr, currentSize, i); + } + } else { + LLVM_DEBUG( + { splatOp.emitRemark("PtrAnalysis: unsupported splat pattern"); }); + return failure(); + } + state.updatePtrState(newStateInfo, newSizes, state.source, state.offset, loc, + builder, state.shouldLinearize); + + LLVM_DEBUG({ + llvm::dbgs() << "After SPLAT ptrState: \n"; + state.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + return success(); +} + +LogicalResult +PtrAnalysis::visitOperandExpandDims(triton::ExpandDimsOp expandDimsOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + expandDimsOp.emitError( + "PtrAnalysis: PtrState should be empty when visiting expand_dims"); + }); + return failure(); + } + + if (visitOperand(expandDimsOp.getSrc(), state, loc, builder).failed()) { + return failure(); + } + + auto dstShape = + cast(expandDimsOp.getResult().getType()).getShape(); + auto axis = expandDimsOp.getAxis(); + + SmallVector newStateInfo(state.stateInfo); + SmallVector newSizes(state.sizes); + size_t insertPos = 0; + for (auto &info : newStateInfo) { + if (info.dimIndex >= axis) + ++info.dimIndex; + if (info.dimIndex < axis) + ++insertPos; + } + auto zeroAttr = builder.getIndexAttr(0); + auto oneAttr = builder.getIndexAttr(1); + StateInfo insertInfo(zeroAttr, oneAttr, axis); + + newStateInfo.insert(newStateInfo.begin() + insertPos, insertInfo); + newSizes.insert(newSizes.begin() + axis, oneAttr); + + state.updatePtrState(newStateInfo, newSizes, state.source, state.offset, loc, + builder, state.shouldLinearize); + return success(); +} + +LogicalResult PtrAnalysis::visitOperandConstSplat(arith::ConstantOp op, + PtrState &state, + const Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + op->emitError( + "PtrAnalysis: PtrState should be empty when visiting const_splat"); + }); + return failure(); + } + + auto attr = cast(op.getValue()); + auto elementType = attr.getElementType(); + if (!attr.isSplat() || !isa(elementType)) { + LLVM_DEBUG( + { op->emitError("PtrAnalysis: only support splat integer constant"); }); + return failure(); + } + + auto value = attr.getValues()[0].getValue(); + auto constAttr = builder.getIndexAttr(value.getSExtValue()); + + auto resultShape = cast(op.getResult().getType()).getShape(); + + SmallVector sizes; + SmallVector stateInfo; + auto defaultAttr = builder.getIndexAttr(0); + + for (auto [i, shape] : llvm::enumerate(resultShape)) { + auto shapeAttr = builder.getIndexAttr(shape); + sizes.emplace_back(shapeAttr); + stateInfo.emplace_back(defaultAttr, shapeAttr, i); + } + + state.updatePtrState(stateInfo, sizes, nullptr, constAttr, loc, builder); + return success(); +} + +LogicalResult PtrAnalysis::visitOperandExtSI(arith::ExtSIOp extOp, + PtrState &state, + const Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + extOp.emitError( + "PtrAnalysis: PtrState should be empty when visiting extsi"); + }); + return failure(); + } + + if (visitOperand(extOp.getIn(), state, loc, builder).failed()) { + return failure(); + } + + return success(); +} + +LogicalResult PtrAnalysis::visitOperandRem(arith::RemSIOp remOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + remOp.emitError( + "PtrAnalysis: PtrState should be empty when visiting remsi"); + }); + return failure(); + } + LLVM_DEBUG({ + llvm::dbgs() << "before VisitRemOperands \n"; + state.dump(); + }); + + PtrState rhsState; + if (visitOperand(remOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + if (!rhsState.isScalar() || rhsState.hasSource()) { + LLVM_DEBUG({ + remOp.emitRemark("PtrAnalysis: only support cases when rhs of remainder " + "contains scalar"); + }); + return failure(); + } + + if (visitOperand(remOp.getLhs(), state, loc, builder).failed()) { + return failure(); + } + + bool hasAnnotation = optimizeDynamicOffset; + + auto zeroAttr = builder.getIndexAttr(0); + auto oneAttr = builder.getIndexAttr(1); + auto divisorAttr = rhsState.offset; + + auto staticOffset = getIntAttr(state.offset); + if ((!staticOffset.has_value() || !isMultiple(state.offset, divisorAttr)) && + !hasAnnotation) { + LLVM_DEBUG({ + remOp.emitRemark( + "PtrAnalysis: dynamic offset before REMSI, adding annotation"); + }); + return failure(); + } + + if (!getIntAttr(divisorAttr).has_value()) { + LLVM_DEBUG({ + remOp.emitError("PtrAnalysis: do not support dynamix divisor in REMSI."); + }); + return failure(); + } + + SmallVector newStateInfo; + for (auto info : state.stateInfo) { + if (!getIntAttr(info.stride).has_value()) { + LLVM_DEBUG({ + remOp.emitError( + "PtrAnalysis: do not support dynamix stride before REMSI."); + }); + return failure(); + } + + if (!isMultiple(divisorAttr, info.shape) && + !isMultiple(info.shape, divisorAttr)) { + LLVM_DEBUG({ + remOp.emitError( + "PtrAnalysis: do not support dynamix stride before REMSI."); + }); + } + + if (isMultiple(info.stride, divisorAttr)) { + newStateInfo.emplace_back(zeroAttr, info.shape, info.dimIndex); + } else if (isMultiple(divisorAttr, info.stride)) { + auto contiguousSize = + divOpFoldResult(divisorAttr, info.stride, loc, builder); + contiguousSize = + minOpFoldResult(contiguousSize, info.shape, loc, builder); + auto nonContiguousSize = + divOpFoldResult(info.shape, contiguousSize, loc, builder); + + auto staticNonContiguousSize = getIntAttr(nonContiguousSize); + if (!staticNonContiguousSize.has_value()) { + LLVM_DEBUG({ + remOp.emitError( + "PtrAnalysis: do not support dynamix size before REMSI."); + }); + return failure(); + } + + if (staticNonContiguousSize.value() > 1) + newStateInfo.emplace_back(zeroAttr, nonContiguousSize, info.dimIndex); + + newStateInfo.emplace_back(info.stride, contiguousSize, info.dimIndex); + } else { + LLVM_DEBUG({ + remOp.emitError("PtrAnalysis: stride that are not divisible by REMSI " + "are not allowed " + "to precede REMSI"); + }); + return failure(); + } + } + + auto newOffset = remOpFoldResult(state.offset, divisorAttr, loc, builder); + state.updatePtrState(newStateInfo, state.sizes, state.source, newOffset, loc, + builder, true); + + LLVM_DEBUG({ + llvm::dbgs() << "after VisitRemOperands \n"; + state.dump(); + }); + + return success(); +} + +LogicalResult PtrAnalysis::visitOperandDiv(arith::DivSIOp divOp, + PtrState &state, const Location loc, + OpBuilder &builder) { + if (!state.isEmpty()) { + LLVM_DEBUG({ + divOp.emitError( + "PtrAnalysis: PtrState should be empty when visiting divsi"); + }); + return failure(); + } + LLVM_DEBUG({ + llvm::dbgs() << "----------------------------------------------\n"; + llvm::dbgs() << "Visit DIVSI operation: " << divOp << "\n"; + }); + + PtrState rhsState; + if (visitOperand(divOp.getRhs(), rhsState, loc, builder).failed()) { + return failure(); + } + + if (!rhsState.isScalar() || rhsState.hasSource()) { + LLVM_DEBUG({ + divOp.emitRemark("PtrAnalysis: only support cases when rhs of remainder " + "contains scalar"); + }); + return failure(); + } + + if (visitOperand(divOp.getLhs(), state, loc, builder).failed()) { + return failure(); + } + + bool hasAnnotation = optimizeDynamicOffset; + + auto staticMultipleOf = extractDivisibilityFromOpFoldResult(state.offset); + if (!hasAnnotation && staticMultipleOf.has_value()) { + auto attr = builder.getIndexAttr(staticMultipleOf.value()); + hasAnnotation = isMultiple(attr, rhsState.offset); + } + + // add divState + auto zeroAttr = builder.getIndexAttr(0); + auto oneAttr = builder.getIndexAttr(1); + auto divisorAttr = rhsState.offset; + + auto staticOffset = getIntAttr(state.offset); + if ((!staticOffset.has_value() || !isMultiple(state.offset, divisorAttr)) && + !hasAnnotation) { + LLVM_DEBUG({ + divOp.emitRemark( + "PtrAnalysis: dynamic offset before DIVSI, adding annotation"); + }); + return failure(); + } + + if (!getIntAttr(divisorAttr).has_value()) { + LLVM_DEBUG({ + divOp.emitError("PtrAnalysis: do not support dynamix divisor in DIVSI."); + }); + return failure(); + } + + SmallVector newStateInfo; + for (auto info : state.stateInfo) { + auto staticStride = getIntAttr(info.stride); + if (!staticStride.has_value()) { + LLVM_DEBUG({ + divOp.emitError( + "PtrAnalysis: do not support dynamix stride before DIVSI."); + }); + return failure(); + } + + if (!isMultiple(divisorAttr, info.shape) && + !isMultiple(info.shape, divisorAttr)) { + LLVM_DEBUG({ + divOp.emitError( + "PtrAnalysis: do not support dynamix stride before DivSI."); + }); + } + + if (isMultiple(info.stride, divisorAttr)) { + auto newStride = divOpFoldResult(info.stride, divisorAttr, loc, builder); + newStateInfo.emplace_back(newStride, info.shape, info.dimIndex); + } else if (isMultiple(divisorAttr, info.stride)) { + auto nonContiguousSize = + divOpFoldResult(divisorAttr, info.stride, loc, builder); + nonContiguousSize = + minOpFoldResult(nonContiguousSize, info.shape, loc, builder); + auto contiguousSize = + divOpFoldResult(info.shape, nonContiguousSize, loc, builder); + + auto staticContiguousSize = getIntAttr(contiguousSize); + if (!staticContiguousSize.has_value()) { + LLVM_DEBUG({ + divOp.emitError( + "PtrAnalysis: do not support dynamix size before DIVSI."); + }); + return failure(); + } + + if (staticContiguousSize.value() != 0) + newStateInfo.emplace_back(oneAttr, contiguousSize, info.dimIndex); + + newStateInfo.emplace_back(zeroAttr, nonContiguousSize, info.dimIndex); + } else { + LLVM_DEBUG({ + divOp.emitError("PtrAnalysis: stride that are not divisible by DIVSI " + "are not allowed " + "to precede DIVSI"); + }); + return failure(); + } + } + + auto newOffset = divOpFoldResult(state.offset, divisorAttr, loc, builder); + state.updatePtrState(newStateInfo, state.sizes, state.source, newOffset, loc, + builder, true); + + LLVM_DEBUG({ + llvm::dbgs() << "after VisitDivOperands \n"; + state.dump(); + }); + return success(); +} + +LogicalResult PtrAnalysis::visitOperandAdd(arith::AddIOp addOp, PtrState &state, + const Location loc, + OpBuilder &builder) { + PtrState lhsState; + if (visitOperand(addOp.getLhs(), lhsState, loc, builder).failed()) { + return failure(); + } + + PtrState rhsState; + if (visitOperand(addOp.getRhs(), rhsState, loc, builder).failed()) + return failure(); + return state.addState(lhsState, rhsState, addOp, builder); +} + +LogicalResult PtrAnalysis::visitOperand(Value operand, PtrState &state, + const Location loc, + OpBuilder &builder) { + if (knownPtrs.find(operand) != knownPtrs.end()) { + state = knownPtrs.lookup(operand); + return success(); + } + + if (operandIsScalar(operand)) { + return initStateByScalar(operand, state, loc, builder); + } + + if (isa(operand.getType())) { + return initStateByPointer(operand, state, loc, builder); + } + + if (auto op = operand.getDefiningOp()) { + return visitOperandAdd(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandMul(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandSub(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandMakeRange(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandBroadcast(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandSplat(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandExpandDims(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandAddptr(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandConstSplat(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandRem(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandDiv(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + return visitOperandExtSI(op, state, loc, builder); + } else if (auto op = operand.getDefiningOp()) { + LLVM_DEBUG({ + op.emitRemark("TritonToStructured: Invalid dynamic offset" + "The load operation's offset cannot be derived from " + "another load result."); + }); + return failure(); + } else if (auto op = operand.getDefiningOp()) { + LLVM_DEBUG({ + op.emitWarning("IllegalTypeConversionInAddressCalculation" + "float-to-int precision conversion is not supported " + "during address computation."); + llvm::dbgs() << "Operand: \n"; + operand.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + return failure(); + } else if (!operand.getDefiningOp()) { + if (!knownPtrs.contains(operand)) { + llvm::dbgs() << "TritonToStructured: Pointer analysis is not supported " + "for input parameters\n"; + return failure(); + } + + // This operand must be an iter-arg of an inner-loop in a multiple-level + // nested loop, which means its PtrState must have already been populated + // during rewriteForOp of the parent loop. + state = knownPtrs[operand]; + return success(); + } else { + auto op = operand.getDefiningOp(); + LLVM_DEBUG({ + op->emitWarning("TritonToStructured: encountered addptr operand produced " + "by an unsupported operation"); + llvm::dbgs() << "Operand: \n"; + operand.dump(); + llvm::dbgs() << "----------------------------------------------\n"; + }); + return failure(); + } + return success(); +} + +LogicalResult PtrAnalysis::rewriteAddptrOp(triton::AddPtrOp op) { + OpBuilder builder(op); + auto loc = op.getLoc(); + + PtrState state; + if (visitOperandAddptr(op, state, op.getLoc(), builder).failed()) { + return failure(); + } + + auto maketptrOp = state.createAddPtrOp(builder, op.getLoc()); + knownPtrs[op.getResult()] = state; + ptrMap.map(op.getResult(), maketptrOp.getResult()); + + return success(); +} + +void PtrState::analyzePermute() { + for (size_t i = 0; i < stateInfo.size(); ++i) { + permuteIds.emplace_back(i); + } + + // Generate dimension permutation following triton::TransOp convention: + // out[i] = in[permute[i]] where permute maps logical to physical + // dimensions + // Example: + // logicalAxes: [0, 1] (original order) + // physicalAxes: [1, 0] (memory layout order) + // permute: [1, 0] (out[0] = in[1], out[1] = in[0]) + + std::stable_sort(permuteIds.begin(), permuteIds.end(), + [&](const size_t &a, const size_t &b) { + return isGreater(stateInfo[a].stride, stateInfo[b].stride); + }); + + for (size_t i = 0; !isPermuted && i < permuteIds.size(); ++i) { + if (i != permuteIds[i]) + isPermuted = true; + } + + return; +} + +std::optional +extractDivisibilityFromOpFoldResult(mlir::OpFoldResult ofr) { + auto value = dyn_cast(ofr); + if (!value) { + return std::nullopt; + } + auto defOp = value.getDefiningOp(); + if (!defOp) { + return std::nullopt; + } + + auto divisibilityAttr = defOp->getAttr("tt.divisibility"); + if (!divisibilityAttr) { + return std::nullopt; + } + + auto denseAttr = dyn_cast(divisibilityAttr); + if (!denseAttr || denseAttr.empty()) { + return std::nullopt; + } + + return denseAttr.getValues()[0]; +} +} // namespace TritonToStructuredIncubated diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.cpp new file mode 100755 index 00000000..d8338328 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.cpp @@ -0,0 +1,149 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.h" + +#include +#include +#include + +#include "incubated/Conversion/UtilsIncubated/InterleaveOptimization.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/Builders.h" +#if __has_include("bishengir/Dialect/HIVM/IR/HIVM.h") +#include "bishengir/Dialect/HIVM/IR/HIVM.h" +#endif +#include "mlir/IR/Operation.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#if __has_include("bishengir/Dialect/Annotation/IR/Annotation.h") +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#endif +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Transforms/Transforms.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Visitors.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "mlir/Transforms/Passes.h" + +#include "incubated/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.h" +#include "incubated/Conversion/TritonToStructuredIncubated/MemOpConverter.h" +#include "incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h" +#include "llvm/ADT/BitVector.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/Twine.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/LogicalResult.h" + +#define DEBUG_TYPE "triton-to-structured" + +using namespace mlir; +using namespace triton; + +void TritonToStructuredIncubatedPass::getDependentDialects( + DialectRegistry ®istry) const { + registry.insert(); +} + +void TritonToStructuredIncubatedPass:: + populateTritonToStructuredCanonicalizationPatterns( + RewritePatternSet &patterns) { + // TODO enable this optimization after fixing the bisheng bug it causes in + // current version + // patterns.add(patterns.getContext()); + patterns.add( + patterns.getContext()); +} + +void TritonToStructuredIncubatedPass::populateTritonToStructuredPatterns( + RewritePatternSet &patterns, bool optimizeDynamicOffset, + bool enableMaskFallbackConversion, bool compileOn91095) { + patterns.add( + patterns.getContext(), optimizeDynamicOffset, + enableMaskFallbackConversion, compileOn91095); + patterns.add( + patterns.getContext(), optimizeDynamicOffset, + enableMaskFallbackConversion, false); +} + +void TritonToStructuredIncubatedPass::runOnOperation() { + auto moduleOp = getOperation(); + ConversionTarget target(getContext()); + RewritePatternSet canonicalizerPatterns(&getContext()); + + this->populateTritonToStructuredCanonicalizationPatterns( + canonicalizerPatterns); + if (failed(applyPatternsAndFoldGreedily(moduleOp, + std::move(canonicalizerPatterns)))) { + moduleOp.emitWarning("Canonicalize failed"); + } + + RewritePatternSet tritonToStructuredPatterns(&getContext()); + populateTritonToStructuredPatterns( + tritonToStructuredPatterns, optimizeDynamicOffset, + enableMaskFallbackConversion, compileOn91095); + + if (failed(applyPatternsAndFoldGreedily( + moduleOp, std::move(tritonToStructuredPatterns)))) { + LLVM_DEBUG({ moduleOp->emitRemark("PtrAnalysis: rewrite MemOp failed"); }); + } + + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + if (failed(runPipeline(pm, getOperation()))) { + moduleOp->emitWarning("Canonicalize failed"); + } +} + +std::unique_ptr> +triton::createTritonToStructuredIncubatedPass() { + return std::make_unique(); +} + +std::unique_ptr> +triton::createTritonToStructuredIncubatedPass(bool enableMaskFallbackConversion, + bool optimizeDynamicOffset, + bool compileOn91095) { + return std::make_unique( + enableMaskFallbackConversion, optimizeDynamicOffset, compileOn91095); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.cpp new file mode 100755 index 00000000..76ca7d13 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.cpp @@ -0,0 +1,502 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "mlir/Transforms/Passes.h" + +#define DEBUG_TYPE "triton-bubble-up-operation" + +template +BubbleUpExtract::BubbleUpExtract(MLIRContext *context, + bool enableAggressiveMode) + : OpRewritePattern(context), + enableAggressiveMode(enableAggressiveMode) {} + +template +LogicalResult +BubbleUpExtract::matchAndRewrite(ExtractOpTy op, + PatternRewriter &rewriter) const { + Value tensorValue; + if constexpr (std::is_same_v) { + tensorValue = op.getTensor(); + } else if constexpr (std::is_same_v) { + tensorValue = op.getSource(); + if (tensorValue.getType() == op.getResult().getType()) { + rewriter.replaceAllUsesWith(op.getResult(), tensorValue); + rewriter.eraseOp(op); + return success(); + } + } else { + llvm_unreachable("Unhandled case"); + } + auto funcOp = op->template getParentOfType(); + auto parentOp = tensorValue.getDefiningOp(); + auto loc = op.getLoc(); + + if (!parentOp || (!enableAggressiveMode && !parentOp->hasOneUse())) { + return failure(); + } + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Before bubble up\n" << op << '\n' << funcOp << "\n"; + }); + + if (auto extsiOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, extsiOp, loc, rewriter); + } else if (auto addIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, addIOp, loc, rewriter); + } else if (auto subIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, subIOp, loc, rewriter); + } else if (auto mulIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, mulIOp, loc, rewriter); + } else if (auto divSIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, divSIOp, loc, rewriter); + } else if (auto remSIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, remSIOp, loc, rewriter); + } else if (auto maxSIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, maxSIOp, loc, rewriter); + } else if (auto minSIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, minSIOp, loc, rewriter); + } else if (auto andIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, andIOp, loc, rewriter); + } else if (auto orIOp = dyn_cast(parentOp)) { + bubbleUpIntBinaryOp(op, orIOp, loc, rewriter); + } else if (auto cmpIOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, cmpIOp, loc, rewriter); + } else if (auto truncFOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, truncFOp, loc, rewriter); + } else if (auto extFOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, extFOp, loc, rewriter); + } else if (auto fpTosiOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, fpTosiOp, loc, rewriter); + } else if (auto siTofpOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, siTofpOp, loc, rewriter); + } else if (auto clampFOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, clampFOp, loc, rewriter); + } else if (auto addFOp = dyn_cast(parentOp)) { + bubbleUpFloatBinaryOp(op, addFOp, loc, rewriter); + } else if (auto subFOp = dyn_cast(parentOp)) { + bubbleUpFloatBinaryOp(op, subFOp, loc, rewriter); + } else if (auto mulFOp = dyn_cast(parentOp)) { + bubbleUpFloatBinaryOp(op, mulFOp, loc, rewriter); + } else if (auto divFOp = dyn_cast(parentOp)) { + bubbleUpFloatBinaryOp(op, divFOp, loc, rewriter); + } else if (auto minNumFOp = dyn_cast(parentOp)) { + bubbleUpFloatBinaryOp(op, minNumFOp, loc, rewriter); + } else if (auto maxNumFOp = dyn_cast(parentOp)) { + bubbleUpFloatBinaryOp(op, maxNumFOp, loc, rewriter); + } else if (auto cmpFOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, cmpFOp, loc, rewriter); + } else if (auto broadCastOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, broadCastOp, loc, rewriter); + } else if (auto expandDimsOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, expandDimsOp, loc, rewriter); + } else if (auto splatOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, splatOp, loc, rewriter); + } else if (auto makeRangeOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, makeRangeOp, loc, rewriter); + } else if (auto addPtrOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, addPtrOp, loc, rewriter); + } else if (auto floorOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, floorOp, loc, rewriter); + } else if (auto ceilOp = dyn_cast(parentOp)) { + bubbleUpOperation(op, ceilOp, loc, rewriter); + } else if (auto extractSliceOp = dyn_cast(parentOp)) { + if constexpr (std::is_same_v) { + bubbleUpOperation(op, extractSliceOp, loc, rewriter); + } else { + return failure(); + } + } else { + return failure(); + } + if (parentOp->use_empty()) + rewriter.eraseOp(parentOp); + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "After bubble up\n" << funcOp << '\n'; + }); + + return success(); +} + +template +Value BubbleUpExtract::createExtractOp( + ExtractOpTy op, Value value, Location loc, + PatternRewriter &rewriter) const { + llvm_unreachable("Unhandled extract operation"); +} + +template <> +Value BubbleUpExtract::createExtractOp( + tensor::ExtractOp op, Value value, Location loc, + PatternRewriter &rewriter) const { + auto extractedOp = + rewriter.create(loc, value, op.getIndices()); + extractedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + return extractedOp; +} + +template <> +Value BubbleUpExtract::createExtractOp( + tensor::ExtractSliceOp op, Value value, Location loc, + PatternRewriter &rewriter) const { + auto extractedOp = rewriter.create( + loc, value, op.getMixedOffsets(), op.getMixedSizes(), + op.getMixedStrides()); + extractedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + return extractedOp; +} + +template +template +void BubbleUpExtract::bubbleUpIntBinaryOp( + ExtractOpTy op, BinOpTy binOp, Location loc, + PatternRewriter &rewriter) const { + auto lhs = createExtractOp(op, binOp.getLhs(), loc, rewriter); + auto rhs = createExtractOp(op, binOp.getRhs(), loc, rewriter); + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Binary\n" << *op << '\n' << binOp << '\n'; + }); + rewriter.replaceOpWithNewOp(op, lhs, rhs); +} + +template +template +void BubbleUpExtract::bubbleUpFloatBinaryOp( + ExtractOpTy op, BinOpTy binOp, Location loc, + PatternRewriter &rewriter) const { + auto lhs = createExtractOp(op, binOp.getLhs(), loc, rewriter); + auto rhs = createExtractOp(op, binOp.getRhs(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, lhs, rhs); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, arith::ExtSIOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto in = createExtractOp(op, parentOp.getIn(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, op.getResult().getType(), in); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, arith::CmpIOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto lhs = createExtractOp(op, parentOp.getLhs(), loc, rewriter); + auto rhs = createExtractOp(op, parentOp.getRhs(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, parentOp.getPredicateAttr(), + lhs, rhs); +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractOp op, triton::BroadcastOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto src = parentOp.getSrc(); + auto srcShape = cast(src.getType()).getShape(); + SmallVector newIndices; + for (const auto &[index, shape] : + llvm::zip_equal(op.getIndices(), srcShape)) { + if (shape == 1) { + newIndices.push_back( + rewriter.create(loc, rewriter.getIndexAttr(0))); + } else { + newIndices.push_back(index); + } + } + auto extractedOp = rewriter.create(loc, src, newIndices); + extractedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + rewriter.replaceOp(op, extractedOp); +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractSliceOp op, triton::BroadcastOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto src = parentOp.getSrc(); + auto srcShape = cast(src.getType()).getShape(); + SmallVector newOffsets; + SmallVector newSizes; + bool isScalarLikeSrc = true; + for (const auto &[offset, size, shape] : + llvm::zip_equal(op.getMixedOffsets(), op.getMixedSizes(), srcShape)) { + if (shape == 1) { + newOffsets.push_back(rewriter.getIndexAttr(0)); + newSizes.push_back(rewriter.getIndexAttr(1)); + } else { + newOffsets.push_back(offset); + newSizes.push_back(size); + } + if (getConstantIntValue(newSizes.back()).value_or(-1) != 1) + isScalarLikeSrc = false; + } + auto extractedOp = rewriter.create( + loc, src, newOffsets, newSizes, op.getMixedStrides()); + extractedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + if (isScalarLikeSrc) { + SmallVector indices( + srcShape.size(), + rewriter.create(loc, rewriter.getIndexAttr(0))); + auto extractedValue = + rewriter.create(loc, extractedOp, indices); + rewriter.replaceOpWithNewOp(op, op.getResult().getType(), + extractedValue); + } else { + rewriter.replaceOpWithNewOp( + op, op.getResult().getType(), extractedOp); + } +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractOp op, triton::ExpandDimsOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto src = parentOp.getSrc(); + SmallVector newIndices; + for (const auto index : llvm::enumerate(op.getIndices())) { + if (index.index() != parentOp.getAxis()) + newIndices.push_back(index.value()); + } + auto extractedOp = rewriter.create(loc, src, newIndices); + extractedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + rewriter.replaceOp(op, extractedOp); +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractSliceOp op, triton::ExpandDimsOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto src = parentOp.getSrc(); + auto srcShape = cast(src.getType()).getShape(); + SmallVector newOffsets; + SmallVector newSizes; + SmallVector newStrides; + for (size_t i = 0; i <= srcShape.size(); i++) { + if (i != parentOp.getAxis()) { + newOffsets.push_back(op.getMixedOffsets()[i]); + newSizes.push_back(op.getMixedSizes()[i]); + newStrides.push_back(op.getMixedStrides()[i]); + } + } + auto extractedOp = rewriter.create( + loc, src, newOffsets, newSizes, newStrides); + extractedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + rewriter.replaceOpWithNewOp(op, extractedOp, + parentOp.getAxisAttr()); +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractOp op, triton::SplatOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto src = parentOp.getSrc(); + rewriter.replaceOp(op, src); +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractSliceOp op, triton::SplatOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto src = parentOp.getSrc(); + rewriter.replaceOpWithNewOp( + op, cast(op.getResult().getType()), src); +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractOp op, triton::MakeRangeOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto resultType = cast(parentOp.getResult().getType()); + rewriter.replaceOpWithNewOp( + op, resultType.getElementType(), op.getIndices()[0]); +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractSliceOp op, triton::MakeRangeOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto resultType = cast(parentOp.getResult().getType()); + Value idx; + if (auto offsetVal = dyn_cast(op.getMixedOffsets()[0])) { + idx = offsetVal; + } else { + idx = rewriter.create( + op.getLoc(), rewriter.getIndexAttr( + getConstantIntValue(op.getMixedOffsets()[0]).value())); + } + idx = rewriter.create(op.getLoc(), + resultType.getElementType(), idx); + rewriter.replaceOpWithNewOp(op, op.getResult().getType(), + idx); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, triton::AddPtrOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto ptr = createExtractOp(op, parentOp.getPtr(), loc, rewriter); + auto offset = createExtractOp(op, parentOp.getOffset(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, ptr.getType(), ptr, offset); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, arith::TruncFOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto in = createExtractOp(op, parentOp.getIn(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, op.getResult().getType(), + in); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, arith::ExtFOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto in = createExtractOp(op, parentOp.getIn(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, op.getResult().getType(), in); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, arith::FPToSIOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto in = createExtractOp(op, parentOp.getIn(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, op.getResult().getType(), + in); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, arith::SIToFPOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto in = createExtractOp(op, parentOp.getIn(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, op.getResult().getType(), + in); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, triton::ClampFOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto x = createExtractOp(op, parentOp.getX(), loc, rewriter); + auto min = createExtractOp(op, parentOp.getMin(), loc, rewriter); + auto max = createExtractOp(op, parentOp.getMax(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, x, min, max, + parentOp.getPropagateNan()); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, arith::CmpFOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto lhs = createExtractOp(op, parentOp.getLhs(), loc, rewriter); + auto rhs = createExtractOp(op, parentOp.getRhs(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, parentOp.getPredicateAttr(), + lhs, rhs); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, math::FloorOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto operand = createExtractOp(op, parentOp.getOperand(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, operand, + parentOp.getFastmath()); +} + +template +void BubbleUpExtract::bubbleUpOperation( + ExtractOpTy op, math::CeilOp parentOp, Location loc, + PatternRewriter &rewriter) const { + auto operand = createExtractOp(op, parentOp.getOperand(), loc, rewriter); + rewriter.replaceOpWithNewOp(op, operand, + parentOp.getFastmath()); +} + +template <> +void BubbleUpExtract::bubbleUpOperation( + tensor::ExtractOp op, tensor::ExtractSliceOp parentOp, Location loc, + PatternRewriter &rewriter) const { + SmallVector newIndices; + for (const auto &[offset, index] : + llvm::zip_equal(parentOp.getMixedOffsets(), op.getIndices())) { + Value offsetVal; + if (isa(offset)) { + offsetVal = offset.template get(); + } else { + offsetVal = rewriter.create( + op.getLoc(), rewriter.getIndexAttr(*getConstantIntValue(offset))); + } + newIndices.push_back( + rewriter.create(op.getLoc(), offsetVal, index)); + } + rewriter + .replaceOpWithNewOp(op, parentOp.getSource(), + newIndices) + ->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); +} + +BubbleUpOperationPass::BubbleUpOperationPass( + const BubbleUpOperationOptions &options) + : BubbleUpOperationBase(options) {} + +void BubbleUpOperationPass::runOnOperation() { + ModuleOp moduleOp = getOperation(); + MLIRContext *ctx = &getContext(); + + RewritePatternSet patterns(ctx); + patterns.add, + BubbleUpExtract>(ctx, + enableAggressiveMode); + + if (failed(applyPatternsAndFoldGreedily(moduleOp, std::move(patterns)))) { + moduleOp->emitError("failed to apply Patterns"); + signalPassFailure(); + } + + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } +} + +std::unique_ptr> +triton::createBubbleUpOperationPass(const BubbleUpOperationOptions &options) { + return std::make_unique(options); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/CMakeLists.txt new file mode 100755 index 00000000..cb0684c7 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/CMakeLists.txt @@ -0,0 +1,20 @@ +add_triton_library(TritonToUnstructureIncubated + UnstructureConversionPass.cpp + OffsetAnalysis.cpp + BubbleUpOperation.cpp + + DEPENDS + TritonToUnstructureConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonAnalysis + MLIRSCFTransforms +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.cpp new file mode 100755 index 00000000..52cc14e9 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.cpp @@ -0,0 +1,989 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#include "mlir/Dialect/Linalg/IR/Linalg.h" + +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "triton-offset-analysis" + +namespace mlir { +namespace triton { + +PtrOffsetInfo::PtrOffsetInfo() : ptr(nullptr), offset(nullptr) {} + +PtrOffsetInfo::PtrOffsetInfo(const PtrOffsetInfo &other) { *this = other; } + +PtrOffsetInfo::PtrOffsetInfo(const Value &ptr) : ptr(ptr) { setZeroOffset(); } + +PtrOffsetInfo::PtrOffsetInfo(ArrayRef structured) + : ptr(nullptr), offset(nullptr) { + setStructured(structured); +} + +PtrOffsetInfo::PtrOffsetInfo(const Value &ptr, bool structured) : ptr(ptr) { + setZeroOffset(); + if (auto tensorType = dyn_cast(ptr.getType())) + this->structured.resize(tensorType.getRank(), structured); +} + +PtrOffsetInfo::PtrOffsetInfo(const Value &ptr, ArrayRef structured) + : ptr(ptr) { + setStructured(structured); +} + +PtrOffsetInfo::PtrOffsetInfo(const Value &ptr, const Value &offset, + bool structured) + : ptr(ptr), offset(offset) { + if (auto tensorType = dyn_cast(ptr.getType())) + this->structured.resize(tensorType.getRank(), structured); +} + +PtrOffsetInfo::PtrOffsetInfo(const Value &ptr, const Value &offset, + ArrayRef structured) + : ptr(ptr), offset(offset) { + setStructured(structured); +} + +PtrOffsetInfo &PtrOffsetInfo::operator=(const PtrOffsetInfo &other) { + setPtr(other.getPtr()); + setOffset(other.getOffset()); + setOffsets(other.getOffsets()); + setStructured(other.getStructured()); + setScalarLike(other.isScalarLike()); + return *this; +} + +Value PtrOffsetInfo::getPtr() const { return this->ptr; } +Value PtrOffsetInfo::getOffset() const { return this->offset; } +SmallVector PtrOffsetInfo::getOffsets() const { + return this->tptOffsets; +} +SmallVector &PtrOffsetInfo::getOffsetsRef() { return this->tptOffsets; } + +bool PtrOffsetInfo::isScalarLike() const { return this->scalarLike; } + +SmallVector &PtrOffsetInfo::getStructuredRef() { + return this->structured; +} +const SmallVector &PtrOffsetInfo::getStructured() const { + return this->structured; +} + +int PtrOffsetInfo::getRank() const { return structured.size(); } + +void PtrOffsetInfo::setPtr(const Value &ptr) { this->ptr = ptr; } +void PtrOffsetInfo::setOffset(const Value &offset) { this->offset = offset; } + +void PtrOffsetInfo::setOffsets(ValueRange offsets) { + tptOffsets.clear(); + for (auto offset : offsets) + tptOffsets.push_back(offset); +} + +void PtrOffsetInfo::setStructured() { + assert(ptr && "ptr Should be to infer rank"); + this->structured.clear(); + if (auto tensorType = dyn_cast(ptr.getType())) + this->structured.resize(tensorType.getRank(), true); +} + +void PtrOffsetInfo::setStructured(int rank) { + this->structured.clear(); + this->structured.resize(rank, true); +} + +void PtrOffsetInfo::setUnstructured() { + assert(ptr && "ptr Should be to infer rank"); + this->structured.clear(); + if (auto tensorType = dyn_cast(ptr.getType())) + this->structured.resize(tensorType.getRank(), false); +} + +void PtrOffsetInfo::setUnstructured(int rank) { + this->structured.clear(); + this->structured.resize(rank, false); +} + +void PtrOffsetInfo::setStructured(ArrayRef structured) { + this->structured.resize(structured.size()); + for (size_t i = 0; i < structured.size(); i++) + this->structured[i] = structured[i]; +} + +void PtrOffsetInfo::setStructured(const PtrOffsetInfo &other) { + this->setStructured(other.getStructured()); +} + +void PtrOffsetInfo::setScalarLike(bool scalarLike) { + this->scalarLike = scalarLike; +} + +bool PtrOffsetInfo::isStructured(int dim) const { + return this->scalarLike || structured[dim]; +} + +bool PtrOffsetInfo::isStructured() const { + return this->scalarLike || + llvm::all_of(structured, [](auto dim) { return dim; }); +} + +bool PtrOffsetInfo::isUnstructured() const { + return llvm::all_of(structured, [](auto dim) { return !dim; }); +} + +void PtrOffsetInfo::setZeroOffset() { + if (!ptr) + return; + Value offset; + OpBuilder builder(ptr.getContext()); + builder.setInsertionPointToStart(ptr.getParentBlock()); + if (auto tensorType = dyn_cast(ptr.getType())) { + offset = builder.create( + ptr.getLoc(), DenseElementsAttr::get( + RankedTensorType::get(tensorType.getShape(), + builder.getIntegerType(64)), + builder.getZeroAttr(builder.getIntegerType(64)))); + } else { + offset = builder.create(ptr.getLoc(), + builder.getI64IntegerAttr(0)); + } + setOffset(offset); +} + +PtrOffsetInfo combineInfo(const PtrOffsetInfo &lhs, const PtrOffsetInfo &rhs) { + PtrOffsetInfo info; + assert(lhs.getRank() == rhs.getRank() && "Rank must be same to be combined"); + + info.setScalarLike(lhs.isScalarLike() && rhs.isScalarLike()); + SmallVector &structuredRef = info.getStructuredRef(); + structuredRef.resize(lhs.getRank()); + for (size_t i = 0; i < structuredRef.size(); i++) + structuredRef[i] = lhs.isStructured(i) && rhs.isStructured(i); + return info; +} + +void parse(Value operand, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + if (offsetMap.contains(operand)) { + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "found\n" << operand << '\n'; + }); + return; + } + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "parse\n" << operand << '\n'; + }); + + if (auto *defOp = operand.getDefiningOp()) { + if (isa(defOp->getDialect())) { + parseArithOp(defOp, loc, rewriter, offsetMap); + } else if (isa(defOp->getDialect())) { + parseTritonOp(defOp, loc, rewriter, offsetMap); + } else { + if (auto ifOp = dyn_cast(defOp)) { + parseIf(ifOp, loc, rewriter, offsetMap, operand); + } else if (auto yieldOp = dyn_cast(defOp)) { + parseYield(yieldOp, loc, rewriter, offsetMap); + } else if (auto loopOp = dyn_cast(defOp)) { + parseLoopOp(loopOp, loc, rewriter, offsetMap, operand); + } else if (auto extractOp = dyn_cast(defOp)) { + parseExtract(extractOp, loc, rewriter, offsetMap); + } + } + } else if (auto blockArgument = dyn_cast(operand)) { + auto parentOp = blockArgument.getOwner()->getParentOp(); + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Handling block argument\n" << *blockArgument.getOwner() << '\n'; + }); + if (isa(parentOp)) { + if (auto ptrType = dyn_cast(operand.getType())) { + offsetMap[operand] = PtrOffsetInfo(operand, true); + } else { + offsetMap[operand] = PtrOffsetInfo(); + } + } else if (auto loopOp = dyn_cast(parentOp)) { + parseLoopRegionIterArg(loopOp, loc, rewriter, offsetMap, blockArgument); + } + } else { + llvm_unreachable("Unreachable"); + } + + if (!offsetMap.contains(operand)) { + offsetMap[operand] = PtrOffsetInfo(); + if (auto tensorType = dyn_cast(operand.getType())) + offsetMap[operand].setUnstructured(tensorType.getRank()); + } + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "finish parse\n" << operand << '\n'; + auto data = offsetMap.at(operand); + for (auto s : data.getStructuredRef()) + os << s; + os << "\n"; + }); +} + +void parseLoopRegionIterArg(LoopLikeOpInterface loopOp, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap, + BlockArgument regionIterArg) { + if (auto whileOp = dyn_cast(loopOp.getOperation()); + whileOp && whileOp.getAfterBody() == regionIterArg.getOwner()) { + auto argNum = regionIterArg.getArgNumber(); + auto conditionArg = whileOp.getConditionOp().getArgs()[argNum]; + parse(conditionArg, loc, rewriter, offsetMap); + offsetMap[regionIterArg] = offsetMap[conditionArg]; + return; + } + OpOperand *initArgOperand = loopOp.getTiedLoopInit(regionIterArg); + if (!initArgOperand) + return; + Value initArg = initArgOperand->get(); + parse(initArg, loc, rewriter, offsetMap); + offsetMap[regionIterArg] = offsetMap[initArg]; +} + +void parseArithOp(Operation *arithOp, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + assert(isa(arithOp->getDialect())); + if (auto addIOp = dyn_cast(arithOp)) { + parseAddI(addIOp, loc, rewriter, offsetMap); + } else if (auto subIOp = dyn_cast(arithOp)) { + parseSubI(subIOp, loc, rewriter, offsetMap); + } else if (auto indexCastOp = dyn_cast(arithOp)) { + parseIndexCast(indexCastOp, loc, rewriter, offsetMap); + } else if (auto constantFloatOp = dyn_cast(arithOp)) { + parseConstantOp(constantFloatOp, loc, rewriter, offsetMap); + } else if (auto constantIntOp = dyn_cast(arithOp)) { + parseConstantOp(constantIntOp, loc, rewriter, offsetMap); + } else if (auto constantOp = dyn_cast(arithOp)) { + parseConstantOp(constantOp, loc, rewriter, offsetMap); + } else if (auto extSIOp = dyn_cast(arithOp)) { + parseExtSI(extSIOp, loc, rewriter, offsetMap); + } else if (auto mulIOp = dyn_cast(arithOp)) { + parseMulI(mulIOp, loc, rewriter, offsetMap); + } else if (auto remSIOp = dyn_cast(arithOp)) { + parseBinaryOp(remSIOp, loc, rewriter, offsetMap); + } else if (auto divSIOp = dyn_cast(arithOp)) { + parseBinaryOp(divSIOp, loc, rewriter, offsetMap); + } else if (auto selectOp = dyn_cast(arithOp)) { + parseSelect(selectOp, loc, rewriter, offsetMap); + } else if (auto fPToSIOp = dyn_cast(arithOp)) { + parseFPToSI(fPToSIOp, loc, rewriter, offsetMap); + } else if (auto sIToFPOp = dyn_cast(arithOp)) { + parseSIToFP(sIToFPOp, loc, rewriter, offsetMap); + } else if (auto mulFOp = dyn_cast(arithOp)) { + parseBinaryOp(mulFOp, loc, rewriter, offsetMap); + } else if (auto divFOp = dyn_cast(arithOp)) { + parseBinaryOp(divFOp, loc, rewriter, offsetMap); + } else if (auto addFOp = dyn_cast(arithOp)) { + parseBinaryOp(addFOp, loc, rewriter, offsetMap); + } else if (auto subFOp = dyn_cast(arithOp)) { + parseBinaryOp(subFOp, loc, rewriter, offsetMap); + } else if (auto minNumFOp = dyn_cast(arithOp)) { + parseBinaryOp(minNumFOp, loc, rewriter, offsetMap); + } else if (auto maxNumFOp = dyn_cast(arithOp)) { + parseBinaryOp(maxNumFOp, loc, rewriter, offsetMap); + } else if (auto maxSIOp = dyn_cast(arithOp)) { + parseBinaryOp(maxSIOp, loc, rewriter, offsetMap); + } else if (auto minSIOp = dyn_cast(arithOp)) { + parseBinaryOp(minSIOp, loc, rewriter, offsetMap); + } else if (auto cmpIOp = dyn_cast(arithOp)) { + parseBinaryOp(cmpIOp, loc, rewriter, offsetMap); + } else if (auto andIOp = dyn_cast(arithOp)) { + parseBinaryOp(andIOp, loc, rewriter, offsetMap); + } else if (auto orIOp = dyn_cast(arithOp)) { + parseBinaryOp(orIOp, loc, rewriter, offsetMap); + } +} + +void parseTritonOp(Operation *tritonOp, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + assert(isa(tritonOp->getDialect())); + if (auto addPtrOp = dyn_cast(tritonOp)) { + parseAddPtr(addPtrOp, loc, rewriter, offsetMap); + } else if (auto splatOp = dyn_cast(tritonOp)) { + parseSplat(splatOp, loc, rewriter, offsetMap); + } else if (auto getProgramIdOp = dyn_cast(tritonOp)) { + parseConstantOp(getProgramIdOp, loc, rewriter, offsetMap); + } else if (auto getNumProgramsOp = + dyn_cast(tritonOp)) { + parseConstantOp(getNumProgramsOp, loc, rewriter, offsetMap); + } else if (auto makeRangeOp = dyn_cast(tritonOp)) { + parseMakeRange(makeRangeOp, loc, rewriter, offsetMap); + } else if (auto bitcastOp = dyn_cast(tritonOp)) { + parseBitcast(bitcastOp, loc, rewriter, offsetMap); + } else if (auto loadOp = dyn_cast(tritonOp)) { + parseLoad(loadOp, loc, rewriter, offsetMap); + } else if (auto broadcastOp = dyn_cast(tritonOp)) { + parseBroadcast(broadcastOp, loc, rewriter, offsetMap); + } else if (auto expandDimsOp = dyn_cast(tritonOp)) { + parseExpandDims(expandDimsOp, loc, rewriter, offsetMap); + } else if (auto clampFOp = dyn_cast(tritonOp)) { + parseClampF(clampFOp, loc, rewriter, offsetMap); + } + // FIXME:Z|wait triton version upgrade to 3.4 + // else if (auto makeTensorDescOp = + // dyn_cast(tritonOp)) { + // parseMakeTensorDesc(makeTensorDescOp, loc, rewriter, offsetMap); + // } + else if (auto makeTensorPtrOp = dyn_cast(tritonOp)) { + parseMakeTensorPtr(makeTensorPtrOp, loc, rewriter, offsetMap); + } else if (auto reduceOp = dyn_cast(tritonOp)) { + parseReduce(reduceOp, loc, rewriter, offsetMap); + } else if (auto reduceReturnOp = dyn_cast(tritonOp)) { + parseReduceReturn(reduceReturnOp, loc, rewriter, offsetMap); + } else if (auto advanceOp = dyn_cast(tritonOp)) { + parseAdvance(advanceOp, loc, rewriter, offsetMap); + } else if (auto intToPtrOp = dyn_cast(tritonOp)) { + parseIntToPtr(intToPtrOp, loc, rewriter, offsetMap); + } +} + +void parseAddPtr(triton::AddPtrOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get addPtr base_ptr + Value ptr = op.getPtr(); + parse(ptr, op.getLoc(), rewriter, offsetMap); + // Get addPtr offset + Value offsetValue = op.getOffset(); + parse(offsetValue, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo ptrOffsetInfo = offsetMap.at(ptr); + PtrOffsetInfo offsetOffsetInfo = offsetMap.at(offsetValue); + // Modify IR + + RewriterBase::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(op); + if (auto offsetType = dyn_cast(offsetValue.getType())) { + auto offsetElementType = cast(offsetType.getElementType()); + if (offsetElementType.getWidth() != 64) { + auto newOffsetType = RankedTensorType::get(offsetType.getShape(), + rewriter.getIntegerType(64)); + offsetValue = rewriter.create(op.getLoc(), newOffsetType, + offsetValue); + } + } else { + auto offsetIntType = cast(offsetValue.getType()); + if (offsetIntType.getWidth() != 64) { + offsetValue = rewriter.create( + op.getLoc(), rewriter.getIntegerType(64), offsetValue); + } + } + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "[parseAddPtr] Adding offset\n"; + os << ptrOffsetInfo.getOffset() << '\n' << offsetValue << '\n'; + }); + Value offset = rewriter.create( + op.getLoc(), ptrOffsetInfo.getOffset(), offsetValue); + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "[parseAddPtr] offset is\n" << offset << '\n'; + }); + // Set addPtr offset map + auto dst = op.getResult(); + auto dstOffsetInfo = combineInfo(ptrOffsetInfo, offsetOffsetInfo); + dstOffsetInfo.setPtr(ptrOffsetInfo.getPtr()); + dstOffsetInfo.setOffset(offset); + offsetMap[dst] = dstOffsetInfo; + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + SmallVector &ptrStructured = ptrOffsetInfo.getStructuredRef(); + SmallVector &offsetStructured = offsetOffsetInfo.getStructuredRef(); + os << "[parseAddPtr] ptrStructured: "; + for (size_t i = 0; i < ptrStructured.size(); i++) + os << ptrStructured[i]; + os << "\n"; + os << "[parseAddPtr] offsetStructured: "; + for (size_t i = 0; i < offsetStructured.size(); i++) + os << offsetStructured[i]; + os << "\n"; + }); +} + +void parseSplat(triton::SplatOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get splat src + auto src = op.getSrc(); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + auto dst = op.getResult(); + auto dstType = cast(dst.getType()); + PtrOffsetInfo dstOffsetInfo(srcOffsetInfo.getPtr()); + // Modify IR + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "[parseSplat] dst is\n" << dst << '\n'; + }); + if (isa(dstType.getElementType())) { + RewriterBase::InsertionGuard guard(rewriter); + auto dstShape = dstType.getShape(); + rewriter.setInsertionPoint(op); + Value valueOffset = srcOffsetInfo.getOffset(); + Value offset = rewriter.create( + loc, RankedTensorType::get(dstShape, rewriter.getIntegerType(64)), + valueOffset); + dstOffsetInfo.setOffset(offset); + } + // Set addPtr offset map + + dstOffsetInfo.setStructured(dstType.getRank()); + dstOffsetInfo.setScalarLike(true); + offsetMap[dst] = dstOffsetInfo; +} + +template +void parseBinaryOp(BinOpTy op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + auto lhs = op.getLhs(); + parse(lhs, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo lhsOffsetInfo = offsetMap.at(lhs); + SmallVector &lhsStructured = lhsOffsetInfo.getStructuredRef(); + auto rhs = op.getRhs(); + parse(rhs, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo rhsOffsetInfo = offsetMap.at(rhs); + SmallVector &rhsStructured = rhsOffsetInfo.getStructuredRef(); + auto dst = op->getResult(0); + PtrOffsetInfo dstOffsetInfo; + dstOffsetInfo.setScalarLike(lhsOffsetInfo.isScalarLike() && + rhsOffsetInfo.isScalarLike()); + if (dstOffsetInfo.isScalarLike()) + dstOffsetInfo.setStructured(lhsStructured.size()); + else + dstOffsetInfo.setUnstructured(lhsStructured.size()); + offsetMap[dst] = dstOffsetInfo; +} + +void parseAddI(arith::AddIOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get addi lhs + auto lhs = op.getLhs(); + parse(lhs, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo lhsOffsetInfo = offsetMap.at(lhs); + // Get addi rhs + auto rhs = op.getRhs(); + parse(rhs, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo rhsOffsetInfo = offsetMap.at(rhs); + // Set addi offset map + auto dst = op.getResult(); + offsetMap[dst] = combineInfo(lhsOffsetInfo, rhsOffsetInfo); +} + +void parseSubI(arith::SubIOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get addi lhs + auto lhs = op.getLhs(); + parse(lhs, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo lhsOffsetInfo = offsetMap.at(lhs); + // Get addi rhs + auto rhs = op.getRhs(); + parse(rhs, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo rhsOffsetInfo = offsetMap.at(rhs); + // Set addi offset map + auto dst = op.getResult(); + offsetMap[dst] = combineInfo(lhsOffsetInfo, rhsOffsetInfo); + if (!(lhsOffsetInfo.isStructured() && rhsOffsetInfo.isScalarLike())) { + offsetMap[dst].setUnstructured(offsetMap[dst].getRank()); + } +} + +void parseIndexCast(arith::IndexCastOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get indexCast input + auto src = op.getIn(); + parse(src, op.getLoc(), rewriter, offsetMap); + // Set indexCast offset map + auto dst = op.getOut(); + offsetMap[dst] = offsetMap.at(src); +} + +template +void parseConstantOp(ConstOpTy dst, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Set constant offset map + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(true); + if (auto tensorType = dyn_cast(dst->getResult(0).getType())) + offsetMap[dst].setStructured(tensorType.getRank()); +} + +void parseMakeRange(triton::MakeRangeOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Set makeRange offset map + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setStructured(1); +} + +void parseExtSI(arith::ExtSIOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get extSI input + auto src = op.getIn(); + parse(src, op.getLoc(), rewriter, offsetMap); + // Set extSI offset map + auto dst = op.getOut(); + offsetMap[dst] = offsetMap.at(src); +} + +void parseBitcast(triton::BitcastOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get bitcast src + auto src = op.getSrc(); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + SmallVector &srcStructured = srcOffsetInfo.getStructuredRef(); + // Set extSI offset map + auto dst = op.getResult(); + if (auto ptr = srcOffsetInfo.getPtr()) { + Type ptrType = dst.getType(); + if (auto tensorType = dyn_cast(ptrType)) + ptrType = tensorType.getElementType(); + rewriter.setInsertionPoint(op); + ptr = rewriter.create(loc, ptrType, ptr); + offsetMap[dst] = + PtrOffsetInfo(ptr, srcOffsetInfo.getOffset(), srcStructured); + } else { + offsetMap[dst] = PtrOffsetInfo(srcStructured); + } + offsetMap[dst].setScalarLike(srcOffsetInfo.isScalarLike()); +} + +void parseLoad(triton::LoadOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get load ptr + auto ptr = op.getPtr(); + parse(ptr, op.getLoc(), rewriter, offsetMap); + // Set load offset map + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(offsetMap[ptr].isScalarLike()); + auto tensorType = dyn_cast(dst.getType()); + if (!tensorType) + return; + offsetMap[dst].setUnstructured(tensorType.getRank()); +} + +void parseMulI(arith::MulIOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get muli lhs + auto lhs = op.getLhs(); + parse(lhs, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo lhsOffsetInfo = offsetMap.at(lhs); + SmallVector &lhsStructured = lhsOffsetInfo.getStructuredRef(); + bool lhsScalarLike = lhsOffsetInfo.isScalarLike(); + // Get muli rhs + auto rhs = op.getRhs(); + parse(rhs, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo rhsOffsetInfo = offsetMap.at(rhs); + SmallVector &rhsStructured = rhsOffsetInfo.getStructuredRef(); + bool rhsScalarLike = rhsOffsetInfo.isScalarLike(); + // Set muli offset map + size_t maxSize = std::max(lhsStructured.size(), rhsStructured.size()); + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(lhsScalarLike && rhsScalarLike); + SmallVector &dstStructured = offsetMap[dst].getStructuredRef(); + dstStructured.resize(maxSize); + for (size_t i = 0; i < maxSize; i++) + if (lhsScalarLike) + dstStructured[i] = rhsStructured[i]; + else if (rhsScalarLike) + dstStructured[i] = lhsStructured[i]; + else + dstStructured[i] = false; +} + +void parseBroadcast(triton::BroadcastOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get broadcast src + auto src = op.getSrcMutable().get(); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + SmallVector &srcStructured = srcOffsetInfo.getStructuredRef(); + // Get broadcast dim + auto dst = op.getResult(); + assert(isa(src.getType()) && + "tt.broadcast's input should be a tensor"); + auto srcType = cast(src.getType()); + auto dstType = cast(dst.getType()); + assert(srcType.getRank() == dstType.getRank() && + "rank of source shoule be equal to destnation"); + auto broadcastDim = ConverterUtils::getBroadcastDims(srcType, dstType); + // Set broadcast offset map + offsetMap[dst] = PtrOffsetInfo(srcOffsetInfo.getPtr()); + offsetMap[dst].setScalarLike(srcOffsetInfo.isScalarLike()); + + if (srcOffsetInfo.getPtr()) { + RewriterBase::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(op); + Value valueOffset = srcOffsetInfo.getOffset(); + Value offset = rewriter.create( + loc, + RankedTensorType::get(dstType.getShape(), rewriter.getIntegerType(64)), + valueOffset); + + offsetMap[dst].setOffset(offset); + } + + SmallVector &dstStructured = offsetMap[dst].getStructuredRef(); + dstStructured.resize(srcStructured.size()); + for (size_t i = 0; i < dstStructured.size(); i++) + if (llvm::find(broadcastDim, i) != broadcastDim.end()) + dstStructured[i] = true; + else + dstStructured[i] = srcStructured[i]; +} + +void parseExpandDims(triton::ExpandDimsOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get expandDims src + auto src = op.getSrcMutable().get(); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + SmallVector &srcStructured = srcOffsetInfo.getStructuredRef(); + // Set expandDims offset map + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(srcOffsetInfo.getPtr()); + offsetMap[dst].setScalarLike(srcOffsetInfo.isScalarLike()); + if (srcOffsetInfo.getPtr()) { + RewriterBase::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(op); + Value valueOffset = srcOffsetInfo.getOffset(); + Value offset = rewriter.create(loc, valueOffset, + op.getAxisAttr()); + + offsetMap[dst].setOffset(offset); + } + SmallVector &dstStructured = offsetMap[dst].getStructuredRef(); + dstStructured.resize(srcStructured.size() + 1); + size_t j = 0; + for (size_t i = 0; i < dstStructured.size(); i++) + if (i == op.getAxis()) { + dstStructured[i] = true; + } else { + dstStructured[i] = srcStructured[j]; + j++; + } +} + +void parseClampF(triton::ClampFOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get clampF src + auto src = op.getX(); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + // Get clampF min + auto clampMin = op.getX(); + parse(clampMin, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo minOffsetInfo = offsetMap.at(clampMin); + // Get clampF max + auto clampMax = op.getX(); + parse(clampMax, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo maxOffsetInfo = offsetMap.at(clampMax); + // Set clampF offset map + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(srcOffsetInfo.isScalarLike() && + minOffsetInfo.isScalarLike() && + maxOffsetInfo.isScalarLike()); + auto dstType = dyn_cast(dst.getType()); + if (!dstType) + return; + offsetMap[dst].setUnstructured(dstType.getRank()); +} + +void parseSelect(arith::SelectOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get select condition + auto condition = op.getCondition(); + parse(condition, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo conditionOffsetInfo = offsetMap.at(condition); + bool conditionScalarLike = conditionOffsetInfo.isScalarLike(); + // Get select trueValue + auto trueValue = op.getTrueValue(); + parse(trueValue, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo trueValueOffsetInfo = offsetMap.at(trueValue); + SmallVector &trueValueStructured = + trueValueOffsetInfo.getStructuredRef(); + bool trueValueScalarLike = trueValueOffsetInfo.isScalarLike(); + // Get select falseValue + auto falseValue = op.getFalseValue(); + parse(falseValue, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo falseValueOffsetInfo = offsetMap.at(falseValue); + SmallVector &falseValueStructured = + falseValueOffsetInfo.getStructuredRef(); + bool falseValueScalarLike = falseValueOffsetInfo.isScalarLike(); + // Set select offset map + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(); + auto dstType = dyn_cast(dst.getType()); + if (!dstType) + return; + offsetMap[dst].setUnstructured(dstType.getRank()); +} + +void parseFPToSI(arith::FPToSIOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get FPToSI src + auto src = op.getIn(); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + // Set FPToSI offset map + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(srcOffsetInfo.isScalarLike()); + auto dstType = dyn_cast(dst.getType()); + if (!dstType) + return; + if (offsetMap[dst].isScalarLike()) + offsetMap[dst].setStructured(dstType.getRank()); + else + offsetMap[dst].setUnstructured(dstType.getRank()); +} + +void parseSIToFP(arith::SIToFPOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get SIToFP src + auto src = op.getIn(); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + // Set SIToFP offset map + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(srcOffsetInfo.isScalarLike()); + auto dstType = dyn_cast(dst.getType()); + if (!dstType) + return; + if (offsetMap[dst].isScalarLike()) + offsetMap[dst].setStructured(dstType.getRank()); + else + offsetMap[dst].setUnstructured(dstType.getRank()); +} + +// FIXME:Z|wait triton version upgrade to 3.4 +// void parseMakeTensorDesc(triton::MakeTensorDescOp op, const Location &loc, +// RewriterBase &rewriter, +// llvm::DenseMap &offsetMap) { +// // Set MakeTensorDesc offset map +// auto dst = op.getResult(); +// offsetMap[dst] = PtrOffsetInfo(); +// auto dstType = dyn_cast(dst.getType()); +// if (!dstType) +// return; +// offsetMap[dst].setStructured(dstType.getRank()); +// } + +void parseMakeTensorPtr(triton::MakeTensorPtrOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Set MakeTensorPtr offset map + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(dst); + auto dstType = dyn_cast( + cast(dst.getType()).getPointeeType()); + if (!dstType) + return; + offsetMap[dst].setStructured(dstType.getRank()); + offsetMap[dst].setOffsets(op.getOffsets()); +} + +void parseAdvance(triton::AdvanceOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Set Advance offset map + auto ptr = op.getPtr(); + parse(ptr, op.getLoc(), rewriter, offsetMap); + auto dst = op.getResult(); + offsetMap[dst] = offsetMap.at(ptr); + auto dstType = dyn_cast( + cast(dst.getType()).getPointeeType()); + if (!dstType) + return; + offsetMap[dst].setStructured(dstType.getRank()); + auto &offsets = offsetMap[dst].getOffsetsRef(); + + RewriterBase::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(op); + for (auto [curOffset, opOffset] : llvm::zip(offsets, op.getOffsets())) { + curOffset = + rewriter.create(op.getLoc(), curOffset, opOffset); + } +} + +void parseReduce(triton::ReduceOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get reduce src + Value src = op->getOperand(0); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + SmallVector &srcStructured = srcOffsetInfo.getStructuredRef(); + // Set reduce offset map + Value dst = op->getResult(0); + auto dstType = dyn_cast(dst.getType()); + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(srcOffsetInfo.isScalarLike()); + if (!dstType) + return; + SmallVector &dstStructured = offsetMap[dst].getStructuredRef(); + auto dstShape = dstType.getShape(); + dstStructured.resize(dstShape.size()); + for (size_t i = 0; i < dstStructured.size(); i++) + if (dstShape[i] == 1) + dstStructured[i] = true; + else + dstStructured[i] = srcStructured[i]; +} + +void parseReduceReturn(triton::ReduceReturnOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get reduce src + Value src = op->getOperand(0); + parse(src, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo srcOffsetInfo = offsetMap.at(src); + SmallVector &srcStructured = srcOffsetInfo.getStructuredRef(); + // Set reduce offset map + Value dst = op->getResult(0); + auto dstType = dyn_cast(dst.getType()); + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(srcOffsetInfo.isScalarLike()); + if (!dstType) + return; + SmallVector &dstStructured = offsetMap[dst].getStructuredRef(); + auto dstShape = dstType.getShape(); + dstStructured.resize(dstShape.size()); + for (size_t i = 0; i < dstStructured.size(); i++) + if (dstShape[i] == 1) + dstStructured[i] = true; + else + dstStructured[i] = srcStructured[i]; +} + +void parseIf(scf::IfOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap, Value dst) { + const unsigned int index = cast(dst).getResultNumber(); + // Get if then region + Block &thenBlock = op.getThenRegion().front(); + Value thenYieldedValue = thenBlock.getTerminator()->getOperand(index); + parse(thenYieldedValue, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo thenOffsetInfo = offsetMap.at(thenYieldedValue); + SmallVector &thenStructured = thenOffsetInfo.getStructuredRef(); + // Get if else region + bool dstIsScalar = thenOffsetInfo.isScalarLike(); + SmallVector elseStructured; + if (op.elseBlock()) { + Block &elseBlock = op.getElseRegion().front(); + Value elseYieldedValue = elseBlock.getTerminator()->getOperand(index); + parse(elseYieldedValue, op.getLoc(), rewriter, offsetMap); + PtrOffsetInfo elseOffsetInfo = offsetMap.at(elseYieldedValue); + elseStructured = elseOffsetInfo.getStructuredRef(); + dstIsScalar = dstIsScalar && elseOffsetInfo.isScalarLike(); + } + // Set if offset map + offsetMap[dst] = PtrOffsetInfo(); + offsetMap[dst].setScalarLike(dstIsScalar); + SmallVector &dstStructured = offsetMap[dst].getStructuredRef(); + dstStructured.resize(thenStructured.size()); + for (size_t i = 0; i < dstStructured.size(); i++) + if (op.elseBlock()) + dstStructured[i] = thenStructured[i] && elseStructured[i]; + else + dstStructured[i] = thenStructured[i]; +} + +void parseYield(scf::YieldOp op, const Location &loc, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get yield src + for (auto src : op->getOperands()) + parse(src, op.getLoc(), rewriter, offsetMap); +} + +void parseLoopOp(LoopLikeOpInterface op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap, Value dst) { + auto resNum = cast(dst).getResultNumber(); + Value yieldedValue = nullptr; + if (auto whileOp = dyn_cast(op.getOperation())) { + yieldedValue = whileOp.getConditionOp().getArgs()[resNum]; + } else { + yieldedValue = op.getYieldedValues()[resNum]; + } + parse(yieldedValue, op.getLoc(), rewriter, offsetMap); + offsetMap[dst] = offsetMap.at(yieldedValue); +} + +void parseExtractSlice(tensor::ExtractSliceOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + // Get extractSlice src + auto src = op.getOperand(0); + parse(src, op.getLoc(), rewriter, offsetMap); + // Set extractSlice offset map + auto dst = op.getResult(); + offsetMap[dst] = offsetMap.at(src); +} + +void parseExtract(tensor::ExtractOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + auto parentValue = op.getTensor(); + parse(parentValue, op.getLoc(), rewriter, offsetMap); + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(); + if (isa(dst.getType())) { + offsetMap[dst].setPtr(dst); + } + offsetMap[dst].setScalarLike(true); +} + +void parseIntToPtr(triton::IntToPtrOp op, const Location &loc, + RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + auto dst = op.getResult(); + offsetMap[dst] = PtrOffsetInfo(dst); + offsetMap[dst].setScalarLike(true); +} + +} // namespace triton +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.cpp new file mode 100755 index 00000000..b181d58c --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.cpp @@ -0,0 +1,964 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.h" +#include "incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "bishengir/Dialect/Annotation/IR/Annotation.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "mlir/Transforms/Passes.h" + +#include "llvm/ADT/STLExtras.h" + +#define DEBUG_TYPE "triton-unstructure-converter" + +using namespace mlir; +using namespace triton; + +#include "llvm/Support/Debug.h" + +bool forceSimtTemplateFlag = false; + +template +bool UnstructuredMemAccessConverter::checkUnstructureAnnotated( + MemAccOpTy op, PatternRewriter &rewriter) const { + return llvm::any_of(op->getUsers(), [&rewriter](Operation *user) { + auto annotationOp = dyn_cast(user); + if (annotationOp && annotationOp->hasAttr("mayDiscretememaccess")) { + rewriter.eraseOp(annotationOp); + return true; + } + return false; + }); +} + +template <> +bool UnstructuredMemAccessConverter::checkUnstructureAnnotated( + triton::StoreOp op, PatternRewriter &rewriter) const { + return llvm::any_of(op.getValue().getUsers(), [&rewriter](Operation *user) { + auto annotationOp = dyn_cast(user); + if (annotationOp && annotationOp->hasAttr("mayDiscretememaccess")) { + rewriter.eraseOp(annotationOp); + return true; + } + return false; + }); +} + +template +Value UnstructuredMemAccessConverter::createExtractOp( + Location loc, Value value, PatternRewriter &rewriter, + ArrayRef iterIdx) const { + if (!value) + return value; + SmallVector indices; + for (auto idx : iterIdx) { + if (auto val = dyn_cast(idx)) { + indices.push_back(val); + } else { + auto idxVal = rewriter.create( + loc, rewriter.getIndexAttr(*getConstantIntValue(idx))); + indices.push_back(idxVal); + } + } + auto extractedOp = rewriter.create(loc, value, indices); + extractedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + return extractedOp; +} + +template +Value UnstructuredMemAccessConverter::createExtractOp( + Location loc, Value value, PatternRewriter &rewriter, + ArrayRef offsets, ArrayRef sizes, + ArrayRef strides) const { + if (!value) + return value; + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Extracting\n"; + os << value << "\n"; + }); + auto extractedOp = rewriter.create( + loc, value, offsets, sizes, strides); + extractedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + return extractedOp; +} + +template <> +template +triton::LoadOp UnstructuredMemAccessConverter::createMemAccOp( + triton::LoadOp op, Value ptrToAccess, Location loc, + PatternRewriter &rewriter, Args &&...args) const { + return rewriter.create(loc, ptrToAccess, op.getCache(), + op.getEvict(), op.getIsVolatile()); +} + +template <> +template +triton::AtomicRMWOp +UnstructuredMemAccessConverter::createMemAccOp( + triton::AtomicRMWOp op, Value ptrToAccess, Location loc, + PatternRewriter &rewriter, Args &&...args) const { + auto extractedValue = + createExtractOp(loc, op.getVal(), rewriter, std::forward(args)...); + auto extractedMask = + createExtractOp(loc, op.getMask(), rewriter, std::forward(args)...); + Type targetType = ptrToAccess.getType(); + if (auto tensorType = dyn_cast(targetType)) { + auto ptrType = cast(tensorType.getElementType()); + targetType = + RankedTensorType::get(tensorType.getShape(), ptrType.getPointeeType()); + } else { + auto resultType = cast(op.getResult().getType()); + SmallVector scalarLikeShape(resultType.getRank(), 1); + targetType = + RankedTensorType::get(scalarLikeShape, resultType.getElementType()); + ptrToAccess = rewriter.create( + loc, RankedTensorType::get(scalarLikeShape, ptrToAccess.getType()), + ptrToAccess); + extractedValue = rewriter.create( + loc, RankedTensorType::get(scalarLikeShape, extractedValue.getType()), + extractedValue); + if (extractedMask) { + extractedMask = rewriter.create( + loc, RankedTensorType::get(scalarLikeShape, extractedMask.getType()), + extractedMask); + } + } + return rewriter.create( + loc, targetType, op.getAtomicRmwOpAttr(), ptrToAccess, extractedValue, + extractedMask, op.getSemAttr(), op.getScopeAttr()); +} + +template <> +template +triton::AtomicCASOp +UnstructuredMemAccessConverter::createMemAccOp( + triton::AtomicCASOp op, Value ptrToAccess, Location loc, + PatternRewriter &rewriter, Args &&...args) const { + auto extractedCmp = + createExtractOp(loc, op.getCmp(), rewriter, std::forward(args)...); + auto extractedValue = + createExtractOp(loc, op.getVal(), rewriter, std::forward(args)...); + Type targetType = ptrToAccess.getType(); + if (auto tensorType = dyn_cast(targetType)) { + auto ptrType = cast(tensorType.getElementType()); + targetType = + RankedTensorType::get(tensorType.getShape(), ptrType.getPointeeType()); + } else { + auto resultType = cast(op.getResult().getType()); + SmallVector scalarLikeShape(resultType.getRank(), 1); + targetType = + RankedTensorType::get(scalarLikeShape, resultType.getElementType()); + ptrToAccess = rewriter.create( + loc, RankedTensorType::get(scalarLikeShape, ptrToAccess.getType()), + ptrToAccess); + extractedCmp = rewriter.create( + loc, RankedTensorType::get(scalarLikeShape, extractedCmp.getType()), + extractedCmp); + extractedValue = rewriter.create( + loc, RankedTensorType::get(scalarLikeShape, extractedValue.getType()), + extractedValue); + } + return rewriter.create( + loc, targetType, ptrToAccess, extractedCmp, extractedValue, + op.getSemAttr(), op.getScopeAttr()); +} + +template <> +template +triton::StoreOp UnstructuredMemAccessConverter::createMemAccOp( + triton::StoreOp op, Value ptrToAccess, Location loc, + PatternRewriter &rewriter, Args &&...args) const { + auto extractedValue = createExtractOp(loc, op.getValue(), rewriter, + std::forward(args)...); + auto extractedMask = + createExtractOp(loc, op.getMask(), rewriter, std::forward(args)...); + return rewriter.create(loc, ptrToAccess, extractedValue, + extractedMask); +} + +template <> +template <> +void UnstructuredMemAccessConverter::splatAndLoadScenario< + triton::LoadOp>(triton::LoadOp op, int rank, + PatternRewriter &rewriter) const { + auto loc = op.getLoc(); + SmallVector idx(rank, rewriter.getIndexAttr(0)); + auto extractedPtr = createExtractOp(loc, op.getPtr(), rewriter, idx); + Value mask = op.getMask(); + Value other = op.getOther(); + Value loadedValue = rewriter.create( + loc, extractedPtr, /*mask=*/nullptr, /*other=*/nullptr, + /*boundaryCheck=*/ArrayRef(), + /*PaddingOptionAttr=*/nullptr); + loadedValue = rewriter.create(loc, op.getResult().getType(), + loadedValue); + if (mask) + rewriter.replaceOpWithNewOp(op, mask, loadedValue, other); + else + rewriter.replaceOp(op, loadedValue); +} + +template +UnstructuredMemAccessConverter::UnstructuredMemAccessConverter( + MLIRContext *context, bool forceScalarizeMode, + const llvm::DenseMap &offsetMap, + const llvm::SmallDenseMap &fromTensorArg) + : OpRewritePattern(context), + forceScalarizeMode(forceScalarizeMode), offsetMap(offsetMap), + fromTensorArg(fromTensorArg) {} + +template +LogicalResult UnstructuredMemAccessConverter::matchAndRewrite( + MemAccOpTy op, PatternRewriter &rewriter) const { + auto loc = op.getLoc(); + + auto ptr = op.getPtr(); + auto ptrType = dyn_cast(ptr.getType()); + + if (auto ptrPtrType = dyn_cast(ptr.getType())) { + if (auto ptrTensorType = + dyn_cast_or_null(ptrPtrType.getPointeeType())) + ptrType = ptrTensorType; + } + + if (!ptrType || op->hasAttr(ConverterUtils::discreteAttrName)) + return failure(); + if (!offsetMap.contains(ptr)) + return op.emitError() << "PtrOffsetInfo should be computed\n" << ptr; + + auto ptrOffsetInfo = offsetMap.at(ptr); + + if (checkUnstructureAnnotated(op, rewriter)) + ptrOffsetInfo.setUnstructured(ptrOffsetInfo.getRank()); + + if (ptrOffsetInfo.isStructured() && + (!ptrOffsetInfo.isScalarLike() || + llvm::all_of(ptrType.getShape(), [](int64_t dim) { return dim == 1; }))) + return failure(); + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Converting " << op->getName() << "\n"; + os << op << "\n"; + os << ptrOffsetInfo.isStructured() << "\n"; + os << ptrOffsetInfo.isScalarLike() << "\n"; + }); + + if constexpr (std::is_same_v) { + if (ptrOffsetInfo.isScalarLike()) { + splatAndLoadScenario(op, ptrOffsetInfo.getRank(), rewriter); + return success(); + } + } + + std::optional mstate = + Incubated::runMaskAnalysis(op, static_cast(rewriter)); + + if (op->hasAttr(ConverterUtils::discreteMaskAttrName)) { + if constexpr (std::is_same_v) { + auto selectOp = op.getValue().template getDefiningOp(); + op = rewriter.replaceOpWithNewOp( + op, op.getPtr(), selectOp.getTrueValue(), selectOp.getCondition(), + op.getCache(), op.getEvict()); + rewriter.setInsertionPoint(op); + ptrOffsetInfo.setUnstructured(ptrOffsetInfo.getRank()); + } else if constexpr (std::is_same_v) { + auto selectOp = op.getVal().template getDefiningOp(); + op = rewriter.replaceOpWithNewOp( + op, op.getType(), op.getAtomicRmwOp(), op.getPtr(), + selectOp.getTrueValue(), selectOp.getCondition(), op.getSem(), + op.getScope()); + } + rewriter.setInsertionPoint(op); + ptrOffsetInfo.setUnstructured(ptrOffsetInfo.getRank()); + } + + if (forceScalarizeMode || ptrOffsetInfo.isScalarLike() || + fromTensorArg.at(ptr)) { + ptrOffsetInfo.setUnstructured(ptrOffsetInfo.getRank()); + } + + auto srcPtr = ptrOffsetInfo.getPtr(); + auto ptrOffset = ptrOffsetInfo.getOffset(); + + // LoadLike is operation with result + bool isLoadLike = !op->use_empty(); + + Value zeroIdx = + rewriter.create(loc, rewriter.getIndexAttr(0)); + Value oneIdx = + rewriter.create(loc, rewriter.getIndexAttr(1)); + auto resultShape = ptrType.getShape(); + auto resultElementType = ptrType.getElementType(); + if (auto pointerType = + dyn_cast(ptrType.getElementType())) { + resultElementType = pointerType.getPointeeType(); + } + + int64_t sizeInByte; + if (auto intType = dyn_cast(resultElementType)) { + sizeInByte = intType.getWidth() / 8; + } else if (auto floatType = dyn_cast(resultElementType)) { + sizeInByte = floatType.getWidth() / 8; + } else { + llvm_unreachable("Unhandled element type of tensor"); + } + + for (int i = ptrOffsetInfo.getRank() - 1; i >= 0; i--) { + if (!ptrOffsetInfo.isStructured(i)) + break; + sizeInByte *= resultShape[i]; + } + + // Force scalarize if memory is not aligned + if (sizeInByte % 32 != 0) + ptrOffsetInfo.setUnstructured(ptrOffsetInfo.getRank()); + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "UnStructured Flag check:\n"; + os << "ptrOffsetInfo.isStructured: " << ptrOffsetInfo.isStructured() + << "\n"; + os << "compileOn91095Flag: " << compileOn91095Flag << "\n"; + os << "forceSimtTemplateFlag: " << forceSimtTemplateFlag << "\n"; + }); + + // Fast path on A5: rewrite tt.load/store to tt.indirect_load/store directly. + if (compileOn91095Flag && forceSimtTemplateFlag && + !ptrOffsetInfo.isStructured()) { + if constexpr (std::is_same_v) { + assert(isa(srcPtr.getType()) && + "src must be ptr type"); + Value mask = op.getMask(); + Value other = op.getOther(); + auto resultType = op.getType(); + auto indirect = rewriter.create( + loc, resultType, srcPtr, ptrOffset, mask, other); + rewriter.replaceOp(op, indirect.getResult()); + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Rewriting tt.load to tt.indirect_load\n"; + os << indirect << "\n"; + }); + return success(); + } else if constexpr (std::is_same_v) { + assert(isa(srcPtr.getType()) && + "src must be ptr type"); + Value value = op.getValue(); + Value mask = op.getMask(); + auto indirect = rewriter.create( + loc, srcPtr, ptrOffset, value, mask); + rewriter.eraseOp(op); + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Rewriting tt.store to tt.indirect_store\n"; + os << indirect << "\n"; + }); + return success(); + } + } + + Value iterArg = nullptr; + + // Only load case + if (isLoadLike) { + iterArg = + rewriter.create(loc, resultShape, resultElementType); + } + Value newOpResult = nullptr; + + auto insertPoint = rewriter.saveInsertionPoint(); + + SmallVector offsets; + SmallVector sizes; + SmallVector strides; + SmallVector extractedShape; + + for (size_t i = 0; i < resultShape.size(); i++) { + auto size = resultShape[i]; + auto structured = ptrOffsetInfo.getStructuredRef()[i]; + // handle indirect dimension + strides.push_back(rewriter.getIndexAttr(1)); + Value sizeVal = + rewriter.create(loc, rewriter.getIndexAttr(size)); + if (structured) { + offsets.push_back(rewriter.getIndexAttr(0)); + sizes.push_back(rewriter.getIndexAttr(size)); + extractedShape.push_back(size); + } else { + scf::ForOp forOp; + if (auto mtptOp = + srcPtr.template getDefiningOp()) { + auto tptShape = mtptOp.getShape()[i]; + if (tptShape.getType() != rewriter.getIndexType()) { + tptShape = rewriter.create( + loc, rewriter.getIndexType(), tptShape); + } + sizeVal = rewriter.create(loc, sizeVal, tptShape); + } else if (mstate) { + sizeVal = + getValueOrCreateConstantIndexOp(rewriter, loc, mstate->dims[i]); + } + if (isLoadLike) { + forOp = rewriter.create(loc, zeroIdx, sizeVal, oneIdx, + ValueRange({iterArg})); + if (!newOpResult) { + newOpResult = forOp->getResult(0); + } else { + rewriter.create(loc, forOp->getResult(0)); + } + iterArg = forOp.getRegionIterArg(0); + } else { + forOp = rewriter.create(loc, zeroIdx, sizeVal, oneIdx); + } + sizes.push_back(rewriter.getIndexAttr(1)); + offsets.push_back(forOp.getInductionVar()); + extractedShape.push_back(1); + forOp->setAttr("ExtractedLoadOrStore", + UnitAttr::get(rewriter.getContext())); + rewriter.setInsertionPointToStart(forOp.getBody()); + } + } + + bool fullyUnstructured = ptrOffsetInfo.isUnstructured(); + auto extractedType = RankedTensorType::get(extractedShape, resultElementType); + + Value extractedOffset; + if (fullyUnstructured) { + if (auto mtptOp = + srcPtr.template getDefiningOp()) { + auto I64Type = rewriter.getIntegerType(64); + srcPtr = mtptOp.getBase(); + extractedOffset = rewriter.create(loc, 0, 64); + for (auto [indVar, offset, stride] : llvm::zip_equal( + offsets, ptrOffsetInfo.getOffsets(), mtptOp.getStrides())) { + Value inductionVar = rewriter.create( + loc, I64Type, cast(indVar)); + Value tptOffset = rewriter.create(loc, I64Type, offset); + Value tptStride = rewriter.create(loc, I64Type, stride); + tptOffset = rewriter.create(loc, tptStride, tptOffset); + tptStride = + rewriter.create(loc, tptStride, inductionVar); + extractedOffset = + rewriter.create(loc, extractedOffset, tptOffset); + extractedOffset = + rewriter.create(loc, extractedOffset, tptStride); + } + } else { + extractedOffset = createExtractOp(loc, ptrOffset, rewriter, offsets); + } + } else { + extractedOffset = + createExtractOp(loc, ptrOffset, rewriter, offsets, sizes, strides); + } + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Extracted offset\n"; + os << extractedOffset << "\n"; + }); + + assert(isa(srcPtr.getType()) && "src must be ptr type"); + if (!fullyUnstructured) { + srcPtr = rewriter.create( + loc, RankedTensorType::get(extractedShape, srcPtr.getType()), srcPtr); + } + Value ptrToAccess = rewriter.create( + loc, srcPtr.getType(), srcPtr, extractedOffset); + + MemAccOpTy accessedOp; + if (fullyUnstructured) { + accessedOp = createMemAccOp(op, ptrToAccess, loc, rewriter, offsets); + } else { + accessedOp = + createMemAccOp(op, ptrToAccess, loc, rewriter, offsets, sizes, strides); + } + + accessedOp->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + + if (isLoadLike) { + assert(iterArg && "Load case must have iterArg in for loop"); + + Value value = accessedOp->getResult(0); + Value result; + if (!isa(value.getType()) && + (std::is_same_v || + std::is_same_v)) { + value = rewriter.create(loc, extractedType, value); + } + if (!isa(value.getType())) { + SmallVector indices; + for (auto idx : offsets) { + if (auto val = dyn_cast(idx)) { + indices.push_back(val); + } else { + auto idxVal = rewriter.create( + loc, rewriter.getIndexAttr(*getConstantIntValue(idx))); + indices.push_back(idxVal); + } + } + result = rewriter.create(loc, value, iterArg, indices); + } else { + result = rewriter.create(loc, value, iterArg, + offsets, sizes, strides); + } + rewriter.create(loc, result) + ->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + rewriter.restoreInsertionPoint(insertPoint); + if constexpr (std::is_same_v) { + if (op.getMask() && op.getOther()) { + rewriter + .replaceOpWithNewOp(op, op.getMask(), newOpResult, + op.getOther()) + ->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + } else { + rewriter.replaceOp(op, newOpResult); + } + } else { + rewriter.replaceOp(op, newOpResult); + } + } else { + if constexpr (std::is_same_v) { + if (fullyUnstructured && accessedOp.getMask()) { + auto mask = createExtractOp( + loc, accessedOp.getMask(), rewriter, + SmallVector(ptrOffsetInfo.getRank(), + rewriter.getIndexAttr(0))); + rewriter.create(loc, mask, [&](OpBuilder &b, Location loc) { + b.create( + loc, accessedOp.getType(), accessedOp.getAtomicRmwOp(), + accessedOp.getPtr(), accessedOp.getVal(), nullptr, + accessedOp.getSem(), accessedOp.getScope()) + ->setAttr(ConverterUtils::discreteAttrName, + UnitAttr::get(rewriter.getContext())); + b.create(loc); + }); + rewriter.eraseOp(accessedOp); + } + } + rewriter.eraseOp(op); + } + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "After conversion\n" + << ptrToAccess.getDefiningOp() + ->template getParentOfType() + << "\n"; + }); + return success(); +} + +void replaceOperands(MutableArrayRef oprs, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + for (auto it = oprs.begin(); it != oprs.end(); ++it) { + auto &opr = *it; + auto operand = opr.get(); + if (auto tensorType = dyn_cast(operand.getType()); + tensorType && isa(tensorType.getElementType())) { + parse(operand, operand.getLoc(), rewriter, offsetMap); + opr.set(offsetMap.at(operand).getOffset()); + } else if (auto ptrType = + dyn_cast(operand.getType())) { + parse(operand, operand.getLoc(), rewriter, offsetMap); + if (auto tensorType = + dyn_cast(ptrType.getPointeeType())) { + for (auto offset : offsetMap.at(operand).getOffsets()) { + it->set(offset); + ++it; + } + --it; + } else { + opr.set(offsetMap.at(operand).getOffset()); + } + } + } +} + +void replaceArgs(ValueRange args, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + for (auto it = args.begin(); it != args.end(); ++it) { + auto arg = *it; + if (auto tensorType = dyn_cast(arg.getType()); + tensorType && isa(tensorType.getElementType())) { + RewriterBase::InsertionGuard guard(rewriter); + if (auto blockArg = dyn_cast(arg)) { + rewriter.setInsertionPointToStart(blockArg.getOwner()); + } else { + rewriter.setInsertionPointAfterValue(arg); + } + auto tempVar = rewriter + .create( + arg.getLoc(), arg.getType(), ValueRange({})) + ->getResult(0); + parse(arg, arg.getLoc(), rewriter, offsetMap); + auto src = offsetMap.at(arg).getPtr(); + rewriter.replaceAllUsesWith(arg, tempVar); + arg.setType(RankedTensorType::get(tensorType.getShape(), + rewriter.getIntegerType(64))); + src = rewriter.create(arg.getLoc(), tempVar.getType(), + src); + rewriter.replaceOpWithNewOp( + tempVar.getDefiningOp(), tempVar.getType(), src, arg); + } else if (auto ptrType = dyn_cast(arg.getType())) { + RewriterBase::InsertionGuard guard(rewriter); + if (auto blockArg = dyn_cast(arg)) { + rewriter.setInsertionPointToStart(blockArg.getOwner()); + } else { + rewriter.setInsertionPointAfterValue(arg); + } + auto tempVar = rewriter + .create( + arg.getLoc(), arg.getType(), ValueRange({})) + ->getResult(0); + parse(arg, arg.getLoc(), rewriter, offsetMap); + rewriter.replaceAllUsesWith(arg, tempVar); + if (auto tensorType = + dyn_cast(ptrType.getPointeeType())) { + auto srcOp = + offsetMap.at(arg).getPtr().getDefiningOp(); + arg.setType(rewriter.getIntegerType(32)); + SmallVector newOffsets; + for (auto offset : offsetMap.at(arg).getOffsets()) { + newOffsets.push_back(*it); + ++it; + } + --it; + rewriter.replaceOpWithNewOp( + tempVar.getDefiningOp(), tempVar.getType(), srcOp.getBase(), + srcOp.getShape(), srcOp.getStrides(), newOffsets, srcOp.getOrder()); + } else { + auto src = offsetMap.at(arg).getPtr(); + arg.setType(rewriter.getIntegerType(64)); + rewriter.replaceOpWithNewOp( + tempVar.getDefiningOp(), tempVar.getType(), src, arg); + } + } + } +} + +void convertTensorPtrPre(LoopLikeOpInterface op, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + if (auto whileOp = dyn_cast(op.getOperation())) { + replaceArgs(whileOp.getBeforeArguments(), rewriter, offsetMap); + replaceOperands(whileOp.getInitsMutable(), rewriter, offsetMap); + replaceArgs(whileOp.getAfterArguments(), rewriter, offsetMap); + replaceArgs(whileOp->getResults(), rewriter, offsetMap); + replaceOperands(whileOp.getConditionOp().getArgsMutable(), rewriter, + offsetMap); + } else { + replaceArgs(op.getRegionIterArgs(), rewriter, offsetMap); + replaceOperands(op.getInitsMutable(), rewriter, offsetMap); + } +} + +void convertTensorPtrPost(LoopLikeOpInterface op, RewriterBase &rewriter, + llvm::DenseMap &offsetMap) { + if (auto whileOp = dyn_cast(op.getOperation())) { + replaceOperands(whileOp.getYieldOp()->getOpOperands(), rewriter, offsetMap); + } else { + replaceArgs(op->getResults(), rewriter, offsetMap); + replaceOperands(*op.getYieldedValuesMutable(), rewriter, offsetMap); + } +} + +int getPtrTensorRank(Type type) { + if (auto ptrType = dyn_cast(type)) { + if (auto tensorType = + dyn_cast(ptrType.getPointeeType())) { + return tensorType.getRank(); + } + } + return 0; +} + +SmallVector constructOperands(ValueRange operands, Value tempVar, + IRMapping mapping) { + SmallVector newOperands; + for (auto opr : operands) { + opr = mapping.lookupOrDefault(opr); + newOperands.push_back(opr); + auto numAppend = getPtrTensorRank(opr.getType()) - 1; + if (numAppend > 0) + newOperands.append(numAppend, tempVar); + } + return newOperands; +} + +SmallVector constructTypes(TypeRange types) { + SmallVector newTypes; + for (auto type : types) { + newTypes.push_back(type); + if (auto ptrType = dyn_cast(type)) { + if (auto tensorType = + dyn_cast(ptrType.getPointeeType())) { + if (tensorType.getRank() > 0) + newTypes.append(tensorType.getRank() - 1, + IntegerType::get(type.getContext(), 32)); + } + } + } + return newTypes; +} + +void replacePtrLoopArguments(Operation *rootOp, + llvm::DenseMap &offsetMap) { + std::function convertTensorPtr = + [&](LoopLikeOpInterface op) { + IRRewriter rewriter(op.getContext()); + IRMapping mapping; + LoopLikeOpInterface newOp; + rewriter.setInsertionPointAfter(op); + Value tempVar = + rewriter + .create( + op.getLoc(), rewriter.getI32Type(), ValueRange({})) + ->getResult(0); + if (auto forOp = dyn_cast(op.getOperation())) { + newOp = rewriter.create( + forOp.getLoc(), forOp.getLowerBound(), forOp.getUpperBound(), + forOp.getStep(), + constructOperands(forOp.getInitArgs(), tempVar, mapping), + [&](OpBuilder &b, Location loc, Value iv, ValueRange args) { + mapping.map(forOp.getInductionVar(), iv); + auto newArgIter = args.begin(); + for (auto oldArg : forOp.getRegionIterArgs()) { + mapping.map(oldArg, *newArgIter); + std::advance(newArgIter, + std::max(getPtrTensorRank(oldArg.getType()), 1)); + } + for (auto &bodyOp : forOp.getBody()->without_terminator()) { + b.clone(bodyOp, mapping); + } + auto yieldOp = + cast(forOp.getBody()->getTerminator()); + b.create( + yieldOp.getLoc(), + constructOperands(yieldOp.getOperands(), tempVar, mapping)); + }); + newOp->setAttrs(op->getAttrs()); + } else if (auto whileOp = dyn_cast(op.getOperation())) { + newOp = rewriter.create( + whileOp.getLoc(), constructTypes(whileOp->getResultTypes()), + constructOperands(whileOp.getInits(), tempVar, mapping), + [&](OpBuilder &b, Location loc, ValueRange args) { + auto newArgIter = args.begin(); + for (auto oldArg : whileOp.getBeforeArguments()) { + mapping.map(oldArg, *newArgIter); + std::advance(newArgIter, + std::max(getPtrTensorRank(oldArg.getType()), 1)); + } + for (auto &bodyOp : + whileOp.getBeforeBody()->without_terminator()) { + b.clone(bodyOp, mapping); + } + auto conditionOp = whileOp.getConditionOp(); + b.create( + conditionOp.getLoc(), + mapping.lookup(conditionOp.getCondition()), + constructOperands(conditionOp.getArgs(), tempVar, mapping)); + }, + [&](OpBuilder &b, Location loc, ValueRange args) { + auto newArgIter = args.begin(); + for (auto oldArg : whileOp.getAfterArguments()) { + mapping.map(oldArg, *newArgIter); + std::advance(newArgIter, + std::max(getPtrTensorRank(oldArg.getType()), 1)); + } + for (auto &bodyOp : + whileOp.getAfterBody()->without_terminator()) { + b.clone(bodyOp, mapping); + } + auto yieldOp = whileOp.getYieldOp(); + b.create( + yieldOp.getLoc(), + constructOperands(yieldOp.getOperands(), tempVar, mapping)); + }); + } else { + llvm_unreachable("Unsupported loop op"); + } + auto resIter = newOp->result_begin(); + for (auto res : op->getResults()) { + rewriter.replaceAllUsesWith(res, *resIter); + std::advance(resIter, std::max(getPtrTensorRank(res.getType()), 1)); + } + rewriter.eraseOp(op); + op = newOp; + convertTensorPtrPre(op, rewriter, offsetMap); + for (auto *region : op.getLoopRegions()) + region->walk(convertTensorPtr); + convertTensorPtrPost(op, rewriter, offsetMap); + return WalkResult::skip(); + }; + + rootOp->walk(convertTensorPtr); +} + +void TritonToUnstructureIncubatedPass::runPreparse(LoopLikeOpInterface op) { + IRRewriter rewriter(&getContext()); + auto loc = op.getLoc(); + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Pre-parsing " << op->getName() << "\n" << op << "\n"; + }); + + Block::BlockArgListType args; + ValueRange yields; + if (auto whileOp = dyn_cast(op.getOperation())) { + args = whileOp.getBeforeArguments(); + yields = whileOp.getYieldOp().getOperands(); + } else { + args = op.getRegionIterArgs(); + yields = op.getYieldedValues(); + } + + for (auto [arg, yield] : llvm::zip_equal(args, yields)) { + if (auto tensorType = dyn_cast(yield.getType())) { + parse(yield, loc, rewriter, offsetMapForLoopArgs); + offsetMap[arg] = offsetMapForLoopArgs.at(yield); + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Pre-parsing result of\n" << arg << "\nis "; + for (auto structured : offsetMap[arg].getStructuredRef()) + os << structured; + os << '\n'; + }); + } + } +} + +static bool isFromTensorArg(Value v, + llvm::SmallDenseMap &fromTensorArg) { + if (fromTensorArg.contains(v)) + return fromTensorArg.at(v); + auto *defOp = v.getDefiningOp(); + if (!defOp) { + fromTensorArg[v] = isa(v.getType()); + return isa(v.getType()); + } + for (auto opr : defOp->getOperands()) { + if (isFromTensorArg(opr, fromTensorArg)) { + fromTensorArg[v] = true; + return true; + } + } + fromTensorArg[v] = false; + return false; +} + +template +void TritonToUnstructureIncubatedPass::runParse(MemAccOpTy op) { + IRRewriter rewriter(&getContext()); + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Parsing " << op->getName() << "\n" << op << "\n"; + }); + parse(op.getPtr(), op.getLoc(), rewriter, offsetMap); + isFromTensorArg(op.getPtr(), fromTensorArg); +} + +TritonToUnstructureIncubatedPass::TritonToUnstructureIncubatedPass( + const TritonToUnstructureIncubatedOptions &options) + : TritonToUnstructureIncubatedBase(options) {} + +void TritonToUnstructureIncubatedPass::runOnOperation() { + compileOn91095Flag = this->compileOn91095; + forceSimtTemplateFlag = this->forceSimtTemplate; + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "TritonToUnstructureIncubatedPass started with options:\n"; + os << " compileOn91095: " << compileOn91095Flag << "\n"; + os << " forceSimtTemplate: " << forceSimtTemplateFlag << "\n"; + }); + + ModuleOp moduleOp = getOperation(); + MLIRContext *ctx = &getContext(); + + replacePtrLoopArguments(moduleOp, offsetMapForLoopArgs); + offsetMapForLoopArgs.clear(); + moduleOp->walk([this](LoopLikeOpInterface op) { runPreparse(op); }); + moduleOp->walk([this](Operation *op) { + if (auto loadOp = dyn_cast(op)) { + runParse(loadOp); + } else if (auto storeOp = dyn_cast(op)) { + runParse(storeOp); + } else if (auto atomicRMWOp = dyn_cast(op)) { + runParse(atomicRMWOp); + } else if (auto atomicCASOp = dyn_cast(op)) { + runParse(atomicCASOp); + } + }); + + RewritePatternSet patterns(ctx); + + patterns.add, + UnstructuredMemAccessConverter, + UnstructuredMemAccessConverter, + UnstructuredMemAccessConverter>( + ctx, forceScalarizeMode, offsetMap, fromTensorArg); + + LLVM_DEBUG({ + auto &os = llvm::dbgs(); + os << "Parsing done\n"; + }); + + if (failed(applyPatternsAndFoldGreedily(moduleOp, std::move(patterns)))) { + moduleOp->emitError("failed to apply Patterns"); + signalPassFailure(); + } + + PassManager pm(&getContext(), moduleOp.getOperationName()); + pm.addPass(createCSEPass()); + pm.addPass(createCanonicalizerPass()); + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } +} + +void TritonToUnstructureIncubatedPass::getDependentDialects( + DialectRegistry ®istry) const { + registry.insert(); +} + +std::unique_ptr> +triton::createTritonToUnstructureIncubatedPass( + const TritonToUnstructureIncubatedOptions &options) { + return std::make_unique(options); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructured/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructured/CMakeLists.txt new file mode 100755 index 00000000..d665048b --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructured/CMakeLists.txt @@ -0,0 +1,23 @@ +add_triton_library(TritonToUnstructured + TritonToUnstructuredPass.cpp + + DEPENDS + TritonStructuredTableGen + TritonToUnstructuredConversionPassIncGen + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + MLIRReconcileUnrealizedCasts + TritonIR + TritonTransforms + TritonSharedAnalysisStructured + TritonStructuredIR + TritonSharedUtils +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructured/TritonToUnstructuredPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructured/TritonToUnstructuredPass.cpp new file mode 100755 index 00000000..d0eccc73 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructured/TritonToUnstructuredPass.cpp @@ -0,0 +1,777 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// +// +//////////////////////////////////////////////////////////////////////////////// +// Overview +//////////////////////////////////////////////////////////////////////////////// +// +// This pass attempts to lower all loads and stores of unstructured pointers to +// tts.gather or tts.scatter that take a single base, a tensor of offsets, an +// optional tensor of mask values, and a default value in case of load. +// +// In addition, all pointer-producing ops will be eliminated and replaced by +// offset-producing ops. tts.gather and tts.scatter will use the pointer +// directly from the kernel arguments as opposed to pointer produced by ops such +// as tt.addptr and tt.splat. +// +// Example: +// +// %12 = tts.gather %arg0[%10] : (, tensor<64xi64>) -> tensor<64xf32> +// tts.scatter %12 into %arg1[%arg3] : tensor<64xf32> into (, +// tensor<64xi32>) +// +// Current assumptions and limitations: +// - For simplicity, the pass assumes that gather / scatter operations load / +// store from / to a single base with a tensor of random offsets. As a +// result, the following triton program would not work: +// +// @triton.jit +// def gather_simple(in0, in1, out0): +// offs = tl.arange(0, 8) +// in0_ptrs = in0 + offs +// in1_ptrs = in1 + offs +// ptrs = tl.cat(in0_ptrs, in1_ptrs, can_reorder=True) +// c = tl.load(ptrs) +// out_offs = tl.arange(0, 16) +// tl.store(out0 + out_offs, c) +// +// In the above program, `ptrs` contains 2 bases: `in0` and `in1` after the +// `cat` operation. +// +//////////////////////////////////////////////////////////////////////////////// +// Future work +//////////////////////////////////////////////////////////////////////////////// +// +// Future work may include scaling the algorithm to support such cases -- one +// possible solution is to let tts.gather and tts.scatter take in an additional +// tensor of base pointers corresponding to the tensor of offsets. But because +// we do not want pointer-producing ops to be present after this pass, we can +// use a tensor of index where each element indicates the index of the pointer +// argument to be used. The drawback is a gather or scatter operation now needs +// one extract lookup to get the base which will affect performance. +// +//////////////////////////////////////////////////////////////////////////////// +// Algorithm +//////////////////////////////////////////////////////////////////////////////// +// +// Because the goal of triton-shared is to eventually lower all triton ops and +// types to mlir, we want to transform the IR such that the usages of triton +// pointers are as limited as possible. Doing so will help simplify conversion +// to mlir dialects in subsequent passes. In a familiar fashion to the +// triton-to-structured pass, we want triton pointers to only appear in +// tts.gather and tts.scatter only. +// +// With that goal in mind, we want to revisit the triton pointer type. +// +// Triton pointers are created and manipulated through a sequence of ops such as +// tt.addptr, tt.splat, or tt.broadcast. If a triton pointer is created +// through `tt.addptr %ptr %offset`, the new pointer will contain the same base +// pointer as the original pointer; its offset will also be accumulated. +// +// Triton pointers created through tt.splat and tt.broadcast retain their base +// pointers and offsets. Tensors of pointers, however, may have different bases +// when tl.cat is present. For simplicity, we assume tl.cat isn't present as +// mentioned in the overview section. +// +// Therefore, a single triton pointer (tt.ptr) has two pieces of info that is +// implicit: +// - a base pointer which comes from the kernel arguments +// - an offset which could be either a tensor of offset or a single integer +// offset +// +// Leveraging this insight, in order to limit the usages of triton pointer, we +// can explicitly compute and split the above two pieces of info. So chains of +// tt.addptr, tt.splat, and tt.broadcast which produce triton pointers can be +// transformed to sequences of offset (of integer type) manipulation ops and a +// base pointer which comes from the kernel arguments. With this approach, only +// tts.gather and tts.scatter need to be aware of the pointer type. +// +// In essence, this pass transforms all sequences of tt.addptr into sequences of +// offset accumulation ops which are then fed into a single op +// tts.gather or tts.scatter that takes: +// +// - a base pointer from the kernel arguments +// - a tensor of offsets (or single offset) that indicates the offsets from +// the base pointer +// +// All intermediate tt.addptr ops are converted to arith.addi ops that compute +// the offsets. Offsets start at 0 with the provided bit-width. All pointer +// shape manipulation ops such as tt.splat and tt.broadcast will instead operate +// on the offsets and will be converted to linalg in triton-arith-to-linalg. +// +// By default, the pass uses i32 for the initial offsets of all pointers +// (configurable via offset-bit-width=width). If any intermediate tt.addptr +// introduces a larger bitwidth offset, the offsets will be sign-extended to the +// larger bitwidth. +// +//////////////////////////////////////////////////////////////////////////////// +// Algorithm +//////////////////////////////////////////////////////////////////////////////// +// +// This pass uses a standard worklist-based algorithm to walk the use-def chains +// of all pointer arguments and create replacement ops that operate on offsets +// instead of tt.ptr types. +// +// In cases such as tt.addptr, tt.splat, and tt.broadcast, we create +// corresponding replacement ops which will then be used to map the results +// at the end of the algorithm. We do not want to modify these ops in-place +// because the use-def chains may be changed. In special cases like scf.for, we +// also set the type of the iter-arg and result directly which is usually frown +// upon (but justified). +// +// This approach is used in favor of the traditional ConversionPatternRewriter +// which converts all pointer type into an offset integer type because +// TypeConverter does not support dynamic type based on value. This limitation +// means we have to decide the same bitwidth for all tt.addptr sequences which +// is not ideal. +// +// For instance, assuming we have two sequences of tt.addptr: one operates on +// 32-bit offsets while the other operates on 64-bit offsets. If we set the +// default bitwidth to 64, the 32-bit sequence will require unncessary +// sign-extending when computing the offsets. Contrast this with the manual +// approach, we will only sign-extend where necessary. + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/Operation.h" +#include "mlir/IR/TypeRange.h" +#include "mlir/IR/Types.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/Support/LLVM.h" +#include "mlir/Transforms/Passes.h" +#include "triton-shared/Analysis/OpFoldResultUtils.h" +#include "triton-shared/AnalysisStructured/PtrAnalysis.h" +#include "triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Utils/Utils.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Pass/PassManager.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/DenseSet.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/LogicalResult.h" + +#include +#include + +#define DEBUG_TYPE "triton-to-unstructured" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/TritonToUnstructured/Passes.h.inc" + +namespace { + +// Given a type, return the offset type corresponding to that type with the +// specified width. +// If the type is a tensor, return a tensor of offsets of the same shape. If the +// type is a pointer, return a single offset type. +static Type getPtrOffsetType(Type type, unsigned int bitWidth) { + if (auto tensorType = dyn_cast(type)) { + if (auto ptrType = + dyn_cast(tensorType.getElementType())) { + return RankedTensorType::get( + tensorType.getShape(), IntegerType::get(type.getContext(), bitWidth)); + } + } + + if (auto ptrType = dyn_cast(type)) { + return IntegerType::get(type.getContext(), bitWidth); + } + + llvm_unreachable("unexpected type"); + return nullptr; +} + +static unsigned int getBitWidth(Type type) { + if (auto tensorType = dyn_cast(type)) { + if (auto integerType = dyn_cast(tensorType.getElementType())) { + return integerType.getWidth(); + } + } else if (auto integerType = dyn_cast(type)) { + return integerType.getWidth(); + } + + llvm_unreachable("unexpected type"); + return 0; +} + +class TritonToUnstructuredPass + : public TritonToUnstructuredBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + struct PtrOffset { + // the source pointer which comes from the kernel argument + Value ptr; + // the pointer type that corresponds to this offset; used when + // creating tts.make_unstructured_tptr + Type ptrType; + // bitwidth that is used for this offset, used to track if sign-extension is + // necessary + unsigned int bitWidth; + // the offset value + Value offset; + }; + + LogicalResult processUnstructuredPtrs(unsigned int defaultBitWidth = 32) { + llvm::SmallDenseSet ptrArgs; + llvm::DenseMap offsetMap; + std::queue workList; + + getOperation().walk([&](FunctionOpInterface func) { + for (auto arg : func.getArguments()) { + if (!triton::isPtrTypeLike(arg.getType())) { + continue; + } + + OpBuilder b(func->getRegion(0)); + Value zero = b.create( + arg.getLoc(), + b.getIntegerAttr(IntegerType::get(&getContext(), defaultBitWidth), + 0)); + + ptrArgs.insert(arg); + offsetMap.insert({arg, {arg, arg.getType(), defaultBitWidth, zero}}); + workList.push(arg); + } + }); + + getOperation().walk([&](triton::IntToPtrOp op) { + // We only want to handle single source pointer, + // skip if this op produces tensor of pointers + if (isa(op.getType())) { + return; + } + auto res = op.getResult(); + OpBuilder b(op); + Value zero = b.create( + op.getLoc(), + b.getIntegerAttr(IntegerType::get(&getContext(), defaultBitWidth), + 0)); + + offsetMap.insert({res, {res, res.getType(), defaultBitWidth, zero}}); + workList.push(res); + }); + + llvm::SmallVector toDelete; + llvm::SmallVector ptrUsers; + + while (!workList.empty()) { + auto val = workList.front(); + workList.pop(); + + for (auto &use : val.getUses()) { + auto user = use.getOwner(); + + auto res = + llvm::TypeSwitch(user) + .Case([&](arith::SelectOp op) { + auto ptr = op->getOperand(0); + + if (!offsetMap.contains(op.getTrueValue()) || + !offsetMap.contains(op.getFalseValue())) + return success(); + auto TrueValue = offsetMap.at(op.getTrueValue()); + auto FalseValue = offsetMap.at(op.getFalseValue()); + assert(TrueValue.bitWidth == FalseValue.bitWidth && + "arith.select op should have the same bitwidth " + "for both true and false values"); + auto res = op.getResult(); + auto resType = op.getType(); + + OpBuilder b{op}; + auto newOffset = b.create( + op->getLoc(), + getPtrOffsetType(resType, TrueValue.bitWidth), + op.getCondition(), TrueValue.offset, FalseValue.offset); + PtrOffset newOffsetInfo{res, resType, TrueValue.bitWidth, + newOffset}; + + offsetMap.insert({ + res, + newOffsetInfo, + }); + workList.push(res); + return success(); + }) + .Case([&](triton::PtrToIntOp op) { + auto offsetInfo = offsetMap.at(op.getSrc()); + + OpBuilder b{op}; + // We are converting a pointer to an integer here, + // materialized the pointer using the accumulated offset + // that we have stored so far. + auto materializedAddPtr = b.create( + op->getLoc(), offsetInfo.ptrType, offsetInfo.ptr, + offsetInfo.offset); + + // Change the op to use the "simplified" pointer above. + // This should not affect the traversal of uses, but hacky. + // We will need to revisit how we process the IRs in this pass + // later. + op->setOperand(0, materializedAddPtr); + + return success(); + }) + .Case([&](triton::BitcastOp bitcast) { + auto resPtrType = bitcast.getType(); + OpBuilder b{bitcast}; + auto loc = bitcast->getLoc(); + + auto offsetInfo = offsetMap.at(bitcast.getOperand()); + auto srcPtrType = offsetInfo.ptrType; + assert((triton::isPtrTypeLike(resPtrType) && + triton::isPtrTypeLike(srcPtrType)) && + "unexpected bitcast type"); + Type resType = + isa(resPtrType) + ? cast(resPtrType).getElementType() + : resPtrType; + Type srcType = + isa(srcPtrType) + ? cast(srcPtrType).getElementType() + : srcPtrType; + assert(((cast(srcType) + .getPointeeType() + .isInteger(1) && + cast(resType) + .getPointeeType() + .isInteger(8)) || + (cast(srcType) + .getPointeeType() + .isInteger(8) && + cast(resType) + .getPointeeType() + .isInteger(1))) && + "only bitcast between i1 and i8 pointer is supported"); + auto newBitcast = + b.create(loc, resType, offsetInfo.ptr); + bitcast->replaceAllUsesWith(newBitcast); + PtrOffset newOffsetInfo{newBitcast, resPtrType, + offsetInfo.bitWidth, + offsetInfo.offset}; + offsetMap.insert({newBitcast, newOffsetInfo}); + workList.push(newBitcast); + + return success(); + }) + .Case([&](triton::AddPtrOp addptr) { + OpBuilder b{addptr}; + auto loc = addptr->getLoc(); + + auto offsetInfo = offsetMap.at(addptr.getPtr()); + + auto prevOff = offsetInfo.offset; + auto off = addptr.getOffset(); + + auto lhsWidth = offsetInfo.bitWidth; + auto rhsWidth = getBitWidth(off.getType()); + auto resWidth = std::max(lhsWidth, rhsWidth); + + if (lhsWidth < resWidth) { + prevOff = b.create( + loc, getPtrOffsetType(offsetInfo.ptrType, resWidth), + prevOff); + } + + if (rhsWidth < resWidth) { + off = b.create( + loc, getPtrOffsetType(offsetInfo.ptrType, resWidth), + off); + } + + auto accumulatedOff = b.create( + loc, getPtrOffsetType(addptr.getType(), resWidth), + prevOff, off); + + PtrOffset newOffsetInfo{offsetInfo.ptr, addptr.getType(), + resWidth, accumulatedOff}; + + offsetMap.insert({addptr, newOffsetInfo}); + workList.push(addptr); + toDelete.push_back(addptr); + + return success(); + }) + .Case([&](Operation *op) { + auto res = op->getResult(0); + auto resType = res.getType(); + + if (!triton::isPtrTypeLike(resType)) { + return success(); + } + + auto ptr = op->getOperand(0); + auto offsetInfo = offsetMap.at(ptr); + + OpBuilder b{op}; + auto clone = b.create( + op->getLoc(), op->getName().getIdentifier(), + ValueRange{offsetInfo.offset}, + TypeRange{getPtrOffsetType(resType, offsetInfo.bitWidth)}, + op->getAttrs()); + + PtrOffset newOffsetInfo{offsetInfo.ptr, resType, + offsetInfo.bitWidth, + clone->getResult(0)}; + + offsetMap.insert({ + res, + newOffsetInfo, + }); + workList.push(res); + toDelete.push_back(op); + + return success(); + }) + .Case([&](Operation *op) { + // Special case: + // We do not want to create "unstructured tensor pointer" into + // tts.make_tptr if the base pointer is directly from the + // kernel arguments. + if (auto makeTensorPtr = dyn_cast(op)) { + if (ptrArgs.contains(makeTensorPtr.getBase())) { + return success(); + } + if (auto bitcast = dyn_cast( + makeTensorPtr.getBase().getDefiningOp())) { + if (ptrArgs.contains(bitcast.getSrc())) { + return success(); + } + } + } + + ptrUsers.push_back(op); + return success(); + }) + .Case([&](scf::ForOp forOp) { + // Index of the init-arg corresponding to this use, note that + // we have to subtract by 3 from the operand number because + // scf.for ops always have 3 leading operands for start, end, + // and step. + auto argIndex = use.getOperandNumber() - 3; + auto init = forOp.getInitArgs()[argIndex]; + + auto offsetInfo = offsetMap.at(init); + + auto offsetType = + getPtrOffsetType(offsetInfo.ptrType, offsetInfo.bitWidth); + + + // In order to keep types of the operands consistent, we need + // to replace the base pointer if it is directly from the + // kernel arguments. + if (ptrArgs.contains(init)) { + forOp->setOperand(argIndex, offsetInfo.offset); + } + + // We're setting both the types of the iter-arg and the + // corresponding result directly to the offset type. + // At this point, the IR is in an invalid state because the + // init-args still have tt.ptr. But at the end, we will + // replace all uses of the tt.ptr to offset values. + auto iterArg = forOp.getRegionIterArg(argIndex); + iterArg.setType(offsetType); + + auto res = forOp.getResult(argIndex); + res.setType(offsetType); + + // For other ops, we only need to push the result into the + // worklist. But for scf.for, the iter-arg corresponding to + // the init-arg is used in the op's body instead, we have to + // process uses of the iter-arg. + PtrOffset iterArgOffset{offsetInfo.ptr, offsetInfo.ptrType, + offsetInfo.bitWidth, iterArg}; + offsetMap.insert({ + iterArg, + iterArgOffset, + }); + + PtrOffset resOffset{offsetInfo.ptr, offsetInfo.ptrType, + offsetInfo.bitWidth, res}; + offsetMap.insert({ + res, + resOffset, + }); + workList.push(iterArg); + workList.push(res); + + return success(); + }) + .Case([&](scf::WhileOp whileOp) { + auto argIndex = use.getOperandNumber(); + auto init = whileOp.getInits()[argIndex]; + auto offsetInfo = offsetMap.at(init); + auto offsetType = + getPtrOffsetType(offsetInfo.ptrType, offsetInfo.bitWidth); + + // In order to keep types of the operands consistent, we need + // to replace the base pointer if it is directly from the + // kernel arguments. + if (ptrArgs.contains(init)) { + whileOp->setOperand(argIndex, offsetInfo.offset); + } + auto beforeArg = whileOp.getBeforeArguments()[argIndex]; + beforeArg.setType(offsetType); + + auto afterArg = whileOp.getAfterArguments()[argIndex]; + afterArg.setType(offsetType); + + auto res = whileOp->getOpResult(argIndex); + res.setType(offsetType); + + PtrOffset beforeArgOffset{offsetInfo.ptr, offsetInfo.ptrType, + offsetInfo.bitWidth, beforeArg}; + offsetMap.insert({ + beforeArg, + beforeArgOffset, + }); + + PtrOffset afterArgOffset{offsetInfo.ptr, offsetInfo.ptrType, + offsetInfo.bitWidth, afterArg}; + + offsetMap.insert({ + afterArg, + afterArgOffset, + }); + + PtrOffset resOffset{offsetInfo.ptr, offsetInfo.ptrType, + offsetInfo.bitWidth, res}; + offsetMap.insert({ + res, + resOffset, + }); + workList.push(beforeArg); + workList.push(afterArg); + workList.push(res); + + return success(); + }) + .Case( + [&](Operation *op) { + ptrUsers.push_back(op); + return success(); + }) + .Case([&](triton::BitcastOp bitcast) { + for (auto *user : bitcast->getUsers()) { + if (isa(user)) { + // If this bitcast is used by a tts.make_tptr op, ignore + // it. + return success(); + } + } + + auto offsetInfo = offsetMap.at(bitcast.getSrc()); + + Value castedPtr; + if (ptrArgs.contains(bitcast.getSrc())) { + castedPtr = bitcast.getResult(); + } else { + // Move the bitcast to source pointer. + OpBuilder b{bitcast}; + castedPtr = b.create( + bitcast.getLoc(), bitcast.getType(), offsetInfo.ptr); + + offsetMap.erase(bitcast.getResult()); + bitcast.replaceAllUsesWith(castedPtr); + bitcast.erase(); + } + + PtrOffset updatedOffsetInfo{castedPtr, castedPtr.getType(), + offsetInfo.bitWidth, + offsetInfo.offset}; + + offsetMap.insert({castedPtr, updatedOffsetInfo}); + workList.push(castedPtr); + + return success(); + }) + .Case( + [](auto) { return success(); }) + .Case([](triton::CatOp op) { + op->emitError("Do not support gather / scatter with multiple " + "bases yet"); + return failure(); + }) + .Default([&](Operation *op) { + op->emitError("unexpected op in ptr sequence"); + return failure(); + }); + + if (failed(res)) { + return failure(); + } + } + } + + for (auto op : ptrUsers) { + OpBuilder b{op}; + auto loc = op->getLoc(); + auto res = + llvm::TypeSwitch(op) + .Case([&](triton::LoadOp load) { + auto offsetInfo = offsetMap.at(load.getPtr()); + + auto other = load.getOther(); + + if (other) { + other = tts::utils::getScalarValue(other, loc, b); + if (!other) { + load->emitError("cannot parse `other` value for load"); + return failure(); + } + } + + auto gather = b.create( + loc, load.getType(), offsetInfo.ptr, offsetInfo.offset, + load.getMask(), other); + + load->replaceAllUsesWith(gather->getResults()); + load->erase(); + return success(); + }) + .Case([&](triton::StoreOp store) { + auto offsetInfo = offsetMap.at(store.getPtr()); + b.create(loc, offsetInfo.ptr, offsetInfo.offset, + store.getValue(), store.getMask()); + store->erase(); + return success(); + }) + .Case([&](auto makeTensorPtr) { + // For block pointers, the base could come from a sequence of + // `tt.addptr`. Accumulate the target offset with the offset + // we have saved. + auto offsetInfo = offsetMap.at(makeTensorPtr.getBase()); + auto baseOffset = offsetInfo.offset; + + makeTensorPtr.getBaseMutable().set(offsetInfo.ptr); + + // Add the existing offset from the base to the offset + // operand in the ops. + auto &offsetOpnd = makeTensorPtr.getOffsetsMutable()[0]; + auto currOffset = offsetOpnd.get(); + + auto baseOffType = baseOffset.getType(); + auto currOffType = currOffset.getType(); + + if (baseOffType != currOffType) { + if (currOffType.isIndex()) { + baseOffset = b.create( + loc, b.getIndexType(), baseOffset); + } else if (currOffType.isInteger()) { + if (baseOffType.getIntOrFloatBitWidth() < + currOffType.getIntOrFloatBitWidth()) { + baseOffset = b.create(loc, currOffType, + baseOffset); + } else { + // MakeTensorPtrOp only takes i32 offsets, so we need + // to truncate if the offsets were already in i64 + makeTensorPtr.emitWarning( + "truncating offsets which may result in data loss"); + baseOffset = b.create(loc, currOffType, + baseOffset); + } + } + } + + auto accumulatedOffset = b.create( + loc, currOffset.getType(), baseOffset, currOffset); + + offsetOpnd.set(accumulatedOffset); + + return success(); + }) + .Case([&](triton::AtomicCASOp atomicOp) { + auto offsetInfo = offsetMap.at(atomicOp.getPtr()); + auto casOp = b.create( + loc, atomicOp.getType(), offsetInfo.ptr, atomicOp.getCmp(), + atomicOp.getVal(), offsetInfo.offset, atomicOp.getSemAttr(), + atomicOp.getScopeAttr()); + atomicOp->replaceAllUsesWith(casOp->getResults()); + atomicOp->erase(); + + return success(); + }) + .Case([&](triton::AtomicRMWOp atomicOp) { + auto offsetInfo = offsetMap.at(atomicOp.getPtr()); + auto rmwOp = b.create( + loc, atomicOp.getType(), offsetInfo.ptr, atomicOp.getVal(), + atomicOp.getMask(), offsetInfo.offset, + atomicOp.getAtomicRmwOpAttr(), atomicOp.getSemAttr(), + atomicOp.getScopeAttr()); + atomicOp->replaceAllUsesWith(rmwOp->getResults()); + atomicOp->erase(); + return success(); + }) + + .Default([&](Operation *op) { + op->emitError("unexpected op in ptr sequence"); + return failure(); + }); + + if (failed(res)) { + return failure(); + } + } + + for (auto op : toDelete) { + auto ptrInfo = offsetMap.at(op->getResult(0)); + op->replaceAllUsesWith(ValueRange{ptrInfo.offset}); + op->erase(); + } + + return success(); + } + + void runOnOperation() override { + if (failed(processUnstructuredPtrs(offsetBitWidth))) { + getOperation()->emitWarning( + "Cannot transform tensor of pointers into a single base pointer " + "with tensor of offsets"); + return; + } + + PassManager pm(&getContext(), getOperation().getOperationName()); + pm.addPass(createCanonicalizerPass()); + pm.addPass(createCSEPass()); + if (failed(runPipeline(pm, getOperation()))) { + signalPassFailure(); + } + } +}; +} // namespace + +std::unique_ptr> +triton::createTritonToUnstructuredPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Conversion/UnstructuredToMemref/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Conversion/UnstructuredToMemref/CMakeLists.txt new file mode 100755 index 00000000..a2ba7c65 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/UnstructuredToMemref/CMakeLists.txt @@ -0,0 +1,27 @@ +#===------------------------------------------------------------------------===# +# +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. +# +#===------------------------------------------------------------------------===# + +add_triton_library(UnstructuredToMemref + UnstructuredToMemrefPass.cpp + + DEPENDS + UnstructuredToMemrefConversionPassIncGen + + LINK_LIBS PUBLIC + TritonTilingExtIR + MLIRArithDialect + MLIRDialectUtils + MLIRIR + MLIRMathDialect + MLIRPass + MLIRTensorDialect + MLIRTransforms + MLIRSupport + TritonIR + TritonTransforms + TritonSharedAnalysis +) diff --git a/third_party/wafer/third_party/flir/lib/Conversion/UnstructuredToMemref/UnstructuredToMemrefPass.cpp b/third_party/wafer/third_party/flir/lib/Conversion/UnstructuredToMemref/UnstructuredToMemrefPass.cpp new file mode 100755 index 00000000..c377c069 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Conversion/UnstructuredToMemref/UnstructuredToMemrefPass.cpp @@ -0,0 +1,439 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/ValueRange.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/Support/ErrorHandling.h" +#include + +#define DEBUG_TYPE "unstructured-to-memref" + +using namespace mlir; +using namespace triton; + +#define GEN_PASS_CLASSES +#include "triton-shared/Conversion/UnstructuredToMemref/Passes.h.inc" + +namespace { + +class PtrToUnrankedMemrefConverter : public TypeConverter { +public: + PtrToUnrankedMemrefConverter() { + addConversion([](Type type) { return type; }); + addConversion([](triton::PointerType ptrType) { + return UnrankedMemRefType::get(ptrType.getPointeeType(), 0); + }); + addTargetMaterialization([&](OpBuilder &builder, + UnrankedMemRefType resultType, + ValueRange inputs, Location loc) -> Value { + return builder.create(loc, resultType, inputs) + .getResult(0); + }); + } +}; + +static MemRefType getMemrefTypeForScalarPtr(triton::PointerType ptrType, + MLIRContext *context) { + SmallVector strides{1}; + auto layout = StridedLayoutAttr::get(context, ShapedType::kDynamic, strides); + auto elemType = ptrType.getPointeeType(); + auto memrefType = MemRefType::get({1}, elemType, layout); + return memrefType; +} + +struct ScalarLoadConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + ScalarLoadConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + ScalarLoadConverter(MLIRContext *context) + : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(tts::GatherOp gatherOp, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!gatherOp.getType().isIntOrIndexOrFloat()) { + return failure(); + } + + auto loc = gatherOp->getLoc(); + + auto basePtr = adaptor.getPtr(); + auto offset = adaptor.getOffset(); + + Value loadIndex = rewriter.create( + loc, rewriter.getIndexType(), offset); + + auto memref = rewriter.create( + loc, + getMemrefTypeForScalarPtr( + cast(gatherOp.getPtr().getType()), + rewriter.getContext()), + basePtr, getAsOpFoldResult(loadIndex) /*offset*/, + ArrayRef{rewriter.getIndexAttr(1)} /*sizes*/, + ArrayRef{rewriter.getIndexAttr(1)} /*strides*/); + + auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext()); + + auto scalarLoadOp = rewriter.create( + loc, memref, zeroMap, ValueRange{}); + + rewriter.replaceOp(gatherOp, scalarLoadOp.getResult()); + + return success(); + } +}; + +struct ScalarStoreConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + ScalarStoreConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + ScalarStoreConverter(MLIRContext *context) + : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(tts::ScatterOp scatterOp, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + + if (!scatterOp.getValue().getType().isIntOrIndexOrFloat()) { + return failure(); + } + + auto loc = scatterOp->getLoc(); + + auto basePtr = adaptor.getPtr(); + auto offset = adaptor.getOffset(); + + Value storeIndex = rewriter.create( + loc, rewriter.getIndexType(), offset); + + auto memref = rewriter.create( + loc, + getMemrefTypeForScalarPtr( + cast(scatterOp.getPtr().getType()), + rewriter.getContext()), + basePtr, getAsOpFoldResult(storeIndex) /*offset*/, + ArrayRef{rewriter.getIndexAttr(1)} /*sizes*/, + ArrayRef{rewriter.getIndexAttr(1)} /*strides*/); + + auto storeVal = scatterOp.getValue(); + +#ifdef FLAGTREE_BACKEND_WAFER + rewriter.create( + loc, storeVal, memref, + ValueRange{rewriter.create(loc, 0)}); +#else + auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext()); + + rewriter.create(loc, storeVal, memref, zeroMap, + ValueRange{}); +#endif + rewriter.eraseOp(scatterOp); + + return success(); + } +}; + +// Lowering an unstructured load op (gather) into a linalg.generic op. +struct GatherConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + GatherConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + GatherConverter(MLIRContext *context) + : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(tts::GatherOp gatherOp, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = gatherOp->getLoc(); + + auto ptr = adaptor.getPtr(); + auto offsetTensor = adaptor.getOffset(); + auto offsetType = dyn_cast(offsetTensor.getType()); + + // This must be a scalar load, skip processing. + if (!offsetType) { + return failure(); + } + + auto resultType = + dyn_cast(gatherOp.getResult().getType()); + + // Treat the base pointer (memref) as 1D because the offsets are all + // relative to a single base pointer (already collapsed). + auto baseMemref = rewriter + .create( + loc, + MemRefType::get({ShapedType::kDynamic}, + resultType.getElementType()), + ptr) + .getResult(); + + auto baseTensor = + rewriter + .create( + loc, + RankedTensorType::get( + SmallVector(1, ShapedType::kDynamic), + resultType.getElementType()), + baseMemref, true /* restrict */, false /* writable */) + .getResult(); + + // The linalg.generic op should have the following inputs: + // - the offset tensor. + // - an optional mask tensor if the gather op contains mask. + SmallVector inputs{offsetTensor}; + + if (gatherOp.getMask()) { + inputs.push_back(gatherOp.getMask()); + } + + auto emptyTensor = rewriter + .create(loc, resultType.getShape(), + resultType.getElementType()) + .getResult(); + + // Affine maps for the inputs and one additional output. + SmallVector affineMaps( + inputs.size() + 1, + rewriter.getMultiDimIdentityMap(resultType.getRank())); + + // All iterator types are parallel. + SmallVector iteratorTypes( + resultType.getRank(), utils::IteratorType::parallel); + + auto genericOp = rewriter.create( + loc, TypeRange{resultType}, inputs, ValueRange{emptyTensor}, affineMaps, + iteratorTypes, [&](OpBuilder &b, Location loc, ValueRange args) { +#ifdef FLAGTREE_BACKEND_WAFER + auto getValueAtIndex = [baseMemref](OpBuilder &b, Location loc, + Value index) -> Value { + Value index0 = + b.create(loc, b.getIndexType(), index); + + return b.create(loc, baseMemref, + ValueRange{index0}); +#else + auto getValueAtIndex = [baseTensor](OpBuilder &b, Location loc, + Value index) -> Value { + Value index0 = + b.create(loc, b.getIndexType(), index); + + return b.create(loc, baseTensor, + ValueRange{index0}); +#endif + }; + + auto offset = args[0]; + + if (!gatherOp.getMask()) { + // If there is no mask, simply extract the current element from the + // base tensor and use it as the yield value. + auto loadValue = getValueAtIndex(b, loc, offset); + b.create(loc, loadValue); + } else { + // If the mask value is truthy, the current element is loaded from + // the base tensor using its offset. Otherwise, if `other` is + // present, yield `other`. If `other` is not present, a default + // value of 0 is used. + auto mask = args[1]; + auto ifOp = b.create( + loc, mask, + [&](OpBuilder &b, Location loc) { + // Truthy case, load from the index. + auto value = getValueAtIndex(b, loc, offset); + b.create(loc, value); + }, + [&](OpBuilder &b, Location loc) { + // Falsy case, yield `other` or 0 as the default value. + if (gatherOp.getOther()) { + b.create(loc, gatherOp.getOther()); + } else { + auto elemType = resultType.getElementType(); + auto zeroAttr = b.getZeroAttr(elemType); + assert(zeroAttr && "unexpected element type"); + Value extract = b.create(loc, zeroAttr); + b.create(loc, extract); + } + }); + + b.create(loc, ifOp->getResult(0)); + } + }); + + rewriter.replaceOp(gatherOp, genericOp); + + return success(); + } +}; + +// Lowering an unstructured store op (scatter) into a linalg.generic op. +struct ScatterConverter : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + ScatterConverter(const TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context) {} + + ScatterConverter(MLIRContext *context) + : OpConversionPattern(context) {} + + LogicalResult + matchAndRewrite(tts::ScatterOp scatterOp, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto loc = scatterOp->getLoc(); + + auto ptr = adaptor.getPtr(); + auto offsetTensor = adaptor.getOffset(); + auto valueTensor = adaptor.getValue(); + auto offsetType = dyn_cast(offsetTensor.getType()); + + // This must be a scalar store, skip processing. + if (!offsetType) { + return failure(); + } + + auto valueType = dyn_cast(scatterOp.getValue().getType()); + + // Treat the base pointer (memref) as 1D because the offsets are all + // relative to a single base pointer (already collapsed). + auto baseMemref = + rewriter + .create(loc, + MemRefType::get({ShapedType::kDynamic}, + valueType.getElementType()), + ptr) + .getResult(); + + // The linalg.generic op should have the following inputs: + // - the offset tensor. + // - the value tensor. + // - an optional mask tensor if the scatter op contains mask. + SmallVector inputs{offsetTensor, valueTensor}; + + if (scatterOp.getMask()) { + inputs.push_back(scatterOp.getMask()); + } + + // Affine maps for the inputs. + SmallVector affineMaps( + inputs.size(), rewriter.getMultiDimIdentityMap(valueType.getRank())); + + // All iterator types are parallel. + SmallVector iteratorTypes( + valueType.getRank(), utils::IteratorType::parallel); + + rewriter.setInsertionPoint(scatterOp); + + auto genericOp = rewriter.create( + loc, TypeRange{}, inputs, ValueRange{}, affineMaps, iteratorTypes, + [&](OpBuilder &b, Location loc, ValueRange args) { + auto storeValueAtIndex = [baseMemref](OpBuilder &b, Location loc, + Value index, Value value) { + Value index0 = + b.create(loc, b.getIndexType(), index); + + b.create(loc, value, baseMemref, + ValueRange{index0}); + }; + + auto offset = args[0]; + auto value = args[1]; + + if (!scatterOp.getMask()) { + // If there is no mask, simply insert the current value to the + // base memref using its offset. + storeValueAtIndex(b, loc, offset, value); + } else { + // If the mask value is truthy, insert the current value to the + // the base memref using its offset. Otherwise, noop. + auto mask = args[2]; + auto ifOp = + b.create(loc, mask, [&](OpBuilder &b, Location loc) { + storeValueAtIndex(b, loc, offset, value); + b.create(loc); + }); + } + + b.create(loc); + }); + + rewriter.eraseOp(scatterOp); + + return success(); + } +}; + +class UnstructuredToMemrefPass + : public UnstructuredToMemrefBase { + +public: + void getDependentDialects(DialectRegistry ®istry) const override { + registry + .insert(); + } + + void runOnOperation() override { + auto moduleOp = getOperation(); + + RewritePatternSet patterns(&getContext()); + ConversionTarget target(getContext()); + + target.addLegalDialect< + func::FuncDialect, arith::ArithDialect, math::MathDialect, + linalg::LinalgDialect, affine::AffineDialect, scf::SCFDialect, + cf::ControlFlowDialect, tensor::TensorDialect, + bufferization::BufferizationDialect, memref::MemRefDialect, + ttx::TritonTilingExtDialect>(); + + target.addIllegalOp(); + + PtrToUnrankedMemrefConverter typeConverter; + + patterns.add(typeConverter, patterns.getContext()); + + if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) + signalPassFailure(); + } +}; +} // namespace + +std::unique_ptr> +triton::createUnstructuredToMemrefPass() { + return std::make_unique(); +} diff --git a/third_party/wafer/third_party/flir/lib/Dialect/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/CMakeLists.txt new file mode 100755 index 00000000..36402e12 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/CMakeLists.txt @@ -0,0 +1,8 @@ +add_subdirectory(TritonTilingExt) +add_subdirectory(TPtr) +add_subdirectory(MathExt) +add_subdirectory(TritonStructured) +if (FLIR_BUILD_INCUBATED) + add_subdirectory(TritonAscend) + add_subdirectory(TritonStructuredIncubated) +endif() diff --git a/third_party/wafer/third_party/flir/lib/Dialect/MathExt/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/MathExt/CMakeLists.txt new file mode 100755 index 00000000..7d59dce8 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/MathExt/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/CMakeLists.txt new file mode 100755 index 00000000..121acc0d --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/CMakeLists.txt @@ -0,0 +1,11 @@ +add_mlir_dialect_library(MLIRMathExtDialect + MathExtOps.cpp + MathExtDialect.cpp + + DEPENDS + MLIRMathExtDialectIncGen + MLIRMathExtOpsIncGen + + LINK_LIBS PUBLIC + MLIRIR +) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/MathExtDialect.cpp b/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/MathExtDialect.cpp new file mode 100755 index 00000000..f45b2068 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/MathExtDialect.cpp @@ -0,0 +1,22 @@ +#include "mlir-ext/Dialect/MathExt/IR/MathExt.h" + +#include "mlir/IR/Builders.h" +#include "mlir/IR/DialectImplementation.h" +#include "mlir/IR/OpImplementation.h" +#include "mlir/Transforms/InliningUtils.h" + +using namespace mlir; +using namespace mlir::mathext; + +#include "mlir-ext/Dialect/MathExt/IR/MathExtDialect.cpp.inc" + +//===----------------------------------------------------------------------===// +// MathExt dialect. +//===----------------------------------------------------------------------===// + +void MathExtDialect::initialize() { + addOperations< +#define GET_OP_LIST +#include "mlir-ext/Dialect/MathExt/IR/MathExtOps.cpp.inc" + >(); +} diff --git a/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/MathExtOps.cpp b/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/MathExtOps.cpp new file mode 100755 index 00000000..e858d81f --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/MathExtOps.cpp @@ -0,0 +1,54 @@ +#include "mlir-ext/Dialect/MathExt/IR/MathExt.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/CommonFolders.h" +#include "mlir/Dialect/Math/IR/Math.h" +#include "mlir/Dialect/UB/IR/UBOps.h" +#include "mlir/IR/Builders.h" +#include + +using namespace mlir; +using namespace mlir::mathext; + +#define GET_OP_CLASSES +#include "mlir-ext/Dialect/MathExt/IR/MathExtOps.cpp.inc" + +//===----------------------------------------------------------------------===// +// FModOp folder +//===----------------------------------------------------------------------===// + +OpFoldResult mathext::FModOp::fold(FoldAdaptor adaptor) { + return constFoldBinaryOp(adaptor.getOperands(), + [](const APFloat &a, const APFloat &b) { + APFloat result(a); + // APFloat::mod() offers the remainder + // behavior we want, i.e. the result has + // the sign of LHS operand. + (void)result.mod(b); + return result; + }); +} + +//===----------------------------------------------------------------------===// +// DivRzOp folder +//===----------------------------------------------------------------------===// + +OpFoldResult mathext::DivRzOp::fold(FoldAdaptor adaptor) { + return constFoldBinaryOp(adaptor.getOperands(), + [](const APFloat &a, const APFloat &b) { + APFloat result(a); + result.divide(b, APFloat::rmTowardZero); + return result; + }); +} + +/// Materialize an integer or floating point constant. +Operation *mathext::MathExtDialect::materializeConstant(OpBuilder &builder, + Attribute value, + Type type, + Location loc) { + if (auto poison = dyn_cast(value)) + return builder.create(loc, type, poison); + + return arith::ConstantOp::materialize(builder, value, type, loc); +} \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TPtr/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TPtr/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TPtr/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/CMakeLists.txt new file mode 100755 index 00000000..62ee1d3f --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/CMakeLists.txt @@ -0,0 +1,12 @@ +add_triton_library(TPtrIR + TPtrOps.cpp + TPtrDialect.cpp + + DEPENDS + TPtrTableGen + + LINK_LIBS PUBLIC + TritonIR + MLIRIR + MLIRPtrDialect + ) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/TPtrDialect.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/TPtrDialect.cpp new file mode 100755 index 00000000..5d4df692 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/TPtrDialect.cpp @@ -0,0 +1,54 @@ +#include "mlir/IR/Builders.h" + +#include "triton-shared/Dialect/TPtr/IR/TPtrDialect.h" + +#include "mlir/Dialect/Ptr/IR/PtrDialect.h" +#include "mlir/Dialect/Ptr/IR/PtrTypes.h" + +#define GET_TYPEDEF_CLASSES +#include "triton-shared/Dialect/TPtr/IR/TPtrTypes.cpp.inc" + +using namespace mlir; + +namespace { +ParseResult parseIntType(OpAsmParser &parser, Type &ty) { + if (succeeded(parser.parseOptionalColon()) && parser.parseType(ty)) + return parser.emitError(parser.getNameLoc(), "expected a type"); + if (!ty) + ty = parser.getBuilder().getIndexType(); + return success(); +} +void printIntType(OpAsmPrinter &p, Operation *op, Type ty) { + if (!ty.isIndex()) + p << " : " << ty; +} +} // namespace + +//===----------------------------------------------------------------------===// +// Dialect +//===----------------------------------------------------------------------===// +void mlir::tptr::TPtrDialect::registerTypes() { + addTypes< +#define GET_TYPEDEF_LIST +#include "triton-shared/Dialect/TPtr/IR/TPtrTypes.cpp.inc" + >(); +} + +/// Dialect creation, the instance will be owned by the context. This is the +/// point of registration of custom types and operations for the dialect. +void mlir::tptr::TPtrDialect::initialize() { + registerTypes(); + addOperations< +#define GET_OP_LIST +#include "triton-shared/Dialect/TPtr/IR/TPtrOps.cpp.inc" + >(); +} + +//===----------------------------------------------------------------------===// +// TableGen'd op method definitions +//===----------------------------------------------------------------------===// + +#define GET_OP_CLASSES +#include "triton-shared/Dialect/TPtr/IR/TPtrOps.cpp.inc" + +#include "triton-shared/Dialect/TPtr/IR/TPtrDialect.cpp.inc" diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/TPtrOps.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/TPtrOps.cpp new file mode 100755 index 00000000..7d48ec15 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/TPtrOps.cpp @@ -0,0 +1,44 @@ +#include "mlir/Interfaces/SideEffectInterfaces.h" // Required for IR/TPtrOps.h.inc +#include "mlir/Bytecode/BytecodeOpInterface.h" + +#include "mlir/IR/OpImplementation.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OperationSupport.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Dialect.h" + +#include "mlir/Dialect/Ptr/IR/PtrDialect.h" +#include "mlir/Dialect/Ptr/IR/PtrTypes.h" + +#define GET_OP_CLASSES +#include "triton-shared/Dialect/TPtr/IR/TPtrOps.h.inc" + +using namespace mlir; +using namespace mlir::tptr; + +void LoadOp::getEffects( + SmallVectorImpl> + &effects) { + effects.emplace_back(MemoryEffects::Read::get(), &getAddrMutable(), + SideEffects::DefaultResource::get()); +} + +void StoreOp::getEffects( + SmallVectorImpl> + &effects) { + effects.emplace_back(MemoryEffects::Write::get(), &getAddrMutable(), + SideEffects::DefaultResource::get()); +} + + void TypeOffsetOp::build(OpBuilder &odsBuilder, OperationState &odsState, + TypeAttr baseType, Type resultTy) { + build(odsBuilder, odsState, + resultTy ? resultTy : odsBuilder.getIndexType(), baseType); + } + +OpFoldResult TypeOffsetOp::fold(FoldAdaptor adaptor) { + return adaptor.getBaseTypeAttr(); +} diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/CMakeLists.txt new file mode 100755 index 00000000..07ad62dd --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/CMakeLists.txt @@ -0,0 +1,14 @@ +add_triton_library(TritonAscendIR + TritonAscendAttrs.cpp + TritonAscendDialect.cpp + TritonAscendOps.cpp + + DEPENDS + TritonAscendTableGen + TritonAscendAttrDefsIncGen + + LINK_LIBS PUBLIC + MLIRIR + MLIRLLVMDialect + TritonIR +) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendAttrs.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendAttrs.cpp new file mode 100755 index 00000000..91b0fa62 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendAttrs.cpp @@ -0,0 +1,15 @@ +//===- TritonAscendAttrs.cpp - TritonAscend Attributes Definition +//--------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h" + +#include "llvm/ADT/SmallSet.h" +#include "llvm/ADT/TypeSwitch.h" + +namespace mlir::triton::ascend {} // namespace mlir::triton::ascend diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendDialect.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendDialect.cpp new file mode 100755 index 00000000..0e805408 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendDialect.cpp @@ -0,0 +1,44 @@ +//===- TritonAscendDialect.cpp - TritonAscend Dialect registration +//--------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h" + +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/SPIRV/IR/TargetAndABI.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/DialectImplementation.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/Operation.h" + +#include "llvm/ADT/StringExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/AsmParser/Parser.h" +#include "llvm/IR/Function.h" +#include "llvm/Support/SourceMgr.h" + +using namespace mlir; +using namespace mlir::triton::ascend; + +void TritonAscendDialect::initialize() { + addOperations< +#define GET_OP_LIST +#include "npu/Dialect/TritonAscend/IR/TritonAscendOps.cpp.inc" + >(); + addAttributes< +#define GET_ATTRDEF_LIST +#include "npu/Dialect/TritonAscend/IR/TritonAscendOpsAttrDefs.cpp.inc" + >(); +} + +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.cpp.inc" +#define GET_ATTRDEF_CLASSES +#include "npu/Dialect/TritonAscend/IR/TritonAscendOpsAttrDefs.cpp.inc" +#define GET_OP_CLASSES +#include "npu/Dialect/TritonAscend/IR/TritonAscendOps.cpp.inc" diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendOps.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendOps.cpp new file mode 100755 index 00000000..a2d61491 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendOps.cpp @@ -0,0 +1,137 @@ +//===- TritonAscendOps.cpp - TritonAscend dialect operations +//--------------------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/SPIRV/IR/TargetAndABI.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "npu/Dialect/TritonAscend/IR/TritonAscendDialect.h" +#include "triton/Tools/Sys/GetEnv.hpp" +#include "llvm/ADT/STLExtras.h" +#include + +using namespace mlir; +using namespace mlir::triton; + +namespace mlir::triton::ascend { + +void EmbeddingGatherOp::getEffects( + SmallVectorImpl> + &effects) { + effects.emplace_back(MemoryEffects::Read::get(), &getSrcMutable(), + triton::GlobalMemory::get()); +} + +void GatherOutToUbOp::getEffects( + SmallVectorImpl> + &effects) { + effects.emplace_back(MemoryEffects::Read::get(), &getSrcMutable(), + triton::GlobalMemory::get()); +} + +void IndirectLoadOp::getEffects( + SmallVectorImpl> + &effects) { + effects.emplace_back(MemoryEffects::Read::get(), &getSrcMutable(), + triton::GlobalMemory::get()); +} + +//-- IndexSelectSimdOp -- +LogicalResult IndexSelectSimdOp::inferReturnTypes( + MLIRContext *context, std::optional location, ValueRange operands, + DictionaryAttr attributes, OpaqueProperties properties, RegionRange regions, + SmallVectorImpl &inferredReturnTypes) { + + // Get operands using adaptor + IndexSelectSimdOpAdaptor adaptor(operands, attributes, properties, regions); + + // Get element type from src pointer + Type elemType; + if (auto ptrType = + dyn_cast(adaptor.getSrc().getType())) { + elemType = ptrType.getPointeeType(); + } else { + return failure(); + } + + // Get index shape to determine the size of dim + auto indicesType = dyn_cast(adaptor.getIndex().getType()); + if (!indicesType) + return failure(); + int64_t numIndices = indicesType.getShape()[0]; + + // Use adaptor to get attributes - this is the compatible way + int32_t dim = adaptor.getDim(); + auto readShapeAttr = adaptor.getReadShape(); + + // Build result shape: read_shape but with dim replaced by numIndices + SmallVector resultShape; + for (size_t i = 0; i < readShapeAttr.size(); ++i) { + if (i == static_cast(dim)) { + resultShape.push_back(numIndices); + } else { + resultShape.push_back(readShapeAttr[i]); + } + } + + // Create result tensor type + inferredReturnTypes.push_back(RankedTensorType::get(resultShape, elemType)); + + return success(); +} + +// FlipOp +LogicalResult +FlipOp::inferReturnTypes(MLIRContext *context, std::optional location, + ValueRange operands, DictionaryAttr attributes, + OpaqueProperties properties, RegionRange regions, + SmallVectorImpl &inferredReturnTypes) { + auto inputTy = dyn_cast(operands[0].getType()); + if (!inputTy) { + if (location) + return emitOptionalError(location, + "expected ranked tensor for flip input"); + return failure(); + } + inferredReturnTypes.push_back(inputTy); + return success(); +} + +//-- SortOp -- +LogicalResult +SortOp::inferReturnTypes(MLIRContext *context, std::optional location, + ValueRange operands, DictionaryAttr attributes, + OpaqueProperties properties, RegionRange regions, + SmallVectorImpl &inferredReturnTypes) { + if (operands.size() != 1) { + return emitOptionalError(location, + "expected exactly one operand for SortOp"); + } + + if (!isa(operands[0].getType())) { + return emitOptionalError(location, + "operand must be a ranked tensor type for SortOp"); + } + + Value src = operands[0]; + auto srcTy = cast(src.getType()); + auto srcShape = srcTy.getShape(); + auto srcEnc = srcTy.getEncoding(); + + if (srcShape.empty()) { + return emitOptionalError(location, "input tensor must have rank >= 1"); + } + + Type sortedTy = + RankedTensorType::get(srcShape, srcTy.getElementType(), srcEnc); + + inferredReturnTypes.push_back(sortedTy); + + return success(); +} + +} // namespace mlir::triton::ascend diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/CMakeLists.txt new file mode 100755 index 00000000..27aac38f --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/CMakeLists.txt @@ -0,0 +1,11 @@ +add_triton_library(TritonStructuredIR + TritonStructuredOps.cpp + TritonStructuredDialect.cpp + + DEPENDS + TritonStructuredTableGen + + LINK_LIBS PUBLIC + TritonIR + MLIRIR + ) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/TritonStructuredDialect.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/TritonStructuredDialect.cpp new file mode 100755 index 00000000..2af19b8a --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/TritonStructuredDialect.cpp @@ -0,0 +1,22 @@ +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" + +using namespace mlir; +using namespace mlir::tts; + +/// Dialect creation, the instance will be owned by the context. This is the +/// point of registration of custom types and operations for the dialect. +void TritonStructuredDialect::initialize() { + addOperations< +#define GET_OP_LIST +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredOps.cpp.inc" + >(); +} + +//===----------------------------------------------------------------------===// +// TableGen'd op method definitions +//===----------------------------------------------------------------------===// + +#define GET_OP_CLASSES +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredOps.cpp.inc" + +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.cpp.inc" diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/TritonStructuredOps.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/TritonStructuredOps.cpp new file mode 100755 index 00000000..85038cae --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/TritonStructuredOps.cpp @@ -0,0 +1,267 @@ +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OperationSupport.h" +#include "mlir/Support/LogicalResult.h" + +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" + +#include +#include +#include + +#define GET_OP_CLASSES +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredOps.h.inc" + +using namespace mlir; +using namespace mlir::tts; + +namespace mlir { +namespace tts { + +namespace utils { +// Extract a scalar value from v. +// If v is a scalar, return that directly. Otherwise, parse through operations +// (currently only support splat, sitofp, and truncf) that produce it to +// extract the underlying scalar value. We then reconstruct the chain of +// operations that can produce this constant with the original type. If no +// scalar value can be extracted, a nullptr is returned. +Value getScalarValue(Value operand, Location loc, OpBuilder &builder) { + SmallVector ops; + + auto reconstructScalarValue = [&](Value src) { + for (auto op = ops.rbegin(); op != ops.rend(); ++op) { + src = TypeSwitch(*op) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return builder.create(loc, resType, src); + }) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return builder.create(loc, resType, src); + }) + .Default([](Operation *op) { + llvm_unreachable("unsupported op in generating "); + return nullptr; + }); + } + return src; + }; + + while (true) { + if (!dyn_cast(operand.getType())) { + return reconstructScalarValue(operand); + } else if (auto op = operand.getDefiningOp()) { + if (auto attr = dyn_cast(op.getValue())) { + if (!attr.isSplat()) { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load " + "produced by unsupported instruction"; + return nullptr; + } + auto elemValue = attr.getSplatValue(); + auto constOp = arith::ConstantOp::materialize( + builder, elemValue, attr.getElementType(), op.getLoc()); + return reconstructScalarValue(constOp.getResult()); + } + } else if (auto op = operand.getDefiningOp()) { + operand = op.getSrc(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load produced " + "by unsupported instruction"; + return nullptr; + } + } + return nullptr; +} + +} // namespace utils + +void MakeTensorPtrOp::build(OpBuilder &b, OperationState &state, Value base, + ArrayRef sizes, + ArrayRef strides, + ArrayRef offsets, + ArrayRef shape, + ArrayRef order) { + SmallVector staticStrides, staticOffsets, staticShape; + SmallVector dynamicStrides, dynamicOffsets, dynamicShape; + + dispatchIndexOpFoldResults(offsets, dynamicOffsets, staticOffsets); + dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides); + dispatchIndexOpFoldResults(shape, dynamicShape, staticShape); + + Type resType; + auto basePtr = cast(base.getType()); + auto elemType = basePtr.getPointeeType(); + // non-block pointer + if (order.empty()) { + resType = RankedTensorType::get(sizes, basePtr); + } + // block pointer + else { + resType = triton::PointerType::get(RankedTensorType::get(sizes, elemType), + basePtr.getAddressSpace()); + } + + build(b, state, resType, base, sizes, dynamicStrides, dynamicOffsets, + dynamicShape, b.getDenseI64ArrayAttr(staticStrides), + b.getDenseI64ArrayAttr(staticOffsets), + b.getDenseI64ArrayAttr(staticShape), order); +} + +void LoadOp::build(OpBuilder &b, OperationState &state, Value ptr, + ArrayRef dims, Value other) { + SmallVector staticDims; + SmallVector dynamicDims; + + dispatchIndexOpFoldResults(dims, dynamicDims, staticDims); + + // non-block pointer type + auto ptrTensorType = dyn_cast(ptr.getType()); + // block pointer type + auto tensorPtrType = dyn_cast(ptr.getType()); + + Type resType; + if (ptrTensorType) { + auto ptrType = cast(ptrTensorType.getElementType()); + auto elemType = ptrType.getPointeeType(); + resType = RankedTensorType::get(ptrTensorType.getShape(), elemType); + + } else if (tensorPtrType) { + auto tensorType = cast(tensorPtrType.getPointeeType()); + resType = RankedTensorType::get(tensorType.getShape(), + tensorType.getElementType()); + } + build(b, state, resType, ptr, dynamicDims, b.getDenseI64ArrayAttr(staticDims), + other); +} + +void StoreOp::build(OpBuilder &b, OperationState &state, Value ptr, Value value, + ArrayRef dims) { + SmallVector staticDims; + SmallVector dynamicDims; + + dispatchIndexOpFoldResults(dims, dynamicDims, staticDims); + + build(b, state, ptr, value, dynamicDims, b.getDenseI64ArrayAttr(staticDims)); +} + +void AtomicRMWOp::build(OpBuilder &b, OperationState &state, mlir::Type result, + Value ptr, Value value, ArrayRef dims, + triton::RMWOpAttr atomicRMWOp, + triton::MemSemanticAttr sem, + triton::MemSyncScopeAttr scope) { + SmallVector staticDims; + SmallVector dynamicDims; + + dispatchIndexOpFoldResults(dims, dynamicDims, staticDims); + + build(b, state, result, ptr, value, dynamicDims, + b.getDenseI64ArrayAttr(staticDims), atomicRMWOp, sem, scope); +} + +LogicalResult GetStructuredStateOp::verify() { + auto expectedOffsetAndStrideTypes = + getOffsetAndStrideTypes(getContext(), getInput().getType()); + + if (!expectedOffsetAndStrideTypes.has_value()) { + return failure(); + } + + auto [expectedOffsetTypes, expectedStrideTypes] = + *expectedOffsetAndStrideTypes; + + return success(expectedOffsetTypes.size() == getOffsets().size() && + llvm::equal(expectedOffsetTypes, getOffsets().getTypes()) && + expectedStrideTypes.size() == getStrides().size() && + llvm::equal(expectedStrideTypes, getStrides().getTypes())); +} + +void GetStructuredStateOp::build(OpBuilder &b, OperationState &state, + Value val) { + auto type = val.getType(); + + // Builder cannot fail, so we default to empty offset and stride types. + // The invalid op will be rejected by the verifier later. + auto [offsetTypes, strideTypes] = + getOffsetAndStrideTypes(b.getContext(), type) + .value_or(std::make_pair(SmallVector{}, SmallVector{})); + + build(b, state, val.getType(), offsetTypes, strideTypes, val); +} + +std::optional, SmallVector>> +GetStructuredStateOp::getOffsetAndStrideTypes(MLIRContext *context, Type type) { + auto sizes = getOffsetAndStrideSegmentSizes(type); + if (!sizes.has_value()) { + return std::nullopt; + } + return std::make_pair( + SmallVector(sizes->first, IndexType::get(context)), + SmallVector(sizes->second, IndexType::get(context))); +} + +std::optional> +GetStructuredStateOp::getOffsetAndStrideSegmentSizes(Type type) { + int32_t offsetSegmentSize = 0; + int32_t strideSegmentSize = 0; + + if (auto tensorType = llvm::dyn_cast(type)) { + if (tensorType.getElementType().isIntOrIndex()) { + // Tensors of offsets + // Important note: + // We only care about tensor of index / int (in addition to pointer type) + // because only values of int and index type can potentially be part of a + // pointer arithmetic sequence. + offsetSegmentSize = strideSegmentSize = tensorType.getRank(); + } else if (auto ptrType = + dyn_cast(tensorType.getElementType())) { + // Unstructured pointers (tensor>) + // Each tensor of rank k gets k values for its offsets and k values for + // its strides, all of which has Index type. + offsetSegmentSize = strideSegmentSize = tensorType.getRank(); + } + } + // Block pointers (!tt.ptr> or !tt.ptr) + else if (auto ptrType = llvm::dyn_cast(type)) { + if (auto tensorType = + llvm::dyn_cast(ptrType.getPointeeType())) { + // Each tensor of rank k gets k values for its offsets and k values for + // its strides, all of which has Index type. + offsetSegmentSize = strideSegmentSize = tensorType.getRank(); + } else { + // The only relevant state that can be updated in loops for scalar + // pointers are offset. No need to include stride here. + offsetSegmentSize = 1; + } + } else { + return std::nullopt; + } + + return std::make_pair(offsetSegmentSize, strideSegmentSize); +} + +} // namespace tts +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/CMakeLists.txt new file mode 100755 index 00000000..13514915 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/CMakeLists.txt @@ -0,0 +1,11 @@ +add_triton_library(TritonStructuredIncubatedIR + TritonStructuredOpsIncubated.cpp + TritonStructuredDialectIncubated.cpp + + DEPENDS + TritonStructuredIncubatedTableGen + + LINK_LIBS PUBLIC + TritonIR + MLIRIR + ) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.cpp new file mode 100755 index 00000000..f0e33404 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.cpp @@ -0,0 +1,22 @@ +#include "incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.h" + +using namespace mlir; +using namespace mlir::tts::Incubated; + +/// Dialect creation, the instance will be owned by the context. This is the +/// point of registration of custom types and operations for the dialect. +void TritonStructuredDialectIncubated::initialize() { + addOperations< +#define GET_OP_LIST +#include "incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredOpsIncubated.cpp.inc" + >(); +} + +//===----------------------------------------------------------------------===// +// TableGen'd op method definitions +//===----------------------------------------------------------------------===// + +#define GET_OP_CLASSES +#include "incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredOpsIncubated.cpp.inc" + +#include "incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.cpp.inc" diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/TritonStructuredOpsIncubated.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/TritonStructuredOpsIncubated.cpp new file mode 100755 index 00000000..1cf5851b --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/TritonStructuredOpsIncubated.cpp @@ -0,0 +1,111 @@ +#include "mlir/Bytecode/BytecodeOpInterface.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/OperationSupport.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" +#include "mlir/Support/LogicalResult.h" +#include "triton/Dialect/Triton/IR/Types.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/LogicalResult.h" +#include +#include +#include + +#define GET_OP_CLASSES +#include "incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.h" +using namespace mlir; +using namespace mlir::tts::Incubated; + +namespace mlir { +namespace tts { +namespace Incubated { + +LogicalResult GetStructuredStateOp::verify() { + auto expectedOffsetAndStrideTypes = + getOffsetAndStrideTypes(getContext(), getInput().getType()); + + if (!expectedOffsetAndStrideTypes.has_value()) { + return failure(); + } + + auto [expectedOffsetTypes, expectedStrideTypes] = + *expectedOffsetAndStrideTypes; + + return success(expectedOffsetTypes.size() == getOffsets().size() && + llvm::equal(expectedOffsetTypes, getOffsets().getTypes()) && + expectedStrideTypes.size() == getStrides().size() && + llvm::equal(expectedStrideTypes, getStrides().getTypes())); +} + +void GetStructuredStateOp::build(OpBuilder &b, OperationState &state, + Value val) { + auto type = val.getType(); + + // Builder cannot fail, so we default to empty offset and stride types. + // The invalid op will be rejected by the verifier later. + auto [offsetTypes, strideTypes] = + getOffsetAndStrideTypes(b.getContext(), type) + .value_or(std::make_pair(SmallVector{}, SmallVector{})); + + build(b, state, val.getType(), offsetTypes, strideTypes, val); +} + +std::optional, SmallVector>> +GetStructuredStateOp::getOffsetAndStrideTypes(MLIRContext *context, Type type) { + auto sizes = getOffsetAndStrideSegmentSizes(type); + if (!sizes.has_value()) { + return std::nullopt; + } + return std::make_pair( + SmallVector(sizes->first, IndexType::get(context)), + SmallVector(sizes->second, IndexType::get(context))); +} + +std::optional> +GetStructuredStateOp::getOffsetAndStrideSegmentSizes(Type type) { + int32_t offsetSegmentSize = 0; + int32_t strideSegmentSize = 0; + + if (auto tensorType = llvm::dyn_cast(type)) { + if (tensorType.getElementType().isIntOrIndex()) { + // Tensors of offsets + // Important note: + // We only care about tensor of index / int (in addition to pointer type) + // because only values of int and index type can potentially be part of a + // pointer arithmetic sequence. + offsetSegmentSize = strideSegmentSize = tensorType.getRank(); + } else if (auto ptrType = + dyn_cast(tensorType.getElementType())) { + // Unstructured pointers (tensor>) + // Each tensor of rank k gets k values for its offsets and k values for + // its strides, all of which has Index type. + offsetSegmentSize = strideSegmentSize = tensorType.getRank(); + } + } + // Block pointers (!tt.ptr> or !tt.ptr) + else if (auto ptrType = llvm::dyn_cast(type)) { + if (auto tensorType = + llvm::dyn_cast(ptrType.getPointeeType())) { + // Each tensor of rank k gets k values for its offsets and k values for + // its strides, all of which has Index type. + offsetSegmentSize = strideSegmentSize = tensorType.getRank(); + } else { + // The only relevant state that can be updated in loops for scalar + // pointers are offset. No need to include stride here. + offsetSegmentSize = 1; + } + } else { + return std::nullopt; + } + + return std::make_pair(offsetSegmentSize, strideSegmentSize); +} + +} // namespace Incubated +} // namespace tts +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/CMakeLists.txt new file mode 100755 index 00000000..f33061b2 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(IR) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/BufferizableOpInterfaceImpl.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/BufferizableOpInterfaceImpl.cpp new file mode 100755 index 00000000..921f7818 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/BufferizableOpInterfaceImpl.cpp @@ -0,0 +1,138 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/Linalg/Transforms/BufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Bufferization/IR/BufferizableOpInterface.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Bufferization/IR/DstBufferizableOpInterfaceImpl.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/Operation.h" + +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +using namespace mlir; +using namespace linalg; +using namespace mlir::bufferization; + +// +// This file implements the bufferizable interface for TritonTilingExtOps. +// The interface is required for bufferization (i.e: converting from tensors to +// memrefs). +// Since the bufferization semantics of TritonTilingExtOps are identical to +// linalg ops, the implementation was borrowed almost verbatim from +// mlir/lib/Dialect/Linalg/Transforms/BufferizableOpInterfaceImpl.cpp +// with the exception that the code to handle linalg's region has been removed. +// (the original implementation is in an anonymous namespace, so we cannot +// reuse) +// +namespace { + +/// Generic conversion for any DestinationStyleOpInterface on tensors. +static LogicalResult bufferizeTritonTilingExtDestinationStyleOpInterface( + RewriterBase &rewriter, DestinationStyleOpInterface op, + const BufferizationOptions &options, + BufferizationState &bufferizationState) { + // Take a guard before anything else. + OpBuilder::InsertionGuard g(rewriter); + rewriter.setInsertionPoint(op); + + // Nothing to do. This op is already bufferized. + if (op.hasPureBufferSemantics()) + return success(); + + // Ensure op has only tensors. Allow mixed tensor-buffer mode on a per-need + // basis. + if (!op.hasPureTensorSemantics()) + return op->emitError() << "op does not have tensor semantics"; + + // New input operands for the cloned op. + SmallVector newInputBuffers; + newInputBuffers.reserve(op.getNumDpsInputs()); + for (OpOperand *opOperand : op.getDpsInputOperands()) { + if (op.isScalar(opOperand)) { + newInputBuffers.push_back(opOperand->get()); + continue; + } + FailureOr buffer = + getBuffer(rewriter, opOperand->get(), options, bufferizationState); + if (failed(buffer)) + return failure(); + newInputBuffers.push_back(*buffer); + } + + // New output operands for the cloned op. + SmallVector newOutputBuffers; + for (OpResult opResult : op->getOpResults()) { + OpOperand *opOperand = op.getDpsInitOperand(opResult.getResultNumber()); + FailureOr resultBuffer = getBuffer( + rewriter, opOperand->get(), options, bufferizationState); + if (failed(resultBuffer)) + return failure(); + newOutputBuffers.push_back(*resultBuffer); + } + + // Merge input/output operands. + SmallVector newOperands = newInputBuffers; + newOperands.append(newOutputBuffers.begin(), newOutputBuffers.end()); + + // Set insertion point now that potential alloc/dealloc are introduced. + rewriter.setInsertionPoint(op); + // Clone the op, but use the new operands. Move the existing block into the + // new op. Since the new op does not have any tensor results, it does not + // return anything. + clone(rewriter, op, /*resultTypes=*/TypeRange{}, newOperands); + + // Replace the results of the old op with the new output buffers. + replaceOpWithBufferizedValues(rewriter, op, newOutputBuffers); + + return success(); +} + +template +struct TritonTilingExtOpInterface + : public DstBufferizableOpInterfaceExternalModel< + TritonTilingExtOpInterface, OpTy> { + bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand, + const AnalysisState &state) const { + // Operand is read if it is used in the computation. + return cast(op).isDpsInput(&opOperand); + } + + bool bufferizesToMemoryWrite(Operation *op, OpOperand &opOperand, + const AnalysisState &state) const { + // Operand is written to if it is not an input/init. + return cast(op).isDpsInit(&opOperand); + } + + LogicalResult bufferize(Operation *op, RewriterBase &rewriter, + const BufferizationOptions &options, + BufferizationState &bufferizationState) const { + return bufferizeTritonTilingExtDestinationStyleOpInterface( + rewriter, cast(op), options, + bufferizationState); + } +}; + +template struct TritonTilingExtOpInterfaceHelper { + static void registerOpInterface(MLIRContext *ctx) { + (Ops::template attachInterface>(*ctx), ...); + } +}; +} // namespace + +void mlir::ttx::registerBufferizableOpInterfaceExternalModels( + DialectRegistry ®istry) { + // clang-format off + registry.addExtension(+[](MLIRContext *ctx, ttx::TritonTilingExtDialect *dialect) { + TritonTilingExtOpInterfaceHelper< + ttx::CumSumOp + >::registerOpInterface(ctx); + }); + // clang-format on +} diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/CMakeLists.txt new file mode 100755 index 00000000..b6b07162 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/CMakeLists.txt @@ -0,0 +1,17 @@ +add_triton_library(TritonTilingExtIR + BufferizableOpInterfaceImpl.cpp + CumSum.cpp + TritonTilingExtDialect.cpp + + DEPENDS + TritonTilingExtInterfacesIncGen + TritonTilingExtOpsIncGen + + LINK_LIBS PUBLIC + TritonIR + MLIRAffineAnalysis + MLIRFuncDialect + MLIRIR + MLIRLinalgDialect + MLIRLinalgUtils + ) diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/CumSum.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/CumSum.cpp new file mode 100755 index 00000000..619b653e --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/CumSum.cpp @@ -0,0 +1,112 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +//===----------------------------------------------------------------------===// +// This file implements cumulative sum (CumSum) using the TilingInterface. Only +// supports tensors of rank 1 & 2 and axis == rank - 1 (i.e: we can split the +// computation of each row and compute them independently). The semantics of +// tiling for other axes are more complex and require working with +// non-contiguous memory. +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/Affine/IR/AffineOps.h" + +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "llvm/Support/Debug.h" + +#define DEBUG_TYPE "ttx-cumsum" + +using namespace mlir; +using namespace mlir::ttx; + +void ttx::CumSumOp::build(OpBuilder &odsBuilder, OperationState &odsState, + Value input, IntegerAttr axis, Value output, + ArrayRef attributes) { + SmallVector inputs{input}; + SmallVector outputs{output}; + odsState.addOperands(inputs); + odsState.addOperands(outputs); + odsState.addAttribute( + "operand_segment_sizes", + odsBuilder.getDenseI32ArrayAttr({static_cast(inputs.size()), + static_cast(outputs.size())})); + + odsState.addAttribute(getAxisAttrStrName(), axis); + odsState.addAttributes(attributes); + odsState.addTypes(SmallVector{output.getType()}); +} + +mlir::LogicalResult ttx::CumSumOp::verify() { + auto inputType = getInput().getType(); + if (!isa(inputType) && !isa(inputType)) { + return emitOpError( + "CumSum op expects input to be either tensor or memref."); + } + + auto outputType = getOutput().getType(); + if (!isa(outputType) && !isa(outputType)) { + return emitOpError( + "CumSum op expects output to be either tensor or memref."); + } + + if (dyn_cast(inputType).getShape() != + dyn_cast(outputType).getShape()) { + return emitOpError("Input and output types must be the same."); + } + + int64_t rank = getRank(); + if (rank != 1 && rank != 2) { + return emitOpError("CumSum op only takes tensors of rank 1 & 2."); + } + + int64_t axis = getAxis(); + if (axis != rank - 1) { + return emitOpError("CumSum computation only supports axis == rank - 1"); + } + + return success(); +} + +AffineMap ttx::CumSumOp::getInputIndexingMap(MLIRContext *context, + unsigned int index, + ArrayRef sizes) { + assert(index == 0); + return AffineMap::getMultiDimIdentityMap(getRank(), context); +} + +AffineMap ttx::CumSumOp::getOutputIndexingMap(MLIRContext *context, + unsigned int index, + ArrayRef sizes) { + assert(index == 0); + return AffineMap::getMultiDimIdentityMap(getRank(), context); +} + +SmallVector ttx::CumSumOp::getLoopIteratorTypes() { + SmallVector iterators; + iterators.append(getRank() - 1, utils::IteratorType::parallel); + iterators.push_back(utils::IteratorType::reduction); + return iterators; +} + +SmallVector ttx::CumSumOp::getIterationDomain(OpBuilder &b) { + OpBuilder::InsertionGuard g(b); + b.setInsertionPoint(*this); + auto loc = getLoc(); + auto zero = b.getIndexAttr(0); + auto one = b.getIndexAttr(1); + SmallVector iterationDomain; + + // Return the bounds for all dimensions. The caller is responsible for not + // tiling the inner most dimension, otherwise the semantic of the resulting op + // is incorrect. + for (auto i = 0; i < getRank(); i++) { + OpFoldResult upperbound = linalg::createFoldedDimOp(b, loc, getInput(), i); + iterationDomain.push_back(Range{zero, upperbound, one}); + } + return iterationDomain; +} diff --git a/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.cpp b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.cpp new file mode 100755 index 00000000..da4d9767 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.cpp @@ -0,0 +1,404 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Linalg/IR/LinalgInterfaces.h" +#include "mlir/Dialect/Linalg/Utils/Utils.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/IR/Value.h" + +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "llvm/ADT/ArrayRef.h" +#include "llvm/ADT/TypeSwitch.h" + +using namespace mlir; +using namespace mlir::ttx; +using namespace mlir::linalg; + +namespace mlir { +namespace ttx { + +Value getSlice(OpBuilder &b, Location loc, Value source, + ArrayRef offsets, ArrayRef sizes, + ArrayRef strides) { + return TypeSwitch(source.getType()) + .Case([&](RankedTensorType t) -> Value { + return b.create(loc, source, offsets, sizes, + strides); + }) + .Case([&](MemRefType type) -> Value { + return b.create(loc, source, offsets, sizes, + strides); + }) + .Default([&](Type t) { return nullptr; }); +} + +// +// getTiledImplementation +// +// Given an array of offsets and sizes, return the corresponding tiled version +// of the current op. +// +// This method is responsible for creating the extract slice ops for each +// operand of the op (including input and output operand). +// +// As an example, assuming we tile a linalg.matmul ins(%0, %1) out(%out) +// +// This method then generate: +// +// %in_slice_0 = extract_slice from %0 +// %in_slice_1 = extract_slice from %1 +// %out_slice = extract_slice from %out +// %tile = linalg.matmul ins(%in_slice_0, %in_slice_1) out(%out_slice) +// +// To generate these extract slice, we go over each operand, get the +// corresponding affine map to compute the correct offsets and sizes. +// +// Now let's describe how we compute the correct offsets and sizes from +// an affine map. +// +// - Offsets: +// An affine map describes how to access a tensor (i.e: the indicies into a +// tensor), so getting the offsets (also indices) from an affine map is just +// simply "applying" the sub-map on the offset (calling +// makeComposedFoldedAffineApply which also does constant folding +// automatically). +// +// For example: +// Let's assume we have the following nested loops: +// for i in range(0, 10): +// for j in range(0, 20): +// for k in range(0, 30): +// dst[i][j][k] = src[i * 2][j + k] +// +// Assume that we describe the iteration space based on dst. So: +// - dst's affine map is (d0, d1, d2) -> (d0, d1, d2) +// - src's affine map is (d0, d1, d2) -> (d0 * 2, d1 + d2) +// +// Now let's say we want to tile the operator with offset (0, 1, 2). +// +// For dst, we apply this (0, 1, 2) to its affine map and get (0, 1, 2) +// +// For src, we have to plug in the offsets into the affine map to get: +// +// (0 * 2, 1 + 2) = (0, 3) +// +// This is exactly what the implementation does as well. +// The call to getSubMap gets the i'th result expression, then the call to +// makeComposedFoldedAffineApply apply the `offsets` array to the i'th result +// expression in the affine map. +// +// +// - Sizes: +// Size is slightly more complex, notice that there are 3 steps to compute +// sizes: +// +// 1) call linalg::computeTileSizes on the provided `sizes` +// 2) apply the affine map +// 3) add 1 to the result +// +// The reason for this complexity is because the affine maps describe indices +// iteration space with a half open interval (i.e.: we always from 0 until right +// before the upper bound). So if we simply apply the affine map on the sizes, +// we will have incorrect results. +// +// Consider this snippet again: +// for i in range(0, 16): +// for j in range(0, 32): +// for k in range(0, 64): +// dst[i][j][k] = src[i * 2][j + k] +// +// Assume we want the operator to have a tile size of (16, 32, 64) -- so no +// tiling at all. If we apply the affine map of src (d0, d1, d2) -> (d0 * 2, d1 +// + d2), we have +// +// (16 * 2, 32 + 64) -> (32, 96) +// +// However, consider the second dimension of source: +// - j goes from 0 till 31 inclusive +// - k goes from 0 till 63 inclusive +// +// So the max index of src's second dimension is 31 + 63 = 94. Since index +// starts from 0, this means the second dimension has 95 elements. But the +// formula gives us a tile size of 96!!! The same argument can be applied for +// the first dimension as well, the number of elements is 15 * 2 + 1 = 31, but +// computed tile size is 32. +// +// So simply applying the indexing map to compute tile size is INCORRECT!! +// This happens because the indexing map operates on [0, size), while tile sizes +// are inclusive. +// +// The correct formula is: +// ((d0 - 1) * 2 + 1), (d1 - 1) + (d2 - 1) + 1 which gives +// (15 * 2 + 1, 32 - 1 + 64 - 1 + 1) -> (31, 95) +// +// So again, the steps are: +// - Subtract 1 from the sizes (what linalg::computeTileSizes does) +// - Apply the affine map +// - Add 1 to the result +// +template +FailureOr getTiledImplementation(TritonTilingExtOpTy op, + OpBuilder &b, + ArrayRef offsets, + ArrayRef sizes) { + Location loc = op->getLoc(); + SmallVector valuesToTile = op->getOperands(); + SmallVector tiledValues; + auto oneAttr = b.getI64IntegerAttr(1); + + for (OpOperand &opOperand : op->getOpOperands()) { + unsigned int index = opOperand.getOperandNumber(); + auto val = valuesToTile[index]; + auto type = dyn_cast(val.getType()); + + if (!type) { + tiledValues.push_back(val); + continue; + } + + auto rank = type.getRank(); + SmallVector newOffsets; + SmallVector newSizes; + SmallVector newStrides(rank, oneAttr); + + llvm::SmallVector composedTileSizes = + linalg::computeTileSizes(b, loc, sizes, {}); + + AffineMap map = op.getIndexingMap(b.getContext(), index, sizes); + for (int64_t i = 0; i < rank; i++) { + AffineMap m = map.getSubMap(i); + { + OpFoldResult upperboundClosed = + affine::makeComposedFoldedAffineApply(b, loc, m, composedTileSizes); + AffineExpr s0 = getAffineSymbolExpr(0, b.getContext()); + OpFoldResult size = affine::makeComposedFoldedAffineApply( + b, loc, s0 + 1, upperboundClosed); + newSizes.push_back(size); + } + { + OpFoldResult offset = + affine::makeComposedFoldedAffineApply(b, loc, m, offsets); + newOffsets.push_back(offset); + } + } + + tiledValues.push_back( + getSlice(b, loc, val, newOffsets, newSizes, newStrides)); + } + + SmallVector resultTensorTypes = llvm::to_vector( + llvm::map_range(op.getDpsInitsMutable(), [&](OpOperand &opOperand) { + return tiledValues[opOperand.getOperandNumber()].getType(); + })); + + Operation *tiledOp = clone(b, op, resultTensorTypes, tiledValues); + + return TilingResult{{tiledOp}, SmallVector(tiledOp->getResults())}; +} + +// +// getResultTilePosition +// This method returns the resultOffsets and resultSizes through references +// for the tiled operator. While `getTiledImplementation` is responsible for +// generating the extract slice for all operands, `getResultTilePosition` is +// responsible for returning the offsets and sizes which the tiling engine will +// then use to generate the corresponding insert_slice ops. +// +// Because the slice we insert back to the output tensor is the same as the +// slice that we extracted from the output tensor, this method just repeats the +// offset and size computation in `getTiledImplementation`. +// +template +LogicalResult getResultTilePosition(TritonTilingExtOpTy op, OpBuilder &b, + unsigned resultNumber, + ArrayRef offsets, + ArrayRef sizes, + SmallVector &resultOffsets, + SmallVector &resultSizes) { + Location loc = op.getLoc(); + + AffineMap outputMap = + op.getOutputIndexingMap(b.getContext(), resultNumber, sizes); + + Value result = op.getDpsInitOperand(resultNumber)->get(); + auto rank = dyn_cast(result.getType()).getRank(); + + llvm::SmallVector composedTileSizes = + linalg::computeTileSizes(b, loc, sizes, {}); + for (int64_t i = 0; i < rank; i++) { + AffineMap m = outputMap.getSubMap(i); + { + OpFoldResult upperboundClosed = + affine::makeComposedFoldedAffineApply(b, loc, m, composedTileSizes); + AffineExpr s0 = getAffineSymbolExpr(0, b.getContext()); + OpFoldResult size = affine::makeComposedFoldedAffineApply( + b, loc, s0 + 1, upperboundClosed); + resultSizes.push_back(size); + } + { + OpFoldResult offset = + affine::makeComposedFoldedAffineApply(b, loc, m, offsets); + resultOffsets.push_back(offset); + } + } + + return success(); +} + +// This method is borrowed verbatim from +// mlir/lib/Dialect/Linalg/Transforms/TilingInterfaceImpl.cpp +// +// This is invoked when the current op produces a result that is used +// as an input to another op that is being tiled. The method essentially handles +// producing a new op where the result matches the given offsets and sizes. +// If the method succeeds, the two new operators will be fused in the same loop. +// +// As an example, consider the following IR where the linalg.generic is being +// tiled (unnecessary detailed omitted for brevity): +// +// clang-format: off +// +// func.func @some_op_1( +// %arg0: tensor<8x2x256x512xbf16>, +// %arg1: tensor<8x256x1024xbf16> +// ) -> tensor<8x256x1024xbf16> +// %1 = linalg.init_tensor [8, 256, 1024] : tensor<8x256x1024xbf16> +// %2 = linalg.init_tensor [8, 256, 1024] : tensor<8x256x1024xbf16> +// %3 = ttx.some_op +// ins(%arg0 : tensor<8x2x256x512xbf16>) +// outs(%1 : tensor<8x256x1024xbf16>) -> tensor<8x256x1024xbf16> +// %4 = linalg.generic +// ins(%3, %arg1 : tensor<8x256x1024xbf16>, tensor<8x256x1024xbf16>) +// outs(%2 : tensor<8x256x1024xbf16>) { +// ^bb0(%arg2: bf16, %arg3: bf16, %arg4: bf16): +// %add = arith.addf %arg2, %arg3 : bf16 +// linalg.yield %add : bf16 +// } -> tensor<8x256x1024xbf16> +// return %4 : tensor<8x256x1024xbf16> +// } +// +// clang-format: on +// +// We tile linalg.generic, but one of its inputs is %3 which is the result of +// ttx.some_op. So the tiling engine will invoke +// generateResultTileValue of ttx.some_op to determine if it's +// possible to create a tiled version of it, thereby making it possible to fuse +// both operators together in a loop. +template +FailureOr +generateResultTileValue(TritonTilingExtOpTy op, OpBuilder &b, + unsigned resultNumber, ArrayRef offsets, + ArrayRef sizes) { + + // Check that the indexing map used for the output is a projected + // permutation. This could be relaxed with a more general approach that can + // map the offsets and sizes from the result to iteration space tiles + // (filling in full extent for dimensions not used to access the result). + AffineMap indexingMap = op.getOutputIndexingMap(b.getContext(), 0, sizes); + if (!indexingMap.isProjectedPermutation()) { + return op.emitOpError( + "unhandled tiled implementation generation when result is not " + "accessed using a permuted projection"); + } + + auto numLoops = op.getLoopIteratorTypes().size(); + SmallVector iterationTileOffsets(numLoops), + iterationTileSizes(numLoops); + if (!indexingMap.isPermutation()) { + SmallVector iterationDomain = op.getIterationDomain(b); + for (auto range : llvm::enumerate(iterationDomain)) { + iterationTileOffsets[range.index()] = range.value().offset; + iterationTileSizes[range.index()] = range.value().size; + } + } + for (auto resultExpr : llvm::enumerate(indexingMap.getResults())) { + assert(resultExpr.value().getKind() == AffineExprKind::DimId); + // HACK: LLVM casting utilities do not work here for out-of-tree builds, + // as there is no template specialization for this cast in the base + // build. + AffineDimExpr affineDimExpr(static_cast( + const_cast(resultExpr.value().getAsOpaquePointer()))); + unsigned dimPosition = affineDimExpr.getPosition(); + iterationTileOffsets[dimPosition] = offsets[resultExpr.index()]; + iterationTileSizes[dimPosition] = sizes[resultExpr.index()]; + } + + FailureOr tilingResult = + op.getTiledImplementation(b, iterationTileOffsets, iterationTileSizes); + if (tilingResult->tiledOps.size() != 1) + return op.emitOpError("failed to generate tiled implementation"); + + return TilingResult{ + tilingResult->tiledOps, + SmallVector{tilingResult->tiledValues[resultNumber]}}; +} + +// This method is borrowed directly from linalg.generic's implementation +// in mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp +// This marks all operands that are part of the input group to have read +// effect, while all other operands that are part of the output group +// to have both read and write effects. +static void getTritonTilingExtEffectsImpl( + SmallVectorImpl> + &effects, + ValueRange results, ArrayRef inputOperands, + const MutableOperandRange &outputOperands) { + for (auto operand : inputOperands) { + if (!llvm::isa(operand->get().getType())) + continue; + effects.emplace_back(MemoryEffects::Read::get(), operand, /*stage=*/0, + /*effectOnFullRegion=*/true, + SideEffects::DefaultResource::get()); + } + for (auto &operand : outputOperands) { + if (!llvm::isa(operand.get().getType())) + continue; + + effects.emplace_back(MemoryEffects::Read::get(), &operand, /*stage=*/0, + /*effectOnFullRegion=*/true, + SideEffects::DefaultResource::get()); + effects.emplace_back(MemoryEffects::Write::get(), &operand, /*stage=*/0, + /*effectOnFullRegion=*/true, + SideEffects::DefaultResource::get()); + } +} + +template +void getEffects( + TritonTilingExtOpTy op, + SmallVectorImpl> + &effects) { + getTritonTilingExtEffectsImpl(effects, op.getOperation()->getResults(), + op.getDpsInputOperands(), + op.getDpsInitsMutable()); +} + +} // namespace ttx +} // namespace mlir + +/// Dialect creation, the instance will be owned by the context. This is the +/// point of registration of custom types and operations for the dialect. +void TritonTilingExtDialect::initialize() { + addOperations< +#define GET_OP_LIST +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOps.cpp.inc" + >(); +} + +//===----------------------------------------------------------------------===// +// TableGen'd op method definitions +//===----------------------------------------------------------------------===// + +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtInterfaces.cpp.inc" + +#define GET_OP_CLASSES +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOps.cpp.inc" + +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOpsDialect.cpp.inc" diff --git a/third_party/wafer/third_party/flir/lib/Utils/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/Utils/CMakeLists.txt new file mode 100755 index 00000000..9dde3c83 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Utils/CMakeLists.txt @@ -0,0 +1,6 @@ +add_triton_library(TritonSharedUtils + Utils.cpp + + LINK_LIBS PUBLIC + TritonIR +) diff --git a/third_party/wafer/third_party/flir/lib/Utils/Utils.cpp b/third_party/wafer/third_party/flir/lib/Utils/Utils.cpp new file mode 100755 index 00000000..642668c4 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/Utils/Utils.cpp @@ -0,0 +1,305 @@ +#include "triton-shared/Utils/Utils.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" +#include "triton-shared/Utils/ReduceScanCommon.h" +#include "llvm/ADT/TypeSwitch.h" + +namespace mlir { +namespace triton { +bool isPtrTypeLike(Type t) { + if (auto tensorType = dyn_cast(t)) { + return isa(tensorType.getElementType()); + } + return isa(t); +} + +Value getScalarValue(Value operand, Location loc, OpBuilder &builder) { + SmallVector ops; + + auto reconstructScalarValue = [&](Value src) { + for (auto op = ops.rbegin(); op != ops.rend(); ++op) { + src = TypeSwitch(*op) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return builder.create(loc, resType, src); + }) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return builder.create(loc, resType, src); + }) + .Default([](Operation *op) { + llvm_unreachable("unsupported op in generating "); + return nullptr; + }); + } + return src; + }; + + while (true) { + if (!dyn_cast(operand.getType())) { + return reconstructScalarValue(operand); + } else if (auto op = operand.getDefiningOp()) { + if (auto attr = dyn_cast(op.getValue())) { + if (!attr.isSplat()) { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load " + "produced by unsupported instruction"; + return nullptr; + } + auto elemValue = attr.getSplatValue(); + auto constOp = arith::ConstantOp::materialize( + builder, elemValue, attr.getElementType(), op.getLoc()); + return reconstructScalarValue(constOp.getResult()); + } + } else if (auto op = operand.getDefiningOp()) { + operand = op.getSrc(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load produced " + "by unsupported instruction"; + return nullptr; + } + } + return nullptr; +} + +bool isOperandMemorySpaceSPM(Value operand) { + Operation *lastOp = operand.getDefiningOp(); + Operation *op = lastOp; + // May be nested scf::ForOp block arguments + if (!op && isa(operand)) { + auto argBlock = operand.getParentBlock()->getParentOp(); + if (auto funcOp = dyn_cast(argBlock)) { + return false; + } + auto forOp = dyn_cast(argBlock); + assert(forOp && "BlockArgument should be in a scf::ForOp"); + + auto initArgs = forOp.getInitArgs(); + auto arguments = forOp.getBody()->getArguments(); + + auto idx = + std::distance(arguments.begin(), + std::find(arguments.begin(), arguments.end(), operand)); + assert(initArgs.size() + forOp.getNumInductionVars() == arguments.size() && + "InitArgs and InductionVars should match the arguments size"); + + int initArgIdx = idx - forOp.getNumInductionVars(); + assert(initArgIdx >= 0 && initArgIdx < initArgs.size() && + "Index out of bounds for initArgs"); + operand = initArgs[idx - forOp.getNumInductionVars()]; + return isOperandMemorySpaceSPM(operand); + } + + do { + if (isa(op)) + return true; + else if (isa(op)) + return false; + else if (auto forOp = dyn_cast(op)) { + // Here we assume that yieldResults (inner loop region) and + // loopResults (outer loop region) correspond one-to-one to obtain the + // inner loop region definingOp of the outer loop region value. + // FIXME: Need reference the standard loop analysis to refactor this. + + auto yieldResults = forOp.getYieldedValues(); + mlir::ResultRange loopResults = forOp.getLoopResults().value(); + assert(yieldResults.size() == loopResults.size()); + auto idx = std::distance( + loopResults.begin(), + std::find(loopResults.begin(), loopResults.end(), operand)); + operand = yieldResults[idx]; + if (operand.getDefiningOp() == nullptr) { + operand = forOp.getInitArgs()[idx]; + } + } else if (auto ifOp = dyn_cast(op)) { + bool thenResult = isOperandMemorySpaceSPM(ifOp.thenYield().getOperand(0)); + bool elseResult = isOperandMemorySpaceSPM(ifOp.elseYield().getOperand(0)); + assert(thenResult == elseResult && + "Inconsistent memory space for IfOp results: " + "one branch uses SPM, another branch does not"); + return thenResult; + } else if (auto selectOp = dyn_cast(op)) { + // Assuming that the selectOp is used to select between two pointers with + // same memory space, we can check the memory space of the first operand. + operand = op->getOperand(1); + } else { + operand = op->getOperand(0); + } + lastOp = op; + op = operand.getDefiningOp(); + } while (op); + return false; +} + +// Function to declare Wafer runtime function +Value declareWaferRuntimeFunction(ModuleOp module, OpBuilder &builder, Location loc, + StringRef name, Type resultType, + ArrayRef argumentTypes) { + // Check if the function already exists + Operation *funcOp = module.lookupSymbol(name); + if (funcOp) + return builder.create( + loc, LLVM::LLVMPointerType::get(builder.getContext()), name); + + // Create function type + Type funcType = LLVM::LLVMFunctionType::get(resultType, argumentTypes, + /*isVarArg=*/false); + + // Create a function declaration + auto ip = builder.saveInsertionPoint(); + builder.setInsertionPointToStart(module.getBody()); + + builder.create(loc, name, funcType, + LLVM::Linkage::External); + + builder.restoreInsertionPoint(ip); + + // Return function pointer + return builder.create( + loc, LLVM::LLVMPointerType::get(builder.getContext()), name); +} + +TypedAttr getRedBaseAttr(OpBuilder &builder, Operation *redOp, + Type constantType) { + const int64_t bitWidth = constantType.getIntOrFloatBitWidth(); + + auto attr = llvm::TypeSwitch(redOp) + .Case([&](arith::AddFOp) { + return builder.getFloatAttr(constantType, 0.f); + }) + .Case([&](arith::AddIOp) { + return builder.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::MulFOp) { + return builder.getFloatAttr(constantType, 1.f); + }) + .Case([&](arith::MulIOp) { + return builder.getIntegerAttr(constantType, 1); + }) + .Case([&](auto) { + return builder.getFloatAttr( + constantType, -std::numeric_limits::infinity()); + }) + .Case([&](auto) { + return builder.getFloatAttr( + constantType, std::numeric_limits::infinity()); + }) + .Case([&](arith::MinSIOp) { + return builder.getIntegerAttr(constantType, + llvm::maxIntN(bitWidth)); + }) + .Case([&](arith::MinUIOp) { + return builder.getIntegerAttr(constantType, + llvm::maxUIntN(bitWidth)); + }) + .Case([&](arith::MaxSIOp) { + return builder.getIntegerAttr(constantType, + llvm::minIntN(bitWidth)); + }) + .Case([&](arith::MaxUIOp) { + return builder.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::OrIOp) { + return builder.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::XOrIOp) { + return builder.getIntegerAttr(constantType, 0); + }) + .Case([&](arith::AndIOp) { + return builder.getIntegerAttr(constantType, + llvm::maxUIntN(bitWidth)); + }) + .Default([](Operation *op) { + op->dump(); + llvm_unreachable("Reduction op not yet supported"); + return nullptr; + }); + return attr; +} + +arith::ConstantOp getRedBaseConstOp(ConversionPatternRewriter &rewriter, + Operation *redOp, Type constantType) { + auto attr = getRedBaseAttr(rewriter, redOp, constantType); + return rewriter.create(redOp->getLoc(), constantType, + attr); +} + +bool isTypeRestrictedTargetSupportedReductionOp(mlir::Operation *redOp) { + return isa(redOp); +} + +bool isTargetSupportedReductionOp(mlir::Operation *redOp) { + return isTypeRestrictedTargetSupportedReductionOp(redOp) || + isa(redOp); +} + +bool isReduceLogicOp(mlir::Operation *redOp) { + return isa(redOp); +} + +bool isTypeRestrictedTargetSupportedReduceToElementWiseOp( + mlir::Operation *redOp) { + return isa(redOp) || isReduceLogicOp(redOp); +} + +bool isTargetSupportedReduceToElementWiseOp(mlir::Operation *redOp) { + return isTypeRestrictedTargetSupportedReduceToElementWiseOp(redOp) || + isa(redOp); +} + +bool isTritonAllowedReductionOp(Operation *redOp) { + return isTargetSupportedReductionOp(redOp) || + isTargetSupportedReduceToElementWiseOp(redOp) || + isa(redOp); +} + +bool isTargetSupportedFloatType(Type elementType) { + // Check if the operation is a supported type. + return elementType.isBF16() || elementType.isF16() || elementType.isF32() || + elementType.isTF32(); +} + +bool isTargetSupportedType(Type elementType) { + // Check if the operation is a supported type. + return isTargetSupportedFloatType(elementType) || elementType.isInteger(8); +} + +bool isReductionOpAndTypeSupportedByTarget(mlir::Operation *redOp, + Type elementType) { + return isTypeRestrictedTargetSupportedReductionOp(redOp) && + isTargetSupportedFloatType(elementType); +} + +bool isReduceToElementWiseOpAndTypeSupportedByTarget(mlir::Operation *redOp, + Type elementType, + int64_t elemCount, + int64_t rank) { + + // I1 type need special handle (Memref type i1 need 8bit alignment) + // NOTE: Here can optimize 64 to 8 element + return (isa(redOp) && + isTargetSupportedFloatType(elementType)) || + (isReduceLogicOp(redOp) && + !(elementType.isInteger(1) && (rank > 1 || elemCount <= 64))); +} + +} // namespace triton + +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/UtilsIncubated/CMakeLists.txt b/third_party/wafer/third_party/flir/lib/UtilsIncubated/CMakeLists.txt new file mode 100755 index 00000000..b6aa5164 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/UtilsIncubated/CMakeLists.txt @@ -0,0 +1,8 @@ +add_triton_library(MLIRTritonNPUUtils + Utils.cpp + InterleaveOptimization.cpp + + LINK_LIBS PUBLIC + MLIRIR + TritonIR +) diff --git a/third_party/wafer/third_party/flir/lib/UtilsIncubated/InterleaveOptimization.cpp b/third_party/wafer/third_party/flir/lib/UtilsIncubated/InterleaveOptimization.cpp new file mode 100755 index 00000000..1c421706 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/UtilsIncubated/InterleaveOptimization.cpp @@ -0,0 +1,677 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/UtilsIncubated/InterleaveOptimization.h" +#include "incubated/Conversion/UtilsIncubated/Utils.h" +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/Interfaces/ViewLikeInterface.h" +#include "mlir/Support/LogicalResult.h" + +#include "mlir/IR/Operation.h" +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "llvm/ADT/SmallVector.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include +#include + +namespace mlir { +namespace triton { +// For origin MemRefType of ReinterpretCastOp under interleave state, here wanna +// adjust its shape info by expanding last dimension double. +MemRefType expandInterleaveMemRefType(MemRefType originType) { + // Double the last dimension shape + SmallVector shape(originType.getShape()); + shape.back() = shape.back() * 2; + + // Adjuest layout attribute + StridedLayoutAttr originLayout = + llvm::dyn_cast(originType.getLayout()); + // If offset is static, just reset it to 0 + auto offset = originLayout.getOffset() == ShapedType::kDynamic + ? originLayout.getOffset() + : 0; + // Set last dimension stride to 1 + SmallVector stride(originLayout.getStrides()); + stride.back() = 1; + + return MemRefType::get( + shape, originType.getElementType(), + StridedLayoutAttr::get(originType.getContext(), offset, stride)); +} + +// ********************* +// ** NOTE ** +// ********************* +// How to determine new offset is a little tricky and specific +// Here just consider this state in triton language: +// +// dim_range = tl.arange(0, BLOCK // 2) +// last_dim_even_range = dim_range * 2 +// last_dim_odd_range = dim_range * 2 + 1 +// +// Here `multiply two` represents that last dimension stride is 2, and +// `add constant one` represents whether it's odd index part of +// deinterleave result. +// +// Therefore, how to distinguish interleave/deinterleave on even index or odd +// index is whether last dimension range explicitly `add constant one` without +// any other operation. In IR it's shown that whether defining op of +// `castOffset` is an arith::addOp, as this arith::addOp would contain above +// `add constant one` opeartion after LegacyAddPtrConverter. +// +// Well, index mode should be passed to interleave/deinterleave, in other words, +// `add constant one` should work on offset of next insert_slice/extract_slic. +// The new reinterpretcast just wanna describe whole tensor, so new castOffset +// is just from non-last diemsnion accumulation and remove `add constant one` +std::pair +recountReinterpretCastOffset(OpFoldResult originOffset, Builder &builder) { + // To trace value type offset + std::function traceOffset = [&](Operation *op) -> bool { + // Consider constant one in `add constant one` operation + if (llvm::isa(op)) + return false; + + if (llvm::isa(op)) { + auto addOp = llvm::cast(op); + if (auto constLHS = addOp.getLhs().getDefiningOp()) { + assert(dyn_cast(constLHS.getValueAttr()).getInt() == 1 && + "Arith::constant value of addi's operand must be 1 when " + "calculate deinterleave offset"); + return false; + } + if (auto constRHS = addOp.getRhs().getDefiningOp()) { + assert(dyn_cast(constRHS.getValueAttr()).getInt() == 1 && + "Arith::constant value of addi's operand must be 1 when " + "calculate deinterleave offset"); + return false; + } + } + return true; + }; + + IndexMode evenOrOdd = IndexMode::EVEN_MODE; + // Reuse origin offset if there's no 'add constant one' + OpFoldResult newOffset = originOffset; + if (llvm::isa(originOffset)) { + // If offset is constant int(IndexAttr), + // the int value could only be 0 or 1 + int64_t intOffset = getConstantIntValue(originOffset).value(); + assert((intOffset == 0 || intOffset == 1)); + if (intOffset == 1) { + evenOrOdd = IndexMode::ODD_MODE; + newOffset = builder.getIndexAttr(0); + } + } else if (llvm::isa(originOffset)) { + if (!traceOffset(originOffset.get().getDefiningOp())) { + evenOrOdd = IndexMode::ODD_MODE; + Operation *traceResult = findFirstMatchingOperandDef( + originOffset.get().getDefiningOp(), traceOffset); + assert(traceResult->getNumResults() == 1 && + "Offset defining operation must have one result"); + newOffset = traceResult->getResult(0); + } + } + + return {newOffset, evenOrOdd}; +} + +LogicalResult +DeinterleaveStatusOptimization(triton::LoadOp op, + triton::LoadOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) { + auto ptr = adaptor.getPtr(); + if (auto reinterpretCast = ptr.getDefiningOp()) { + auto loc = op.getLoc(); + + // 1. Get new source memref type + auto srcType = expandInterleaveMemRefType(reinterpretCast.getType()); + + // 2. Create new ReinterpretCastOp + auto originCastOffset = reinterpretCast.getConstifiedMixedOffset(); + auto castSize = reinterpretCast.getConstifiedMixedSizes(); + auto castStride = reinterpretCast.getConstifiedMixedStrides(); + // Actually, `castSize` is always constant value as `MemRefType` result + if (auto lastDimSize = getConstantIntValue(castSize.back())) { + castSize.back() = rewriter.getIndexAttr(lastDimSize.value() * 2); + } else { + return failure(); + } + // Last element of castStride is also constant value as prerequisite + // is that last dimension stride of casted memref type is always 2. + castStride.back() = rewriter.getIndexAttr(1); + auto [castOffset, indexMode] = + recountReinterpretCastOffset(originCastOffset, rewriter); + auto newCastOp = rewriter.create( + loc, srcType, reinterpretCast.getViewSource(), castOffset, castSize, + castStride); + + // 3. Create new memref allocOp + auto newAllocOp = rewriter.create( + loc, MemRefType::get(srcType.getShape(), srcType.getElementType())); + + // 4. Implement memref copy and bufferization back to tensor + rewriter.create(loc, newCastOp.getResult(), newAllocOp); + Value newTensor = rewriter.create( + loc, + RankedTensorType::get(srcType.getShape(), srcType.getElementType()), + newAllocOp, true /* restrict */, true /* writable */); + + // 5. Implement tensor extract_slice to represent deinterleave + // Here use `castOffset` to determine whether even index deinterleave or + // odd index. + SmallVector extractOffsets(srcType.getRank(), + rewriter.getIndexAttr(0)); + SmallVector extractStrides(srcType.getRank(), + rewriter.getIndexAttr(1)); + SmallVector extractSizes = llvm::to_vector( + llvm::map_range(srcType.getShape(), [&](int64_t dim) -> OpFoldResult { + return rewriter.getIndexAttr(dim); + })); + + // Adjust extract_slice shape + switch (indexMode) { + case IndexMode::EVEN_MODE: + extractOffsets.back() = rewriter.getIndexAttr(0); + break; + case IndexMode::ODD_MODE: + extractOffsets.back() = rewriter.getIndexAttr(1); + break; + } + extractStrides.back() = rewriter.getIndexAttr(2); + extractSizes.back() = rewriter.getIndexAttr(srcType.getShape().back() / 2); + + Value deinterleaveSlice = rewriter.create( + loc, newTensor, extractOffsets, extractSizes, extractStrides); + + rewriter.replaceOp(op, deinterleaveSlice); + return success(); + } + + return failure(); +} + +LogicalResult DeinterleaveStatusWithMaskOptimization( + triton::LoadOp op, triton::LoadOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter, + mlir::triton::Incubated::MaskState &mstate, Value localMem) { + auto ptr = adaptor.getPtr(); + if (auto reinterpretCast = ptr.getDefiningOp()) { + auto loc = op.getLoc(); + + // 1. Get new source memref type + auto srcType = expandInterleaveMemRefType(reinterpretCast.getType()); + + // 2. Create new ReinterpretCastOp + auto originCastOffset = reinterpretCast.getConstifiedMixedOffset(); + auto castSize = reinterpretCast.getConstifiedMixedSizes(); + auto castStride = reinterpretCast.getConstifiedMixedStrides(); + + if (auto lastDimSize = getConstantIntValue(castSize.back())) { + castSize.back() = rewriter.getIndexAttr(lastDimSize.value() * 2); + } else { + return failure(); + } + castStride.back() = rewriter.getIndexAttr(1); + auto [castOffset, indexMode] = + recountReinterpretCastOffset(originCastOffset, rewriter); + + auto newCastOp = rewriter.create( + loc, srcType, reinterpretCast.getViewSource(), castOffset, castSize, + castStride); + + // 3. Create new memref allocOp + // To reuse existing linalg::fill, here need to change insertion point + auto savedInsertPoint = rewriter.saveInsertionPoint(); + rewriter.setInsertionPointAfterValue(localMem); + auto newAllocOp = rewriter.create( + loc, MemRefType::get(srcType.getShape(), srcType.getElementType())); + rewriter.restoreInsertionPoint(savedInsertPoint); + + // 4. Broadcast other value by linalg.fill if necessary + auto other = op.getOther(); + // While deinterleave optimization will just adjust last dimension info + // and origin mask state wouldn't involve last dimension. Therefore in + // current `scf.if + linalg.fill` combination, condition of `if` could be + // kept and just replace linalg.fill' + if (other) { + assert(localMem.hasOneUse() && + llvm::isa(*(localMem.getUsers().begin()))); + auto originFillOp = + llvm::dyn_cast(*(localMem.getUsers().begin())); + + assert(llvm::isa(originFillOp->getParentOp())); + auto ifOp = llvm::dyn_cast(originFillOp->getParentOp()); + + auto newFillOp = ifOp.getThenBodyBuilder().create( + originFillOp.getLoc(), originFillOp.getInputs(), + ValueRange{newAllocOp}); + rewriter.replaceOp(originFillOp, newFillOp); + } + + // 5. Implement new subview, memref copy and bufferization back to tensor + SmallVector subviewStrides(srcType.getRank(), + rewriter.getIndexAttr(1)); + SmallVector subviewOffsets = mstate.offsets; + SmallVector subviewSizes = mstate.dims; + // Just adjust last dimension size to double + std::optional originSubviewLastDim = + getConstantIntValue(subviewSizes.back()); + assert(originSubviewLastDim.has_value()); + subviewSizes.back() = + rewriter.getIndexAttr(originSubviewLastDim.value() * 2); + + auto argSubviewType = memref::SubViewOp::inferResultType( + srcType, subviewOffsets, subviewSizes, subviewStrides); + // alloca subview type doesn't carry layout attribute + auto allocSubviewType = memref::SubViewOp::inferResultType( + newAllocOp.getType(), subviewOffsets, subviewSizes, subviewStrides); + + memref::SubViewOp srcSubview = rewriter.create( + loc, llvm::cast(argSubviewType), newCastOp, subviewOffsets, + subviewSizes, subviewStrides); + memref::SubViewOp dstSubview = rewriter.create( + loc, llvm::cast(allocSubviewType), newAllocOp, + subviewOffsets, subviewSizes, subviewStrides); + rewriter.create(loc, srcSubview, dstSubview); + Value newTensor = rewriter.create( + loc, + RankedTensorType::get(srcType.getShape(), srcType.getElementType()), + newAllocOp, true /* restrict */, true /* writable */); + + // 6. Implement tensor extract_slice to represent deinterleave + // Here use `castOffset` to determine whether even index deinterleave or + // odd index. + SmallVector extractOffsets(srcType.getRank(), + rewriter.getIndexAttr(0)); + SmallVector extractStrides(srcType.getRank(), + rewriter.getIndexAttr(1)); + SmallVector extractSizes = llvm::to_vector( + llvm::map_range(srcType.getShape(), [&](int64_t dim) -> OpFoldResult { + return rewriter.getIndexAttr(dim); + })); + + switch (indexMode) { + case IndexMode::EVEN_MODE: + extractOffsets.back() = rewriter.getIndexAttr(0); + break; + case IndexMode::ODD_MODE: + extractOffsets.back() = rewriter.getIndexAttr(1); + break; + } + extractStrides.back() = rewriter.getIndexAttr(2); + extractSizes.back() = rewriter.getIndexAttr(srcType.getShape().back() / 2); + + Value deinterleaveSlice = rewriter.create( + loc, newTensor, extractOffsets, extractSizes, extractStrides); + + rewriter.replaceOp(op, deinterleaveSlice); + return success(); + } + return failure(); +} + +LogicalResult +InterleaveStatusOptimization(SmallVector materializeVec) { + OpBuilder builder(materializeVec[1]); + auto loc = materializeVec[1]->getLoc(); + + auto firstReinterpretCastOp = + llvm::dyn_cast( + materializeVec[0]) + .getDest() + .getDefiningOp(); + auto secondReinterpretCastOp = + llvm::dyn_cast( + materializeVec[1]) + .getDest() + .getDefiningOp(); + + assert(firstReinterpretCastOp && secondReinterpretCastOp); + + // Judge whether two `ReinterpretCastOp` shape satisfy interleave state + // a. both size are equal + if (!isEqualConstantIntOrValueArray( + firstReinterpretCastOp.getConstifiedMixedSizes(), + secondReinterpretCastOp.getConstifiedMixedSizes())) { + return failure(); + } + // b. both strides are equal + if (!isEqualConstantIntOrValueArray( + firstReinterpretCastOp.getConstifiedMixedStrides(), + secondReinterpretCastOp.getConstifiedMixedStrides())) { + return failure(); + } + // c. both offsets should satisfy tricky rule + auto firstOriginCastOffset = + firstReinterpretCastOp.getConstifiedMixedOffset(); + auto secondOriginCastOffset = + secondReinterpretCastOp.getConstifiedMixedOffset(); + std::pair indexModeRecord; + OpFoldResult newCastOffset; + if (llvm::isa(firstOriginCastOffset) && + llvm::isa(secondOriginCastOffset)) { + auto [firstCastOffset, firstIndexMode] = + recountReinterpretCastOffset(firstOriginCastOffset, builder); + auto [secondCastOffset, secondIndexMode] = + recountReinterpretCastOffset(secondOriginCastOffset, builder); + + if (!(static_cast(firstIndexMode) ^ static_cast(secondIndexMode))) + return failure(); + newCastOffset = builder.getIndexAttr(0); + indexModeRecord = {firstIndexMode, secondIndexMode}; + + } else if (llvm::isa(firstOriginCastOffset) && + llvm::isa(secondOriginCastOffset)) { + auto [firstCastOffset, firstIndexMode] = + recountReinterpretCastOffset(firstOriginCastOffset, builder); + auto [secondCastOffset, secondIndexMode] = + recountReinterpretCastOffset(secondOriginCastOffset, builder); + + if (!(static_cast(firstIndexMode) ^ + static_cast(secondIndexMode)) || + (llvm::dyn_cast(firstCastOffset) != + llvm::dyn_cast(secondCastOffset))) + return failure(); + + if (firstIndexMode == IndexMode::EVEN_MODE) { + newCastOffset = llvm::dyn_cast(firstCastOffset); + } + if (secondIndexMode == IndexMode::EVEN_MODE) { + newCastOffset = llvm::dyn_cast(secondCastOffset); + } + indexModeRecord = {firstIndexMode, secondIndexMode}; + + } else { + return failure(); + } + + // Create new op + // 1. Get new destination memref type + auto dstType = expandInterleaveMemRefType(firstReinterpretCastOp.getType()); + + // 2. New tensor::EmptyOp + auto emptyTensor = builder.create(loc, dstType.getShape(), + dstType.getElementType()); + + // 3. New insert_slice from materialization source into new empty tensor + SmallVector insertOffsets(dstType.getRank(), + builder.getIndexAttr(0)); + SmallVector insertStrides(dstType.getRank(), + builder.getIndexAttr(1)); + SmallVector insertSizes = llvm::to_vector( + llvm::map_range(dstType.getShape(), [&](int64_t dim) -> OpFoldResult { + return builder.getIndexAttr(dim); + })); + insertStrides.back() = builder.getIndexAttr(2); + insertSizes.back() = builder.getIndexAttr(dstType.getShape().back() / 2); + if (indexModeRecord.first == IndexMode::ODD_MODE) { + insertOffsets.back() = builder.getIndexAttr(1); + } else { + insertOffsets.back() = builder.getIndexAttr(0); + } + auto insertFirst = builder.create( + loc, + llvm::dyn_cast( + materializeVec[0]) + .getSource(), + emptyTensor.getResult(), insertOffsets, insertSizes, insertStrides); + + if (indexModeRecord.second == IndexMode::ODD_MODE) { + insertOffsets.back() = builder.getIndexAttr(1); + } else { + insertOffsets.back() = builder.getIndexAttr(0); + } + auto insertSecond = builder.create( + loc, + llvm::dyn_cast( + materializeVec[1]) + .getSource(), + insertFirst.getResult(), insertOffsets, insertSizes, insertStrides); + + // 4. Reinterpret_cast block arg + auto newCastSize = firstReinterpretCastOp.getConstifiedMixedSizes(); + auto newCastStride = firstReinterpretCastOp.getConstifiedMixedStrides(); + newCastSize.back() = builder.getIndexAttr(dstType.getShape().back()); + newCastStride.back() = builder.getIndexAttr(1); + auto newCastOp = builder.create( + loc, dstType, firstReinterpretCastOp.getViewSource(), newCastOffset, + newCastSize, newCastStride); + + // 5. Create new bufferization::MaterializeInDestinationOp + auto newStoreOp = builder.create( + loc, insertSecond.getResult(), newCastOp.getResult()); + // Setting writable is necessary as dst is memref type + newStoreOp.setWritable(true); + + // 6. Erase origin materialization + materializeVec[0]->erase(); + materializeVec[1]->erase(); + + return success(); +} + +LogicalResult +InterleaveStatusWithMaskOptimization(SmallVector materializeVec) { + OpBuilder builder(materializeVec[1]); + + auto firstSubviewOpOfReCast = + llvm::dyn_cast( + materializeVec[0]) + .getDest() + .getDefiningOp(); + auto firstSrcExtractSlice = + llvm::dyn_cast( + materializeVec[0]) + .getSource() + .getDefiningOp(); + auto firstReinterpretCastOp = firstSubviewOpOfReCast.getSource() + .getDefiningOp(); + + auto secondSubviewOpOfReCast = + llvm::dyn_cast( + materializeVec[1]) + .getDest() + .getDefiningOp(); + auto secondSrcExtractSlice = + llvm::dyn_cast( + materializeVec[1]) + .getSource() + .getDefiningOp(); + auto secondReinterpretCastOp = + secondSubviewOpOfReCast.getSource() + .getDefiningOp(); + + // 1. Both source shapes of subview and extract_slice are equal + if (firstSubviewOpOfReCast.getSourceType().getShape() != + firstSrcExtractSlice.getSourceType().getShape()) + return failure(); + if (secondSubviewOpOfReCast.getSourceType().getShape() != + secondSrcExtractSlice.getSourceType().getShape()) + return failure(); + if (firstSubviewOpOfReCast.getSourceType().getShape() != + secondSubviewOpOfReCast.getSourceType().getShape()) + return failure(); + + // 2. both mask state are equal + std::function cmpFunc = + mlir::isEqualConstantIntOrValue; + if (!mlir::detail::sameOffsetsSizesAndStrides(firstSubviewOpOfReCast, + firstSrcExtractSlice, cmpFunc)) + return failure(); + if (!mlir::detail::sameOffsetsSizesAndStrides(secondSubviewOpOfReCast, + secondSrcExtractSlice, cmpFunc)) + return failure(); + if (!mlir::detail::sameOffsetsSizesAndStrides( + firstSubviewOpOfReCast, secondSubviewOpOfReCast, cmpFunc)) + return failure(); + + // 3. Still judge whether two `ReinterpretCastOp` shape satisfy request + // a. both size are equal + if (!isEqualConstantIntOrValueArray( + firstReinterpretCastOp.getConstifiedMixedSizes(), + secondReinterpretCastOp.getConstifiedMixedSizes())) + return failure(); + // b. both strides are equal + if (!isEqualConstantIntOrValueArray( + firstReinterpretCastOp.getConstifiedMixedStrides(), + secondReinterpretCastOp.getConstifiedMixedStrides())) + return failure(); + // c. both offsets should satisfy tricky rule + auto firstOriginCastOffset = + firstReinterpretCastOp.getConstifiedMixedOffset(); + auto secondOriginCastOffset = + secondReinterpretCastOp.getConstifiedMixedOffset(); + std::pair indexModeRecord; + OpFoldResult newCastOffset; + if (llvm::isa(firstOriginCastOffset) && + llvm::isa(secondOriginCastOffset)) { + auto [firstCastOffset, firstIndexMode] = + recountReinterpretCastOffset(firstOriginCastOffset, builder); + auto [secondCastOffset, secondIndexMode] = + recountReinterpretCastOffset(secondOriginCastOffset, builder); + + if (!(static_cast(firstIndexMode) ^ static_cast(secondIndexMode))) + return failure(); + newCastOffset = builder.getIndexAttr(0); + indexModeRecord = {firstIndexMode, secondIndexMode}; + + } else if (llvm::isa(firstOriginCastOffset) && + llvm::isa(secondOriginCastOffset)) { + auto [firstCastOffset, firstIndexMode] = + recountReinterpretCastOffset(firstOriginCastOffset, builder); + auto [secondCastOffset, secondIndexMode] = + recountReinterpretCastOffset(secondOriginCastOffset, builder); + + if (!(static_cast(firstIndexMode) ^ + static_cast(secondIndexMode)) || + (llvm::dyn_cast(firstCastOffset) != + llvm::dyn_cast(secondCastOffset))) + return failure(); + + if (firstIndexMode == IndexMode::EVEN_MODE) { + newCastOffset = llvm::dyn_cast(firstCastOffset); + } + if (secondIndexMode == IndexMode::EVEN_MODE) { + newCastOffset = llvm::dyn_cast(secondCastOffset); + } + indexModeRecord = {firstIndexMode, secondIndexMode}; + + } else { + return failure(); + } + auto loc = materializeVec[1]->getLoc(); + + // Create new op + // 1. Get new destination memref type + auto dstType = expandInterleaveMemRefType(firstReinterpretCastOp.getType()); + + // 2. New tensor::EmptyOp + auto emptyTensor = builder.create(loc, dstType.getShape(), + dstType.getElementType()); + + // 3. New insert_slice from extract_slice source into new empty tensor + SmallVector insertOffsets(dstType.getRank(), + builder.getIndexAttr(0)); + SmallVector insertStrides(dstType.getRank(), + builder.getIndexAttr(1)); + SmallVector insertSizes = llvm::to_vector( + llvm::map_range(dstType.getShape(), [&](int64_t dim) -> OpFoldResult { + return builder.getIndexAttr(dim); + })); + insertStrides.back() = builder.getIndexAttr(2); + insertSizes.back() = builder.getIndexAttr(dstType.getShape().back() / 2); + if (indexModeRecord.first == IndexMode::ODD_MODE) { + insertOffsets.back() = builder.getIndexAttr(1); + } else { + insertOffsets.back() = builder.getIndexAttr(0); + } + auto insertFirst = builder.create( + loc, firstSrcExtractSlice.getSource(), emptyTensor.getResult(), + insertOffsets, insertSizes, insertStrides); + + if (indexModeRecord.second == IndexMode::ODD_MODE) { + insertOffsets.back() = builder.getIndexAttr(1); + } else { + insertOffsets.back() = builder.getIndexAttr(0); + } + auto insertSecond = builder.create( + loc, secondSrcExtractSlice.getSource(), insertFirst.getResult(), + insertOffsets, insertSizes, insertStrides); + + // 4. To enable store with mask, create new extract_slice + SmallVector extractOffsets = + firstSrcExtractSlice.getMixedOffsets(); + SmallVector extractStrides = + firstSrcExtractSlice.getMixedStrides(); + SmallVector extractSizes = firstSrcExtractSlice.getMixedSizes(); + assert(llvm::isa(extractSizes.back())); + extractSizes.back() = builder.getIndexAttr( + getConstantIntValue(extractSizes.back()).value() * 2); + auto newSrcExtractSlice = builder.create( + loc, insertSecond.getResult(), extractOffsets, extractSizes, + extractStrides); + + // 5. Reinterpret_cast block arg + auto newCastSize = firstReinterpretCastOp.getConstifiedMixedSizes(); + auto newCastStride = firstReinterpretCastOp.getConstifiedMixedStrides(); + newCastSize.back() = builder.getIndexAttr(dstType.getShape().back()); + newCastStride.back() = builder.getIndexAttr(1); + auto newCastOp = builder.create( + loc, dstType, firstReinterpretCastOp.getViewSource(), newCastOffset, + newCastSize, newCastStride); + + // 6. Create new memref::SubViewOp of above new reinterpret_cast + // Here could reuse shape info of new extract_slice + auto dstSubviewType = memref::SubViewOp::inferResultType( + dstType, extractOffsets, extractSizes, extractStrides); + auto newSubviewOpOfReCast = builder.create( + loc, llvm::cast(dstSubviewType), newCastOp, extractOffsets, + extractSizes, extractStrides); + + // 7. Create new bufferization::MaterializeInDestinationOp + auto newStoreOp = builder.create( + loc, newSrcExtractSlice.getResult(), newSubviewOpOfReCast.getResult()); + // Setting writable is necessary as dst is memref type + newStoreOp.setWritable(true); + + // 8. Erase origin operation + materializeVec[0]->erase(); + materializeVec[1]->erase(); + firstSubviewOpOfReCast->erase(); + firstSrcExtractSlice->erase(); + secondSubviewOpOfReCast->erase(); + secondSrcExtractSlice->erase(); + + return success(); +} + +} // namespace triton +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/lib/UtilsIncubated/Utils.cpp b/third_party/wafer/third_party/flir/lib/UtilsIncubated/Utils.cpp new file mode 100755 index 00000000..34d1e701 --- /dev/null +++ b/third_party/wafer/third_party/flir/lib/UtilsIncubated/Utils.cpp @@ -0,0 +1,1254 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2025. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN + * THE SOFTWARE. + */ + +#include "incubated/Conversion/UtilsIncubated/Utils.h" + +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Dialect/Utils/StaticValueUtils.h" +#include "mlir/IR/Attributes.h" +#include "mlir/IR/BuiltinAttributes.h" +#include "mlir/IR/BuiltinTypeInterfaces.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Diagnostics.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/IR/Operation.h" +#include "mlir/IR/Value.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "llvm/ADT/SmallVector.h" +#include "llvm/ADT/SmallVectorExtras.h" +#include "llvm/ADT/TypeSwitch.h" +#include "llvm/Support/Casting.h" +#include "llvm/Support/Debug.h" +#include "llvm/Support/ErrorHandling.h" +#include "llvm/Support/LogicalResult.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#define DEBUG_TYPE "TritonNPU-Utils" + +namespace mlir { + +static Value createConstIndexValueOp(const Location &loc, OpBuilder &b, + int64_t value) { + return b.create(loc, b.getIndexAttr(value)).getResult(); +} + +static std::optional getConstantOfAttr(const OpFoldResult &arg) { + if (isa(arg)) { + return getConstantIntValue(arg); + } + + return std::nullopt; +} + +namespace ConverterUtils { + +std::optional +getLastStrideOfReinterpretCastOp(memref::ReinterpretCastOp op) { + SmallVector mixedStrides = op.getMixedStrides(); + if (mixedStrides.empty()) { + op->emitError("ReinterpretCastOp has no strides"); + return std::nullopt; + } + + OpFoldResult lastStride = mixedStrides.back(); + + if (op.getStaticStrides().back() > 0) { + return op.getStaticStrides().back(); + } else if (isa(op.getStrides().back())) { + auto u = op.getStrides().back(); + while (auto blkArg = dyn_cast(u)) { + if (auto forOp = dyn_cast(blkArg.getOwner()->getParentOp())) { + auto prt = forOp->getOperand(3 + blkArg.getArgNumber() - 1); + u = prt; + } else { + u = nullptr; + break; + } + } + if (!u) + return std::nullopt; + lastStride = u; + } + + if (auto attr = lastStride.dyn_cast()) { + return getConstantOfAttr(lastStride); + } else if (auto value = lastStride.dyn_cast()) { + auto defOp = value.getDefiningOp(); + if (auto constIndexOp = dyn_cast(defOp)) { + int64_t constValue = constIndexOp.value(); + return constValue; + } else if (auto constIntOp = dyn_cast(defOp)) { + int64_t constValue = constIntOp.value(); + return constValue; + } + } + return std::nullopt; +} + +bool isaPermutedMemRefType(MemRefType memRefType) { +#if LLVM_VERSION_MAJOR < 21 + auto [ptrStrides, ptrOffsets] = getStridesAndOffset(memRefType); +#else // triton_v3.3.x + auto [ptrStrides, ptrOffsets] = memRefType.getStridesAndOffset(); +#endif + LLVM_DEBUG({ + llvm::dbgs() << "---------- [BEG] ptrStrides ----------\n"; + for (auto stride : ptrStrides) + llvm::dbgs() << stride << " "; + llvm::dbgs() << "\n"; + llvm::dbgs() << "---------- [END] ptrStrides ----------\n"; + }); + + switch (ptrStrides.size()) { + case 0: + return false; + case 1: + return false; + default: { + return ptrStrides[ptrStrides.size() - 1] != 1; + } + } +} + +Value getTransposedValue(Value source, const Location loc, + ConversionPatternRewriter &rewriter, + llvm::ArrayRef order) { + auto sourceType = cast(source.getType()); + auto sourceRank = sourceType.getRank(); + + SmallVector perm(order); + SmallVector originalShape(sourceType.getShape()); + SmallVector transposedShape(sourceRank); + for (size_t i = 0; i < sourceRank; i++) { + transposedShape[i] = originalShape[perm[i]]; + } + + Value transposeInit = rewriter.create( + loc, transposedShape, sourceType.getElementType()); + + Value transpose = + rewriter.create(loc, source, transposeInit, perm) + .getResults()[0]; + + return transpose; +} + +SmallVector getNParallelLoopsAttrs(unsigned n) { + return SmallVector(n, utils::IteratorType::parallel); +} + +Value getScalarValue(Value operand, Location loc, + ConversionPatternRewriter &rewriter) { + SmallVector ops; + auto reconstructScalarValue = [&](Value src) { + for (auto op = ops.rbegin(); op != ops.rend(); ++op) { + src = mlir::TypeSwitch(*op) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return rewriter.create(loc, resType, src); + }) + .Case([&](Operation *op) { + auto resType = op->getResults()[0].getType(); + if (auto shapedType = dyn_cast(resType)) { + resType = shapedType.getElementType(); + } + return rewriter.create(loc, resType, src); + }) + .Default([](Operation *op) { + llvm_unreachable("unsupported op in generating "); + return nullptr; + }); + } + return src; + }; + + while (true) { + if (!dyn_cast(operand.getType())) { + return reconstructScalarValue(operand); + } else if (auto op = operand.getDefiningOp()) { + if (auto attr = dyn_cast(op.getValue())) { + if (!attr.isSplat()) { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load " + "produced by unsupported instruction"; + return nullptr; + } + auto elemValue = attr.getSplatValue(); + auto constOp = arith::ConstantOp::materialize( + rewriter, elemValue, attr.getElementType(), op.getLoc()); + return reconstructScalarValue(constOp.getResult()); + } + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load produced " + "by unsupported instruction"; + return nullptr; + } else if (auto op = operand.getDefiningOp()) { + operand = op.getSrc(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else if (auto op = operand.getDefiningOp()) { + ops.push_back(op.getOperation()); + operand = op.getIn(); + } else { + InFlightDiagnostic diag = emitError(loc) + << "other value used in masked load produced " + "by unsupported instruction"; + return nullptr; + } + } + return nullptr; +} + +memref::SubViewOp makeSubViewOp(Value src, + const llvm::SmallVector &sizes, + const Location &loc, + ConversionPatternRewriter &rewriter) { + auto srcType = cast(src.getType()); + SmallVector offsets(srcType.getRank(), + rewriter.getIndexAttr(0)); + SmallVector strides(srcType.getRank(), + rewriter.getIndexAttr(1)); + auto dstType = + memref::SubViewOp::inferResultType(srcType, offsets, sizes, strides); + return rewriter.create(loc, dyn_cast(dstType), + src, offsets, sizes, strides); +} + +tensor::ExtractSliceOp +makeExtractSliceOp(Value src, const llvm::SmallVector &sizes, + const Location &loc, ConversionPatternRewriter &rewriter) { + auto srcType = cast(src.getType()); + SmallVector offsets(srcType.getRank(), + rewriter.getIndexAttr(0)); + SmallVector strides(srcType.getRank(), + rewriter.getIndexAttr(1)); + auto dstType = + tensor::ExtractSliceOp::inferResultType(srcType, offsets, sizes, strides); + return rewriter.create(loc, dstType, src, offsets, + sizes, strides); +} + +std::optional getFullShapeOp(Value val, + ConversionPatternRewriter &rewriter) { + assert(isa(val.getType())); + + if (isa(val)) { + auto blockArg = dyn_cast(val); + auto blockOp = blockArg.getOwner()->getParentOp(); + if (isa(blockOp)) { + auto forOp = dyn_cast(blockOp); + auto operand = forOp.getTiedLoopInit(blockArg)->get(); + return getFullShapeOp(operand, rewriter); + } else { + emitError(val.getLoc()) + << "getFullShapeOp() only support ReinterpretCastOp " + "and scf.for's block argument, but got : " + << val << "\n"; + } + return std::nullopt; + } + + if (!isa(val.getDefiningOp())) { + emitError(val.getLoc()) + << "getFullShapeOp() only support ReinterpretCastOp " + "and scf.for's block argument, but got : " + << val << "\n"; + return std::nullopt; + } + + auto reCastOp = val.getDefiningOp(); + if (reCastOp->hasAttr("tensor_ptr_full_shape")) + return reCastOp; + + return getFullShapeOp(reCastOp.getSource(), rewriter); +} + +SmallVector +getBoundarySizes(llvm::ArrayRef boundaryCheck, Value ptr, + const Location &loc, ConversionPatternRewriter &rewriter) { + if (isa(ptr.getType())) + ptr = rewriter.getRemappedValue(ptr); + + auto shapedType = dyn_cast_if_present(ptr.getType()); + assert(shapedType && shapedType.hasStaticShape()); + + auto fullShapeOp = getFullShapeOp(ptr, rewriter); + + assert(fullShapeOp.has_value()); + SmallVector boundarySize = + getAsIndexOpFoldResult(rewriter.getContext(), shapedType.getShape()); + + auto fullShapeReCast = + dyn_cast(fullShapeOp.value()); + OpFoldResult curPtrOffset; + if (auto curReCast = ptr.getDefiningOp()) { + curPtrOffset = curReCast.getConstifiedMixedOffset(); + } else if (isa(ptr) && + isa(ptr.getParentBlock()->getParentOp())) { + // Here's to process loop state where ptr is just from loop interator. + // Following assertion corresponds to conversion result from `rewriteFor` + auto blockArg = dyn_cast(ptr); + auto forOp = dyn_cast(ptr.getParentBlock()->getParentOp()); + auto initReCastOfLoop = forOp.getTiedLoopInit(blockArg) + ->get() + .getDefiningOp(); + assert(initReCastOfLoop && initReCastOfLoop.getOffsets().size() == 1); + Value initReCastOffset = initReCastOfLoop.getOffsets()[0]; + + for (OpOperand &use : initReCastOffset.getUses()) { + if (use.getOwner() == initReCastOfLoop) + continue; + else if (isa(use.getOwner())) + continue; + else if (use.getOwner() == forOp) + curPtrOffset = OpFoldResult(forOp.getTiedLoopRegionIterArg(&use)); + else + llvm_unreachable("Illegal interation offset after rewriteFor"); + } + } else { + llvm_unreachable("Unsupported state when check tensor_ptr boundary"); + } + + assert(curPtrOffset); + + OpFoldResult offsetShift = subOpFoldResult( + curPtrOffset, fullShapeReCast.getConstifiedMixedOffset(), loc, rewriter); + + for (int i = 0; i < shapedType.getRank(); ++i) { + if (llvm::find(boundaryCheck, i) != boundaryCheck.end()) { + auto fullShape = fullShapeReCast.getConstifiedMixedSizes()[i]; + + OpFoldResult curOffset = divOpFoldResult( + offsetShift, fullShapeReCast.getConstifiedMixedStrides()[i], loc, + rewriter); + OpFoldResult curLeftSize = + maxOpFoldResult(subOpFoldResult(fullShape, curOffset, loc, rewriter), + rewriter.getIndexAttr(0), loc, rewriter); + + boundarySize[i] = + minOpFoldResult(boundarySize[i], curLeftSize, loc, rewriter); + + offsetShift = remOpFoldResult( + offsetShift, fullShapeReCast.getConstifiedMixedStrides()[i], loc, + rewriter); + } + } + + return boundarySize; +} + +SmallVector getBroadcastDims(RankedTensorType src, + RankedTensorType dst) { + SmallVector broadcastDims; + auto srcShape = src.getShape(); + auto dstShape = dst.getShape(); + + for (size_t i = 0; i < srcShape.size(); ++i) { + if (dstShape[i] != srcShape[i]) { + assert(srcShape[i] == 1 && + "Size of source broadcast dimension must be 1"); + broadcastDims.push_back(i); + } + } + assert(!broadcastDims.empty() && "Cannot identify broadcast dimension"); + return broadcastDims; +} + +// Dimensions of collapesd tensor is all unbroadcast dims +SmallVector getUnbroadcastDims(RankedTensorType src, + RankedTensorType dst) { + SmallVector unbroadcastDims; + auto srcShape = src.getShape(); + auto dstShape = dst.getShape(); + + for (size_t i = 0; i < srcShape.size(); ++i) { + if (dstShape[i] == srcShape[i]) { + unbroadcastDims.emplace_back(srcShape[i]); + } + } + return unbroadcastDims; +} + +} // namespace ConverterUtils + +namespace triton { + +mlir::Operation * +findFirstMatchingOperandDef(mlir::Operation *rootOp, + const std::function &condFn) { + LLVM_DEBUG(llvm::dbgs() << "[findFirstMatchingOperandDef] Current op: " + << *rootOp << "\n"); + mlir::Value lhs = nullptr; + mlir::Value rhs = nullptr; + if (auto op = dyn_cast(rootOp)) { + lhs = op.getPtr(); + rhs = op.getOffset(); + } else if (auto op = dyn_cast(rootOp)) { + lhs = op.getLhs(); + rhs = op.getRhs(); + } else if (auto op = dyn_cast(rootOp)) { + lhs = op.getLhs(); + rhs = op.getRhs(); + } else if (auto op = dyn_cast(rootOp)) { + lhs = op.getLhs(); + rhs = op.getRhs(); + } else if (auto op = dyn_cast(rootOp)) { + lhs = op.getLhs(); + rhs = op.getRhs(); + } else if (auto op = dyn_cast(rootOp)) { + lhs = op.getLhs(); + rhs = op.getRhs(); + } else if (auto op = dyn_cast(rootOp)) { + lhs = op.getSrc(); + } else if (auto op = dyn_cast(rootOp)) { + } else { + rootOp->emitRemark("Backtracing encounters unsupported Operation"); + return nullptr; + } + // Backtrace operands + if (!lhs) { + return nullptr; + } + auto lhsDef = lhs.getDefiningOp(); + mlir::Operation *targetOp; + if (lhsDef) { + if (condFn(lhsDef)) { + targetOp = lhsDef; + } else { + targetOp = findFirstMatchingOperandDef(lhsDef, condFn); + } + if (targetOp) { + return targetOp; + } + } + if (!rhs) { + return nullptr; + } + auto rhsDef = rhs.getDefiningOp(); + if (rhsDef) { + if (condFn(rhsDef)) { + targetOp = rhsDef; + } else { + targetOp = findFirstMatchingOperandDef(rhsDef, condFn); + } + if (targetOp) { + return targetOp; + } + } + return nullptr; +} + +void traverseBackwardUpdateOperandChainIf( + Operation *op, std::function conditionFn, + std::function stopFn, + std::function actionFn, OpBuilder &builder, + DenseSet &handledOperation) { + + if (!op || handledOperation.contains(op)) + return; + + handledOperation.insert(op); + + if (stopFn(op)) + return; + + if (conditionFn(op)) + actionFn(builder, op); + + DenseSet handledOperand; + + std::function handler = [&](Value operand) { + if (handledOperand.contains(operand)) + return; + handledOperand.insert(operand); + if (Operation *defOp = operand.getDefiningOp()) { + traverseBackwardUpdateOperandChainIf(defOp, conditionFn, stopFn, actionFn, + builder, handledOperation); + } else { + auto blockArgument = cast(operand); + auto parentOp = blockArgument.getOwner()->getParentOp(); + if (auto whileOp = dyn_cast(parentOp); + whileOp && whileOp.getAfterBody() == blockArgument.getOwner()) { + auto argNum = blockArgument.getArgNumber(); + auto conditionArg = whileOp.getConditionOp().getArgs()[argNum]; + handler(conditionArg); + } else if (auto loopOp = dyn_cast(parentOp)) { + OpOperand *initArgOperand = loopOp.getTiedLoopInit(blockArgument); + if (!initArgOperand) + return; + Value initArg = initArgOperand->get(); + handler(initArg); + Value yieldedValue = + loopOp.getTiedLoopYieldedValue(blockArgument)->get(); + if (yieldedValue != blockArgument) + handler(yieldedValue); + } + } + }; + + for (Value operand : op->getOperands()) { + handler(operand); + } + + if (auto loopOp = dyn_cast(op)) { + for (auto yieldedValue : loopOp.getYieldedValues()) + handler(yieldedValue); + } +} + +// Note: rootOp will also be processed. +void traverseBackwardUpdateOperandChainIf( + Operation *rootOp, std::function conditionFn, + std::function stopFn, + std::function actionFn) { + + OpBuilder builder(rootOp->getContext()); + DenseSet handledOperation; + + traverseBackwardUpdateOperandChainIf(rootOp, conditionFn, stopFn, actionFn, + builder, handledOperation); +} + +void traverseForwardUpdateUserChainIf( + Operation *op, std::function conditionFn, + std::function stopFn, + std::function actionFn, OpBuilder &builder, + llvm::SmallPtrSet &stopOps) { + + if (!op) { + return; + } + + if (stopFn(op)) { + stopOps.insert(op); + return; + } + + if (conditionFn(op)) { + actionFn(builder, op); + } + + for (auto res : op->getResults()) { + for (auto userOp : res.getUsers()) { + traverseForwardUpdateUserChainIf(userOp, conditionFn, stopFn, actionFn, + builder, stopOps); + } + } +} + +// Note: rootOp will also be processed. +void traverseForwardUpdateUserChainIf( + Operation *rootOp, std::function conditionFn, + std::function stopFn, + std::function actionFn, + llvm::SmallPtrSet &stopOps) { + + OpBuilder builder(rootOp->getContext()); + + traverseForwardUpdateUserChainIf(rootOp, conditionFn, stopFn, actionFn, + builder, stopOps); +} + +bool isMetaUse(Operation *op) { return op->hasAttr("MetaUse"); } + +bool isMixUse(Operation *op) { return op->hasAttr("MixUse"); } + +IndirectLoadInterfaceOpType getIndirectLoadInterfaceOpType(Operation *op) { + auto ty = IndirectLoadInterfaceOpType::Undefined; + if (isMetaUse(op)) { + if (isa(op)) { + ty = IndirectLoadInterfaceOpType::Load; + } else if (isa(op)) { + ty = IndirectLoadInterfaceOpType::Calc; + } + } + return ty; +} + +bool opIsIndirectLoad(Operation *op) { + auto opType = getIndirectLoadInterfaceOpType(op); + return opType == IndirectLoadInterfaceOpType::Load; +} + +bool opIsIndirectCalc(Operation *op) { + auto opType = getIndirectLoadInterfaceOpType(op); + return opType == IndirectLoadInterfaceOpType::Calc; +} + +scf::ForOp createNestedLoops( + OpBuilder &builder, Location loc, unsigned currentDim, unsigned totalDims, + ValueRange LBs, ValueRange UBs, ValueRange steps, SmallVector &ivs, + ValueRange initArgs, + function_ref &, ValueRange)> + bodyBuilder) { + + if (currentDim >= totalDims) { + bodyBuilder(builder, loc, ivs, initArgs); + return nullptr; + } + + auto loop = builder.create( + loc, LBs[currentDim], UBs[currentDim], steps[currentDim], initArgs, + [&](OpBuilder &nestedBuilder, Location nestedLoc, Value iv, + ValueRange iterArgs) { + ivs.push_back(iv); + auto innerLoop = createNestedLoops(nestedBuilder, nestedLoc, + currentDim + 1, totalDims, LBs, UBs, + steps, ivs, iterArgs, bodyBuilder); + if (innerLoop) { + nestedBuilder.create(loc, innerLoop.getResults()); + } + }); + + return loop; +} + +ModuleOp getModuleOpFromOperation(Operation *op) { + Operation *parent = op; + while (parent != nullptr && !isa(parent)) { + parent = parent->getParentOp(); // 向上查找 + } + return cast(parent); // 如果没找到会抛出异常 +} + +} // namespace triton + +// TODO: imply these function below +OpFoldResult addOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b) { + auto lhsInt = getConstantOfAttr(lhs); + auto rhsInt = getConstantOfAttr(rhs); + + if (lhsInt && rhsInt) + return b.getIndexAttr(lhsInt.value() + rhsInt.value()); + + if (!lhsInt && rhsInt && rhsInt.value() == 0) + return lhs; + if (!rhsInt && lhsInt && lhsInt.value() == 0) + return rhs; + + auto lhsValue = dyn_cast(lhs); + if (lhsInt) { + lhsValue = createConstIndexValueOp(loc, b, lhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + auto rhsValue = dyn_cast(rhs); + if (rhsInt) { + rhsValue = createConstIndexValueOp(loc, b, rhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + return b.create(loc, lhsValue, rhsValue).getResult(); +} + +OpFoldResult subOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b) { + auto lhsInt = getConstantOfAttr(lhs); + auto rhsInt = getConstantOfAttr(rhs); + + if (lhsInt && rhsInt) + return b.getIndexAttr(lhsInt.value() - rhsInt.value()); + + if (!lhsInt && rhsInt && rhsInt.value() == 0) + return lhs; + + auto lhsValue = dyn_cast(lhs), rhsValue = dyn_cast(rhs); + if (lhsInt) { + lhsValue = createConstIndexValueOp(loc, b, lhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + if (rhsInt) { + rhsValue = createConstIndexValueOp(loc, b, rhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + return b.create(loc, lhsValue, rhsValue).getResult(); +} + +OpFoldResult mulOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b) { + auto lhsInt = getConstantOfAttr(lhs); + auto rhsInt = getConstantOfAttr(rhs); + + if (lhsInt && rhsInt) + return b.getIndexAttr(lhsInt.value() * rhsInt.value()); + + if (lhsInt) { + if (lhsInt.value() == 0) + return lhs; + if (lhsInt.value() == 1) + return rhs; + } + if (rhsInt) { + if (rhsInt.value() == 0) + return rhs; + if (rhsInt.value() == 1) + return lhs; + } + + auto lhsValue = dyn_cast(lhs), rhsValue = dyn_cast(rhs); + if (lhsInt) { + lhsValue = createConstIndexValueOp(loc, b, lhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + if (rhsInt) { + rhsValue = createConstIndexValueOp(loc, b, rhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + return b.create(loc, lhsValue, rhsValue).getResult(); +} + +OpFoldResult divOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b) { + auto lhsInt = getConstantOfAttr(lhs); + auto rhsInt = getConstantOfAttr(rhs); + + if (rhsInt && rhsInt.value() == 0) { + emitError(loc) << "cannot div 0!"; + return OpFoldResult(); + } + + if (lhsInt && rhsInt) + return b.getIndexAttr(lhsInt.value() / rhsInt.value()); + + if (lhsInt) { + if (lhsInt.value() == 0) + return lhs; + } + + if (rhsInt) { + if (rhsInt.value() == 1) + return lhs; + } + + auto lhsValue = dyn_cast(lhs), rhsValue = dyn_cast(rhs); + if (lhsInt) { + lhsValue = createConstIndexValueOp(loc, b, lhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + if (rhsInt) { + rhsValue = createConstIndexValueOp(loc, b, rhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + return b.create(loc, lhsValue, rhsValue).getResult(); +} + +OpFoldResult remOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b) { + auto lhsInt = getConstantOfAttr(lhs); + auto rhsInt = getConstantOfAttr(rhs); + + if (rhsInt && rhsInt.value() == 0) { + emitError(loc) << "cannot remainder by 0!"; + return OpFoldResult(); + } + + if (lhsInt && rhsInt) + return b.getIndexAttr(lhsInt.value() % rhsInt.value()); + + if (lhsInt) { + if (lhsInt.value() == 0) + return lhs; + } + + auto lhsValue = dyn_cast(lhs), rhsValue = dyn_cast(rhs); + if (lhsInt) { + lhsValue = createConstIndexValueOp(loc, b, lhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + if (rhsInt) { + rhsValue = createConstIndexValueOp(loc, b, rhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + return b.create(loc, lhsValue, rhsValue).getResult(); +} + +OpFoldResult minOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b) { + auto lhsInt = getConstantOfAttr(lhs); + auto rhsInt = getConstantOfAttr(rhs); + if (lhsInt && rhsInt) + return b.getIndexAttr(std::min(lhsInt.value(), rhsInt.value())); + + auto lhsValue = dyn_cast(lhs), rhsValue = dyn_cast(rhs); + if (lhsInt) { + lhsValue = createConstIndexValueOp(loc, b, lhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + if (rhsInt) { + rhsValue = createConstIndexValueOp(loc, b, rhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + return b.create(loc, lhsValue, rhsValue).getResult(); +} + +OpFoldResult maxOpFoldResult(const OpFoldResult &lhs, const OpFoldResult &rhs, + const Location &loc, OpBuilder &b) { + auto lhsInt = getConstantOfAttr(lhs); + auto rhsInt = getConstantOfAttr(rhs); + if (lhsInt && rhsInt) + return b.getIndexAttr(std::max(lhsInt.value(), rhsInt.value())); + + auto lhsValue = dyn_cast(lhs), rhsValue = dyn_cast(rhs); + if (lhsInt) { + lhsValue = createConstIndexValueOp(loc, b, lhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + if (rhsInt) { + rhsValue = createConstIndexValueOp(loc, b, rhsInt.value()); + } else { + lhsValue = convertToIndexIfNeeded(lhsValue, loc, b); + assert(isa(lhsValue.getType())); + } + + return b.create(loc, lhsValue, rhsValue).getResult(); +} + +void addReduceWithIndexAttr(ReduceWithIndexParams params, + ConversionPatternRewriter &rewriter, + linalg::ReduceOp reduceOp) { + const StringRef reduceRef = "reduce_mode"; + const StringRef tieBreakLeftRef = "tie_break_left"; + const StringRef unsignedSrcRef = "unsigned_src"; + + const StringRef tieBreakStr = + params.tieBreakType == TieBreakType::LEFT ? "true" : "false"; + const StringRef withIndexStr = + params.withIndexType == ReduceWithIndexType::MAX ? "max_with_index" + : "min_with_index"; + const StringRef unsignedSrcStr = params.isUnsignedSrc ? "true" : "false"; + + reduceOp->setAttr(reduceRef, rewriter.getStringAttr(withIndexStr)); + reduceOp->setAttr(tieBreakLeftRef, rewriter.getStringAttr(tieBreakStr)); + reduceOp->setAttr(unsignedSrcRef, rewriter.getStringAttr(unsignedSrcStr)); +} + +std::optional +getReduceWithIndexParams(triton::ReduceOp reduceOp) { + auto tritonReduceBlock = reduceOp.getBody(); + auto *tritonYield = tritonReduceBlock->getTerminator(); + auto yieldVelues = tritonYield->getOperands(); + constexpr int yieldValuesNum = 2; + if (yieldVelues.size() != yieldValuesNum) { + return {}; + } + + // Unify signed/unsigned and int/float predicate + enum class Predicate { Undefined = 0, lt = 1, gt = 2, eq = 3 }; + enum class Signedness { NotApplicable = 0, Signed = 1, Unsigned = 2 }; + auto unifyPredicateI = + [](arith::CmpIPredicate p) -> std::pair { + switch (p) { + case arith::CmpIPredicate::slt: + return {Predicate::lt, Signedness::Signed}; + case arith::CmpIPredicate::ult: + return {Predicate::lt, Signedness::Unsigned}; + case arith::CmpIPredicate::sgt: + return {Predicate::gt, Signedness::Signed}; + case arith::CmpIPredicate::ugt: + return {Predicate::gt, Signedness::Unsigned}; + case arith::CmpIPredicate::eq: + return {Predicate::eq, Signedness::NotApplicable}; + default: + return {Predicate::Undefined, Signedness::NotApplicable}; + } + }; + auto unifyPredicateF = + [](arith::CmpFPredicate p) -> std::pair { + switch (p) { + case arith::CmpFPredicate::OLT: + return {Predicate::lt, Signedness::Signed}; + case arith::CmpFPredicate::ULT: + return {Predicate::lt, Signedness::Unsigned}; + case arith::CmpFPredicate::OGT: + return {Predicate::gt, Signedness::Signed}; + case arith::CmpFPredicate::UGT: + return {Predicate::gt, Signedness::Unsigned}; + case arith::CmpFPredicate::OEQ: + return {Predicate::eq, Signedness::Signed}; + case arith::CmpFPredicate::UEQ: + return {Predicate::eq, Signedness::Unsigned}; + default: + return {Predicate::Undefined, Signedness::NotApplicable}; + } + }; + + // Composite predicate to pick index of min (or max) element have to be + // written in following form: (v means value and i means index) + // For leftmost element: + // (new_v == old_v and new_i < old_i) or new_v < old_v + // new_v < old_v or (new_v == old_v and new_i < old_i) + // new_v < old_v // python3.11 ttir + // (new_v == old_v and new_i < old_i) or new_v > old_v + // new_v > old_v or (new_v == old_v and new_i < old_i) + // new_v > old_v // python3.11 ttir + // For rightmost element: + // (new_v == old_v and new_i > old_i) or new_v < old_v + // new_v < old_v or (new_v == old_v and new_i > old_i) + // (new_v == old_v and new_i > old_i) or new_v > old_v + // new_v > old_v or (new_v == old_v and new_i > old_i) + + std::map, std::pair> + m{ + // leftmost + {{Predicate::eq, Predicate::lt, Predicate::lt}, + {ReduceWithIndexType::MIN, TieBreakType::LEFT}}, + {{Predicate::lt, Predicate::eq, Predicate::lt}, + {ReduceWithIndexType::MIN, TieBreakType::LEFT}}, + {{Predicate::lt}, {ReduceWithIndexType::MIN, TieBreakType::LEFT}}, + {{Predicate::eq, Predicate::lt, Predicate::gt}, + {ReduceWithIndexType::MAX, TieBreakType::LEFT}}, + {{Predicate::gt, Predicate::eq, Predicate::lt}, + {ReduceWithIndexType::MAX, TieBreakType::LEFT}}, + {{Predicate::gt}, {ReduceWithIndexType::MAX, TieBreakType::LEFT}}, + // rightmost + {{Predicate::eq, Predicate::gt, Predicate::lt}, + {ReduceWithIndexType::MIN, TieBreakType::RIGHT}}, + {{Predicate::lt, Predicate::eq, Predicate::gt}, + {ReduceWithIndexType::MIN, TieBreakType::RIGHT}}, + {{Predicate::eq, Predicate::gt, Predicate::gt}, + {ReduceWithIndexType::MAX, TieBreakType::RIGHT}}, + {{Predicate::gt, Predicate::eq, Predicate::gt}, + {ReduceWithIndexType::MAX, TieBreakType::RIGHT}}, + }; + + std::vector preds; + std::vector signednesses; + // A better way is to trace the arith.select + // Checking the operations one by one is hacky :( + for (auto &op : tritonReduceBlock->without_terminator()) { + Predicate pred = Predicate::Undefined; + Signedness signedness = Signedness::NotApplicable; + if (auto cmpiOp = dyn_cast(op)) { + auto predi = cmpiOp.getPredicate(); + std::tie(pred, signedness) = unifyPredicateI(predi); + } + if (auto cmpfOp = dyn_cast(op)) { + auto predf = cmpfOp.getPredicate(); + std::tie(pred, signedness) = unifyPredicateF(predf); + } + if (pred != Predicate::Undefined) { + preds.push_back(pred); + signednesses.push_back(signedness); + } + } + // check if sequence of predicates matches any sequence for min/max + // leftmost/rightmost + if (m.find(preds) == m.end()) { + return {}; + } + + assert(!signednesses.empty()); + const bool isUnsignedSrc = + signednesses[0] == Signedness::Unsigned || + signednesses[signednesses.size() - 1] == Signedness::Unsigned; + return ReduceWithIndexParams{.withIndexType = m.at(preds).first, + .tieBreakType = m.at(preds).second, + .isUnsignedSrc = isUnsignedSrc}; +} + +// Fold layout constant info to attr, otherwise convert to index type value +OpFoldResult getOpFoldResultOfLayoutInfo(Value value, OpBuilder &builder) { + OpFoldResult constantFold = getAsOpFoldResult(value); + if (llvm::isa(constantFold)) { + assert(isa(constantFold.get())); + return constantFold; + } + + if (!isa(value.getType())) + llvm_unreachable("Illegal data type when parse block data layout info"); + + if (!isa(value.getType())) { + if (value.getType().isInteger(/*width*/ 1)) + value = builder.create( + value.getLoc(), builder.getIndexType(), value); + else + value = builder.create(value.getLoc(), + builder.getIndexType(), value); + } + + return value; +} + +// Specialize the Typeless Value (Zero, Min, Max) into a mlir TypedAttr +FailureOr specializeTypelessValueToAttr(TypelessValue value, + Type type, OpBuilder &b) { + mlir::Type f16Ty = Float16Type::get(b.getContext()); + mlir::Type f32Ty = Float32Type::get(b.getContext()); + mlir::Type i8TySL = IntegerType::get( + b.getContext(), 8, IntegerType::SignednessSemantics::Signless); + mlir::Type i8TyS = IntegerType::get(b.getContext(), 8, + IntegerType::SignednessSemantics::Signed); + mlir::Type i8TyU = IntegerType::get( + b.getContext(), 8, IntegerType::SignednessSemantics::Unsigned); + mlir::Type i16TySL = IntegerType::get( + b.getContext(), 16, IntegerType::SignednessSemantics::Signless); + mlir::Type i16TyS = IntegerType::get( + b.getContext(), 16, IntegerType::SignednessSemantics::Signed); + mlir::Type i16TyU = IntegerType::get( + b.getContext(), 16, IntegerType::SignednessSemantics::Unsigned); + mlir::Type i32TySL = IntegerType::get( + b.getContext(), 32, IntegerType::SignednessSemantics::Signless); + mlir::Type i32TyS = IntegerType::get( + b.getContext(), 32, IntegerType::SignednessSemantics::Signed); + mlir::Type i32TyU = IntegerType::get( + b.getContext(), 32, IntegerType::SignednessSemantics::Unsigned); + mlir::Type i64TySL = IntegerType::get( + b.getContext(), 64, IntegerType::SignednessSemantics::Signless); + mlir::Type i64TyS = IntegerType::get( + b.getContext(), 64, IntegerType::SignednessSemantics::Signed); + mlir::Type i64TyU = IntegerType::get( + b.getContext(), 64, IntegerType::SignednessSemantics::Unsigned); + llvm::APFloat halfZero = llvm::APFloat::getZero(llvm::APFloat::IEEEhalf()); + llvm::APFloat halfOne(llvm::APFloat::IEEEhalf(), 1); + llvm::APFloat halfMax = llvm::APFloat::getInf(llvm::APFloat::IEEEhalf()); + llvm::APFloat halfMin = + llvm::APFloat::getInf(llvm::APFloat::IEEEhalf(), true); + llvm::APFloat floatZero = llvm::APFloat::getZero(llvm::APFloat::IEEEsingle()); + llvm::APFloat floatOne(llvm::APFloat::IEEEsingle(), 1); + llvm::APFloat floatMax = llvm::APFloat::getInf(llvm::APFloat::IEEEsingle()); + llvm::APFloat floatMin = + llvm::APFloat::getInf(llvm::APFloat::IEEEsingle(), true); + auto toPtr = [](mlir::Type ty) { return ty.getAsOpaquePointer(); }; + + std::map, + std::variant> + initMap = { + {{TypelessValue::Zero, toPtr(f16Ty)}, halfZero}, + {{TypelessValue::Zero, toPtr(f32Ty)}, floatZero}, + {{TypelessValue::Zero, toPtr(i16TySL)}, (int16_t)0}, + {{TypelessValue::Zero, toPtr(i16TyS)}, (int16_t)0}, + {{TypelessValue::Zero, toPtr(i16TyU)}, (int16_t)0}, + {{TypelessValue::Zero, toPtr(i32TySL)}, 0}, + {{TypelessValue::Zero, toPtr(i32TyS)}, 0}, + {{TypelessValue::Zero, toPtr(i32TyU)}, 0}, + {{TypelessValue::Zero, toPtr(i64TySL)}, (int64_t)0}, + {{TypelessValue::Zero, toPtr(i64TyS)}, (int64_t)0}, + {{TypelessValue::Zero, toPtr(i64TyU)}, (int64_t)0}, + {{TypelessValue::Min, toPtr(f16Ty)}, halfMin}, + {{TypelessValue::Min, toPtr(f32Ty)}, floatMin}, + {{TypelessValue::Min, toPtr(i16TySL)}, + std::numeric_limits::min()}, + {{TypelessValue::Min, toPtr(i16TyS)}, + std::numeric_limits::min()}, + {{TypelessValue::Min, toPtr(i16TyU)}, + std::numeric_limits::min()}, + {{TypelessValue::Min, toPtr(i32TySL)}, + std::numeric_limits::min()}, + {{TypelessValue::Min, toPtr(i32TyS)}, + std::numeric_limits::min()}, + {{TypelessValue::Min, toPtr(i32TyU)}, + std::numeric_limits::min()}, + {{TypelessValue::Min, toPtr(i64TySL)}, + std::numeric_limits::min()}, + {{TypelessValue::Min, toPtr(i64TyS)}, + std::numeric_limits::min()}, + {{TypelessValue::Min, toPtr(i64TyU)}, + std::numeric_limits::min()}, + {{TypelessValue::Max, toPtr(f16Ty)}, halfMax}, + {{TypelessValue::Max, toPtr(f32Ty)}, floatMax}, + {{TypelessValue::Max, toPtr(i16TySL)}, + std::numeric_limits::max()}, + {{TypelessValue::Max, toPtr(i16TyS)}, + std::numeric_limits::max()}, + {{TypelessValue::Max, toPtr(i16TyU)}, + std::numeric_limits::max()}, + {{TypelessValue::Max, toPtr(i32TySL)}, + std::numeric_limits::max()}, + {{TypelessValue::Max, toPtr(i32TyS)}, + std::numeric_limits::max()}, + {{TypelessValue::Max, toPtr(i32TyU)}, + std::numeric_limits::max()}, + {{TypelessValue::Max, toPtr(i64TySL)}, + std::numeric_limits::max()}, + {{TypelessValue::Max, toPtr(i64TyS)}, + std::numeric_limits::max()}, + {{TypelessValue::Max, toPtr(i64TyU)}, + std::numeric_limits::max()}, + }; + + std::pair key = + std::make_pair(value, toPtr(type)); + if (initMap.find(key) == initMap.end()) + return failure(); + if (type.isInteger(8)) + return success(IntegerAttr::get(IntegerType::get(b.getContext(), 8), + std::get(initMap.at(key)))); + if (type.isInteger(16)) + return success(IntegerAttr::get(IntegerType::get(b.getContext(), 16), + std::get(initMap.at(key)))); + if (type.isInteger(32)) + return success(IntegerAttr::get(IntegerType::get(b.getContext(), 32), + std::get(initMap.at(key)))); + if (type.isInteger(64)) + return success(IntegerAttr::get(IntegerType::get(b.getContext(), 64), + std::get(initMap.at(key)))); + if (isa(type)) + return success( + FloatAttr::get(f16Ty, std::get(initMap.at(key)))); + if (isa(type)) + return success( + FloatAttr::get(f32Ty, std::get(initMap.at(key)))); + return failure(); +} + +// Specialize the Typeless Value (Zero, Min, Max) into a mlir constant value +FailureOr specializeTypelessValueToConstant(TypelessValue value, + Type type, Location loc, + OpBuilder &b) { + std::function getElemType = [&](mlir::Type ty) { + if (auto ptrType = dyn_cast(getElementTypeOrSelf(ty))) + return getElemType(ptrType.getPointeeType()); + if (auto tensorType = mlir::dyn_cast(ty)) + return getElemType(tensorType.getElementType()); + return ty; + }; + + if (value == TypelessValue::Undefined) + return failure(); + if (auto tensorType = mlir::dyn_cast(type)) { + auto elemType = getElemType(tensorType); + FailureOr typedAttr = + specializeTypelessValueToAttr(value, elemType, b); + if (failed(typedAttr)) + return failure(); + auto otherTensorType = + RankedTensorType::get(tensorType.getShape(), elemType); + auto denseAttr = DenseElementsAttr::get(otherTensorType, *typedAttr); + return b.create(loc, denseAttr).getResult(); + } + if (mlir::isa(type) || mlir::isa(type)) { + FailureOr typedAttr = + specializeTypelessValueToAttr(value, type, b); + if (failed(typedAttr)) + return failure(); + return b.create(loc, *typedAttr).getResult(); + } + return failure(); +} + +std::optional getIntAttr(const OpFoldResult ofr) { + Attribute attr; + if (auto val = dyn_cast(ofr)) { + if (!val.getDefiningOp()) + return std::nullopt; + attr = cast(val.getDefiningOp().getValue()); + } else { + attr = dyn_cast(ofr); + } + if (attr && isa(attr)) + return dyn_cast(attr).getInt(); + return std::nullopt; +} + +Value materializeValue(OpBuilder &builder, Location loc, OpFoldResult ofr) { + if (auto val = ofr.dyn_cast()) { + return val; + } + + auto intVal = getIntAttr(ofr); + if (intVal.has_value()) { + return builder.create( + loc, builder.getI32IntegerAttr(intVal.value())); + } + assert(intVal.has_value()); + return Value(); + + // return builder.create( + // loc, dyn_cast(attr).getInt()); +} + +bool isZero(const OpFoldResult ofr) { + auto staticOfr = getIntAttr(ofr); + return staticOfr.has_value() && staticOfr.value() == 0; +} + +Value convertToIndexIfNeeded(Value input, const Location &loc, OpBuilder &b) { + auto inputType = input.getType(); + if (auto intType = dyn_cast(inputType)) { + if (intType.isInteger(32) || intType.isInteger(64)) { + return b.create(loc, b.getIndexType(), input); + } + } + return input; +} + +} // namespace mlir diff --git a/third_party/wafer/third_party/flir/python/examples/bare_matmul.py b/third_party/wafer/third_party/flir/python/examples/bare_matmul.py new file mode 100755 index 00000000..eaa408b1 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/bare_matmul.py @@ -0,0 +1,45 @@ +# this is a benchmark which multiplies square matrices with maximum block size +# to check the performance of tl.dot operation + +import torch +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def bare_matmul(X, Y, Z, M, N, K, BLOCK_SIZE: tl.constexpr): + pid_x = tl.program_id(0) # block row id + pid_y = tl.program_id(1) # block column id + + offs_x = pid_x * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + offs_y = pid_y * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + + x = tl.load(X + offs_x[:, None] * K + offs_y[None, :]) + y = tl.load(Y + offs_x[:, None] * N + offs_y[None, :]) + + z = tl.dot(x, y) + + tl.store(Z + offs_x[:, None] * N + offs_y[None, :], z) + + +@benchmark.measure() +def bench_matmul(N, provider): + device = 'cpu' + dtype = torch.float32 + a = torch.randn((N, N), device=device, dtype=dtype) + b = torch.randn((N, N), device=device, dtype=dtype) + c = torch.empty((N, N), device=device, dtype=dtype) + if provider == 'torch' or provider == 'test': + c_ref = torch.matmul(a, b) + if provider == 'triton' or provider == 'test': + bare_matmul[(1,)](a, b, c, N, N, N, N) + if provider == 'test': + torch.testing.assert_close(c, c_ref, atol=1e-2, rtol=0) + + +if __name__ == "__main__": + benchmark.select_cpu_backend() + for X in [2**i for i in range(7, 10, 1)]: + for provider in ['test', 'torch', 'triton']: + bench_matmul(X, provider) diff --git a/third_party/wafer/third_party/flir/python/examples/benchmark.py b/third_party/wafer/third_party/flir/python/examples/benchmark.py new file mode 100755 index 00000000..f82d3441 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/benchmark.py @@ -0,0 +1,66 @@ +import time +import numpy as np +from functools import wraps +import triton +from triton.backends.triton_shared.driver import CPUDriver + + +def select_cpu_backend(): + triton.runtime.driver.set_active(CPUDriver()) + + +# Unfortunately, we can't use triton.testing.perf_report and triton.testing.do_bench for CPU backend because +# they are very specific to cuda + +def measure(repeats=20, percentiles=(), timers={'Wall':time.perf_counter, 'CPU':time.process_time}): + """ + Decorator to benchmark a function. + + Parameters: + - repeats (int): The number of times the function should be executed for each set of parameters. + - percentiles (tuple): The percentiles to compute on the execution times (e.g., (50, 90, 99)). + - timers (dict): A dictionary where keys are timer names (e.g., 'Wall', 'CPU') and values are timer functions + that measure elapsed time. By default: + * 'Wall': Uses time.perf_counter for high-resolution wall-clock time. + * 'CPU': Uses time.process_time for CPU time spent by the process. + + Returns: + - A decorated function that prints: + * Average execution time. + * Standard deviation time. + * Minimum and maximum times. + * Computed percentiles for each timer. + """ + def decorator(func): + @wraps(func) + def wrapper(*args, **kwargs): + print(f"{func.__name__}{args} {kwargs}, {repeats} times, all results in seconds") + times = {} + for t, _ in timers.items(): + times[t] = [] + + for _ in range(repeats): + starts = {} + for t, f in timers.items(): + starts[t] = f() + + result = func(*args, **kwargs) + + for t, f in timers.items(): + times[t].append(f() - starts[t]) + + for t, _ in timers.items(): + average_time = np.mean(times[t]) + min_time = np.min(times[t]) + max_time = np.max(times[t]) + computed_percentiles = np.percentile(times[t], percentiles) + std_dev_time = np.std(times[t]) + + print(f"{t}: Avg={average_time:.6f}, min={min_time:.6f}, std={std_dev_time:.6f},", end=" ") + for p, value in zip(percentiles, computed_percentiles): + print(f"{p}pp={value:.6f},", end=" ") + print(f"max={max_time:.6f}") + + return result + return wrapper + return decorator \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/python/examples/conftest.py b/third_party/wafer/third_party/flir/python/examples/conftest.py new file mode 100755 index 00000000..dc8c6bee --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/conftest.py @@ -0,0 +1,82 @@ +import pytest +import os +import tempfile +import triton +from triton.backends.triton_shared.driver import CPUDriver + +triton.runtime.driver.set_active(CPUDriver()) + + +def empty_decorator(func): + return func + +pytest.mark.interpreter = empty_decorator + + +@pytest.fixture +def device(request): + return "cpu" + + +tests_supported = { + "test_store_eviction_policy", + "test_unary_op", + "test_umulhi", + "test_for_iv", + "test_trans_2d", + "test_math_op", + "test_math_fma_op", + "test_abs", + "test_call", + "test_vectorization", + "test_convert_float16_to_float32", + "test_index1d", + "test_shift_op", + "test_full", + "test_floordiv", + "test_empty_kernel", + "test_if_return", + "test_value_specialization", + "test_clamp", + "test_store_cache_modifier", + "test_permute", + "test_broadcast", + "test_precise_math", + "test_vectorization_hints", + "test_dot", + "test_value_specialization_overflow", + "test_bitwise_op", + "test_const", + "test_unary_math", + "test_dot_mulbroadcasted", + "test_masked_load_scalar", + "test_enable_fp_fusion", + "test_load_cache_modifier", + "test_dot_without_load", + "test_cat", + "test_addptr" +} + +def pytest_collection_modifyitems(config, items): + skip_marker = pytest.mark.skip(reason="CPU backend does not support it yet") + # There is a dependency issue on build machine which breaks bfloat16 + skip_marker_bfloat = pytest.mark.skip(reason="bfloat16 linking issue") + skip_marker_tf32 = pytest.mark.skip(reason="tf32 is not supported on CPU") + skip_marker_float8 = pytest.mark.skip(reason="float8 is not supported on CPU") + + for item in items: + test_func_name = item.originalname if item.originalname else item.name + + test_file = str(item.fspath) + if test_file.endswith("test_core.py") and test_func_name not in tests_supported: + item.add_marker(skip_marker) + continue + + if "parametrize" in item.keywords: + for param_name, param_value in item.callspec.params.items(): + if (param_name.startswith('dtype') or param_name.endswith('dtype')) and param_value == 'bfloat16': + item.add_marker(skip_marker_bfloat) + if param_name.startswith('input_precision') and param_value.startswith('tf32'): + item.add_marker(skip_marker_tf32) + if (param_name.startswith('dtype') or param_name.endswith('dtype')) and ('float8' in str(param_value)): + item.add_marker(skip_marker_float8) diff --git a/third_party/wafer/third_party/flir/python/examples/test_addptr.py b/third_party/wafer/third_party/flir/python/examples/test_addptr.py new file mode 100755 index 00000000..1804496f --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_addptr.py @@ -0,0 +1,43 @@ +import torch + +import triton +import triton.language as tl + + +@triton.jit +def addptr(in0, out0): + for i in range(0, 10, 2): + in1 = in0 + 1 + i + in2 = in1 + 1 + + out1 = out0 + 1 + i + out2 = out1 + 1 + + a1 = tl.load(in1) + a2 = tl.load(in2) + + tl.store(out1, a1) + tl.store(out2, a2) + + + +def test(device): + input = torch.arange(0, 11, device=device, dtype=torch.float32) + output = torch.full((11,), 0, device=device, dtype=torch.float32) + grid = lambda meta: (1,) + + print(output) + addptr[grid](input, output) + print(input) + print(output) + assert torch.equal(input, output) + + # TODO: need to check some conditions otherwise the code below does not make any difference for the test + src = triton.compiler.ASTSource( + fn=addptr, + signature="*fp32,*fp32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) diff --git a/third_party/wafer/third_party/flir/python/examples/test_blockptr_complex_offset.py b/third_party/wafer/third_party/flir/python/examples/test_blockptr_complex_offset.py new file mode 100755 index 00000000..94c0cd5f --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_blockptr_complex_offset.py @@ -0,0 +1,37 @@ +import torch + +import triton +import triton.language as tl + + +@triton.jit +def block_copy_kernel(a_ptr, b_ptr): + a_block_ptr = tl.make_block_ptr( + base=a_ptr + 8, + shape=(2, 2), + strides=(2, 1), + offsets=(0, 0), + block_shape=(2, 2), + order=(1, 0), + ) + b_block_ptr = tl.make_block_ptr( + base=b_ptr, + shape=(2, 2), + strides=(2, 1), + offsets=(0, 0), + block_shape=(2, 2), + order=(1, 0), + ) + a = tl.load(a_block_ptr, boundary_check=(0,)) + tl.store(b_block_ptr, a, boundary_check=(0,)) + + + +def test(device): + input = torch.arange(0, 16, device=device, dtype=torch.float32) + output = torch.full((4,), -1, device=device, dtype=torch.float32) + expected = torch.arange(8, 12, device=device) + grid = lambda meta: (1,) + + block_copy_kernel[grid](input, output) + torch.equal(expected, output) diff --git a/third_party/wafer/third_party/flir/python/examples/test_early_return.py b/third_party/wafer/third_party/flir/python/examples/test_early_return.py new file mode 100755 index 00000000..2bb59be0 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_early_return.py @@ -0,0 +1,55 @@ +import torch + +import triton +import triton.language as tl + +from triton.backends.triton_shared.driver import CPUDriver + +@triton.jit +def early_return(in0, out0): + pid = tl.program_id(0) + id = tl.load(in0 + pid) + if id == -1: + return + offs = 1 + tl.arange(0, 4) + out_offs = tl.arange(0, 4) + tl.store(out0 + out_offs, offs) + +def compile(device): + src = triton.compiler.ASTSource( + fn=early_return, + signature="*fp32,*fp32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) + + +def test_return_case(device): + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + SIZE = 8 + input = torch.full((SIZE, ), -1, device=device, dtype=torch.int32) + output = torch.full((SIZE,), -1, device=device, dtype=torch.int32) + grid = lambda meta: (1,) + print(output) + early_return[grid](input, output) + print(input) + print(output) + torch.testing.assert_close(torch.tensor([ -1, -1, -1, -1, -1, -1, -1, -1], dtype=torch.int32), output) + +def test_normal_case(device): + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + SIZE = 8 + input = torch.arange(0, SIZE, device=device, dtype=torch.int32) + output = torch.full((SIZE,), -1, device=device, dtype=torch.int32) + grid = lambda meta: (1,) + print(output) + early_return[grid](input, output) + print(input) + print(output) + torch.testing.assert_close(torch.tensor([ 1, 2, 3, 4, -1, -1, -1, -1], dtype=torch.int32), output) diff --git a/third_party/wafer/third_party/flir/python/examples/test_gather_scatter.py b/third_party/wafer/third_party/flir/python/examples/test_gather_scatter.py new file mode 100755 index 00000000..db08736f --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_gather_scatter.py @@ -0,0 +1,163 @@ +import torch + +import triton +import triton.language as tl + +from triton.backends.triton_shared.driver import CPUDriver + +@triton.jit +def gather_simple_no_mask(in0, out0): + offs = tl.arange(0, 64) + out_offs = tl.arange(0, 64) + for i in range(0, 2): + offs = offs // 10 + (i + 1 * 5) % 64 + a = tl.load(in0 + offs) + tl.store(out0 + out_offs, a) + offs += 64 + out_offs += 64 + + +@triton.jit +def gather_simple_mask_no_other(in0, out0): + offs = tl.arange(0, 64) + out_offs = tl.arange(0, 64) + mask_bound = 8 + for i in range(0, 2): + gather_offs = offs // 4 + a = tl.load(in0 + gather_offs, mask=gather_offs < mask_bound) + tl.store(out0 + out_offs, a) + mask_bound += 16 + offs += 64 + out_offs += 64 + + +@triton.jit +def gather_simple_mask_with_other(in0, out0): + offs = tl.arange(0, 64) + out_offs = tl.arange(0, 64) + mask_bound = 8 + for i in range(0, 2): + gather_offs = offs // 4 + a = tl.load(in0 + gather_offs, mask=gather_offs < mask_bound, other=-1) + tl.store(out0 + out_offs, a) + mask_bound += 16 + offs += 64 + out_offs += 64 + + +@triton.jit +def masked_gather_scatter(in0, out0): + offs = tl.arange(0, 64) + out_offs = tl.arange(0, 64) + for i in range(0, 2): + offs = offs // 12 + mask = offs < 64 + store_mask = out_offs < 77 + a = tl.load(in0 + offs, mask=mask, other=99) + tl.store(out0 + out_offs, a, mask=store_mask) + offs += 64 + out_offs += 64 + + +@triton.jit +def complex_gather_scatter(in0, out0): + offs = tl.arange(0, 32) + out_offs = tl.arange(0, 32) + for i in range(0, 2): + offs = offs // 3 + i + mask = offs < 32 + a = tl.load(in0 + offs, mask=mask) + tl.store(out0 + offs, a) + offs += 32 + out_offs += 32 + + for j in range(0, 2): + offs = offs // ((i + 1) * (j + 1)) + i + mask = offs < 32 + a = tl.load(in0 + offs, mask=mask) + tl.store(out0 + offs, a) + offs += 32 + out_offs += 32 + + offs = offs // 3 + i + mask = offs < 32 + a = tl.load(in0 + offs, mask=mask) + tl.store(out0 + offs, a) + offs += 32 + out_offs += 32 + + +def run_test(triton_kernel, device, expected_output): + SIZE = 128 + input = torch.arange(2, SIZE + 2, device=device, dtype=torch.int32) + output = torch.full((SIZE,), -1, device=device, dtype=torch.int32) + + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + grid = lambda meta: (1,) + + print(output) + triton_kernel[grid](input, output) + print(input) + print(output) + torch.testing.assert_close(output, expected_output) + + +def test_gather_simple_no_mask(device): + expected_output = torch.tensor([ 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 8, 8, 8, 8, 8, 8, 8, 8, + 8, 8, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 10, 10, 10, 10, 10, 10, + 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 12, 12, 12, 12, + 12, 12, 12, 12, 12, 12, 13, 13, 13, 13, 14, 14, 14, 14, 14, 14, 14, 14, + 14, 14, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, + 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, + 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, + 15, 15], device=device, dtype=torch.int32) + run_test(gather_simple_no_mask, device, expected_output) + +def test_gather_simple_mask_no_other(device): + expected_output = torch.tensor([ 2, 2, 2, 2, 3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 6, 6, + 6, 6, 7, 7, 7, 7, 8, 8, 8, 8, 9, 9, 9, 9, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 18, 18, 18, 18, 19, 19, 19, 19, + 20, 20, 20, 20, 21, 21, 21, 21, 22, 22, 22, 22, 23, 23, 23, 23, 24, 24, + 24, 24, 25, 25, 25, 25, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0], device=device, dtype=torch.int32) + run_test(gather_simple_mask_no_other, device, expected_output) + + +def test_gather_simple_mask_with_other(device): + expected_output = torch.tensor([ 2, 2, 2, 2, 3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 6, 6, + 6, 6, 7, 7, 7, 7, 8, 8, 8, 8, 9, 9, 9, 9, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 18, 18, 18, 18, 19, 19, 19, 19, + 20, 20, 20, 20, 21, 21, 21, 21, 22, 22, 22, 22, 23, 23, 23, 23, 24, 24, + 24, 24, 25, 25, 25, 25, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1], device=device, dtype=torch.int32) + run_test(gather_simple_mask_with_other, device, expected_output) + + +def test_masked_gather_scatter(device): + expected_output = torch.tensor([ 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, + 3, 3, 3, 3, 3, 3, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, + 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 6, 6, 6, 6, 6, 6, + 6, 6, 6, 6, 6, 6, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, + 7, 7, 7, 7, 7, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1], device=device, dtype=torch.int32) + run_test(masked_gather_scatter, device, expected_output) + + +def test_complex_gather_scatter(device): + expected_output = torch.tensor([ 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, -1, -1, -1, -1, 17, 18, -1, + 20, 21, -1, 23, 24, 25, -1, -1, 28, -1, -1, -1, -1, -1, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1], device=device, dtype=torch.int32) + run_test(complex_gather_scatter, device, expected_output) diff --git a/third_party/wafer/third_party/flir/python/examples/test_layernorm.py b/third_party/wafer/third_party/flir/python/examples/test_layernorm.py new file mode 100755 index 00000000..dc9b8471 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_layernorm.py @@ -0,0 +1,160 @@ +# This is the Layer Norm forward pass from the Triton tutorial found here: +# https://github.com/triton-lang/triton/blob/main/python/tutorials/05-layer-norm.py + +# %% +# Motivations +# ----------- +# +# The *LayerNorm* operator was first introduced in [BA2016]_ as a way to improve the performance +# of sequential models (e.g., Transformers) or neural networks with small batch size. +# It takes a vector :math:`x` as input and produces a vector :math:`y` of the same shape as output. +# The normalization is performed by subtracting the mean and dividing by the standard deviation of :math:`x`. +# After the normalization, a learnable linear transformation with weights :math:`w` and biases :math:`b` is applied. +# The forward pass can be expressed as follows: +# +# .. math:: +# y = \frac{ x - \text{E}[x] }{ \sqrt{\text{Var}(x) + \epsilon} } * w + b +# +# where :math:`\epsilon` is a small constant added to the denominator for numerical stability. +# Let’s first take a look at the forward pass implementation. + +import torch + +import triton +import triton.language as tl +import pytest +import benchmark + + +@triton.jit +def _layer_norm_fwd_fused( + X, # pointer to the input + Y, # pointer to the output + W, # pointer to the weights + B, # pointer to the biases + Mean, # pointer to the mean + Rstd, # pointer to the 1/std + stride, # how much to increase the pointer when moving by 1 row + N, # number of columns in X + eps, # epsilon to avoid division by zero + BLOCK_SIZE: tl.constexpr, +): + # Map the program id to the row of X and Y it should compute. + row = tl.program_id(0) + Y += row * stride + X += row * stride + # Compute mean + mean = 0 + _mean = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + a = tl.load(X + cols, mask=cols < N, other=0.).to(tl.float32) + _mean += a + mean = tl.sum(_mean, axis=0) / N + # Compute variance + _var = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + x = tl.load(X + cols, mask=cols < N, other=0.).to(tl.float32) + x = tl.where(cols < N, x - mean, 0.) + _var += x * x + var = tl.sum(_var, axis=0) / N + rstd = 1 / tl.sqrt(var + eps) + # Write mean / rstd + tl.store(Mean + row, mean) + tl.store(Rstd + row, rstd) + # Normalize and apply linear transformation + for off in range(0, N, BLOCK_SIZE): + cols = off + tl.arange(0, BLOCK_SIZE) + mask = cols < N + w = tl.load(W + cols, mask=mask) + b = tl.load(B + cols, mask=mask) + x = tl.load(X + cols, mask=mask, other=0.).to(tl.float32) + x_hat = (x - mean) * rstd + y = x_hat * w + b + # Write output + tl.store(Y + cols, y, mask=mask) + + +class LayerNorm(torch.autograd.Function): + + @staticmethod + def forward(ctx, x, normalized_shape, weight, bias, eps, device): + # allocate output + y = torch.empty_like(x) + # reshape input data into 2D tensor + x_arg = x.reshape(-1, x.shape[-1]) + M, N = x_arg.shape + mean = torch.empty((M, ), dtype=torch.float32, device=device) + rstd = torch.empty((M, ), dtype=torch.float32, device=device) + # Less than 64KB per feature: enqueue fused kernel + MAX_FUSED_SIZE = 65536 // x.element_size() + BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(N)) + if N > BLOCK_SIZE: + raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.") + # heuristics for number of warps + num_warps = min(max(BLOCK_SIZE // 256, 1), 8) + # enqueue kernel + _layer_norm_fwd_fused[(M, )]( # + x_arg, y, weight, bias, mean, rstd, # + x_arg.stride(0), N, eps, # + BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps, num_ctas=1) + ctx.save_for_backward(x, weight, bias, mean, rstd) + ctx.BLOCK_SIZE = BLOCK_SIZE + ctx.num_warps = num_warps + ctx.eps = eps + return y + + +@pytest.mark.parametrize("M, N, dtype, eps", [ # + (M, N, dtype, eps) + for M in [1151] + for N in [8192] + for dtype in [torch.float16] + for eps in [1e-5] +]) +def test_layer_norm(M, N, dtype, eps, device): + layer_norm = LayerNorm.apply + # create data + x_shape = (M, N) + w_shape = (x_shape[-1], ) + weight = torch.rand(w_shape, dtype=dtype, device=device, requires_grad=False) + bias = torch.rand(w_shape, dtype=dtype, device=device, requires_grad=False) + x = -2.3 + 0.5 * torch.randn(x_shape, dtype=dtype, device=device) + dy = .1 * torch.randn_like(x) + x.requires_grad_(False) + + # forward pass + y_tri = layer_norm(x, w_shape, weight, bias, eps, device) + # TODO We can't compare against Torch layer_norm since it doesn't support float16 on CPU + #y_ref = torch.nn.functional.layer_norm(x, w_shape, weight, bias, eps).to(dtype) + + print(y_tri) + #print(y_ref) + + # compare + #assert torch.allclose(y_tri, y_ref, atol=1e-2, rtol=0) + + +@benchmark.measure() +def bench_layernorm(size, provider): + layer_norm = LayerNorm.apply + device = 'cpu' + eps = 1e-5 + dtype = torch.float16 + x_shape = (size, size) + w_shape = (x_shape[-1], ) + weight = torch.rand(w_shape, dtype=dtype, device=device, requires_grad=False) + bias = torch.rand(w_shape, dtype=dtype, device=device, requires_grad=False) + x = -2.3 + 0.5 * torch.randn(x_shape, dtype=dtype, device=device) + dy = .1 * torch.randn_like(x) + x.requires_grad_(False) + # forward pass + y_tri = layer_norm(x, w_shape, weight, bias, eps, device) + + +if __name__ == "__main__": + benchmark.select_cpu_backend() + for X in [2**i for i in range(10, 13, 1)]: + for provider in ['triton']: + bench_layernorm(X, provider) \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/python/examples/test_load_2d_tensor_block.py b/third_party/wafer/third_party/flir/python/examples/test_load_2d_tensor_block.py new file mode 100755 index 00000000..ec7d53d0 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_load_2d_tensor_block.py @@ -0,0 +1,78 @@ +import torch + +import triton +import triton.language as tl + + +""" + +|-----|-----|-----|-----| +| | | | | +|-----|-----|-----|-----| +| | | | | +|-----|-----|-----|-----| + +Each instance loads BLOCK_SIZE_ROW * BLOCK_SIZE_COL +""" + + +@triton.jit +def kernel( + x_ptr, + y_ptr, + n_rows, + n_cols, + stride_0, + stride_1, + BLOCK_SIZE_ROW: tl.constexpr, + BLOCK_SIZE_COL: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + pid1 = tl.program_id(axis=1) + + input_ptr = tl.make_block_ptr( + base=x_ptr, + shape=[n_rows, n_cols], + strides=[stride_0, stride_1], + offsets=[pid0 * BLOCK_SIZE_ROW, pid1 * BLOCK_SIZE_COL], + block_shape=[BLOCK_SIZE_ROW, BLOCK_SIZE_COL], + order=[1, 0], + ) + x = tl.load(input_ptr) + x = (2 * x) + 1 + output_ptr = tl.make_block_ptr( + base=y_ptr, + shape=[n_rows, n_cols], + strides=[stride_0, stride_1], + offsets=[pid0 * BLOCK_SIZE_ROW, pid1 * BLOCK_SIZE_COL], + block_shape=[BLOCK_SIZE_ROW, BLOCK_SIZE_COL], + order=[1, 0], + ) + tl.store(output_ptr, x) + + +def test(device): + n_rows = 512 + n_cols = 256 + x = torch.arange(0, n_rows * n_cols, 1, device=device, dtype=torch.float32).reshape( + [n_rows, n_cols] + ) + output = torch.full([n_rows, n_cols], -1, device=device, dtype=x.dtype) + BLOCK_SIZE_ROW = 4 + BLOCK_SIZE_COL = 2 + + grid = lambda meta: (n_rows // BLOCK_SIZE_ROW, n_cols // BLOCK_SIZE_COL) + + kernel[grid]( + x, + output, + n_rows, + n_cols, + x.stride(0), + x.stride(1), + BLOCK_SIZE_ROW=BLOCK_SIZE_ROW, + BLOCK_SIZE_COL=BLOCK_SIZE_COL, + ) + expected = (2 * x) + 1 + + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/third_party/flir/python/examples/test_load_2d_tensor_col.py b/third_party/wafer/third_party/flir/python/examples/test_load_2d_tensor_col.py new file mode 100755 index 00000000..f1ed8ada --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_load_2d_tensor_col.py @@ -0,0 +1,70 @@ +import torch + +import triton +import triton.language as tl + + +""" + +|-----|-----|-----|-----| +| | | | | +|-----|-----|-----|-----| +| | | | | +|-----|-----|-----|-----| + +Each instance loads the entire column +""" + + +@triton.jit +def kernel( + x_ptr, + y_ptr, + n_rows, + n_cols, + BLOCK_SIZE_ROW: tl.constexpr, + BLOCK_SIZE_COL: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + input_ptr = tl.make_block_ptr( + base=x_ptr, + shape=[n_rows, n_cols], + strides=[BLOCK_SIZE_COL, 1], + offsets=[0, pid0], + block_shape=[BLOCK_SIZE_ROW, 1], + order=[1, 0], + ) + x = tl.load(input_ptr) + output_ptr = tl.make_block_ptr( + base=y_ptr, + shape=[n_rows, n_cols], + strides=[BLOCK_SIZE_COL, 1], + offsets=[0, pid0], + block_shape=[BLOCK_SIZE_ROW, 1], + order=[1, 0], + ) + tl.store(output_ptr, x) + + +def test(device): + n_rows = 4 + n_cols = 2 + x = torch.arange(0, n_rows * n_cols, 1, device=device, dtype=torch.float32).reshape( + [n_rows, n_cols] + ) + output = torch.full([n_rows, n_cols], -1, device=device, dtype=x.dtype) + BLOCK_SIZE_ROW = n_rows + BLOCK_SIZE_COL = n_cols + + grid = lambda meta: (n_cols,) + + kernel[grid]( + x, + output, + n_rows, + n_cols, + BLOCK_SIZE_ROW=BLOCK_SIZE_ROW, + BLOCK_SIZE_COL=BLOCK_SIZE_COL, + ) + + torch.testing.assert_close(output, x, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/third_party/flir/python/examples/test_mask.py b/third_party/wafer/third_party/flir/python/examples/test_mask.py new file mode 100755 index 00000000..499dfa1d --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_mask.py @@ -0,0 +1,39 @@ +import torch + +import triton +import triton.language as tl + +from triton.backends.triton_shared.driver import CPUDriver + + +def test_mask(device): + @triton.jit + def test(in0, out0): + offs = 100 + tl.arange(0, 4) + out_offs = tl.arange(0, 4) + a = tl.load(in0 + offs, mask=offs < 4, other=-1) + tl.store(out0 + out_offs, a) + + SIZE = 8 + input = torch.arange(0, SIZE, device=device, dtype=torch.int32) + output = torch.full((SIZE,), -2, device=device, dtype=torch.int32) + + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + grid = lambda meta: (1,) + + src = triton.compiler.ASTSource( + fn=test, + signature="*fp32,*fp32,i32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) + + print(output) + test[grid](input, output) + print(input) + print(output) + torch.testing.assert_close(output, torch.tensor([-1, -1, -1, -1, -2, -2, -2, -2], device=device, dtype=torch.int32)) diff --git a/third_party/wafer/third_party/flir/python/examples/test_matmul.py b/third_party/wafer/third_party/flir/python/examples/test_matmul.py new file mode 100755 index 00000000..470006a5 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_matmul.py @@ -0,0 +1,174 @@ +import torch + +import triton +import triton.language as tl +import benchmark + +# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes: +# - A list of `triton.Config` objects that define different configurations of +# meta-parameters (e.g., `BLOCK_SIZE_M`) and compilation options (e.g., `num_warps`) to try +# - An auto-tuning *key* whose change in values will trigger evaluation of all the +# provided configs +# @triton.autotune( +# configs=[ +# triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=3, +# num_warps=8), +# triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=4, +# num_warps=4), +# triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=5, +# num_warps=2), +# triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8}, num_stages=5, +# num_warps=2), +# ], +# key=['M', 'N', 'K'], +# ) +@triton.jit +def matmul_kernel( + # Pointers to matrices + a_ptr, b_ptr, c_ptr, + # Matrix dimensions + M, N, K, + # The stride variables represent how much to increase the ptr by when moving by 1 + # element in a particular dimension. E.g. `stride_am` is how much to increase `a_ptr` + # by to get the element one row down (A has M rows). + stride_am, stride_ak, # + stride_bk, stride_bn, # + stride_cm, stride_cn, + # Meta-parameters + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, # + GROUP_SIZE_M: tl.constexpr, # + ACTIVATION: tl.constexpr # +): + """Kernel for computing the matmul C = A x B. + A has shape (M, K), B has shape (K, N) and C has shape (M, N) + """ + # ----------------------------------------------------------- + # Map program ids `pid` to the block of C it should compute. + # This is done in a grouped ordering to promote L2 data reuse. + # See above `L2 Cache Optimizations` section for details. + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + # ---------------------------------------------------------- + # Create pointers for the first blocks of A and B. + # We will advance this pointer as we move in the K direction + # and accumulate + # `a_ptrs` is a block of [BLOCK_SIZE_M, BLOCK_SIZE_K] pointers + # `b_ptrs` is a block of [BLOCK_SIZE_K, BLOCK_SIZE_N] pointers + # See above `Pointer Arithmetics` section for details + offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + + # ----------------------------------------------------------- + # Iterate to compute a block of the C matrix. + # We accumulate into a `[BLOCK_SIZE_M, BLOCK_SIZE_N]` block + # of fp32 values for higher accuracy. + # `accumulator` will be converted back to fp16 after the loop. + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + # Load the next block of A and B, generate a mask by checking the K dimension. + # If it is out of bounds, set it to 0. + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) + # We accumulate along the K dimension. + accumulator += tl.dot(a, b) + # Advance the ptrs to the next K block. + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + # You can fuse arbitrary activation functions here + # while the accumulator is still in FP32! + if ACTIVATION == "leaky_relu": + accumulator = leaky_relu(accumulator) + c = accumulator.to(tl.float32) + + # ----------------------------------------------------------- + # Write back the block of the output matrix C with masks. + offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, c, mask=c_mask) + + + +# We can fuse `leaky_relu` by providing it as an `ACTIVATION` meta-parameter in `_matmul`. +@triton.jit +def leaky_relu(x): + x = x + 1 + return tl.where(x >= 0, x, 0.01 * x) + + +def matmul(a, b, activation=""): + # Check constraints. + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + assert a.is_contiguous(), "Matrix A must be contiguous" + assert b.is_contiguous(), "Matrix B must be contiguous" + M, K = a.shape + K, N = b.shape + # Allocates output. + c = torch.empty((M, N), device=a.device, dtype=a.dtype) + # 1D launch kernel where each block gets its own program. + grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), ) + matmul_kernel[grid]( + a, b, c, # + M, N, K, # + a.stride(0), a.stride(1), # + b.stride(0), b.stride(1), # + c.stride(0), c.stride(1), # + ACTIVATION=activation, # + BLOCK_SIZE_M=32, + BLOCK_SIZE_N=64, + BLOCK_SIZE_K=16, + GROUP_SIZE_M=8 + ) + return c + + +def test_matmul(device): + torch.manual_seed(0) + rows1 = 179 + cols1 = 167 + rows2 = 167 + cols2 = 321 + a = torch.randn((rows1, cols1), device=device, dtype=torch.float32) + b = torch.randn((rows2, cols2), device=device, dtype=torch.float32) + # a = torch.full((rows1, cols1), 1, device='cpu', dtype=torch.float32) + # b = torch.full((rows2, cols2), 1, device='cpu', dtype=torch.float32) + triton_output = matmul(a, b) + torch_output = torch.matmul(a, b) + torch.testing.assert_close(triton_output, torch_output, atol=1e-2, rtol=0) + + +@benchmark.measure() +def bench_matmul(M, N, K, provider): + a = torch.randn((M, K), device='cpu', dtype=torch.float32) + b = torch.randn((K, N), device='cpu', dtype=torch.float32) + if provider == 'torch': + torch.matmul(a, b) + if provider == 'triton': + matmul(a, b) + + +if __name__ == "__main__": + benchmark.select_cpu_backend() + for X in [128 * i for i in range(2, 7)]: + for provider in ['torch', 'triton']: + bench_matmul(X, X, X, provider) diff --git a/third_party/wafer/third_party/flir/python/examples/test_modulo.py b/third_party/wafer/third_party/flir/python/examples/test_modulo.py new file mode 100755 index 00000000..f7e1775e --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_modulo.py @@ -0,0 +1,386 @@ +import torch + +import triton +import triton.language as tl + + +def test_wrap_stacked(device): + + @triton.jit + def wrap_stacked(a_ptr, c_ptr, M, N, stride_am, stride_an, stride_cm, + stride_cn, BLOCK_SIZE_K: tl.constexpr): + offs_am = (2 + tl.arange(0, 4)) % M + offs_an = tl.arange(0, 4) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + + offs_an[None, :] * stride_an) + + offs_cm = tl.arange(0, 4) + offs_cn = tl.arange(0, 4) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[ + None, :] + + for k in range(0, 2): + a = tl.load(a_ptrs) + tl.store(c_ptrs, a) + a_ptrs += BLOCK_SIZE_K * stride_an + c_ptrs += BLOCK_SIZE_K * stride_an + + M = 4 + N = 8 + A = torch.arange(0, M * N, device=device, dtype=torch.float32).reshape( + (M, N)) + out = torch.full((M, N), 88888, device=device, dtype=torch.float32) + grid = lambda meta: (1, ) + + wrap_stacked[grid](A, + out, + M, + N, + A.stride(0), + A.stride(1), + out.stride(0), + out.stride(1), + BLOCK_SIZE_K=4) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[16, 17, 18, 19, 20, 21, 22, 23], [24, 25, 26, 27, 28, 29, 30, 31], + [0, 1, 2, 3, 4, 5, 6, 7], [8, 9, 10, 11, 12, 13, 14, 15]], + device=device) + + assert torch.equal(expected_out.int(), out.int()) + + +def test_1d(device): + + @triton.jit + def mod_1d(a_ptr, c_ptr, M, N, stride_am, stride_an, stride_cm, stride_cn, + BLOCK_SIZE_K: tl.constexpr): + row = 7 + offs_an = (6 + tl.arange(0, 4)) % N + a_ptrs = a_ptr + (row * stride_am) + offs_an[None, :] * stride_an + + offs_cn = tl.arange(0, 4) + c_ptrs = c_ptr + stride_cn * offs_cn[None, :] + + a = tl.load(a_ptrs) + tl.store(c_ptrs, a) + + M = 8 + N = 8 + A = torch.arange(0, M * N, device=device, dtype=torch.float32).reshape( + (M, N)) + out = torch.full((M, N), 88888, device=device, dtype=torch.float32) + grid = lambda meta: (1, ) + + mod_1d[grid](A, + out, + M, + N, + A.stride(0), + A.stride(1), + out.stride(0), + out.stride(1), + BLOCK_SIZE_K=4) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[62, 63, 56, 57, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888]], + device=device) + + assert torch.equal(expected_out.int(), out.int()) + + +def test_2d(device): + + @triton.jit + def mod_2d(a_ptr, c_ptr, M, N, stride_am, stride_an, stride_cm, stride_cn, + BLOCK_SIZE_K: tl.constexpr): + offs_am = 2 + tl.arange(0, 4) + offs_an = (6 + tl.arange(0, 4)) % N + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + + offs_an[None, :] * stride_an) + + offs_cm = tl.arange(0, 4) + offs_cn = tl.arange(0, 4) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[ + None, :] + + a = tl.load(a_ptrs) + tl.store(c_ptrs, a) + + M = 8 + N = 8 + A = torch.arange(0, M * N, device=device, dtype=torch.float32).reshape( + (M, N)) + out = torch.full((M, N), 88888, device=device, dtype=torch.float32) + grid = lambda meta: (1, ) + + mod_2d[grid](A, + out, + M, + N, + A.stride(0), + A.stride(1), + out.stride(0), + out.stride(1), + BLOCK_SIZE_K=4) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[22, 23, 16, 17, 88888, 88888, 88888, 88888], + [30, 31, 24, 25, 88888, 88888, 88888, 88888], + [38, 39, 32, 33, 88888, 88888, 88888, 88888], + [46, 47, 40, 41, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888]], + device=device) + + assert torch.equal(expected_out.int(), out.int()) + + +def test_side_by_side_masked_loop(device): + + @triton.jit + def wrap_side_by_side_masked_loop(a_ptr, c_ptr, M, N, stride_am, stride_an, + stride_cm, stride_cn, + BLOCK_SIZE_K: tl.constexpr): + offs_am = 2 + tl.arange(0, BLOCK_SIZE_K) + offs_an = (6 + tl.arange(0, BLOCK_SIZE_K)) % N + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + + offs_an[None, :] * stride_an) + + offs_k = tl.arange(0, BLOCK_SIZE_K) + + offs_cm = tl.arange(0, BLOCK_SIZE_K) + offs_cn = tl.arange(0, BLOCK_SIZE_K) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[ + None, :] + + for k in range(0, 2): + a = tl.load(a_ptrs, mask=offs_k[:, None] < 2, other=-99) + tl.store(c_ptrs, a) + a_ptrs += BLOCK_SIZE_K * stride_am + c_ptrs += BLOCK_SIZE_K * stride_an + + M = 12 + N = 8 + A = torch.arange(0, M * N, device=device, dtype=torch.float32).reshape( + (M, N)) + out = torch.full((M, N), 88888, device=device, dtype=torch.float32) + print(out) + grid = lambda meta: (1, ) + + wrap_side_by_side_masked_loop[grid](A, + out, + M, + N, + A.stride(0), + A.stride(1), + out.stride(0), + out.stride(1), + BLOCK_SIZE_K=4) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[22, 23, 16, 17, 54, 55, 48, 49], [30, 31, 24, 25, 62, 63, 56, 57], + [-99, -99, -99, -99, -99, -99, -99, -99], + [-99, -99, -99, -99, -99, -99, -99, -99], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888], + [88888, 88888, 88888, 88888, 88888, 88888, 88888, 88888]], + dtype=torch.int32) + + assert torch.equal(expected_out.int(), out.int()) + + +def test_stacked_masked_loop(device): + + @triton.jit + def wrap_stacked_masked_loop(a_ptr, c_ptr, M, N, stride_am, stride_an, + stride_cm, stride_cn, + BLOCK_SIZE_K: tl.constexpr): + offs_am = (2 + tl.arange(0, BLOCK_SIZE_K)) % M + offs_an = 3 + tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + + offs_an[None, :] * stride_an) + + offs_cm = tl.arange(0, BLOCK_SIZE_K) + offs_cn = tl.arange(0, BLOCK_SIZE_K) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[ + None, :] + + offs_k = tl.arange(0, BLOCK_SIZE_K) + + for k in range(0, 2): + a = tl.load(a_ptrs, mask=offs_k[None, :] < 3, other=-99) + tl.store(c_ptrs, a) + a_ptrs += BLOCK_SIZE_K * stride_an + c_ptrs += BLOCK_SIZE_K * stride_an + + M = 4 + N = 12 + BLOCK_SIZE_M = 4 + BLOCK_SIZE_N = 4 + A = torch.arange(0, M * N, device=device, dtype=torch.float32).reshape( + (M, N)) + out = torch.full((BLOCK_SIZE_M, N), + 88888, + device=device, + dtype=torch.float32) + print(out) + grid = lambda meta: (1, ) + + wrap_stacked_masked_loop[grid](A, + out, + M, + N, + A.stride(0), + A.stride(1), + out.stride(0), + out.stride(1), + BLOCK_SIZE_K=4) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor([ + [ + 27.0, + 28.0, + 29.0, + -99.0, + 31.0, + 32.0, + 33.0, + -99.0, + 88888, + 88888, + 88888, + 88888, + ], + [ + 39.0, + 40.0, + 41.0, + -99.0, + 43.0, + 44.0, + 45.0, + -99.0, + 88888, + 88888, + 88888, + 88888, + ], + [ + 3.0, + 4.0, + 5.0, + -99.0, + 7.0, + 8.0, + 9.0, + -99.0, + 88888, + 88888, + 88888, + 88888, + ], + [ + 15.0, + 16.0, + 17.0, + -99.0, + 19.0, + 20.0, + 21.0, + -99.0, + 88888, + 88888, + 88888, + 88888, + ], + ], ) + + assert torch.equal(expected_out.int(), out.int()) + + +def test_torch_inductor_pattern(): + + @triton.jit + def triton_(in_ptr2, out_ptr2, rnumel, XBLOCK: tl.constexpr, + RBLOCK: tl.constexpr): + xnumel = 128 + rnumel = 32 + xoffset = tl.program_id(0) * XBLOCK + xindex = xoffset + tl.arange(0, XBLOCK)[:, None] + rbase = tl.arange(0, RBLOCK)[None, :] + x0 = xindex % 7 + x0 = xindex + roffset = 0 + rindex = roffset + rbase + rmask = rindex < rnumel + r2 = rindex + tmp3 = tl.load(in_ptr2 + (r2 + (xnumel * x0)), rmask, other=77) + tl.store( + out_ptr2 + (XBLOCK * tl.arange(0, RBLOCK)[None, :] + + tl.arange(0, XBLOCK)[:, None]), tmp3) + + device = "cpu" + xnumel = 128 + rnumel = 32 + + XBLOCK = 4 + RBLOCK = 64 + A = torch.arange(0, xnumel * rnumel, device=device, + dtype=torch.int32).reshape((xnumel, rnumel)) + out = torch.full((XBLOCK, RBLOCK), 88888, device=device, dtype=torch.int32) + grid = lambda meta: (1, ) + + triton_[grid](A, out, rnumel, XBLOCK=XBLOCK, RBLOCK=RBLOCK) + + # Expected output copied from running triton on NVDIA gpu + expected_out = torch.tensor( + [[ + 0, 128, 256, 384, 1, 129, 257, 385, 2, 130, 258, 386, 3, 131, 259, + 387, 4, 132, 260, 388, 5, 133, 261, 389, 6, 134, 262, 390, 7, 135, + 263, 391, 8, 136, 264, 392, 9, 137, 265, 393, 10, 138, 266, 394, + 11, 139, 267, 395, 12, 140, 268, 396, 13, 141, 269, 397, 14, 142, + 270, 398, 15, 143, 271, 399 + ], + [ + 16, 144, 272, 400, 17, 145, 273, 401, 18, 146, 274, 402, 19, 147, + 275, 403, 20, 148, 276, 404, 21, 149, 277, 405, 22, 150, 278, 406, + 23, 151, 279, 407, 24, 152, 280, 408, 25, 153, 281, 409, 26, 154, + 282, 410, 27, 155, 283, 411, 28, 156, 284, 412, 29, 157, 285, 413, + 30, 158, 286, 414, 31, 159, 287, 415 + ], + [ + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77 + ], + [ + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, + 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77, 77 + ]], + device=device, + dtype=torch.int32) + + assert torch.equal(expected_out.int(), out.int()) diff --git a/third_party/wafer/third_party/flir/python/examples/test_nested_loops.py b/third_party/wafer/third_party/flir/python/examples/test_nested_loops.py new file mode 100755 index 00000000..3df04bd6 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_nested_loops.py @@ -0,0 +1,413 @@ +import torch + +import triton +from triton.backends.compiler import GPUTarget +from triton.backends.triton_shared.driver import CPUDriver +import triton.language as tl + + +# Not used for testing but serves as a template to generate the lit test at +# test/Conversion/TritonToStructured/ridiculously_nested_loops.mlir +@triton.jit +def nested_who_knows_how_many_levels(in_ptr, out_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = in_ptr + (offs_am[:, None] * stride_m + + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = out_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[ + None, :] + + for i1 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j1 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k1 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + for i2 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j2 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k2 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + for i3 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j3 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k3 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + for i4 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j4 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k4 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + for i5 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j5 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k5 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + a_ptrs += 2 * stride_n + + for i6 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j6 in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k6 in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + a_ptrs += 2 * stride_n + + + a_ptrs += 2 * stride_n + + +@triton.jit +def nested_use_same_level_loop_results(in_ptr, out_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = in_ptr + (offs_am[:, None] * stride_m + + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = out_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[ + None, :] + + for i1 in range(0, 2): + a1 = tl.load(a_ptrs) + + for j1 in range(0, 2): + a_ptrs += 2 * stride_n + + for i6 in range(0, 2): + a1 = tl.load(a_ptrs) + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + a_ptrs += 2 * stride_n + + a_ptrs += 2 * stride_n + +@triton.jit +def nested2_complex_body(a_ptr, c_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = a_ptr + (offs_am[:, None] * stride_m + + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = c_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[ + None, :] + + for i in range(0, 2): + a_ptrs_copy = a_ptrs + c_ptrs_copy = c_ptrs + + a_ptrs += 1 + c_ptrs += 1 + + for j in range(0, 2): + a2 = tl.load(a_ptrs) + tl.store(c_ptrs, a2) + a_ptrs += 3 + c_ptrs += 3 + + a_ptrs = a_ptrs_copy + 2 * stride_m + 1 + c_ptrs = c_ptrs_copy + 2 * stride_m + 1 + + + +@triton.jit +def nested2_use_loop_results(in_ptr, out_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = in_ptr + (offs_am[:, None] * stride_m + + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = out_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[ + None, :] + + for i in range(0, 2): + a2 = tl.load(a_ptrs) + tl.store(c_ptrs, a2) + + a_ptrs += 4 * stride_n + c_ptrs += 4 * stride_n + + + for j in range(0, 2): + a2 = tl.load(a_ptrs) + tl.store(c_ptrs, a2) + a_ptrs += 4 * stride_n + c_ptrs += 4 * stride_n + + +@triton.jit +def nested3(in_ptr, out_ptr, stride_m, stride_n): + offs_am = tl.arange(0, 2) + offs_an = tl.arange(0, 2) + a_ptrs = in_ptr + (offs_am[:, None] * stride_m + + offs_an[None, :] * stride_n) + + offs_cm = tl.arange(0, 2) + offs_cn = tl.arange(0, 2) + c_ptrs = out_ptr + stride_m * offs_cm[:, None] + stride_n * offs_cn[ + None, :] + + for i in range(0, 2): + a1 = tl.load(a_ptrs) + + for j in range(0, 2): + a_ptrs += 2 * stride_n + a2 = tl.load(a_ptrs) + + for k in range(0, 2): + a_ptrs += 2 * stride_n + a3 = tl.load(a_ptrs) + tl.store(c_ptrs, a1) + c_ptrs += 2 * stride_n + + tl.store(c_ptrs, a2) + c_ptrs += 2 * stride_n + tl.store(c_ptrs, a3) + c_ptrs += 2 * stride_n + + + a_ptrs += 2 * stride_n + +def test_nested3(): + n_rows = 4 + n_cols = 48 + expected = torch.tensor([[ 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 6, 7, 0, 1, + 8, 9, 10, 11, 0, 1, 8, 9, 12, 13, 14, 15, 16, 17, + 18, 19, 14, 15, 16, 17, 20, 21, 14, 15, 22, 23, 24, 25, + 14, 15, 22, 23, 26, 27], + [48, 49, 50, 51, 52, 53, 48, 49, 50, 51, 54, 55, 48, 49, + 56, 57, 58, 59, 48, 49, 56, 57, 60, 61, 62, 63, 64, 65, + 66, 67, 62, 63, 64, 65, 68, 69, 62, 63, 70, 71, 72, 73, + 62, 63, 70, 71, 74, 75], + [ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0], + [ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0]], dtype=torch.int32, device='cpu') + triton.runtime.driver.set_active(CPUDriver()) + x = torch.arange(0, n_rows * n_cols, device="cpu", dtype=torch.int32).reshape([n_rows, n_cols]) + output = torch.zeros([n_rows, n_cols], device=x.device, dtype=x.dtype) + grid = lambda meta: (n_cols // 4,) + + print('before:') + print(x) + print(output) + + nested3[grid](x, output, x.stride(0), x.stride(1)) + print(output) + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) + print("Pass!") + + src = triton.compiler.ASTSource( + fn=nested3, + signature="*fp32,*fp32,i32,i32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) + print('Pass') + + +def test_nested2_use_loop_results(): + n_rows = 4 + n_cols = 32 + expected = torch.tensor([[ 0, 1, 0, 0, 4, 5, 0, 0, 8, 9, 0, 0, 12, 13, 0, 0, 16, 17, + 0, 0, 20, 21, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [32, 33, 0, 0, 36, 37, 0, 0, 40, 41, 0, 0, 44, 45, 0, 0, 48, 49, + 0, 0, 52, 53, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]], + device='cpu', dtype=torch.int32) + # x = torch.arange(0, n_rows * n_cols, device="cuda", dtype=torch.int32).reshape([n_rows, n_cols]) + triton.runtime.driver.set_active(CPUDriver()) + x = torch.arange(0, n_rows * n_cols, device="cpu", dtype=torch.int32).reshape([n_rows, n_cols]) + output = torch.zeros([n_rows, n_cols], device=x.device, dtype=x.dtype) + grid = lambda meta: (n_cols // 4,) + + print('before:') + print(x) + print(output) + + nested2_use_loop_results[grid](x, output, x.stride(0), x.stride(1)) + print(output) + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) + print("Pass!") + + src = triton.compiler.ASTSource( + fn=nested2_use_loop_results, + signature="*fp32,*fp32,i32,i32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) + print('Pass') + + +def test_nested2_complex_body(): + n_rows = 4 + n_cols = 8 + grid = lambda meta: (n_cols // 4,) + expected = torch.tensor([[ 0, 1, 2, 0, 4, 5, 0, 0], + [ 0, 9, 10, 0, 12, 13, 0, 0], + [ 0, 0, 18, 19, 0, 21, 22, 0], + [ 0, 0, 26, 27, 0, 29, 30, 0]], device='cpu', dtype=torch.int32) + + + x = torch.arange(0, n_rows * n_cols, device="cpu", dtype=torch.int32).reshape([n_rows, n_cols]) + triton.runtime.driver.set_active(CPUDriver()) + output = torch.zeros([n_rows, n_cols], device=x.device, dtype=x.dtype) + + + print('before:') + print(x) + print(output) + + nested2_complex_body[grid](x, output, x.stride(0), x.stride(1)) + print(output) + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) + print("Pass!") + + src = triton.compiler.ASTSource( + fn=nested2_complex_body, + signature="*fp32,*fp32,i32,i32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) + print('Pass') + +def test_nested2_use_same_level_loop_result(): + n_rows = 4 + n_cols = 32 + grid = lambda meta: (n_cols // 4,) + expected = torch.tensor([[ 4, 5, 0, 0, 6, 7, 8, 9, 0, 0, 10, 11, 18, 19, 0, 0, 20, 21, + 22, 23, 0, 0, 24, 25, 0, 0, 0, 0, 0, 0, 0, 0], + [36, 37, 0, 0, 38, 39, 40, 41, 0, 0, 42, 43, 50, 51, 0, 0, 52, 53, + 54, 55, 0, 0, 56, 57, 0, 0, 0, 0, 0, 0, 0, 0], + [ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]], + device='cpu', dtype=torch.int32) + + + x = torch.arange(0, n_rows * n_cols, device="cpu", dtype=torch.int32).reshape([n_rows, n_cols]) + triton.runtime.driver.set_active(CPUDriver()) + output = torch.zeros([n_rows, n_cols], device=x.device, dtype=x.dtype) + + + print('before:') + print(x) + print(output) + + nested_use_same_level_loop_results[grid](x, output, x.stride(0), x.stride(1)) + print(output) + torch.testing.assert_close(output, expected, rtol=0.001, atol=1e-5) + print("Pass!") + + src = triton.compiler.ASTSource( + fn=nested_use_same_level_loop_results, + signature="*fp32,*fp32,i32,i32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) + print('Pass') diff --git a/third_party/wafer/third_party/flir/python/examples/test_reduce.py b/third_party/wafer/third_party/flir/python/examples/test_reduce.py new file mode 100755 index 00000000..3757975f --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_reduce.py @@ -0,0 +1,61 @@ +import torch + +import triton +from triton.backends.compiler import GPUTarget +import triton.language as tl + + +@triton.jit +def reduce_kernel_2d( + x_ptr, + output_ptr, + stride, + n_elements, + BLOCK_SIZE: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + x = tl.load( + tl.make_block_ptr( + base=x_ptr, + shape=[n_elements * tl.num_programs(0)], + strides=[1], + offsets=[stride * pid0], + block_shape=[BLOCK_SIZE], + order=[0], + ), + boundary_check=[0], + ) + output = triton.language.sum(x, axis=0).to(dtype=x.dtype) + tl.store(output_ptr + pid0, output) + + +def test(device): + n_rows = 16 + n_cols = 32 + x = torch.rand([n_cols, n_rows], device=device, dtype=torch.float32) + output = torch.empty([n_cols], device=device, dtype=x.dtype) + BLOCK_SIZE = n_rows + grid = lambda meta: (n_cols,) + + reduce_kernel_2d[grid](x, output, x.stride(0), n_rows, BLOCK_SIZE=BLOCK_SIZE) + ans = torch.sum(x, dim=1) + torch.testing.assert_close(output, ans, rtol=0.001, atol=1e-5) + + # TODO: need to check some conditions otherwise the code below does not make any difference for the test + src = triton.compiler.ASTSource( + fn=reduce_kernel_2d, + signature={"x_ptr": "*fp32", + "output_ptr": "*fp32", + "stride": "i32", + "n_elements": "i32", + "BLOCK_SIZE": "constexpr"}, + constexprs={"BLOCK_SIZE": 32} + ) + ret = triton.compile( + src, + target=GPUTarget(device, 0, 0) + ) + print(ret.asm["ttir"]) + print(ret.asm["ttsharedir"]) + print(ret.asm["llir"]) + print(ret.asm["obj"]) diff --git a/third_party/wafer/third_party/flir/python/examples/test_scalar_store.py b/third_party/wafer/third_party/flir/python/examples/test_scalar_store.py new file mode 100755 index 00000000..7ef14d46 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_scalar_store.py @@ -0,0 +1,54 @@ +import torch + +import triton +import triton.language as tl + +from triton.backends.triton_shared.driver import CPUDriver + +@triton.jit +def test_scalar_store( + output_ptr, + BLOCK_SIZE: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + base_ptr = output_ptr + pid0 + for i in range(0, BLOCK_SIZE // 2): + output = i * 2 + for j in range(0, BLOCK_SIZE // 4): + output += j + tl.store(base_ptr, output) + base_ptr += 1 + + +def compile(): + src = triton.compiler.ASTSource( + fn=test_scalar_store, + signature="*fp32", + constexprs={ + "BLOCK_SIZE": 8 + } + ) + ret = triton.compile( + src + ) + print(ret.asm["ttir"]) + + + +def test(device): + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + BLOCK_SIZE = 8 + x = torch.full([BLOCK_SIZE], -1, device=device, dtype=torch.float32) + output = torch.full((BLOCK_SIZE,), -99, device=device, dtype=x.dtype) + grid = lambda meta: (1,) + + print(x) + print(output) + + test_scalar_store[grid](output, BLOCK_SIZE=BLOCK_SIZE) + print('---') + print(output) + ans = torch.arange(BLOCK_SIZE, device=device, dtype=torch.float32) + torch.testing.assert_close(output, ans, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/third_party/flir/python/examples/test_sign_extend.py b/third_party/wafer/third_party/flir/python/examples/test_sign_extend.py new file mode 100755 index 00000000..21726e70 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_sign_extend.py @@ -0,0 +1,39 @@ +import torch + +import triton + +import triton.language as tl + +from triton.backends.triton_shared.driver import CPUDriver + +@triton.jit +def sign_extend(off, in0, out0, in0_size): + offset = tl.load(off).to(tl.int64) + offsets = offset + tl.arange(0, 4) + a = tl.load(in0 + offsets, mask=offsets < in0_size, other=11) + tl.store(out0 + tl.arange(0, 4), a) + +def compile(): + src = triton.compiler.ASTSource( + fn=sign_extend, + signature="*i32,*fp32,*fp32,i32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) + +def test_sign_extend(device): + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + SIZE = 4 + offsets = torch.full((1, ), 1, device=device, dtype=torch.int32) + input = torch.arange(0, SIZE, device=device, dtype=torch.int32) + output = torch.full((SIZE,), -1, device=device, dtype=torch.int32) + grid = lambda meta: (1,) + print(output) + sign_extend[grid](offsets, input, output, SIZE) + print(input) + print(output) + torch.testing.assert_close(torch.tensor([1, 2, 3, 11], device=device, dtype=torch.int32), output) diff --git a/third_party/wafer/third_party/flir/python/examples/test_softmax.py b/third_party/wafer/third_party/flir/python/examples/test_softmax.py new file mode 100755 index 00000000..b3d43dfa --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_softmax.py @@ -0,0 +1,82 @@ +import torch + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def softmax_kernel(output_ptr, input_ptr, input_row_stride, output_row_stride, n_cols, BLOCK_SIZE: tl.constexpr): + # The rows of the softmax are independent, so we parallelize across those + row_idx = tl.program_id(0) + # The stride represents how much we need to increase the pointer to advance 1 row + row_start_ptr = input_ptr + row_idx * input_row_stride + # The block size is the next power of two greater than n_cols, so we can fit each + # row in a single block + col_offsets = tl.arange(0, BLOCK_SIZE) + input_ptrs = row_start_ptr + col_offsets + # Load the row into SRAM, using a mask since BLOCK_SIZE may be > than n_cols + row = tl.load(input_ptrs, mask=col_offsets < n_cols, other=-float('inf')) + # Subtract maximum for numerical stability + row_minus_max = row - tl.max(row, axis=0) + # Note that exponentiation in Triton is fast but approximate (i.e., think __expf in CUDA) + numerator = tl.exp(row_minus_max) + denominator = tl.sum(numerator, axis=0) + softmax_output = numerator / denominator + # Write back output to DRAM + output_row_start_ptr = output_ptr + row_idx * output_row_stride + output_ptrs = output_row_start_ptr + col_offsets + tl.store(output_ptrs, softmax_output, mask=col_offsets < n_cols) + + +def softmax(x): + n_rows, n_cols = x.shape + # The block size is the smallest power of two greater than the number of columns in `x` + BLOCK_SIZE = triton.next_power_of_2(n_cols) + # Another trick we can use is to ask the compiler to use more threads per row by + # increasing the number of warps (`num_warps`) over which each row is distributed. + # You will see in the next tutorial how to auto-tune this value in a more natural + # way so you don't have to come up with manual heuristics yourself. + num_warps = 4 + if BLOCK_SIZE >= 2048: + num_warps = 8 + if BLOCK_SIZE >= 4096: + num_warps = 16 + # Allocate output + y = torch.empty_like(x) + # Enqueue kernel. The 1D launch grid is simple: we have one kernel instance per row o + # f the input matrix + softmax_kernel[(n_rows, )]( + y, + x, + x.stride(0), + y.stride(0), + n_cols, + num_warps=num_warps, + BLOCK_SIZE=BLOCK_SIZE, + ) + return y + +def test_softmax(device): + torch.manual_seed(0) + x = torch.randn(1823, 781, device=device) + y_triton = softmax(x) + y_torch = torch.softmax(x, axis=1) + assert torch.allclose(y_triton, y_torch), (y_triton, y_torch) + + +@benchmark.measure() +def bench_softmax(size, provider): + torch.manual_seed(0) + x = torch.randn(size, size, device='cpu') + if provider == 'torch': + torch.softmax(x, axis=1) + if provider == 'triton': + softmax(x) + + +if __name__ == "__main__": + benchmark.select_cpu_backend() + for X in [2**i for i in range(10, 14, 1)]: + for provider in ['torch', 'triton']: + bench_softmax(X, provider) \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/python/examples/test_splat.py b/third_party/wafer/third_party/flir/python/examples/test_splat.py new file mode 100755 index 00000000..a396b5ef --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_splat.py @@ -0,0 +1,41 @@ +import torch + +import triton +import triton.language as tl + + +@triton.jit +def splat( + f32_val, + f32_out, + stride_row, + stride_col, + BLOCK_SIZE_ROW: tl.constexpr, + BLOCK_SIZE_COL: tl.constexpr, +): + pid0 = tl.program_id(axis=0) + x = tl.full((2, BLOCK_SIZE_COL), f32_val, dtype=tl.float32) + offs_row = 2 * pid0 + tl.arange(0, 2) + offs_col = tl.arange(0, BLOCK_SIZE_COL) + a_ptrs = f32_out + (offs_row[:, None] * stride_row + offs_col[None, :] * stride_col) + tl.store(a_ptrs, x) + + +def test(device): + n_rows = 256 + n_cols = 512 + fill_value = 123.456 + expected_result = torch.full((n_rows, n_cols), fill_value, dtype=torch.float32) + output = torch.empty([n_rows, n_cols], device=device, dtype=expected_result.dtype) + grid = lambda meta: (n_rows // 2,) + + splat[grid]( + fill_value, + output, + output.stride(0), + output.stride(1), + BLOCK_SIZE_ROW=n_rows, + BLOCK_SIZE_COL=n_cols, + ) + + torch.testing.assert_close(output, expected_result, rtol=0.001, atol=1e-5) diff --git a/third_party/wafer/third_party/flir/python/examples/test_swap.py b/third_party/wafer/third_party/flir/python/examples/test_swap.py new file mode 100755 index 00000000..2693ab60 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_swap.py @@ -0,0 +1,44 @@ +import torch + +import triton +import triton.language as tl + +# The purpose of this kernel and test is to catch incorrectly optimized kernels +# where copy elimination happens erroneously in the absence of explicit memory allocation. +# Such optimization bugs can result in incorrect behavior when swapping two arrays, +# particularly when both arrays unintentionally end up with the same data due to +# missing intermediate storage or mismanaged memory access. + +@triton.jit +def swap_kernel( + x_ptr, # *Pointer* to first inout vector. + y_ptr, # *Pointer* to second inout vector. + BLOCK_SIZE: tl.constexpr, # Number of elements each program should process. + # NOTE: `constexpr` so it can be used as a shape value. +): + pid = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0. + block_start = pid * BLOCK_SIZE + offsets = block_start + tl.arange(0, BLOCK_SIZE) + x = tl.load(x_ptr + offsets) + y = tl.load(y_ptr + offsets) + tl.store(x_ptr + offsets, y) + tl.store(y_ptr + offsets, x) + + +def swap(x: torch.Tensor, y: torch.Tensor): + n_elements = x.numel() + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) + swap_kernel[grid](x, y, BLOCK_SIZE=1024) + + +def test(device): + torch.manual_seed(0) + size = 10240 + x = torch.rand(size, device=device) + y = torch.rand(size, device=device) + assert not torch.equal(x, y) + x_ = x.clone() + y_ = y.clone() + swap(x, y) + assert torch.equal(x, y_) + assert torch.equal(y, x_) diff --git a/third_party/wafer/third_party/flir/python/examples/test_tensor_index_iterargs.py b/third_party/wafer/third_party/flir/python/examples/test_tensor_index_iterargs.py new file mode 100755 index 00000000..3f5878f0 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_tensor_index_iterargs.py @@ -0,0 +1,114 @@ +import torch + +import triton +import triton.language as tl + +from triton.backends.triton_shared.driver import CPUDriver + +def test_tensor_indices_nested_with_mask(device): + @triton.jit + def addptr_with_masks(in0, out0, mask_bound): + offs = tl.arange(0, 4) + out_offs = tl.arange(0, 4) + # We're loading 16 elements here, the bound is set to 14 so that + # the mask only applies to the last iteration's load + # TODO: The current mask implementation in triton-shared does not seem + # to work when the mask applies to the entire tensor load, perhaps + # the lowerings for subviews with 0-dimensions do not work? + for i in range(0, 4): + mask = offs < mask_bound + a = tl.load(in0 + offs, mask=mask, other=-11) + tl.store(out0 + out_offs, a) + offs += 4 + out_offs += 4 + + + SIZE = 17 + input = torch.arange(0, SIZE, device=device, dtype=torch.int32) + output = torch.full((SIZE,), -1, device=device, dtype=torch.int32) + + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + grid = lambda meta: (1,) + + print(output) + addptr_with_masks[grid](input, output, 14) + expected_output = torch.tensor([ 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, + -11, -11, -1], dtype=torch.int32, device=device) + torch.testing.assert_close(output, expected_output) + print(input) + print(output) + + +def test_tensor_indices_nested(device): + @triton.jit + def tensor_indices_nested(in0, out0): + offs = tl.arange(0, 4) + out_offs = tl.arange(0, 4) + for i in range(0, 2): + offs += i * 2 + a = tl.load(in0 + offs) + tl.store(out0 + out_offs, a) + offs += 4 + out_offs += 4 + for j in range(0, 3): + offs += j * 3 + a = tl.load(in0 + offs) + tl.store(out0 + out_offs, a) + offs += 4 + out_offs += 4 + + SIZE = 64 + input = torch.arange(0, SIZE, device=device, dtype=torch.int32) + output = torch.full((SIZE,), -1, device=device, dtype=torch.int32) + + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + grid = lambda meta: (1,) + + print(output) + tensor_indices_nested[grid](input, output) + expected_output = torch.tensor([ 0, 1, 2, 3, 4, 5, 6, 7, 11, 12, 13, 14, 21, 22, 23, 24, 27, 28, + 29, 30, 31, 32, 33, 34, 38, 39, 40, 41, 48, 49, 50, 51, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1], device=device, + dtype=torch.int32) + torch.testing.assert_close(output, expected_output) + print(input) + print(output) + +def test_integer_tensor(device): + @triton.jit + def test_1(out0): + offs = tl.arange(0, 4) + out_offs = tl.arange(0, 4) + for i in range(0, 2): + tl.store(out0 + out_offs, offs) + out_offs += 4 + offs += 4 + + + SIZE = 8 + input = torch.arange(0, SIZE, device=device, dtype=torch.int32) + output = torch.full((SIZE,), -1, device=device, dtype=torch.int32) + + if device == 'cpu': + triton.runtime.driver.set_active(CPUDriver()) + + grid = lambda meta: (1,) + + print(output) + test_1[grid](output) + print(input) + print(output) + torch.testing.assert_close(input, output) + src = triton.compiler.ASTSource( + fn=test_1, + signature="*fp32", + ) + ret = triton.compile( + src, + ) + print(ret.asm["ttir"]) diff --git a/third_party/wafer/third_party/flir/python/examples/test_vec_add.py b/third_party/wafer/third_party/flir/python/examples/test_vec_add.py new file mode 100755 index 00000000..db2fa098 --- /dev/null +++ b/third_party/wafer/third_party/flir/python/examples/test_vec_add.py @@ -0,0 +1,85 @@ +import torch + +import triton +import triton.language as tl +import benchmark + + +@triton.jit +def add_kernel( + x_ptr, # *Pointer* to first input vector. + y_ptr, # *Pointer* to second input vector. + output_ptr, # *Pointer* to output vector. + n_elements, # Size of the vector. + BLOCK_SIZE: tl.constexpr, # Number of elements each program should process. + # NOTE: `constexpr` so it can be used as a shape value. +): + # There are multiple 'programs' processing different data. We identify which program + # we are here: + pid = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0. + # This program will process inputs that are offset from the initial data. + # For instance, if you had a vector of length 256 and block_size of 64, the programs + # would each access the elements [0:64, 64:128, 128:192, 192:256]. + # Note that offsets is a list of pointers: + block_start = pid * BLOCK_SIZE + offsets = block_start + tl.arange(0, BLOCK_SIZE) + # Create a mask to guard memory operations against out-of-bounds accesses. + mask = offsets < n_elements + # Load x and y from DRAM, masking out any extra elements in case the input is not a + # multiple of the block size. + x = tl.load(x_ptr + offsets, mask=mask) + y = tl.load(y_ptr + offsets, mask=mask) + output = x + y + # Write x + y back to DRAM. + tl.store(output_ptr + offsets, output, mask=mask) + + +def add(x: torch.Tensor, y: torch.Tensor): + # We need to preallocate the output. + output = torch.empty_like(x) + # assert x.is_cuda and y.is_cuda and output.is_cuda + n_elements = output.numel() + # The SPMD launch grid denotes the number of kernel instances that run in parallel. + # It is analogous to CUDA launch grids. It can be either Tuple[int], or Callable(metaparameters) -> Tuple[int]. + # In this case, we use a 1D grid where the size is the number of blocks: + grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) + # NOTE: + # - Each torch.tensor object is implicitly converted into a pointer to its first element. + # - `triton.jit`'ed functions can be indexed with a launch grid to obtain a callable GPU kernel. + # - Don't forget to pass meta-parameters as keywords arguments. + add_kernel[grid](x, y, output, n_elements, BLOCK_SIZE=1024) + # We return a handle to z but, since `torch.cuda.synchronize()` hasn't been called, the kernel is still + # running asynchronously at this point. + return output + + +def test(device): + torch.manual_seed(0) + size = 98432 + x = torch.rand(size, device=device) + y = torch.rand(size, device=device) + output_torch = x + y + output_triton = add(x, y) + # TODO: need to check some conditions otherwise the code below does not make any difference for the test + print("expected", output_torch) + print("actual", output_triton) + print( + f"The maximum difference between torch and triton is " + f"{torch.max(torch.abs(output_torch - output_triton))}" + ) + +@benchmark.measure() +def bench_vecadd(size, provider): + a = torch.rand(size, device='cpu', dtype=torch.float32) + b = torch.rand(size, device='cpu', dtype=torch.float32) + if provider == 'torch': + a + b + if provider == 'triton': + add(a, b) + + +if __name__ == "__main__": + benchmark.select_cpu_backend() + for X in [2**i for i in range(22, 25, 1)]: + for provider in ['torch', 'triton']: + bench_vecadd(X, provider) \ No newline at end of file diff --git a/third_party/wafer/third_party/flir/test/CMakeLists.txt b/third_party/wafer/third_party/flir/test/CMakeLists.txt new file mode 100755 index 00000000..33308aa8 --- /dev/null +++ b/third_party/wafer/third_party/flir/test/CMakeLists.txt @@ -0,0 +1,27 @@ + +llvm_canonicalize_cmake_booleans( + MLIR_ENABLE_BINDINGS_PYTHON +) + +configure_lit_site_cfg( + ${CMAKE_CURRENT_SOURCE_DIR}/lit.site.cfg.py.in + ${CMAKE_CURRENT_BINARY_DIR}/lit.site.cfg.py + MAIN_CONFIG + ${CMAKE_CURRENT_SOURCe_DIR}/lit.cfg.py +) + +set(TRITON_SHARED_TEST_DEPENDS + triton-shared-opt +) + +set(FILECHECK_PATH "${LLVM_LIBRARY_DIR}/../bin/FileCheck") +set(LIT_ARGS "-Dfilecheck=${FILECHECK_PATH}") +add_lit_testsuite(check-triton-shared-lit-tests "Running the triton-shared regression tests" + ${CMAKE_CURRENT_BINARY_DIR} + ARGS ${LIT_ARGS} + DEPENDS ${TRITON_SHARED_TEST_DEPENDS} + ) + +set_target_properties(check-triton-shared-lit-tests PROPERTIES FOLDER "Tests") + +add_lit_testsuites(TRITON-SHARED-LIT-TESTS ${CMAKE_CURRENT_SOURCE_DIR} DEPENDS ${TRITON_SHARED_TEST_DEPENDS}) diff --git a/third_party/wafer/third_party/flir/test/README.md b/third_party/wafer/third_party/flir/test/README.md new file mode 100755 index 00000000..e9100cac --- /dev/null +++ b/third_party/wafer/third_party/flir/test/README.md @@ -0,0 +1,2 @@ +# triton-shared +shared middle layer for Triton as a submodule of Triton diff --git a/third_party/wafer/third_party/flir/test/lit.cfg.py b/third_party/wafer/third_party/flir/test/lit.cfg.py new file mode 100755 index 00000000..1ce9eb9c --- /dev/null +++ b/third_party/wafer/third_party/flir/test/lit.cfg.py @@ -0,0 +1,74 @@ +# -*- Python -*- + +import os +import platform +import re +import subprocess +import tempfile + +import lit.formats +import lit.util +from lit.llvm import llvm_config +from lit.llvm.subst import FindTool, ToolSubst + +# Configuration file for the 'lit' test runner + +# name: The name of this test suite +config.name = 'TRITON-SHARED' + +config.test_format = lit.formats.ShTest(not llvm_config.use_lit_shell) + +# suffixes: A list of file extensions to treat as test files. +config.suffixes = ['.mlir'] + +# test_source_root: The root path where tests are located. +config.test_source_root = os.path.dirname(__file__) + +# test_exec_root: The root path where tests should be run. +config.test_exec_root = os.path.join(config.triton_obj_root, 'test') + +config.substitutions.append(('%PATH%', config.environment['PATH'])) +config.substitutions.append(('%shlibext', config.llvm_shlib_ext)) + +llvm_config.with_system_environment( + ['HOME', 'INCLUDE', 'LIB', 'TMP', 'TEMP']) + +# llvm_config.use_default_substitutions() + +# excludes: A list of directories to exclude from the testsuite. The 'Inputs' +# subdirectories contain auxiliary inputs for various tests in their parent +# directories. +config.excludes = [ + 'Inputs', + 'Examples', + 'CMakeLists.txt', + 'README.txt', + 'LICENSE.txt'] + +# test_source_root: The root path where tests are located. +config.test_source_root = os.path.dirname(__file__) + +# test_exec_root: The root path where tests should be run. +config.test_exec_root = os.path.join(config.triton_shared_obj_root, 'test') +config.triton_tools_dir = os.path.join(config.triton_shared_obj_root, 'tools/triton-shared-opt') +config.filecheck_dir = os.path.join(config.triton_obj_root, 'bin', 'FileCheck') + +tool_dirs = [ + config.triton_tools_dir, + config.llvm_tools_dir, + config.filecheck_dir] + +# Tweak the PATH to include the tools dir. +for d in tool_dirs: + llvm_config.with_environment('PATH', d, append_path=True) +tools = [ + 'triton-shared-opt', + ToolSubst('%PYTHON', config.python_executable, unresolved='ignore'), +] + +llvm_config.add_tool_substitutions(tools, tool_dirs) + +# TODO: what's this? +llvm_config.with_environment('PYTHONPATH', [ + os.path.join(config.mlir_binary_dir, 'python_packages', 'triton'), +], append_path=True) diff --git a/third_party/wafer/third_party/flir/test/lit.site.cfg.py.in b/third_party/wafer/third_party/flir/test/lit.site.cfg.py.in new file mode 100755 index 00000000..8c00903e --- /dev/null +++ b/third_party/wafer/third_party/flir/test/lit.site.cfg.py.in @@ -0,0 +1,24 @@ +@LIT_SITE_CFG_IN_HEADER@ + +import sys + +config.triton_obj_root = "@TRITON_BINARY_DIR@" +config.triton_shared_obj_root = "@TRITON_SHARED_BINARY_DIR@" +config.llvm_src_root = "@LLVM_SOURCE_DIR@" +config.llvm_obj_root = "@LLVM_BINARY_DIR@" +config.llvm_tools_dir = "@LLVM_TOOLS_DIR@" +config.llvm_lib_dir = "@LLVM_LIBS_DIR@" +config.llvm_shlib_dir = "@SHLIBDIR@" +config.llvm_shlib_ext = "@SHLIBEXT@" +config.llvm_exe_ext = "@EXEEXT@" +config.lit_tools_dir = "@LLVM_LIT_TOOLS_DIR@" +config.mlir_binary_dir = "@MLIR_BINARY_DIR@" +config.python_executable = "@Python3_EXECUTABLE@" +config.enable_bindings_python = @MLIR_ENABLE_BINDINGS_PYTHON@ + + +import lit.llvm +lit.llvm.initialize(lit_config, config) + +# Let the main config do the real work +lit_config.load_config(config, "@TRITON_SHARED_SOURCE_DIR@/test/lit.cfg.py") diff --git a/third_party/wafer/third_party/flir/tools/CMakeLists.txt b/third_party/wafer/third_party/flir/tools/CMakeLists.txt new file mode 100755 index 00000000..3cdf7432 --- /dev/null +++ b/third_party/wafer/third_party/flir/tools/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(triton-shared-opt) diff --git a/third_party/wafer/third_party/flir/tools/RegisterTritonSharedDialects.h b/third_party/wafer/third_party/flir/tools/RegisterTritonSharedDialects.h new file mode 100755 index 00000000..8ee2d46d --- /dev/null +++ b/third_party/wafer/third_party/flir/tools/RegisterTritonSharedDialects.h @@ -0,0 +1,70 @@ +#pragma once +#include "mlir/Dialect/Bufferization/IR/Bufferization.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" +#include "mlir/Dialect/Linalg/Passes.h" +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Ptr/IR/PtrDialect.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "triton-shared/Conversion/StructuredToMemref/Passes.h" +#include "triton-shared/Conversion/ReconcilePtrCasts/Passes.h" +#include "triton/Dialect/Triton/IR/Dialect.h" + +#include "triton/Dialect/Triton/Transforms/Passes.h" + +#include "triton-shared/Conversion/StructuredToMemref/Passes.h" +#include "triton-shared/Conversion/TritonArithToLinalg/Passes.h" +#include "triton-shared/Conversion/TritonPtrToMemref/Passes.h" +#include "triton-shared/Conversion/TritonToLinalg/Passes.h" +#include "triton-shared/Conversion/TritonToLinalgExperimental/Passes.h" +#include "triton-shared/Conversion/TritonToStructured/Passes.h" +#include "triton-shared/Conversion/TritonToUnstructured/Passes.h" +#include "triton-shared/Conversion/UnstructuredToMemref/Passes.h" +#include "triton-shared/Dialect/TPtr/IR/TPtrDialect.h" +#include "triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h" +#include "triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h" + +#include "mlir/InitAllPasses.h" + +namespace mlir { +namespace test { +void registerTestAliasPass(); +void registerTestAlignmentPass(); +void registerTestAllocationPass(); +#ifdef __NVIDIA__ +void registerTestMembarPass(); +#endif +} // namespace test +} // namespace mlir + +inline void registerTritonSharedDialects(mlir::DialectRegistry ®istry) { + mlir::registerAllPasses(); + mlir::registerTritonPasses(); + mlir::registerLinalgPasses(); + mlir::test::registerTestAliasPass(); + mlir::test::registerTestAlignmentPass(); + mlir::test::registerTestAllocationPass(); +#ifdef __NVIDIA__ + mlir::test::registerTestMembarPass(); +#endif + mlir::triton::registerTritonToLinalgPass(); + mlir::triton::registerTritonToLinalgExperimentalPass(); + mlir::triton::registerTritonToStructuredPass(); + mlir::triton::registerTritonPtrToMemref(); + mlir::triton::registerReconcilePtrCasts(); + mlir::triton::registerTritonToPtr(); + mlir::triton::registerUnstructuredToMemref(); + mlir::triton::registerTritonToUnstructuredPasses(); + mlir::triton::registerTritonArithToLinalgPasses(); + mlir::triton::registerStructuredToMemrefPasses(); + + // TODO: register Triton & TritonGPU passes + registry.insert< + mlir::tptr::TPtrDialect, mlir::ptr::PtrDialect, + mlir::ttx::TritonTilingExtDialect, mlir::tts::TritonStructuredDialect, + mlir::triton::TritonDialect, mlir::cf::ControlFlowDialect, + mlir::math::MathDialect, mlir::arith::ArithDialect, mlir::scf::SCFDialect, + mlir::gpu::GPUDialect, mlir::linalg::LinalgDialect, + mlir::func::FuncDialect, mlir::tensor::TensorDialect, + mlir::memref::MemRefDialect, mlir::bufferization::BufferizationDialect>(); +} diff --git a/third_party/wafer/third_party/flir/tools/triton-shared-opt/CMakeLists.txt b/third_party/wafer/third_party/flir/tools/triton-shared-opt/CMakeLists.txt new file mode 100755 index 00000000..21ac037f --- /dev/null +++ b/third_party/wafer/third_party/flir/tools/triton-shared-opt/CMakeLists.txt @@ -0,0 +1,21 @@ +get_property(dialect_libs GLOBAL PROPERTY MLIR_DIALECT_LIBS) +get_property(conversion_libs GLOBAL PROPERTY MLIR_CONVERSION_LIBS) + +add_llvm_executable(triton-shared-opt triton-shared-opt.cpp PARTIAL_SOURCES_INTENDED) + +# TODO: what's this? +llvm_update_compile_flags(triton-shared-opt) +target_link_libraries(triton-shared-opt PRIVATE + TritonTransforms + TritonSharedAnalysis + ${dialect_libs} + ${conversion_libs} + # tests + TritonTestAnalysis + # MLIR core + MLIROptLib + MLIRPass + MLIRTransforms +) + +mlir_check_all_link_libraries(triton-shared-opt) diff --git a/third_party/wafer/third_party/flir/tools/triton-shared-opt/triton-shared-opt.cpp b/third_party/wafer/third_party/flir/tools/triton-shared-opt/triton-shared-opt.cpp new file mode 100755 index 00000000..9a6869c0 --- /dev/null +++ b/third_party/wafer/third_party/flir/tools/triton-shared-opt/triton-shared-opt.cpp @@ -0,0 +1,18 @@ +//===----------------------------------------------------------------------===// +// +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. +// +//===----------------------------------------------------------------------===// + +#include "../RegisterTritonSharedDialects.h" + +#include "mlir/Tools/mlir-opt/MlirOptMain.h" + +int main(int argc, char **argv) { + mlir::DialectRegistry registry; + registerTritonSharedDialects(registry); + + return mlir::asMainReturnCode(mlir::MlirOptMain( + argc, argv, "Triton-Shared test driver\n", registry)); +} diff --git a/third_party/wafer/third_party/flir/triton_shared.cc b/third_party/wafer/third_party/flir/triton_shared.cc new file mode 100755 index 00000000..7688a456 --- /dev/null +++ b/third_party/wafer/third_party/flir/triton_shared.cc @@ -0,0 +1,8 @@ +#include + +namespace py = pybind11; + +// The CPU backend with triton_shared doesn't do compilation from within python +// but rather externally through triton-shared-opt, so we leave this function +// blank. +void init_triton_flir(py::module &&m) {} diff --git a/third_party/wafer/third_party/tle/CMakeLists.txt b/third_party/wafer/third_party/tle/CMakeLists.txt new file mode 100755 index 00000000..08d4c270 --- /dev/null +++ b/third_party/wafer/third_party/tle/CMakeLists.txt @@ -0,0 +1,9 @@ +# Include and generated-file paths for all subdirs +set(TLE_INCLUDE_SOURCE "${CMAKE_CURRENT_SOURCE_DIR}/include") +set(TLE_INCLUDE_BINARY "${CMAKE_CURRENT_BINARY_DIR}/include") + +include_directories(${TLE_INCLUDE_SOURCE}) +include_directories(${TLE_INCLUDE_BINARY}) +add_subdirectory(include) +add_subdirectory(lib) +add_subdirectory(python) diff --git a/third_party/wafer/third_party/tle/REANME.md b/third_party/wafer/third_party/tle/REANME.md new file mode 100755 index 00000000..e69de29b diff --git a/third_party/wafer/third_party/tle/include/CMakeLists.txt b/third_party/wafer/third_party/tle/include/CMakeLists.txt new file mode 100755 index 00000000..16882549 --- /dev/null +++ b/third_party/wafer/third_party/tle/include/CMakeLists.txt @@ -0,0 +1 @@ +add_subdirectory(tle-dsa/Dialect/IR) diff --git a/third_party/wafer/third_party/tle/include/tle-dsa/Conversion/DsaToCore/DsaToCore.h b/third_party/wafer/third_party/tle/include/tle-dsa/Conversion/DsaToCore/DsaToCore.h new file mode 100755 index 00000000..f121aa6d --- /dev/null +++ b/third_party/wafer/third_party/tle/include/tle-dsa/Conversion/DsaToCore/DsaToCore.h @@ -0,0 +1,17 @@ +#ifndef TLE_DSA_CONVERSION_DSATOMCORE_H +#define TLE_DSA_CONVERSION_DSATOMCORE_H + +#include + +namespace mlir { +class Pass; +} // namespace mlir + +namespace mlir::dsa { + +std::unique_ptr createDsaMemoryToCorePass(); +void registerDsaMemoryToCorePass(); + +} // namespace mlir::dsa + +#endif // TLE_DSA_CONVERSION_DSATOMCORE_H diff --git a/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/CMakeLists.txt b/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/CMakeLists.txt new file mode 100755 index 00000000..92f59801 --- /dev/null +++ b/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/CMakeLists.txt @@ -0,0 +1,14 @@ +set(LLVM_TARGET_DEFINITIONS DsaDialect.td) +mlir_tablegen(DsaOpsDialect.h.inc -gen-dialect-decls -dialect=dsa) +mlir_tablegen(DsaOpsDialect.cpp.inc -gen-dialect-defs -dialect=dsa) +add_public_tablegen_target(TleDsaDialectIncGen) + +set(LLVM_TARGET_DEFINITIONS DsaDialect.td) +mlir_tablegen(DsaOpsTypes.h.inc -gen-typedef-decls -typedefs-dialect=dsa) +mlir_tablegen(DsaOpsTypes.cpp.inc -gen-typedef-defs -typedefs-dialect=dsa) +add_public_tablegen_target(TleDsaTypesIncGen) + +set(LLVM_TARGET_DEFINITIONS DsaOps.td) +mlir_tablegen(DsaOps.h.inc -gen-op-decls) +mlir_tablegen(DsaOps.cpp.inc -gen-op-defs) +add_public_tablegen_target(TleDsaOpsIncGen) diff --git a/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.h b/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.h new file mode 100755 index 00000000..4dd49464 --- /dev/null +++ b/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.h @@ -0,0 +1,33 @@ +//===- DsaDialect.h - TLE DSA dialect ---------------------------*- C++ -*-===// +// +// Template dialect for TLE-Struct style DSA extensions. +// +//===----------------------------------------------------------------------===// + +#ifndef TLE_DSA_DIALECT_IR_DSADIALECT_H +#define TLE_DSA_DIALECT_IR_DSADIALECT_H + +#include "mlir/Bytecode/BytecodeOpInterface.h" +#include "mlir/IR/Dialect.h" +#include "mlir/IR/OpDefinition.h" +#include "mlir/Interfaces/DestinationStyleOpInterface.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" + +// DsaOps.td uses TT_Tensor / TT_Ptr / TT_Int type constraints from the +// Triton dialect, so the generated verifiers need these types visible. +#include "triton/Dialect/Triton/IR/Dialect.h" +#include "triton/Dialect/Triton/IR/Types.h" + +#include "tle-dsa/Dialect/IR/DsaOpsDialect.h.inc" + +namespace mlir { +class PatternRewriter; +} // namespace mlir + +#define GET_TYPEDEF_CLASSES +#include "tle-dsa/Dialect/IR/DsaOpsTypes.h.inc" + +#define GET_OP_CLASSES +#include "tle-dsa/Dialect/IR/DsaOps.h.inc" + +#endif // TLE_DSA_DIALECT_IR_DSADIALECT_H diff --git a/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.td b/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.td new file mode 100755 index 00000000..b3144704 --- /dev/null +++ b/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.td @@ -0,0 +1,68 @@ +//===- DsaDialect.td - TLE DSA dialect --------------------*- tablegen -*-===// +// +// Template dialect for TLE-Struct style DSA extensions. +// +//===----------------------------------------------------------------------===// + +#ifndef TLE_DSA_DIALECT +#define TLE_DSA_DIALECT + +include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/BuiltinTypeInterfaces.td" +include "mlir/IR/OpBase.td" + +//===----------------------------------------------------------------------===// +// Dialect definition. +//===----------------------------------------------------------------------===// + +def Dsa_Dialect : Dialect { + let name = "dsa"; + let summary = "TLE DSA struct dialect (template)"; + let cppNamespace = "::mlir::dsa"; + let useDefaultTypePrinterParser = 1; + let extraClassDeclaration = [{ + void registerTypes(); + }]; +} + +//===----------------------------------------------------------------------===// +// Type definitions. +//===----------------------------------------------------------------------===// + +class Dsa_Type traits = []> + : TypeDef { + let mnemonic = typeMnemonic; +} + +def DsaBufferType : Dsa_Type<"Buffer", "buffer"> { + let summary = "Opaque local buffer handle"; + let description = [{ + This is a minimal, parameterized buffer handle type for DSA extensions. + The backend defines the lowering semantics (e.g. mapping to SRAM/UB/NZ, etc). + }]; + + let parameters = (ins + "Type":$elementType, + OptionalParameter<"Attribute">:$memorySpace + ); + + let builders = [ + TypeBuilder<(ins + "Type":$elementType, + CArg<"Attribute", "nullptr">:$memorySpace), [{ + return Base::get($_ctxt, elementType, memorySpace); + }]> + ]; + let skipDefaultBuilders = 1; + + let assemblyFormat = "`<` $elementType (`,` $memorySpace^)? `>`"; +} + +//===----------------------------------------------------------------------===// +// Base op definition. +//===----------------------------------------------------------------------===// + +class Dsa_Op traits = []> : + Op; + +#endif // TLE_DSA_DIALECT diff --git a/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaOps.td b/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaOps.td new file mode 100755 index 00000000..7c5c4e14 --- /dev/null +++ b/third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaOps.td @@ -0,0 +1,200 @@ +//===- DsaOps.td - TLE DSA dialect ops --------------------*- tablegen -*-===// +// +// Template ops for TLE-Struct style DSA extensions. +// +//===----------------------------------------------------------------------===// + +#ifndef TLE_DSA_OPS +#define TLE_DSA_OPS + +include "tle-dsa/Dialect/IR/DsaDialect.td" +include "mlir/Dialect/LLVMIR/LLVMTypes.td" +include "mlir/Interfaces/SideEffectInterfaces.td" +include "mlir/IR/CommonTypeConstraints.td" +include "triton/Dialect/Triton/IR/TritonTypes.td" + +//===----------------------------------------------------------------------===// +// dsa.alloc +//===----------------------------------------------------------------------===// + +def Dsa_AllocOp : Dsa_Op<"alloc", [MemoryEffects<[MemAlloc]>]> { + let summary = "Allocate a DSA local buffer"; + let arguments = (ins DenseI64ArrayAttr:$shape); + let results = (outs AnyRankedOrUnrankedMemRef:$result); + let assemblyFormat = [{ + $shape attr-dict `:` type($result) + }]; +} + +//===----------------------------------------------------------------------===// +// dsa.copy +//===----------------------------------------------------------------------===// + +def Dsa_CopyOp : Dsa_Op<"copy", + [MemoryEffects<[MemRead, MemWrite]>]> { + let summary = "Copy between DSA local buffers"; + let arguments = (ins AnyRankedOrUnrankedMemRef:$src, AnyRankedOrUnrankedMemRef:$dst); + let results = (outs); + let assemblyFormat = [{ + $src `,` $dst attr-dict `:` type($src) `,` type($dst) + }]; +} + +def DsaLocalPointerResultType : AnyTypeOf<[TT_Tensor, TT_Ptr]>; +def DsaLocalPointerIndexType : AnyTypeOf<[TT_Tensor, TT_Int]>; +def DsaRemotePointerType : AnyTypeOf<[TT_Tensor, TT_Ptr]>; +def DsaRemoteShardIdType : AnyTypeOf<[TT_Tensor, TT_Int]>; + +//===----------------------------------------------------------------------===// +// dsa.local_pointers / dsa.remote_pointers / dsa.distributed_barrier +// These are DSA-side structural counterparts of the official TLE distributed +// and pointer-building ops. They let the frontend move toward the same +// structured programming model without binding the design to GPU dialect names. +//===----------------------------------------------------------------------===// + +def Dsa_LocalPointersOp : Dsa_Op<"local_pointers", [Pure]> { + let arguments = (ins AnyRankedOrUnrankedMemRef:$src, + Variadic:$indices); + let results = (outs DsaLocalPointerResultType:$result); +} + +def Dsa_DistributedBarrierOp : Dsa_Op<"distributed_barrier", + [MemoryEffects<[MemRead, MemWrite]>]> { + let arguments = (ins + OptionalAttr:$group_kind, + OptionalAttr:$group_rank, + OptionalAttr:$group_shape, + OptionalAttr:$group_axes, + OptionalAttr:$group_mask + ); + let assemblyFormat = "attr-dict"; +} + +def Dsa_RemotePointersOp : Dsa_Op<"remote_pointers", [Pure]> { + let arguments = (ins DsaRemotePointerType:$src, DsaRemoteShardIdType:$shard_id); + let results = (outs DsaRemotePointerType:$result); +} + +class Dsa_BinaryOp : Dsa_Op]> { + let arguments = (ins AnyRankedOrUnrankedMemRef:$lhs, + AnyRankedOrUnrankedMemRef:$rhs, + AnyRankedOrUnrankedMemRef:$out); + let results = (outs); + let assemblyFormat = [{ + $lhs `,` $rhs `,` $out attr-dict `:` type($lhs) `,` type($rhs) `,` type($out) + }]; +} + +def Dsa_AddOp : Dsa_BinaryOp<"add"> { + let summary = "elementwise add on local memrefs (out = lhs + rhs)"; +} +def Dsa_SubOp : Dsa_BinaryOp<"sub"> { + let summary = "elementwise subtract on local memrefs (out = lhs - rhs)"; +} +def Dsa_MulOp : Dsa_BinaryOp<"mul"> { + let summary = "elementwise multiply on local memrefs (out = lhs * rhs)"; +} +def Dsa_MaximumOp : Dsa_BinaryOp<"maximum"> { + let summary = "elementwise maximum on local memrefs (out = max(lhs, rhs))"; +} +def Dsa_MinimumOp : Dsa_BinaryOp<"minimum"> { + let summary = "elementwise minimum on local memrefs (out = min(lhs, rhs))"; +} +def Dsa_DivOp : Dsa_BinaryOp<"div"> { + let summary = "elementwise divide on local memrefs (out = lhs / rhs)"; +} + +def Dsa_ToTensorOp : Dsa_Op<"to_tensor", [MemoryEffects<[MemRead]>]> { + let summary = "Convert a DSA local buffer (memref) to a tl.tensor view"; + let arguments = (ins AnyRankedOrUnrankedMemRef:$src, BoolAttr:$writable); + let results = (outs TT_Tensor:$result); + let assemblyFormat = [{ + $src attr-dict `:` type($src) `->` type($result) + }]; +} + +def Dsa_ToBufferOp : Dsa_Op<"to_buffer", [MemoryEffects<[MemRead, MemWrite]>]> { + let summary = "Copy a tl.tensor value into an existing DSA local buffer (memref)"; + let arguments = (ins TT_Tensor:$src, AnyRankedOrUnrankedMemRef:$dst); + let results = (outs); + let assemblyFormat = [{ + $src `,` $dst attr-dict `:` type($src) `,` type($dst) + }]; +} + +def DsaSliceOffsetType : AnyTypeOf<[TT_Tensor, TT_Int]>; + +def Dsa_ExtractSliceOp : Dsa_Op<"extract_slice", [Pure]> { + let summary = "Extract a strided slice from a tl.tensor (mixed static/dynamic offsets, static sizes/strides)"; + let arguments = (ins TT_Tensor:$src, + Variadic:$offsets, + DenseI64ArrayAttr:$static_offsets, + DenseI64ArrayAttr:$sizes, + DenseI64ArrayAttr:$strides); + let results = (outs TT_Tensor:$result); +} + +def Dsa_InsertSliceOp : Dsa_Op<"insert_slice", [Pure]> { + let summary = "Insert a strided slice into a tl.tensor (mixed static/dynamic offsets, static sizes/strides)"; + let arguments = (ins TT_Tensor:$src, + TT_Tensor:$tile, + Variadic:$offsets, + DenseI64ArrayAttr:$static_offsets, + DenseI64ArrayAttr:$sizes, + DenseI64ArrayAttr:$strides); + let results = (outs TT_Tensor:$result); +} + + +def DsaRandGenSeedType : RankedTensorOf<[I64]>; +def DsaRandGenOutType : RankedTensorOf<[I64]>; + +def Dsa_RandGenOp : Dsa_Op<"randgen", + [MemoryEffects<[MemRead, MemWrite]>]> { + let summary = "randgen for wafer"; + let description = [{ + Emits a single Wafer peri RandGen call. `seed0`/`seed1` are length-16 i64 + seed vectors; `byte_count` is the random output size in bytes and must be + a multiple of 128. Results are `(out, seed0_out, seed1_out)`. + }]; + let arguments = (ins + DsaRandGenSeedType:$seed0, + DsaRandGenSeedType:$seed1, + I32Attr:$byte_count, + I16Attr:$fmt + ); + let results = (outs + DsaRandGenOutType:$out, + DsaRandGenSeedType:$seed0_out, + DsaRandGenSeedType:$seed1_out + ); + let assemblyFormat = [{ + $seed0 `,` $seed1 attr-dict `:` type($seed0) `,` type($seed1) + `->` type($out) `,` type($seed0_out) `,` type($seed1_out) + }]; +} + +//===----------------------------------------------------------------------===// +// dsa.bitcast +//===----------------------------------------------------------------------===// + +def Dsa_BitcastOp : Dsa_Op<"bitcast", [Pure]> { + let summary = "Same-nbytes type/shape reinterpret (vendor-neutral)"; + let description = [{ + Reinterpret `src` as `result` when both tensors have the same total bit + size. Element type and/or shape may change (e.g. `tensor` → + `tensor<(2N)x i32>`). This is a pure view: no data movement. + + Backends lower this to a zero-cost alias when possible (e.g. Tsingmicro + `mk.bitcast`); a portable fallback is `tensor.bitcast`. + }]; + let arguments = (ins AnyRankedTensor:$src); + let results = (outs AnyRankedTensor:$result); + let assemblyFormat = [{ + $src attr-dict `:` type($src) `->` type($result) + }]; + let hasVerifier = 1; +} + +#endif // TLE_DSA_OPS diff --git a/third_party/wafer/third_party/tle/lib/CMakeLists.txt b/third_party/wafer/third_party/tle/lib/CMakeLists.txt new file mode 100755 index 00000000..cf09d435 --- /dev/null +++ b/third_party/wafer/third_party/tle/lib/CMakeLists.txt @@ -0,0 +1,4 @@ +add_subdirectory(Dialect/IR) +if(NOT DEFINED WAFER_TLE_BUILD_CONVERSIONS OR WAFER_TLE_BUILD_CONVERSIONS) + add_subdirectory(Conversion/DsaToCore) +endif() diff --git a/third_party/wafer/third_party/tle/lib/Conversion/DsaToCore/CMakeLists.txt b/third_party/wafer/third_party/tle/lib/Conversion/DsaToCore/CMakeLists.txt new file mode 100755 index 00000000..66763862 --- /dev/null +++ b/third_party/wafer/third_party/tle/lib/Conversion/DsaToCore/CMakeLists.txt @@ -0,0 +1,17 @@ +add_triton_library(TleDsaToCore + DsaToCore.cpp + + LINK_LIBS PUBLIC + MLIRIR + MLIRPass + MLIRTransformUtils + MLIRRewrite + TleDsaIR + TritonIR +) + +target_include_directories(TleDsaToCore + PUBLIC + ${TLE_INCLUDE_SOURCE} + ${TLE_INCLUDE_BINARY} +) diff --git a/third_party/wafer/third_party/tle/lib/Conversion/DsaToCore/DsaToCore.cpp b/third_party/wafer/third_party/tle/lib/Conversion/DsaToCore/DsaToCore.cpp new file mode 100755 index 00000000..4837e7f7 --- /dev/null +++ b/third_party/wafer/third_party/tle/lib/Conversion/DsaToCore/DsaToCore.cpp @@ -0,0 +1,72 @@ +//===- DsaToCore.cpp - Lower dsa memory ops to core dialects ----*- C++ -*-===// +// +// Lower dsa.alloc/copy into standard memref ops +// ops before one-shot-bufferize. +// +//===----------------------------------------------------------------------===// + +#include "tle-dsa/Conversion/DsaToCore/DsaToCore.h" + +#include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" +#include "tle-dsa/Dialect/IR/DsaDialect.h" + +using namespace mlir; + +namespace { + +struct DsaAllocToMemRefPattern : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(mlir::dsa::AllocOp op, + PatternRewriter &rewriter) const override { + auto memrefTy = dyn_cast(op.getResult().getType()); + if (!memrefTy) + return failure(); + // wafer-memref-to-llvm expects integer/default memref address spaces. + // Canonicalize away non-integer memory-space attrs (e.g. "local"). + if (Attribute ms = memrefTy.getMemorySpace(); ms && !isa(ms)) { + memrefTy = MemRefType::get(memrefTy.getShape(), memrefTy.getElementType(), + memrefTy.getLayout()); + } + rewriter.replaceOpWithNewOp(op, memrefTy); + return success(); + } +}; + +struct DsaCopyToMemRefPattern : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(mlir::dsa::CopyOp op, + PatternRewriter &rewriter) const override { + rewriter.create(op.getLoc(), op.getSrc(), op.getDst()); + rewriter.eraseOp(op); + return success(); + } +}; + +struct DsaMemoryToCorePass + : public PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(DsaMemoryToCorePass) + StringRef getArgument() const final { return "dsa-memory-to-core"; } + StringRef getDescription() const final { + return "Lower dsa.alloc/copy to memref"; + } + void runOnOperation() override { + RewritePatternSet patterns(&getContext()); + patterns.add( + &getContext()); + if (failed(applyPatternsGreedily(getOperation(), std::move(patterns)))) + signalPassFailure(); + } +}; + +} // namespace + +namespace mlir::dsa { +std::unique_ptr createDsaMemoryToCorePass() { + return std::make_unique(); +} + +void registerDsaMemoryToCorePass() { PassRegistration(); } +} // namespace mlir::dsa diff --git a/third_party/wafer/third_party/tle/lib/Dialect/IR/CMakeLists.txt b/third_party/wafer/third_party/tle/lib/Dialect/IR/CMakeLists.txt new file mode 100755 index 00000000..870dc009 --- /dev/null +++ b/third_party/wafer/third_party/tle/lib/Dialect/IR/CMakeLists.txt @@ -0,0 +1,19 @@ +add_triton_library(TleDsaIR + DsaDialect.cpp + + DEPENDS + TleDsaDialectIncGen + TleDsaTypesIncGen + TleDsaOpsIncGen + + LINK_LIBS PUBLIC + MLIRIR + MLIRSupport + TritonIR +) + +target_include_directories(TleDsaIR + PUBLIC + ${TLE_INCLUDE_SOURCE} + ${TLE_INCLUDE_BINARY} +) diff --git a/third_party/wafer/third_party/tle/lib/Dialect/IR/DsaDialect.cpp b/third_party/wafer/third_party/tle/lib/Dialect/IR/DsaDialect.cpp new file mode 100755 index 00000000..55e99666 --- /dev/null +++ b/third_party/wafer/third_party/tle/lib/Dialect/IR/DsaDialect.cpp @@ -0,0 +1,56 @@ +//===- DsaDialect.cpp - TLE DSA dialect -------------------------*- C++ -*-===// +// +// Template dialect for TLE-Struct style DSA extensions. +// +//===----------------------------------------------------------------------===// + +#include "tle-dsa/Dialect/IR/DsaDialect.h" + +#include "mlir/IR/DialectImplementation.h" +#include "llvm/ADT/TypeSwitch.h" + +using namespace mlir; + +namespace mlir::dsa { + +void DsaDialect::initialize() { + addOperations< +#define GET_OP_LIST +#include "tle-dsa/Dialect/IR/DsaOps.cpp.inc" + >(); + registerTypes(); +} + +void DsaDialect::registerTypes() { + addTypes< +#define GET_TYPEDEF_LIST +#include "tle-dsa/Dialect/IR/DsaOpsTypes.cpp.inc" + >(); +} + +} // namespace mlir::dsa + +#include "tle-dsa/Dialect/IR/DsaOpsDialect.cpp.inc" + +#define GET_OP_CLASSES +#include "tle-dsa/Dialect/IR/DsaOps.cpp.inc" + +#define GET_TYPEDEF_CLASSES +#include "tle-dsa/Dialect/IR/DsaOpsTypes.cpp.inc" + +LogicalResult mlir::dsa::BitcastOp::verify() { + auto srcTy = dyn_cast(getSrc().getType()); + auto dstTy = dyn_cast(getResult().getType()); + if (!srcTy || !dstTy) + return emitOpError("expects ranked tensor src/result"); + auto srcElem = srcTy.getElementType(); + auto dstElem = dstTy.getElementType(); + if (!srcElem.isIntOrFloat() || !dstElem.isIntOrFloat()) + return emitOpError("element types must be int or float"); + int64_t srcBits = srcTy.getNumElements() * srcElem.getIntOrFloatBitWidth(); + int64_t dstBits = dstTy.getNumElements() * dstElem.getIntOrFloatBitWidth(); + if (srcBits != dstBits) + return emitOpError("src and result must have the same total bit size, got ") + << srcBits << " vs " << dstBits; + return success(); +} diff --git a/third_party/wafer/third_party/tle/python/CMakeLists.txt b/third_party/wafer/third_party/tle/python/CMakeLists.txt new file mode 100755 index 00000000..ea37f185 --- /dev/null +++ b/third_party/wafer/third_party/tle/python/CMakeLists.txt @@ -0,0 +1,7 @@ +if(TRITON_BUILD_PYTHON_MODULE) + add_triton_plugin(TritonTleDsaTemplate + ${CMAKE_CURRENT_SOURCE_DIR}/triton_tle_dsa.cc + LINK_LIBS TleDsaIR + ) + target_link_libraries(TritonTleDsaTemplate PRIVATE Python3::Module pybind11::headers) +endif() diff --git a/third_party/wafer/third_party/tle/python/ir.h b/third_party/wafer/third_party/tle/python/ir.h new file mode 100755 index 00000000..cdb5257c --- /dev/null +++ b/third_party/wafer/third_party/tle/python/ir.h @@ -0,0 +1,91 @@ +#pragma once + +#include "mlir/IR/Builders.h" +#include "triton/Tools/Sys/GetEnv.hpp" + +#include +#include + +// A custom op builder that keeps track of the last location. +class TritonOpBuilder { +public: + TritonOpBuilder(mlir::MLIRContext *context) { + builder = std::make_unique(context); + lastLoc = std::make_unique(builder->getUnknownLoc()); + } + + mlir::OpBuilder &getBuilder() { return *builder; } + mlir::MLIRContext *getContext() { return builder->getContext(); } + + bool isLineInfoEnabled() { return lineInfoEnabled; } + + void setLastLoc(mlir::Location loc) { + if (lineInfoEnabled) + lastLoc = std::make_unique(loc); + } + + void setLastLoc(const std::string &fileName, int line, int column) { + auto context = builder->getContext(); + setLastLoc(mlir::FileLineColLoc::get(context, fileName, line, column)); + } + + mlir::Location getLastLoc() { + assert(lastLoc); + return *lastLoc; + } + + void setInsertionPointToStart(mlir::Block &block) { + if (!block.empty()) + setLastLoc(block.begin()->getLoc()); + else + setLastLoc(builder->getUnknownLoc()); + builder->setInsertionPointToStart(&block); + } + + void setInsertionPointToEnd(mlir::Block &block) { + if (!block.empty()) + setLastLoc(block.back().getLoc()); + else + setLastLoc(builder->getUnknownLoc()); + builder->setInsertionPointToEnd(&block); + } + + void setInsertionPointAfter(mlir::Operation &op) { + setLastLoc(op.getLoc()); + builder->setInsertionPointAfter(&op); + } + + void restoreInsertionPoint(mlir::OpBuilder::InsertPoint pt) { + if (pt.isSet() && pt.getPoint() != pt.getBlock()->end()) + setLastLoc(pt.getPoint()->getLoc()); + else + setLastLoc(builder->getUnknownLoc()); + builder->restoreInsertionPoint(pt); + } + + template OpTy create(Args &&...args) { + auto loc = getLastLoc(); + return builder->create(loc, std::forward(args)...); + } + + template + std::enable_if_t(), + mlir::Value> + createOrFold(Args &&...args) { + auto loc = getLastLoc(); + return builder->createOrFold(loc, std::forward(args)...); + } + + template + std::enable_if_t(), OpTy> + createOrFold(Args &&...args) { + auto loc = getLastLoc(); + return builder->createOrFold(loc, std::forward(args)...); + } + +private: + std::unique_ptr builder; + std::unique_ptr lastLoc; + bool lineInfoEnabled = + !mlir::triton::tools::getBoolEnv("TRITON_DISABLE_LINE_INFO"); +}; diff --git a/third_party/wafer/third_party/tle/python/triton_tle_dsa.cc b/third_party/wafer/third_party/tle/python/triton_tle_dsa.cc new file mode 100755 index 00000000..14c33495 --- /dev/null +++ b/third_party/wafer/third_party/tle/python/triton_tle_dsa.cc @@ -0,0 +1,241 @@ +//===- triton_tle_dsa.cc - TLE DSA builder injection -------------*- C++ +//-*-===// +// +// Template pybind that injects DSA dialect ops into TritonOpBuilder. +// +//===----------------------------------------------------------------------===// + +#include +#include + +#include "mlir/IR/Attributes.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/MLIRContext.h" +#include "mlir/IR/Value.h" +#include "llvm/ADT/SmallVector.h" + +#include "ir.h" +#include "tle-dsa/Dialect/IR/DsaDialect.h" + +namespace py = pybind11; +using namespace mlir; + +namespace dsa = mlir::dsa; + +template +static void defBinaryOp(py::module &builderCls, + const char *name) { + builderCls.def( + name, [](TritonOpBuilder &self, Value lhs, Value rhs, Value out) -> void { + self.getContext()->getOrLoadDialect(); + self.getBuilder().create(self.getLastLoc(), lhs, rhs, out); + }); +} + +// Cast a Python list of dynamic-offset values to a SmallVector. +static llvm::SmallVector castDynOffsets(py::list dynOffsets) { + llvm::SmallVector dyn; + dyn.reserve(py::len(dynOffsets)); + for (py::handle arg : dynOffsets) + dyn.push_back(py::cast(arg)); + return dyn; +} + + +static Value createDsaAlloc(TritonOpBuilder &self, py::object shapeObj, + py::object elementTyObj) { + self.getContext()->getOrLoadDialect(); + auto &b = self.getBuilder(); + std::vector dims; + if (py::isinstance(shapeObj)) { + dims.push_back(py::cast(shapeObj)); + } else { + py::iterable shape = py::reinterpret_borrow(shapeObj); + dims.reserve(py::len(shape)); + for (py::handle dim : shape) + dims.push_back(py::cast(dim)); + } + auto shapeAttr = DenseI64ArrayAttr::get(b.getContext(), dims); + Type elementTy = py::cast(elementTyObj); + auto bufTy = MemRefType::get(dims, elementTy); + auto op = + self.getBuilder().create(self.getLastLoc(), bufTy, shapeAttr); + return op.getResult(); +} + +static void createDsaCopy(TritonOpBuilder &self, Value src, Value dst) { + self.getContext()->getOrLoadDialect(); + self.getBuilder().create(self.getLastLoc(), src, dst); +} + +static OpState createDsaLocalPointers(TritonOpBuilder &self, Type resultTy, + Value src, py::args args) { + self.getContext()->getOrLoadDialect(); + llvm::SmallVector indices; + indices.reserve(args.size()); + for (const auto &arg : args) + indices.push_back(py::cast(arg)); + return self.create(resultTy, src, indices); +} + +static OpState createDsaRemotePointers(TritonOpBuilder &self, Type resultTy, + Value src, Value shardId) { + self.getContext()->getOrLoadDialect(); + return self.create(resultTy, src, shardId); +} + +static void createDsaDistributedBarrier(TritonOpBuilder &self, + const std::string &groupKind, + const std::vector &groupShape, + const std::vector &groupAxes, + const std::vector &groupMask) { + self.getContext()->getOrLoadDialect(); + auto &builder = self.getBuilder(); + auto *ctx = builder.getContext(); + StringAttr kindAttr; + IntegerAttr rankAttr; + DenseI32ArrayAttr shapeAttr; + DenseI32ArrayAttr axesAttr; + DenseI32ArrayAttr maskAttr; + + if (!groupKind.empty()) { + kindAttr = builder.getStringAttr(groupKind); + rankAttr = + builder.getI32IntegerAttr(static_cast(groupShape.size())); + shapeAttr = DenseI32ArrayAttr::get(ctx, groupShape); + axesAttr = DenseI32ArrayAttr::get(ctx, groupAxes); + if (!groupMask.empty()) + maskAttr = DenseI32ArrayAttr::get(ctx, groupMask); + } + + self.create(kindAttr, rankAttr, shapeAttr, + axesAttr, maskAttr); +} + +static void init_triton_tle_ir(py::module m) { + (void)m; + auto core_ir = py::module::import("triton._C.libtriton.ir"); + auto builder_cls = core_ir.attr("builder"); + + builder_cls.attr("create_dsa_alloc") = py::cpp_function( + &createDsaAlloc, + py::is_method(builder_cls)); + + builder_cls.attr("create_dsa_copy") = py::cpp_function( + &createDsaCopy, + py::is_method(builder_cls)); + + builder_cls.attr("create_dsa_local_pointers") = py::cpp_function( + &createDsaLocalPointers, + py::is_method(builder_cls)); + + builder_cls.attr("create_dsa_remote_pointers") = py::cpp_function( + &createDsaRemotePointers, + py::is_method(builder_cls)); + + builder_cls.attr("create_dsa_distributed_barrier") = py::cpp_function( + &createDsaDistributedBarrier, + py::is_method(builder_cls)); +} + +// void init_triton_tle(py::module &&m, const char *submodule_name = nullptr) { +// if (submodule_name && *submodule_name != '\0') +// m = m.def_submodule(submodule_name); +// py::module local_m = std::move(m); +// local_m.def("load_dialects", [](mlir::MLIRContext &context) { +// context.getOrLoadDialect(); +// }); +// init_triton_tle_ir(std::move(local_m)); +// } + +void init_triton_tle(py::module &&m) { + py::module local_m = std::move(m); + + local_m.def("load_dialects", [](mlir::MLIRContext &context) { + context.getOrLoadDialect(); + }); + local_m + .def("create_dsa_to_tensor", + [](TritonOpBuilder &self, Type resultTy, Value src, + bool writable) -> Value { + self.getContext()->getOrLoadDialect(); + auto &b = self.getBuilder(); + auto writableAttr = b.getBoolAttr(writable); + auto op = + self.create(resultTy, src, writableAttr); + return op.getResult(); + }) + .def("create_dsa_to_buffer", + [](TritonOpBuilder &self, Value src, Value dst) -> void { + self.getContext()->getOrLoadDialect(); + self.create(src, dst); + }); + + local_m + .def( + "create_dsa_extract_slice", + [](TritonOpBuilder &self, Type resultTy, Value src, + const std::vector &staticOffsets, py::list dynOffsets, + const std::vector &sizes, + const std::vector &strides) -> Value { + self.getContext()->getOrLoadDialect(); + auto &builder = self.getBuilder(); + auto *ctx = builder.getContext(); + auto dyn = castDynOffsets(dynOffsets); + auto op = self.create( + resultTy, src, dyn, DenseI64ArrayAttr::get(ctx, staticOffsets), + DenseI64ArrayAttr::get(ctx, sizes), + DenseI64ArrayAttr::get(ctx, strides)); + return op.getResult(); + }) + .def( + "create_dsa_insert_slice", + [](TritonOpBuilder &self, Type resultTy, Value src, Value tile, + const std::vector &staticOffsets, py::list dynOffsets, + const std::vector &sizes, + const std::vector &strides) -> Value { + self.getContext()->getOrLoadDialect(); + auto &builder = self.getBuilder(); + auto *ctx = builder.getContext(); + auto dyn = castDynOffsets(dynOffsets); + auto op = self.create( + resultTy, src, tile, dyn, + DenseI64ArrayAttr::get(ctx, staticOffsets), + DenseI64ArrayAttr::get(ctx, sizes), + DenseI64ArrayAttr::get(ctx, strides)); + return op.getResult(); + }); + // Three-operand binary arithmetic (out = lhs OP rhs). + defBinaryOp(local_m, "create_dsa_add"); + defBinaryOp(local_m, "create_dsa_sub"); + defBinaryOp(local_m, "create_dsa_mul"); + defBinaryOp(local_m, "create_dsa_maximum"); + defBinaryOp(local_m, "create_dsa_minimum"); + defBinaryOp(local_m, "create_dsa_div"); + local_m + .def("create_dsa_randgen", + [](TritonOpBuilder &self, Type outTy, Type seed0OutTy, + Type seed1OutTy, Value seed0, Value seed1, int32_t byteCount, + int16_t fmt) -> OpState { + self.getContext()->getOrLoadDialect(); + auto &builder = self.getBuilder(); + return builder.create( + self.getLastLoc(), TypeRange{outTy, seed0OutTy, seed1OutTy}, + seed0, seed1, builder.getI32IntegerAttr(byteCount), + builder.getI16IntegerAttr(fmt)); + }) + // Vendor-neutral same-nbytes type/shape reinterpret (e.g. + // i64[N]→i32[2N]). Backends lower this (Tsingmicro: mk.bitcast alias; + // others: tensor.bitcast). + .def("create_dsa_bitcast", + [](TritonOpBuilder &self, Type dstTy, Value src) -> Value { + self.getContext()->getOrLoadDialect(); + return self.create(dstTy, src); + }); + local_m.def("create_dsa_alloc", &createDsaAlloc); + local_m.def("create_dsa_copy", &createDsaCopy); + local_m.def("create_dsa_local_pointers", &createDsaLocalPointers); + local_m.def("create_dsa_remote_pointers", &createDsaRemotePointers); + local_m.def("create_dsa_distributed_barrier", &createDsaDistributedBarrier); +}