Download code/kernels/sp_desc/sm_compute.cpp from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 5.2 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_desc/sm_compute.cpp
- Command line
-
hf download hf://changh95/superpoint-p150/code/kernels/sp_desc/sm_compute.cpp
-
curl -L -o sm_compute.cpp https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_desc/sm_compute.cpp
5.2 kB
| // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. | |
| // SPDX-License-Identifier: Apache-2.0 | |
| // | |
| // Score softmax on the otherwise idle cores of the merged head op (models/tt/desc_head.py, SP_HEAD_SM=1). | |
| // The per-row math is ttnn's sharded attention softmax kernel (softmax_sharded.cpp, NUMERIC_STABLE, no mask) | |
| // verbatim, with the compute config ttnn.softmax uses (HiFi4, fp32 DST off, approx mode on -> bf16 | |
| // intermediates): max reduce, x - max, exp, sum reduce + recip, x * (1 / sum). The score logits arrive with | |
| // -inf in the padding columns (65..95), so this equals ttnn.softmax's padded-column-masked path bit for bit | |
| // (verified: scratch/r8_sm_test.py, test_superpoint_merged_head_softmax). | |
| template < | |
| std::uint32_t block_w, | |
| std::uint32_t num_subblocks_w, | |
| std::uint32_t subblock_w, | |
| std::uint32_t dfb_in_id, | |
| std::uint32_t dfb_max_scaler_id, | |
| std::uint32_t dfb_max_id, | |
| std::uint32_t dfb_out_id> | |
| ALWI void calc_numeric_stable() { | |
| DataflowBuffer dfb_in_obj(dfb_in_id); | |
| DataflowBuffer dfb_max_obj(dfb_max_id); | |
| DataflowBuffer dfb_out_obj(dfb_out_id); | |
| compute_kernel_lib::reduce< | |
| PoolType::MAX, | |
| ReduceDim::REDUCE_ROW, | |
| dfb_in_id, | |
| dfb_max_scaler_id, | |
| dfb_max_id, | |
| compute_kernel_lib::ReduceInputPolicy::NoWaitNoPop>(compute_kernel_lib::ReduceInputBlockShape::row(block_w)); | |
| exp_tile_init<EXP_APPROX>(); | |
| reconfig_data_format_srcb(dfb_max_id); | |
| dfb_max_obj.wait_front(1); | |
| sub_bcast_cols_init(dfb_in_id, dfb_max_id); | |
| std::uint32_t index_subblock_w_offset = 0; | |
| for (std::uint32_t j = 0; j < num_subblocks_w; j++) { | |
| tile_regs_acquire(); | |
| dfb_out_obj.reserve_back(subblock_w); | |
| for (std::uint32_t w = 0; w < subblock_w; w++) { | |
| std::uint32_t index = w + index_subblock_w_offset; | |
| sub_tiles_bcast_cols(dfb_in_id, dfb_max_id, index, 0, w); | |
| } | |
| dfb_out_obj.reserve_back(subblock_w); | |
| for (std::uint32_t w = 0; w < subblock_w; w++) { | |
| exp_tile<EXP_APPROX>(w); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| for (std::uint32_t w = 0; w < subblock_w; w++) { | |
| pack_tile(w, dfb_out_id); | |
| } | |
| tile_regs_release(); | |
| dfb_out_obj.push_back(subblock_w); | |
| index_subblock_w_offset += subblock_w; | |
| } | |
| dfb_in_obj.pop_front(block_w); | |
| dfb_max_obj.pop_front(1); | |
| dfb_out_obj.wait_front(block_w); | |
| } | |
| void kernel_main() { | |
| const std::uint32_t nrows = get_arg_val<uint32_t>(0); | |
| constexpr std::uint32_t dfb_in0 = get_compile_time_arg_val(0); | |
| constexpr std::uint32_t dfb_max_scaler = get_compile_time_arg_val(1); | |
| constexpr std::uint32_t dfb_sum_scaler = get_compile_time_arg_val(2); | |
| constexpr std::uint32_t dfb_exps = get_compile_time_arg_val(3); | |
| constexpr std::uint32_t dfb_recipsumexps = get_compile_time_arg_val(4); | |
| constexpr std::uint32_t dfb_out0 = get_compile_time_arg_val(5); | |
| constexpr std::uint32_t dfb_max = get_compile_time_arg_val(6); | |
| constexpr std::uint32_t block_w = 3, subblock_w = 3, num_subblocks_w = 1; | |
| compute_kernel_hw_startup(dfb_in0, dfb_max_scaler, dfb_exps); | |
| DataflowBuffer dfb_in0_obj(dfb_in0); | |
| DataflowBuffer dfb_exps_obj(dfb_exps); | |
| DataflowBuffer dfb_recipsumexps_obj(dfb_recipsumexps); | |
| DataflowBuffer dfb_out0_obj(dfb_out0); | |
| for (std::uint32_t i = 0; i < nrows; i++) { | |
| dfb_in0_obj.wait_front(block_w); | |
| calc_numeric_stable<block_w, num_subblocks_w, subblock_w, dfb_in0, dfb_max_scaler, dfb_max, dfb_exps>(); | |
| dfb_exps_obj.wait_front(block_w); | |
| compute_kernel_lib::reduce< | |
| PoolType::SUM, | |
| ReduceDim::REDUCE_ROW, | |
| dfb_exps, | |
| dfb_sum_scaler, | |
| dfb_recipsumexps, | |
| compute_kernel_lib::ReduceInputPolicy::NoWaitNoPop>( | |
| compute_kernel_lib::ReduceInputBlockShape::row(block_w), | |
| compute_kernel_lib::ReduceInputMemoryLayout::contiguous(), | |
| compute_kernel_lib::NoAccumulation{}, | |
| [](std::uint32_t) { | |
| recip_tile_init(); | |
| recip_tile(0); | |
| }); | |
| reconfig_data_format(dfb_exps, dfb_recipsumexps); | |
| pack_reconfig_data_format(dfb_out0); | |
| dfb_recipsumexps_obj.wait_front(1); | |
| mul_bcast_cols_init(dfb_exps, dfb_recipsumexps); | |
| tile_regs_acquire(); | |
| dfb_out0_obj.reserve_back(subblock_w); | |
| for (std::uint32_t w = 0; w < subblock_w; w++) { | |
| mul_tiles_bcast<BroadcastType::COL>(dfb_exps, dfb_recipsumexps, w, 0, w); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| for (std::uint32_t w = 0; w < subblock_w; w++) { | |
| pack_tile(w, dfb_out0); | |
| } | |
| tile_regs_release(); | |
| dfb_out0_obj.push_back(subblock_w); | |
| dfb_recipsumexps_obj.pop_front(1); | |
| dfb_exps_obj.pop_front(block_w); | |
| } | |
| } | |