changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
2.89 kB
// SPDX-License-Identifier: Apache-2.0
// Split-matmul operand build (tt/kcat_kernel.py), compute: per fp32 x tile, x_hi = bf16(x) (typecast_tile
// Float32 -> Float16_b, the stock ttnn.typecast LLK) kept as fp32, and x_lo = x - x_hi (sub_binary_tile, exact).
// With act (KCAT_ACT) the previous linear's activation is applied first, as the stock unary program does it
// (math_approx_mode false): 1 = gelu_tile<0> (ttnn.gelu, accurate), 2 = gelu_tanh_tile (GeluVariant.Tanh).
// CT args: [0] per-core RT-arg count (P3), [1] act, [2] once (KCAT_ACT_ONCE: the activation on one DEST copy, then
// copy_dest_values to the second; the same values). Per-core RT args: [n, 0].
#include <cstdint>
#include "api/compute/common.h"
#include "api/compute/compute_kernel_api.h"
#include "api/compute/copy_dest_values.h"
#include "api/compute/eltwise_binary_sfpu.h"
#include "api/compute/eltwise_unary/eltwise_unary.h"
#include "api/compute/eltwise_unary/gelu.h"
#include "api/compute/eltwise_unary/typecast.h"
#include "api/compute/tile_move_copy.h"
using namespace ckernel;
constexpr uint32_t act = get_compile_time_arg_val(1);
constexpr uint32_t once = get_compile_time_arg_val(2);
constexpr uint32_t cb_x = 0, cb_hi = 16, cb_lo = 17;
void kernel_main() {
const uint32_t n = get_arg_val<uint32_t>(0);
if (n == 0) {
return;
}
compute_kernel_hw_startup(cb_x, cb_hi);
for (uint32_t t = 0; t < n; ++t) {
cb_wait_front(cb_x, 1);
tile_regs_acquire();
copy_init(cb_x);
copy_tile(cb_x, 0, 0);
if constexpr (once && act != 0) {
if constexpr (act == 1) {
gelu_tile_init<0u>();
gelu_tile<0u>(0);
} else {
gelu_tanh_tile_init();
gelu_tanh_tile(0);
}
copy_dest_values_init();
copy_dest_values<DataFormat::Float32>(0, 1);
} else {
copy_tile(cb_x, 0, 1);
}
if constexpr (once && act != 0) {
} else if constexpr (act == 1) {
gelu_tile_init<0u>();
gelu_tile<0u>(0);
gelu_tile<0u>(1);
} else if constexpr (act == 2) {
gelu_tanh_tile_init();
gelu_tanh_tile(0);
gelu_tanh_tile(1);
}
typecast_tile_init<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>();
typecast_tile<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(0);
sub_binary_tile_init();
sub_binary_tile<ckernel::DstRoundingMode::NearestEven>(1, 0, 1);
cb_reserve_back(cb_hi, 1);
cb_reserve_back(cb_lo, 1);
tile_regs_commit();
tile_regs_wait();
pack_tile(0, cb_hi);
pack_tile(1, cb_lo);
tile_regs_release();
cb_push_back(cb_hi, 1);
cb_push_back(cb_lo, 1);
cb_pop_front(cb_x, 1);
}
}