changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
2.92 kB
// 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<NearestEven>; 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 <cstdint>
#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<uint32_t>(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<DataFormat::Int32>(0, 0xFFFFE000u);
bitwise_and_tile<DataFormat::Int32>(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<RNE>(0, 3, 0);
add_binary_tile<RNE>(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);
}
}