superpoint-p150 / code /kernels /sp_desc /dh_compute.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
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.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/compute_kernel_hw_startup.h"
#include "api/compute/matmul.h"
#include "api/compute/pack.h"
#include "api/compute/eltwise_binary.h"
#include "api/compute/bcast.h"
#include "api/compute/reduce.h"
#include "api/compute/eltwise_unary/rsqrt.h"
#include "api/compute/pack_untilize.h"
#include "api/compute/reconfig_data_format.h"
#ifndef DH_SKIP
#define DH_SKIP 0
#endif
#ifndef DH_KIN
#define DH_KIN 8 // input tiles per tile row (16: merged head conv output, descriptor channels first)
#endif
#ifndef DH_NW
#define DH_NW 72
#endif
#ifndef MM_FID
#define MM_FID ckernel::MathFidelity::HiFi2
#endif
// 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);
#ifdef DH_SCORE_CB
// ---- 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);
}
}
#endif
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);
}
}