changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
9.92 kB
// SPDX-License-Identifier: Apache-2.0
// Fused fp32 matmul attention (tt/fattn_kernel.py, ATTN_FUSED), compute. Per unit (head h, query tile row r; Wt key
// tiles), the five stock programs of tt/attention.py (ATTN_FAST=1 + ATTN_SMSM) in one pass, the scores and the
// probabilities kept in L1:
// phase Q: s_j = Q_r K_j^T in DEST (matmul_block, in1 transposed by the unpacker: the stock Q K^T matmul with its
// K = 32 in one block), two key tiles at a time;
// phase A: x_j = trunc_tf32(s_j * scale) + mask_j on the SFPU (kernels/smsm_compute.cpp, scale mode 1), packed to
// cb_x;
// phase B: the stock numeric-stable softmax of the row (kernels/smsm_compute.cpp phase B) -> cb_p;
// phase C: o = sum_j P_j V_j accumulated in DEST in key order (the stock P V matmul: the whole K in one block),
// packed to cb_o; the writer stores it as tile (r, h) of the merged-heads output.
// (kcat_out, KCAT_EMIT) the output tile as the operand of the next K-concatenated split linear: o_hi = bf16(o)
// and o_lo = o - o_hi (kernels/kcat_compute.cpp's LLK calls on the packed fp32 tile) -> cb_o / cb_lo.
// CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] trunc_tf32, [3] ndst, [4] P pad tiles per row (to a multiple
// of ndst), [5] scale bits (fp32), [6] kcat_out. Per-core RT args: [nunits, 0].
#include <cstdint>
#include "api/compute/common.h"
#include "api/compute/compute_kernel_api.h"
#include "api/compute/eltwise_binary.h"
#include "api/compute/eltwise_binary_sfpu.h"
#include "api/compute/eltwise_unary/binop_with_scalar.h"
#include "api/compute/eltwise_unary/bitwise.h"
#include "api/compute/eltwise_unary/eltwise_unary.h"
#include "api/compute/eltwise_unary/typecast.h"
#include "api/compute/tile_move_copy.h"
#include "api/compute/matmul.h"
#include "api/compute/bcast.h"
#include "api/compute/softmax.h"
#include "api/compute/reduce.h"
#include "ttnn/cpp/ttnn/kernel_lib/reduce_helpers_compute.hpp"
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/api/chain.hpp"
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/api/convenience.hpp"
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/unary/math.hpp"
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/core/optional.hpp"
namespace ckl = compute_kernel_lib;
using namespace ckernel;
constexpr uint32_t Wt = get_compile_time_arg_val(0);
constexpr uint32_t trunc_tf32 = get_compile_time_arg_val(2);
constexpr uint32_t ndst = get_compile_time_arg_val(3);
constexpr uint32_t p_pad = get_compile_time_arg_val(4);
constexpr uint32_t scale_bits = get_compile_time_arg_val(5);
constexpr uint32_t kcat_out = get_compile_time_arg_val(6);
constexpr uint32_t cb_q = 0, cb_m = 1, cb_k = 2, cb_max_scaler = 3, cb_sum_scaler = 4, cb_v = 5;
constexpr uint32_t cb_o = 16, cb_lo = 17, cb_y = 29;
constexpr uint32_t cb_x = 24, cb_max = 25, cb_exps = 26, cb_recip = 27, cb_p = 28;
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;
constexpr uint32_t Wp = (Wt + 1) / 2;
constexpr uint32_t Wr = Wt + p_pad;
template <std::uint32_t dfb_in, std::uint32_t dfb_max_scaler, std::uint32_t dfb_max, std::uint32_t dfb_out>
void calc_numeric_stable(std::uint32_t W, std::uint32_t nd) {
compute_kernel_lib::reduce<
PoolType::MAX,
ReduceDim::REDUCE_ROW,
dfb_in,
dfb_max_scaler,
dfb_max,
compute_kernel_lib::ReduceInputPolicy::WaitUpfrontNoPop,
compute_kernel_lib::ReduceDataFormatReconfigMode::INPUT>(compute_kernel_lib::ReduceInputBlockShape::row(W));
ckl::eltwise_chain(
ckl::IterationShape::tiles(W).block_size(nd),
ckl::BinaryFpu<
ckl::BinaryFpuOp::Sub,
ckl::input(
dfb_in,
ckl::WaitPolicy::Upfront,
ckl::PopPolicy::AtEnd,
ckl::InputTileMapping::Block,
ckl::DataFormatReconfig::Disabled),
ckl::input(dfb_max, ckl::BroadcastDim::Col, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd)>{},
ckl::Exp<ckl::Approx::Exact, ckl::Dst::D0>{},
ckl::PackTile<ckl::output(
dfb_out,
ckl::ReservePolicy::PerBlockSize,
ckl::PushPolicy::PerBlockSize,
ckl::DataFormatReconfig::Disabled)>{});
cb_wait_front(dfb_out, W);
}
void kernel_main() {
const uint32_t nunits = get_arg_val<uint32_t>(0);
if (nunits == 0) {
return;
}
compute_kernel_hw_startup(cb_q, cb_k, cb_x);
cb_wait_front(cb_max_scaler, 1);
cb_wait_front(cb_sum_scaler, 1);
for (uint32_t u = 0; u < nunits; ++u) {
// ---- phases Q + A: x_j = trunc(Q K_j^T * scale) + mask_j, two key tiles at a time ----
cb_wait_front(cb_q, 1);
cb_wait_front(cb_m, 2 * Wp);
pack_reconfig_data_format(cb_x);
for (uint32_t p = 0; p < Wp; ++p) {
const uint32_t j = 2 * p;
const bool two = (j + 1) < Wt;
cb_wait_front(cb_k, 2);
reconfig_data_format(cb_k, cb_q);
matmul_block_init(cb_q, cb_k, 1, 1, 1, 1);
tile_regs_acquire();
matmul_block(cb_q, cb_k, 0, 0, 0, 1, 1, 1, 1);
matmul_block(cb_q, cb_k, 0, 1, 1, 1, 1, 1, 1);
binop_with_scalar_tile_init();
mul_unary_tile(0, scale_bits);
mul_unary_tile(1, scale_bits);
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_k, cb_m);
copy_init(cb_m);
copy_tile(cb_m, j, 2);
copy_tile(cb_m, j + 1, 3);
add_binary_tile_init();
add_binary_tile<RNE>(0, 2, 0);
add_binary_tile<RNE>(1, 3, 1);
const uint32_t np = two ? 2 : 1;
cb_reserve_back(cb_x, np);
tile_regs_commit();
tile_regs_wait();
pack_tile(0, cb_x);
if (two) {
pack_tile(1, cb_x);
}
tile_regs_release();
cb_push_back(cb_x, np);
cb_pop_front(cb_k, 2);
}
cb_pop_front(cb_m, 2 * Wp);
cb_pop_front(cb_q, 1);
// ---- phase B: the stock numeric-stable softmax of the row -> cb_p ----
reconfig_data_format(cb_x, cb_x);
pack_reconfig_data_format(cb_exps);
copy_init(cb_x);
calc_numeric_stable<cb_x, cb_max_scaler, cb_max, cb_exps>(Wt, ndst);
reconfig_data_format(cb_exps, cb_sum_scaler);
compute_kernel_lib::reduce<
PoolType::SUM,
ReduceDim::REDUCE_ROW,
cb_exps,
cb_sum_scaler,
cb_recip,
compute_kernel_lib::ReduceInputPolicy::WaitUpfrontNoPop>(
compute_kernel_lib::ReduceInputBlockShape::row(Wt),
compute_kernel_lib::ReduceInputMemoryLayout::contiguous(),
compute_kernel_lib::NoAccumulation{},
[](std::uint32_t) {
if constexpr (DST_ACCUM_MODE) {
recip_tile_init<ReciprocalDestAcc::FP32, ReciprocalApproxMode::Precise>();
recip_tile<ReciprocalDestAcc::FP32, ReciprocalApproxMode::Precise>(0);
} else {
recip_tile_init();
recip_tile(0);
}
});
ckl::mul<
ckl::input(cb_exps, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd, ckl::InputTileMapping::Block),
ckl::input(cb_recip, ckl::BroadcastDim::Col, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd),
ckl::output(cb_p, ckl::ReservePolicy::PerBlockSize, ckl::PushPolicy::PerBlockSize)>(
ckl::IterationShape::tiles(Wt).block_size(ndst));
if constexpr (p_pad > 0) {
cb_reserve_back(cb_p, p_pad);
cb_push_back(cb_p, p_pad);
}
// ---- phase C: o = sum_j P_j V_j (key order, one DEST accumulation) ----
cb_wait_front(cb_p, Wr);
reconfig_data_format(cb_v, cb_p);
pack_reconfig_data_format(cb_o);
matmul_block_init(cb_p, cb_v, 0, 1, 1, 1);
tile_regs_acquire();
for (uint32_t j = 0; j < Wt; ++j) {
cb_wait_front(cb_v, 1);
matmul_block(cb_p, cb_v, j, 0, 0, 0, 1, 1, 1);
cb_pop_front(cb_v, 1);
}
if constexpr (kcat_out) {
cb_reserve_back(cb_y, 1);
tile_regs_commit();
tile_regs_wait();
pack_tile(0, cb_y);
tile_regs_release();
cb_push_back(cb_y, 1);
cb_pop_front(cb_p, Wr);
// ---- the split operand of the next linear (kernels/kcat_compute.cpp) ----
cb_wait_front(cb_y, 1);
reconfig_data_format_srca(cb_y);
tile_regs_acquire();
copy_init(cb_y);
copy_tile(cb_y, 0, 0);
copy_tile(cb_y, 0, 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_o, 1);
cb_reserve_back(cb_lo, 1);
tile_regs_commit();
tile_regs_wait();
pack_tile(0, cb_o);
pack_tile(1, cb_lo);
tile_regs_release();
cb_push_back(cb_o, 1);
cb_push_back(cb_lo, 1);
cb_pop_front(cb_y, 1);
} else {
cb_reserve_back(cb_o, 1);
tile_regs_commit();
tile_regs_wait();
pack_tile(0, cb_o);
tile_regs_release();
cb_push_back(cb_o, 1);
cb_pop_front(cb_p, Wr);
}
}
cb_pop_front(cb_max_scaler, 1);
cb_pop_front(cb_sum_scaler, 1);
}