changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
8.18 kB
// SPDX-License-Identifier: Apache-2.0
// Attention score scale + mask + softmax (tt/smsm_kernel.py, ATTN_SMSM), compute. Per tile row (one head, 32 query
// rows, Wt key tiles):
// phase A: x = trunc_tf32(s * scale) + mask, the smask kernel's SFPU sequence (kernels/smask_compute.cpp: the two
// stock binary_ng programs bit for bit), packed to cb_x (L1) instead of DRAM;
// phase B: the stock ttnn.softmax(numeric_stable=True) of that row, i.e. the no-mask path of
// ttnn/.../softmax/device/kernels/attention/compute/softmax.cpp with the same kernel_lib calls: row max
// (FPU reduce), exp(x - max) (FPU bcast sub + SFPU exp), row sum (FPU reduce) + precise fp32
// reciprocal, x * 1/sum (FPU bcast mul).
// The stock softmax unpacks its fp32 input to SrcA (TF32); x holds TF32 values already (phase A truncates them and
// the mask is 0 / -inf), so the stock program saw exactly these operands.
// CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] trunc_tf32, [3] ndst (block), [4] out pad tiles per row,
// [5] scale mode: 0 = multiply by a tile filled with the scale (mul_binary_tile, as the stock binary_ng
// program), 1 = mul_unary_tile with the scale bits (the same SFPU fp32 multiply, no scale tile copy).
// Per-core RT args: [nrows, 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/tile_move_copy.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 out_pad = get_compile_time_arg_val(4);
constexpr uint32_t scale_mode = get_compile_time_arg_val(5);
constexpr uint32_t scale_bits = get_compile_time_arg_val(6);
constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2, cb_max_scaler = 3, cb_sum_scaler = 4;
constexpr uint32_t cb_out = 16;
constexpr uint32_t cb_x = 24, cb_max = 25, cb_exps = 26, cb_recip = 27;
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;
constexpr uint32_t Wp = (Wt + 1) / 2; // tile pairs per row (the last one half-used when Wt is odd)
// the stock calc_numeric_stable (softmax.cpp), CB ids instead of DFB handles
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 nrows = get_arg_val<uint32_t>(0);
if (nrows == 0) {
return;
}
compute_kernel_hw_startup(cb_s, cb_c, cb_x);
cb_wait_front(cb_c, 1);
cb_wait_front(cb_max_scaler, 1);
cb_wait_front(cb_sum_scaler, 1);
for (uint32_t r = 0; r < nrows; ++r) {
// ---- phase A: x = trunc(s * scale) + mask (pairs of tiles) ----
reconfig_data_format(cb_s, cb_c);
pack_reconfig_data_format(cb_x);
cb_wait_front(cb_m, 2 * Wp);
for (uint32_t p = 0; p < Wp; ++p) {
const uint32_t j = 2 * p;
const bool two = (j + 1) < Wt;
cb_wait_front(cb_s, 2);
tile_regs_acquire();
copy_init(cb_s);
copy_tile(cb_s, 0, 0);
copy_tile(cb_s, 1, 1);
if constexpr (scale_mode == 0) {
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);
} else {
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(scale_mode == 0 ? cb_c : cb_s, cb_m);
copy_init(cb_m);
copy_tile(cb_m, j, 2);
copy_tile(cb_m, j + 1, 3);
reconfig_data_format_srca(cb_m, cb_s);
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_s, 2);
}
cb_pop_front(cb_m, 2 * Wp);
// ---- phase B: the stock numeric-stable softmax of the row ----
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_out, ckl::ReservePolicy::PerBlockSize, ckl::PushPolicy::PerBlockSize)>(
ckl::IterationShape::tiles(Wt).block_size(ndst));
if constexpr (out_pad > 0) {
cb_reserve_back(cb_out, out_pad);
cb_push_back(cb_out, out_pad);
}
}
cb_pop_front(cb_c, 1);
cb_pop_front(cb_max_scaler, 1);
cb_pop_front(cb_sum_scaler, 1);
}