From acf2a2e1aa8335342764ca33a23a35feb200cb3c Mon Sep 17 00:00:00 2001 From: zhuwenrui Date: Sun, 20 Sep 2026 14:40:45 +0800 Subject: [PATCH] Integrate Wafer compiler, runtime and local build workflow Introduce the Wafer frontend and lowering passes, Kuiper launcher and CRT, isolated build and wheel packaging, local dependency configuration, and native regression coverage. Preserve backend isolation and bundled third-party attribution. Group build tools under scripts/wafer and acceptance programs under test/wafer. Dependency acquisition is outside the Wafer setup interface; vendor SDK ABI and internal filenames remain unchanged. --- .github/workflows/wafer-isolation.yml | 46 + backend/compiler.py | 33 +- backend/driver.py | 59 +- backend/utils.py | 7 + backend/wafer.py | 619 +++ backend/wafer_cache.py | 28 + backend/wafer_runtime.py | 486 ++ scripts/wafer/apply_triton_profile.py | 130 + scripts/wafer/apply_wafer_triton_patches.sh | 4 + scripts/wafer/audit_wafer_elf.py | 110 + scripts/wafer/build_wafer_isolated.py | 159 + scripts/wafer/compile_wafer.sh | 4 + scripts/wafer/init_wafer_env.sh | 73 + scripts/wafer/install_wafer.sh | 88 + scripts/wafer/inventory_wafer_tests.py | 134 + scripts/wafer/migrate_wafer_env.sh | 26 + scripts/wafer/package_wafer.py | 78 + scripts/wafer/run_wafer_example_suite.py | 131 + scripts/wafer/setup_llvm22_env.sh | 109 + scripts/wafer/setup_wafer_env.sh | 44 + scripts/wafer/wafer_artifacts.py | 33 + scripts/wafer/wafer_manifest.py | 94 + scripts/wafer/wafer_pytest.py | 123 + setup_on_wafer.py | 42 + test/wafer/conftest.py | 68 + test/wafer/ir/interfaces-argmax2d.mlir | 35 + test/wafer/ir/interfaces-flip.mlir | 48 + test/wafer/ir/interfaces-sort.mlir | 125 + test/wafer/ir/pointer-state-modulo.mlir | 81 + test/wafer/ir/pointer-state-nested_loops.mlir | 103 + test/wafer/ir/pointer-state-scalar_store.mlir | 30 + .../pointer-state-tensor_index_iterargs.mlir | 48 + test/wafer/native_math/__init__.py | 1 + test/wafer/native_math/conftest.py | 7 + test/wafer/native_math/test_log1p.py | 70 + test/wafer/native_math/test_multi_return.py | 338 ++ test/wafer/native_math/test_relu.py | 78 + test/wafer/native_math/test_unary.py | 98 + test/wafer/ops/conftest.py | 13 + test/wafer/ops/test_2d_permute.py | 59 + test/wafer/ops/test_3Dgrid.py | 118 + test/wafer/ops/test_abs.py | 65 + test/wafer/ops/test_abs_2.py | 70 + test/wafer/ops/test_add.py | 97 + test/wafer/ops/test_add_multi_return.py | 99 + test/wafer/ops/test_advance.py | 231 + test/wafer/ops/test_and.py | 67 + test/wafer/ops/test_arange.py | 159 + test/wafer/ops/test_associative_scan.py | 222 + .../ops/test_associative_scan_multi_input.py | 143 + test/wafer/ops/test_block_ptr.py | 93 + test/wafer/ops/test_broadcast_op.py | 60 + test/wafer/ops/test_cat_dim.py | 123 + test/wafer/ops/test_cdiv.py | 66 + test/wafer/ops/test_ceil.py | 68 + test/wafer/ops/test_clamp.py | 64 + test/wafer/ops/test_common.py | 247 + test/wafer/ops/test_conv.py | 257 + test/wafer/ops/test_cos.py | 102 + test/wafer/ops/test_cos_2.py | 64 + test/wafer/ops/test_count_dim0.py | 214 + test/wafer/ops/test_count_dim1.py | 222 + test/wafer/ops/test_cumprod.py | 108 + test/wafer/ops/test_cumsum.py | 95 + test/wafer/ops/test_debug_barrier.py | 67 + test/wafer/ops/test_device_print.py | 155 + test/wafer/ops/test_div.py | 71 + test/wafer/ops/test_elementwise_ceil.py | 91 + test/wafer/ops/test_elementwise_clip.py | 89 + test/wafer/ops/test_elementwise_f2i.py | 133 + test/wafer/ops/test_elementwise_floor.py | 91 + test/wafer/ops/test_elementwise_i2f.py | 128 + test/wafer/ops/test_elementwise_round.py | 95 + test/wafer/ops/test_eq.py | 86 + test/wafer/ops/test_eq_2.py | 66 + test/wafer/ops/test_exp.py | 63 + test/wafer/ops/test_exp2.py | 64 + test/wafer/ops/test_exp_.py | 106 + test/wafer/ops/test_expand_dims.py | 76 + test/wafer/ops/test_extract_slice.py | 44 + test/wafer/ops/test_fdiv.py | 69 + test/wafer/ops/test_floor.py | 66 + test/wafer/ops/test_floordiv.py | 75 + test/wafer/ops/test_full.py | 97 + test/wafer/ops/test_ge.py | 86 + test/wafer/ops/test_ge_2.py | 69 + test/wafer/ops/test_gelu.py | 102 + test/wafer/ops/test_gt.py | 67 + test/wafer/ops/test_hd_permute.py | 65 + test/wafer/ops/test_if_tensor.py | 61 + test/wafer/ops/test_insert_slice.py | 60 + test/wafer/ops/test_interleave.py | 89 + test/wafer/ops/test_invert.py | 72 + test/wafer/ops/test_join.py | 299 ++ test/wafer/ops/test_lanzcos.py | 284 ++ .../ops/test_launcher_empty_signature.py | 37 + test/wafer/ops/test_layernorm.py | 146 + test/wafer/ops/test_ldst.py | 650 +++ test/wafer/ops/test_le.py | 86 + test/wafer/ops/test_load.py | 158 + test/wafer/ops/test_load_store.py | 242 + test/wafer/ops/test_log.py | 106 + test/wafer/ops/test_log2.py | 66 + test/wafer/ops/test_log_2.py | 66 + test/wafer/ops/test_logical_and.py | 71 + test/wafer/ops/test_logical_or.py | 71 + test/wafer/ops/test_lshift.py | 105 + test/wafer/ops/test_lt.py | 70 + test/wafer/ops/test_max_dim0.py | 139 + test/wafer/ops/test_max_dim1.py | 149 + test/wafer/ops/test_max_vector.py | 94 + test/wafer/ops/test_maximum.py | 72 + test/wafer/ops/test_mean_dim0.py | 162 + test/wafer/ops/test_mean_dim1.py | 162 + test/wafer/ops/test_mean_vector.py | 110 + test/wafer/ops/test_min_dim0.py | 171 + test/wafer/ops/test_min_dim1.py | 147 + test/wafer/ops/test_min_vector.py | 93 + test/wafer/ops/test_minimum.py | 72 + test/wafer/ops/test_mod.py | 94 + test/wafer/ops/test_mul.py | 71 + test/wafer/ops/test_nearest.py | 171 + test/wafer/ops/test_neg.py | 66 + test/wafer/ops/test_npu_indexing.py | 197 + test/wafer/ops/test_npu_indexing2.py | 121 + test/wafer/ops/test_or.py | 71 + test/wafer/ops/test_permute.py | 172 + test/wafer/ops/test_permute_full.py | 203 + test/wafer/ops/test_permute_reshape.py | 102 + test/wafer/ops/test_precise_div.py | 71 + test/wafer/ops/test_precise_sqrt.py | 60 + test/wafer/ops/test_ravel.py | 75 + test/wafer/ops/test_reduce_count_vector.py | 181 + test/wafer/ops/test_reduce_mean.py | 112 + test/wafer/ops/test_reduce_sum.py | 109 + test/wafer/ops/test_reshape.py | 75 + test/wafer/ops/test_rms_norm.py | 154 + test/wafer/ops/test_rotary_embedding.py | 191 + test/wafer/ops/test_rotatry_gpt.py | 202 + test/wafer/ops/test_rotaty_embedding_gpt.py | 190 + test/wafer/ops/test_rshift.py | 105 + test/wafer/ops/test_rsqrt.py | 80 + test/wafer/ops/test_scalar_calc.py | 774 +++ test/wafer/ops/test_sigmoid.py | 71 + test/wafer/ops/test_silu.py | 106 + test/wafer/ops/test_silu_and_mul.py | 211 + test/wafer/ops/test_sin.py | 105 + test/wafer/ops/test_softmax.py | 142 + test/wafer/ops/test_split.py | 86 + test/wafer/ops/test_sqrt.py | 106 + test/wafer/ops/test_store_scalar.py | 50 + test/wafer/ops/test_sub.py | 72 + test/wafer/ops/test_sum.py | 160 + test/wafer/ops/test_sum_dim0.py | 150 + test/wafer/ops/test_sum_dim1.py | 150 + test/wafer/ops/test_sum_vector.py | 92 + test/wafer/ops/test_swap.py | 66 + test/wafer/ops/test_swiglu.py | 109 + test/wafer/ops/test_swizzle2d.py | 66 + test/wafer/ops/test_template.py | 71 + test/wafer/ops/test_tensor_get_item.py | 57 + test/wafer/ops/test_trans_3d.py | 94 + test/wafer/ops/test_triton_eq.py | 70 + test/wafer/ops/test_triton_le.py | 70 + test/wafer/ops/test_triton_lt.py | 70 + test/wafer/ops/test_triton_neq.py | 70 + test/wafer/ops/test_umulhi.py | 67 + test/wafer/ops/test_unlign_sum.py | 82 + test/wafer/ops/test_unused_func_arg.py | 95 + test/wafer/ops/test_view.py | 70 + test/wafer/ops/test_where_lt.py | 66 + test/wafer/ops/test_where_mask.py | 64 + test/wafer/ops/test_where_var.py | 65 + test/wafer/ops/test_xor.py | 102 + test/wafer/ops/test_xor_sum.py | 89 + test/wafer/ops/test_zeros.py | 125 + test/wafer/ops/test_zeroslike.py | 73 + test/wafer/runtime/conftest.py | 7 + test/wafer/runtime/test_autotune.py | 33 + test/wafer/runtime/test_elementwise.py | 29 + test/wafer/runtime/test_matmul.py | 33 + test/wafer/runtime/test_reduction.py | 34 + test/wafer/runtime/test_tensor_runtime.py | 52 + test/wafer/suites/accepted.txt | 2866 +++++++++++ test/wafer/test_benchmark_interface.py | 31 + test/wafer/test_cache_and_runtime.py | 223 + test/wafer/test_cluster_launch.py | 52 + test/wafer/test_commonir_abi.py | 51 + test/wafer/test_crt_diagnostics.py | 94 + test/wafer/test_crt_fp8.py | 47 + test/wafer/test_crt_randgen.py | 41 + test/wafer/test_dense_constants.py | 41 + test/wafer/test_device_print_abi.py | 35 + test/wafer/test_elf_audit.py | 71 + test/wafer/test_external_interfaces.py | 26 + test/wafer/test_frontend_isolation.py | 180 + test/wafer/test_loader_isolation.py | 58 + test/wafer/test_log1p_lowering.py | 21 + test/wafer/test_module_runtime.py | 68 + test/wafer/test_mxfp_reference.py | 40 + test/wafer/test_naming_compat.py | 75 + test/wafer/test_native_errors.py | 49 + test/wafer/test_noc_sync.py | 89 + test/wafer/test_patch_profiles.py | 160 + test/wafer/test_precision_modes.py | 64 + test/wafer/test_scalar_copy.py | 37 + test/wafer/test_tle_frontend.py | 96 + test/wafer/verify_wafer_acceptance.py | 143 + test/wafer/verify_wafer_examples.py | 196 + test/wafer/verify_wafer_installation.py | 81 + test/wafer/verify_wafer_runtime.py | 431 ++ test/wafer/verify_wafer_runtime.sh | 18 + test/wafer/verify_wafer_torch_stack.py | 51 + third_party/wafer/CMakeLists.txt | 161 + third_party/wafer/README.md | 80 + third_party/wafer/backend/__init__.py | 8 + third_party/wafer/backend/compiler.py | 27 + third_party/wafer/backend/driver.py | 20 + third_party/wafer/backend/include/logger.h | 106 + third_party/wafer/backend/logger_config.py | 136 + third_party/wafer/backend/name.conf | 1 + third_party/wafer/backend/txda_tools.py | 2 + third_party/wafer/backend/wafer_tools.py | 195 + third_party/wafer/benchmark/benchmark.py | 1668 +++++++ third_party/wafer/bin/CMakeLists.txt | 110 + .../wafer/bin/RegisterTritonDialects.h | 132 + third_party/wafer/bin/wafer-llvm-opt.cpp | 121 + third_party/wafer/bin/wafer-lsp.cpp | 10 + third_party/wafer/bin/wafer-opt.cpp | 11 + third_party/wafer/bin/wafer-reduce.cpp | 11 + third_party/wafer/bin/wafer-tensor-layout.cpp | 232 + third_party/wafer/cmake/WaferConfig.cmake | 5 + third_party/wafer/crt/CMakeLists.txt | 171 + third_party/wafer/crt/README.md | 2 + third_party/wafer/crt/include/Wafer/op_gelu.h | 19 + .../crt/include/Wafer/op_reduce_mul_impl.h | 10 + third_party/wafer/crt/include/Wafer/wafer.h | 97 + third_party/wafer/crt/lib/Wafer/abs.c | 33 + third_party/wafer/crt/lib/Wafer/argmax.c | 73 + third_party/wafer/crt/lib/Wafer/argmin.c | 73 + third_party/wafer/crt/lib/Wafer/arith.c | 240 + third_party/wafer/crt/lib/Wafer/assert.c | 42 + .../wafer/crt/lib/Wafer/atomic_barrier_in.c | 21 + .../wafer/crt/lib/Wafer/atomic_barrier_out.c | 20 + third_party/wafer/crt/lib/Wafer/barrier.c | 17 + third_party/wafer/crt/lib/Wafer/bf16_fp16.c | 32 + third_party/wafer/crt/lib/Wafer/bf16_fp32.c | 31 + third_party/wafer/crt/lib/Wafer/bf16_int16.c | 33 + third_party/wafer/crt/lib/Wafer/bf16_int32.c | 33 + third_party/wafer/crt/lib/Wafer/bf16_int8.c | 32 + third_party/wafer/crt/lib/Wafer/bf16_tf32.c | 32 + third_party/wafer/crt/lib/Wafer/bilinear.c | 40 + third_party/wafer/crt/lib/Wafer/bit2fp.c | 37 + third_party/wafer/crt/lib/Wafer/channelnorm.c | 156 + third_party/wafer/crt/lib/Wafer/common.c | 19 + third_party/wafer/crt/lib/Wafer/concat.c | 41 + third_party/wafer/crt/lib/Wafer/conv.c | 67 + third_party/wafer/crt/lib/Wafer/cos.c | 33 + third_party/wafer/crt/lib/Wafer/count.c | 34 + third_party/wafer/crt/lib/Wafer/empty.c | 7 + third_party/wafer/crt/lib/Wafer/exp.c | 33 + third_party/wafer/crt/lib/Wafer/explp.c | 33 + third_party/wafer/crt/lib/Wafer/fp16_bf16.c | 33 + third_party/wafer/crt/lib/Wafer/fp16_fp32.c | 31 + third_party/wafer/crt/lib/Wafer/fp16_int16.c | 33 + third_party/wafer/crt/lib/Wafer/fp16_int32.c | 33 + third_party/wafer/crt/lib/Wafer/fp16_int8.c | 32 + third_party/wafer/crt/lib/Wafer/fp16_tf32.c | 32 + third_party/wafer/crt/lib/Wafer/fp32_bf16.c | 33 + third_party/wafer/crt/lib/Wafer/fp32_fp16.c | 32 + third_party/wafer/crt/lib/Wafer/fp32_int16.c | 33 + third_party/wafer/crt/lib/Wafer/fp32_int32.c | 32 + third_party/wafer/crt/lib/Wafer/fp32_int8.c | 33 + third_party/wafer/crt/lib/Wafer/fp32_tf32.c | 32 + .../wafer/crt/lib/Wafer/gatherscatter.c | 44 + third_party/wafer/crt/lib/Wafer/gelu_none.c | 20 + third_party/wafer/crt/lib/Wafer/gelu_tanh.c | 20 + third_party/wafer/crt/lib/Wafer/gemm.c | 58 + third_party/wafer/crt/lib/Wafer/img2col.c | 43 + third_party/wafer/crt/lib/Wafer/int16_bf16.c | 33 + third_party/wafer/crt/lib/Wafer/int16_fp16.c | 32 + third_party/wafer/crt/lib/Wafer/int16_fp32.c | 32 + third_party/wafer/crt/lib/Wafer/int16_tf32.c | 32 + third_party/wafer/crt/lib/Wafer/int32_bf16.c | 32 + third_party/wafer/crt/lib/Wafer/int32_fp16.c | 33 + third_party/wafer/crt/lib/Wafer/int32_fp32.c | 32 + third_party/wafer/crt/lib/Wafer/int32_tf32.c | 32 + third_party/wafer/crt/lib/Wafer/int8_bf16.c | 33 + third_party/wafer/crt/lib/Wafer/int8_fp16.c | 33 + third_party/wafer/crt/lib/Wafer/int8_fp32.c | 32 + third_party/wafer/crt/lib/Wafer/int8_tf32.c | 33 + third_party/wafer/crt/lib/Wafer/leakyrelu.c | 34 + third_party/wafer/crt/lib/Wafer/ln.c | 32 + third_party/wafer/crt/lib/Wafer/log2.c | 32 + third_party/wafer/crt/lib/Wafer/logic.c | 162 + third_party/wafer/crt/lib/Wafer/lut16.c | 35 + third_party/wafer/crt/lib/Wafer/lut32.c | 35 + third_party/wafer/crt/lib/Wafer/mask_move.c | 31 + third_party/wafer/crt/lib/Wafer/memcpy.c | 76 + third_party/wafer/crt/lib/Wafer/memset.c | 49 + third_party/wafer/crt/lib/Wafer/mirror.c | 37 + third_party/wafer/crt/lib/Wafer/mxfp_bf16.c | 291 ++ third_party/wafer/crt/lib/Wafer/mxfp_fp16.c | 232 + .../wafer/crt/lib/Wafer/mxfp_scale_bf16.c | 68 + .../wafer/crt/lib/Wafer/mxfp_scale_fp16.c | 91 + third_party/wafer/crt/lib/Wafer/nchw2nhwc.c | 37 + third_party/wafer/crt/lib/Wafer/neg.c | 32 + third_party/wafer/crt/lib/Wafer/nhwc2nchw.c | 37 + third_party/wafer/crt/lib/Wafer/noc_init.c | 16 + third_party/wafer/crt/lib/Wafer/op_gelu.c | 524 ++ .../wafer/crt/lib/Wafer/op_reduce_mul_impl.c | 161 + third_party/wafer/crt/lib/Wafer/pad.c | 40 + third_party/wafer/crt/lib/Wafer/pow.c | 467 ++ third_party/wafer/crt/lib/Wafer/pow2.c | 33 + third_party/wafer/crt/lib/Wafer/print.c | 27 + third_party/wafer/crt/lib/Wafer/randgen.c | 39 + third_party/wafer/crt/lib/Wafer/rdma.c | 140 + third_party/wafer/crt/lib/Wafer/recip.c | 33 + third_party/wafer/crt/lib/Wafer/recv.c | 26 + third_party/wafer/crt/lib/Wafer/reduce.c | 112 + third_party/wafer/crt/lib/Wafer/relation.c | 562 +++ third_party/wafer/crt/lib/Wafer/relu.c | 33 + third_party/wafer/crt/lib/Wafer/rotate180.c | 38 + third_party/wafer/crt/lib/Wafer/rotate270.c | 38 + third_party/wafer/crt/lib/Wafer/rotate90.c | 38 + third_party/wafer/crt/lib/Wafer/rsqrt.c | 35 + third_party/wafer/crt/lib/Wafer/satrelu.c | 35 + third_party/wafer/crt/lib/Wafer/send.c | 199 + third_party/wafer/crt/lib/Wafer/sigmoid.c | 35 + third_party/wafer/crt/lib/Wafer/sin.c | 33 + third_party/wafer/crt/lib/Wafer/softplus.c | 36 + third_party/wafer/crt/lib/Wafer/sqrt.c | 34 + third_party/wafer/crt/lib/Wafer/tanh.c | 33 + third_party/wafer/crt/lib/Wafer/tensornorm.c | 38 + third_party/wafer/crt/lib/Wafer/tf32_bf16.c | 33 + third_party/wafer/crt/lib/Wafer/tf32_fp16.c | 32 + third_party/wafer/crt/lib/Wafer/tf32_fp32.c | 32 + third_party/wafer/crt/lib/Wafer/tf32_int16.c | 33 + third_party/wafer/crt/lib/Wafer/tf32_int32.c | 33 + third_party/wafer/crt/lib/Wafer/tf32_int8.c | 33 + third_party/wafer/crt/lib/Wafer/transpose.c | 37 + third_party/wafer/crt/lib/Wafer/wafer.c | 178 + third_party/wafer/crt/lib/Wafer/wdma.c | 139 + .../wafer/examples/_wafer_reference.py | 27 + third_party/wafer/examples/bare_matmul.py | 52 + third_party/wafer/examples/bare_matmul_acc.py | 46 + .../wafer/examples/bare_matmul_autotune.py | 69 + third_party/wafer/examples/benchmark.py | 65 + third_party/wafer/examples/conftest.py | 12 + third_party/wafer/examples/dump_vec_add_ir.sh | 47 + third_party/wafer/examples/embedding.py | 89 + .../examples/flagtree/test_tle_cumsum.py | 135 + .../examples/flagtree/test_tle_dsa_arith.py | 102 + .../examples/flagtree/test_tle_dsa_bridge.py | 84 + .../flagtree/test_tle_dsa_pipeline_e2e.py | 234 + .../examples/flagtree/test_tle_dsa_rand.py | 155 + .../examples/flagtree/test_tle_dsa_slice.py | 338 ++ third_party/wafer/examples/mult_ir.py | 194 + third_party/wafer/examples/profile_matmul.py | 193 + third_party/wafer/examples/quant_gptq.py | 280 ++ third_party/wafer/examples/quant_kernel.py | 623 +++ third_party/wafer/examples/single_conv2d.py | 261 + third_party/wafer/examples/test_abs.py | 103 + third_party/wafer/examples/test_addptr.py | 45 + third_party/wafer/examples/test_argmax2d.py | 53 + third_party/wafer/examples/test_assert.py | 56 + third_party/wafer/examples/test_autotune.py | 39 + .../wafer/examples/test_bare_matmul.py | 75 + .../wafer/examples/test_bare_matmul_acc.py | 77 + .../examples/test_blockptr_complex_offset.py | 41 + third_party/wafer/examples/test_cdiv.py | 111 + third_party/wafer/examples/test_ceil.py | 100 + third_party/wafer/examples/test_clamp.py | 106 + third_party/wafer/examples/test_cos.py | 103 + .../wafer/examples/test_debug_barrier.py | 30 + third_party/wafer/examples/test_div_rn.py | 107 + third_party/wafer/examples/test_dot_scaled.py | 393 ++ .../wafer/examples/test_early_return.py | 64 + third_party/wafer/examples/test_embedding.py | 94 + third_party/wafer/examples/test_exp.py | 96 + third_party/wafer/examples/test_exp2.py | 98 + third_party/wafer/examples/test_fdiv.py | 104 + third_party/wafer/examples/test_flip.py | 51 + third_party/wafer/examples/test_floor.py | 100 + third_party/wafer/examples/test_fma.py | 112 + .../wafer/examples/test_fp8_conversion.py | 29 + third_party/wafer/examples/test_gather.py | 53 + third_party/wafer/examples/test_histogram.py | 47 + third_party/wafer/examples/test_layernorm.py | 183 + third_party/wafer/examples/test_libdevice.py | 207 + third_party/wafer/examples/test_load6d.py | 76 + .../examples/test_load_2d_tensor_block.py | 79 + .../wafer/examples/test_load_2d_tensor_col.py | 71 + .../wafer/examples/test_load_store_mod.py | 76 + third_party/wafer/examples/test_log.py | 96 + third_party/wafer/examples/test_log2.py | 96 + third_party/wafer/examples/test_mask.py | 42 + .../wafer/examples/test_math_erf_op.py | 85 + third_party/wafer/examples/test_matmul.py | 187 + third_party/wafer/examples/test_maximum.py | 102 + third_party/wafer/examples/test_minimum.py | 102 + third_party/wafer/examples/test_modulo.py | 329 ++ .../wafer/examples/test_nested_loops.py | 399 ++ third_party/wafer/examples/test_pipeline.py | 36 + .../wafer/examples/test_precision_modes.py | 38 + third_party/wafer/examples/test_print.py | 30 + third_party/wafer/examples/test_reduce.py | 55 + third_party/wafer/examples/test_reduce1d.py | 45 + third_party/wafer/examples/test_rsqrt.py | 96 + .../wafer/examples/test_scalar_store.py | 32 + third_party/wafer/examples/test_scan.py | 72 + third_party/wafer/examples/test_scan2d.py | 222 + third_party/wafer/examples/test_scan3d.py | 168 + third_party/wafer/examples/test_sigmoid.py | 96 + .../wafer/examples/test_sign_extend.py | 45 + third_party/wafer/examples/test_sin.py | 104 + third_party/wafer/examples/test_softmax.py | 93 + third_party/wafer/examples/test_sort.py | 36 + third_party/wafer/examples/test_splat.py | 45 + third_party/wafer/examples/test_sqrt.py | 98 + third_party/wafer/examples/test_sqrt_rn.py | 97 + third_party/wafer/examples/test_swap.py | 50 + third_party/wafer/examples/test_swizzle2d.py | 26 + .../examples/test_tensor_index_iterargs.py | 126 + third_party/wafer/examples/test_trans2d.py | 38 + third_party/wafer/examples/test_umulhi.py | 82 + third_party/wafer/examples/test_vec_add.py | 98 + third_party/wafer/examples/test_where.py | 69 + third_party/wafer/examples/time1.py | 210 + third_party/wafer/examples/time_zs_opt2.py | 146 + .../tle/test_tle_dsa_noc_gemm_4096.py | 153 + third_party/wafer/examples/util.py | 52 + third_party/wafer/examples/view_vec_add_ir.sh | 41 + third_party/wafer/experimental/README.md | 13 + .../wafer/experimental/tle/__init__.py | 14 + .../experimental/tle/language/__init__.py | 40 + .../wafer/experimental/tle/language/core.py | 140 + .../experimental/tle/language/distributed.py | 669 +++ .../experimental/tle/language/dsa/__init__.py | 36 + .../experimental/tle/language/dsa/core.py | 586 +++ .../experimental/tle/language/dsa/semantic.py | 179 + .../experimental/tle/language/dsa/types.py | 126 + .../tle/language/dsa/wafer/__init__.py | 4 + .../tle/language/dsa/wafer/core.py | 194 + .../wafer/include/Address/CMakeLists.txt | 2 + .../include/Address/Dialect/CMakeLists.txt | 1 + .../Address/Dialect/IR/AddressDialect.h | 36 + .../Address/Dialect/IR/AddressDialect.td | 77 + .../include/Address/Dialect/IR/AddressOps.td | 247 + .../include/Address/Dialect/IR/CMakeLists.txt | 20 + .../include/Address/Transforms/CMakeLists.txt | 5 + .../wafer/include/Address/Transforms/Passes.h | 33 + .../include/Address/Transforms/Passes.td | 44 + third_party/wafer/include/Analysis/Alias.h | 96 + .../wafer/include/Analysis/Allocation.h | 219 + third_party/wafer/include/Analysis/Membar.h | 121 + third_party/wafer/include/Analysis/Utility.h | 154 + third_party/wafer/include/CMakeLists.txt | 7 + .../include/ExecutionEngine/CRunnerUtils.cpp | 189 + .../include/ExecutionEngine/CRunnerUtils.h | 482 ++ .../wafer/include/ExecutionEngine/Msan.h | 35 + .../wafer/include/ExecutionEngine/version.txt | 1 + .../include/flagtree/Common/UnifiedHardware.h | 31 + .../include/magic-kernel-func/CMakeLists.txt | 2 + .../magic-kernel-func/Dialect/CMakeLists.txt | 1 + .../Dialect/IR/MagicKernelFuncOps.td | 19 + .../include/magic-kernel-instr/CMakeLists.txt | 2 + .../magic-kernel-instr/Dialect/CMakeLists.txt | 1 + .../Dialect/IR/MagicKernelInstrOps.td | 13 + .../wafer/include/magic-kernel/CMakeLists.txt | 3 + .../magic-kernel/Conversion/CMakeLists.txt | 6 + .../CoreDialectsToMK/CMakeLists.txt | 10 + .../CoreDialectsToMK/CoreDialectsToMK.h | 27 + .../Conversion/CoreDialectsToMK/Passes.h | 26 + .../Conversion/CoreDialectsToMK/Passes.td | 24 + .../LegalizeTensorFormLoops/CMakeLists.txt | 3 + .../LegalizeTensorFormLoops/Passes.h | 23 + .../LegalizeTensorFormLoops/Passes.td | 17 + .../Conversion/LinalgToMK/CMakeLists.txt | 3 + .../Conversion/LinalgToMK/LinalgToMK.h | 83 + .../Conversion/LinalgToMK/Passes.h | 22 + .../Conversion/LinalgToMK/Passes.td | 24 + .../Conversion/MKPipeline/CMakeLists.txt | 3 + .../Conversion/MKPipeline/Passes.h | 19 + .../Conversion/MKPipeline/Passes.td | 64 + .../Conversion/TLEToMK/CMakeLists.txt | 3 + .../magic-kernel/Conversion/TLEToMK/Passes.h | 22 + .../magic-kernel/Conversion/TLEToMK/Passes.td | 19 + .../magic-kernel/Conversion/TLEToMK/TLEToMK.h | 32 + .../magic-kernel/Dialect/CMakeLists.txt | 1 + .../magic-kernel/Dialect/IR/CMakeLists.txt | 11 + .../Dialect/IR/MagicKernelAttrDefs.td | 15 + .../Dialect/IR/MagicKernelDialect.h | 32 + .../Dialect/IR/MagicKernelDialect.td | 44 + .../magic-kernel/Dialect/IR/MagicKernelOps.td | 909 ++++ .../Dialect/IR/MagicKernelTypes.td | 102 + .../Transforms/BufferizableOpInterfaceImpl.h | 26 + .../magic-kernel/Transforms/CMakeLists.txt | 3 + .../include/magic-kernel/Transforms/Passes.h | 20 + .../include/magic-kernel/Transforms/Passes.td | 49 + .../include/triton-shared/CMakeLists.txt | 1 + .../triton-shared/Conversion/CMakeLists.txt | 12 + .../ConvertTritonPtr/CMakeLists.txt | 9 + .../Conversion/ConvertTritonPtr/Passes.h | 22 + .../Conversion/ConvertTritonPtr/Passes.td | 18 + .../ConvertTritonPtr/TritonPtrToAddress.h | 22 + .../ReconcilePtrCasts/CMakeLists.txt | 3 + .../Conversion/ReconcilePtrCasts/Passes.h | 15 + .../Conversion/ReconcilePtrCasts/Passes.td | 18 + .../ReconcilePtrCasts/ReconcilePtrCasts.h | 22 + .../Conversion/StructuredToMK/CMakeLists.txt | 3 + .../Conversion/StructuredToMK/Passes.h | 15 + .../Conversion/StructuredToMK/Passes.td | 10 + .../StructuredToMK/StructuredToMK.h | 24 + .../StructuredToMemref/CMakeLists.txt | 3 + .../Conversion/StructuredToMemref/Passes.h | 15 + .../Conversion/StructuredToMemref/Passes.td | 10 + .../StructuredToMemref/StructuredToMemref.h | 24 + .../TritonArithToLinalg/CMakeLists.txt | 3 + .../TritonArithToLinalg/ConversionPatterns.h | 2572 ++++++++++ .../Conversion/TritonArithToLinalg/Passes.h | 15 + .../Conversion/TritonArithToLinalg/Passes.td | 22 + .../TritonArithToLinalg/TritonArithToLinalg.h | 31 + .../TritonPtrToMemref/CMakeLists.txt | 3 + .../Conversion/TritonPtrToMemref/Passes.h | 15 + .../Conversion/TritonPtrToMemref/Passes.td | 11 + .../TritonPtrToMemref/TritonPtrToMemref.h | 17 + .../TritonToCoreDialects/CMakeLists.txt | 9 + .../Conversion/TritonToCoreDialects/Passes.h | 22 + .../Conversion/TritonToCoreDialects/Passes.td | 18 + .../TritonToCoreDialects.h | 27 + .../Conversion/TritonToMK/CMakeLists.txt | 3 + .../Conversion/TritonToMK/Passes.h | 15 + .../Conversion/TritonToMK/Passes.td | 10 + .../TritonToMK/TritonToMKPatterns.hpp | 194 + .../TritonToUnstructured/CMakeLists.txt | 3 + .../Conversion/TritonToUnstructured/Passes.h | 15 + .../Conversion/TritonToUnstructured/Passes.td | 15 + .../TritonToUnstructured.h | 17 + .../UnstructuredToMK/CMakeLists.txt | 10 + .../Conversion/UnstructuredToMK/Passes.h | 22 + .../Conversion/UnstructuredToMK/Passes.td | 18 + .../UnstructuredToMK/UnstructuredToMK.h | 21 + .../UnstructuredToMemref/CMakeLists.txt | 10 + .../Conversion/UnstructuredToMemref/Passes.h | 22 + .../Conversion/UnstructuredToMemref/Passes.td | 18 + .../UnstructuredToMemref.h | 21 + .../include/utils/LinalgOpBuilderHelper.h | 50 + .../wafer/include/utils/TypeConvertor.h | 30 + .../wafer/include/wafer/CMakeLists.txt | 3 + .../AllocateSharedMemory/CMakeLists.txt | 3 + .../Conversion/AllocateSharedMemory/Passes.h | 25 + .../Conversion/AllocateSharedMemory/Passes.td | 24 + .../include/wafer/Conversion/CMakeLists.txt | 7 + .../ExportKernelSymbols/CMakeLists.txt | 3 + .../ExportKernelSymbols/ExportKernelSymbols.h | 26 + .../Conversion/ExportKernelSymbols/Passes.h | 22 + .../Conversion/ExportKernelSymbols/Passes.td | 20 + .../Conversion/LinalgFusion/CMakeLists.txt | 3 + .../Conversion/LinalgFusion/LinalgFusion.h | 27 + .../wafer/Conversion/LinalgFusion/Passes.h | 22 + .../wafer/Conversion/LinalgFusion/Passes.td | 23 + .../Conversion/LinalgTiling/CMakeLists.txt | 3 + .../Conversion/LinalgTiling/LinalgTiling.h | 34 + .../wafer/Conversion/LinalgTiling/Passes.h | 22 + .../wafer/Conversion/LinalgTiling/Passes.td | 22 + .../wafer/Conversion/MKToWafer/CMakeLists.txt | 3 + .../wafer/Conversion/MKToWafer/MKToWafer.h | 36 + .../wafer/Conversion/MKToWafer/Passes.h | 22 + .../wafer/Conversion/MKToWafer/Passes.td | 18 + .../WaferMemrefToLLVM/CMakeLists.txt | 3 + .../Conversion/WaferMemrefToLLVM/Passes.h | 22 + .../Conversion/WaferMemrefToLLVM/Passes.td | 19 + .../WaferMemrefToLLVM/WaferMemrefToLLVM.h | 40 + .../Conversion/WaferToLLVM/CMakeLists.txt | 7 + .../WaferToLLVM/KernelArgBufferPass.h | 35 + .../WaferToLLVM/KernelArgBufferPass.td | 32 + .../wafer/Conversion/WaferToLLVM/Passes.h | 22 + .../wafer/Conversion/WaferToLLVM/Passes.td | 38 + .../Conversion/WaferToLLVM/WaferToLLVM.h | 33 + .../include/wafer/Dialect/CMakeLists.txt | 1 + .../include/wafer/Dialect/IR/CMakeLists.txt | 14 + .../include/wafer/Dialect/IR/WaferAttrDefs.td | 24 + .../include/wafer/Dialect/IR/WaferDialect.h | 33 + .../include/wafer/Dialect/IR/WaferDialect.td | 43 + .../wafer/include/wafer/Dialect/IR/WaferOps.h | 26 + .../include/wafer/Dialect/IR/WaferOps.td | 1278 +++++ .../include/wafer/Dialect/IR/WaferTypes.td | 107 + .../include/wafer/Transforms/CMakeLists.txt | 3 + .../wafer/include/wafer/Transforms/Passes.h | 10 + .../wafer/include/wafer/Transforms/Passes.td | 5 + third_party/wafer/language/cpu/__init__.py | 3 + third_party/wafer/language/cpu/libdevice.py | 1496 ++++++ third_party/wafer/language/txda/__init__.py | 4 + third_party/wafer/language/txda/libdevice.py | 2 + third_party/wafer/language/wafer/__init__.py | 3 + third_party/wafer/language/wafer/libdevice.py | 1497 ++++++ third_party/wafer/language/wafer/slicing.py | 42 + third_party/wafer/lib/Analysis/Alias.cpp | 79 + third_party/wafer/lib/Analysis/Allocation.cpp | 756 +++ third_party/wafer/lib/Analysis/CMakeLists.txt | 13 + third_party/wafer/lib/Analysis/Membar.cpp | 411 ++ third_party/wafer/lib/CMakeLists.txt | 9 + third_party/wafer/lib/Common/CMakeLists.txt | 1 + .../wafer/lib/Common/UnifiedHardware.cc | 32 + .../AllocateSharedMemoryPass.cpp | 135 + .../AllocateSharedMemory/CMakeLists.txt | 20 + .../wafer/lib/Conversion/CMakeLists.txt | 26 + .../ConvertTritonPtr/CMakeLists.txt | 25 + .../TritonPtrToAddressPass.cpp | 181 + .../CoreDialectsToMK/CMakeLists.txt | 25 + .../CoreDialectsToMK/CoreDialectsToMKPass.cpp | 63 + .../ExportKernelSymbols/CMakeLists.txt | 14 + .../ExportKernelSymbols.cpp | 168 + .../LegalizeTensorFormLoops/CMakeLists.txt | 23 + .../LegalizeTensorFormLoops.cpp | 79 + .../Conversion/LinalgFusion/CMakeLists.txt | 17 + .../Conversion/LinalgFusion/LinalgFusion.cpp | 275 ++ .../LinalgFusion/LinalgFusionPass.cpp | 66 + .../Conversion/LinalgTiling/CMakeLists.txt | 17 + .../Conversion/LinalgTiling/LinalgTiling.cpp | 87 + .../LinalgTiling/LinalgTilingPass.cpp | 64 + .../lib/Conversion/LinalgToMK/CMakeLists.txt | 21 + .../lib/Conversion/LinalgToMK/LinalgToMK.cpp | 4244 +++++++++++++++++ .../Conversion/LinalgToMK/LinalgToMKPass.cpp | 156 + .../lib/Conversion/MKPipeline/CMakeLists.txt | 18 + .../MKLoopBoundCanonicalizePass.cpp | 229 + .../Conversion/MKPipeline/MKPipelinePass.cpp | 1965 ++++++++ .../lib/Conversion/MKToWafer/CMakeLists.txt | 20 + .../lib/Conversion/MKToWafer/MKToWafer.cpp | 2438 ++++++++++ .../Conversion/MKToWafer/MKToWaferPass.cpp | 123 + .../ReconcilePtrCasts/CMakeLists.txt | 19 + .../ReconcilePtrCastsPass.cpp | 161 + .../Conversion/StructuredToMK/CMakeLists.txt | 23 + .../StructuredToMK/StructuredToMK.cpp | 151 + .../StructuredToMK/StructuredToMKPass.cpp | 151 + .../StructuredToMemref/CMakeLists.txt | 23 + .../StructuredToMemref/StructuredToMemref.cpp | 966 ++++ .../StructuredToMemrefPass.cpp | 153 + .../lib/Conversion/TLEToMK/CMakeLists.txt | 21 + .../TLEToMK/MKCommonBufferPlanningPass.cpp | 183 + .../wafer/lib/Conversion/TLEToMK/TLEToMK.cpp | 717 +++ .../lib/Conversion/TLEToMK/TLEToMKPass.cpp | 59 + .../TritonArithToLinalg/CMakeLists.txt | 22 + .../TritonArithToLinalg.cpp | 107 + .../TritonArithToLinalgPass.cpp | 250 + .../TritonToCoreDialects/CMakeLists.txt | 38 + .../TritonToCoreDialectsPass.cpp | 102 + .../UnstructuredToMK/CMakeLists.txt | 19 + .../UnstructuredToMK/UnstructuredToMKPass.cpp | 293 ++ .../WaferMemrefToLLVM/CMakeLists.txt | 19 + .../WaferMemrefToLLVM/WaferMemrefToLLVM.cpp | 523 ++ .../WaferMemrefToLLVMPass.cpp | 88 + .../lib/Conversion/WaferToLLVM/CMakeLists.txt | 25 + .../WaferToLLVM/KernelArgBufferPass.cpp | 161 + .../Conversion/WaferToLLVM/WaferToLLVM.cpp | 2716 +++++++++++ .../WaferToLLVM/WaferToLLVMPass.cpp | 78 + .../wafer/lib/Dialect/Address/CMakeLists.txt | 2 + .../lib/Dialect/Address/IR/AddressDialect.cpp | 303 ++ .../lib/Dialect/Address/IR/CMakeLists.txt | 13 + .../Dialect/Address/Transforms/AddrToLLVM.cpp | 366 ++ .../Dialect/Address/Transforms/CMakeLists.txt | 16 + third_party/wafer/lib/Dialect/CMakeLists.txt | 5 + .../lib/Dialect/MagicKernel/CMakeLists.txt | 12 + .../MagicKernel/IR/MagicKernelDialect.cpp | 51 + .../BufferizableOpInterfaceImpl.cpp | 338 ++ .../MagicKernel/Transforms/CMakeLists.txt | 18 + .../MaterializeStridedLinalgInputsPass.cpp | 232 + .../wafer/lib/Dialect/Wafer/CMakeLists.txt | 12 + .../lib/Dialect/Wafer/IR/WaferDialect.cpp | 30 + .../wafer/lib/Dialect/Wafer/IR/WaferOps.cpp | 10 + .../Dialect/Wafer/Transforms/CMakeLists.txt | 5 + .../Wafer/Transforms/InsertBarrierPass.cpp | 553 +++ .../wafer/lib/Registrar/CMakeLists.txt | 1 + third_party/wafer/lib/Registrar/Registrar.cc | 15 + third_party/wafer/name.conf | 1 + .../triton/cache_without_vendor_imports.patch | 51 + .../wafer/patches/triton/profiles.json | 32 + ...on_triton_compiler_optional_gluon_py.patch | 49 + .../python_triton_jit_optional_gluon.patch | 19 + .../triton/wafer_builder_optional_gluon.patch | 116 + .../triton/wafer_proton_backend_filter.patch | 82 + third_party/wafer/profiler/CMakeLists.txt | 70 + third_party/wafer/profiler/profiler.cpp | 232 + third_party/wafer/python/triton_wafer.cc | 98 + .../wafer/python/triton_wafer_frontend.cc | 48 + third_party/wafer/requirements-build.txt | 7 + third_party/wafer/scripts/base/base_run.sh | 110 + third_party/wafer/scripts/build_llvm.sh | 31 + third_party/wafer/scripts/build_wafer.sh | 142 + .../publish/run_flaggems_on_multicards.sh | 97 + .../wafer/scripts/publish/run_wafer.sh | 10 + third_party/wafer/scripts/requirements_ts.txt | 7 + third_party/wafer/scripts/run_wafer.sh | 11 + third_party/wafer/scripts/tools/suuplement.sh | 42 + .../wafer/third_party/flir/.clang-format | 1 + .../flir/.github/PULL_REQUEST_TEMPLATE.md | 16 + .../.github/workflows/code-format-check.yml | 23 + third_party/wafer/third_party/flir/.gitignore | 5 + .../wafer/third_party/flir/.gitmodules | 0 .../third_party/flir/.pre-commit-config.yaml | 29 + .../wafer/third_party/flir/CMakeLists.txt | 27 + third_party/wafer/third_party/flir/LICENSE | 23 + third_party/wafer/third_party/flir/README.md | 12 + .../third_party/flir/backend/compiler.py | 219 + .../wafer/third_party/flir/backend/driver.py | 397 ++ .../include/ExecutionEngine/CRunnerUtils.cpp | 192 + .../include/ExecutionEngine/CRunnerUtils.h | 499 ++ .../backend/include/ExecutionEngine/Msan.h | 35 + .../include/ExecutionEngine/version.txt | 1 + .../wafer/third_party/flir/backend/name.conf | 1 + .../third_party/flir/include/CMakeLists.txt | 6 + .../flir/include/incubated/CMakeLists.txt | 2 + .../incubated/Conversion/CMakeLists.txt | 5 + .../CMakeLists.txt | 3 + .../DiscreteMaskAccessConversionPass.h | 66 + .../DiscreteMaskAccessConversion/Passes.h | 37 + .../DiscreteMaskAccessConversion/Passes.td | 24 + .../TritonToAnnotation/CMakeLists.txt | 3 + .../Conversion/TritonToAnnotation/Passes.h | 43 + .../Conversion/TritonToAnnotation/Passes.td | 17 + .../ArgMinMaxConverter.h | 360 ++ .../BlockPtrAnalysis.h | 296 ++ .../TritonToLinalgIncubated/CMakeLists.txt | 3 + .../ConversionPatterns.h | 133 + .../DescriptorConverter.h | 74 + .../FunctionConverter.h | 60 + .../TritonToLinalgIncubated/HoistBroadcast.h | 82 + .../LoadStoreConverter.h | 278 ++ .../TritonToLinalgIncubated/MaskAnalysis.h | 151 + .../TritonToLinalgIncubated/Passes.h | 39 + .../TritonToLinalgIncubated/Passes.td | 28 + .../TritonOpConverter.h | 696 +++ .../TritonToLinalgIncubatedPass.h | 126 + .../TritonToLinalgIncubated/UseAnalysis.h | 146 + .../CMakeLists.txt | 3 + .../CannonicalizerConverter.h | 162 + .../MaskAnalysis.h | 131 + .../MemOpConverter.h | 138 + .../TritonToStructuredIncubated/Passes.h | 37 + .../TritonToStructuredIncubated/Passes.td | 21 + .../TritonToStructuredIncubated/PtrAnalysis.h | 177 + .../TritonToStructuredIncubatedPass.h | 74 + .../BubbleUpOperation.h | 112 + .../CMakeLists.txt | 3 + .../OffsetAnalysis.h | 260 + .../TritonToUnstructureIncubated/Passes.h | 38 + .../TritonToUnstructureIncubated/Passes.td | 29 + .../UnstructureConversionPass.h | 150 + .../Conversion/UtilsIncubated/CMakeLists.txt | 0 .../UtilsIncubated/InterleaveOptimization.h | 93 + .../Conversion/UtilsIncubated/Utils.h | 252 + .../include/incubated/Dialect/CMakeLists.txt | 1 + .../TritonStructuredIncubated/CMakeLists.txt | 1 + .../IR/CMakeLists.txt | 8 + .../IR/TritonStructuredDialectIncubated.h | 30 + .../IR/TritonStructuredDialectIncubated.td | 66 + .../flir/include/mlir-ext/CMakeLists.txt | 1 + .../include/mlir-ext/Dialect/CMakeLists.txt | 1 + .../mlir-ext/Dialect/MathExt/CMakeLists.txt | 1 + .../Dialect/MathExt/IR/CMakeLists.txt | 10 + .../mlir-ext/Dialect/MathExt/IR/MathExt.h | 22 + .../Dialect/MathExt/IR/MathExtBase.td | 25 + .../mlir-ext/Dialect/MathExt/IR/MathExtOps.td | 128 + .../flir/include/npu/CMakeLists.txt | 1 + .../flir/include/npu/Dialect/CMakeLists.txt | 1 + .../npu/Dialect/TritonAscend/CMakeLists.txt | 1 + .../Dialect/TritonAscend/IR/CMakeLists.txt | 15 + .../TritonAscend/IR/TritonAscendAttrDefs.td | 25 + .../TritonAscend/IR/TritonAscendDialect.h | 34 + .../TritonAscend/IR/TritonAscendDialect.td | 31 + .../TritonAscend/IR/TritonAscendOps.td | 488 ++ .../triton-shared/Analysis/MaskAnalysis.h | 238 + .../Analysis/OpFoldResultUtils.h | 77 + .../triton-shared/Analysis/PtrAnalysis.h | 271 ++ .../triton-shared/Analysis/UseAnalysis.h | 119 + .../AnalysisStructured/PtrAnalysis.h | 312 ++ .../flir/include/triton-shared/CMakeLists.txt | 2 + .../triton-shared/Conversion/CMakeLists.txt | 13 + .../MemrefCopyToDMA_FlagTree/CMakeLists.txt | 3 + .../MemrefCopyToDMAFlagTree.h | 24 + .../MemrefCopyToDMA_FlagTree/Passes.h | 15 + .../MemrefCopyToDMA_FlagTree/Passes.td | 10 + .../NoBufferize_FlagTree/CMakeLists.txt | 3 + .../NoBufferizeFlagTree.h | 24 + .../Conversion/NoBufferize_FlagTree/Passes.h | 15 + .../Conversion/NoBufferize_FlagTree/Passes.td | 10 + .../ReconcilePtrCasts/CMakeLists.txt | 3 + .../Conversion/ReconcilePtrCasts/Passes.h | 15 + .../Conversion/ReconcilePtrCasts/Passes.td | 18 + .../ReconcilePtrCasts/ReconcilePtrCasts.h | 22 + .../StructuredToMemref/CMakeLists.txt | 3 + .../Conversion/StructuredToMemref/Passes.h | 15 + .../Conversion/StructuredToMemref/Passes.td | 10 + .../StructuredToMemref/StructuredToMemref.h | 24 + .../TritonArithToLinalg/CMakeLists.txt | 3 + .../ConversionPatterns.hpp | 2690 +++++++++++ .../ConversionPatterns_FlagTree.hpp | 261 + .../Conversion/TritonArithToLinalg/Passes.h | 15 + .../Conversion/TritonArithToLinalg/Passes.td | 22 + .../TritonArithToLinalg/TritonArithToLinalg.h | 32 + .../TritonPtrToMemref/CMakeLists.txt | 3 + .../Conversion/TritonPtrToMemref/Passes.h | 15 + .../Conversion/TritonPtrToMemref/Passes.td | 11 + .../TritonPtrToMemref/TritonPtrToMemref.h | 17 + .../Conversion/TritonToLinalg/CMakeLists.txt | 9 + .../Conversion/TritonToLinalg/Passes.h | 22 + .../Conversion/TritonToLinalg/Passes.td | 18 + .../TritonToLinalg/TritonToLinalg.h | 33 + .../TritonToLinalgExperimental/CMakeLists.txt | 9 + .../TritonToLinalgExperimental/Passes.h | 24 + .../TritonToLinalgExperimental/Passes.td | 23 + .../TritonToLinalgExperimental.h | 22 + .../TritonToLinalgExperimental/TritonToPtr.h | 22 + .../TritonToStructured/CMakeLists.txt | 3 + .../Conversion/TritonToStructured/Passes.h | 15 + .../Conversion/TritonToStructured/Passes.td | 19 + .../TritonToStructured/TritonToStructured.h | 17 + .../TritonToUnstructured/CMakeLists.txt | 3 + .../Conversion/TritonToUnstructured/Passes.h | 15 + .../Conversion/TritonToUnstructured/Passes.td | 15 + .../TritonToUnstructured.h | 17 + .../UnstructuredToMemref/CMakeLists.txt | 10 + .../Conversion/UnstructuredToMemref/Passes.h | 22 + .../Conversion/UnstructuredToMemref/Passes.td | 18 + .../UnstructuredToMemref.h | 21 + .../triton-shared/Dialect/CMakeLists.txt | 3 + .../triton-shared/Dialect/TPtr/CMakeLists.txt | 1 + .../Dialect/TPtr/IR/CMakeLists.txt | 11 + .../Dialect/TPtr/IR/TPtrDialect.h | 23 + .../Dialect/TPtr/IR/TPtrDialect.td | 200 + .../Dialect/TritonStructured/CMakeLists.txt | 1 + .../TritonStructured/IR/CMakeLists.txt | 8 + .../IR/TritonStructuredDialect.h | 29 + .../IR/TritonStructuredDialect.td | 338 ++ .../Dialect/TritonTilingExt/CMakeLists.txt | 1 + .../Dialect/TritonTilingExt/IR/CMakeLists.txt | 11 + .../IR/TritonTilingExtDialect.h | 107 + .../IR/TritonTilingExtInterfaces.td | 102 + .../TritonTilingExt/IR/TritonTilingExtOps.td | 242 + .../triton-shared/Utils/FusionHelper.h | 207 + .../triton-shared/Utils/ReduceScanCommon.h | 353 ++ .../flir/include/triton-shared/Utils/Utils.h | 69 + .../flir/lib/Analysis/CMakeLists.txt | 14 + .../flir/lib/Analysis/MaskAnalysis.cpp | 821 ++++ .../flir/lib/Analysis/OpFoldResultUtils.cpp | 396 ++ .../flir/lib/Analysis/PtrAnalysis.cpp | 1375 ++++++ .../flir/lib/Analysis/UseAnalysis.cpp | 220 + .../lib/AnalysisStructured/CMakeLists.txt | 26 + .../lib/AnalysisStructured/PtrAnalysis.cpp | 1396 ++++++ .../lib/AnalysisStructured/PtrAnalysisTS.cpp | 1618 +++++++ .../wafer/third_party/flir/lib/CMakeLists.txt | 9 + .../flir/lib/Conversion/CMakeLists.txt | 20 + .../CMakeLists.txt | 14 + .../DiscreteMaskAccessConversionPass.cpp | 196 + .../MemrefCopyToDMA_FlagTree/CMakeLists.txt | 24 + .../MemrefCopyToDMAFlagTree.cpp | 190 + .../MemrefCopyToDMAFlagTreePass.cpp | 145 + .../NoBufferize_FlagTree/CMakeLists.txt | 24 + .../NoBufferizeFlagTree.cpp | 65 + .../NoBufferizeFlagTreePass.cpp | 64 + .../ReconcilePtrCasts/CMakeLists.txt | 18 + .../ReconcilePtrCastsPass.cpp | 167 + .../StructuredToMemref/CMakeLists.txt | 22 + .../StructuredToMemref/StructuredToMemref.cpp | 903 ++++ .../StructuredToMemrefPass.cpp | 183 + .../TritonArithToLinalg/CMakeLists.txt | 24 + .../TritonArithToLinalg.cpp | 106 + .../TritonArithToLinalgPass.cpp | 255 + .../TritonPtrToMemref/CMakeLists.txt | 18 + .../TritonPtrToMemrefPass.cpp | 145 + .../TritonToAnnotation/CMakeLists.txt | 15 + .../TritonToAnnotation/TritonToAnnotation.cpp | 78 + .../Conversion/TritonToLinalg/CMakeLists.txt | 23 + .../TritonToLinalg/TritonToLinalg.cpp | 96 + .../TritonToLinalg/TritonToLinalgPass.cpp | 229 + .../TritonToLinalgExperimental/CMakeLists.txt | 40 + .../TritonToLinalgExperimentalPass.cpp | 107 + .../TritonToPtrPass.cpp | 490 ++ .../ArgMinMaxConverter.cpp | 113 + .../BlockPtrAnalysis.cpp | 2176 +++++++++ .../TritonToLinalgIncubated/CMakeLists.txt | 32 + .../DescriptorConverter.cpp | 196 + .../FunctionConverter.cpp | 56 + .../HoistBroadcast.cpp | 229 + .../LoadStoreConverter.cpp | 1281 +++++ .../TritonToLinalgIncubated/MaskAnalysis.cpp | 667 +++ .../TritonOpConverter.cpp | 2670 +++++++++++ .../TritonToLinalgIncubatedPass.cpp | 1213 +++++ .../TritonToLinalgIncubated/UseAnalysis.cpp | 531 +++ .../TritonToStructured/CMakeLists.txt | 22 + .../TritonToStructuredPass.cpp | 388 ++ .../CMakeLists.txt | 27 + .../CannonicalizerConverter.cpp | 605 +++ .../MaskAnalysis.cpp | 965 ++++ .../MemOpConverter.cpp | 583 +++ .../PtrAnalysis.cpp | 1361 ++++++ .../TritonToStructuredIncubatedPass.cpp | 149 + .../BubbleUpOperation.cpp | 502 ++ .../CMakeLists.txt | 20 + .../OffsetAnalysis.cpp | 989 ++++ .../UnstructureConversionPass.cpp | 964 ++++ .../TritonToUnstructured/CMakeLists.txt | 23 + .../TritonToUnstructuredPass.cpp | 777 +++ .../UnstructuredToMemref/CMakeLists.txt | 27 + .../UnstructuredToMemrefPass.cpp | 439 ++ .../flir/lib/Dialect/CMakeLists.txt | 8 + .../flir/lib/Dialect/MathExt/CMakeLists.txt | 1 + .../lib/Dialect/MathExt/IR/CMakeLists.txt | 11 + .../lib/Dialect/MathExt/IR/MathExtDialect.cpp | 22 + .../lib/Dialect/MathExt/IR/MathExtOps.cpp | 54 + .../flir/lib/Dialect/TPtr/CMakeLists.txt | 1 + .../flir/lib/Dialect/TPtr/IR/CMakeLists.txt | 12 + .../flir/lib/Dialect/TPtr/IR/TPtrDialect.cpp | 54 + .../flir/lib/Dialect/TPtr/IR/TPtrOps.cpp | 44 + .../lib/Dialect/TritonAscend/CMakeLists.txt | 1 + .../Dialect/TritonAscend/IR/CMakeLists.txt | 14 + .../TritonAscend/IR/TritonAscendAttrs.cpp | 15 + .../TritonAscend/IR/TritonAscendDialect.cpp | 44 + .../TritonAscend/IR/TritonAscendOps.cpp | 137 + .../Dialect/TritonStructured/CMakeLists.txt | 1 + .../TritonStructured/IR/CMakeLists.txt | 11 + .../IR/TritonStructuredDialect.cpp | 22 + .../IR/TritonStructuredOps.cpp | 267 ++ .../TritonStructuredIncubated/CMakeLists.txt | 1 + .../IR/CMakeLists.txt | 11 + .../IR/TritonStructuredDialectIncubated.cpp | 22 + .../IR/TritonStructuredOpsIncubated.cpp | 111 + .../Dialect/TritonTilingExt/CMakeLists.txt | 1 + .../IR/BufferizableOpInterfaceImpl.cpp | 138 + .../Dialect/TritonTilingExt/IR/CMakeLists.txt | 17 + .../lib/Dialect/TritonTilingExt/IR/CumSum.cpp | 112 + .../IR/TritonTilingExtDialect.cpp | 404 ++ .../third_party/flir/lib/Utils/CMakeLists.txt | 6 + .../third_party/flir/lib/Utils/Utils.cpp | 305 ++ .../flir/lib/UtilsIncubated/CMakeLists.txt | 8 + .../UtilsIncubated/InterleaveOptimization.cpp | 677 +++ .../flir/lib/UtilsIncubated/Utils.cpp | 1254 +++++ .../flir/python/examples/bare_matmul.py | 45 + .../flir/python/examples/benchmark.py | 66 + .../flir/python/examples/conftest.py | 82 + .../flir/python/examples/test_addptr.py | 43 + .../examples/test_blockptr_complex_offset.py | 37 + .../flir/python/examples/test_early_return.py | 55 + .../python/examples/test_gather_scatter.py | 163 + .../flir/python/examples/test_layernorm.py | 160 + .../examples/test_load_2d_tensor_block.py | 78 + .../examples/test_load_2d_tensor_col.py | 70 + .../flir/python/examples/test_mask.py | 39 + .../flir/python/examples/test_matmul.py | 174 + .../flir/python/examples/test_modulo.py | 386 ++ .../flir/python/examples/test_nested_loops.py | 413 ++ .../flir/python/examples/test_reduce.py | 61 + .../flir/python/examples/test_scalar_store.py | 54 + .../flir/python/examples/test_sign_extend.py | 39 + .../flir/python/examples/test_softmax.py | 82 + .../flir/python/examples/test_splat.py | 41 + .../flir/python/examples/test_swap.py | 44 + .../examples/test_tensor_index_iterargs.py | 114 + .../flir/python/examples/test_vec_add.py | 85 + .../third_party/flir/test/CMakeLists.txt | 27 + .../wafer/third_party/flir/test/README.md | 2 + .../wafer/third_party/flir/test/lit.cfg.py | 74 + .../third_party/flir/test/lit.site.cfg.py.in | 24 + .../third_party/flir/tools/CMakeLists.txt | 1 + .../flir/tools/RegisterTritonSharedDialects.h | 70 + .../tools/triton-shared-opt/CMakeLists.txt | 21 + .../triton-shared-opt/triton-shared-opt.cpp | 18 + .../wafer/third_party/flir/triton_shared.cc | 8 + .../wafer/third_party/tle/CMakeLists.txt | 9 + third_party/wafer/third_party/tle/REANME.md | 0 .../third_party/tle/include/CMakeLists.txt | 1 + .../tle-dsa/Conversion/DsaToCore/DsaToCore.h | 17 + .../include/tle-dsa/Dialect/IR/CMakeLists.txt | 14 + .../include/tle-dsa/Dialect/IR/DsaDialect.h | 33 + .../include/tle-dsa/Dialect/IR/DsaDialect.td | 68 + .../tle/include/tle-dsa/Dialect/IR/DsaOps.td | 200 + .../wafer/third_party/tle/lib/CMakeLists.txt | 4 + .../lib/Conversion/DsaToCore/CMakeLists.txt | 17 + .../lib/Conversion/DsaToCore/DsaToCore.cpp | 72 + .../tle/lib/Dialect/IR/CMakeLists.txt | 19 + .../tle/lib/Dialect/IR/DsaDialect.cpp | 56 + .../third_party/tle/python/CMakeLists.txt | 7 + third_party/wafer/third_party/tle/python/ir.h | 91 + .../third_party/tle/python/triton_tle_dsa.cc | 241 + 985 files changed, 128809 insertions(+), 10 deletions(-) create mode 100644 .github/workflows/wafer-isolation.yml create mode 100644 backend/wafer.py create mode 100644 backend/wafer_cache.py create mode 100644 backend/wafer_runtime.py create mode 100644 scripts/wafer/apply_triton_profile.py create mode 100755 scripts/wafer/apply_wafer_triton_patches.sh create mode 100755 scripts/wafer/audit_wafer_elf.py create mode 100644 scripts/wafer/build_wafer_isolated.py create mode 100755 scripts/wafer/compile_wafer.sh create mode 100755 scripts/wafer/init_wafer_env.sh create mode 100755 scripts/wafer/install_wafer.sh create mode 100644 scripts/wafer/inventory_wafer_tests.py create mode 100755 scripts/wafer/migrate_wafer_env.sh create mode 100644 scripts/wafer/package_wafer.py create mode 100644 scripts/wafer/run_wafer_example_suite.py create mode 100755 scripts/wafer/setup_llvm22_env.sh create mode 100755 scripts/wafer/setup_wafer_env.sh create mode 100644 scripts/wafer/wafer_artifacts.py create mode 100755 scripts/wafer/wafer_manifest.py create mode 100644 scripts/wafer/wafer_pytest.py create mode 100755 setup_on_wafer.py create mode 100644 test/wafer/conftest.py create mode 100644 test/wafer/ir/interfaces-argmax2d.mlir create mode 100644 test/wafer/ir/interfaces-flip.mlir create mode 100644 test/wafer/ir/interfaces-sort.mlir create mode 100644 test/wafer/ir/pointer-state-modulo.mlir create mode 100644 test/wafer/ir/pointer-state-nested_loops.mlir create mode 100644 test/wafer/ir/pointer-state-scalar_store.mlir create mode 100644 test/wafer/ir/pointer-state-tensor_index_iterargs.mlir create mode 100644 test/wafer/native_math/__init__.py create mode 100644 test/wafer/native_math/conftest.py create mode 100644 test/wafer/native_math/test_log1p.py create mode 100644 test/wafer/native_math/test_multi_return.py create mode 100644 test/wafer/native_math/test_relu.py create mode 100644 test/wafer/native_math/test_unary.py create mode 100644 test/wafer/ops/conftest.py create mode 100644 test/wafer/ops/test_2d_permute.py create mode 100644 test/wafer/ops/test_3Dgrid.py create mode 100644 test/wafer/ops/test_abs.py create mode 100644 test/wafer/ops/test_abs_2.py create mode 100644 test/wafer/ops/test_add.py create mode 100644 test/wafer/ops/test_add_multi_return.py create mode 100644 test/wafer/ops/test_advance.py create mode 100644 test/wafer/ops/test_and.py create mode 100644 test/wafer/ops/test_arange.py create mode 100644 test/wafer/ops/test_associative_scan.py create mode 100644 test/wafer/ops/test_associative_scan_multi_input.py create mode 100644 test/wafer/ops/test_block_ptr.py create mode 100644 test/wafer/ops/test_broadcast_op.py create mode 100644 test/wafer/ops/test_cat_dim.py create mode 100644 test/wafer/ops/test_cdiv.py create mode 100644 test/wafer/ops/test_ceil.py create mode 100644 test/wafer/ops/test_clamp.py create mode 100644 test/wafer/ops/test_common.py create mode 100644 test/wafer/ops/test_conv.py create mode 100644 test/wafer/ops/test_cos.py create mode 100644 test/wafer/ops/test_cos_2.py create mode 100644 test/wafer/ops/test_count_dim0.py create mode 100644 test/wafer/ops/test_count_dim1.py create mode 100644 test/wafer/ops/test_cumprod.py create mode 100644 test/wafer/ops/test_cumsum.py create mode 100644 test/wafer/ops/test_debug_barrier.py create mode 100644 test/wafer/ops/test_device_print.py create mode 100644 test/wafer/ops/test_div.py create mode 100644 test/wafer/ops/test_elementwise_ceil.py create mode 100644 test/wafer/ops/test_elementwise_clip.py create mode 100644 test/wafer/ops/test_elementwise_f2i.py create mode 100644 test/wafer/ops/test_elementwise_floor.py create mode 100644 test/wafer/ops/test_elementwise_i2f.py create mode 100644 test/wafer/ops/test_elementwise_round.py create mode 100644 test/wafer/ops/test_eq.py create mode 100644 test/wafer/ops/test_eq_2.py create mode 100644 test/wafer/ops/test_exp.py create mode 100644 test/wafer/ops/test_exp2.py create mode 100644 test/wafer/ops/test_exp_.py create mode 100644 test/wafer/ops/test_expand_dims.py create mode 100644 test/wafer/ops/test_extract_slice.py create mode 100644 test/wafer/ops/test_fdiv.py create mode 100644 test/wafer/ops/test_floor.py create mode 100644 test/wafer/ops/test_floordiv.py create mode 100644 test/wafer/ops/test_full.py create mode 100644 test/wafer/ops/test_ge.py create mode 100644 test/wafer/ops/test_ge_2.py create mode 100644 test/wafer/ops/test_gelu.py create mode 100644 test/wafer/ops/test_gt.py create mode 100644 test/wafer/ops/test_hd_permute.py create mode 100644 test/wafer/ops/test_if_tensor.py create mode 100644 test/wafer/ops/test_insert_slice.py create mode 100644 test/wafer/ops/test_interleave.py create mode 100644 test/wafer/ops/test_invert.py create mode 100644 test/wafer/ops/test_join.py create mode 100644 test/wafer/ops/test_lanzcos.py create mode 100644 test/wafer/ops/test_launcher_empty_signature.py create mode 100644 test/wafer/ops/test_layernorm.py create mode 100644 test/wafer/ops/test_ldst.py create mode 100644 test/wafer/ops/test_le.py create mode 100644 test/wafer/ops/test_load.py create mode 100644 test/wafer/ops/test_load_store.py create mode 100644 test/wafer/ops/test_log.py create mode 100644 test/wafer/ops/test_log2.py create mode 100644 test/wafer/ops/test_log_2.py create mode 100644 test/wafer/ops/test_logical_and.py create mode 100644 test/wafer/ops/test_logical_or.py create mode 100644 test/wafer/ops/test_lshift.py create mode 100644 test/wafer/ops/test_lt.py create mode 100644 test/wafer/ops/test_max_dim0.py create mode 100644 test/wafer/ops/test_max_dim1.py create mode 100644 test/wafer/ops/test_max_vector.py create mode 100644 test/wafer/ops/test_maximum.py create mode 100644 test/wafer/ops/test_mean_dim0.py create mode 100644 test/wafer/ops/test_mean_dim1.py create mode 100644 test/wafer/ops/test_mean_vector.py create mode 100644 test/wafer/ops/test_min_dim0.py create mode 100644 test/wafer/ops/test_min_dim1.py create mode 100644 test/wafer/ops/test_min_vector.py create mode 100644 test/wafer/ops/test_minimum.py create mode 100644 test/wafer/ops/test_mod.py create mode 100644 test/wafer/ops/test_mul.py create mode 100644 test/wafer/ops/test_nearest.py create mode 100644 test/wafer/ops/test_neg.py create mode 100644 test/wafer/ops/test_npu_indexing.py create mode 100644 test/wafer/ops/test_npu_indexing2.py create mode 100644 test/wafer/ops/test_or.py create mode 100644 test/wafer/ops/test_permute.py create mode 100644 test/wafer/ops/test_permute_full.py create mode 100644 test/wafer/ops/test_permute_reshape.py create mode 100644 test/wafer/ops/test_precise_div.py create mode 100644 test/wafer/ops/test_precise_sqrt.py create mode 100644 test/wafer/ops/test_ravel.py create mode 100644 test/wafer/ops/test_reduce_count_vector.py create mode 100644 test/wafer/ops/test_reduce_mean.py create mode 100644 test/wafer/ops/test_reduce_sum.py create mode 100644 test/wafer/ops/test_reshape.py create mode 100644 test/wafer/ops/test_rms_norm.py create mode 100644 test/wafer/ops/test_rotary_embedding.py create mode 100644 test/wafer/ops/test_rotatry_gpt.py create mode 100644 test/wafer/ops/test_rotaty_embedding_gpt.py create mode 100644 test/wafer/ops/test_rshift.py create mode 100644 test/wafer/ops/test_rsqrt.py create mode 100644 test/wafer/ops/test_scalar_calc.py create mode 100644 test/wafer/ops/test_sigmoid.py create mode 100644 test/wafer/ops/test_silu.py create mode 100644 test/wafer/ops/test_silu_and_mul.py create mode 100644 test/wafer/ops/test_sin.py create mode 100644 test/wafer/ops/test_softmax.py create mode 100644 test/wafer/ops/test_split.py create mode 100644 test/wafer/ops/test_sqrt.py create mode 100644 test/wafer/ops/test_store_scalar.py create mode 100644 test/wafer/ops/test_sub.py create mode 100644 test/wafer/ops/test_sum.py create mode 100644 test/wafer/ops/test_sum_dim0.py create mode 100644 test/wafer/ops/test_sum_dim1.py create mode 100644 test/wafer/ops/test_sum_vector.py create mode 100644 test/wafer/ops/test_swap.py create mode 100644 test/wafer/ops/test_swiglu.py create mode 100644 test/wafer/ops/test_swizzle2d.py create mode 100644 test/wafer/ops/test_template.py create mode 100644 test/wafer/ops/test_tensor_get_item.py create mode 100644 test/wafer/ops/test_trans_3d.py create mode 100644 test/wafer/ops/test_triton_eq.py create mode 100644 test/wafer/ops/test_triton_le.py create mode 100644 test/wafer/ops/test_triton_lt.py create mode 100644 test/wafer/ops/test_triton_neq.py create mode 100644 test/wafer/ops/test_umulhi.py create mode 100644 test/wafer/ops/test_unlign_sum.py create mode 100644 test/wafer/ops/test_unused_func_arg.py create mode 100644 test/wafer/ops/test_view.py create mode 100644 test/wafer/ops/test_where_lt.py create mode 100644 test/wafer/ops/test_where_mask.py create mode 100644 test/wafer/ops/test_where_var.py create mode 100644 test/wafer/ops/test_xor.py create mode 100644 test/wafer/ops/test_xor_sum.py create mode 100644 test/wafer/ops/test_zeros.py create mode 100644 test/wafer/ops/test_zeroslike.py create mode 100644 test/wafer/runtime/conftest.py create mode 100644 test/wafer/runtime/test_autotune.py create mode 100644 test/wafer/runtime/test_elementwise.py create mode 100644 test/wafer/runtime/test_matmul.py create mode 100644 test/wafer/runtime/test_reduction.py create mode 100644 test/wafer/runtime/test_tensor_runtime.py create mode 100644 test/wafer/suites/accepted.txt create mode 100644 test/wafer/test_benchmark_interface.py create mode 100644 test/wafer/test_cache_and_runtime.py create mode 100644 test/wafer/test_cluster_launch.py create mode 100644 test/wafer/test_commonir_abi.py create mode 100644 test/wafer/test_crt_diagnostics.py create mode 100644 test/wafer/test_crt_fp8.py create mode 100644 test/wafer/test_crt_randgen.py create mode 100644 test/wafer/test_dense_constants.py create mode 100644 test/wafer/test_device_print_abi.py create mode 100644 test/wafer/test_elf_audit.py create mode 100644 test/wafer/test_external_interfaces.py create mode 100644 test/wafer/test_frontend_isolation.py create mode 100644 test/wafer/test_loader_isolation.py create mode 100644 test/wafer/test_log1p_lowering.py create mode 100644 test/wafer/test_module_runtime.py create mode 100644 test/wafer/test_mxfp_reference.py create mode 100644 test/wafer/test_naming_compat.py create mode 100644 test/wafer/test_native_errors.py create mode 100644 test/wafer/test_noc_sync.py create mode 100644 test/wafer/test_patch_profiles.py create mode 100644 test/wafer/test_precision_modes.py create mode 100644 test/wafer/test_scalar_copy.py create mode 100644 test/wafer/test_tle_frontend.py create mode 100644 test/wafer/verify_wafer_acceptance.py create mode 100644 test/wafer/verify_wafer_examples.py create mode 100644 test/wafer/verify_wafer_installation.py create mode 100755 test/wafer/verify_wafer_runtime.py create mode 100755 test/wafer/verify_wafer_runtime.sh create mode 100755 test/wafer/verify_wafer_torch_stack.py create mode 100755 third_party/wafer/CMakeLists.txt create mode 100755 third_party/wafer/README.md create mode 100755 third_party/wafer/backend/__init__.py create mode 100755 third_party/wafer/backend/compiler.py create mode 100755 third_party/wafer/backend/driver.py create mode 100755 third_party/wafer/backend/include/logger.h create mode 100755 third_party/wafer/backend/logger_config.py create mode 100755 third_party/wafer/backend/name.conf create mode 100644 third_party/wafer/backend/txda_tools.py create mode 100755 third_party/wafer/backend/wafer_tools.py create mode 100755 third_party/wafer/benchmark/benchmark.py create mode 100755 third_party/wafer/bin/CMakeLists.txt create mode 100755 third_party/wafer/bin/RegisterTritonDialects.h create mode 100755 third_party/wafer/bin/wafer-llvm-opt.cpp create mode 100755 third_party/wafer/bin/wafer-lsp.cpp create mode 100755 third_party/wafer/bin/wafer-opt.cpp create mode 100755 third_party/wafer/bin/wafer-reduce.cpp create mode 100755 third_party/wafer/bin/wafer-tensor-layout.cpp create mode 100644 third_party/wafer/cmake/WaferConfig.cmake create mode 100755 third_party/wafer/crt/CMakeLists.txt create mode 100755 third_party/wafer/crt/README.md create mode 100755 third_party/wafer/crt/include/Wafer/op_gelu.h create mode 100755 third_party/wafer/crt/include/Wafer/op_reduce_mul_impl.h create mode 100755 third_party/wafer/crt/include/Wafer/wafer.h create mode 100755 third_party/wafer/crt/lib/Wafer/abs.c create mode 100755 third_party/wafer/crt/lib/Wafer/argmax.c create mode 100755 third_party/wafer/crt/lib/Wafer/argmin.c create mode 100755 third_party/wafer/crt/lib/Wafer/arith.c create mode 100755 third_party/wafer/crt/lib/Wafer/assert.c create mode 100755 third_party/wafer/crt/lib/Wafer/atomic_barrier_in.c create mode 100755 third_party/wafer/crt/lib/Wafer/atomic_barrier_out.c create mode 100755 third_party/wafer/crt/lib/Wafer/barrier.c create mode 100755 third_party/wafer/crt/lib/Wafer/bf16_fp16.c create mode 100755 third_party/wafer/crt/lib/Wafer/bf16_fp32.c create mode 100755 third_party/wafer/crt/lib/Wafer/bf16_int16.c create mode 100755 third_party/wafer/crt/lib/Wafer/bf16_int32.c create mode 100755 third_party/wafer/crt/lib/Wafer/bf16_int8.c create mode 100755 third_party/wafer/crt/lib/Wafer/bf16_tf32.c create mode 100755 third_party/wafer/crt/lib/Wafer/bilinear.c create mode 100755 third_party/wafer/crt/lib/Wafer/bit2fp.c create mode 100755 third_party/wafer/crt/lib/Wafer/channelnorm.c create mode 100755 third_party/wafer/crt/lib/Wafer/common.c create mode 100755 third_party/wafer/crt/lib/Wafer/concat.c create mode 100755 third_party/wafer/crt/lib/Wafer/conv.c create mode 100755 third_party/wafer/crt/lib/Wafer/cos.c create mode 100755 third_party/wafer/crt/lib/Wafer/count.c create mode 100755 third_party/wafer/crt/lib/Wafer/empty.c create mode 100755 third_party/wafer/crt/lib/Wafer/exp.c create mode 100755 third_party/wafer/crt/lib/Wafer/explp.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp16_bf16.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp16_fp32.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp16_int16.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp16_int32.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp16_int8.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp16_tf32.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp32_bf16.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp32_fp16.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp32_int16.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp32_int32.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp32_int8.c create mode 100755 third_party/wafer/crt/lib/Wafer/fp32_tf32.c create mode 100755 third_party/wafer/crt/lib/Wafer/gatherscatter.c create mode 100755 third_party/wafer/crt/lib/Wafer/gelu_none.c create mode 100755 third_party/wafer/crt/lib/Wafer/gelu_tanh.c create mode 100755 third_party/wafer/crt/lib/Wafer/gemm.c create mode 100755 third_party/wafer/crt/lib/Wafer/img2col.c create mode 100755 third_party/wafer/crt/lib/Wafer/int16_bf16.c create mode 100755 third_party/wafer/crt/lib/Wafer/int16_fp16.c create mode 100755 third_party/wafer/crt/lib/Wafer/int16_fp32.c create mode 100755 third_party/wafer/crt/lib/Wafer/int16_tf32.c create mode 100755 third_party/wafer/crt/lib/Wafer/int32_bf16.c create mode 100755 third_party/wafer/crt/lib/Wafer/int32_fp16.c create mode 100755 third_party/wafer/crt/lib/Wafer/int32_fp32.c create mode 100755 third_party/wafer/crt/lib/Wafer/int32_tf32.c create mode 100755 third_party/wafer/crt/lib/Wafer/int8_bf16.c create mode 100755 third_party/wafer/crt/lib/Wafer/int8_fp16.c create mode 100755 third_party/wafer/crt/lib/Wafer/int8_fp32.c create mode 100755 third_party/wafer/crt/lib/Wafer/int8_tf32.c create mode 100755 third_party/wafer/crt/lib/Wafer/leakyrelu.c create mode 100755 third_party/wafer/crt/lib/Wafer/ln.c create mode 100755 third_party/wafer/crt/lib/Wafer/log2.c create mode 100755 third_party/wafer/crt/lib/Wafer/logic.c create mode 100755 third_party/wafer/crt/lib/Wafer/lut16.c create mode 100755 third_party/wafer/crt/lib/Wafer/lut32.c create mode 100755 third_party/wafer/crt/lib/Wafer/mask_move.c create mode 100755 third_party/wafer/crt/lib/Wafer/memcpy.c create mode 100755 third_party/wafer/crt/lib/Wafer/memset.c create mode 100755 third_party/wafer/crt/lib/Wafer/mirror.c create mode 100755 third_party/wafer/crt/lib/Wafer/mxfp_bf16.c create mode 100755 third_party/wafer/crt/lib/Wafer/mxfp_fp16.c create mode 100755 third_party/wafer/crt/lib/Wafer/mxfp_scale_bf16.c create mode 100755 third_party/wafer/crt/lib/Wafer/mxfp_scale_fp16.c create mode 100755 third_party/wafer/crt/lib/Wafer/nchw2nhwc.c create mode 100755 third_party/wafer/crt/lib/Wafer/neg.c create mode 100755 third_party/wafer/crt/lib/Wafer/nhwc2nchw.c create mode 100644 third_party/wafer/crt/lib/Wafer/noc_init.c create mode 100755 third_party/wafer/crt/lib/Wafer/op_gelu.c create mode 100755 third_party/wafer/crt/lib/Wafer/op_reduce_mul_impl.c create mode 100755 third_party/wafer/crt/lib/Wafer/pad.c create mode 100755 third_party/wafer/crt/lib/Wafer/pow.c create mode 100755 third_party/wafer/crt/lib/Wafer/pow2.c create mode 100755 third_party/wafer/crt/lib/Wafer/print.c create mode 100755 third_party/wafer/crt/lib/Wafer/randgen.c create mode 100755 third_party/wafer/crt/lib/Wafer/rdma.c create mode 100755 third_party/wafer/crt/lib/Wafer/recip.c create mode 100755 third_party/wafer/crt/lib/Wafer/recv.c create mode 100755 third_party/wafer/crt/lib/Wafer/reduce.c create mode 100755 third_party/wafer/crt/lib/Wafer/relation.c create mode 100755 third_party/wafer/crt/lib/Wafer/relu.c create mode 100755 third_party/wafer/crt/lib/Wafer/rotate180.c create mode 100755 third_party/wafer/crt/lib/Wafer/rotate270.c create mode 100755 third_party/wafer/crt/lib/Wafer/rotate90.c create mode 100755 third_party/wafer/crt/lib/Wafer/rsqrt.c create mode 100755 third_party/wafer/crt/lib/Wafer/satrelu.c create mode 100755 third_party/wafer/crt/lib/Wafer/send.c create mode 100755 third_party/wafer/crt/lib/Wafer/sigmoid.c create mode 100755 third_party/wafer/crt/lib/Wafer/sin.c create mode 100755 third_party/wafer/crt/lib/Wafer/softplus.c create mode 100755 third_party/wafer/crt/lib/Wafer/sqrt.c create mode 100755 third_party/wafer/crt/lib/Wafer/tanh.c create mode 100755 third_party/wafer/crt/lib/Wafer/tensornorm.c create mode 100755 third_party/wafer/crt/lib/Wafer/tf32_bf16.c create mode 100755 third_party/wafer/crt/lib/Wafer/tf32_fp16.c create mode 100755 third_party/wafer/crt/lib/Wafer/tf32_fp32.c create mode 100755 third_party/wafer/crt/lib/Wafer/tf32_int16.c create mode 100755 third_party/wafer/crt/lib/Wafer/tf32_int32.c create mode 100755 third_party/wafer/crt/lib/Wafer/tf32_int8.c create mode 100755 third_party/wafer/crt/lib/Wafer/transpose.c create mode 100755 third_party/wafer/crt/lib/Wafer/wafer.c create mode 100755 third_party/wafer/crt/lib/Wafer/wdma.c create mode 100644 third_party/wafer/examples/_wafer_reference.py create mode 100755 third_party/wafer/examples/bare_matmul.py create mode 100755 third_party/wafer/examples/bare_matmul_acc.py create mode 100755 third_party/wafer/examples/bare_matmul_autotune.py create mode 100755 third_party/wafer/examples/benchmark.py create mode 100755 third_party/wafer/examples/conftest.py create mode 100755 third_party/wafer/examples/dump_vec_add_ir.sh create mode 100755 third_party/wafer/examples/embedding.py create mode 100644 third_party/wafer/examples/flagtree/test_tle_cumsum.py create mode 100644 third_party/wafer/examples/flagtree/test_tle_dsa_arith.py create mode 100644 third_party/wafer/examples/flagtree/test_tle_dsa_bridge.py create mode 100644 third_party/wafer/examples/flagtree/test_tle_dsa_pipeline_e2e.py create mode 100644 third_party/wafer/examples/flagtree/test_tle_dsa_rand.py create mode 100644 third_party/wafer/examples/flagtree/test_tle_dsa_slice.py create mode 100755 third_party/wafer/examples/mult_ir.py create mode 100755 third_party/wafer/examples/profile_matmul.py create mode 100755 third_party/wafer/examples/quant_gptq.py create mode 100755 third_party/wafer/examples/quant_kernel.py create mode 100755 third_party/wafer/examples/single_conv2d.py create mode 100755 third_party/wafer/examples/test_abs.py create mode 100755 third_party/wafer/examples/test_addptr.py create mode 100755 third_party/wafer/examples/test_argmax2d.py create mode 100755 third_party/wafer/examples/test_assert.py create mode 100644 third_party/wafer/examples/test_autotune.py create mode 100755 third_party/wafer/examples/test_bare_matmul.py create mode 100755 third_party/wafer/examples/test_bare_matmul_acc.py create mode 100755 third_party/wafer/examples/test_blockptr_complex_offset.py create mode 100755 third_party/wafer/examples/test_cdiv.py create mode 100755 third_party/wafer/examples/test_ceil.py create mode 100755 third_party/wafer/examples/test_clamp.py create mode 100755 third_party/wafer/examples/test_cos.py create mode 100755 third_party/wafer/examples/test_debug_barrier.py create mode 100755 third_party/wafer/examples/test_div_rn.py create mode 100755 third_party/wafer/examples/test_dot_scaled.py create mode 100755 third_party/wafer/examples/test_early_return.py create mode 100755 third_party/wafer/examples/test_embedding.py create mode 100755 third_party/wafer/examples/test_exp.py create mode 100755 third_party/wafer/examples/test_exp2.py create mode 100755 third_party/wafer/examples/test_fdiv.py create mode 100755 third_party/wafer/examples/test_flip.py create mode 100755 third_party/wafer/examples/test_floor.py create mode 100755 third_party/wafer/examples/test_fma.py create mode 100644 third_party/wafer/examples/test_fp8_conversion.py create mode 100755 third_party/wafer/examples/test_gather.py create mode 100755 third_party/wafer/examples/test_histogram.py create mode 100755 third_party/wafer/examples/test_layernorm.py create mode 100755 third_party/wafer/examples/test_libdevice.py create mode 100755 third_party/wafer/examples/test_load6d.py create mode 100755 third_party/wafer/examples/test_load_2d_tensor_block.py create mode 100755 third_party/wafer/examples/test_load_2d_tensor_col.py create mode 100755 third_party/wafer/examples/test_load_store_mod.py create mode 100755 third_party/wafer/examples/test_log.py create mode 100755 third_party/wafer/examples/test_log2.py create mode 100755 third_party/wafer/examples/test_mask.py create mode 100755 third_party/wafer/examples/test_math_erf_op.py create mode 100755 third_party/wafer/examples/test_matmul.py create mode 100755 third_party/wafer/examples/test_maximum.py create mode 100755 third_party/wafer/examples/test_minimum.py create mode 100755 third_party/wafer/examples/test_modulo.py create mode 100755 third_party/wafer/examples/test_nested_loops.py create mode 100644 third_party/wafer/examples/test_pipeline.py create mode 100644 third_party/wafer/examples/test_precision_modes.py create mode 100755 third_party/wafer/examples/test_print.py create mode 100755 third_party/wafer/examples/test_reduce.py create mode 100755 third_party/wafer/examples/test_reduce1d.py create mode 100755 third_party/wafer/examples/test_rsqrt.py create mode 100755 third_party/wafer/examples/test_scalar_store.py create mode 100755 third_party/wafer/examples/test_scan.py create mode 100755 third_party/wafer/examples/test_scan2d.py create mode 100755 third_party/wafer/examples/test_scan3d.py create mode 100755 third_party/wafer/examples/test_sigmoid.py create mode 100755 third_party/wafer/examples/test_sign_extend.py create mode 100755 third_party/wafer/examples/test_sin.py create mode 100755 third_party/wafer/examples/test_softmax.py create mode 100755 third_party/wafer/examples/test_sort.py create mode 100755 third_party/wafer/examples/test_splat.py create mode 100755 third_party/wafer/examples/test_sqrt.py create mode 100755 third_party/wafer/examples/test_sqrt_rn.py create mode 100755 third_party/wafer/examples/test_swap.py create mode 100755 third_party/wafer/examples/test_swizzle2d.py create mode 100755 third_party/wafer/examples/test_tensor_index_iterargs.py create mode 100755 third_party/wafer/examples/test_trans2d.py create mode 100755 third_party/wafer/examples/test_umulhi.py create mode 100755 third_party/wafer/examples/test_vec_add.py create mode 100755 third_party/wafer/examples/test_where.py create mode 100755 third_party/wafer/examples/time1.py create mode 100755 third_party/wafer/examples/time_zs_opt2.py create mode 100755 third_party/wafer/examples/tle/test_tle_dsa_noc_gemm_4096.py create mode 100755 third_party/wafer/examples/util.py create mode 100755 third_party/wafer/examples/view_vec_add_ir.sh create mode 100644 third_party/wafer/experimental/README.md create mode 100644 third_party/wafer/experimental/tle/__init__.py create mode 100644 third_party/wafer/experimental/tle/language/__init__.py create mode 100644 third_party/wafer/experimental/tle/language/core.py create mode 100644 third_party/wafer/experimental/tle/language/distributed.py create mode 100644 third_party/wafer/experimental/tle/language/dsa/__init__.py create mode 100644 third_party/wafer/experimental/tle/language/dsa/core.py create mode 100644 third_party/wafer/experimental/tle/language/dsa/semantic.py create mode 100644 third_party/wafer/experimental/tle/language/dsa/types.py create mode 100644 third_party/wafer/experimental/tle/language/dsa/wafer/__init__.py create mode 100644 third_party/wafer/experimental/tle/language/dsa/wafer/core.py create mode 100755 third_party/wafer/include/Address/CMakeLists.txt create mode 100755 third_party/wafer/include/Address/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/include/Address/Dialect/IR/AddressDialect.h create mode 100755 third_party/wafer/include/Address/Dialect/IR/AddressDialect.td create mode 100755 third_party/wafer/include/Address/Dialect/IR/AddressOps.td create mode 100755 third_party/wafer/include/Address/Dialect/IR/CMakeLists.txt create mode 100755 third_party/wafer/include/Address/Transforms/CMakeLists.txt create mode 100755 third_party/wafer/include/Address/Transforms/Passes.h create mode 100755 third_party/wafer/include/Address/Transforms/Passes.td create mode 100755 third_party/wafer/include/Analysis/Alias.h create mode 100755 third_party/wafer/include/Analysis/Allocation.h create mode 100644 third_party/wafer/include/Analysis/Membar.h create mode 100755 third_party/wafer/include/Analysis/Utility.h create mode 100755 third_party/wafer/include/CMakeLists.txt create mode 100755 third_party/wafer/include/ExecutionEngine/CRunnerUtils.cpp create mode 100755 third_party/wafer/include/ExecutionEngine/CRunnerUtils.h create mode 100755 third_party/wafer/include/ExecutionEngine/Msan.h create mode 100755 third_party/wafer/include/ExecutionEngine/version.txt create mode 100755 third_party/wafer/include/flagtree/Common/UnifiedHardware.h create mode 100755 third_party/wafer/include/magic-kernel-func/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel-func/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel-func/Dialect/IR/MagicKernelFuncOps.td create mode 100755 third_party/wafer/include/magic-kernel-instr/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel-instr/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel-instr/Dialect/IR/MagicKernelInstrOps.td create mode 100755 third_party/wafer/include/magic-kernel/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel/Conversion/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/CoreDialectsToMK.h create mode 100755 third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/Passes.h create mode 100755 third_party/wafer/include/magic-kernel/Conversion/CoreDialectsToMK/Passes.td create mode 100755 third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.h create mode 100755 third_party/wafer/include/magic-kernel/Conversion/LegalizeTensorFormLoops/Passes.td create mode 100755 third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/LinalgToMK.h create mode 100755 third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/Passes.h create mode 100755 third_party/wafer/include/magic-kernel/Conversion/LinalgToMK/Passes.td create mode 100644 third_party/wafer/include/magic-kernel/Conversion/MKPipeline/CMakeLists.txt create mode 100644 third_party/wafer/include/magic-kernel/Conversion/MKPipeline/Passes.h create mode 100644 third_party/wafer/include/magic-kernel/Conversion/MKPipeline/Passes.td create mode 100755 third_party/wafer/include/magic-kernel/Conversion/TLEToMK/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel/Conversion/TLEToMK/Passes.h create mode 100755 third_party/wafer/include/magic-kernel/Conversion/TLEToMK/Passes.td create mode 100755 third_party/wafer/include/magic-kernel/Conversion/TLEToMK/TLEToMK.h create mode 100755 third_party/wafer/include/magic-kernel/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel/Dialect/IR/CMakeLists.txt create mode 100755 third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelAttrDefs.td create mode 100755 third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelDialect.h create mode 100755 third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelDialect.td create mode 100755 third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelOps.td create mode 100755 third_party/wafer/include/magic-kernel/Dialect/IR/MagicKernelTypes.td create mode 100755 third_party/wafer/include/magic-kernel/Transforms/BufferizableOpInterfaceImpl.h create mode 100644 third_party/wafer/include/magic-kernel/Transforms/CMakeLists.txt create mode 100644 third_party/wafer/include/magic-kernel/Transforms/Passes.h create mode 100644 third_party/wafer/include/magic-kernel/Transforms/Passes.td create mode 100755 third_party/wafer/include/triton-shared/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/ConvertTritonPtr/TritonPtrToAddress.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/StructuredToMK/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/StructuredToMK/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/StructuredToMK/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/StructuredToMK/StructuredToMK.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToCoreDialects/TritonToCoreDialects.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToMK/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToMK/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToMK/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToMK/TritonToMKPatterns.hpp create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/UnstructuredToMK/UnstructuredToMK.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/CMakeLists.txt create mode 100755 third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/Passes.h create mode 100755 third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/Passes.td create mode 100755 third_party/wafer/include/triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h create mode 100755 third_party/wafer/include/utils/LinalgOpBuilderHelper.h create mode 100755 third_party/wafer/include/utils/TypeConvertor.h create mode 100755 third_party/wafer/include/wafer/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/Passes.h create mode 100755 third_party/wafer/include/wafer/Conversion/AllocateSharedMemory/Passes.td create mode 100755 third_party/wafer/include/wafer/Conversion/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/ExportKernelSymbols.h create mode 100755 third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/Passes.h create mode 100755 third_party/wafer/include/wafer/Conversion/ExportKernelSymbols/Passes.td create mode 100755 third_party/wafer/include/wafer/Conversion/LinalgFusion/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/LinalgFusion/LinalgFusion.h create mode 100755 third_party/wafer/include/wafer/Conversion/LinalgFusion/Passes.h create mode 100755 third_party/wafer/include/wafer/Conversion/LinalgFusion/Passes.td create mode 100755 third_party/wafer/include/wafer/Conversion/LinalgTiling/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/LinalgTiling/LinalgTiling.h create mode 100755 third_party/wafer/include/wafer/Conversion/LinalgTiling/Passes.h create mode 100755 third_party/wafer/include/wafer/Conversion/LinalgTiling/Passes.td create mode 100755 third_party/wafer/include/wafer/Conversion/MKToWafer/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/MKToWafer/MKToWafer.h create mode 100755 third_party/wafer/include/wafer/Conversion/MKToWafer/Passes.h create mode 100755 third_party/wafer/include/wafer/Conversion/MKToWafer/Passes.td create mode 100755 third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/Passes.h create mode 100755 third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/Passes.td create mode 100755 third_party/wafer/include/wafer/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.h create mode 100755 third_party/wafer/include/wafer/Conversion/WaferToLLVM/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Conversion/WaferToLLVM/KernelArgBufferPass.h create mode 100755 third_party/wafer/include/wafer/Conversion/WaferToLLVM/KernelArgBufferPass.td create mode 100755 third_party/wafer/include/wafer/Conversion/WaferToLLVM/Passes.h create mode 100755 third_party/wafer/include/wafer/Conversion/WaferToLLVM/Passes.td create mode 100755 third_party/wafer/include/wafer/Conversion/WaferToLLVM/WaferToLLVM.h create mode 100755 third_party/wafer/include/wafer/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Dialect/IR/CMakeLists.txt create mode 100755 third_party/wafer/include/wafer/Dialect/IR/WaferAttrDefs.td create mode 100755 third_party/wafer/include/wafer/Dialect/IR/WaferDialect.h create mode 100755 third_party/wafer/include/wafer/Dialect/IR/WaferDialect.td create mode 100755 third_party/wafer/include/wafer/Dialect/IR/WaferOps.h create mode 100755 third_party/wafer/include/wafer/Dialect/IR/WaferOps.td create mode 100755 third_party/wafer/include/wafer/Dialect/IR/WaferTypes.td create mode 100644 third_party/wafer/include/wafer/Transforms/CMakeLists.txt create mode 100644 third_party/wafer/include/wafer/Transforms/Passes.h create mode 100644 third_party/wafer/include/wafer/Transforms/Passes.td create mode 100755 third_party/wafer/language/cpu/__init__.py create mode 100755 third_party/wafer/language/cpu/libdevice.py create mode 100644 third_party/wafer/language/txda/__init__.py create mode 100644 third_party/wafer/language/txda/libdevice.py create mode 100755 third_party/wafer/language/wafer/__init__.py create mode 100755 third_party/wafer/language/wafer/libdevice.py create mode 100644 third_party/wafer/language/wafer/slicing.py create mode 100755 third_party/wafer/lib/Analysis/Alias.cpp create mode 100755 third_party/wafer/lib/Analysis/Allocation.cpp create mode 100755 third_party/wafer/lib/Analysis/CMakeLists.txt create mode 100644 third_party/wafer/lib/Analysis/Membar.cpp create mode 100755 third_party/wafer/lib/CMakeLists.txt create mode 100755 third_party/wafer/lib/Common/CMakeLists.txt create mode 100755 third_party/wafer/lib/Common/UnifiedHardware.cc create mode 100755 third_party/wafer/lib/Conversion/AllocateSharedMemory/AllocateSharedMemoryPass.cpp create mode 100755 third_party/wafer/lib/Conversion/AllocateSharedMemory/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/ConvertTritonPtr/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/ConvertTritonPtr/TritonPtrToAddressPass.cpp create mode 100755 third_party/wafer/lib/Conversion/CoreDialectsToMK/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/CoreDialectsToMK/CoreDialectsToMKPass.cpp create mode 100755 third_party/wafer/lib/Conversion/ExportKernelSymbols/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/ExportKernelSymbols/ExportKernelSymbols.cpp create mode 100755 third_party/wafer/lib/Conversion/LegalizeTensorFormLoops/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/LegalizeTensorFormLoops/LegalizeTensorFormLoops.cpp create mode 100755 third_party/wafer/lib/Conversion/LinalgFusion/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/LinalgFusion/LinalgFusion.cpp create mode 100755 third_party/wafer/lib/Conversion/LinalgFusion/LinalgFusionPass.cpp create mode 100755 third_party/wafer/lib/Conversion/LinalgTiling/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/LinalgTiling/LinalgTiling.cpp create mode 100755 third_party/wafer/lib/Conversion/LinalgTiling/LinalgTilingPass.cpp create mode 100755 third_party/wafer/lib/Conversion/LinalgToMK/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/LinalgToMK/LinalgToMK.cpp create mode 100755 third_party/wafer/lib/Conversion/LinalgToMK/LinalgToMKPass.cpp create mode 100644 third_party/wafer/lib/Conversion/MKPipeline/CMakeLists.txt create mode 100644 third_party/wafer/lib/Conversion/MKPipeline/MKLoopBoundCanonicalizePass.cpp create mode 100644 third_party/wafer/lib/Conversion/MKPipeline/MKPipelinePass.cpp create mode 100755 third_party/wafer/lib/Conversion/MKToWafer/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/MKToWafer/MKToWafer.cpp create mode 100755 third_party/wafer/lib/Conversion/MKToWafer/MKToWaferPass.cpp create mode 100755 third_party/wafer/lib/Conversion/ReconcilePtrCasts/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/ReconcilePtrCasts/ReconcilePtrCastsPass.cpp create mode 100755 third_party/wafer/lib/Conversion/StructuredToMK/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/StructuredToMK/StructuredToMK.cpp create mode 100755 third_party/wafer/lib/Conversion/StructuredToMK/StructuredToMKPass.cpp create mode 100755 third_party/wafer/lib/Conversion/StructuredToMemref/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/StructuredToMemref/StructuredToMemref.cpp create mode 100755 third_party/wafer/lib/Conversion/StructuredToMemref/StructuredToMemrefPass.cpp create mode 100755 third_party/wafer/lib/Conversion/TLEToMK/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/TLEToMK/MKCommonBufferPlanningPass.cpp create mode 100755 third_party/wafer/lib/Conversion/TLEToMK/TLEToMK.cpp create mode 100755 third_party/wafer/lib/Conversion/TLEToMK/TLEToMKPass.cpp create mode 100755 third_party/wafer/lib/Conversion/TritonArithToLinalg/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp create mode 100755 third_party/wafer/lib/Conversion/TritonArithToLinalg/TritonArithToLinalgPass.cpp create mode 100755 third_party/wafer/lib/Conversion/TritonToCoreDialects/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/TritonToCoreDialects/TritonToCoreDialectsPass.cpp create mode 100755 third_party/wafer/lib/Conversion/UnstructuredToMK/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/UnstructuredToMK/UnstructuredToMKPass.cpp create mode 100755 third_party/wafer/lib/Conversion/WaferMemrefToLLVM/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVM.cpp create mode 100755 third_party/wafer/lib/Conversion/WaferMemrefToLLVM/WaferMemrefToLLVMPass.cpp create mode 100755 third_party/wafer/lib/Conversion/WaferToLLVM/CMakeLists.txt create mode 100755 third_party/wafer/lib/Conversion/WaferToLLVM/KernelArgBufferPass.cpp create mode 100755 third_party/wafer/lib/Conversion/WaferToLLVM/WaferToLLVM.cpp create mode 100755 third_party/wafer/lib/Conversion/WaferToLLVM/WaferToLLVMPass.cpp create mode 100755 third_party/wafer/lib/Dialect/Address/CMakeLists.txt create mode 100755 third_party/wafer/lib/Dialect/Address/IR/AddressDialect.cpp create mode 100755 third_party/wafer/lib/Dialect/Address/IR/CMakeLists.txt create mode 100755 third_party/wafer/lib/Dialect/Address/Transforms/AddrToLLVM.cpp create mode 100755 third_party/wafer/lib/Dialect/Address/Transforms/CMakeLists.txt create mode 100755 third_party/wafer/lib/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/lib/Dialect/MagicKernel/CMakeLists.txt create mode 100755 third_party/wafer/lib/Dialect/MagicKernel/IR/MagicKernelDialect.cpp create mode 100755 third_party/wafer/lib/Dialect/MagicKernel/Transforms/BufferizableOpInterfaceImpl.cpp create mode 100644 third_party/wafer/lib/Dialect/MagicKernel/Transforms/CMakeLists.txt create mode 100644 third_party/wafer/lib/Dialect/MagicKernel/Transforms/MaterializeStridedLinalgInputsPass.cpp create mode 100755 third_party/wafer/lib/Dialect/Wafer/CMakeLists.txt create mode 100755 third_party/wafer/lib/Dialect/Wafer/IR/WaferDialect.cpp create mode 100755 third_party/wafer/lib/Dialect/Wafer/IR/WaferOps.cpp create mode 100644 third_party/wafer/lib/Dialect/Wafer/Transforms/CMakeLists.txt create mode 100644 third_party/wafer/lib/Dialect/Wafer/Transforms/InsertBarrierPass.cpp create mode 100755 third_party/wafer/lib/Registrar/CMakeLists.txt create mode 100755 third_party/wafer/lib/Registrar/Registrar.cc create mode 100755 third_party/wafer/name.conf create mode 100644 third_party/wafer/patches/triton/cache_without_vendor_imports.patch create mode 100644 third_party/wafer/patches/triton/profiles.json create mode 100644 third_party/wafer/patches/triton/python_triton_compiler_optional_gluon_py.patch create mode 100644 third_party/wafer/patches/triton/python_triton_jit_optional_gluon.patch create mode 100644 third_party/wafer/patches/triton/wafer_builder_optional_gluon.patch create mode 100644 third_party/wafer/patches/triton/wafer_proton_backend_filter.patch create mode 100755 third_party/wafer/profiler/CMakeLists.txt create mode 100755 third_party/wafer/profiler/profiler.cpp create mode 100755 third_party/wafer/python/triton_wafer.cc create mode 100644 third_party/wafer/python/triton_wafer_frontend.cc create mode 100644 third_party/wafer/requirements-build.txt create mode 100755 third_party/wafer/scripts/base/base_run.sh create mode 100755 third_party/wafer/scripts/build_llvm.sh create mode 100755 third_party/wafer/scripts/build_wafer.sh create mode 100755 third_party/wafer/scripts/publish/run_flaggems_on_multicards.sh create mode 100755 third_party/wafer/scripts/publish/run_wafer.sh create mode 100755 third_party/wafer/scripts/requirements_ts.txt create mode 100755 third_party/wafer/scripts/run_wafer.sh create mode 100755 third_party/wafer/scripts/tools/suuplement.sh create mode 100755 third_party/wafer/third_party/flir/.clang-format create mode 100755 third_party/wafer/third_party/flir/.github/PULL_REQUEST_TEMPLATE.md create mode 100755 third_party/wafer/third_party/flir/.github/workflows/code-format-check.yml create mode 100755 third_party/wafer/third_party/flir/.gitignore create mode 100755 third_party/wafer/third_party/flir/.gitmodules create mode 100755 third_party/wafer/third_party/flir/.pre-commit-config.yaml create mode 100755 third_party/wafer/third_party/flir/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/LICENSE create mode 100755 third_party/wafer/third_party/flir/README.md create mode 100755 third_party/wafer/third_party/flir/backend/compiler.py create mode 100755 third_party/wafer/third_party/flir/backend/driver.py create mode 100755 third_party/wafer/third_party/flir/backend/include/ExecutionEngine/CRunnerUtils.cpp create mode 100755 third_party/wafer/third_party/flir/backend/include/ExecutionEngine/CRunnerUtils.h create mode 100755 third_party/wafer/third_party/flir/backend/include/ExecutionEngine/Msan.h create mode 100755 third_party/wafer/third_party/flir/backend/include/ExecutionEngine/version.txt create mode 100755 third_party/wafer/third_party/flir/backend/name.conf create mode 100755 third_party/wafer/third_party/flir/include/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/DiscreteMaskAccessConversion/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToAnnotation/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/ConversionPatterns.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/DescriptorConverter.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/FunctionConverter.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/HoistBroadcast.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/LoadStoreConverter.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/MaskAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/TritonOpConverter.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToLinalgIncubated/UseAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/MaskAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/MemOpConverter.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/PtrAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/InterleaveOptimization.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Conversion/UtilsIncubated/Utils.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.h create mode 100755 third_party/wafer/third_party/flir/include/incubated/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.td create mode 100755 third_party/wafer/third_party/flir/include/mlir-ext/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/mlir-ext/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExt.h create mode 100755 third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExtBase.td create mode 100755 third_party/wafer/third_party/flir/include/mlir-ext/Dialect/MathExt/IR/MathExtOps.td create mode 100755 third_party/wafer/third_party/flir/include/npu/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/npu/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendAttrDefs.td create mode 100755 third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendDialect.h create mode 100755 third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendDialect.td create mode 100755 third_party/wafer/third_party/flir/include/npu/Dialect/TritonAscend/IR/TritonAscendOps.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Analysis/MaskAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Analysis/OpFoldResultUtils.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Analysis/PtrAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Analysis/UseAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/AnalysisStructured/PtrAnalysis.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/MemrefCopyToDMA_FlagTree/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/NoBufferize_FlagTree/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/ReconcilePtrCasts/ReconcilePtrCasts.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/StructuredToMemref/StructuredToMemref.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.hpp create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns_FlagTree.hpp create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonArithToLinalg/TritonArithToLinalg.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonPtrToMemref/TritonPtrToMemref.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalg/TritonToLinalg.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimental.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToLinalgExperimental/TritonToPtr.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToStructured/TritonToStructured.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/TritonToUnstructured/TritonToUnstructured.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/Passes.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/Passes.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Conversion/UnstructuredToMemref/UnstructuredToMemref.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/TPtrDialect.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TPtr/IR/TPtrDialect.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtInterfaces.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Dialect/TritonTilingExt/IR/TritonTilingExtOps.td create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Utils/FusionHelper.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Utils/ReduceScanCommon.h create mode 100755 third_party/wafer/third_party/flir/include/triton-shared/Utils/Utils.h create mode 100755 third_party/wafer/third_party/flir/lib/Analysis/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Analysis/MaskAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Analysis/OpFoldResultUtils.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Analysis/PtrAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Analysis/UseAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/AnalysisStructured/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/AnalysisStructured/PtrAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/AnalysisStructured/PtrAnalysisTS.cpp create mode 100755 third_party/wafer/third_party/flir/lib/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/DiscreteMaskAccessConversion/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTree.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/MemrefCopyToDMA_FlagTree/MemrefCopyToDMAFlagTreePass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTree.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/NoBufferize_FlagTree/NoBufferizeFlagTreePass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/ReconcilePtrCasts/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/ReconcilePtrCasts/ReconcilePtrCastsPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/StructuredToMemref.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/StructuredToMemref/StructuredToMemrefPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonArithToLinalg/TritonArithToLinalgPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonPtrToMemref/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonPtrToMemref/TritonPtrToMemrefPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToAnnotation/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToAnnotation/TritonToAnnotation.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/TritonToLinalg.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalg/TritonToLinalgPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimentalPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgExperimental/TritonToPtrPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/ArgMinMaxConverter.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/BlockPtrAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/DescriptorConverter.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/FunctionConverter.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/HoistBroadcast.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/LoadStoreConverter.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/MaskAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/TritonOpConverter.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/TritonToLinalgIncubatedPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToLinalgIncubated/UseAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToStructured/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToStructured/TritonToStructuredPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/CannonicalizerConverter.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/MaskAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/MemOpConverter.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/PtrAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToStructuredIncubated/TritonToStructuredIncubatedPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/BubbleUpOperation.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/OffsetAnalysis.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructureIncubated/UnstructureConversionPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructured/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/TritonToUnstructured/TritonToUnstructuredPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/UnstructuredToMemref/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Conversion/UnstructuredToMemref/UnstructuredToMemrefPass.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/MathExt/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/MathExtDialect.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/MathExt/IR/MathExtOps.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TPtr/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/TPtrDialect.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TPtr/IR/TPtrOps.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendAttrs.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendDialect.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonAscend/IR/TritonAscendOps.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/TritonStructuredDialect.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonStructured/IR/TritonStructuredOps.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/TritonStructuredDialectIncubated.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonStructuredIncubated/IR/TritonStructuredOpsIncubated.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/BufferizableOpInterfaceImpl.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/CumSum.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Dialect/TritonTilingExt/IR/TritonTilingExtDialect.cpp create mode 100755 third_party/wafer/third_party/flir/lib/Utils/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/Utils/Utils.cpp create mode 100755 third_party/wafer/third_party/flir/lib/UtilsIncubated/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/lib/UtilsIncubated/InterleaveOptimization.cpp create mode 100755 third_party/wafer/third_party/flir/lib/UtilsIncubated/Utils.cpp create mode 100755 third_party/wafer/third_party/flir/python/examples/bare_matmul.py create mode 100755 third_party/wafer/third_party/flir/python/examples/benchmark.py create mode 100755 third_party/wafer/third_party/flir/python/examples/conftest.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_addptr.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_blockptr_complex_offset.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_early_return.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_gather_scatter.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_layernorm.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_load_2d_tensor_block.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_load_2d_tensor_col.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_mask.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_matmul.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_modulo.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_nested_loops.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_reduce.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_scalar_store.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_sign_extend.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_softmax.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_splat.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_swap.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_tensor_index_iterargs.py create mode 100755 third_party/wafer/third_party/flir/python/examples/test_vec_add.py create mode 100755 third_party/wafer/third_party/flir/test/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/test/README.md create mode 100755 third_party/wafer/third_party/flir/test/lit.cfg.py create mode 100755 third_party/wafer/third_party/flir/test/lit.site.cfg.py.in create mode 100755 third_party/wafer/third_party/flir/tools/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/tools/RegisterTritonSharedDialects.h create mode 100755 third_party/wafer/third_party/flir/tools/triton-shared-opt/CMakeLists.txt create mode 100755 third_party/wafer/third_party/flir/tools/triton-shared-opt/triton-shared-opt.cpp create mode 100755 third_party/wafer/third_party/flir/triton_shared.cc create mode 100755 third_party/wafer/third_party/tle/CMakeLists.txt create mode 100755 third_party/wafer/third_party/tle/REANME.md create mode 100755 third_party/wafer/third_party/tle/include/CMakeLists.txt create mode 100755 third_party/wafer/third_party/tle/include/tle-dsa/Conversion/DsaToCore/DsaToCore.h create mode 100755 third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.h create mode 100755 third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaDialect.td create mode 100755 third_party/wafer/third_party/tle/include/tle-dsa/Dialect/IR/DsaOps.td create mode 100755 third_party/wafer/third_party/tle/lib/CMakeLists.txt create mode 100755 third_party/wafer/third_party/tle/lib/Conversion/DsaToCore/CMakeLists.txt create mode 100755 third_party/wafer/third_party/tle/lib/Conversion/DsaToCore/DsaToCore.cpp create mode 100755 third_party/wafer/third_party/tle/lib/Dialect/IR/CMakeLists.txt create mode 100755 third_party/wafer/third_party/tle/lib/Dialect/IR/DsaDialect.cpp create mode 100755 third_party/wafer/third_party/tle/python/CMakeLists.txt create mode 100755 third_party/wafer/third_party/tle/python/ir.h create mode 100755 third_party/wafer/third_party/tle/python/triton_tle_dsa.cc 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); +}