Download code/kernels/sp_desc/dh_compute.cpp from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.64 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_desc/dh_compute.cpp
- Command line
-
hf download hf://changh95/superpoint-p150/code/kernels/sp_desc/dh_compute.cpp
-
curl -L -o dh_compute.cpp https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_desc/dh_compute.cpp
8.64 kB
| // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. | |
| // SPDX-License-Identifier: Apache-2.0 | |
| // | |
| // Descriptor head in one op (models/tt/desc_head.py): 1x1 conv (256 -> 256) + bias, then the L2 norm and | |
| // untilize of DescNormRM (dn_compute.cpp), per tile row of the height-sharded [rows, 256] TILE input. | |
| // 1) z_f32 = x @ W, fp32 DST, K in ttnn order (k = 0..7), at HiFi2 (the fidelity of the ttnn 1x1 conv; | |
| // the kernel itself is compiled at HiFi4 for the norm steps) -> CB_P (fp32) | |
| // 2) z = bf16(z_f32 + bias) (row broadcast; ttnn's fused-bias epilogue) -> CB_Z (bf16, the ttnn conv output) | |
| // 3) S = sum_c z^2 (fp32), r = rsqrt(S), y = z * r, pack-untilized -> CB_OUT (exactly dn_compute.cpp) | |
| // CB_W tile t = nb * 32 + k * 4 + j holds W tile (k, 4 nb + j); tiles 64 + n hold the bias (row 0) of | |
| // output tile n. | |
| // matmul_block_init / matmul_block with an explicit math fidelity (no dynamic throttle) | |
| ALWI void mmf_init(uint32_t in0, uint32_t in1, uint32_t ct) { | |
| state_configure(in1, in0, __builtin_LINE()); | |
| UNPACK((llk_unpack_AB_matmul_init(in0, in1, 0, ct, 1, 1))); | |
| MATH((llk_math_matmul_init<MM_FID, MM_THROTTLE>(in0, in1, 0, ct, 1))); | |
| } | |
| ALWI void mmf_block(uint32_t in0, uint32_t in1, uint32_t i0, uint32_t i1, uint32_t d, uint32_t ct) { | |
| state_configure(in1, in0, __builtin_LINE()); | |
| UNPACK((llk_unpack_AB_matmul(in0, in1, i0, i1, ct, 1, 1))); | |
| MATH((llk_math_matmul<MM_FID, MM_THROTTLE>(d, ct, 1))); | |
| } | |
| void kernel_main() { | |
| constexpr uint32_t cb_in = get_compile_time_arg_val(0); | |
| constexpr uint32_t cb_one = get_compile_time_arg_val(1); | |
| constexpr uint32_t cb_sq = get_compile_time_arg_val(2); | |
| constexpr uint32_t cb_rs = get_compile_time_arg_val(3); | |
| constexpr uint32_t cb_out = get_compile_time_arg_val(4); | |
| constexpr uint32_t TR = get_compile_time_arg_val(5); // tile rows per core | |
| constexpr uint32_t cb_w = get_compile_time_arg_val(6); | |
| constexpr uint32_t cb_p = get_compile_time_arg_val(7); | |
| constexpr uint32_t cb_z = get_compile_time_arg_val(8); | |
| constexpr uint32_t WT = 8, HB = 4, KT = 8, NW = DH_NW, BIAS0 = 64, KIN = DH_KIN; | |
| compute_kernel_hw_startup<SrcOrder::Reverse>(cb_in, cb_w, cb_p); | |
| cb_wait_front(cb_in, TR * KIN); | |
| constexpr uint32_t PS = 0; | |
| cb_wait_front(cb_w, NW); | |
| cb_wait_front(cb_one, 1); | |
| // ---- merged head: score 1x1 first (s = bf16(x[:, 256:] @ Ws + bs), the ttnn 1x1 conv steps: fp32 DST over K | |
| // in order at MM_FID, bias row-broadcast on the fp32 partials); one tile row at a time packed in place into the | |
| // logits shard (CB bound to it) and pushed, so the writer can signal the softmax cores early | |
| { | |
| constexpr uint32_t cb_s = DH_SCORE_CB, cb_sp = DH_SP_CB, ST = 3, SW0 = 72, SB0 = 96; | |
| for (uint32_t r = 0; r < TR; ++r) { | |
| reconfig_data_format<SrcOrder::Reverse>(cb_in, cb_w); | |
| pack_reconfig_data_format(cb_sp); | |
| mmf_init(cb_in, cb_w, ST); | |
| cb_reserve_back(cb_sp, ST); | |
| tile_regs_acquire(); | |
| for (uint32_t k = 0; k < KT; ++k) { | |
| mmf_block(cb_in, cb_w, r * KIN + KT + k, SW0 + k * ST, 0, ST); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| for (uint32_t j = 0; j < ST; ++j) { | |
| pack_tile<true>(j, cb_sp, j); | |
| } | |
| tile_regs_release(); | |
| cb_push_back(cb_sp, ST); | |
| cb_wait_front(cb_sp, ST); | |
| reconfig_data_format(cb_sp, cb_w); | |
| pack_reconfig_data_format(cb_s); | |
| add_bcast_rows_init(cb_sp, cb_w); | |
| cb_reserve_back(cb_s, ST); | |
| tile_regs_acquire(); | |
| for (uint32_t j = 0; j < ST; ++j) { | |
| add_tiles_bcast_rows(cb_sp, cb_w, j, SB0 + j, j); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| for (uint32_t j = 0; j < ST; ++j) { | |
| pack_tile(j, cb_s); | |
| } | |
| tile_regs_release(); | |
| cb_push_back(cb_s, ST); | |
| cb_pop_front(cb_sp, ST); | |
| } | |
| } | |
| for (uint32_t r = 0; r < TR; ++r) { | |
| // ---- 1) matmul -> CB_P (fp32) | |
| reconfig_data_format<SrcOrder::Reverse>(cb_in, cb_w); | |
| pack_reconfig_data_format(cb_p); | |
| mmf_init(cb_in, cb_w, HB); | |
| cb_reserve_back(cb_p, WT + PS); | |
| for (uint32_t nb = 0; nb < WT / HB; ++nb) { | |
| tile_regs_acquire(); | |
| for (uint32_t k = 0; k < KT; ++k) { | |
| if constexpr (!(DH_SKIP & 1)) mmf_block(cb_in, cb_w, r * KIN + k, nb * (KT * HB) + k * HB, 0, HB); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| for (uint32_t j = 0; j < HB; ++j) { | |
| pack_tile<true>(j, cb_p, nb * HB + j); | |
| } | |
| tile_regs_release(); | |
| } | |
| cb_push_back(cb_p, WT + PS); | |
| // ---- 2) + bias -> CB_Z (bf16) | |
| cb_wait_front(cb_p, WT + PS); | |
| reconfig_data_format(cb_p, cb_w); | |
| pack_reconfig_data_format(cb_z); | |
| add_bcast_rows_init(cb_p, cb_w); | |
| cb_reserve_back(cb_z, WT); | |
| for (uint32_t b = 0; b < WT / HB; ++b) { | |
| tile_regs_acquire(); | |
| for (uint32_t k = 0; k < HB; ++k) { | |
| if constexpr (!(DH_SKIP & 2)) add_tiles_bcast_rows(cb_p, cb_w, b * HB + k, BIAS0 + b * HB + k, k); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| for (uint32_t k = 0; k < HB; ++k) { | |
| pack_tile(k, cb_z); | |
| } | |
| tile_regs_release(); | |
| } | |
| cb_push_back(cb_z, WT); | |
| cb_pop_front(cb_p, WT + PS); | |
| cb_wait_front(cb_z, WT); | |
| // ---- 3) x^2 -> CB_SQ (fp32) | |
| reconfig_data_format(cb_z, cb_z); | |
| pack_reconfig_data_format(cb_sq); | |
| mul_init(cb_z, cb_z); | |
| cb_reserve_back(cb_sq, WT); | |
| for (uint32_t b = 0; b < WT / HB; ++b) { | |
| tile_regs_acquire(); | |
| for (uint32_t k = 0; k < HB; ++k) { | |
| if constexpr (!(DH_SKIP & 4)) mul_tiles(cb_z, cb_z, b * HB + k, b * HB + k, k); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| for (uint32_t k = 0; k < HB; ++k) { | |
| pack_tile(k, cb_sq); | |
| } | |
| tile_regs_release(); | |
| } | |
| cb_push_back(cb_sq, WT); | |
| // ---- row sum, rsqrt -> CB_RS (fp32, column 0) | |
| cb_wait_front(cb_sq, WT); | |
| reconfig_data_format(cb_one, cb_sq); | |
| pack_reconfig_data_format(cb_rs); | |
| reduce_init<PoolType::SUM, ReduceDim::REDUCE_ROW>(cb_sq, cb_one, cb_rs); | |
| cb_reserve_back(cb_rs, 1); | |
| tile_regs_acquire(); | |
| for (uint32_t t = 0; t < WT; ++t) { | |
| if constexpr (!(DH_SKIP & 8)) reduce_tile<PoolType::SUM, ReduceDim::REDUCE_ROW>(cb_sq, cb_one, t, 0, 0); | |
| } | |
| rsqrt_tile_init(); | |
| if constexpr (!(DH_SKIP & 16)) rsqrt_tile(0); | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| pack_tile(0, cb_rs); | |
| tile_regs_release(); | |
| reduce_uninit(cb_sq); | |
| cb_push_back(cb_rs, 1); | |
| cb_pop_front(cb_sq, WT); | |
| // ---- x * r (column broadcast), pack-untilize -> CB_OUT | |
| cb_wait_front(cb_rs, 1); | |
| reconfig_data_format(cb_z, cb_rs); | |
| pack_reconfig_data_format(cb_out); | |
| mul_bcast_cols_init(cb_z, cb_rs); | |
| pack_untilize_dest_init<HB, WT>(cb_out); | |
| cb_reserve_back(cb_out, WT); | |
| for (uint32_t b = 0; b < WT / HB; ++b) { | |
| tile_regs_acquire(); | |
| for (uint32_t k = 0; k < HB; ++k) { | |
| if constexpr (!(DH_SKIP & 32)) mul_tiles_bcast_cols(cb_z, cb_rs, b * HB + k, 0, k); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| pack_untilize_dest<HB, WT>(cb_out, 1, b); | |
| tile_regs_release(); | |
| } | |
| pack_untilize_uninit(cb_out); | |
| cb_push_back(cb_out, WT); | |
| cb_pop_front(cb_rs, 1); | |
| cb_pop_front(cb_z, WT); | |
| } | |
| } | |