File size: 2,159 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 | // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// Score softmax core (SP_HEAD_SM=1), PROC 0: build the MAX / SUM reduce scaler tiles (bf16 1.0 in row 0 of
// every face, as ttnn's calculate_and_prepare_reduce_scaler), wait until every producer core of this core's
// score tile rows has signalled (local semaphore 0 reaches nrows; one noc_semaphore_inc per tile row), then
// read each [32, 96] logits tile row (3 tiles, 6 KB) from the producer's L1 shard into CB_IN.
// RT args: src_addr (score logits shard base, same on every producer core), nrows, then per row
// (noc_x << 16 | noc_y) of the producer and the byte offset of the row inside its shard.
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"
FORCE_INLINE void scaler_tile(uint32_t cb) {
cb_reserve_back(cb, 1);
volatile tt_l1_ptr uint32_t* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb));
for (uint32_t i = 0; i < 512; ++i) {
p[i] = 0;
}
for (uint32_t f = 0; f < 4; ++f) {
for (uint32_t j = 0; j < 8; ++j) {
p[f * 128 + j] = 0x3F803F80u;
}
}
(void)p[511];
cb_push_back(cb, 1);
}
void kernel_main() {
const uint32_t src_addr = get_arg_val<uint32_t>(0);
const uint32_t nrows = get_arg_val<uint32_t>(1);
constexpr uint32_t cb_in = get_compile_time_arg_val(0);
constexpr uint32_t cb_maxs = get_compile_time_arg_val(1);
constexpr uint32_t cb_sums = get_compile_time_arg_val(2);
constexpr uint32_t ROWB = 3 * 2048;
scaler_tile(cb_maxs);
scaler_tile(cb_sums);
volatile tt_l1_ptr uint32_t* sem = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(0));
noc_semaphore_wait_min(sem, nrows);
for (uint32_t i = 0; i < nrows; ++i) {
const uint32_t xy = get_arg_val<uint32_t>(2 + 2 * i);
const uint32_t off = get_arg_val<uint32_t>(3 + 2 * i);
cb_reserve_back(cb_in, 3);
noc_async_read(get_noc_addr(xy >> 16, xy & 0xFFFF, src_addr + off), get_write_ptr(cb_in), ROWB);
noc_async_read_barrier();
cb_push_back(cb_in, 3);
}
noc_semaphore_set(sem, 0);
}
|