superpoint-p150 / code /kernels /sp_conv /rc_compute.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
3.61 kB
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// RowConv compute (models/tt/row_conv.py): ttnn conv_bmm_tilize.cpp's height-sharded order with
// packer_l1_acc and fused bias. Per K block ky (3*CT tiles in K order kx, c): fp32 DST accumulates the
// tile products of output tile row m and NL output channel tiles; the block result is packed to fp32
// partials (ky = 0 overwrite, else packer L1 accumulate). Then partials + bias (row broadcast) with
// ReLU on pack -> bf16 output tiles (m, n) at m * NL + n.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/matmul.h"
#include "api/compute/compute_kernel_hw_startup.h"
#include "api/compute/pack.h"
#include "api/compute/bcast.h"
#include "api/compute/reconfig_data_format.h"
void kernel_main() {
constexpr uint32_t cb_a0 = get_compile_time_arg_val(0);
constexpr uint32_t cb_a1 = get_compile_time_arg_val(1);
constexpr uint32_t cb_w = get_compile_time_arg_val(2);
constexpr uint32_t cb_b = get_compile_time_arg_val(3);
constexpr uint32_t cb_p = get_compile_time_arg_val(4);
constexpr uint32_t cb_o = get_compile_time_arg_val(5);
constexpr uint32_t NL = get_compile_time_arg_val(6);
constexpr uint32_t CH = get_compile_time_arg_val(7);
constexpr uint32_t CT = get_compile_time_arg_val(8);
constexpr uint32_t BLK_H = get_compile_time_arg_val(9);
constexpr bool RELU = get_compile_time_arg_val(10) == 1;
constexpr uint32_t KB = 3 * CT; // K tiles per block
static_assert(NL <= 4, "fp32 half-sync DST holds 4 tiles");
compute_kernel_hw_startup<SrcOrder::Reverse>(cb_a0, cb_w, cb_p);
matmul_block_init(cb_a0, cb_w, false, NL, 1, 1);
cb_reserve_back(cb_p, 3 * NL);
for (uint32_t ky = 0; ky < 3; ++ky) {
cb_wait_front(cb_a0, (ky + 1) * BLK_H);
cb_wait_front(cb_a1, (ky + 1) * BLK_H);
cb_wait_front(cb_w, (ky + 1) * KB * NL);
pack_reconfig_l1_acc(ky > 0 ? 1 : 0);
for (uint32_t m = 0; m < 3; ++m) {
tile_regs_acquire();
for (uint32_t kx = 0; kx < 3; ++kx) {
for (uint32_t c = 0; c < CT; ++c) {
const uint32_t cb = c < CH ? cb_a0 : cb_a1;
const uint32_t idx = ky * BLK_H + (m * 3 + kx) * CH + (c < CH ? c : c - CH);
matmul_block(cb, cb_w, idx, (ky * KB + kx * CT + c) * NL, 0, false, NL, 1, 1);
}
}
tile_regs_commit();
tile_regs_wait();
for (uint32_t n = 0; n < NL; ++n) {
pack_tile<true>(n, cb_p, m * NL + n);
}
tile_regs_release();
}
}
pack_reconfig_l1_acc(0);
cb_push_back(cb_p, 3 * NL);
cb_wait_front(cb_p, 3 * NL);
cb_wait_front(cb_b, NL);
cb_reserve_back(cb_o, 3 * NL);
if constexpr (RELU) {
PACK((llk_pack_relu_config(ReluConfig::zero())));
}
pack_reconfig_data_format(cb_p, cb_o);
reconfig_data_format(cb_w, cb_p, cb_a0, cb_b);
add_bcast_rows_init(cb_p, cb_b);
for (uint32_t m = 0; m < 3; ++m) {
tile_regs_acquire();
for (uint32_t n = 0; n < NL; ++n) {
add_tiles_bcast_rows(cb_p, cb_b, m * NL + n, n, n);
}
tile_regs_commit();
tile_regs_wait();
for (uint32_t n = 0; n < NL; ++n) {
pack_tile<true>(n, cb_o, m * NL + n);
}
tile_regs_release();
}
if constexpr (RELU) {
PACK((llk_pack_relu_config(ReluConfig::none())));
}
cb_pop_front(cb_p, 3 * NL);
cb_push_back(cb_o, 3 * NL);
}