Download code/tt_diffusion_planner/tt/kernels/fattn_compute.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 9.92 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/fattn_compute.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/fattn_compute.cpp
-
curl -L -o fattn_compute.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/fattn_compute.cpp
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]. | |
| 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); | |
| } | |