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