File size: 8,642 Bytes
c699c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 | // 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);
}
}
|