File size: 3,672 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 | // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// Descriptor L2 normalisation + untilize in one op (models/tt/desc_norm.py). Per tile row of the
// height-sharded [rows, 256] TILE descriptor map: S = sum_c x^2 (x*x on the FPU, packed fp32, row-reduced
// with a ones scaler), r = rsqrt(S) (SFPU), y = x * r (column broadcast), pack-untilized into a
// [32 rows, 256] row-major block for the writer.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/compute_kernel_hw_startup.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"
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 WT = 8, HB = 4; // tiles per row (256 ch), tiles per DST block
compute_kernel_hw_startup(cb_in, cb_in, cb_sq);
cb_wait_front(cb_in, TR * WT);
cb_wait_front(cb_one, 1);
for (uint32_t r = 0; r < TR; ++r) {
// ---- x^2 -> CB_SQ (fp32)
reconfig_data_format(cb_in, cb_in);
pack_reconfig_data_format(cb_sq);
mul_init(cb_in, cb_in);
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) {
const uint32_t t = r * WT + b * HB + k;
mul_tiles(cb_in, cb_in, t, t, 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) {
reduce_tile<PoolType::SUM, ReduceDim::REDUCE_ROW>(cb_sq, cb_one, t, 0, 0);
}
rsqrt_tile_init();
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_in, cb_rs);
pack_reconfig_data_format(cb_out);
mul_bcast_cols_init(cb_in, 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) {
mul_tiles_bcast_cols(cb_in, cb_rs, r * WT + 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);
}
}
|