// SPDX-License-Identifier: Apache-2.0 // Attention score scale + mask (tt/smask_kernel.py), compute: out = s * scale + mask, the two stock binary_ng fp32 // programs `ttnn.multiply(s, scale)` (SFPU mul_binary_tile against a tile filled with the fp32 scale) and // `ttnn.add(., mask)`, in one pass. With an fp32 mask the stock add is the SFPU add_binary_tile; with a // bf16 mask (trunc_tf32 = 1) binary_ng takes its FPU path, whose SrcA truncates the fp32 scores to TF32 (measured: // logs/diffusion-planner/opt_r2/smask_diag.log, the stock result == (exact & 0xFFFFE000) for every finite element), // so the scaled scores are truncated the same way (SFPU bitwise AND on the raw bits) before the exact add of the // 0 / -inf mask. Bit-identical to the two programs in both cases. // CT args: [0] H (even), [1] per-core RT-arg count (P3), [2] trunc_tf32. Per-core RT args: [n, 0]. #include #include "api/compute/common.h" #include "api/compute/compute_kernel_api.h" #include "api/compute/eltwise_binary_sfpu.h" #include "api/compute/eltwise_unary/bitwise.h" #include "api/compute/eltwise_unary/eltwise_unary.h" #include "api/compute/tile_move_copy.h" using namespace ckernel; constexpr uint32_t H = get_compile_time_arg_val(0); constexpr uint32_t trunc_tf32 = get_compile_time_arg_val(2); constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2, cb_out = 16; constexpr auto RNE = ckernel::DstRoundingMode::NearestEven; void kernel_main() { const uint32_t n = get_arg_val(0); if (n == 0) { return; } compute_kernel_hw_startup(cb_s, cb_out); cb_wait_front(cb_c, 1); for (uint32_t t = 0; t < n; ++t) { cb_wait_front(cb_m, 1); cb_wait_front(cb_s, H); for (uint32_t h = 0; h < H; h += 2) { tile_regs_acquire(); copy_init(cb_s); copy_tile(cb_s, h, 0); copy_tile(cb_s, h + 1, 1); copy_init(cb_c); copy_tile(cb_c, 0, 2); mul_binary_tile_init(); mul_binary_tile(0, 2, 0); mul_binary_tile(1, 2, 1); if constexpr (trunc_tf32) { bitwise_and_tile_init(); bitwise_and_tile(0, 0xFFFFE000u); bitwise_and_tile(1, 0xFFFFE000u); } reconfig_data_format_srca(cb_c, cb_m); copy_init(cb_m); copy_tile(cb_m, 0, 3); reconfig_data_format_srca(cb_m, cb_s); add_binary_tile_init(); add_binary_tile(0, 3, 0); add_binary_tile(1, 3, 1); cb_reserve_back(cb_out, 2); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_out); pack_tile(1, cb_out); tile_regs_release(); cb_push_back(cb_out, 2); } cb_pop_front(cb_s, H); cb_pop_front(cb_m, 1); } }